"""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