Cccccz's picture
Upload code and configuration only
ae8ade0 verified
Raw History Blame Contribute Delete
11.2 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
InitializationMethod = Literal[
"teacher_full",
"random_full",
"teacher_identity",
"random_identity",
"full_zero",
]
GateMode = Literal["baseline", "learned", "constant"]
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,
)
)
)
)
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,
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__()
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
def fusion_parameters(self) -> list[nn.Parameter]:
return list(self.fusion.parameters()) + list(self.residual_out.parameters())
def gate_parameters(self) -> list[nn.Parameter]:
return list(self.fusion.gate.parameters()) if self.fusion.gate 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 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,
) -> torch.Tensor:
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)
return anchor_hidden + self.residual_out(transformed)