File size: 4,772 Bytes
3e936b2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import torch
import torch.distributed as dist
from wan.modules.attention import attention
from wan.modules.causal_model import causal_rope_apply

class SinkMemory:

    def __init__(self, num_blocks, num_heads, head_dim, tokens_per_frame, sink_size, hidden_dim):
        self.num_blocks = num_blocks
        self.num_heads = num_heads
        self.head_dim = head_dim
        self.tokens_per_frame = tokens_per_frame
        self.sink_size = sink_size
        self.M = sink_size * tokens_per_frame
        self.hidden_dim = hidden_dim
        self.memory_q = [None] * num_blocks
        self.memory_k = [None] * num_blocks
        self.memory_v = [None] * num_blocks
        self.sink_hidden = None
        self.initialized = False
        self.has_history = False
        self._update_count = 0
        self._pending_sink_captures = {}
        self._pending_sink_hidden = None

    def reset(self):
        for i in range(self.num_blocks):
            self.memory_q[i] = None
            self.memory_k[i] = None
            self.memory_v[i] = None
        self.sink_hidden = None
        self.initialized = False
        self.has_history = False
        self._update_count = 0
        self._pending_sink_captures = {}
        self._pending_sink_hidden = None

    def initialize_block(self, block_idx, q_roped, k_pre_rope, v):
        self.memory_q[block_idx] = q_roped.clone().detach()
        self.memory_k[block_idx] = k_pre_rope.clone().detach()
        self.memory_v[block_idx] = v.clone().detach()
        if block_idx == self.num_blocks - 1:
            self.initialized = True

    def initialize_hidden(self, hidden):
        self.sink_hidden = hidden.clone().detach()

    def get_kv(self, block_idx):
        if not self.has_history:
            return None
        return (self.memory_k[block_idx], self.memory_v[block_idx])

    @torch.no_grad()
    def update(self, blocks, evicted_kv_all, e0, context, context_lens, freqs, grid_sizes, crossattn_cache_template=None, evicted_k_is_pre_rope=False):
        if not self.initialized or self.sink_hidden is None:
            return
        x_mem = self.sink_hidden
        B = x_mem.shape[0]
        Q_frames = self.sink_size
        tpf = self.tokens_per_frame
        n, d = (self.num_heads, self.head_dim)
        mem_grid = grid_sizes.clone()
        mem_grid[:, 0] = Q_frames
        for i in range(self.num_blocks):
            block = blocks[i]
            evicted_k_i, evicted_v_i = evicted_kv_all[i]
            e_block = (block.modulation.unsqueeze(1) + e0).chunk(6, dim=2)
            h = block.norm1(x_mem)
            h = (h.unflatten(1, (Q_frames, tpf)) * (1 + e_block[1]) + e_block[0]).flatten(1, 2)
            mem_k_roped = causal_rope_apply(self.memory_k[i], mem_grid, freqs, start_frame=0)
            if evicted_k_is_pre_rope and evicted_k_i is not None and (evicted_k_i.shape[1] > 0):
                ev_frames = evicted_k_i.shape[1] // self.tokens_per_frame
                if ev_frames > 0:
                    ev_grid = grid_sizes.clone()
                    ev_grid[:, 0] = ev_frames
                    evicted_k_i = causal_rope_apply(evicted_k_i, ev_grid, freqs, start_frame=self.sink_size)
            ctx_k = torch.cat([mem_k_roped, evicted_k_i], dim=1)
            ctx_v = torch.cat([self.memory_v[i], evicted_v_i], dim=1)
            attn_out = attention(self.memory_q[i], ctx_k, ctx_v)
            y = block.self_attn.o(attn_out.flatten(2))
            x_mem = x_mem + (y.unflatten(1, (Q_frames, tpf)) * e_block[2]).flatten(1, 2)
            x_mem = x_mem + block.cross_attn(block.norm3(x_mem), context, context_lens)
            y_ffn = block.ffn((block.norm2(x_mem).unflatten(1, (Q_frames, tpf)) * (1 + e_block[4]) + e_block[3]).flatten(1, 2))
            x_mem = x_mem + (y_ffn.unflatten(1, (Q_frames, tpf)) * e_block[5]).flatten(1, 2)
            if i < self.num_blocks - 1:
                next_block = blocks[i + 1]
                next_e = (next_block.modulation.unsqueeze(1) + e0).chunk(6, dim=2)
                h_next = next_block.norm1(x_mem)
                h_next_mod = (h_next.unflatten(1, (Q_frames, tpf)) * (1 + next_e[1]) + next_e[0]).flatten(1, 2)
                self.memory_k[i + 1] = next_block.self_attn.norm_k(next_block.self_attn.k(h_next_mod)).view(B, self.M, n, d)
                self.memory_v[i + 1] = next_block.self_attn.v(h_next_mod).view(B, self.M, n, d)
        self.has_history = True
        self._update_count += 1
        if self._update_count <= 100 or self._update_count % 10 == 0:
            sample_layers = [0, self.num_blocks // 2, self.num_blocks - 1]
            norms_k = [f'L{i}={self.memory_k[i].norm().item():.1f}' for i in sample_layers]
            norms_v = [f'L{i}={self.memory_v[i].norm().item():.1f}' for i in sample_layers]