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())
|