"""Local-only live path and prompt replanning for the SEED MEI138 demo runtime. No source or model patch is installed. Convert an unstarted Runtime with ``WindowRuntime.from_runtime(runtime)``, then use its normal start/recover/close API. Before each step, simulate ``plan_count`` native path rows from the last COMMITTED path state using the latest controls, and call step_replanned. Advance the caller's committed path state by one row only when a frame returns. """ from __future__ import annotations import numpy as np import torch from space.inference import Runtime class WindowRuntime(Runtime): """Rewrite all pending path/text conditions before the next denoiser call. The fixed 30-step / 30-frame scheduler has at most 29 already submitted, uncommitted rows. A plan includes those rows followed by ONE new row. Rows are native [yaw radians/frame, local X metres/frame, local Z metres/frame], beginning at absolute output frame ``self._frames``. Calls must be serialized just like Runtime. Existing pose state, scheduler, committed text history, CFG and the denoiser graph wrapper remain intact. By default a new prompt applies to every uncommitted frame immediately. This does not redo earlier denoising steps; it changes the conditions read by the next step. Use replan_text=False for append-only text conditioning. """ window_frames = 30 @classmethod def from_runtime(cls, runtime): """Take the model from an unstarted base runtime; discard that wrapper. The constructor does not reload weights or allocate another model. Refuse a started wrapper so no live graphs or recovery state can be lost. """ if runtime._ready or runtime._graph is not None: raise RuntimeError("Convert the runtime before calling start") return cls(runtime.model, runtime._recovery_type, runtime.metadata) @property def pending_count(self): """Already submitted path rows that have not produced output yet.""" return self._steps - self._frames @property def plan_count(self): """Required row count for the next step_replanned call (1 through 30).""" return self.pending_count + 1 def _pending_bounds(self): if not self._ready: raise RuntimeError("Call start inside the GPU allocation first") model = self.model scheduler = model.time_scheduler if scheduler.chunk_size != self.window_frames or scheduler.steps != self.window_frames: raise RuntimeError("Window replanning requires the 30-step / 30-frame schedule") start, end = model.current_commit, model.condition_frames if model.batch_size != 1 or model._stream_position.shape != (1, model.buf_len, 3): raise RuntimeError("Window replanning requires one MEI138 stream with three root channels") if not 0 <= start <= end <= model.buf_len: raise RuntimeError("Invalid position-buffer bounds") if end - start != self.pending_count or not 0 <= self.pending_count < self.window_frames: raise RuntimeError("Runtime and model pending-frame counters disagree") # Ring rollback subtracts seq_len from both local counters, but not the # wrapper's absolute counters. These two offsets must remain identical. offset = self._steps - end if offset < 0 or offset != self._frames - start or offset % model.seq_len: raise RuntimeError("Runtime and model ring-buffer offsets disagree") # This runtime never leaves denoised-but-unemitted frames between calls. # Refuse other scheduler/flush states rather than editing clean history. if model.current_step != end or start != max(0, end - self.window_frames + 1): raise RuntimeError("Stream is not at a completed single-frame step boundary") if model._pending_stream_position is not None: raise RuntimeError("A stream-condition transaction is already in progress") if not bool(model._stream_position_valid[:, :end].all()): raise RuntimeError("An accumulated path condition is missing") return start, end def _replanned_text_segments(self, prompt, start, end, replan_text): """Validate the text state and stage an edit without mutating it. Segment ends are implicit. Keep starts belonging to committed history verbatim, including negative starts retained for text RoPE by rollback. Extending an unchanged final prompt must not move its original start. """ text = self.model.text_module if not isinstance(replan_text, bool): raise TypeError("replan_text must be a bool") if text.stream_frames != [end] or len(text.stream_segments) != 1: raise RuntimeError("Text and motion stream counters disagree") segments = text.stream_segments[0] if "" not in text.text_cache: raise RuntimeError("The unconditional text features are missing") if (end == 0 and segments) or (end > 0 and not segments): raise RuntimeError("Text segments do not cover the accumulated frames") previous = None for segment in segments: if not isinstance(segment, (tuple, list)) or len(segment) != 2: raise RuntimeError("Invalid text segment") name, index = segment if (not isinstance(name, str) or name not in text.text_cache or not isinstance(index, (int, np.integer)) or isinstance(index, bool) or index >= end or (previous is not None and index <= previous)): raise RuntimeError("Text segment ordering or encoded features are invalid") previous = index if segments and segments[0][1] > 0: raise RuntimeError("Text segments leave accumulated frames uncovered") if not replan_text or start == end: return None revised = [segment for segment in segments if segment[1] < start] if not revised or revised[-1][0] != prompt: revised.append((prompt, start)) return revised @torch.inference_mode() def step_replanned(self, prompt: str, rows, *, replan_text: bool = True): """Replace pending path/text, append one frame, and denoise once. ``rows`` must be finite CPU data of shape (plan_count, 3). Its first pending_count rows replace buffer[current_commit:condition_frames]; the last row is submitted through the official stream transaction. The original frame-zero convention is preserved until frame zero is committed. Returns one completed (138,) numpy frame or None at startup. With replan_text=True, prompt also replaces every pending frame's text. Committed text segments retain their masks and original RoPE starts. False keeps prior text on pending frames and appends prompt only once. Validation completes before mutation. As with the base runtime, a model execution failure requires restarting the session before reuse. """ start, end = self._pending_bounds() if not isinstance(prompt, str) or prompt not in self.model.text_module.text_cache: raise ValueError("Prompt has no encoded features; call add_prompt_features first") if isinstance(rows, torch.Tensor): if rows.device.type != "cpu": raise ValueError("Replanned rows must be CPU data") rows = rows.detach().to(dtype=torch.float32).numpy() plan = np.array(rows, dtype=np.float32, copy=True) if plan.shape != (self.plan_count, 3) or not np.isfinite(plan).all(): raise ValueError(f"Replanned root control must be finite ({self.plan_count}, 3) rows") if self._frames == 0 and getattr(self.model, "representation", "mei138") != "humanml3d263": plan[0] = 0.0 model = self.model revised_text = self._replanned_text_segments(prompt, start, end, replan_text) device_plan = torch.as_tensor(plan, device=model._stream_position.device) inputs = {model.input_keys["text"]: [prompt], "position": device_plan[-1:].contiguous()} # In-place values preserve the buffer identity and validity flags. If # the next call wraps the ring, its official rollback shifts these # freshly revised values with the pose/text state before denoising. model._stream_position[:, start:end].copy_(device_plan[:-1].unsqueeze(0)) if revised_text is not None: model.text_module.stream_segments[0] = revised_text result = model.stream_generate_step(inputs)["generated"][0] self._steps += 1 self._action = None self._prompt = prompt if len(result) == 0: return None feature_dim = getattr(model, "model_input_dim", 138) if tuple(result.shape) != (1, feature_dim): raise RuntimeError(f"Expected one committed {feature_dim}D frame, got {tuple(result.shape)}") frame = result[0].float().cpu().numpy().copy() if not np.isfinite(frame).all(): raise RuntimeError("Model produced a non-finite motion frame") self._frames += 1 return frame def status(self): status = super().status() status.update(path_control="replan-uncommitted-window", replanning_window_frames=self.window_frames, plan_rows=self.plan_count, text_replanning_default=True) return status