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