Diffusers
Safetensors
HY / predictor_data /prefeature_hooks.py
Cccccz's picture
Upload batch 64: 500 files (0.40 GiB)
5f0e4a2 verified
Raw History Blame Contribute Delete
7.61 kB
"""Read-only Full-DiT capture of targets and joint-window Context prefeatures."""
from __future__ import annotations
from typing import Any, Callable
import torch
from .prefeature_schema import CONTEXT_BLOCK_IDS
from .v2_hooks import PredictorV2TeacherCapture, _clone_cpu
class PredictorPrefeatureTeacherCapture(PredictorV2TeacherCapture):
"""Capture ``img_modulated`` before selected Teacher K/V projections."""
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 = 4,
context_block_ids: tuple[int, ...] = CONTEXT_BLOCK_IDS,
capture_direct_kv: bool = False,
) -> None:
super().__init__(
transformer,
on_chunk=on_chunk,
num_steps=num_steps,
context_block_ids=context_block_ids,
)
self.capture_direct_kv = bool(capture_direct_kv)
self.pending_context: dict[int, dict[str, torch.Tensor]] | None = None
self.pending_direct_kv: dict[int, dict[str, torch.Tensor]] | None = None
self._prefill_active = False
self._prefill_features: dict[int, torch.Tensor] = {}
self._prefill_metadata: dict[str, torch.Tensor] = {}
def __enter__(self) -> "PredictorPrefeatureTeacherCapture":
super().__enter__()
for block_id in self.context_block_ids:
block = self.transformer.double_blocks[block_id]
self._handles.append(
block.img_attn_k.register_forward_pre_hook(
self._make_k_pre_hook(block_id), with_kwargs=True
)
)
return self
def _make_k_pre_hook(self, block_id: int):
def hook(module, args, kwargs) -> None:
if not self._prefill_active:
return
if block_id in self._prefill_features:
raise RuntimeError(f"Context block {block_id} was captured twice")
if not args or not torch.is_tensor(args[0]):
raise RuntimeError(f"Missing img_modulated input for block {block_id}")
self._prefill_features[block_id] = _clone_cpu(args[0])
return hook
def _transformer_pre(self, module, args, kwargs) -> None:
if kwargs.get("ar_vision_inference", False) and kwargs.get("cache_vision", False):
if self.pending_context is not None or self._prefill_active:
raise RuntimeError("A Context prefill was not consumed before the next prefill")
frame_indices = kwargs.get("context_frame_indices")
if frame_indices is None:
raise RuntimeError("Context prefill lacks context_frame_indices metadata")
if not torch.is_tensor(frame_indices):
frame_indices = torch.tensor(frame_indices, dtype=torch.int64)
frame_indices = frame_indices.detach().to(device="cpu", dtype=torch.int64).reshape(-1)
frames = int(kwargs["hidden_states"].shape[2])
if frame_indices.numel() != frames:
raise ValueError("Context frame metadata does not match prefill tensor")
self._prefill_metadata = {
"selected_frame_indices": frame_indices.contiguous(),
"context_viewmats": _clone_cpu(kwargs["viewmats"]),
"context_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._prefill_features = {}
self._prefill_active = True
return
super()._transformer_pre(module, args, kwargs)
if self.active is not None and self.active.step_id == 0:
self.chunk_context = self._consume_context()
def _transformer_post(self, module, args, kwargs, output) -> None:
if self._prefill_active:
self._prefill_active = False
missing = set(self.context_block_ids).difference(self._prefill_features)
if missing:
raise RuntimeError(f"Missing prefeatures for blocks {sorted(missing)}")
token_counts = {int(value.shape[1]) for value in self._prefill_features.values()}
if len(token_counts) != 1:
raise RuntimeError("Selected Context blocks have different window lengths")
tokens = next(iter(token_counts))
mask = torch.ones((1, tokens), dtype=torch.bool)
self.pending_context = {
block_id: {
"img_modulated": feature,
"context_valid_mask": mask.clone(),
**{name: value.clone() for name, value in self._prefill_metadata.items()},
}
for block_id, feature in self._prefill_features.items()
}
if self.capture_direct_kv:
self.pending_direct_kv = {
block_id: {
"k_vision": _clone_cpu(output[block_id]["k_vision"]),
"v_vision": _clone_cpu(output[block_id]["v_vision"]),
}
for block_id in self.context_block_ids
}
self._prefill_features = {}
self._prefill_metadata = {}
return
super()._transformer_post(module, args, kwargs, output)
def _consume_context(self) -> dict[int, dict[str, torch.Tensor]]:
if self.pending_context is not None:
context = self.pending_context
self.pending_context = None
return context
# The first chunk has no history prefill.
empty_metadata = {
"context_valid_mask": torch.ones((1, 0), dtype=torch.bool),
"selected_frame_indices": torch.empty((0,), dtype=torch.int64),
"context_viewmats": torch.empty((1, 0, 4, 4), dtype=torch.bfloat16),
"context_Ks": torch.empty((1, 0, 3, 3), dtype=torch.bfloat16),
"rope_temporal_size": torch.tensor([0], dtype=torch.int64),
"start_rope_start_idx": torch.tensor([0], dtype=torch.int64),
}
return {
block_id: {
"img_modulated": torch.empty((1, 0, 2048), dtype=torch.bfloat16),
**{name: value.clone() for name, value in empty_metadata.items()},
}
for block_id in self.context_block_ids
}
def _capture_context(self, kv_cache):
"""Capture only case-level text KV; Context features come from block hooks."""
if self.text_context is None:
text_context: dict[int, dict[str, torch.Tensor]] = {}
for block_id in self.context_block_ids:
cache = kv_cache[block_id]
k_txt, v_txt = cache.get("k_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")
text_context[block_id] = {
"k_txt": _clone_cpu(k_txt),
"v_txt": _clone_cpu(v_txt),
}
self.text_context = text_context
return {block_id: {} for block_id in self.context_block_ids}
def take_direct_kv(self) -> dict[int, dict[str, torch.Tensor]] | None:
value = self.pending_direct_kv
self.pending_direct_kv = None
return value