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)