Transformers 中的 AXK2 架构解析:SK Telecom A.X-K2 稀疏注意力 MoE 大模型集成指南
本文以 axk2 模型文档 为核心骨架,结合本仓库
src/transformers/models/axk2/下的真实源码实现,系统讲解 SK Telecom A.X-K2 旗舰大语言模型在 Hugging Face Transformers 中的集成方式。你将了解到它的 DeepSeek-V3.2 系 MoE 底座、稀疏门控注意力(SGA)与门控 RMSNorm 等关键改进的内部原理、完整配置项含义,以及如何用 Pipeline 与 AutoModel 开箱即用地加载和生成文本。
背景:A.X-K2 是什么
A.X-K2 是韩国 SK Telecom(SKT)的旗舰大语言模型,由 SK Telecom 于 2026-07-24 贡献到 Hugging Face Transformers。从模型结构上看,它是一款 Mixture-of-Experts(MoE)Decoder,架构底座建立在 DeepSeek-V3.2 之上,核心组件是:
- Multi-head Latent Attention(MLA):低秩潜变量注意力,显著压缩 KV 缓存;
- DeepSeek Sparse Attention(DSA):稀疏注意力机制。
在此基础上,A.X-K2 加入了三项 SK Telecom 自研改进(详见下文"三大改进"),并采用了非分组的 sigmoid top-k 路由(带 correction bias),且首层为 dense 层、其余层为 MoE 层(含一个共享 expert)。
本仓库将其实现为
axk2模型家族,核心文件集中在 src/transformers/models/axk2/ 目录:
- configuration_axk2.py:
AXK2Config配置类及默认值;- modeling_axk2.py:模型前向实现(由 modular_axk2.py 自动生成,文件头部的警告注明:任何改动都应落在 modular 源文件上,CI 会强制校验一致性);
- init.py:模块懒加载入口。
快速上手:文本生成
A.X-K2 支持两种等价的加载方式,可直接体验韩语/多语言文本生成能力。
方式一:Pipeline
from transformers import pipeline
pipe = pipeline(task="text-generation", model="skt/A.X-K2")
print(pipe("대한민국의 수도는", max_new_tokens=32)[0]["generated_text"])
方式二:AutoModel
from transformers import AutoModelForCausalLM, AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained("skt/A.X-K2")
model = AutoModelForCausalLM.from_pretrained("skt/A.X-K2", device_map="auto")
inputs = tokenizer("대한민국의 수도는", return_tensors="pt").to(model.device)
outputs = model.generate(**inputs, max_new_tokens=32, do_sample=False)
print(tokenizer.decode(outputs[0], skip_special_tokens=True))
其中 skt/A.X-K2 是官方文档给出的检查点名称。加载大模型时建议像 AutoModel 示例那样配合 device_map="auto"(需安装 accelerate),Pipeline 也可通过传入 device_map 或 torch_dtype 进一步控制设备与精度。
注意力后端选择:为什么默认推荐 SDPA
官方文档用 [!TIP] 特别强调:
A.X-K2 依赖显式的加性稀疏掩码,因此它在
eager与sdpa两种注意力实现下运行,attn_implementation="sdpa"是默认且推荐的 backend。
这一点在源码中有直接证据。modeling_axk2.py 中 AXK2PreTrainedModel 的能力开关明确写着:
_supports_flash_attn = False # flash-mla kernels need a bit more work in the way we enable them!
_supports_sdpa = True
也就是说,该模型当前不开放标准的 Flash Attention,flash-mla 专用 kernel 的接入还在推进中(代码中 indices=sparse_indices 注释提到该参数会被 flash_mla_with_kvcache 消费,但当前主要由 eager/SDPA 路径消费)。日常推理请直接使用默认的 sdpa,无需额外指定。
注意力前向的骨架:MLA + 显式稀疏掩码
在 AXK2Attention.forward(modeling_axk2.py)中,MLA 的完整计算流如下:
- 通过
q_a_proj压缩 query 得到低秩瓶颈q_resid,再经q_gate_proj同时产出 query 与 attention 输出 gate; kv_a_proj_with_mqa压缩 K/V 为kv_lora_rank潜变量 + 共享的 RoPE key 切片;expand_kv将压缩潜变量展开成完整的 key/value 状态;- Indexer 选出 top-k 位置并构造加性稀疏掩码,与因果掩码一起
masked_fill进attention_mask(modeling_axk2.py)——每个位置只对索引器选中的位置做注意力,其余 token 被置为-inf; - 注意力输出乘上输入相关的 sigmoid gate 后经
o_proj输出。
需要特别留意缓存差异(modeling_axk2.py):由于存在稀疏注意力,该模型在 cache 中缓存的是展开后的完整 K/V(expand_kv 的输出),而非压缩潜变量;同时 indexer 拥有自己独立的 key cache。
SK Telecom 三大改进
这是 A.X-K2 相对 DeepSeek-V3.2 最核心的差异,逐条对照源码展开。
1. Sparse Gated Attention(SGA)——轻量 lightning indexer
定义:每一层都运行一个轻量的 lightning indexer,它对每个 query 与所有 keys 打分,只保留 top-index_topk 的位置,并将其作为加性稀疏掩码折叠进 MLA 注意力中。indexer 在主 KV cache 之外维护一份自己的 key cache(对应 DynamicIndexedLayer / StaticIndexedLayer)。
源码实现:AXK2Indexer(modeling_axk2.py)是 DSA 的完整落地:
- 拥有独立于主 MLA 注意力的轻量投影:
wq_b(作用于 query LoRA 瓶颈)、wk+k_norm(对 key 做 LayerNorm); - 打分逻辑为
index_score[b,s,t] = Σ_h (weight[b,s,h] · softmax_scale · q[b,s,h,:] · k[b,t,:]),即对每个 head 的得分按可学习的weights_proj加权求和,再经过 ReLU; - 因果性通过在打分时叠加
attention_mask或手动构造上三角掩码保证; - 最终返回形状为
[B, S, topk]的 top-k token 索引(int32)。
实现注释还揭示了与参考实现(如 Triton/Flash 版 DSA)的数值等价关系:参考 Indexer 使用 Hadamard 变换(rotate_activation)+ FP8 量化打分 kernel(fp8_index),而本实现因 Hadamard 正交性(Hq·Hk = q·k)与 FP8 仅属精度优化,直接以 bf16/fp32 计算得分,两者数学等价。
AXK2Indexer 的 key cache 通过 past_key_values.update_indexer(k, self.layer_idx)(modeling_axk2.py)更新——这与 DeepSeek-V3 系列中 cache 上新增的 indexer 接口对应,注释说明索引器 key cache 存放在共享 cache 内部、按层索引。
2. Gated RMSNorm —— 低秩输入相关门控
定义:input_layernorm(每一层)与 post_attention_layernorm(仅 MoE 层)都被包装为一个低秩、输入相关的 sigmoid 门控,数学形式为:
RMSNorm(x) * sigmoid(gate_mlp(RMSNorm(x)))
源码实现:AXK2GatedRMSNorm(modeling_axk2.py):
class AXK2GatedRMSNorm(nn.Module):
"""RMSNorm followed by a low-rank input-dependent sigmoid gate (Megatron `GatedNormWrapper`):
y = RMSNorm(x)
return y * sigmoid(gate_mlp(y))
"""
def forward(self, x):
y = self.norm(x)
return (y * torch.sigmoid(self.mlp(y).float())).to(y.dtype)
其中 AXK2GateMLP(modeling_axk2.py)是一个隐藏维度为 gated_norm_rank(默认 16)的双层 SiLU 瓶颈 MLP。逐层差异体现在 AXK2DecoderLayer.__init__(modeling_axk2.py)中:input_layernorm 恒为 AXK2GatedRMSNorm;而 post_attention_layernorm 只有在当前层是 sparse(MoE)层时才用门控版本,dense 层仍使用普通 AXK2RMSNorm。
3. Attention output gate —— 注意力输出门控
定义:注意力输出在进入输出投影前,会乘以一个输入相关的 sigmoid 门(g_proj)。在已发布的检查点中,该门被融合进 q_b_proj(vLLM 布局),权重转换器在加载时再将其拆分开来。
源码实现:实现上,query 与 gate 使用同一个融合投影 q_gate_proj(modeling_axk2.py),输入为 [q_resid, q_compressed] 的拼接,输出切分为 qk_head_dim 的 query 与 v_head_dim 的 gate 两段;前向末尾执行:
attn_output = (attn_output * torch.sigmoid(gate_states.float())).to(attn_output.dtype)
attn_output = self.o_proj(attn_output)
代码注释特别强调 q_gate_proj 必须保持融合状态("needs to be kept fused as the FP8 scales won't match otherwise when split"),即拆分会导致 FP8 量化 scale 失配,这解释了为何发布的 checkpoint 采用 vLLM 融合布局、由转换器在加载期拆分。
AXK2Config:核心配置与默认值
AXK2Config 定义于 configuration_axk2.py,继承 PreTrainedConfig,model_type = "axk2"。它默认给出的是 A.X-K2-Light 规模配置。以下是源码中的完整默认值:
| 配置项 | 默认值 | 说明 |
|---|---|---|
vocab_size |
163840 | 词表大小 |
hidden_size |
2048 | 隐藏维度 |
intermediate_size |
5120 | dense MLP 中间维度 |
moe_intermediate_size |
512 | 单个 expert 的中间维度 |
num_hidden_layers |
48 | 解码层数 |
num_attention_heads |
32 | 注意力头数 |
num_key_value_heads |
32 | KV 头数 |
n_shared_experts |
1 | 共享 expert 数量 |
n_routed_experts |
128 | 路由 expert 数量 |
num_experts_per_tok |
8 | 每个 token 激活的 expert 数 |
routed_scaling_factor |
2.5 | 路由权重缩放因子 |
norm_topk_prob |
True | 是否归一化 top-k 路由权重 |
kv_lora_rank |
128 | K/V 潜变量低秩 |
q_lora_rank |
384 | Query 低秩瓶颈 |
qk_rope_head_dim |
32 | RoPE 部分 head 维度 |
qk_nope_head_dim |
64 | 无 RoPE 部分 head 维度 |
v_head_dim |
64 | value head 维度 |
max_position_embeddings |
131072 | 最大序列长度(128K) |
rms_norm_eps |
1e-6 | RMSNorm epsilon |
bos_token_id / eos_token_id |
163691 | BOS/EOS token id |
tie_word_embeddings |
False | 不共享词嵌入 |
attention_dropout |
0.0 | 注意力 dropout |
hidden_act |
"silu" | 激活函数 |
initializer_range |
0.02 | 初始化范围 |
A.X-K2 特有的稀疏/门控参数
这些参数在标准 DeepSeek 配置上不存在,是本模型文档逐条列出、需重点掌握的部分:
n_group(int | None,默认None):分组路由的专家组数,供更大的 A.X-K2 版本使用。None(A.X-K2-Light 的默认值)表示不分组、在所有专家上直接路由;配置类注释明确说明大版本会设置n_group/topk_group走 DeepSeek-V3 风格的分组路由,因此两种模式都被支持。topk_group(int | None,默认None):当n_group被设置时,top-k 选择被限制在多少个组内。mlp_layer_types(list,默认自动推导):每层的 MLP 类型模式("dense"或"sparse")。未提供时,由旧式 kwargsfirst_k_dense_replace(默认 1,即首层 dense)与moe_layer_freq(默认 1)推导得到(configuration_axk2.py)。这正对应文档中"第一个层为 dense、其余为 MoE"的描述。index_topk(int,默认 2048):索引器为稀疏注意力挑选的 top token 数量。AXK2Indexer中topk = min(self.index_topk, index_scores.shape[-1]),序列变短时会自动回落。index_head_dim(int,默认 128):索引器投影(DSA)的 head 维度。index_n_heads(int,默认 16):索引器投影(DSA)的 head 数量。gated_norm_rank(int,默认 16):AXK2GatedRMSNorm使用的低秩输入相关门控的瓶颈秩。
__post_init__ 中的派生与校验
configuration_axk2.py 在初始化后还会做以下处理:
- 派生 head 维度:
qk_head_dim = qk_nope_head_dim + qk_rope_head_dim(= 96);由于 RoPE 只作用于 rope 切片,head_dim被改写为qk_rope_head_dim(= 32),供继承的旋转位置编码读取。 layer_types:默认全层设置为["deepseek_sparse_attention"],这是为了让 DSA 的 indexer cache 与主 cache 能正确对齐。validate_architecture约束:q_lora_rank必须为正(indexer 与 output gate 都读取 query LoRA 瓶颈);n_group与topk_group必须同时设置或同时为None;设置分组时要求n_routed_experts % n_group == 0且topk_group <= n_group。
AXK2Config 的 attribute_map 把 num_local_experts 映射到 n_routed_experts,以兼容通用代码路径。
模型 API 与支持的任务
模型文档依次列出以下公开类(均以 AXK2 为前缀,可直接从 transformers 顶层导入,并已注册进 auto 映射,见 auto_mappings.py 与 modeling_auto.py):
AXK2Config:模型配置(含from transformers import AXK2Config的可运行示例)。AXK2Model:裸 transformer 主体,输出last_hidden_state与past_key_values(BaseModelOutputWithPast)。它会根据layer_types[i]为每一层分发对应掩码(modeling_axk2.py)。AXK2ForCausalLM:因果语言建模头,继承GenerationMixin以支持generate;提供logits_to_keep以只计算必要 logits,并在提供labels时计算损失(modeling_axk2.py)。AXK2ForSequenceClassification:基于GenericForSequenceClassification的序列分类头。AXK2ForTokenClassification:基于GenericForTokenClassification的 token 分类头。
对直接使用 AXK2ForCausalLM.from_pretrained(...) 的调用,可参考其 docstring 中"加载 → 编码 → generate → batch_decode"的完整示例流程。
模型实现细节与工程要点
路由:非分组 sigmoid top-k + correction bias
AXK2TopkRouter(modeling_axk2.py)实现了文档所说的"plain(non-grouped)sigmoid top-k with a correction bias":
- 路由 logits 经
sigmoid得到分数,再加上一个可学习的e_score_correction_biasbuffer 后做 top-k; - 该 buffer 在 fp32 中维护(
_keep_in_fp32_modules_strict = ["e_score_correction_bias"]); norm_topk_prob=True时对 top-k 权重归一化,最后统一乘以routed_scaling_factor;- 仅当
n_group非空时才进入apply_group_scoring的 DeepSeek 风格分组打分路径——A.X-K2-Light 直接跳过。
MoE 主体:路由专家 + 共享专家
AXK2MoE(modeling_axk2.py)把路由部分与共享专家组合:AXK2Experts 将 expert 权重以 3D 张量存储(gate_up_proj/down_proj),按命中的 expert 逐一分批执行 SwiGLU;AXK2MoE 前向末尾把共享专家输出加到路由专家输出上。
两套 RoPE 布局的差异
实现中同时存在两种旋转位置编码布局,值得注意:
- 主 MLA 注意力使用 DeepSeek 风格的 interleaved(交错) 布局,
apply_rotary_pos_emb_interleave(modeling_axk2.py)直接在奇偶切片上计算旋转,避免view/transpose/reshape的额外拷贝,且与参考实现的位级结果一致; - 索引器则使用非交错(half-split) 布局(见 modeling_axk2.py 注释:"The indexer uses NON-interleaved (half-split) RoPE — unlike the main MLA attention")。
RoPE 支持通过 rope_parameters 配置 rope_theta 与 rope_type(默认 1e6 级别长上下文外推参数由 max_position_embeddings=131072 支撑),动态 rope 更新由 dynamic_rope_update 装饰器处理。
分布式并行支持
从 AXK2Config 可见该模型面向大规模推理/训练的并行方案已内置:base_model_tp_plan(张量并行,含 mla_kv_a_proj、packed_colwise、moe_tp_experts 等切分策略)、base_model_pp_plan(流水线并行)以及 base_model_ep_plan(专家并行,使用 ep_router 与 grouped_gemm)。AXK2ForCausalLM 也声明了 lm_head 的 TP/PP/FSDP 策略。
使用注意事项小结
- 注意力实现:保持默认
sdpa(或eager),本模型基于显式稀疏掩码工作,且_supports_flash_attn = False。 - 缓存行为:生成时缓存的是展开后的 K/V 以及每层索引器的独立 key cache,因此对
past_key_values的处理与常规 MLA 模型不同,属预期设计。 - 配置校验:若自行修改
AXK2Config,需同时遵守q_lora_rank > 0、n_group/topk_group成对出现且满足整除约束等硬性规则,否则初始化会直接抛ValueError。 - 模型规模:仓库默认配置对应 A.X-K2-Light;更大的 A.X-K2 发布版通过设置
n_group/topk_group启用分组路由,加载对应 checkpoint 时配置会自动带上这些字段,无需手工干预。 - API 兼容:
AXK2ForSequenceClassification/AXK2ForTokenClassification走通用分类/序列标注 mixin,可直接沿用 Transformers 标准的微调与评测流程。
如需继续深入,推荐直接阅读 configuration_axk2.py 中的默认值与校验逻辑、modeling_axk2.py 中的 AXK2Indexer/AXK2Attention/AXK2GatedRMSNorm 实现,以及其生成源 modular_axk2.py。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0624
Hy4-previewHy4 preview 是由腾讯混元团队研发的新一代混合专家(MoE)旗舰模型。模型总参数量 770B,每个 token 激活 49B,主干共包含78层,第一层采用标准 FFN,其余 77 层均为 MoE 结构,每层包含 256 个路由专家与 1 个共享专家,每个 token 激活 top-8 路由专家及共享专家。主干之外原生内置 1 层 MTP(总参数量 10B,激活 0.7B)以支持投机解码。Python00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
GLM-5.3-FlashGLM-5.3-Flash (320B-A18B),是GLM-5系列的首个原生多模态模型。320B总参数,能力超过GLM-5.2Jinja00
Spark-X2.5-4BSpark-X2.5-4B 旨在让强大的 AI 更实用、更高效、更易获得。在广泛日常任务中表现强劲,涵盖对话、写作、翻译、推理、编码、工具调用以及智能体工作流,并在同等规模的开源模型中取得领先成绩。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00