副标题: 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 用了两个关键优化:

  1. UT 变换减少非 matmul FLOPs:传统的 DPLR chunkwise 实现需要矩阵求逆(O(C³)),UT 变换将其简化为前向替换(O(C²)),且把更多计算转化为 Tensor Core 友好的 matmul。

  2. 绑定两个 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 的控制方法:

  1. Block 级粒度:检索只在 Block 边界进行,不是每步都做
  2. 特征压缩:存储的不是完整的 k/v,而是投影到更低维度的"残差特征"
  3. 选择性检索:通过一个可学习的门控决定"当前层需要从前面拿多少信息"

五、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 深度可分离因果卷积 给线性注意力提供局部感知能力

一些值得关注的工程视角

  1. KDA 的 DPLR 三步法是一个很好的"数学推导 → 工程实现"案例——同一个公式,朴素实现是 O(d_k²·d_v),拆解后是 O(d_k·d_v),差了 d_k/2 倍。

  2. Chunkwise 算法的核心挑战不是并行计算本身,而是最小化非 matmul FLOPs——Tensor Core 只擅长 matmul,任何非 matmul 操作都是瓶颈。FlashKDA 的 UT 变换就是为了把尽可能多的计算转成 matmul。

  3. SiTU-GLU 的 tanh 限幅思路值得在其他高稀疏度场景借鉴——当计算路径上的信号非常稀疏时,给激活值加一个有界约束可以大幅提高稳定性。

  4. **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%+ 利用率

附录:进一步阅读

更多推荐