File size: 9,726 Bytes
09ccad2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
"""H3 文本编码器替换适配器(接口蒸馏):source_projection + CrossAttention Resampler + TokenRefiner。

前向:
    h3_ids [B, S_T] --QueryEmbedding--> Q [B, S_T, 5376]
    student_hidden [B, S_S, 2560] --source_projection--> KV [B, S_S, 5376]
    CrossAttentionBlock(Q, KV) -> [B, S_T, 5376]
    TokenRefiner(2 层, 原权重初始化) -> [B, S_T, 5376]   # 与教师 target 同坐标系

TokenRefiner 结构严格复刻原始 checkpoint(已源码核实):
    - fused qkv (chunk(3)) + per-head qk_norm + 双向注意力 + out_proj,全部 bias=False
    - SwiGLU MLP: fc1 为 fused [gate; value](gate 在前),fc2(silu(gate)*value)
    - 2 个 pre-norm block + final RMSNorm,eps 全部 1e-5
"""
from __future__ import annotations

import math

import torch
import torch.nn as nn
import torch.nn.functional as F

H3_VOCAB = 151936
H3_HIDDEN = 5376
REFINER_HEADS = 56
REFINER_HEAD_DIM = 128
REFINER_FFN = 14336
STUDENT_HIDDEN = 2560


class RMSNorm(nn.Module):
    def __init__(self, dim: int, eps: float = 1e-5):
        super().__init__()
        self.weight = nn.Parameter(torch.ones(dim))
        self.eps = eps

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) * self.weight


class SwiGLUFFN(nn.Module):
    def __init__(self, hidden: int, ffn: int):
        super().__init__()
        self.fc1 = nn.Linear(hidden, 2 * ffn, bias=False)
        self.fc2 = nn.Linear(ffn, hidden, bias=False)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        gate, value = self.fc1(x).chunk(2, dim=-1)
        return self.fc2(F.silu(gate) * value)


class TokenRefinerAttention(nn.Module):
    def __init__(self, hidden: int, heads: int, dim_head: int):
        super().__init__()
        self.heads = heads
        self.head_dim = dim_head
        self.inner_dim = heads * dim_head
        self.qkv_proj = nn.Linear(hidden, 3 * self.inner_dim, bias=False)
        self.q_norm = RMSNorm(dim_head)
        self.k_norm = RMSNorm(dim_head)
        self.out_proj = nn.Linear(self.inner_dim, hidden, bias=False)
        self.use_sdpa = True  # SDPA flash: 注意力内存 O(S²)->O(S),数值与手写注意力差 ~1e-3(蒸馏噪声级)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        q, k, v = self.qkv_proj(x).chunk(3, dim=-1)
        q = self.q_norm(q.unflatten(-1, (self.heads, self.head_dim)).transpose(1, 2))
        k = self.k_norm(k.unflatten(-1, (self.heads, self.head_dim)).transpose(1, 2))
        v = v.unflatten(-1, (self.heads, self.head_dim)).transpose(1, 2)
        if self.use_sdpa:
            out = F.scaled_dot_product_attention(q, k, v, scale=self.head_dim ** -0.5)
        else:
            attn = torch.softmax((q @ k.transpose(-2, -1)) * (self.head_dim ** -0.5), dim=-1)
            out = attn @ v
        return self.out_proj(out.transpose(1, 2).flatten(2))


class TokenRefinerBlock(nn.Module):
    def __init__(self, hidden: int, heads: int, dim_head: int, ffn: int):
        super().__init__()
        self.norm1 = RMSNorm(hidden)
        self.attn = TokenRefinerAttention(hidden, heads, dim_head)
        self.norm2 = RMSNorm(hidden)
        self.mlp = SwiGLUFFN(hidden, ffn)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        x = x + self.attn(self.norm1(x))
        x = x + self.mlp(self.norm2(x))
        return x


class TokenRefiner(nn.Module):
    def __init__(self, num_layers: int = 2, hidden: int = H3_HIDDEN,
                 heads: int = REFINER_HEADS, dim_head: int = REFINER_HEAD_DIM, ffn: int = REFINER_FFN):
        super().__init__()
        self.blocks = nn.ModuleList([TokenRefinerBlock(hidden, heads, dim_head, ffn) for _ in range(num_layers)])
        self.final_norm = RMSNorm(hidden)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        for block in self.blocks:
            x = block(x)
        return self.final_norm(x)


class CrossAttentionBlock(nn.Module):
    """跨空间交叉注意(学生 2560 投影空间 <-> 教师 5376 表示空间)。

    训练稳定性升级(兼容旧权重加载):
      - QK-norm: to_q/to_k 投影后对每头做 RMSNorm(Llama 3.2 vision 交叉注意惯例,
        稳定跨空间注意力的 q/k 尺度失配)
      - gated tanh: out_proj 输出经 tanh(gate) 门控(Flamingo 惯例),gate 初始 0 => 恒等
      - attn dropout: 交叉注意 dropout 0.1(Emu3 防后期 collapse)
    """

    def __init__(self, hidden: int = H3_HIDDEN, heads: int = 32, dim_head: int = 128, ffn: int = REFINER_FFN,
                 attn_dropout: float = 0.1, use_qk_norm: bool = True, use_gate: bool = True):
        super().__init__()
        self.heads = heads
        self.head_dim = dim_head
        self.inner_dim = heads * dim_head
        self.norm_q = RMSNorm(hidden)
        self.norm_kv = RMSNorm(hidden)
        self.to_q = nn.Linear(hidden, self.inner_dim, bias=False)
        self.to_k = nn.Linear(hidden, self.inner_dim, bias=False)
        self.to_v = nn.Linear(hidden, self.inner_dim, bias=False)
        self.out_proj = nn.Linear(self.inner_dim, hidden, bias=False)
        self.norm2 = RMSNorm(hidden)
        self.mlp = SwiGLUFFN(hidden, ffn)
        self.use_sdpa = True
        self.attn_dropout = attn_dropout
        self.use_qk_norm = use_qk_norm
        if use_qk_norm:
            self.qk_norm = RMSNorm(dim_head)
        self.use_gate = use_gate
        if use_gate:
            self.gate = nn.Parameter(torch.zeros(1))  # tanh(0)=0 -> 恒等,兼容旧权重

    def forward(self, q: torch.Tensor, kv: torch.Tensor) -> torch.Tensor:
        qn = self.norm_q(q)
        kvn = self.norm_kv(kv)
        qh = self.to_q(qn).unflatten(-1, (self.heads, self.head_dim)).transpose(1, 2)
        kh = self.to_k(kvn).unflatten(-1, (self.heads, self.head_dim)).transpose(1, 2)
        vh = self.to_v(kvn).unflatten(-1, (self.heads, self.head_dim)).transpose(1, 2)
        if self.use_qk_norm:
            qh = self.qk_norm(qh)
            kh = self.qk_norm(kh)
        if self.use_sdpa:
            out = F.scaled_dot_product_attention(qh, kh, vh, scale=self.head_dim ** -0.5,
                                                 dropout_p=self.attn_dropout if self.training else 0.0)
        else:
            attn = torch.softmax((qh @ kh.transpose(-2, -1)) * (self.head_dim ** -0.5), dim=-1)
            out = attn @ vh
        proj = self.out_proj(out.transpose(1, 2).flatten(2))
        x = q + (torch.tanh(self.gate) * proj if self.use_gate else proj)
        x = x + self.mlp(self.norm2(x))
        return x


class QueryEmbedding(nn.Module):
    def __init__(self, vocab: int = H3_VOCAB, dim: int = 256, out: int = H3_HIDDEN):
        super().__init__()
        self.embed = nn.Embedding(vocab, dim)
        self.proj = nn.Linear(dim, out, bias=True)

    def forward(self, ids: torch.Tensor) -> torch.Tensor:
        return self.proj(self.embed(ids))


class H3Adapter(nn.Module):
    """完整适配器。param 约 1.14B(source_proj 13.8M + query_embed 40.3M + crossattn 319M + refiner 751M)。"""

    def __init__(self):
        super().__init__()
        self.source_projection = nn.Linear(STUDENT_HIDDEN, H3_HIDDEN, bias=True)
        self.query_embedding = QueryEmbedding()
        self.cross_attention = CrossAttentionBlock()
        self.token_refiner = TokenRefiner()

    def forward(self, h3_ids: torch.Tensor, student_hidden: torch.Tensor) -> torch.Tensor:
        kv = self.source_projection(student_hidden)
        q = self.query_embedding(h3_ids)
        x = self.cross_attention(q, kv)
        x = self.token_refiner(x)
        return x

    def load_token_refiner(self, state_dict: dict[str, torch.Tensor], strict: bool = True) -> None:
        refiner_state = {
            k[len("token_refiner."):]: v
            for k, v in state_dict.items()
            if k.startswith("token_refiner.")
        }
        missing, unexpected = self.token_refiner.load_state_dict(refiner_state, strict=strict)
        assert not missing and not unexpected, f"refiner load: missing={missing} unexpected={unexpected}"

    def trainable_modules(self, stage: int) -> list[nn.Parameter]:
        if stage == 1:
            return (
                list(self.source_projection.parameters())
                + list(self.query_embedding.parameters())
                + list(self.cross_attention.parameters())
            )
        return list(self.parameters())


class TeacherHead(nn.Module):
    """教师 target 计算: h50 -> condition_proj(5120->5376) -> token_refiner -> [B, S_T, 5376]。"""

    def __init__(self):
        super().__init__()
        self.condition_proj = nn.Linear(5120, H3_HIDDEN, bias=True)
        self.token_refiner = TokenRefiner()

    def forward(self, h50: torch.Tensor) -> torch.Tensor:
        return self.token_refiner(self.condition_proj(h50))


def compute_query_embed_init(embed_tokens: torch.Tensor, dim: int = 256, seed: int = 0) -> torch.Tensor:
    """query embedding 初始化: embed.weight = E @ P,P 为 5120->dim 的随机投影。

    仅用于一次性初始化,后续可训练。embed_tokens 为 [151936, 5120](BF16 或 FP32)。
    """
    assert embed_tokens.dim() == 2 and embed_tokens.shape[1] == 5120, embed_tokens.shape
    rng = torch.Generator().manual_seed(seed)
    p = torch.randn(5120, dim, generator=rng) * (1.0 / math.sqrt(5120))
    embed_tokens = embed_tokens.to(torch.float32)
    chunks = 8
    out = torch.empty(embed_tokens.shape[0], dim, dtype=torch.float32)
    for i in range(chunks):
        lo = i * embed_tokens.shape[0] // chunks
        hi = (i + 1) * embed_tokens.shape[0] // chunks
        out[lo:hi] = embed_tokens[lo:hi] @ p
    return out.to(torch.bfloat16)