"""One-block feature Predictor used by the initialization sweep.""" from __future__ import annotations import copy from typing import Literal import torch from torch import nn from torch.utils.checkpoint import checkpoint from wan.modules.causal_model import CausalWanAttentionBlock from predictor_training.atc_fusion import ATCFusion InitializationMethod = Literal[ "teacher_full", "random_full", "teacher_identity", "random_identity", "full_zero", ] GateMode = Literal["baseline", "learned", "constant"] InputVariant = Literal["self_forcing", "disca", "atc"] class PreviousFeatureGate(nn.Module): """Per-token scalar reliability gate for the normalized previous feature.""" def __init__( self, dim: int, hidden_dim: int = 128, initial_bias: float = 4.6, ) -> None: super().__init__() self.net = nn.Sequential( nn.Linear(2 * dim, hidden_dim), nn.SiLU(), nn.Linear(hidden_dim, 1), ) nn.init.zeros_(self.net[-1].weight) nn.init.constant_(self.net[-1].bias, initial_bias) def forward( self, anchor_normalized: torch.Tensor, previous_normalized: torch.Tensor ) -> torch.Tensor: return torch.sigmoid( self.net(torch.cat([anchor_normalized, previous_normalized], dim=-1)) ) class TripleFeatureFusion(nn.Module): """Fuse target latent tokens, anchor hidden, and previous-chunk hidden.""" def __init__( self, dim: int = 1536, gate_mode: GateMode = "baseline", gate_hidden_dim: int = 128, gate_initial_bias: float = 4.6, gate_floor: float = 0.0, constant_gate: float = 1.0, ) -> None: super().__init__() if gate_mode not in {"baseline", "learned", "constant"}: raise ValueError(f"Unknown gate mode: {gate_mode}") if not 0.0 <= constant_gate <= 1.0: raise ValueError("constant_gate must be in [0, 1]") if not 0.0 <= gate_floor < 1.0: raise ValueError("gate_floor must be in [0, 1)") self.current_norm = nn.LayerNorm(dim, eps=1e-6) self.anchor_norm = nn.LayerNorm(dim, eps=1e-6) self.previous_norm = nn.LayerNorm(dim, eps=1e-6) self.gate_mode = gate_mode self.constant_gate = float(constant_gate) self.gate_floor = float(gate_floor) self.gate_override: float | None = None self.gate = ( PreviousFeatureGate(dim, gate_hidden_dim, gate_initial_bias) if gate_mode == "learned" else None ) self.last_gate: torch.Tensor | None = None self.proj_in = nn.Linear(3 * dim, 2 * dim) self.activation = nn.SiLU() self.proj_out = nn.Linear(2 * dim, dim) def forward( self, current: torch.Tensor, anchor: torch.Tensor, previous: torch.Tensor, ) -> torch.Tensor: if current.shape != anchor.shape or anchor.shape != previous.shape: raise ValueError( "TripleFeatureFusion inputs must have identical shapes: " f"{current.shape}, {anchor.shape}, {previous.shape}" ) anchor_normalized = self.anchor_norm(anchor) previous_normalized = self.previous_norm(previous) if self.gate_override is not None: gate = previous_normalized.new_full( (*previous_normalized.shape[:-1], 1), self.gate_override ) elif self.gate_mode == "learned": assert self.gate is not None raw_gate = self.gate(anchor_normalized, previous_normalized) gate = self.gate_floor + (1.0 - self.gate_floor) * raw_gate elif self.gate_mode == "constant": gate = previous_normalized.new_full( (*previous_normalized.shape[:-1], 1), self.constant_gate ) else: gate = previous_normalized.new_ones( *previous_normalized.shape[:-1], 1 ) # Gate after normalization. A positive per-token scalar applied before # LayerNorm would be normalized away and could not test reliability. self.last_gate = gate.detach() return self.proj_out( self.activation( self.proj_in( torch.cat( [ self.current_norm(current), anchor_normalized, gate * previous_normalized, ], dim=-1, ) ) ) ) class DualFeatureFusion(nn.Module): """DisCa input fusion without a previous-chunk feature channel.""" def __init__(self, dim: int = 1536) -> None: super().__init__() self.current_norm = nn.LayerNorm(dim, eps=1e-6) self.anchor_norm = nn.LayerNorm(dim, eps=1e-6) self.proj_in = nn.Linear(2 * dim, 2 * dim) self.activation = nn.SiLU() self.proj_out = nn.Linear(2 * dim, dim) self.last_gate: torch.Tensor | None = None def forward( self, current: torch.Tensor, anchor: torch.Tensor, ) -> torch.Tensor: if current.shape != anchor.shape: raise ValueError( "DualFeatureFusion inputs must have identical shapes: " f"{current.shape}, {anchor.shape}" ) return self.proj_out( self.activation( self.proj_in( torch.cat( [self.current_norm(current), self.anchor_norm(anchor)], dim=-1, ) ) ) ) def _new_random_block( teacher_block: CausalWanAttentionBlock, ) -> CausalWanAttentionBlock: block = CausalWanAttentionBlock( cross_attn_type="t2v_cross_attn", dim=teacher_block.dim, ffn_dim=teacher_block.ffn_dim, num_heads=teacher_block.num_heads, local_attn_size=teacher_block.local_attn_size, sink_size=teacher_block.self_attn.sink_size, qk_norm=teacher_block.qk_norm, cross_attn_norm=teacher_block.cross_attn_norm, eps=teacher_block.eps, ) # This matches CausalWanModel.init_weights rather than PyTorch Linear's # Kaiming default. for module in block.modules(): if isinstance(module, nn.Linear): nn.init.xavier_uniform_(module.weight) if module.bias is not None: nn.init.zeros_(module.bias) return block def _zero_residual_branch_outputs(block: CausalWanAttentionBlock) -> None: for projection in ( block.self_attn.o, block.cross_attn.o, block.ffn[2], ): nn.init.zeros_(projection.weight) if projection.bias is not None: nn.init.zeros_(projection.bias) def initialize_predictor_block( teacher_block: CausalWanAttentionBlock, method: InitializationMethod, ) -> CausalWanAttentionBlock: """Create a trainable FP32 block under one controlled initialization.""" if method in {"teacher_full", "teacher_identity", "full_zero"}: block = copy.deepcopy(teacher_block).float() elif method in {"random_full", "random_identity"}: block = _new_random_block(teacher_block).float() else: raise ValueError(f"Unknown initialization method: {method}") if method in {"teacher_identity", "random_identity"}: _zero_residual_branch_outputs(block) elif method == "full_zero": for parameter in block.parameters(): nn.init.zeros_(parameter) return block class SingleBlockPredictor(nn.Module): """Predict final hidden as an anchor residual through one causal Wan block.""" def __init__( self, block: CausalWanAttentionBlock, dim: int = 1536, gradient_checkpointing: bool = True, input_variant: InputVariant = "self_forcing", gate_mode: GateMode = "baseline", gate_hidden_dim: int = 128, gate_initial_bias: float = 4.6, gate_floor: float = 0.0, constant_gate: float = 1.0, atc_previous_scope: str = "chunk", atc_freq_dim: int = 256, atc_mlp_hidden_dim: int = 3072, atc_gate_hidden_dim: int = 512, atc_transport_residual_scale: float = 0.1, atc_gate_initial_probability: float = 0.3, atc_collect_diagnostics: bool = True, ) -> None: super().__init__() if input_variant not in {"self_forcing", "disca", "atc"}: raise ValueError(f"Unknown input variant: {input_variant}") if input_variant == "disca" and gate_mode != "baseline": raise ValueError("DisCa dual-input fusion does not use a gate") self.input_variant = input_variant if input_variant == "disca": self.fusion = DualFeatureFusion(dim) elif input_variant == "atc": self.fusion = ATCFusion( dim=dim, num_heads=block.num_heads, freq_dim=atc_freq_dim, mlp_hidden_dim=atc_mlp_hidden_dim, gate_hidden_dim=atc_gate_hidden_dim, previous_scope=atc_previous_scope, transport_residual_scale=atc_transport_residual_scale, gate_initial_probability=atc_gate_initial_probability, gradient_checkpointing=gradient_checkpointing, collect_diagnostics=atc_collect_diagnostics, ) else: self.fusion = TripleFeatureFusion( dim, gate_mode=gate_mode, gate_hidden_dim=gate_hidden_dim, gate_initial_bias=gate_initial_bias, gate_floor=gate_floor, constant_gate=constant_gate, ) self.block = block 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 self.last_atc_diagnostics: dict[str, torch.Tensor] = {} def fusion_parameters(self) -> list[nn.Parameter]: return list(self.fusion.parameters()) + list(self.residual_out.parameters()) def gate_parameters(self) -> list[nn.Parameter]: # This method is only for the legacy optional previous-feature gate. # ATC's TokenGate is part of the stage-1 input module and must train. if not isinstance(self.fusion, TripleFeatureFusion): return [] gate = getattr(self.fusion, "gate", None) return list(gate.parameters()) if gate is not None else [] def fusion_parameters_without_gate(self) -> list[nn.Parameter]: gate_ids = {id(parameter) for parameter in self.gate_parameters()} return [ parameter for parameter in self.fusion_parameters() if id(parameter) not in gate_ids ] def block_parameters(self) -> list[nn.Parameter]: return list(self.block.parameters()) def stage1_input_parameters(self) -> list[nn.Parameter]: """All stage-1 input parameters, including ATC's TokenGate.""" if self.input_variant == "atc": return self.fusion_parameters() return self.fusion_parameters_without_gate() def set_block_trainable(self, enabled: bool) -> None: self.block.requires_grad_(enabled) 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_k: torch.Tensor, history_v: torch.Tensor, cross_k: torch.Tensor, cross_v: torch.Tensor, current_start: int, return_features: bool = False, condition_tokens: torch.Tensor | None = None, anchor_distance: torch.Tensor | None = None, ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: gated_delta_cross: torch.Tensor | None = None if self.input_variant == "disca": fused = self.fusion(current_tokens, anchor_hidden) elif self.input_variant == "atc": if condition_tokens is None or anchor_distance is None: raise ValueError( "ATC requires condition_tokens and anchor_distance" ) fused, gated_delta_cross = self.fusion( current_tokens, anchor_hidden, previous_hidden, condition_tokens, anchor_distance, grid_sizes, freqs, current_start, ) else: fused = self.fusion(current_tokens, anchor_hidden, previous_hidden) batch, sequence, dim = fused.shape history_length = history_k.shape[1] if history_k.shape[0] != batch: raise ValueError( f"history_k {history_k.shape} incompatible with " f"batch={batch}" ) if history_v.shape != history_k.shape: raise ValueError("History K/V shapes differ") seq_lens = torch.full( (batch,), sequence, dtype=torch.long, device="cpu", ) context = fused.new_zeros(batch, 1, dim) def run_block(value: torch.Tensor) -> torch.Tensor: # A fresh cache is required for checkpoint recomputation because # CausalWanSelfAttention updates cache indices and current K/V. 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( [history_length], device=value.device, dtype=torch.long ), } crossattn_cache = { "k": cross_k, "v": cross_v, "is_init": True, } return self.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, ) if self.training and self.gradient_checkpointing: transformed = checkpoint( run_block, fused, use_reentrant=False, ) else: transformed = run_block(fused) delta_denoise = self.residual_out(transformed) pred_hidden = anchor_hidden + delta_denoise if gated_delta_cross is not None: pred_hidden = pred_hidden + gated_delta_cross self.last_atc_diagnostics = {} if self.fusion.collect_diagnostics: self.last_atc_diagnostics = { **self.fusion.last_diagnostics, "delta_h_d_norm": ( delta_denoise.detach().float().square().mean().sqrt() ), } else: self.last_atc_diagnostics = {} if return_features: return pred_hidden, transformed return pred_hidden