"""Read-only capture of exact v2 trajectories and selected AR context KV.""" from __future__ import annotations from dataclasses import dataclass from typing import Any, Callable, Mapping import torch from .schema import LATENT_HEIGHT, LATENT_WIDTH, NUM_STEPS from .v2_schema import CONTEXT_BLOCK_IDS, block_key def _clone_cpu(tensor: torch.Tensor) -> torch.Tensor: return tensor.detach().to(device="cpu").contiguous() @dataclass class _ActiveStep: chunk_id: int step_id: int tensors: dict[str, torch.Tensor] shared: dict[str, torch.Tensor] class PredictorV2TeacherCapture: """Capture Full-DiT targets plus step-invariant selected-layer AR memory.""" def __init__( self, transformer: torch.nn.Module, *, on_chunk: Callable[ [int, dict[str, torch.Tensor], dict[int, dict[str, torch.Tensor]]], None ], num_steps: int = NUM_STEPS, context_block_ids: tuple[int, ...] = CONTEXT_BLOCK_IDS, ) -> None: self.transformer = transformer self.on_chunk = on_chunk self.num_steps = num_steps self.context_block_ids = tuple(context_block_ids) self.call_index = 0 self.active: _ActiveStep | None = None self.chunk_steps: list[_ActiveStep] = [] self.image_condition_latent: torch.Tensor | None = None self.text_context: dict[int, dict[str, torch.Tensor]] | None = None self.chunk_context: dict[int, dict[str, torch.Tensor]] | None = None self._handles: list[Any] = [] def __enter__(self) -> "PredictorV2TeacherCapture": self._handles.append( self.transformer.register_forward_pre_hook(self._transformer_pre, with_kwargs=True) ) self._handles.append( self.transformer.register_forward_hook(self._transformer_post, with_kwargs=True) ) self._handles.append( self.transformer.final_layer.register_forward_pre_hook(self._final_pre, with_kwargs=True) ) return self def __exit__(self, exc_type, exc, traceback) -> bool: for handle in self._handles: handle.remove() self._handles.clear() self.active = None return False def _capture_context( self, kv_cache: list[Mapping[str, torch.Tensor | None]], ) -> dict[int, dict[str, torch.Tensor]]: result: dict[int, dict[str, torch.Tensor]] = {} for block_id in self.context_block_ids: cache = kv_cache[block_id] k_txt = cache.get("k_txt") v_txt = cache.get("v_txt") if k_txt is None or v_txt is None: raise RuntimeError(f"Block {block_id} text KV is unavailable") k_vision = cache.get("k_vision") v_vision = cache.get("v_vision") if (k_vision is None) != (v_vision is None): raise RuntimeError(f"Block {block_id} has incomplete vision KV") if k_vision is None: empty_shape = (2, int(k_txt.shape[1]), 0, int(k_txt.shape[3])) k_vision_cpu = torch.empty(empty_shape, dtype=k_txt.dtype, device="cpu") v_vision_cpu = torch.empty(empty_shape, dtype=v_txt.dtype, device="cpu") else: k_vision_cpu = _clone_cpu(k_vision) v_vision_cpu = _clone_cpu(v_vision) result[block_id] = { "k_vision": k_vision_cpu, "v_vision": v_vision_cpu, } if self.text_context is None: result[block_id]["k_txt"] = _clone_cpu(k_txt) result[block_id]["v_txt"] = _clone_cpu(v_txt) if self.text_context is None: self.text_context = { block_id: { "k_txt": result[block_id].pop("k_txt"), "v_txt": result[block_id].pop("v_txt"), } for block_id in self.context_block_ids } return result def _transformer_pre(self, module, args, kwargs) -> None: is_denoise = ( kwargs.get("ar_vision_inference", False) and not kwargs.get("cache_vision", False) ) if not is_denoise: return if self.active is not None: raise RuntimeError("Nested denoising capture is not supported") chunk_id, step_id = divmod(self.call_index, self.num_steps) model_input = kwargs["hidden_states"] if model_input.shape[1] != 65: raise ValueError(f"Teacher denoising input must have 65 channels, got {model_input.shape}") if self.image_condition_latent is None: self.image_condition_latent = _clone_cpu(model_input[:, 32:64, 0:1]) mask = model_input[:, 64:65] if not torch.all(mask[:, :, 0] == 1) or not torch.all(mask[:, :, 1:] == 0): raise ValueError("Unexpected I2V condition mask in first chunk") if step_id == 0: if self.chunk_context is not None: raise RuntimeError("Previous chunk context was not flushed") self.chunk_context = self._capture_context(kwargs["kv_cache"]) timestep = kwargs["timestep"].reshape(-1)[0:1] shared = { "action_labels": _clone_cpu(kwargs["action"].reshape(1, -1).round().long()), "target_viewmats": _clone_cpu(kwargs["viewmats"]), "target_Ks": _clone_cpu(kwargs["Ks"]), "rope_temporal_size": torch.tensor([int(kwargs["rope_temporal_size"])], dtype=torch.int64), "start_rope_start_idx": torch.tensor( [int(kwargs["start_rope_start_idx"])], dtype=torch.int64 ), } self.active = _ActiveStep( chunk_id=chunk_id, step_id=step_id, tensors={ "timestep": _clone_cpu(timestep.float()), "noisy_sample": _clone_cpu(model_input[:, :32]), }, shared=shared, ) def _final_pre(self, module, args, kwargs) -> None: if self.active is None: return hidden, condition = args[0], args[1] batch, tokens, hidden_size = hidden.shape spatial_tokens = LATENT_HEIGHT * LATENT_WIDTH if tokens % spatial_tokens: raise ValueError(f"Final hidden token count {tokens} is not frame-aligned") frames = tokens // spatial_tokens compact = condition.reshape(batch, frames, spatial_tokens, hidden_size)[:, :, 0] expanded = compact[:, :, None].expand(batch, frames, spatial_tokens, hidden_size) if not torch.equal( expanded.reshape(batch, tokens, hidden_size), condition.reshape(batch, tokens, hidden_size), ): raise ValueError("Final-layer condition varies inside a latent frame") self.active.tensors["frame_condition"] = _clone_cpu(compact) self.active.tensors["final_hidden"] = _clone_cpu(hidden) def _transformer_post(self, module, args, kwargs, output) -> None: if self.active is None: return velocity = output[0] if isinstance(output, tuple) else output self.active.tensors["velocity"] = _clone_cpu(velocity) required = {"timestep", "noisy_sample", "frame_condition", "final_hidden", "velocity"} missing = required.difference(self.active.tensors) if missing: raise RuntimeError(f"Incomplete v2 teacher step capture: {sorted(missing)}") self.chunk_steps.append(self.active) completed_step = self.active.step_id self.active = None self.call_index += 1 if completed_step == self.num_steps - 1: self._flush_chunk() def _flush_chunk(self) -> None: if len(self.chunk_steps) != self.num_steps: raise RuntimeError(f"Expected {self.num_steps} steps, got {len(self.chunk_steps)}") if self.chunk_context is None: raise RuntimeError("Chunk context was not captured") chunk_id = self.chunk_steps[0].chunk_id if any(step.chunk_id != chunk_id for step in self.chunk_steps): raise RuntimeError("Captured steps cross a chunk boundary") tensors = dict(self.chunk_steps[0].shared) for step in self.chunk_steps: for name, tensor in step.tensors.items(): tensors[f"step_{step.step_id}_{name}"] = tensor self.on_chunk(chunk_id, tensors, self.chunk_context) self.chunk_steps.clear() self.chunk_context = None def case_tensors(self) -> dict[str, torch.Tensor]: if self.image_condition_latent is None or self.text_context is None: raise RuntimeError("Case condition or text context has not been captured") result = {"image_condition_latent": self.image_condition_latent} for block_id, cache in self.text_context.items(): result[block_key(block_id, "k_txt")] = cache["k_txt"] result[block_key(block_id, "v_txt")] = cache["v_txt"] return result @property def text_token_count(self) -> int: if self.text_context is None: raise RuntimeError("Text context has not been captured") first = self.text_context[self.context_block_ids[0]]["k_txt"] return int(first.shape[2])