latent_backtrack / verify_star.py
Avra98's picture
Add training code (same as GitHub reasoning-by-superposition-latent)
8f46582 verified
Raw
History Blame Contribute Delete
2.05 kB
import json, collections, random
from dataset import expand_data
from stokenizer import STokenizer
from graph_metrics import _distances, _classify
tok = STokenizer()
d = json.load(open("data/star_2arm_L6_valid_coconut.json"))
L = 6
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
bad = 0
for s in d:
r = reach(s["edges"], s["root"])
# target reachable, neg_target NOT reachable, unique path depth L
if s["target"] not in r: bad += 1; continue
if s["neg_target"] in r: bad += 1; continue
fdist, bdist, Ld = _distances(s["edges"], s["root"], s["target"])
if Ld != L: bad += 1; continue
# neighbor_k path must be a real chain root->target on the target arm
if any(s["neighbor_k"][str(k)][0] not in fdist or fdist[s["neighbor_k"][str(k)][0]] != k for k in range(1, L+1)):
bad += 1
print(f"structural check: {len(d)-bad}/{len(d)} valid (target reachable, neg unreachable, depth=={L}, unique path)")
# tokenization check on one sample, all hops + final answer
s = d[0]
print("\nsample root/target/neg:", s["root"], s["target"], s["neg_target"])
for k in range(1, L + 2): # 1..L intermediate, L+1 = [A] answer
q, cont = expand_data(s, k, len(s["steps"]))
ids_q = tok.encode(q, add_special_tokens=False)
ids_c = tok.encode(cont, add_special_tokens=False)
n_lat = q.count("<|latent|>")
tag = "[A]answer" if k == L + 1 else f"hop{k}"
print(f" {tag:9s}: {n_lat} latents -> target '{cont}' (q_tokens={len(ids_q)}, ok)")
# category-metric sanity: decoy arm inflates frontier vs optimal
fdist, bdist, Ld = _distances(s["edges"], s["root"], s["target"])
frontier_nodes = [g for g in fdist if fdist[g] == 1]
print("\nhop-1 frontier (both arms' first nodes):", frontier_nodes,
"| optimal (target arm only):", [g for g in frontier_nodes if _classify(g,1,fdist,bdist,Ld)[2]])