Self-Forcing / predictor_training /three_block.py
Cccccz's picture
Upload predictor training source
20962c9 verified
Raw History Blame Contribute Delete
5.73 kB
"""Three-block feature Predictor for consecutive Teacher-layer 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 ThreeBlockPredictor(nn.Module):
"""Fuse three inputs, run three causal Wan blocks, and predict a residual."""
def __init__(
self,
teacher_blocks: list[CausalWanAttentionBlock],
dim: int = 1536,
gradient_checkpointing: bool = True,
) -> None:
super().__init__()
if len(teacher_blocks) != 3:
raise ValueError(
f"Expected exactly three 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:
# Each checkpoint recomputation needs fresh mutable cache indices
# and current-token K/V storage.
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) != 3 for values in cache_lists):
raise ValueError("Three 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)