Download predictor_training/single_block.py from Cccccz/Self-Forcing: direct link, hf CLI and curl.
- Browser
- Download file 16 kB
-
https://huggingface.co/Cccccz/Self-Forcing/resolve/main/predictor_training/single_block.py
- Command line
-
hf download hf://Cccccz/Self-Forcing/predictor_training/single_block.py
-
curl -L -o single_block.py https://huggingface.co/Cccccz/Self-Forcing/resolve/main/predictor_training/single_block.py
16 kB
| """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 | |