File size: 5,571 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 | """Two-block feature Predictor for Teacher-layer pair 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 TwoBlockPredictor(nn.Module):
"""Fuse three inputs, run two causal Wan blocks, and predict an anchor residual."""
def __init__(
self,
teacher_blocks: list[CausalWanAttentionBlock],
dim: int = 1536,
gradient_checkpointing: bool = True,
) -> None:
super().__init__()
if len(teacher_blocks) != 2:
raise ValueError(f"Expected exactly two 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:
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) != 2 for values in cache_lists):
raise ValueError("Two 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)
|