Diffusers
Safetensors
HY / predictor_data /v2_hooks.py
Cccccz's picture
Upload batch 64: 500 files (0.40 GiB)
5f0e4a2 verified
Raw History Blame Contribute Delete
9.31 kB
"""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])