Canonical orbit search for symmetric pruning masks / experiment.py
Failed on benchmark
1import json
2from time import perf_counter
3from canonical_orbit import (all_permutations, product_group, exhaustive,
4 canonical, canonical_dfs, orbit)
5
6
7def run():
8 n, k = 8, 4
9 blocks = [tuple(range(4)), tuple(range(4, 8))]
10 group = product_group(blocks)
11 masks = exhaustive(n, k)
12 # Exact symmetry: every channel in block 0 has importance 3, and every
13 # channel in block 1 has importance 1. Thus orbit-related masks score alike.
14 score = lambda s: sum(3 if i < 4 else 1 for i in s)
15 t0 = perf_counter()
16 exhaustive_best = max((score(s), s) for s in masks)
17 exhaustive_sec = perf_counter() - t0
18 t0 = perf_counter()
19 reps = canonical_dfs(n, k, group)
20 canonical_sec = perf_counter() - t0
21 canonical_best = max((score(s), s) for s in reps)
22 reps_set = set(reps)
23 covered = set().union(*(orbit(r, group) for r in reps_set))
24 orbit_invariance = all(score(s) == score(canonical(s, group)) for s in masks)
25 return {
26 "n": n, "k": k, "group_size": len(group),
27 "all_masks": len(masks), "canonical_representatives": len(reps),
28 "evaluation_reduction": len(masks) / len(reps),
29 "exhaustive_best_score": exhaustive_best[0],
30 "canonical_best_score": canonical_best[0],
31 "same_best_score": exhaustive_best[0] == canonical_best[0],
32 "orbit_coverage": covered == set(masks),
33 "score_orbit_invariant": orbit_invariance,
34 "exhaustive_scan_sec": exhaustive_sec,
35 "canonical_dfs_sec": canonical_sec,
36 }
37
38if __name__ == '__main__':
39 print(json.dumps(run(), sort_keys=True))