Download code/planner.py from SaltedLemon/lejepa-control-pusht: direct link, hf CLI and curl.
- Browser
- Download file 21.4 kB
-
https://huggingface.co/SaltedLemon/lejepa-control-pusht/resolve/main/code/planner.py
- Command line
-
hf download hf://SaltedLemon/lejepa-control-pusht/code/planner.py
-
curl -L -o planner.py https://huggingface.co/SaltedLemon/lejepa-control-pusht/resolve/main/code/planner.py
21.4 kB
| """Recursive latent planner over the frozen LeWM PushT world model. | |
| Implements the HRM/TRM-style nested recursion of ``recursive_planner_design.pdf`` | |
| section 9, placed inside an MPC rollout:: | |
| for k = 1 .. H: # outer - imagined environment steps | |
| for j = 1 .. T: # middle - answer (action) improvement | |
| for i = 1 .. n: # inner - latent reasoning refinement | |
| z = f(z, h_G, h_{k-1}, y, c) | |
| y = g(y, z, h_G) | |
| c = consequence(M(C, psi(y)), h_G) # Change 2 | |
| b_k = psi(y); h_k = M(h_{k-1}, b_k) # frozen model advances | |
| ``f`` and ``g`` are weight-tied across every ``i``, ``j`` and ``k`` — the | |
| recursion buys depth and compute, not parameters. | |
| Three things in here are load-bearing and easy to get silently wrong: | |
| * **The gradient policy (Change 1).** Cycles ``1..T-1`` run under ``no_grad``; | |
| only cycle ``T`` is differentiated, and ``(y, z)`` are detached when crossing | |
| from horizon step ``k`` to ``k+1``. The chain of world-model states ``h_k`` is | |
| *never* detached — that chain is the entire planning signal, and cutting it | |
| leaves a loss that still falls while the controller quietly becomes greedy. | |
| * **Action/frame alignment.** Block ``k`` is the block *leaving* context frame | |
| ``k``. With ``N`` context frames there are ``N-1`` past blocks between them, | |
| and the current frame pairs with the first block of the plan. This mirrors | |
| ``lejepa_control.rollout.rollout_plan`` and ``LeWM.rollout`` exactly. | |
| * **Recursion stability (Change 11).** RMSNorm on ``z`` and ``y`` before every | |
| ``f``/``g``, a bounded residual gate on each update, and ``h_{k-1}``/``h_G`` | |
| re-injected at *every* ``f`` application rather than only the first. | |
| """ | |
| import torch | |
| from torch import nn | |
| __all__ = [ | |
| 'ActionEmbedding', | |
| 'FTheta', | |
| 'GTheta', | |
| 'RMSNorm', | |
| 'RecursivePlanner', | |
| 'consequence_features', | |
| 'm_step', | |
| ] | |
| class RMSNorm(nn.Module): | |
| """Root-mean-square layer norm, no mean subtraction. | |
| Change 11's first guard. A weight-tied map applied ~90 times has no reason | |
| to be norm-preserving; a mild 5% growth per step compounds to 80x over one | |
| training step's recursion, and long before that ``f`` sees inputs outside | |
| the range its weights were fit for. | |
| """ | |
| def __init__(self, dim, eps=1e-6): | |
| super().__init__() | |
| self.eps = eps | |
| self.weight = nn.Parameter(torch.ones(dim)) | |
| def forward(self, x): | |
| scale = torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) | |
| return x * scale * self.weight | |
| class ActionEmbedding(nn.Module): | |
| """The ``phi`` / ``psi`` pair: action block <-> answer space. | |
| ``phi`` encodes a real 10-dim block into the ``W``-dim answer space, ``psi`` | |
| decodes an answer back to a tanh-bounded block. Pre-train the pair as a | |
| plain autoencoder on dataset blocks (Change 7), then keep both trainable | |
| with the round-trip anchor holding them consistent. | |
| Args: | |
| block_dim: ``frameskip * action_dim`` (10 for PushT). | |
| width: Answer-space width ``W``. | |
| hidden: MLP hidden width. | |
| action_dim: Native env action dim (2). | |
| frameskip: Env actions per world-model transition (5). | |
| action_center / action_scale: Per-dim tanh bounds expressed in the | |
| *normalized* action units the world model was trained on. For raw | |
| PushT actions in ``[-1, 1]`` these are ``-mean/std`` and ``1/std``. | |
| """ | |
| def __init__( | |
| self, | |
| block_dim=10, | |
| width=256, | |
| hidden=256, | |
| action_dim=2, | |
| frameskip=5, | |
| action_center=0.0, | |
| action_scale=1.0, | |
| ): | |
| super().__init__() | |
| self.block_dim = block_dim | |
| self.width = width | |
| self.action_dim = action_dim | |
| self.frameskip = frameskip | |
| center = torch.as_tensor(action_center).float().expand(action_dim) | |
| scale = torch.as_tensor(action_scale).float().expand(action_dim) | |
| self.register_buffer('action_center', center.clone()) | |
| self.register_buffer('action_scale', scale.clone()) | |
| self.encode_net = nn.Sequential( | |
| nn.Linear(block_dim, hidden), | |
| nn.GELU(), | |
| nn.Linear(hidden, width), | |
| ) | |
| self.decode_net = nn.Sequential( | |
| RMSNorm(width), | |
| nn.Linear(width, hidden), | |
| nn.GELU(), | |
| nn.Linear(hidden, block_dim), | |
| ) | |
| def encode(self, block): | |
| """``phi``: ``(..., A)`` normalized block -> ``(..., W)`` answer.""" | |
| return self.encode_net(block) | |
| def bound(self, raw): | |
| """Map pre-tanh activations to a valid normalized action block.""" | |
| r = raw.unflatten(-1, (self.frameskip, self.action_dim)) | |
| return (self.action_center + self.action_scale * torch.tanh(r)).flatten(-2) | |
| def decode(self, y, return_raw=False): | |
| """``psi``: ``(..., W)`` answer -> ``(..., A)`` bounded block. | |
| ``return_raw`` also yields the pre-tanh activations, which is what the | |
| saturation barrier (Change 6) penalizes — the barrier has to act before | |
| the tanh or it cannot reach a dimension that has already frozen. | |
| """ | |
| raw = self.decode_net(y) | |
| block = self.bound(raw) | |
| return (block, raw) if return_raw else block | |
| def round_trip(self, y): | |
| """``phi(psi(y))`` — the manifold anchor's prediction of ``y``.""" | |
| return self.encode(self.decode(y)) | |
| class _GatedUpdate(nn.Module): | |
| """Shared body for ``f`` and ``g``: normalize, condition, gated residual. | |
| The update is ``x <- x + sigmoid(eta) * Delta(...)`` with ``Delta``'s output | |
| layer initialized small, so the recursion starts near-identity and cannot | |
| destroy a good answer in early training (Change 11's second guard). | |
| """ | |
| def __init__(self, width, cond_dim, hidden, gate_init=0.0, out_std=0.01): | |
| super().__init__() | |
| self.norm_state = RMSNorm(width) | |
| self.norm_other = RMSNorm(width) | |
| self.cond_proj = nn.Linear(cond_dim, width) | |
| self.net = nn.Sequential( | |
| nn.Linear(3 * width, hidden), | |
| nn.GELU(), | |
| nn.Linear(hidden, hidden), | |
| nn.GELU(), | |
| ) | |
| self.out = nn.Linear(hidden, width) | |
| nn.init.normal_(self.out.weight, std=out_std) | |
| nn.init.zeros_(self.out.bias) | |
| self.gate = nn.Parameter(torch.tensor(float(gate_init))) | |
| # counts applications made with grad enabled; the Change-1 detach | |
| # schedule is verified against this, see RecursivePlanner.grad_calls | |
| self.grad_calls = 0 | |
| def forward(self, state, other, cond): | |
| if torch.is_grad_enabled(): | |
| self.grad_calls += 1 | |
| x = torch.cat( | |
| [ | |
| self.norm_state(state), | |
| self.norm_other(other), | |
| self.cond_proj(cond), | |
| ], | |
| dim=-1, | |
| ) | |
| return state + torch.sigmoid(self.gate) * self.out(self.net(x)) | |
| def gate_value(self): | |
| with torch.no_grad(): | |
| return torch.sigmoid(self.gate).item() | |
| class FTheta(nn.Module): | |
| """Inner-loop latent reasoning update ``z <- f(z, h_prev, h_G, y, c)``. | |
| Conditioning is ``[h_{k-1}, h_G, c]`` where ``c`` is the consequence | |
| feature from Change 2 — ``[h_hat, h_hat - h_G, ||h_hat - h_G||^2 / D]``. | |
| All of it is re-supplied at every application, not just the first, so the | |
| recursion cannot drift away from the question it was asked. | |
| """ | |
| def __init__(self, width=256, latent_dim=192, hidden=512, gate_init=0.0): | |
| super().__init__() | |
| self.latent_dim = latent_dim | |
| # h_prev (D) + h_G (D) + consequence (2D + 1) | |
| cond_dim = 4 * latent_dim + 1 | |
| self.body = _GatedUpdate(width, cond_dim, hidden, gate_init) | |
| def forward(self, z, h_prev, h_goal, y, consequence): | |
| cond = torch.cat([h_prev, h_goal, consequence], dim=-1) | |
| return self.body(z, y, cond) | |
| class GTheta(nn.Module): | |
| """Middle-loop answer update ``y <- g(y, z, h_G)``. | |
| ``g`` deliberately does not see ``h_{k-1}``: the current state reaches the | |
| answer only through ``z``. That is TRM's convention and it is what makes | |
| ``z`` a scratchpad rather than a redundant conditioning path. | |
| """ | |
| def __init__(self, width=256, latent_dim=192, hidden=512, gate_init=0.0): | |
| super().__init__() | |
| self.body = _GatedUpdate(width, latent_dim, hidden, gate_init) | |
| def forward(self, y, z, h_goal): | |
| return self.body(y, z, h_goal) | |
| def consequence_features(h_hat, h_goal): | |
| """Change 2's feedback vector: ``[h_hat, h_hat - h_G, ||.||^2 / D]``. | |
| Gives ``f`` an error vector in the same space it is trying to shrink, which | |
| is what turns the middle loop from an open-loop guesser into a corrector. | |
| """ | |
| delta = h_hat - h_goal | |
| dist = delta.pow(2).mean(dim=-1, keepdim=True) | |
| return torch.cat([h_hat, delta, dist], dim=-1) | |
| def m_step(model, frames, blocks, num_context): | |
| """One frozen-world-model transition, matching ``rollout_plan``'s windows. | |
| At rollout step ``t`` the predictor consumes frames ``[t, t+N)`` and the | |
| action blocks leaving those same frames. Because ``frames`` has ``N+t`` | |
| entries and ``blocks`` has ``N+t`` entries once the step's block is | |
| appended, both are just the trailing ``N``. | |
| Args: | |
| model: The frozen ``LeWM``. Gradients flow *through* it, never into it. | |
| frames: List of ``(B, D)`` latents, oldest first, length ``N + t``. | |
| blocks: List of ``(B, A)`` normalized blocks, length ``N + t``, where | |
| entry ``i`` is the block leaving ``frames[i]``. | |
| num_context: ``N``. | |
| Returns: | |
| ``(B, D)`` the predicted next latent. | |
| """ | |
| n = num_context | |
| assert len(blocks) == len(frames), ( | |
| f'alignment: {len(frames)} frames but {len(blocks)} blocks; block i ' | |
| f'must be the block leaving frame i' | |
| ) | |
| emb_win = torch.stack(frames[-n:], dim=1) | |
| act_win = model.action_encoder(torch.stack(blocks[-n:], dim=1)) | |
| return model.predict(emb_win, act_win)[:, -1] | |
| class RecursivePlanner(nn.Module): | |
| """Three-loop driver: inner ``n``, middle ``T``, outer ``H``. | |
| Args: | |
| latent_dim: World-model latent width ``D`` (192). | |
| num_context: Context frames the predictor consumes ``N`` (3). | |
| action_dim / frameskip: Native action dim and env steps per transition. | |
| width: Recursion width ``W`` (256). | |
| hidden: MLP hidden width inside ``f`` and ``g``. | |
| inner: ``n``, latent refinements per cycle (6). | |
| cycles: ``T``, answer revisions per horizon step (3; 1 during stage A). | |
| horizon: ``H``, imagined lookahead steps (3 -> 5 curriculum). | |
| use_feedback: Change 2. When off there is no per-cycle lookahead, so | |
| no per-cycle distances are produced and deep supervision has | |
| nothing to score. | |
| warm_start: Change 9. Start ``y`` from the last executed block and | |
| carry the answer across horizon steps instead of resetting it. | |
| lambda_z: Change 9's ``lambda_z``. ``0`` means a fresh scratchpad | |
| ``z0`` at every horizon step, which is the documented default; | |
| larger values blend in the detached carried state. | |
| learn_lambda_z: Make ``lambda_z`` a learned scalar. | |
| action_center / action_scale: tanh bounds in normalized action units. | |
| """ | |
| def __init__( | |
| self, | |
| latent_dim=192, | |
| num_context=3, | |
| action_dim=2, | |
| frameskip=5, | |
| width=256, | |
| hidden=512, | |
| inner=6, | |
| cycles=3, | |
| horizon=5, | |
| use_feedback=True, | |
| warm_start=True, | |
| lambda_z=0.0, | |
| learn_lambda_z=False, | |
| gate_init=0.0, | |
| action_center=0.0, | |
| action_scale=1.0, | |
| detach_schedule='last-cycle', | |
| ): | |
| super().__init__() | |
| assert detach_schedule in ('last-cycle', 'one-step', 'full') | |
| self.detach_schedule = detach_schedule | |
| self.latent_dim = latent_dim | |
| self.num_context = num_context | |
| self.action_dim = action_dim | |
| self.frameskip = frameskip | |
| self.block_dim = frameskip * action_dim | |
| self.width = width | |
| self.inner = inner | |
| self.cycles = cycles | |
| self.horizon = horizon | |
| self.use_feedback = use_feedback | |
| self.warm_start = warm_start | |
| self.f = FTheta(width, latent_dim, hidden, gate_init) | |
| self.g = GTheta(width, latent_dim, hidden, gate_init) | |
| self.action_embed = ActionEmbedding( | |
| block_dim=self.block_dim, | |
| width=width, | |
| hidden=hidden // 2, | |
| action_dim=action_dim, | |
| frameskip=frameskip, | |
| action_center=action_center, | |
| action_scale=action_scale, | |
| ) | |
| # fresh scratchpad, and the cold-start answer when warm start is off | |
| self.z0 = nn.Parameter(torch.randn(1, width) * 0.02) | |
| self.y0 = nn.Parameter(torch.randn(1, width) * 0.02) | |
| self._capture = False | |
| self.taps = {'z_norms': [], 'grad_first': [], 'grad_last': []} | |
| if learn_lambda_z: | |
| self.lambda_z = nn.Parameter(torch.tensor(float(lambda_z))) | |
| else: | |
| self.register_buffer( | |
| 'lambda_z', torch.tensor(float(lambda_z)), persistent=True | |
| ) | |
| # -- diagnostics ------------------------------------------------------- | |
| def reset_call_counts(self): | |
| self.f.body.grad_calls = 0 | |
| self.g.body.grad_calls = 0 | |
| def start_capture(self): | |
| """Begin collecting the section-11 recursion diagnostics. | |
| Populates ``self.taps`` during the next forward/backward with: | |
| ``z_norms`` (should be flat across ``i`` — RMSNorm makes it so) and, | |
| after ``backward()``, ``grad_first`` / ``grad_last``, the gradient | |
| norms at ``f``'s first and last application inside the gradient cycle. | |
| A first/last ratio outside ~10x means the backprop depth is larger | |
| than the detach schedule intends. | |
| """ | |
| self.taps = {'z_norms': [], 'grad_first': [], 'grad_last': []} | |
| self._capture = True | |
| def stop_capture(self): | |
| self._capture = False | |
| def _tap_grad(self, tensor, key): | |
| if tensor.requires_grad: | |
| tensor.register_hook( | |
| lambda g, k=key: self.taps[k].append(g.norm().item()) | |
| ) | |
| def grad_calls(self): | |
| """``f`` and ``g`` applications made with grad enabled. | |
| Under the Change-1 schedule this is ``H * (n + 1)`` — one gradient | |
| cycle per horizon step. Without it, it is ``H * T * (n + 1)``, which | |
| is the failure the activation-count test exists to catch. | |
| """ | |
| return self.f.body.grad_calls + self.g.body.grad_calls | |
| def gate_values(self): | |
| return {'f': self.f.body.gate_value(), 'g': self.g.body.gate_value()} | |
| # -- the recursion ----------------------------------------------------- | |
| def _initial_answer(self, past_actions, batch): | |
| if self.warm_start: | |
| # Change 9: consecutive optimal blocks are highly correlated, so | |
| # phi of the last executed block is a free head start | |
| return self.action_embed.encode(past_actions[:, -1]) | |
| return self.y0.expand(batch, -1) | |
| def forward( | |
| self, | |
| model, | |
| ctx_emb, | |
| past_actions, | |
| goal_emb, | |
| horizon=None, | |
| cycles=None, | |
| inner=None, | |
| ): | |
| """Run one imagined rollout and return everything the loss needs. | |
| Args: | |
| model: Frozen ``LeWM``. | |
| ctx_emb: ``(B, N, D)`` context latents. | |
| past_actions: ``(B, N-1, A)`` normalized executed blocks. | |
| goal_emb: ``(B, D)`` goal latent. | |
| horizon / cycles / inner: Per-call overrides of ``H`` / ``T`` / | |
| ``n``, used by the curriculum and the anytime-inference sweep. | |
| Returns: | |
| Dict with ``distances`` ``(B, H)``, ``cycle_distances`` | |
| ``(B, H, T)`` or ``None``, ``blocks`` ``(B, H, A)``, ``raw`` | |
| ``(B, H, A)`` pre-tanh, ``contexts`` ``(B, H, N, D)``, ``answers`` | |
| ``(B, H, W)`` and ``frames`` ``(B, N+H, D)``. | |
| """ | |
| H = horizon or self.horizon | |
| T = cycles or self.cycles | |
| n = inner or self.inner | |
| N = self.num_context | |
| B = ctx_emb.size(0) | |
| assert ctx_emb.size(1) == N, f'expected {N} context frames' | |
| assert past_actions.size(1) == N - 1, ( | |
| f'expected {N - 1} past blocks between {N} context frames, got ' | |
| f'{past_actions.size(1)}' | |
| ) | |
| frames = list(ctx_emb.unbind(dim=1)) | |
| blocks = list(past_actions.unbind(dim=1)) | |
| y = self._initial_answer(past_actions, B) | |
| z = self.z0.expand(B, -1) | |
| # consequence_features width: [h_hat, h_hat - h_G, ||.||^2/D] = 2D + 1 | |
| zero_c = ctx_emb.new_zeros(B, 2 * self.latent_dim + 1) | |
| distances, cycle_d = [], [] | |
| out_blocks, out_raw, out_ctx, out_y = [], [], [], [] | |
| for _ in range(H): | |
| # the state the recursion plans *from*, read before it advances | |
| h_prev = frames[-1] | |
| context = torch.stack(frames[-N:], dim=1) | |
| # Change 1: no BPTT across horizon steps. The h-chain above stays | |
| # differentiable; only the recursion's own carry is cut. | |
| y = y.detach() | |
| # Change 9: z_k^0 = z0 + lambda_z * sg(z_{k-1}^n); lambda_z = 0 is | |
| # the documented default and gives a fresh scratchpad | |
| z = self.z0.expand(B, -1) + self.lambda_z * z.detach() | |
| c = zero_c | |
| step_cycles = [] | |
| # ---- cycles 1..T-1 : forward only ---------------------------- | |
| # 'full' keeps the graph across every cycle: the ablation-4 control | |
| # showing why Change 1 exists. It may simply not train. | |
| early_grad = self.detach_schedule == 'full' | |
| with torch.set_grad_enabled( | |
| early_grad and torch.is_grad_enabled() | |
| ): | |
| for _ in range(T - 1): | |
| for _ in range(n): | |
| z = self.f(z, h_prev, goal_emb, y, c) | |
| y = self.g(y, z, goal_emb) | |
| if self.use_feedback: | |
| # Change 2: what would this answer actually cause? | |
| b_hat = self.action_embed.decode(y) | |
| h_hat = m_step(model, frames, blocks + [b_hat], N) | |
| c = consequence_features(h_hat, goal_emb) | |
| step_cycles.append( | |
| (h_hat - goal_emb).pow(2).mean(dim=-1) | |
| ) | |
| if not early_grad: | |
| # the lookahead distances above are constants; keep them out of | |
| # the graph explicitly rather than relying on no_grad's scope | |
| step_cycles = [d.detach() for d in step_cycles] | |
| # ---- cycle T : the only one that carries gradient ------------- | |
| # 'one-step' is HRM's original approximation, kept as ablation 4's | |
| # middle rung: only the final f application carries gradient. | |
| head = n - 1 if self.detach_schedule == 'one-step' else 0 | |
| with torch.set_grad_enabled(False): | |
| for _ in range(head): | |
| z = self.f(z, h_prev, goal_emb, y, c) | |
| for i in range(head, n): | |
| z = self.f(z, h_prev, goal_emb, y, c) | |
| if self._capture: | |
| self.taps['z_norms'].append(z.detach().norm(dim=-1).mean().item()) | |
| if i == head: | |
| self._tap_grad(z, 'grad_first') | |
| if i == n - 1: | |
| self._tap_grad(z, 'grad_last') | |
| y = self.g(y, z, goal_emb) | |
| # ---- commit the action and advance the frozen world model ----- | |
| block, raw = self.action_embed.decode(y, return_raw=True) | |
| h = m_step(model, frames, blocks + [block], N) | |
| frames.append(h) | |
| blocks.append(block) | |
| d = (h - goal_emb).pow(2).mean(dim=-1) | |
| distances.append(d) | |
| step_cycles.append(d) | |
| cycle_d.append(torch.stack(step_cycles, dim=-1)) | |
| out_blocks.append(block) | |
| out_raw.append(raw) | |
| out_ctx.append(context) | |
| out_y.append(y) | |
| return { | |
| # pre-stack views. Indexing a stacked tensor produces a node that | |
| # is *downstream* of the committed blocks, so `autograd.grad(d[-1], | |
| # stacked[:, 0])` reports None even on a perfectly intact chain. | |
| # The section-9 invariants are stated on the lists for that reason, | |
| # and the tests assert against these. | |
| 'distance_seq': distances, | |
| 'block_seq': out_blocks, | |
| 'distances': torch.stack(distances, dim=1), | |
| 'cycle_distances': ( | |
| torch.stack(cycle_d, dim=1) if self.use_feedback else None | |
| ), | |
| 'blocks': torch.stack(out_blocks, dim=1), | |
| 'raw': torch.stack(out_raw, dim=1), | |
| 'contexts': torch.stack(out_ctx, dim=1), | |
| 'answers': torch.stack(out_y, dim=1), | |
| 'frames': torch.stack(frames, dim=1), | |
| } | |