Download code/solver.py from nottygian/Leplanner: direct link, hf CLI and curl.
- Browser
- Download file 5.67 kB
-
https://huggingface.co/nottygian/Leplanner/resolve/main/code/solver.py
- Command line
-
hf download hf://nottygian/Leplanner/code/solver.py
-
curl -L -o solver.py https://huggingface.co/nottygian/Leplanner/resolve/main/code/solver.py
5.67 kB
| """Adapter exposing the trained controller as a ``stable_worldmodel`` solver. | |
| Implements the ``Solver`` protocol so the controller drops into | |
| ``WorldModelPolicy`` wherever ``CEMSolver`` goes, which makes the CEM | |
| comparison an apples-to-apples swap: same env, same wrappers, same | |
| preprocessing, same receding-horizon execution. | |
| """ | |
| import gymnasium as gym | |
| import torch | |
| from lejepa_control.controller import IterativeController | |
| class ControllerSolver: | |
| """Runs K amortized refinements instead of CEM's sampling loop. | |
| Args: | |
| model: The frozen ``LeWM``. | |
| controller: A trained :class:`IterativeController`. | |
| device: Device to plan on. | |
| refinements: Override the controller's K at eval time (for the | |
| success-vs-refinement-count ablation). ``None`` keeps the trained | |
| value. | |
| """ | |
| def __init__(self, model, controller, device='cuda', refinements=None): | |
| self.model = model | |
| self.controller = controller.to(device).eval() | |
| self.device = device | |
| self._refinements = refinements | |
| self._n_envs = 1 | |
| self._horizon = controller.horizon | |
| self._action_dim = controller.action_dim | |
| self._action_block = controller.frameskip | |
| def configure(self, *, action_space: gym.Space, n_envs: int, config) -> None: | |
| self._n_envs = n_envs | |
| self._horizon = config.horizon | |
| self._action_block = config.action_block | |
| self._action_dim = int(action_space.shape[-1]) | |
| assert self._horizon == self.controller.horizon, ( | |
| f'plan horizon {self._horizon} != controller horizon ' | |
| f'{self.controller.horizon}' | |
| ) | |
| assert self._action_block == self.controller.frameskip, ( | |
| f'action_block {self._action_block} != controller frameskip ' | |
| f'{self.controller.frameskip}' | |
| ) | |
| def action_dim(self) -> int: | |
| return self._action_dim * self._action_block | |
| def n_envs(self) -> int: | |
| return self._n_envs | |
| def horizon(self) -> int: | |
| return self._horizon | |
| def _encode(self, pixels): | |
| """Encode ``(B, T, C, H, W)`` frames to ``(B, T, D)`` latents.""" | |
| with torch.no_grad(): | |
| return self.model.encode({'pixels': pixels.to(self.device)})['emb'] | |
| def solve(self, info_dict: dict, init_action=None) -> dict: | |
| """Plan for every env in ``info_dict``; returns ``{'actions': ...}``. | |
| Expects ``pixels`` ``(B, T, C, H, W)``, ``goal`` ``(B, ...)`` and, | |
| when the policy carries history, ``action_history`` | |
| ``(B, T-1, block*d_a)``. | |
| """ | |
| pixels = info_dict['pixels'] | |
| if pixels.ndim == 4: # (B, C, H, W) -> single context frame | |
| pixels = pixels.unsqueeze(1) | |
| B, T = pixels.shape[:2] | |
| ctx = self._encode(pixels) | |
| goal = info_dict['goal'] | |
| if goal.ndim == 4: | |
| goal = goal.unsqueeze(1) | |
| goal_emb = self._encode(goal)[:, -1] | |
| num_context = self.controller.num_context | |
| if T < num_context: # early in an episode: repeat the oldest frame | |
| pad = ctx[:, :1].expand(B, num_context - T, -1) | |
| ctx = torch.cat([pad, ctx], dim=1) | |
| elif T > num_context: | |
| ctx = ctx[:, -num_context:] | |
| block_dim = self._action_block * self._action_dim | |
| past = info_dict.get('action_history') | |
| if past is None: | |
| past = ctx.new_zeros(B, num_context - 1, block_dim) | |
| else: | |
| past = past.to(self.device).float() | |
| if past.size(1) < num_context - 1: | |
| pad = past.new_zeros( | |
| B, num_context - 1 - past.size(1), block_dim | |
| ) | |
| past = torch.cat([pad, past], dim=1) | |
| else: | |
| past = past[:, -(num_context - 1) :] | |
| past = torch.nan_to_num(past, 0.0) | |
| k = self._refinements | |
| original = self.controller.refinements | |
| if k is not None: | |
| self.controller.refinements = k | |
| try: | |
| out = self.controller(self.model, ctx, past, goal_emb) | |
| finally: | |
| self.controller.refinements = original | |
| actions = out['plans'][-1] # (B, H, block*d_a) | |
| terminal = out['distances'][-1][:, -1] | |
| return { | |
| 'actions': actions.detach().float().cpu(), | |
| 'costs': terminal.detach().float().cpu(), | |
| 'terminal_distance': terminal.mean().item(), | |
| } | |
| __call__ = solve | |
| def load_controller(path, latent_dim=192, device='cuda', refinements=None): | |
| """Rebuild a controller from a training checkpoint.""" | |
| ckpt = torch.load(path, map_location=device, weights_only=False) | |
| saved = ckpt['args'] | |
| a_mean = torch.tensor(ckpt['action_mean']) | |
| a_std = torch.tensor(ckpt['action_std']) | |
| controller = IterativeController( | |
| latent_dim=latent_dim, | |
| horizon=saved['horizon'], | |
| refinements=saved['refinements'], | |
| width=saved['width'], | |
| depth=saved['depth'], | |
| heads=saved['heads'], | |
| dropout=saved['dropout'], | |
| action_center=(-a_mean / a_std), | |
| action_scale=(1.0 / a_std), | |
| no_latent_proj=saved.get('no_latent_proj', False), | |
| fused=saved.get('fused', False), | |
| ) | |
| controller.load_state_dict(ckpt['state_dict']) | |
| controller.to(device).eval() | |
| if refinements is not None: | |
| controller.refinements = refinements | |
| return controller, ckpt | |