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)