File size: 5,733 Bytes
20962c9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 | """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)
|