Transformers 中的 AXK2 架构解析:SK Telecom A.X-K2 稀疏注意力 MoE 大模型集成指南
【免费下载链接】transformers🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers
本文以 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。
【免费下载链接】transformers🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考