Self-Forcing / predictor_training /single_block.py
Cccccz's picture
Upload predictor training source
20962c9 verified
Raw History Blame Contribute Delete
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