Kimi K3 核心算子拆解——KDA、FlashKDA、Attention Residuals 的实现算法
副标题: Kimi K3(2.8T MoE)今日正式开源,47 页技术报告同步发布。本文脱离架构层面的泛泛介绍,直接切入最硬核的部分——Kimi Delta Attention 的 DPLR 状态更新方程怎么算?Fine-Grained Gating 和 Gated DeltaNet 的区别在哪?FlashKDA 的 Chunkwise 并行算法为何比基线快 2 倍?MoonEP 如何实现数学可证的完美负载均衡?SiTU-GLU 的 tanh 限幅在数值上怎么工作?全部用 PyTorch 风格伪代码 + 公式拆解。
一、从架构说起:KDA + Gated MLA 的 3:1 混合
Kimi K3 有 93 层,但不是所有层都用同一种注意力。技术报告披露了混合比例:
第 0-2 层: KDA × 3
第 3 层: Gated MLA × 1
第 4-6 层: KDA × 3
第 7 层: Gated MLA × 1
...(重复上述模式)
━━━━━━━━━━━━━━━━━━━━━━━━━━━━
总计: KDA × 69 层(占 74%)
Gated MLA × 24 层(占 26%)
为什么这么混?
- KDA(线性注意力):状态大小恒定 O(1),不随序列长度增长。适合处理 100 万 token 的长上下文——KV Cache 几乎为零。
- Gated MLA(门控多头潜在注意力):标准的 MLA(和 DeepSeek V3/V4、GLM-5.2 同族),带门控机制增强局部表征质量。上下文短的时候质量更好,但显存开销随序列增长。
3:1 比例的直觉:大部分计算走 KDA 的低成本路径,每隔 3 层用一次 MLA 做"全局校对",防止信息在 KDA 的循环状态中衰减太多。
二、Kimi Delta Attention——线性注意力的新形态
2.1 从 Linear Attention 到 DeltaNet 到 KDA
KDA 不是凭空出现的。它是线性注意力这条技术路线的最新演进:
传统 Softmax Attention: O = softmax(QK^T)V → O(n²)
↓
Linear Attention (2019): O = Q(K^T V) → O(n),但无遗忘机制
↓
DeltaNet (2024): S_t = β_t k_t v_t^T + (I - β_t k_t k_t^T) S_{t-1}
→ 引入门控,可遗忘旧信息
↓
Gated DeltaNet (2025): + α_t 标量遗忘门 → 每头一个遗忘率
↓
KDA (Kimi Delta): + Diag(α_t) 细粒度门控 → 每特征维度一个遗忘率
KDA 的核心贡献就一句话:从每头一个标量遗忘门,变成每特征维度一个遗忘门。
2.2 KDA 的数学形式
KDA 的循环状态更新方程:
S_t = (I - β_t k_t k_t^T) · Diag(α_t) · S_{t-1} + β_t k_t v_t^T
其中:
S_t ∈ R^(d_k × d_v)— 循环状态(Linear Attention 中的"KV Cache"替代品)k_t ∈ R^(d_k)— 当前 token 的 key(经过 short conv 后)v_t ∈ R^(d_v)— 当前 token 的 valueα_t ∈ (0,1)^(d_k)— 细粒度遗忘门,每个特征维度独立β_t ∈ (0,1)— 写入门(标量)
关键名词解释:
| 符号 | 维度 | 含义 | 和 Gated DeltaNet 的差异 |
|---|---|---|---|
α_t |
(d_k,) 向量 |
每维遗忘率 | GDN 用的是标量 α_t |
Diag(α_t) |
(d_k, d_k) 对角矩阵 |
将向量 α_t 转为对角形式 | 有了它,KDA 对每个特征维度可以独立控制"记住多少" |
β_t |
() 标量 |
写入强度 | 同 GDN |
S_t |
(d_k, d_v) 矩阵 |
循环记忆状态 | 同 GDN |
2.3 三步计算法
KDA 的状态更新虽然看起来是一个公式,但计算上分三步,利用了 Diagonal-Plus-Low-Rank (DPLR) 结构来避免 O(d_k²·d_v) 的昂贵运算:
def kda_step(S_prev: Tensor, k: Tensor, v: Tensor, alpha: Tensor, beta: float) -> Tensor:
"""
KDA 单步状态更新
S_prev: (d_k, d_v) 前一步状态
k: (d_k,) 当前 key(短卷积后)
v: (d_v,) 当前 value
alpha: (d_k,) 逐维遗忘门(sigmoid 后,值域 (0,1))
beta: scalar 写入门
"""
d_k, d_v = S_prev.shape
# Step 1: 对角衰减(Diagonal Decay)
# 每个特征维度独立衰减 → O(d_k × d_v)
S_decayed = S_prev * alpha.unsqueeze(-1) # (d_k, d_v) broadcast
# Step 2: Rank-1 纠偏(Delta Correction)
# 从衰减后的状态中减去 k 方向的分量
# 直觉:避免重复写入已经存在的信息
k_t_S = k @ S_decayed # (d_k,) @ (d_k, d_v) → (d_v,)
S_corrected = S_decayed - beta * k.unsqueeze(-1) @ k_t_S.unsqueeze(0)
# (d_k, d_v) - (d_k, 1) @ (1, d_v) = (d_k, d_v)
# Step 3: KV 写入
S_new = S_corrected + beta * k.unsqueeze(-1) @ v.unsqueeze(0)
# (d_k, d_v) + (d_k, 1) @ (1, d_v) = (d_k, d_v)
return S_new
复杂度对比:
| 实现 | 每步 FLOPs | 说明 |
|---|---|---|
| 朴素 DPLR | O(d_k²·d_v) | 直接计算完整的矩阵乘法 |
| KDA 三步法 | O(d_k·d_v) | 利用 DPLR 结构,只算了两次外积 |
| 加速比 | ~d_k / 2 | d_k=4096 时快 ~2000 倍 |
2.4 Fine-Grained Gating 的生成
α_t 不是独立学习的参数,而是由当前 token 的输入通过一个低秩瓶颈网络生成:
def compute_fine_grained_gate(x: Tensor, W_alpha: Tensor, W_alpha_down: Tensor) -> Tensor:
"""
生成逐维遗忘门 α_t
x: (d_model,) 当前 token 的隐藏状态
W_alpha: (d_gate, d_model) 门控投影(瓶颈层)
W_alpha_down: (d_k, d_gate) 扩展到 key 维度
"""
# 低秩瓶颈:d_model → d_gate → d_k
# d_gate 远小于 d_model 和 d_k(通常是 64-128)
gate_hidden = F.silu(W_alpha @ x) # (d_gate,)
alpha_logits = W_alpha_down @ gate_hidden # (d_k,)
alpha = torch.sigmoid(alpha_logits) # (d_k,),值域 (0,1)
# 强制下界防止数值不稳定
alpha = alpha * 0.99 + 0.01 # 保证 α ∈ [0.01, 1.0]
return alpha
为什么用低秩瓶颈? 如果直接从 d_model → d_k(假设 d_k=4096),参数量是巨大的。低秩瓶颈(d_model → 64 → 4096)把参数量从 d_model×d_k 降到 (d_model+d_k)×64,大约 64x 的节省。
2.5 输出生成
状态 S_t 不是直接输出,还需要做一步内容检索:
def kda_output(q: Tensor, S: Tensor, k: Tensor, W_out_gate: Tensor) -> Tensor:
"""
KDA 输出生成
q: (d_k,) query
S: (d_k, d_v) 当前循环状态
k: (d_k,) 当前 key
W_out_gate: (d_out, d_v) 输出门控投影
"""
# Step 1: 从状态中读取内容
retrieved = q @ S # (d_k,) @ (d_k, d_v) → (d_v,)
# Step 2: RMSNorm 稳定化
retrieved = F.rms_norm(retrieved) # (d_v,)
# Step 3: 输出门控
gate = torch.sigmoid(W_out_gate @ retrieved) # (d_out,)
output = gate * retrieved # 门控后的输出
return output
2.6 Short Convolution on Key
KDA 的 key 在进入循环前要先经过一个深度可分离因果卷积(kernel=4):
def kda_short_conv(k_raw: Tensor, conv_weight: Tensor) -> Tensor:
"""
Key 的短卷积预处理(因果卷积,kernel=4)
k_raw: (d_k,) 原始 key
conv_weight: (4, d_k) 卷积权重(深度可分离,每组独立)
"""
# 需要缓存最近 3 步的 key
# cache: (3, d_k) 之前 3 步的 key
k_cache = update_cache(k_raw)
# 因果卷积:只看过去和当前,不看未来
k_conv = (conv_weight[0] * k_cache[0] + # t-3
conv_weight[1] * k_cache[1] + # t-2
conv_weight[2] * k_cache[2] + # t-1
conv_weight[3] * k_raw) # t
return k_conv
这个 short conv 的作用是给 KDA 提供局部上下文感知能力——循环状态 S 记忆的是全局信息,但 key 本身需要知道附近 token 的语境。
三、Chunkwise 并行算法——FlashKDA 的核心
3.1 为什么需要 Chunkwise 算法?
KDA 是循环的——S_t 依赖 S_{t-1},看起来必须串行计算。但在 Prefill 阶段,所有 token 是同时可用的,我们想利用 GPU 的并行能力一次处理多个 token。
Chunkwise 算法的思路:把序列切成 chunk(如 64-128 tokens),chunk 内部用并行矩阵乘法,chunk 之间串行传递状态。
串行(逐 token):
S_0 → S_1 → S_2 → ... → S_L-1
L 步串行,每步 O(d_k·d_v)
Chunkwise(chunk_size = C):
[S_0 → ... → S_{C-1}] → [S_C → ... → S_{2C-1}] → ...
↑ chunk 内部并行 ↑ chunk 内部并行
L/C 步串行,每步 O(C·d_k·d_v) 的并行 matmul
3.2 WY 表示
Chunkwise 算法的关键数学工具是 WY 表示——把 chunk 内的多个 rank-1 更新打包成紧凑的矩阵形式。
回忆 KDA 的更新中有两个 rank-1 项:
-β_t k_t (k_t^T S_{t-1})— 纠偏项+β_t k_t v_t^T— 写入项
一个 chunk 内有 C 个这样的 rank-1 更新。WY 表示把 C 个更新合并成两个矩阵乘法:
def kda_chunkwise(K_chunk, V_chunk, alpha_chunk, beta_chunk, S_prev):
"""
Chunkwise KDA 前向
K_chunk: (C, d_k) chunk 内所有 key
V_chunk: (C, d_v) chunk 内所有 value
alpha_chunk: (C, d_k) chunk 内所有遗忘门
beta_chunk: (C,) chunk 内所有写入门
S_prev: (d_k, d_v) 上一个 chunk 传过来的状态
"""
C = K_chunk.shape[0]
# Step 1: 计算 chunk 内的累积衰减
# 每个 token 的衰减是累积的——后面的 token 受前面所有衰减影响
cum_alpha = torch.cumprod(alpha_chunk, dim=0) # (C, d_k)
# Step 2: WY 表示——将 rank-1 更新打包
# 这部分是 FlashKDA 的核心优化,用两个矩阵乘法代替 C 个循环
# 具体实现涉及 UT 变换(避免矩阵求逆)
P = compute_wy_representation(K_chunk, V_chunk, alpha_chunk, beta_chunk)
# Step 3: 一次 matmul 完成 chunk 内所有 token 的 attention
outputs = chunkwise_attention(K_chunk, S_prev, P)
# 输入: (C, d_k) @ (d_k, d_v) + WY correction → (C, d_v)
# Step 4: 计算传递给下一 chunk 的状态
S_new = update_state(S_prev, K_chunk, V_chunk, alpha_chunk, beta_chunk)
return outputs, S_new
3.3 FlashKDA 比 FLA 快 1.72-2.22 倍的原因
技术报告指出,FlashKDA 用了两个关键优化:
-
UT 变换减少非 matmul FLOPs:传统的 DPLR chunkwise 实现需要矩阵求逆(O(C³)),UT 变换将其简化为前向替换(O(C²)),且把更多计算转化为 Tensor Core 友好的 matmul。
-
绑定两个 DPLR 变量到 k:因为 KDA 的纠偏项和写入项都用同一个 k(不是像其他线性注意力那样用不同的投影),第二级 chunk 矩阵的计算从 4 个降为 2 个,省了一半。
非 matmul FLOPs 占比(Prefill, 512K 上下文):
FLA baseline: ~18% 非 matmul(大部分在矩阵求逆)
FlashKDA: ~6% 非 matmul(UT 变换 + 变量绑定)
→ Tensor Core 利用率从 ~82% 提升到 ~94%
→ Prefill 端到端加速 1.72-2.22x
四、Attention Residuals——跨层特征检索
4.1 问题:深层网络的信息稀释
标准 Transformer 中,每层只能看到上一层的输出。当网络有 93 层时,底层的特征经过层层变换,到顶层可能已经被稀释了。
Attention Residuals(AttnRes)的思路:让每一层不仅看到上一层的输出,还能选择性检索前面所有层(或前面 Block 内所有层)的特征。
4.2 Block 级设计
Kimi K3 把 93 层分成 9 个 Block,每个 Block 约 10 层:
Block 0: layers 0-9 → 可检索: Block 0 内所有层
Block 1: layers 10-19 → 可检索: Block 0-1 内所有层
Block 2: layers 20-29 → 可检索: Block 0-2 内所有层
...
Block 8: layers 80-92 → 可检索: 前面所有 Block 的层
每层的 attention 输出不是单纯的 Attention(Q, K, V),而是:
def attention_with_residuals(q, k_self, v_self, k_residual_bank, v_residual_bank):
"""
带注意力残差的 attention
q: 当前层的 query
k_self, v_self: 当前层的 key/value
k_residual_bank: 前面所有可供检索的 key 集合
v_residual_bank: 对应的 value 集合
"""
# 当前层的 attention
self_attn = attention(q, k_self, v_self)
# 跨层检索(只对 q 和残差 bank 做 attention)
# 这是一个轻量级的检索——head_dim 可以更小
cross_attn = cross_attention(q, k_residual_bank, v_residual_bank)
# 融合
output = self_attn + cross_attn
return output
注意:跨层检索不是对完整的 k/v 做 attention,而是对每个 Block 末层的某些特征做轻量检索(具体实现细节尚待技术报告披露更多)。
4.3 AttnRes 的代价控制
理论上,如果每层都检索前面所有层,计算量是 O(L²) 的——等于回了 softmax attention 的老路。
Kimi K3 的控制方法:
- Block 级粒度:检索只在 Block 边界进行,不是每步都做
- 特征压缩:存储的不是完整的 k/v,而是投影到更低维度的"残差特征"
- 选择性检索:通过一个可学习的门控决定"当前层需要从前面拿多少信息"
五、Stable LatentMoE + SiTU-GLU
5.1 架构概览
Kimi K3 的 MoE 配置:
路由专家数: 896
共享专家数: 2
每 token 激活: 16(路由)+ 2(共享)= 18 个 expert
稀疏比: 896:16 = 56:1(~1.8% 激活)
激活参数量: ~104B / token
56:1 的稀疏比是什么概念?作为对比:
- Mixtral 8x7B: 8:2 = 4:1
- Qwen3-30B-A3B: 128:8 = 16:1
- DeepSeek V3: 256:8 = 32:1
56:1 是目前开源 MoE 中最高的稀疏比。 高稀疏比的好处是可以用更多 expert 来容纳知识,但代价是训练不稳定——每个 token 只激活极少数 expert,梯度信号稀疏,容易发散。
5.2 SiTU-GLU 激活函数
SiTU-GLU 针对这个稳定性问题做了专门设计。它的全称是 Sigmoid-Tanh-Unit Gated Linear Unit:
def situ_glu(x: Tensor, W_gate: Tensor, W_up: Tensor) -> Tensor:
"""
SiTU-GLU 前向
x: (d_model,) 输入
W_gate: (d_intermediate, d_model) 门控投影
W_up: (d_intermediate, d_model) 值投影
"""
gate_logits = W_gate @ x # (d_intermediate,)
up_logits = W_up @ x # (d_intermediate,)
# SiTU 激活 = 4 · tanh(x/4) · sigmoid(x)
# 对比:SiLU = x · sigmoid(x)
# 对比:GELU = x · Φ(x)
gate = 4 * torch.tanh(gate_logits / 4) * torch.sigmoid(gate_logits)
# Up 分支也限幅:25 · tanh(x/25)
# 标准做法是线性的(SiLU/SwiGLU 中 up 分支不做激活)
up = 25 * torch.tanh(up_logits / 25)
# 门控乘法
hidden = gate * up # (d_intermediate,)
return hidden
SiTU 和 SiLU 的数值对比:
| 函数 | 公式 | 值域 | 梯度范围 |
|---|---|---|---|
| SiLU | x·sigmoid(x) | (-∞, ∞) | 无上界 |
| GELU | x·Φ(x) | 近似 (-∞, ∞) | 无上界 |
| SiTU | 4·tanh(x/4)·sigmoid(x) | (-1, 1) | 有界 |
SiTU 的值域被限制在 (-1, 1),因此:
- 前向传播不会出现极端大的激活值
- 反向传播的梯度也被限制,不会爆炸
- 在高稀疏度(56:1)的 MoE 中,每个 expert 的梯度信号本来就很稀疏,如果还让激活值无界,一个异常 token 就能让对应的 expert 权重产生巨大更新,破坏训练稳定性
5.3 Quantile Balancing——分位数负载均衡
这是 Kimi K3 在负载均衡上的核心创新。
传统方法(Auxiliary Loss):
给 loss 加一个辅助项,惩罚负载不均。问题是辅助 loss 的权重是个超参数,调起来很痛苦——太大了影响模型质量,太小了不管用。
Quantile Balancing 的思路:不用辅助 loss,用分位数直接决定 token 分配。
def quantile_balancing(routing_logits, top_k, expert_capacity=None):
"""
分位数负载均衡
routing_logits: (num_tokens, num_experts) router 输出分数
top_k: 每个 token 选的 expert 数
Returns: token_indices, expert_indices 的分配对
"""
num_tokens, num_experts = routing_logits.shape
# Step 1: 对每个 expert 的分数分布,计算分位数
# 不需要用辅助 loss 来"激励"均衡——直接在分配时强制均衡
for expert in range(num_experts):
scores = routing_logits[:, expert] # 所有 token 对这个 expert 的分数
# Step 2: 每个 expert 选择分数最高的 capacity 个 token
# capacity = num_tokens * top_k / num_experts(理论均衡分配)
# 但 Quantile Balancing 用分位数来决定——哪个 token 应该被哪个 expert 处理
quantile = compute_quantile(scores, capacity_ratio)
mask = scores > quantile # 高于分位数的 token 被这个 expert 选中
# Step 3: 没有"被 drop"的 token
# 因为分位数是动态调整的,保证了每个 expert 恰好选到 capacity 个 token
# 同时也保证了每个 token 恰好被 top_k 个 expert 处理
# -> 完美均衡,数学可证
return assignment
为什么说"数学可证"? 因为分位数方法是确定性的——给定路由分数,分位数阈值就确定了,每个 expert 分配到的 token 数就是固定的。不需要启发式更新,不需要调超参数。
分布式实现: 在 2.8T 规模的训练中,无法把所有 token 的分数集中到单卡上算分位数。MoonEP 团队用分布式直方图近似——每张卡算自己的直方图,然后 all-reduce 合并,再用合并后的直方图估算全局分位数。近似误差可控在 1% 以内。
六、MoonEP——动态冗余专家并行
6.1 传统 EP 的痛点
标准 Expert Parallelism 中,每张卡持有部分 expert。当路由不均衡(某些 expert 特别热门)时:
传统 EP 的问题(EP=8,896 experts,每卡 112 experts):
GPU 0 (expert 0-111): ████████████████████ 350 tokens ← 热点 expert
GPU 1 (expert 112-223): ████████████ 200 tokens
...
GPU 7 (expert 784-895): ████ 80 tokens ← 大量闲置
瓶颈在 GPU 0(最慢的卡决定步时间)
6.2 MoonEP 的解决方案
MoonEP 的核心思想:给热点 expert 动态添加冗余副本,让每张卡处理的 token 数完全相同。
def moonep_dispatch(routing_result, num_redundant_experts, ep_size):
"""
MoonEP 动态冗余 expert 分配
routing_result: 每 token 选中的 expert ID
num_redundant_experts: 在线规划的冗余 expert 数量
ep_size: GPU 数量
"""
# Step 1: 统计每个 expert 的 token 数分布
expert_counts = count_tokens_per_expert(routing_result)
# Step 2: 找到热点 expert,规划冗余
hot_experts = find_hot_experts(expert_counts)
for expert_id in hot_experts:
# 给热点 expert 添加冗余副本到其他 GPU
add_redundant_replica(expert_id, target_gpu)
# Step 3: 重新分配——保证每张卡恰好 S × K 个 token
# S = 每卡 token 数, K = top_k
# 用冗余 expert 填充亏空,使得所有卡完全均衡
balanced_assignment = perfect_balance(routing_result, redundant_plans)
# Step 4: Zero-Copy 调度
# token 通过 NVLink 对称内存直接写到目标 GPU 的 expert 分组位置
# 不需要中间缓冲区,不需要 memcpy
nvlink_dispatch(balanced_assignment)
return balanced_assignment
6.3 MoonEP vs DeepEP v2
| 维度 | DeepEP v2 | MoonEP |
|---|---|---|
| 应对路由不均衡 | 没有特殊处理,会 OOM | 动态冗余 expert,完美均衡 |
| 缓冲区 | 动态分配(不均衡时 OOM) | 静态 S×K 缓冲区,永不 OOM |
| 数据拷贝 | comm→user buffer 拷贝 | Zero-Copy(NVLink 对称内存) |
| 同步开销 | 每层 host sync | 无需 host sync,全设备端 |
| 迭代时间(不均衡 30%) | 增长 40%+ | 不变 |
七、综合数据流:一个 token 在 KDA 层中的完整旅程
把以上所有算子串起来,一个 token 在 KDA 层中的完整路径:
def kda_layer_forward(x, S_prev, params):
"""
KDA 层前向(单 token)
x: (d_model,) 输入
S_prev: (d_k, d_v) 上一层传来的循环状态
"""
# 1. Q/K/V 投影
q = params.W_q @ x # (d_model) → (d_k,)
k_raw = params.W_k @ x # (d_model) → (d_k,)
v = params.W_v @ x # (d_model) → (d_v,)
# 2. Short Conv on K(kernel=4,因果卷积)
k = short_conv(k_raw, params.conv_weight)
# 3. 生成 Fine-Grained Gating
alpha = compute_fine_grained_gate(x, params.W_alpha, params.W_alpha_down)
beta = torch.sigmoid(params.W_beta @ x) # 写入门
# 4. KDA 状态更新(三步 DPLR)
S_new = kda_step(S_prev, k, v, alpha, beta)
# 5. 生成输出(状态检索 + 门控)
output = kda_output(q, S_new, k, params.W_out_gate)
# 6. 如果有跨层残差检索(每 3 层一次)
if params.enable_attn_residuals:
cross = cross_attention(q, params.k_residual_bank, params.v_residual_bank)
output = output + cross
# 7. MoE FFN(SiTU-GLU + 路由)
expert_ids = router(output) # top-16 路由
ffn_output = situ_glu_moe(output, expert_ids) # 16 experts 的 SiTU-GLU
output = output + ffn_output # 残差连接
return output, S_new
八、总结
Kimi K3 核心算子一览
| 算子/技术 | 核心算法贡献 | 量化收益 |
|---|---|---|
| KDA | DPLR 状态更新 + Fine-Grained Gating | O(n) 注意力的 KV Cache 趋近于零 |
| FlashKDA | WY 表示 + UT 变换 + 变量绑定 | Prefill 速度是 FLA 基线的 1.72-2.22x |
| Attention Residuals | Block 级跨层特征检索 | 93 层信息流动改善,等效深度增加 |
| SiTU-GLU | tanh 限幅 + sigmoid 门控 | 56:1 稀疏度下训练稳定收敛 |
| Quantile Balancing | 分位数分配代替辅助 loss | 完美负载均衡,无需调超参 |
| MoonEP | 动态冗余 expert + Zero-Copy | EP 通信时间恒定不随不均衡增长 |
| Short Conv on Key | 深度可分离因果卷积 | 给线性注意力提供局部感知能力 |
一些值得关注的工程视角
-
KDA 的 DPLR 三步法是一个很好的"数学推导 → 工程实现"案例——同一个公式,朴素实现是 O(d_k²·d_v),拆解后是 O(d_k·d_v),差了 d_k/2 倍。
-
Chunkwise 算法的核心挑战不是并行计算本身,而是最小化非 matmul FLOPs——Tensor Core 只擅长 matmul,任何非 matmul 操作都是瓶颈。FlashKDA 的 UT 变换就是为了把尽可能多的计算转成 matmul。
-
SiTU-GLU 的 tanh 限幅思路值得在其他高稀疏度场景借鉴——当计算路径上的信号非常稀疏时,给激活值加一个有界约束可以大幅提高稳定性。
-
**MoonEP 的"用冗余换均衡"**是一个经典的系统设计 tradeoff——多占一点显存(冗余 expert),换来恒定的计算时间和无 OOM 风险。
九、算子开发视角:自研芯片部署的完整算子清单
以下内容从算子开发工程师的视角出发——如果你要在自研芯片(如昇腾、寒武纪、摩尔线程等)上部署 Kimi K3,需要实现哪些算子?每个算子的计算模式、数据流和 shape 特征是什么?哪些需要重点优化?
9.1 全模型算子总表
Kimi K3 的前向推理包含以下算子(按出现频率排序):
| # | 算子 | 出现位置 | 计算模式 | 算力占比 | 访存特征 |
|---|---|---|---|---|---|
| 1 | MatMul (GEMM) | QKV 投影、O 投影、FFN gate/up/down、Router | 矩阵乘 | ~65% | 计算密集,大矩阵 |
| 2 | KDA State Update (DPLR) | KDA 层(69 层) | 逐元素 + 外积 | ~8% | 访存密集,小矩阵 |
| 3 | Grouped GEMM | MoE FFN(896 expert 的 gate/up/down) | 分组矩阵乘 | ~10% | 分组的计算密集 |
| 4 | SiTU / SiLU 激活 | FFN gate 分支 | 逐元素 | ~2% | 访存密集 |
| 5 | RMSNorm | 每层输入、attention 输出 | 逐元素规约 | ~2% | 访存密集 |
| 6 | Short Conv 1D | KDA 的 key 预处理(69 层) | 因果卷积 | ~1% | 访存密集 |
| 7 | Fine-Grained Gate | KDA 的 α_t 生成(69 层) | 低秩 MatMul + sigmoid | ~1% | 计算密集(小) |
| 8 | Softmax | MLA 层(24 层) | 逐元素规约 | ~1% | 访存密集 |
| 9 | FlashAttention / MLA | 24 层 Gated MLA | 矩阵乘 + softmax | ~8% | 混合 |
| 10 | Residual Add | 每层残差连接 | 逐元素加法 | <1% | 访存密集 |
| 11 | All-to-All | EP 模式下 MoE 层间通信 | 通信 | ~2% (通信时间占比) | 带宽密集 |
| 12 | Top-K Router | MoE 的路由选择 | 排序/选择 | <1% | 访存密集 |
| 13 | Cross Attention (AttnRes) | 每 Block 末层 | 矩阵乘 + softmax | ~2% | 混合 |
关键结论:MatMul 占 ~65% 的算力消耗——这是自研芯片最需要优化的算子。其余算子虽然算力占比小,但访存模式各异,可能成为带宽瓶颈。
9.2 各算子的数据流与 Shape 特征
9.2.1 QKV 投影 MatMul(计算密集,GEMM 核心战场)
输入: x (d_model,) = (7168,) ← 假设 d_model=7168
权重: W_q (d_model, d_k) = (7168, 4096)
W_k (d_model, d_k) = (7168, 4096)
W_v (d_model, d_v) = (7168, 4096)
输出: q, k, v 各 (4096,)
计算模式: [7168] × [7168, 4096] → [4096]
M=1, N=4096, K=7168 的 GEMV(batch=1 时)
M=batch, N=4096, K=7168 的 GEMM(prefill 时)
Shape 关键特征:
- Decode (batch=1):GEMV,极度访存密集——瓶颈在 HBM 带宽(读权重),不在计算
- Prefill (batch>>1):GEMM,计算密集——可以打满 Tensor Core
对自研芯片的意义:
- QKV 投影是 最频繁的 GEMM 调用——每个 token 每层做 3 次(Q/K/V),93 层就是 279 次
- 必须支持 GEMV 和 GEMM 的无缝切换(batch=1 和 batch>1 的最优 kernel 不同)
- BF16/FP16 Tensor Core 支持是刚需
9.2.2 KDA DPLR 状态更新(访存密集,线性注意力的独特算子)
这是 KDA 独有的算子,标准 Transformer 和 MLA 都没有。
输入: S_prev (d_k, d_v) = (4096, 4096)
k (d_k,) = (4096,)
v (d_v,) = (4096,)
alpha (d_k,) = (4096,)
beta () = (1,)
三步计算:
Step 1 (对角衰减): S_decayed = S_prev * alpha[:, None]
逐元素乘,broadcast: (4096, 4096) × (4096, 1)
O(16.7M) 元素操作
Step 2 (Rank-1 纠偏): k_t_S = k @ S_decayed → (4096,) @ (4096, 4096) = (4096,)
S_corrected = S_decayed - β·k·k_t_S
一次 GEMV + 一次外积
Step 3 (KV 写入): S_new = S_corrected + β·k·v
一次外积
访存特征:
- 每一步都要读写 S(16.7M 元素 = 32MB FP16)
- 计算量约 50M FLOPs,但访存量 ~100MB——Op:Byte ≈ 0.5:1,极度访存密集
- 这是自研芯片需要特别关注的算子——标准 GPU 的 Tensor Core 在这里帮不上忙(太多逐元素操作)
可能的优化方向:
- 对 S 做分块(tiling),利用片上 SRAM 减少 HBM 读写
- 融合三步计算,减少中间结果的写回
9.2.3 MoE Grouped GEMM(计算密集,但形状特殊)
每层 MoE FFN:
输入: x (d_model,) = (7168,)
路由: 16 experts(每 expert: intermediate=2048)
每个 expert 的 3 个 GEMM:
Gate: [7168] × [2048, 7168] → [2048]
Up: [7168] × [2048, 7168] → [2048]
Down: [2048] × [7168, 2048] → [7168]
Grouped GEMM 合并后:
输入: 16 × [7168] 打包为 [16, 7168]
Gate: [16, 7168] × [16, 2048, 7168] → [16, 2048] (一次 grouped 调用)
Down: [16, 2048] × [16, 7168, 2048] → [16, 7168] (一次 grouped 调用)
Shape 关键特征:
- Grouped GEMM 的核心分歧点:
- 小 batch(≤64):用 grouped 一次算完,省 launch 开销
- 大 batch(≥256):用标准 GEMM + 稀疏 mask,token 越多 grouped 越低效
- 每个 expert 的 intermediate 维度(2048)较小——不是典型的"大 GEMM"
对自研芯片的挑战:
- 需要支持 grouped GEMM 语义——一次 dispatch 处理多个独立的小 GEMM
- 或者退而求其次:高效的小 GEMM(M≤16, K=7168, N=2048) 的 cublas 替代
- 如果芯片不支持 grouped GEMM,串行 16 次小 GEMM 的 launch 开销会很大
9.2.4 SiTU-GLU(访存密集,融合机会大)
输入: gate_logits (2048,) , up_logits (2048,)
SiTU: gate = 4 · tanh(gate_logits/4) · sigmoid(gate_logits)
Up 限幅: up = 25 · tanh(up_logits/25)
乘法: hidden = gate * up
每个元素: 1次除法 + 2次exp(sigmoid) + 2次tanh + 3次乘法 = ~8 FLOPs
总计: 2048 × 8 = 16K FLOPs ← 非常小
访存: 读 2×2048×2 + 写 2048×2 = 12KB ← 但分散在 HBM 中
这个算子的计算量极小,但不能和前后 GEMM 融合的话,每次都是独立的 kernel launch。最佳实践:把 SiTU 融合到 grouped GEMM 的 epilogue 中。
9.2.5 Short Conv on K(访存密集)
输入: k_cache (4, d_k) = (4, 4096)
k_raw (d_k,) = (4096,)
权重: conv_weight (4, d_k) = (4, 4096)
计算: k = Σ w_i · k_cache[i] + w_3 · k_raw
= 4 个向量的加权和 → 4 × 4096 次乘法 + 3 × 4096 次加法
= 28K FLOPs
关键点:这个算子虽然 FLOPs 极低,但必须高效实现深度可分离 1D 卷积。芯片不需要专门的卷积加速器——一个逐元素乘加单元就够了。但不能把它拆成 4 个独立的 vector-scalar 乘 + 3 次加法(那样 launch 开销会吃掉收益)。
9.2.6 Fine-Grained Gate 生成(计算密集,但形状特殊)
输入: x (d_model,) = (7168,)
权重: W_alpha (d_gate, d_model) = (64, 7168)
W_alpha_down (d_k, d_gate) = (4096, 64)
计算:
Step 1: h = W_alpha @ x → [64] = [64, 7168] × [7168]
Step 2: h = SiLU(h) → [64] (逐元素)
Step 3: logits = W_alpha_down @ h → [4096] = [4096, 64] × [64]
Step 4: alpha = sigmoid(logits) → [4096] (逐元素)
shape 特点:这是一个三层瓶颈网络(7168→64→4096),虽然也是 GEMM,但 M=1 时是 GEMV。建议和 QKV 投影融合——它们共享输入 x,可以减少一次 HBM 读。
9.3 自研芯片算子的优先级矩阵
综合出现频率、算力占比和实现难度:
| 优先级 | 算子 | 原因 |
|---|---|---|
| P0(必须有) | MatMul (GEMM/GEMV) | 占 65% 算力,是所有算子的基础。芯片不支持高效 GEMM 就谈不上部署 LLM |
| P0(必须有) | Grouped GEMM | MoE 的核心算子。如果没有,896 expert 的 FFN 需要串行 896 次 GEMM,不可接受 |
| P0(必须有) | Element-wise 向量操作 | RMSNorm、SiTU、SiLU、残差加、逐元素乘——300+ 次/层,融合后可不计成本,分开则 kernel launch 爆炸 |
| P1(关键) | Softmax | MLA 层(24 层)需要。实现难度不高,但数值稳定的 online softmax 需要小心 |
| P1(关键) | KDA DPLR 状态更新 | 独特算子,占 8% 算力但访存模式特殊。建议作为专用 kernel 实现 |
| P1(关键) | All-to-All 通信 | EP 模式下卡间通信的瓶颈。需要片间互联(NVLink 替代品)支持 |
| P2(重要) | Short Conv 1D | 实现简单,但需要深度可分离卷积语义 |
| P2(重要) | Top-K 选择 | MoE Router 需要,实现简单 |
| P2(重要) | Softmax + FlashAttention | MLA 层需要 attention 实现。可以用标准 softmax + GEMM 替代,但效率不如 fused |
9.4 算子的融合策略
对于自研芯片,算子融合可能是最大的优化空间——减少 kernel launch、减少 HBM 中间结果读写。
高价值融合策略:
融合组 1(QKV 投影 + Fine-Grained Gate):
输入 x → [W_q, W_k, W_v, W_alpha] 4 个 GEMM 合并为 1 个 batched GEMM
节省: 3 次 kernel launch + 3 次输入 x 的 HBM 读取
融合组 2(SiTU-GLU + Down Projection):
SiTU-GLU 的输出直接作为 Down GEMM 的输入,不写回 HBM
节省: 1 次写 + 1 次读 = 2 × 2048 × 2 bytes = 8KB 每 expert
但 16 experts × 93 层 = 11.9MB 总节省
融合组 3(KDA 三步融合):
对角衰减 + Rank-1 纠偏 + KV 写入合并为 1 个 kernel
节省: S 矩阵的 2 次中间写回(S_decayed 和 S_corrected 不用落 HBM)
收益: 每层 ~64MB 的 HBM 读写节省(FP16 下)
融合组 4(Residual Add + RMSNorm):
q = rmsnorm(x + attn_output) — 两个逐元素操作融合
节省: 1 次 kernel launch + x 的 1 次读
Kimi K3 总 kernel launch 次数估算(无融合,93 层):
每层: 3(QKV) + 3(KDA三步) + 1(SiTU) + 3(MoE GEMM) + 2(RMSNorm) + 1(Short Conv) + 1(FineGate) + 1(Residual) = ~15 kernels
93 层: ~1395 kernels
+ Embedding + LM Head + Router + LayerNorm: ~1400+
如果 平均每个 kernel launch 开销 10μs → 14ms 纯调度开销
如果能融合到 ~500 kernels → ~5ms
节省的 ~9ms 在 decode 场景中可能就是 15-20% 的加速
9.5 显存带宽需求评估
对于自研芯片,显存带宽决定了 decode 速度的上限。
Kimi K3 的权重加载量估算(MXFP4,1 byte/param):
每层参数加载:
QKV 投影: 3 × 7168 × 4096 × 1 byte = 88 MB
O 投影: 4096 × 7168 × 1 byte = 29 MB
KDA 状态 S: 4096 × 4096 × 2 bytes = 32 MB(FP16 状态,必须高精度)
MoE FFN: 16 × 3 × 7168 × 2048 × 1 byte = 704 MB
───────────────────────────────────────────────────
每 KDA 层: 853 MB
每 MLA 层: 相当(含额外的 attention 参数)
93 层总计: ~78 GB
LM Head: 7168 × 160000 × 1 byte = 1.1 GB
───────────────────────────────────────────────────
总权重加载: ~80 GB / token(MXFP4)
+ KDA 状态 S 的读写: 69 × 64 MB = 4.4 GB / token
| 芯片 | 带宽 | 80GB 加载时间 | 理论 tok/s |
|---|---|---|---|
| H100 80GB | 3.35 TB/s | 23.9 ms | ~42 tok/s |
| H200 141GB | 4.8 TB/s | 16.7 ms | ~60 tok/s |
| B200 192GB | 8.0 TB/s | 10.0 ms | ~100 tok/s |
| Ascend 950DT | 4.0 TB/s | 20.0 ms | ~50 tok/s |
| 自研芯片(目标带宽) | ≥4 TB/s | ≤20 ms | ≥50 tok/s |
结论:对于自研芯片,要达到可用的 decode 速度(≥50 tok/s),HBM 带宽至少需要 4 TB/s。如果带宽只有 1-2 TB/s(如早期国产芯片),即使算子全部实现,decode 速度也会受限在 15-25 tok/s 以下。
9.6 算子开发路线图建议
Phase 1(基础能力,~2 个月)
├─ MatMul GEMM (FP16/BF16, M≥1 通用) ← 最重要,没有就不用谈 LLM
├─ Element-wise 全家桶(add/mul/silu/tanh/sigmoid/rmsnorm)
├─ Softmax(online 版本,数值稳定)
└─ 验证方法: 跑通单层 KDA 的前向
Phase 2(MoE 核心,~1.5 个月)
├─ Grouped GEMM(优先级超过标准 GEMM 的大 batch 优化)
├─ Top-K 选择(router 需要用)
├─ SiTU-GLU 融合(和 GEMM epilogue 融合)
└─ 验证方法: 跑通单层 MoE FFN(含 router)
Phase 3(KDA 专用算子,~1 个月)
├─ DPLR 状态更新(三步融合 kernel)
├─ Short Conv 1D(深度可分离因果卷积)
├─ Fine-Grained Gate(瓶颈网络 + sigmoid)
└─ 验证方法: 跑通完整 KDA 层(对比 PyTorch 输出误差 < 1e-3)
Phase 4(全模型 & 通信,~1.5 个月)
├─ Gated MLA(FlashAttention 风格)
├─ Attention Residuals(跨层检索)
├─ All-to-All 通信 + MoonEP 兼容
└─ 验证方法: 完整模型前向 + 和官方权重对齐
Phase 5(极致优化)
├─ FlashKDA chunkwise 并行(prefill 加速 2x)
├─ 算子融合(1400→500 kernels)
├─ MXFP4 数据格式支持
└─ 目标: 达到理论带宽的 80%+ 利用率
附录:进一步阅读
- 博客 #26:Kimi K3 技术分析——2.8T 参数的 MoE 猛兽
- KDA 论文: Kimi Linear: An Expressive, Efficient Attention Architecture (arXiv:2510.26692)
- Kimi K3 技术报告: 47页,随开源权重同步发布
- 官方 KDA 实现: MoonshotAI/Kimi-Linear
- 教育版 KDA 实现: hwilner/kimi-delta-attention
- 参考 KDA 实现: hkevin01/kimi-linear(含 130+ 测试)
- FlashKDA: MoonshotAI/FlashKDA(CUTLASS 实现)
- MoonEP: MoonshotAI/MoonEP(动态冗余 expert 并行)
- K3 小型复现: cneuralnetwork/smol-kimi-k3(49M 参数,8GB GPU 可训)
- smol-kimi-k3 博文: Training a 49M Kimi K3-inspired Model
更多推荐



所有评论(0)