lejepa-control-pusht / code /planner.py
SaltedLemon's picture
Upload code/planner.py with huggingface_hub
ec0d725 verified
Raw
History Blame Contribute Delete
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())
)
@property
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),
}