File size: 6,358 Bytes
8f46582
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
#!/usr/bin/env python3
"""Latent chain-of-thought interpretability probe (logit-lens).

For a trained coconut model on the 2-arm star task, decode what each intermediate
latent token encodes. Because node id == token id, applying the LM head to a
latent position's hidden state and taking argmax yields a predicted graph node
directly.

In a single forward pass with a full k-latent chain the model predicts the hop-m
node at the position `root + (m-1)` (stage-m training supervises hop-m after m-1
latents). We therefore read, for each latent slot, the model's decoded node and
compare it to the TRUE bfs frontier at that hop, over the whole validation set.

Outputs:
  * per-hop accuracy: does latent slot m decode to a node on the reachable
    frontier at depth m? (frontier = either arm's depth-m node)
  * per-hop "on-target-arm" accuracy: exact match to the shortest-path node.
  * a confusion matrix: latent slot (predicted depth) vs the decoded node's TRUE
    role (reachable depth 0..L, negative-component, or off-graph) -- i.e. which
    depth each latent actually captures and how cleanly.

Usage:
  python scripts/probe_latents.py \
      --ckpt ckpts/star-coconut-L10-bfs-backtrack/checkpoint_650 \
      --val data/star_2arm_L10_valid_fo_bfs.json \
      --model_id configs/symbol-2layer-8head-768dim-L20.json --L 10
"""
import argparse
import json
import torch
from transformers import AutoModelForCausalLM, AutoConfig

from stokenizer import STokenizer
from coconut import Coconut


def build_prefix_tokens(sample, tok):
    """Match dataset.get_prefix (without the random shuffles; order-invariant for a
    trained model). Ends at the root token; the caller appends latent tokens."""
    edges = sample["edges"]
    q = "<eos> " + "|".join([f" {e[0]} {e[1]} " for e in edges]).strip() + " [Q] "
    q += f'{sample["target"]} {sample["neg_target"]}'
    q += " [R] " + str(sample["root"])
    return tok.encode(q, add_special_tokens=False)


def node_depth_maps(sample, L):
    """Map each node id -> ('pos', depth) reachable-arm depth 0..L,
    or ('neg', depth) for the unreachable component, else absent."""
    role = {}
    role[sample["root"]] = ("pos", 0)
    role[sample["neg_root"]] = ("neg", 0)
    for k in range(1, L + 1):
        for n in sample["neighbor_k"].get(str(k), []):
            role[n] = ("pos", k)
        for n in sample["neg_neighbor_k"].get(str(k), []):
            role[n] = ("neg", k)
    return role


@torch.no_grad()
def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--ckpt", required=True)
    ap.add_argument("--val", required=True)
    ap.add_argument("--model_id", default="configs/symbol-2layer-8head-768dim-L20.json")
    ap.add_argument("--L", type=int, required=True)
    ap.add_argument("--batch_size", type=int, default=64)
    ap.add_argument("--device", default="cuda:0")
    args = ap.parse_args()

    tok = STokenizer()
    latent_id = tok.convert_tokens_to_ids("<|latent|>")

    base = AutoModelForCausalLM.from_config(AutoConfig.from_pretrained(args.model_id))
    model = Coconut(base, latent_id,
                    tok.convert_tokens_to_ids("<|start-latent|>"),
                    tok.convert_tokens_to_ids("<|end-latent|>"),
                    tok.eos_token_id)
    sd = torch.load(args.ckpt, map_location="cpu")
    missing, unexpected = model.load_state_dict(sd, strict=False)
    print(f"loaded ckpt: {len(sd)} tensors | missing={len(missing)} unexpected={len(unexpected)}")
    model.to(args.device).eval()

    data = json.load(open(args.val))
    L = args.L

    # per-hop counters
    frontier_correct = [0] * (L + 1)   # index m = hop m
    target_correct = [0] * (L + 1)
    total = 0
    # confusion: rows = latent slot m (1..L), cols index: 0..L reachable depth,
    # L+1 = negative component, L+2 = off-graph / non-node
    NEG, OFF = L + 1, L + 2
    conf = [[0] * (L + 3) for _ in range(L + 1)]

    for i in range(0, len(data), args.batch_size):
        batch = data[i:i + args.batch_size]
        seqs = []
        for s in batch:
            ids = build_prefix_tokens(s, tok) + [latent_id] * L
            seqs.append(ids)
        maxlen = max(len(x) for x in seqs)
        assert all(len(x) == maxlen for x in seqs), "fixed L => equal lengths expected"
        input_ids = torch.tensor(seqs, device=args.device)
        attn = torch.ones_like(input_ids)
        pos = torch.arange(maxlen, device=args.device).unsqueeze(0).expand(len(batch), -1)
        out = model.forward(input_ids, attn, input_ids.clone(), pos)
        logits = out.logits  # (B, T, V)

        for bi, s in enumerate(batch):
            total += 1
            role = node_depth_maps(s, L)
            root_pos = len(build_prefix_tokens(s, tok)) - 1  # position of root token
            for m in range(1, L + 1):
                # hop m decoded at position root_pos + (m-1)
                pred = int(torch.argmax(logits[bi, root_pos + (m - 1)]).item())
                # frontier membership at hop m
                if pred in s["neighbor_k"].get(str(m), []):
                    frontier_correct[m] += 1
                # exact target-arm node (shortest path) at hop m
                if m - 1 < len(s["steps"]) and pred == int(s["steps"][m - 1]):
                    target_correct[m] += 1
                # confusion by true role of the predicted node
                r = role.get(pred)
                if r is None:
                    conf[m][OFF] += 1
                elif r[0] == "neg":
                    conf[m][NEG] += 1
                else:
                    conf[m][r[1]] += 1

    print(f"\n=== samples: {total} | checkpoint: {args.ckpt} ===")
    print("\nPer-hop latent decode accuracy (single forward, full latent chain):")
    print(f"{'hop':>4} {'frontier_acc':>13} {'target_arm_acc':>15}")
    for m in range(1, L + 1):
        print(f"{m:>4} {frontier_correct[m] / total:>13.3f} {target_correct[m] / total:>15.3f}")

    print("\nConfusion: latent slot m (row) vs TRUE reachable depth of decoded node (col).")
    header = "slot\\depth " + " ".join(f"{d:>5}" for d in range(0, L + 1)) + f" {'neg':>5} {'off':>5}"
    print(header)
    for m in range(1, L + 1):
        row = conf[m]
        cells = " ".join(f"{row[d]:>5}" for d in range(0, L + 1))
        print(f"{m:>9} {cells} {row[NEG]:>5} {row[OFF]:>5}")


if __name__ == "__main__":
    main()