File size: 11,487 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
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
"""Tests for tools/history_policy.HistoryPolicy.



Runs the real PushT env pool behind a spy solver that records exactly what the

solver was handed at each replan, so the assertions are on the contract the

planner actually sees rather than on the buffer internals.



    python scripts/test_history_policy.py

"""

import os
import sys
from pathlib import Path

os.environ.setdefault('MUJOCO_GL', 'egl')
os.environ.setdefault('PYTHONUTF8', '1')

import numpy as np  # noqa: E402
import torch  # noqa: E402

REPO = Path(__file__).resolve().parents[1]
os.environ.setdefault('STABLEWM_HOME', str(REPO / 'data' / 'swm_home'))
sys.path.insert(0, str(REPO))

import stable_worldmodel as swm  # noqa: E402
from sklearn.preprocessing import StandardScaler  # noqa: E402

from tools.history_policy import HistoryPolicy  # noqa: E402

N_ENVS, BLOCK, HORIZON, N_FRAMES = 2, 5, 5, 3
FAILURES = []


def check(name, ok, detail=''):
    print(f'  {"PASS" if ok else "FAIL"}  {name}{"  -- " + detail if detail else ""}')
    if not ok:
        FAILURES.append(name)


class SpySolver:
    """Records every solve, returns a deterministic non-zero plan."""

    def __init__(self):
        self.seen = []
        self._n_envs = N_ENVS
        self._horizon = HORIZON
        self._action_block = BLOCK
        self._action_dim = 2

    def configure(self, *, action_space, n_envs, config):
        self._n_envs = n_envs
        self._horizon = config.horizon
        self._action_block = config.action_block
        self._action_dim = int(action_space.shape[-1])

    @property
    def action_dim(self):
        return self._action_dim * self._action_block

    @property
    def n_envs(self):
        return self._n_envs

    @property
    def horizon(self):
        return self._horizon

    def solve(self, info_dict, init_action=None):
        px = info_dict['pixels']
        hist = info_dict.get('action_history')
        b = px.shape[0]
        self.seen.append({
            'pixels_shape': tuple(px.shape),
            'n_frames': int(px.shape[1]) if px.ndim == 5 else 1,
            'frames_identical': bool(
                px.ndim == 5 and px.shape[1] > 1
                and torch.allclose(px[:, 0], px[:, -1])
            ),
            'has_action_history': hist is not None,
            'action_history': None if hist is None else hist.clone(),
        })
        # a distinctive, non-constant plan so the recorded blocks are checkable
        step = len(self.seen)
        plan = torch.full((b, self._horizon, self.action_dim), 0.1 * step)
        return {'actions': plan, 'costs': torch.zeros(b)}

    __call__ = solve


def build(use_frames, use_actions, scaler):
    world = swm.World(
        env_name='swm/PushT-v1', num_envs=N_ENVS, image_shape=(224, 224),
        max_episode_steps=500,
    )
    config = swm.PlanConfig(
        horizon=HORIZON, receding_horizon=1, action_block=BLOCK,
        history_len=N_FRAMES,
    )
    spy = SpySolver()
    policy = HistoryPolicy(
        solver=spy, config=config, process={'action': scaler}, transform={},
        use_frames=use_frames, use_actions=use_actions,
    )
    world.set_policy(policy)
    world.reset(seed=0)
    return world, spy, policy


def run(world, steps):
    for _ in range(steps):
        world.envs.step(world._get_actions())


def main():
    scaler = StandardScaler().fit(np.random.RandomState(0).randn(500, 2) * 0.4)

    print('\n=== history ON (the fix) ===')
    world, spy, policy = build(True, True, scaler)
    run(world, 16)

    n_frames = [s['n_frames'] for s in spy.seen[:3]]
    check('frame history grows 1 -> 2 -> 3 over the first three replans',
          n_frames == [1, 2, 3], f'got {n_frames}')
    check('saturates at history_len and never exceeds it',
          all(s['n_frames'] <= N_FRAMES for s in spy.seen),
          f'max {max(s["n_frames"] for s in spy.seen)}')
    check('the three frames are genuinely different, not one repeated',
          not spy.seen[2]['frames_identical'])
    check('pixels keep the (B, T, H, W, C) layout the solvers expect',
          len(spy.seen[2]['pixels_shape']) == 5,
          str(spy.seen[2]['pixels_shape']))

    check('no action_history at the first replan (nothing executed yet)',
          not spy.seen[0]['has_action_history'])
    check('action_history present from the second replan',
          spy.seen[1]['has_action_history'])
    hist = spy.seen[2]['action_history']
    check('action_history is (B, N-1, block*action_dim)',
          tuple(hist.shape) == (N_ENVS, N_FRAMES - 1, BLOCK * 2),
          str(tuple(hist.shape)))

    # Solvers emit NORMALIZED actions; WorldModelPolicy inverse_transforms them
    # to raw env units before stepping. _record must undo exactly that, so the
    # stored block equals what the solver emitted (0.1 from solve #1). Two
    # assertions: the round trip is identity, and it is not vacuous -- the raw
    # value it passed through is genuinely different.
    got = hist[0, 0].view(BLOCK, 2)
    raw = torch.tensor(
        scaler.inverse_transform(np.full((1, 2), 0.1, np.float32))[0]
    ).float()
    check('recorded blocks round-trip to the normalized units the solver emitted',
          torch.allclose(got, torch.full((BLOCK, 2), 0.1), atol=1e-5),
          f'got {got[0].tolist()} want [0.1, 0.1]')
    check('and are NOT the raw env units (the transform is not a no-op)',
          not torch.allclose(got[0], raw, atol=1e-3),
          f'normalized {got[0].tolist()} vs raw {raw.tolist()}')
    check('the two past blocks differ (they came from different plans)',
          not torch.allclose(hist[0, 0], hist[0, 1]))
    check('no NaN reaches the solver', bool(torch.isfinite(hist).all()))

    check('world.infos was not mutated -- pixels still (B, 1, ...) for video/goal',
          world.infos['pixels'].shape[1] == 1,
          str(tuple(world.infos['pixels'].shape)))
    check('frames sampled on the action_block stride, not the replan cadence',
          policy._t % BLOCK == 16 % BLOCK and len(policy._frames[0]) == N_FRAMES)

    print('\n=== history OFF (--legacy-history) ===')
    world2, spy2, _ = build(False, False, scaler)
    run(world2, 16)
    check('legacy path still hands the solver a single frame',
          all(s['n_frames'] == 1 for s in spy2.seen),
          f'{[s["n_frames"] for s in spy2.seen]}')
    check('legacy path never supplies action_history',
          not any(s['has_action_history'] for s in spy2.seen))

    print('\n=== rh=5: replans are 25 steps apart, frames must still be 5 apart ===')
    world3 = swm.World(env_name='swm/PushT-v1', num_envs=N_ENVS,
                       image_shape=(224, 224), max_episode_steps=500)
    cfg3 = swm.PlanConfig(horizon=HORIZON, receding_horizon=5,
                          action_block=BLOCK, history_len=N_FRAMES)
    spy3 = SpySolver()
    p3 = HistoryPolicy(solver=spy3, config=cfg3, process={'action': scaler},
                       transform={}, use_frames=True, use_actions=True)
    world3.set_policy(p3)
    world3.reset(seed=0)
    run(world3, 30)
    check('only 2 solves in 30 steps at rh=5', len(spy3.seen) == 2,
          f'{len(spy3.seen)} solves')
    check('second solve still gets 3 distinct frames (stride survived)',
          spy3.seen[1]['n_frames'] == N_FRAMES
          and not spy3.seen[1]['frames_identical'],
          f'n_frames={spy3.seen[1]["n_frames"]}')

    integration(scaler)

    print()
    if FAILURES:
        print(f'{len(FAILURES)} FAILED: {FAILURES}')
        return 1
    print('all checks passed')
    return 0


def integration(scaler):
    """The real PlannerSolver on the real LeWM, driven by HistoryPolicy.



    The unit checks above prove the buffers are correct. This proves the

    contract holds where it matters: that PlannerSolver stops hitting its

    padding branch and actually encodes three distinct frames plus real blocks.

    """
    print('\n=== integration: real planner, real world model ===')
    ckpt_path = REPO / 'data/runs/planner_D10/planner.pt'
    if not ckpt_path.exists():
        print(f'  SKIP  no checkpoint at {ckpt_path}')
        return

    from lejepa_control.world_model import load_lewm
    from lejepa_control_2.solver import PlannerSolver, load_planner

    model = load_lewm(device='cpu')
    model.interpolate_pos_encoding = True
    planner, _ = load_planner(str(ckpt_path), device='cpu', horizon=HORIZON)

    seen = []

    class RecordingPlannerSolver(PlannerSolver):
        # PlannerSolver binds `__call__ = solve` at class level, so overriding
        # `solve` alone would never be reached. Same reason eval_planner.py's
        # RecordingSolver overrides __call__.
        def __call__(self, info_dict, init_action=None):
            px = info_dict['pixels']
            ctx = self._encode(px if px.ndim == 5 else px.unsqueeze(1))
            past = info_dict.get('action_history')
            seen.append({
                'n_frames': int(px.shape[1]) if px.ndim == 5 else 1,
                # distance between the oldest and newest encoded frame: zero
                # means the padding branch fired
                'ctx_spread': float((ctx[:, 0] - ctx[:, -1]).pow(2).mean()),
                'past_absmean': None if past is None else float(past.abs().mean()),
            })
            return super().__call__(info_dict, init_action)

    world = swm.World(env_name='swm/PushT-v1', num_envs=N_ENVS,
                      image_shape=(224, 224), max_episode_steps=500)
    config = swm.PlanConfig(horizon=HORIZON, receding_horizon=1,
                            action_block=BLOCK,
                            history_len=model.predictor.num_frames)
    solver = RecordingPlannerSolver(model, planner, device='cpu')
    policy = HistoryPolicy(
        solver=solver, config=config,
        process={'action': scaler},
        transform={'pixels': img_tf(), 'goal': img_tf()},
        use_frames=True, use_actions=True,
    )
    world.set_policy(policy)
    world.reset(seed=0)
    # episodic mode carries no goal; inject a fixed one so the solver can run
    world.infos['goal'] = world.infos['pixels'].copy()

    for _ in range(12):
        world.envs.step(world._get_actions())

    last = seen[-1]
    check('PlannerSolver receives 3 frames', last['n_frames'] == N_FRAMES,
          f'{last["n_frames"]}')
    check('encoded context is NOT degenerate (padding branch did not fire)',
          last['ctx_spread'] > 1e-6, f'||h_first - h_last||^2/D = {last["ctx_spread"]:.6f}')
    check('past_actions reach the planner and are non-zero',
          last['past_absmean'] is not None and last['past_absmean'] > 0,
          f'mean|block| = {last["past_absmean"]}')
    first = seen[0]
    check('first solve of the episode still degrades gracefully',
          first['n_frames'] == 1 and first['past_absmean'] is None)


def img_tf(size=224):
    import stable_pretraining as spt
    from torchvision.transforms import v2 as transforms
    return transforms.Compose([
        transforms.ToImage(),
        transforms.ToDtype(torch.float32, scale=True),
        transforms.Normalize(**spt.data.dataset_stats.ImageNet),
        transforms.Resize(size=size),
    ])


if __name__ == '__main__':
    raise SystemExit(main())