Spaces:
Running on Zero
Running on Zero
File size: 9,575 Bytes
9a25493 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 | """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
|