File size: 9,377 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
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
"""Diagnostics for horizon-reset procrastination and the K=5 regression.



Two questions, both answered in the world model's latent space against the

same held-out latents used for validation:



1. Under receding-horizon execution, does the closed loop have a fixed point

   outside the success radius? Fitting ``D_{n+1} = c*D_n + b`` gives

   ``D* = b/(1-c)`` when ``0 < c < 1``. A positive intercept means each

   replan removes a *fraction* of the remaining distance but adds a floor, so

   the loop stalls short of the goal no matter how long it runs.



2. Does refinement keep improving predicted cost past the trained depth K=3

   even as real success degrades? If predicted cost improves while success

   drops, the extra refinements are exploiting the world model rather than

   planning better.



Question 2 originally measured only the *imagined* side, which cannot separate

"the plan got better" from "the plan left the region where the world model is

trustworthy" — both look like falling predicted cost. The support score

``r(C, b) = -log beta(b | C) / A`` against the fitted GMM answers that directly:

if the violation fraction rises with k while predicted cost falls, the extra

refinements are buying imagined progress with out-of-distribution actions.

"""

import argparse
import json
import sys
from pathlib import Path

import numpy as np
import torch

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

from lejepa_control.data import LatentGoalDataset, split_episodes  # noqa: E402
from lejepa_control.losses import BehaviorDensity  # 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('--density', default='data/runs/density/density.pt')
    p.add_argument('--samples', type=int, default=256)
    p.add_argument('--horizon', type=int, default=5)
    p.add_argument('--replans', type=int, default=12)
    p.add_argument('--max-k', type=int, default=8)
    p.add_argument('--tag', default='main')
    p.add_argument('--out', default='data/runs/diagnostics')
    return p.parse_args()


@torch.no_grad()
def closed_loop_trace(controller, model, batch, device, execute, replans):
    """Latent goal distance after each replan, simulating in latent space.



    The world model is its own simulator here: whatever the predictor says

    the executed blocks caused becomes the next context. That is exactly the

    loop the real MPC runs, minus the env, so it isolates the planner's

    dynamics from simulator mismatch.

    """
    ctx = batch['context'].to(device)
    past = batch['past_actions'].to(device)
    goal = batch['goal'].to(device)
    num_context = controller.num_context

    trace = []
    for _ in range(replans):
        out = controller(model, ctx, past, goal)
        plan = out['plans'][-1]

        # execute the first `execute` blocks and advance the latent state
        pred, frames = rollout_plan(
            model, ctx, past, plan[:, :execute], return_frames=True
        )
        trace.append(goal_distance(pred, goal)[:, execute - 1].cpu().numpy())

        ctx = frames[:, -num_context:]
        past = torch.cat([past, plan[:, :execute]], dim=1)[:, -(num_context - 1):]
    return np.stack(trace, axis=1)  # (B, replans)


def fit_contraction(trace):
    """Least-squares ``D_{n+1} = c*D_n + b`` over all consecutive pairs."""
    x = trace[:, :-1].ravel()
    y = trace[:, 1:].ravel()
    A = np.stack([x, np.ones_like(x)], axis=1)
    (c, b), *_ = np.linalg.lstsq(A, y, rcond=None)
    resid = y - (c * x + b)
    ss = 1 - resid.var() / y.var() if y.var() > 0 else float('nan')
    fixed = b / (1 - c) if abs(1 - c) > 1e-9 else float('inf')
    return float(c), float(b), float(fixed), float(ss)


@torch.no_grad()
def refinement_trace(controller, model, batch, device, max_k,

                     density=None, c95=None):
    """Predicted goal cost after each refinement, out to ``max_k``.



    Also returns the support-violation fraction per refinement when a density

    model is supplied — the imagined cost and the violation fraction moving in

    opposite directions is the exploitation signature.

    """
    saved = controller.refinements
    try:
        controller.refinements = max_k
        out = controller(
            model,
            batch['context'].to(device),
            batch['past_actions'].to(device),
            batch['goal'].to(device),
        )
        q = batch['goal_offset'].to(device).clamp(1, controller.horizon)
        costs, arrivals, moves, violations, scores = [], [], [], [], []
        prev = None
        for k, d in enumerate(out['distances']):
            costs.append(float(d[:, -1].mean()))
            arrivals.append(
                float(d.gather(1, (q - 1).unsqueeze(1)).squeeze(1).mean())
            )
            plan = out['plans'][k]
            moves.append(
                0.0 if prev is None else float((plan - prev).abs().mean())
            )
            prev = plan

            if density is not None and c95 is not None:
                # score every block of the plan at the latent history it is
                # actually proposed from, which is what the support term does
                ctx = out['contexts'][k].flatten(0, 1)  # (B*H, N, D)
                score = density.nll_per_dim(ctx, plan.flatten(0, 1))
                violations.append(float((score > c95).float().mean()))
                scores.append(float(score.mean()))
        return costs, arrivals, moves, violations, scores
    finally:
        controller.refinements = saved


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}')

    _, 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])
        for k in ('context', 'past_actions', 'goal', 'goal_offset')
    }

    density, c95 = None, None
    if Path(args.density).exists():
        d_ckpt = torch.load(args.density, map_location=device, weights_only=False)
        density = BehaviorDensity(
            latent_dim=d_ckpt['latent_dim'], components=d_ckpt['components']
        )
        density.load_state_dict(d_ckpt['state_dict'])
        density.to(device).eval().requires_grad_(False)
        c95 = float(d_ckpt['c95'])
        print(f'support model loaded, c95={c95:.4f}')
    else:
        print(f'no support model at {args.density} — violation column disabled')

    report = {'tag': args.tag, 'checkpoint': args.controller,
              'step': int(ckpt['step'])}

    print('\n=== 1. closed-loop contraction ===')
    print(f'{"execute":>8}{"D_0":>9}{"D_final":>10}{"c":>8}{"b":>9}'
          f'{"D*":>9}{"R^2":>7}')
    report['contraction'] = {}
    for execute in (1, 5):
        trace = closed_loop_trace(
            controller, model, batch, device, execute, args.replans
        )
        c, b, fixed, r2 = fit_contraction(trace)
        m = trace.mean(axis=0)
        print(f'{execute:>8}{m[0]:>9.4f}{m[-1]:>10.4f}'
              f'{c:>8.4f}{b:>9.4f}{fixed:>9.4f}{r2:>7.3f}')
        report['contraction'][f'exec{execute}'] = {
            'c': c, 'b': b, 'fixed_point': fixed, 'r2': r2,
            'mean_trace': [round(float(v), 5) for v in m],
        }

    print('\n=== 2. refinement past trained depth ===')
    costs, arrivals, moves, violations, scores = refinement_trace(
        controller, model, batch, device, args.max_k, density, c95
    )
    print(f'{"k":>3}{"terminal":>11}{"arrival":>10}{"dJ":>10}{"|plan chg|":>12}'
          f'{"support r":>11}{"viol%":>8}')
    for k in range(len(costs)):
        dj = '' if k == 0 else f'{costs[k - 1] - costs[k]:+.5f}'
        sup = f'{scores[k]:>11.4f}' if scores else f'{"-":>11}'
        vio = f'{100 * violations[k]:>8.1f}' if violations else f'{"-":>8}'
        print(f'{k:>3}{costs[k]:>11.5f}{arrivals[k]:>10.5f}{dj:>10}'
              f'{moves[k]:>12.5f}{sup}{vio}')
    report['refinement'] = {
        'terminal': costs, 'arrival': arrivals, 'plan_change': moves,
        'support_score': scores, 'violation': violations, 'c95': c95,
    }

    out = Path(args.out)
    out.mkdir(parents=True, exist_ok=True)
    with (out / 'diagnostics.jsonl').open('a') as f:
        f.write(json.dumps(report) + '\n')
    print(f'\nwrote {out / "diagnostics.jsonl"}')


if __name__ == '__main__':
    main()