Download code/rollout.py from nottygian/Leplanner: direct link, hf CLI and curl.
- Browser
- Download file 2.7 kB
-
https://huggingface.co/nottygian/Leplanner/resolve/main/code/rollout.py
- Command line
-
hf download hf://nottygian/Leplanner/code/rollout.py
-
curl -L -o rollout.py https://huggingface.co/nottygian/Leplanner/resolve/main/code/rollout.py
2.7 kB
| """Differentiable rollout of an action plan through the frozen LeWM predictor. | |
| Mirrors ``LeWM.rollout`` semantics exactly, but operates on cached latents | |
| instead of pixels and keeps the graph intact so gradients reach the plan. | |
| The action/frame alignment is the subtle part and follows ``lewm.py:83-87``: | |
| the action block at index ``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. | |
| """ | |
| import torch | |
| def rollout_plan(model, ctx_emb, past_actions, plan_actions, return_frames=False): | |
| """Roll a plan through the frozen predictor, autoregressively in latent space. | |
| Args: | |
| model: The frozen ``LeWM``. | |
| ctx_emb: ``(B, N, D)`` encoded context frames. | |
| past_actions: ``(B, N-1, 5*d_a)`` normalized blocks between them. | |
| plan_actions: ``(B, H, 5*d_a)`` normalized candidate blocks. | |
| return_frames: Also return the full ``(B, N+H, D)`` frame sequence, | |
| which the support loss slices into per-step context windows. | |
| Returns: | |
| ``(B, H, D)`` predicted latents, and optionally the full sequence. | |
| """ | |
| B, N, _ = ctx_emb.shape | |
| H = plan_actions.size(1) | |
| history_size = model.predictor.num_frames | |
| act_emb = model.action_encoder( | |
| torch.cat([past_actions, plan_actions], dim=1) | |
| ) # (B, N-1+H, A); index k = the block leaving frame k | |
| frames = list(ctx_emb.unbind(dim=1)) | |
| for t in range(H): | |
| lo = max(0, N + t - history_size) | |
| emb_win = torch.stack(frames[lo:], dim=1) | |
| act_win = act_emb[:, lo : N + t] | |
| frames.append(model.predict(emb_win, act_win)[:, -1]) | |
| pred = torch.stack(frames[N:], dim=1) | |
| if return_frames: | |
| return pred, torch.stack(frames, dim=1) | |
| return pred | |
| def rollout_contexts(frames, num_context): | |
| """Context window seen by each plan step, as a single batched tensor. | |
| Args: | |
| frames: ``(B, N+H, D)`` full rollout sequence from ``rollout_plan``. | |
| num_context: Window size ``N``. | |
| Returns: | |
| ``(B, H, N, D)`` where entry ``j`` is the window preceding block ``j``. | |
| """ | |
| H = frames.size(1) - num_context | |
| return torch.stack( | |
| [frames[:, j : j + num_context] for j in range(H)], dim=1 | |
| ) | |
| def goal_distance(pred, goal): | |
| """Per-step latent goal distance ``d_j = ||x_j - x_G||^2 / D``. | |
| Args: | |
| pred: ``(B, H, D)`` predicted latents. | |
| goal: ``(B, D)`` goal latent. | |
| Returns: | |
| ``(B, H)`` distances. | |
| """ | |
| return (pred - goal.unsqueeze(1)).pow(2).mean(dim=-1) | |