latent_backtrack / verify_fo.py
Avra98's picture
Add training code (same as GitHub reasoning-by-superposition-latent)
8f46582 verified
Raw
History Blame Contribute Delete
1.57 kB
import json, collections
def reach(edges, src):
adj = collections.defaultdict(list)
for a, b in edges: adj[a].append(b)
seen, q = {src}, [src]
while q:
u = q.pop()
for v in adj[u]:
if v not in seen: seen.add(v); q.append(v)
return seen
def dist_from(edges, src):
adj = collections.defaultdict(list)
for a, b in edges: adj[a].append(b)
d = {src: 0}; q = collections.deque([src])
while q:
u = q.popleft()
for v in adj[u]:
if v not in d: d[v] = d[u] + 1; q.append(v)
return d
for flavor in ("coconut", "bfs"):
d = json.load(open(f"data/star_2arm_L6_valid_fo_{flavor}.json"))
L = 6; bad = 0; front_sizes = set()
for s in d:
R = s["root"]; rset = reach(s["edges"], R)
fd = dist_from(s["edges"], R)
nd = dist_from(s["edges"], s["neg_root"])
if s["target"] not in rset or s["neg_target"] in rset: bad += 1; continue
for k in range(1, L + 1):
nk = s["neighbor_k"][str(k)]; gk = s["neg_neighbor_k"][str(k)]
front_sizes.add(len(nk))
# reachable frontier: in rset and at forward-distance k from R
if any(x not in rset or fd.get(x) != k for x in nk): bad += 1; break
# neg frontier: NOT reachable from R, and at distance k from neg_root
if any(x in rset or nd.get(x) != k for x in gk): bad += 1; break
print(f"{flavor}: {len(d)-bad}/{len(d)} valid | neighbor_k frontier sizes seen: {sorted(front_sizes)} | keys: {sorted(int(k) for k in d[0]['neighbor_k'])}")