| """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 |
|
|
| 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)) |
|
|
| 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) |
|
|