"""Two-block feature Predictor for Teacher-layer pair experiments.""" from __future__ import annotations import copy import torch from torch import nn from torch.utils.checkpoint import checkpoint from predictor_training.single_block import TripleFeatureFusion from wan.modules.causal_model import CausalWanAttentionBlock class TwoBlockPredictor(nn.Module): """Fuse three inputs, run two causal Wan blocks, and predict an anchor residual.""" def __init__( self, teacher_blocks: list[CausalWanAttentionBlock], dim: int = 1536, gradient_checkpointing: bool = True, ) -> None: super().__init__() if len(teacher_blocks) != 2: raise ValueError(f"Expected exactly two blocks, got {len(teacher_blocks)}") self.fusion = TripleFeatureFusion(dim) self.blocks = nn.ModuleList( [copy.deepcopy(block).float() for block in teacher_blocks] ) self.residual_out = nn.Linear(dim, dim) nn.init.zeros_(self.residual_out.weight) nn.init.zeros_(self.residual_out.bias) self.gradient_checkpointing = gradient_checkpointing def fusion_parameters(self) -> list[nn.Parameter]: return list(self.fusion.parameters()) + list(self.residual_out.parameters()) def block_parameters(self) -> list[nn.Parameter]: return list(self.blocks.parameters()) def set_blocks_trainable(self, enabled: bool) -> None: self.blocks.requires_grad_(enabled) @staticmethod def _block_closure( *, block: CausalWanAttentionBlock, sequence: int, timestep_modulation: torch.Tensor, grid_sizes: torch.Tensor, freqs: torch.Tensor, history_k: torch.Tensor, history_v: torch.Tensor, cross_k: torch.Tensor, cross_v: torch.Tensor, current_start: int, ): batch = history_k.shape[0] dim = block.dim seq_lens = torch.full( (batch,), sequence, dtype=torch.long, device="cpu" ) def run(value: torch.Tensor) -> torch.Tensor: heads = history_k.shape[2] head_dim = history_k.shape[3] current_k = history_k.new_empty(batch, sequence, heads, head_dim) current_v = history_v.new_empty(batch, sequence, heads, head_dim) kv_cache = { "k": torch.cat([history_k, current_k], dim=1), "v": torch.cat([history_v, current_v], dim=1), "global_end_index": torch.tensor( [current_start], device=value.device, dtype=torch.long ), "local_end_index": torch.tensor( [current_start], device=value.device, dtype=torch.long ), } crossattn_cache = { "k": cross_k, "v": cross_v, "is_init": True, } context = value.new_zeros(batch, 1, dim) return block( value, e=timestep_modulation, seq_lens=seq_lens, grid_sizes=grid_sizes, freqs=freqs, context=context, context_lens=None, block_mask=None, kv_cache=kv_cache, crossattn_cache=crossattn_cache, current_start=current_start, cache_start=None, ) return run def forward( self, *, current_tokens: torch.Tensor, anchor_hidden: torch.Tensor, previous_hidden: torch.Tensor, timestep_modulation: torch.Tensor, grid_sizes: torch.Tensor, freqs: torch.Tensor, history_ks: list[torch.Tensor], history_vs: list[torch.Tensor], cross_ks: list[torch.Tensor], cross_vs: list[torch.Tensor], current_start: int, ) -> torch.Tensor: cache_lists = (history_ks, history_vs, cross_ks, cross_vs) if any(len(values) != 2 for values in cache_lists): raise ValueError("Two cache sets are required") transformed = self.fusion(current_tokens, anchor_hidden, previous_hidden) batch, sequence, _ = transformed.shape for position, block in enumerate(self.blocks): history_k = history_ks[position] history_v = history_vs[position] if history_k.shape[:2] != (batch, current_start): raise ValueError( f"history_k[{position}] {history_k.shape} incompatible with " f"batch={batch}, current_start={current_start}" ) if history_v.shape != history_k.shape: raise ValueError(f"History K/V shapes differ at block {position}") run_block = self._block_closure( block=block, sequence=sequence, timestep_modulation=timestep_modulation, grid_sizes=grid_sizes, freqs=freqs, history_k=history_k, history_v=history_v, cross_k=cross_ks[position], cross_v=cross_vs[position], current_start=current_start, ) if self.training and self.gradient_checkpointing: transformed = checkpoint( run_block, transformed, use_reentrant=False ) else: transformed = run_block(transformed) return anchor_hidden + self.residual_out(transformed)