File size: 5,473 Bytes
dc9f917
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""E2 -- gradient alignment probe (exp7 explainability battery, sec 5).



Per refinement k, cosine between the controller's emitted update

``alpha_k * Delta Y_j`` and the true descent direction

``-grad_{Y_j} sum_i d_i``, obtained by autograd through the *frozen rollout

only* (never through the controller's own network). This directly tests the

"amortized descent" story -- does the controller behave like an unrolled

gradient step? -- and works unchanged on both the split and fused

architectures, since it only calls the controller's public forward pieces

(``condition`` / ``initial_plan`` / ``refine`` / ``to_actions``).

"""

import argparse
import json
import sys
from pathlib import Path

import numpy as np
import torch
import torch.nn.functional as F

sys.path.insert(0, str(Path(__file__).resolve().parents[1]))

from lejepa_control.data import LatentGoalDataset, split_episodes  # noqa: E402
from lejepa_control.rollout import goal_distance, rollout_plan  # noqa: E402
from lejepa_control.solver import load_controller  # noqa: E402
from lejepa_control.world_model import load_lewm  # noqa: E402


def parse_args():
    p = argparse.ArgumentParser()
    p.add_argument('--controller', default='data/runs/controller/controller.pt')
    p.add_argument('--latents', default='data/latents')
    p.add_argument('--samples', type=int, default=256)
    p.add_argument('--horizon', type=int, default=5)
    p.add_argument('--tag', default='main')
    p.add_argument('--out', default='data/runs/probe_alignment')
    return p.parse_args()


def descent_direction(controller, model, tokens, ctx, past, goal):
    """``-grad_Y sum_i d_i``: true descent, autograd through the frozen

    rollout only -- ``tokens`` is detached first, so no gradient reaches the

    controller's own network here, only the path token -> action -> rollout.

    """
    y = tokens.detach().clone().requires_grad_(True)
    pred = rollout_plan(model, ctx, past, controller.to_actions(y))
    loss = goal_distance(pred, goal).sum()
    (grad,) = torch.autograd.grad(loss, y)
    return -grad


def alignment_trace(controller, model, ctx, past, goal):
    """Per-refinement cosine between the emitted update and true descent.



    Mirrors ``IterativeController.forward``'s loop exactly (same ``cond``,

    same ``tokens`` recursion), but at each step also scores the emitted

    update against the gradient-through-rollout direction at that same

    ``Y_k``. Returns a ``(K, B, H)`` cosine tensor.

    """
    with torch.no_grad():
        cond = controller.condition(ctx, goal)
        tokens = controller.initial_plan(cond)

    cosines = []
    for k in range(controller.refinements):
        with torch.no_grad():
            actions = controller.to_actions(tokens)
            pred = rollout_plan(model, ctx, past, actions)

        true_dir = descent_direction(controller, model, tokens, ctx, past, goal)

        with torch.no_grad():
            delta = controller.refine(tokens, cond, pred, goal)
            idx = min(k, controller.step_logit.numel() - 1)
            step = torch.sigmoid(controller.step_logit[idx])
            emitted = step * delta  # (B, H, W): the actual applied update

            cos = F.cosine_similarity(emitted, true_dir, dim=-1)  # (B, H)
            cosines.append(cos)

            tokens = tokens + emitted

    return torch.stack(cosines)  # (K, B, H)


def main():
    args = parse_args()
    device = 'cuda' if torch.cuda.is_available() else 'cpu'
    torch.manual_seed(0)

    stats = json.loads((Path(args.latents) / 'stats.json').read_text())
    model = load_lewm(device=device)
    controller, ckpt = load_controller(
        args.controller, latent_dim=stats['latent_dim'], device=device
    )
    controller.eval()
    print(f'controller step {ckpt["step"]}, K={controller.refinements}, '
          f'fused={controller.fused}, no_latent_proj={controller.no_latent_proj}')

    _, val_eps = split_episodes(stats['n_episodes'])
    val = LatentGoalDataset(
        args.latents, max_offset=5, episodes=val_eps, horizon=args.horizon
    )
    idx = np.random.default_rng(0).choice(len(val), args.samples, replace=False)
    batch = {
        k: torch.stack([val[int(i)][k] for i in idx]).to(device)
        for k in ('context', 'past_actions', 'goal')
    }

    cos = alignment_trace(
        controller, model, batch['context'], batch['past_actions'], batch['goal']
    )  # (K, B, H)

    per_k = cos.mean(dim=(1, 2))
    per_k_std = cos.std(dim=(1, 2))
    print('\n=== E2: gradient alignment (emitted update vs true descent) ===')
    print(f'{"k":>3}{"mean cos":>11}{"std":>9}')
    for k in range(cos.size(0)):
        print(f'{k:>3}{per_k[k].item():>11.4f}{per_k_std[k].item():>9.4f}')

    report = {
        'tag': args.tag,
        'checkpoint': args.controller,
        'step': int(ckpt['step']),
        'fused': bool(controller.fused),
        'no_latent_proj': bool(controller.no_latent_proj),
        'mean_cosine_per_k': [round(float(v), 5) for v in per_k],
        'std_cosine_per_k': [round(float(v), 5) for v in per_k_std],
    }
    out = Path(args.out)
    out.mkdir(parents=True, exist_ok=True)
    with (out / 'probe_alignment.jsonl').open('a') as f:
        f.write(json.dumps(report) + '\n')
    print(f'\nwrote {out / "probe_alignment.jsonl"}')


if __name__ == '__main__':
    main()