Transformers 中的 AXK2 架构解析:SK Telecom A.X-K2 稀疏注意力 MoE 大模型集成指南
2026/9/7 1:40:14 网站建设 项目流程

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_maptorch_dtype进一步控制设备与精度。

注意力后端选择:为什么默认推荐 SDPA

官方文档用[!TIP]特别强调:

A.X-K2 依赖显式的加性稀疏掩码,因此它在eagersdpa两种注意力实现下运行,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 Attentionflash-mla专用 kernel 的接入还在推进中(代码中indices=sparse_indices注释提到该参数会被flash_mla_with_kvcache消费,但当前主要由 eager/SDPA 路径消费)。日常推理请直接使用默认的sdpa,无需额外指定。

注意力前向的骨架:MLA + 显式稀疏掩码

AXK2Attention.forward(modeling_axk2.py)中,MLA 的完整计算流如下:

  1. 通过q_a_proj压缩 query 得到低秩瓶颈q_resid,再经q_gate_proj同时产出 query 与 attention 输出 gate;
  2. kv_a_proj_with_mqa压缩 K/V 为kv_lora_rank潜变量 + 共享的 RoPE key 切片;
  3. expand_kv将压缩潜变量展开成完整的 key/value 状态;
  4. Indexer 选出 top-k 位置并构造加性稀疏掩码,与因果掩码一起masked_fillattention_mask(modeling_axk2.py)——每个位置只对索引器选中的位置做注意力,其余 token 被置为-inf
  5. 注意力输出乘上输入相关的 sigmoid gate 后经o_proj输出。

需要特别留意缓存差异(modeling_axk2.py):由于存在稀疏注意力,该模型在 cache 中缓存的是展开后的完整 K/Vexpand_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,继承PreTrainedConfigmodel_type = "axk2"。它默认给出的是A.X-K2-Light规模配置。以下是源码中的完整默认值:

配置项默认值说明
vocab_size163840词表大小
hidden_size2048隐藏维度
intermediate_size5120dense MLP 中间维度
moe_intermediate_size512单个 expert 的中间维度
num_hidden_layers48解码层数
num_attention_heads32注意力头数
num_key_value_heads32KV 头数
n_shared_experts1共享 expert 数量
n_routed_experts128路由 expert 数量
num_experts_per_tok8每个 token 激活的 expert 数
routed_scaling_factor2.5路由权重缩放因子
norm_topk_probTrue是否归一化 top-k 路由权重
kv_lora_rank128K/V 潜变量低秩
q_lora_rank384Query 低秩瓶颈
qk_rope_head_dim32RoPE 部分 head 维度
qk_nope_head_dim64无 RoPE 部分 head 维度
v_head_dim64value head 维度
max_position_embeddings131072最大序列长度(128K)
rms_norm_eps1e-6RMSNorm epsilon
bos_token_id/eos_token_id163691BOS/EOS token id
tie_word_embeddingsFalse不共享词嵌入
attention_dropout0.0注意力 dropout
hidden_act"silu"激活函数
initializer_range0.02初始化范围

A.X-K2 特有的稀疏/门控参数

这些参数在标准 DeepSeek 配置上不存在,是本模型文档逐条列出、需重点掌握的部分:

  • n_groupint | None,默认None):分组路由的专家组数,供更大的 A.X-K2 版本使用。None(A.X-K2-Light 的默认值)表示不分组、在所有专家上直接路由;配置类注释明确说明大版本会设置n_group/topk_group走 DeepSeek-V3 风格的分组路由,因此两种模式都被支持。
  • topk_groupint | None,默认None):当n_group被设置时,top-k 选择被限制在多少个组内。
  • mlp_layer_typeslist,默认自动推导):每层的 MLP 类型模式("dense""sparse")。未提供时,由旧式 kwargsfirst_k_dense_replace(默认 1,即首层 dense)与moe_layer_freq(默认 1)推导得到(configuration_axk2.py)。这正对应文档中"第一个层为 dense、其余为 MoE"的描述。
  • index_topkint,默认 2048):索引器为稀疏注意力挑选的 top token 数量。AXK2Indexertopk = min(self.index_topk, index_scores.shape[-1]),序列变短时会自动回落。
  • index_head_dimint,默认 128):索引器投影(DSA)的 head 维度。
  • index_n_headsint,默认 16):索引器投影(DSA)的 head 数量。
  • gated_norm_rankint,默认 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_grouptopk_group必须同时设置或同时为None;设置分组时要求n_routed_experts % n_group == 0topk_group <= n_group

AXK2Configattribute_mapnum_local_experts映射到n_routed_experts,以兼容通用代码路径。

模型 API 与支持的任务

模型文档依次列出以下公开类(均以AXK2为前缀,可直接从transformers顶层导入,并已注册进auto映射,见 auto_mappings.py 与 modeling_auto.py):

  • AXK2Config:模型配置(含from transformers import AXK2Config的可运行示例)。
  • AXK2Model:裸 transformer 主体,输出last_hidden_statepast_key_valuesBaseModelOutputWithPast)。它会根据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_thetarope_type(默认 1e6 级别长上下文外推参数由max_position_embeddings=131072支撑),动态 rope 更新由dynamic_rope_update装饰器处理。

分布式并行支持

AXK2Config可见该模型面向大规模推理/训练的并行方案已内置:base_model_tp_plan(张量并行,含mla_kv_a_projpacked_colwisemoe_tp_experts等切分策略)、base_model_pp_plan(流水线并行)以及base_model_ep_plan(专家并行,使用ep_routergrouped_gemm)。AXK2ForCausalLM也声明了lm_head的 TP/PP/FSDP 策略。

使用注意事项小结

  1. 注意力实现:保持默认sdpa(或eager),本模型基于显式稀疏掩码工作,且_supports_flash_attn = False
  2. 缓存行为:生成时缓存的是展开后的 K/V 以及每层索引器的独立 key cache,因此对past_key_values的处理与常规 MLA 模型不同,属预期设计。
  3. 配置校验:若自行修改AXK2Config,需同时遵守q_lora_rank > 0n_group/topk_group成对出现且满足整除约束等硬性规则,否则初始化会直接抛ValueError
  4. 模型规模:仓库默认配置对应 A.X-K2-Light;更大的 A.X-K2 发布版通过设置n_group/topk_group启用分组路由,加载对应 checkpoint 时配置会自动带上这些字段,无需手工干预。
  5. 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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询