File size: 5,200 Bytes
b66f552
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang

import torch
from einops import rearrange


def naive_mesa_net_decoding_one_step(q, k, v, g, lamb, beta, prev_h_kk, prev_h_kv, max_CG_iteration=30):
    q = q.float().clone()
    k = k.float().clone()
    v = v.float().clone()
    g = g.float().clone()
    lamb = lamb.float().clone()
    beta = beta.float().clone()
    B, h, d = q.shape
    k_beta = k * beta.unsqueeze(-1)

    h_kk = prev_h_kk * g.exp()[..., None, None] + k_beta.unsqueeze(-1) * k.unsqueeze(-2)
    h_kv = prev_h_kv * g.exp()[..., None, None] + k_beta.unsqueeze(-1) * v.unsqueeze(-2)
    diag_H = torch.diagonal(h_kk, dim1=-2, dim2=-1)
    lamb = lamb.unsqueeze(0)
    x = q / (diag_H + lamb)
    r = q - (x.unsqueeze(-1) * h_kk).sum(-2) - (lamb * x)
    p = r.clone()
    delta_old = (r * r).sum(-1)
    # CG iteration
    for i in range(max_CG_iteration):
        q = (p.unsqueeze(-1) * h_kk).sum(-2) + (lamb * p)
        alpha = (delta_old / ((p * q).sum(-1) + 1e-5))
        x = x + (alpha[..., None] * p)
        r = r - (alpha[..., None] * q)
        delta_new = (r * r).sum(-1)
        beta = delta_new / (delta_old + 1e-5)
        p = r + (beta[..., None] * p)
        delta_old = delta_new
    o = (x.unsqueeze(-1) * h_kv).sum(-2)
    return o, h_kk, h_kv


def naive_mesa_net_exact(q, k, v, g, lamb, beta, h_kk_init=None, h_kv_init=None):
    B, L, h, d = q.shape
    q = q.float()
    k = k.float()
    v = v.float()
    g = g.float()
    lamb = lamb.float()
    beta = beta.float()

    h_kk = h_kk_init.clone() if h_kk_init is not None else torch.zeros(B, h, d, d, device=q.device)
    h_kv = h_kv_init.clone() if h_kv_init is not None else torch.zeros(B, h, d, d, device=q.device)

    h_kk_all = torch.zeros(B, L, h, d, d, device=q.device)
    h_kv_all = torch.zeros(B, L, h, d, d, device=q.device)
    for i in range(L):
        h_kk = h_kk * g[:, i, :, None, None].exp() + (k[:, i, :, :] * beta[:, i, :, None]
                                                      )[..., None] * k[:, i, :, None, :]
        h_kv = h_kv * g[:, i, :, None, None].exp() + (k[:, i, :, :] * beta[:, i, :, None]
                                                      )[..., None] * v[:, i, :, None, :]
        h_kk_all[:, i] = h_kk
        h_kv_all[:, i] = h_kv

    q_star_gold = torch.linalg.solve(h_kk_all + torch.diag_embed(lamb)[None, None, ...], q)
    o_gold = (q_star_gold[..., :, None] * h_kv_all).sum(-2)
    return o_gold, h_kk, h_kv


def naive_mesa_net_CG(q, k, v, g, lamb, beta, chunk_size, max_CG_iteration=30, h_kk_init=None, h_kv_init=None):
    B, L, h, d = q.shape
    C = chunk_size

    def chunk_fn(x): return rearrange(x, 'b (n c) h ... -> b h n c ...', c=C).float()

    q_chunk, k_chunk, v_chunk, g_chunk, beta_chunk = map(chunk_fn, [q, k, v, g, beta])

    g_chunk = g_chunk.cumsum(dim=-1)

    pairwise_decay = (g_chunk[..., None] - g_chunk[..., None, :]).exp().tril() * beta_chunk[..., None, :]

    num_chunks = q_chunk.shape[2]

    h_kv_all = torch.zeros(B, h, num_chunks, d, d, device=q.device)
    h_kk_all = torch.zeros(B, h, num_chunks, d, d, device=q.device)

    h_kv = torch.zeros(B, h, d, d, device=q.device)
    h_kk = torch.zeros(B, h, d, d, device=q.device)

    if h_kk_init is not None:
        h_kk += h_kk_init
    if h_kv_init is not None:
        h_kv += h_kv_init

    chunk_decay_k = (g_chunk[..., -1, None] - g_chunk).exp()
    chunk_decay_q = g_chunk.exp()

    k_chunk_processed = k_chunk * chunk_decay_k[..., None] * beta_chunk[..., None]

    for i in range(num_chunks):
        h_kv_all[:, :, i, :, :] = h_kv
        h_kk_all[:, :, i, :, :] = h_kk

        k_chunk_i = k_chunk[:, :, i, :, :]
        v_chunk_i = v_chunk[:, :, i, :, :]
        k_chunk_i_processed = k_chunk_processed[:, :, i, :, :]

        h_kk = h_kk * g_chunk[:, :, i, -1, None, None].exp() + (k_chunk_i_processed).transpose(-2, -1) @ k_chunk_i
        h_kv = h_kv * g_chunk[:, :, i, -1, None, None].exp() + (k_chunk_i_processed).transpose(-2, -1) @ v_chunk_i

    # CG solver to approximate the matrix inverse solution.
    # diag_H = torch.diagonal(h_kk_all, dim1=-2, dim2=-1)
    lamb = lamb[None, :, None, None, :]
    x = torch.zeros_like(q_chunk)
    r = q_chunk - (x * chunk_decay_q[..., None]) @ h_kk_all - ((x @ k_chunk.transpose(-2, -1))
                                                               * pairwise_decay) @ k_chunk - (lamb * x)
    p = r.clone()
    delta_old = (r * r).sum(-1)

    # CG iteration
    for i in range(max_CG_iteration):
        q = (p * chunk_decay_q[..., None]) @ h_kk_all + ((p @ k_chunk.transpose(-1, -2))
                                                         * pairwise_decay) @ k_chunk + (lamb * p)
        alpha = (delta_old / ((p * q).sum(-1) + 1e-5))
        x = x + (alpha[..., None] * p)
        r = r - (alpha[..., None] * q)
        delta_new = (r * r).sum(-1)
        beta = delta_new / (delta_old + 1e-5)
        p = r + (beta[..., None] * p)
        delta_old = delta_new

    o = (x * chunk_decay_q[..., None]) @ h_kv_all + ((x @ k_chunk.transpose(-1, -2))
                                                     * pairwise_decay) @ v_chunk
    return rearrange(o, 'b h n c d -> b (n c) h d'), h_kk, h_kv