AI infra(1)

前置说明

你这份文档是 SGLang Diffusion 融合算子 fused_inplace_qknorm_rope 深度技术分析,面向大模型推理内核开发;现在我把全文拆解:

  1. 砍掉复杂长句,逐块加通俗注释
  2. 拆成【基础概念 → 算子原理 → CUDA Kernel 代码解读 → 百度昆仑芯移植 → 动手复写教程】
  3. 全部术语大白话,AI 初学者友好;先把核心名词一次性解释清楚。

背景:这个算子是 DiT 图像生成模型(FLUX、Qwen-Image 这类文生图)Attention 前面的预处理融合算子。 融合算子:把多个连续的 GPU 计算步骤合并成 1 个 GPU 核函数 (kernel),减少显存读写,提速。


词汇预习(先看懂这些词,后面就轻松很多)

表格

名词 通俗解释
Kernel / CUDA kernel 在 GPU 上并行执行的一段 C++ 代码,CPU 调用它,GPU 大量线程同时跑
In-place(原地计算) 计算结果直接覆盖原来输入内存,不额外开辟新显存存中间结果,省显存
RMSNorm 归一化算法,把向量缩放,防止数值爆炸,稳定模型推理
RoPE 旋转位置编码,给向量注入位置信息,让 Attention 知道 token 的先后顺序
DiT Diffusion Transformer,现在主流文生图模型架构(FLUX/Qwen-Image)
Q / K / V Attention 机制的三个向量:Query 查询、Key 键、Value 值
GQA / MHA MHA:Q 和 K 头数量完全一样;GQA:Q 头多、K 头少,节省显存
warp GPU 最小调度单位,1 个 warp 固定 32 个线程 (lane),warp 内可以用 shuffle 指令交换寄存器数据
lane warp 内部单个线程,编号 0~31
JIT 编译 运行时动态编译 C++ 代码,不是程序启动前编译;可以根据参数生成定制 kernel
template 模板 (C++) 编译期就固定参数(比如 head_dim=128),编译出来的代码没有 if 分支,运行更快
融合收益 减少 GPU 显存读写。GPU 瓶颈大多是显存带宽,不是计算速度。少读写 = 变快
昆仑芯 P800 百度国产 XPU 芯片;xSGL 是适配昆仑芯的 SGLang 分支
cuda_like 平台 昆仑 XPU 做了兼容层,可以直接跑大部分 CUDA 代码,不用大规模改写(不是完全兼容)
fallback 降级 如果融合算子不能用,自动切回分步慢版本,保证程序不崩溃
tensor 多维数组,深度学习的数据载体(pytorch 里面的数组)
stride 张量内存步长:在内存里,相邻维度元素隔多少字节存放

第一部分:算子源码深度解析 ------ 这个算子干什么

1. 一句话定位(注释版)

在 diffusion 模型(DiT 架构)的每个 attention 层里,Q、K 在送入 attention 计算之前,要先后经过两步数学变换 ------QK RMSNorm(归一化)和 RoPE(旋转位置编码)。这个算子把这两步融合成一个 CUDA kernel,原地(in-place)更新 q/k 张量。

✅ 人话: 文生图模型,每一层 Attention,拿到 Q、K 向量,正常要分开跑两个 GPU 函数:先归一化、再加位置编码。 融合算子:合并成 1 次 GPU 调用,计算结果直接覆盖原来 Q、K,不产生中间显存副本。

复制代码
hidden_states
    │  Linear 投影  // 线性层,把输入向量映射成Q K V
    ▼
q, k, v  ──►  ① q_norm(q), k_norm(k)      ← QK RMSNorm:对Q、K每个头单独归一化
          ──►  ② rope(q), rope(k)          ← RoPE 旋转位置编码
          ──►  ③ attention(q, k, v)        ← 注意力计算
    ▼
下一层

为什么重要:DiT 模型(FLUX、Qwen-Image)生成图片,要循环几十层 Attention,循环上千个去噪步骤。这一段代码被反复执行,属于热点代码,优化这里收益巨大。

2. 数学原理

2.1 QK RMSNorm

设一个 head 的向量是 x ∈ R^head_dim,可学习权重 w ∈ R^head_dim

复制代码
rms(x)   = sqrt( (1/head_dim) · Σᵢ xᵢ² + eps )
x_out[i] = x[i] / rms(x) · w[i]
  • eps:极小值(一般 1e-6),防止分母等于 0,除零报错
  • w:模型 checkpoint 里面保存的可学习参数,逐通道缩放

✅ 人话:

  1. 取出单个注意力头的向量 x;
  2. 每个元素求平方,全部相加求和;
  3. 除以向量长度 head_dim,开平方根,得到 RMS;
  4. 向量每个元素除以 RMS,乘以权重 w;
  5. 作用:把向量数值范围稳定住,防止 Attention 打分数值漂移,现代 DiT 标配。

2.2 RoPE 旋转位置编码

核心思想:把向量每两个数字当成二维平面上的一个坐标点,根据 token 位置旋转这个点。旋转之后向量内积天然自带相对位置信息。

预计算 cos_sin_cache,提前算好三角函数,推理时直接查表,不用实时计算 cos/sin 节省开销。

复制代码
cache[pos] = [ cos(pos·θ₀), cos(pos·θ₁), ..., cos(pos·θ_{r/2-1}) ,   ← 前一半全部cos值
               sin(pos·θ₀), sin(pos·θ₁), ..., sin(pos·θ_{r/2-1}) ]   ← 后一半全部sin值
频率 θⱼ = base^(-2j/rope_dim)

两种配对规则(非常关键,kernel 两套分支)

  1. interleaved(GPT-J 风格)相邻两个一组 (2j,2j+1)

    out[2j] = x[2j]·cosⱼ − x[2j+1]·sinⱼ
    out[2j+1] = x[2j+1]·cosⱼ + x[2j]·sinⱼ

👉 一组两个数字在同一个线程寄存器里面,计算简单。FLUX/Z-Image 使用这个模式。

  1. NeoX(LLaMA 风格):前半段和后半段配对 (d, d+half) half=rope_dim/2。

    out[d] = x[d]·cos_d − x[d+half]·sin_d (d < half)
    out[d+half] = x[d+half]·cos_d + x[d]·sin_d

👉 麻烦点:一对数字不在同一个线程,分散在不同 lane,需要 warp_shuffle 跨线程拿数据。LLaMA 文本模型常用。

部分 RoPE:不是 head_dim 全部维度都旋转,只旋转前rope_dim维,剩下维度原样保留。例 head_dim=128,rope_dim=64:只旋转前 64 维。约束:rope_dim ≤ head_dim

2.3 融合算子完整计算流程(单头)

输入:q 向量、k 向量,权重 w_q/w_k,cos/sin 表,每个 token 的位置 pos

复制代码
1. 读取这个head全部元素,加载进GPU寄存器,转fp32高精度
2. sum_sq = Σ x_i²   //所有元素平方求和
3. scale = rsqrt(sum_sq / head_dim + eps) //RMS倒数,rsqrt是GPU快速求平方根倒数指令
4. x_i = x_i * scale * w_i //RMS归一化完成
5. 前rope_dim维执行RoPE旋转;剩下维度不变
6. 结果写回原来显存地址(in-place原地覆盖)

2.4 融合带来性能收益

GPU 最大瓶颈:显存读写带宽,不是计算。

表格

分步(分开 norm+rope 两个 kernel) 融合单 kernel
GPU 启动次数 2 次 kernel launch 1 次
读取 Q 显存 2 次:读一次给 norm,norm 写完,rope 再读一遍 1 次读
写入 Q 显存 2 次:norm 写中间结果,rope 再写最终结果 1 次写
中间数据 必须写到显存 数据全程保存在寄存器,不落地显存

一句话:显存往返减半,推理速度接近翻倍。这类算子属于带宽受限算子,计算量很小,大量时间浪费在读写显存。

##3. 输入输出契约(Tensor 参数校验表)

契约:调用这个算子,张量必须满足的形状、数据类型、内存排布要求;TensorMatcher 用来自动校验,不满足直接报错。

表格

参数 形状 dtype stride 要求 说明
q [num_tokens, num_qo_heads, head_dim] fp16/bf16 最内层维度 stride 必须等于 1(连续内存);头维度 stride 和 k 保持兼容 由模型 4 维张量[B,S,H,D]reshape 变形得到,B 批次,S 序列长度
k [num_tokens, num_kv_heads, head_dim] 同 q 同上 支持 GQA:Q 头数量和 K 头数量可以不一样
q_weight/k_weight [head_dim] 和 q/k 相同 无 RMSNorm 权重
cos_sin_cache [任意长度, rope_dim] fp32 无 cos 在前半,sin 后半拼接
positions [num_tokens] int32 / int64 无 每个 token 对应的位置编号
返回值 None --- --- 原地修改 q、k,不返回新张量

模板参数(编译期固定,提前实例化):head_dim /rope_dim/is_neox /dtype 运行时参数(每次调用可变):token 数量、head 数量、stride、eps ✅ 设计目的:JIT 缓存的 key 只使用编译期参数,避免每次微小变化都重新编译 kernel。

##4 CUDA Kernel 逐段源码解析

文件:qknorm_rope.cuh,C++ GPU 内核代码,在 GPU 设备上执行

###4.1 参数结构体 QKNormRopeParams

复制代码
struct QKNormRopeParams {
void* q_ptr;
void* k_ptr;                   // k指针做了预偏移,后面单独解释
const void* q_weight_ptr, *k_weight_ptr, *cos_sin_cache_ptr, *positions;
int64_t q_stride_bytes, k_stride_bytes, head_stride_bytes;
uint32_t num_qo_heads, num_kv_heads, num_tokens;
float eps;
};

把所有运行时参数打包放进一个结构体,用__grid_constant__放到 GPU 常量内存。 好处:相比十几个零散入参,常量内存读取更快。

复制代码
constexpr uint32_t kThreadsPerBlock = 256;   // 一个block=256线程 =8个warp ×32lane

GPU 线程层级:Grid(网格)→Block(线程块)→Thread(线程) 这里一个 block 固定 256 线程,拆成 8 个 warp,每个 warp32 线程。

###4.2 线程映射逻辑(重点:warp-per-head)

复制代码
const uint32_t lane_id = threadIdx.x % 32;    // warp内0~31号线程
const uint32_t warp_id = threadIdx.x / 32;    // block内部warp编号0~7
const uint32_t start_worker_id = blockIdx.x * kWarpsPerBlock + warp_id;
const uint32_t num_works = (num_qo_heads + num_kv_heads) * num_tokens;
for (uint32_t idx = start_worker_id; idx < num_works; idx += num_workers)  // grid-stride循环

✅ 设计思路:

  1. 1 个 warp 负责处理 1 个 head(单个注意力头向量)
  2. 总任务量 = 全部 Q 头 + 全部 K 头 × token 数量
  3. idx:任务编号;head_id < num_qo_heads → 当前 warp 处理 Q;否则处理 K。同一个 kernel 同时处理 Q 和 K
  4. grid-stride 循环:任务数量远超 GPUblock 数量时,block 循环反复领取任务,避免启动过多 block,提高 SM 占用率

分配规则:

head_dim 必须被 32 整除(64/128/256),每个 lane 分到 head_dim/32 个元素 例 head_dim=128,128/32=4:每个 lane 负责 4 个数字

👉 为什么 warp-per-head,而不是 block-per-head? 一个 warp32 线程刚好处理一个 head;warp 内 shuffle 归约,不需要共享内存 shared memory,不需要线程同步__syncthreads。 一个 block8 个 warp,并行处理 8 个 head,互相独立,等待少。

###4.3 RMSNorm 主体代码

复制代码
using Packed  = packed_t<DType>;                  // bf16x2 / fp16x2 打包类型,一次读取两个元素
using Storage = AlignedVector<Packed, kVecSize>;  // 128bit向量,16字节对齐,总线一次性读取
auto input_vec  = load_as<Storage>(input, lane_id);       // lane读取对齐向量
const auto weight_vec = load_as<Storage>(weight_ptr, lane_id);
float elems[kElemsPerThread];
float sum_of_squares = 0.0f;
#pragma unroll  //编译器循环展开,消除循环开销
for (uint32_t j = 0; j < kVecSize; ++j) {
const auto [x0, x1] = cast<fp32x2_t>(input_vec[j]);   // bf16/fp16转fp32高精度
elems[2*j] = x0;  elems[2*j+1] = x1;
  sum_of_squares += x0*x0 + x1*x1;                       //平方累加,fp32防止精度丢失
}
sum_of_squares = warp::reduce_sum(sum_of_squares);        //warp内蝶形归约,32lane求和
const float norm_factor = math::rsqrt(sum_of_squares / kHeadDim + eps);
#pragma unroll
for (uint32_t j = 0; j < kVecSize; ++j) {
const auto [w0, w1] = cast<fp32x2_t>(weight_vec[j]);
elems[2*j]   *= norm_factor * w0;
elems[2*j+1] *= norm_factor * w1;                      //RMSNorm计算完成,结果保存在寄存器elems数组
}

四个工程要点注释:

  1. 128bit 对齐向量读取:一次读取 16 字节,充分利用 GPU 内存总线带宽,比逐个读取快很多
  2. 升 fp32 累加平方:bf16 精度很低,大量数字累加误差会越来越大;加载之后立刻转 fp32 计算
  3. warp::reduce_sum:基于__shfl_xor_sync蝶形求和,32lane 把各自的 sum 汇总成总和。不需要 shared 内存
  4. 归一化结果保存在寄存器数组 elems,不写显存,直接进入 RoPE 计算------ 融合算子提速核心!

###4.4 RoPE NeoX 分支(最难的部分:跨 lane 交换寄存器数据)

NeoX 模式下,配对的两个元素不在同一个 lane。必须用__shfl_xor_sync指令,warp 内线程互相交换寄存器的值。

复制代码
constexpr uint32_t kRotaryLanes     = kRopeDim / kElemsPerThread;
constexpr uint32_t kHalfRotaryLanes = kRotaryLanes / 2;
constexpr uint32_t kActiveMask      = active_mask<kRotaryLanes>();
if (lane_id < kRotaryLanes) {
const auto pos = ...;
const auto cos_ptr = cache + pos * rope_dim;
const auto sin_ptr = cos_ptr + rope_dim / 2;
#pragma unroll
for (uint32_t i = 0; i < kElemsPerThread; ++i) {
float swapped = __shfl_xor_sync(kActiveMask, elems[i], kHalfRotaryLanes);  //核心:和搭档lane交换数据
if (lane_id < kHalfRotaryLanes) swapped = -swapped;
int dim_idx = static_cast<int>(lane_id * kElemsPerThread + i);
    dim_idx = (dim_idx * 2) % kRopeDim;
const int half_idx = dim_idx / 2;
elems[i] = elems[i] * cos[half_idx] + swapped * sin[half_idx];
  }
}
三步理解魔法 shfl_xor
  1. 配对 lane 编号 = lane_id ^ kHalfRotaryLanes(异或)。前一半 lane 和后一半 lane 两两配对,互相拿到对方寄存器的值
  2. 符号处理:前半 lane 的公式需要减去 x d+half,所以 swapped 取负;后半 lane 不需要取负
  3. (d*2) % rope_dim /2:一条公式统一 cos/sin 索引,不用 if 分支,减少运行开销

interleaved 分支简单:一对元素在同一个 lane 内部相邻位置,直接计算,不需要 shuffle 交换。

复制代码
for (uint32_t i = 0; i < kElemsPerThread; i += 2) {
const int half_idx = (lane_id * kElemsPerThread + i) / 2;
const float x = elems[i], y = elems[i+1];
elems[i]   = x * cos[half_idx] - y * sin[half_idx];
elems[i+1] = y * cos[half_idx] + x * sin[half_idx];
}

计算完成后,fp32 转回 fp16/bf16,向量对齐 store,原地写回显存。

###4.5 k_ptr 负偏移技巧(host CPU 侧 trick)

复制代码
// host CPU侧预处理
const int64_t k_offset = num_qo_heads * head_stride_bytes;
.k_ptr = pointer::offset(k.data_ptr(), -k_offset),
// kernel内部寻址
input = pointer::offset(k_ptr, token_id*k_stride_bytes, head_id*head_stride_bytes);

问题:kernel 循环统一遍历 0 ~ (Q 头数 + K 头数) 0~Q 头编号:处理 Q Q 头~Q+K 头编号:处理 K 如果不做偏移,Q 和 K 寻址公式需要两套 if 分支判断,代码复杂。

✅ 技巧:CPU 端预先把 K 指针向前偏移一段(负偏移,指向 K 内存起始地址前面)。 kernel 里面 Q、K 可以复用同一套寻址公式,kernel 内部消除分支,简化代码。

###4.6 static_assert 编译期护栏 模板参数编译期静态断言,提前拦截非法参数组合,编译阶段直接报错,不要等到 GPU 运行才崩溃。

表格

static_assert 含义
kHeadDim % kWarpThreads == 0 head_dim 必须被 32 整除,保证 32lane 均匀分配元素
kRopeDim >0 && kRopeDim <=kHeadDim 部分 RoPE 约束,旋转维度不能超过 head_dim
kElemsPerThread %2 ==0 打包向量成对,适配 fp16x2/bf16x2
kRopeDim % kElemsPerThread ==0 参与旋转的 lane 必须完整拥有元素,不能拆分
NeoX:kRotaryLanes 是 2 的幂 shuffle 异或配对逻辑成立的前提

###4.7 host 侧 run () 函数(CPU 端入口,调用 GPU kernel)

复制代码
static void run(q, k, q_weight, k_weight, cos_sin_cache, positions, eps) {
  //① TensorMatcher张量校验,检查形状、stride、设备、dtype,不满足抛异常
auto N/Q/K/D/R/Dq/Dk/Dd = SymbolicSize{...};
D.set_value(kHeadDim);  R.set_value(kRopeDim);
TensorMatcher({N, Q, D}).with_strides({Dq, Dd, 1}).with_dtype<DType>()
      .with_device(device).verify(q);
TensorMatcher({N, K, D}).with_strides({Dk, Dd, 1}).verify(k);
  ...
//② positions支持int32 / int64,两套模板实例
const auto selected_kernel = is_int32 ? kernel<int32_t> : kernel<int64_t>;
//③ 计算合适block数量:不超过GPU SM最大并发,防止过度启动block
static const uint32_t kOccupancyTable[2] = { get_blocks_per_sm(kernel<int32_t>, 256), ... };
const auto num_blocks = std::min(max_blocks, needed_blocks);
//④ LaunchKernel启动GPU kernel,RAII封装,自动检查cuda error
LaunchKernel(num_blocks, kThreadsPerBlock, device.unwrap())
      .enable_pdl(kUsePDL)(selected_kernel, params);
}
  • PDL:Programmatic Dependent Launch。SM90(H100/H200)以上新特性:相邻 kernel 可以重叠执行,消除 kernel launch 间隙,隐藏启动延迟。老显卡 / 昆仑芯不生效,属于性能优化,不影响正确性。
  • occupancy:一个 SM 最多可以同时驻留多少个 block;提前查询,控制 block 数量,最大化 GPU 利用率。
  • LaunchKernel:封装好的启动工具,launch 完成自动检查 GPU 报错,方便调试。

##5 Python 层代码解析(Pytorch 侧包装代码)

Python 层:模型代码调用的 API,底层调用 C++ JIT 编译出来的 kernel。

###5.1 JIT 模块缓存函数

复制代码
@cache_once
def _jit_qknorm_rope_module(head_dim, rope_dim, is_neox, dtype) -> Module:
args = make_cpp_args(head_dim, rope_dim, is_neox, is_arch_support_pdl(), dtype)
return load_jit(
"qknorm_rope", *args,
cuda_files=["diffusion/qknorm_rope.cuh"],
cuda_wrappers=[("qknorm_rope", f"QKNormRopeKernel<{args}>::run")],
)

逻辑:

  1. cache_once 自定义装饰器,缓存编译后的 so;不用 lru_cache,因为 lru_cache 和 torch.compile 冲突。
  2. make_cpp_args 收集编译期模板参数,组成唯一 key;只有模板参数变化才会重新编译。

❗重点:token 数量、eps 这类运行时参数,绝对不能放进缓存 key,否则每次推理尺寸变化,重复编译,内存爆炸。

  1. load_jit 调用 nvcc 编译 cuh 代码,生成动态库 so,加载到 python。
  2. 模板不同,生成不同版本 kernel,缓存起来,第二次调用直接复用,不用编译。

###5.2 can_use_fused_inplace_qknorm_rope 能力检查门控

复制代码
if head_dim not in (64, 128, 256): return False
if rope_dim <= 0 or rope_dim > head_dim: return False
if rope_dim % (head_dim // 32) != 0: return False
if is_neox:
    rotary_lanes = rope_dim // (head_dim // 32)
    if rotary_lanes < 2 or rotary_lanes & (rotary_lanes-1): return False
try:
    _jit_qknorm_rope_module(...); return True
except Exception: return False

功能: Python 层提前检查参数合法性,和 C++ static_assert 一一对应,双层保护 。 最后尝试编译一次:如果环境缺少 nvcc、硬件不支持,返回 False,自动降级到分步实现,不会直接崩溃。 @torch.compiler.assume_constant_result:torch.compile 把这个判断当成常量,图编译时直接折叠。

###5.3 算子主入口函数

复制代码
@register_custom_op(mutates_args=["q", "k"])
def fused_inplace_qknorm_rope(q, k, q_weight, k_weight, cos_sin_cache,
positions, *, is_neox, eps=1e-6,
head_dim=0, rope_dim=0) -> None:
    head_dim = head_dim or q.size(-1)
    rope_dim = rope_dim or cos_sin_cache.size(-1)
    module = _jit_qknorm_rope_module(head_dim, rope_dim, is_neox, q.dtype)
    module.qknorm_rope(q, k, q_weight, k_weight, cos_sin_cache, positions, eps)

@register_custom_op(mutates_args=["q", "k"]) 👉 非常重要:向 PyTorch 声明,这个函数原地修改 q、k 张量。如果不写,torch.compile 计算图会误以为 q/k 没有修改,复用旧张量,计算结果出错。 函数返回 None,结果原地写进 q/k。

##6 四层降级门控(模型调用层) 模型不会直接调用融合算子,会先走条件判断,满足所有条件才走快路径;任意条件不满足,自动 fallback 分步版本。

复制代码
fused_enabled = os.getenv("SGLANG_ENABLE_FUSED_QKNORM_ROPE", "1")
if (fused_enabled
and _is_cuda
and allow_inplace
and (q_eps == k_eps)
and q.dtype in (fp16, bf16)
and q_norm.weight.dtype == q.dtype
and k_norm.weight.dtype == k.dtype
and q.is_contiguous() and k.is_contiguous()
and can_use_fused_inplace_qknorm_rope(...)):
    fused_inplace_qknorm_rope(q.reshape(-1, H, head_dim), ...)
return q, k
# fallback慢路径
q, k = apply_qk_norm(...)
return apply_flashinfer_rope_qk_inplace(...)

四层检查顺序:

  1. 环境变量开关:可以手动关闭融合算子
  2. 平台、数据类型、张量连续性、eps 相等性检查
  3. can_use_fused_inplace_qknorm_rope 能力检查(试编译)
  4. 全部通过才走融合 kernel;否则分步执行:单独 RMSNorm + 单独 RoPE

注意:krea2 模型可以绕过这套封装,直接调用算子,需要提前预处理 cos_sin cache 和权重。


第二部分:百度昆仑芯移植改动分析

##7 移植总览 上游 U:原版 SGLang;百度 B:适配昆仑 P800 的 xSGL 分支

核心结论:CUDA kernel 代码 qknorm_rope.cuh 字节完全一样,一行没改。Python wrapper 只改动 import 路径。 改动只发生在:目录重构、runtime 门控代码快照版本、CI 注册、JIT 底层头文件版本。

表格

移植内容 改动程度 说明
kernel cuh 字节一致 没有针对昆仑芯修改 GPU 代码
python wrapper 仅 import 一行变化 只是文件目录移动,逻辑不变
单测 & benchmark import 路径修改 测试逻辑完全复用
apply_qk_norm_rope 上层门控 快照版本落后 唯一有业务影响的改动:GQA 条件判断限制
JIT 底层头文件 旧版本快照 只删掉 AMD ROCm 相关代码,本算子不受影响

##8 逐项改动解析 ###8.1 目录重组(纯搬家,不影响功能) 原版上游目录:python/sglang/kernels/ 百度分支:python/sglang/jit_kernel/ 只是文件夹改名,导入路径 from xxx 改成 from sglang.jit_kernel.utils,算子逻辑完全不变。

###8.2 runtime 门控快照差异【最重要的缺陷】 上游原版门控:允许 Q、K 头数量不一致(GQA),只要 batch 和 seq 相等。 百度旧版本门控:强制要求 q.shape == k.shape,Q 头数量必须等于 K 头数量。

👉后果: GQA 模型(Q 头≠K 头)无法进入融合算子,直接降级到慢路径 。 但是!kernel 底层代码本身原生支持 GQA,只是上层 python 判断条件卡住了。 好在百度仓库内用到这个算子的模型(Z-Image/FLUX/Qwen-Image)全部是 MHA(Q 头 K 头数量一样),现有模型不受影响,新增 GQA 模型才会踩坑。

上游还额外增加的保护,百度快照没有:

  1. torch.compile 编译期保护,防止图捕获阶段误入融合算子
  2. cos_sin_cache 形状、设备校验
  3. positions 自动转换设备与 dtype(百度要求调用方自己保证)

###8.3 CI 测试注册 API 适配 上游 CI 参数:register_cuda_ci(est_time=44, stage="xxx", runner_config="xxx") 百度 CI:合并成 suite 单参数register_cuda_ci(est_time=44, suite="xxx") 只是 CI 流水线注册语法差异,算子功能完全无关。CI 系统靠 AST 静态扫描收集测试用例,est_time 必须写字面量数字,不能填变量。

###8.4 JIT 底层头文件版本差异 warp.cuh/runtime.cuh/math.cuh 上游新增 AMD ROCm 分支代码。百度版本删掉 ROCm 兼容代码,只保留 CUDA 逻辑。 本算子只用到 warp reduce_sum、rsqrt、向量加载,ROCm 代码完全不会被触发。对 qknorm_rope 无任何影响,所以 kernel 源码可以原封不动搬运。

###8.5 重点:kernel 源码零改动 diff 上游qknorm_rope.cuh 百度qknorm_rope.cuh →无差异。 👉不是百度重写适配昆仑,是直接搬原版 CUDA 代码。

##9 昆仑 P800 上,算子能力盘点

✅ 已具备能力

  1. 数学逻辑完全等价:RMSNorm+RoPE 融合、interleaved/NeoX、部分 RoPE、GQA 内核支持、fp16/bf16、int32/int64 位置、原地计算
  2. 全套单元测试 + 性能 benchmark
  3. Zimage / FLUX / Qwen-Image 在 MHA 场景下,可以成功走到融合快路径
  4. 四层降级兜底,算子不可用时自动切分步,不会崩溃

⚠️ 缺口

  1. runtime 门控不支持 GQA 模型走融合路径
  2. 缺少 torch.compile 编译期保护
  3. PDL 指令:PDL 是英伟达 SM90 专属特性;昆仑芯编译时判定不支持 PDL,编译出来不带 PDL 逻辑,只是少一点启动重叠,不影响正确性。
  4. 没有 Krea2 模型接入代码(属于模型层缺失,不是算子本身)

重要区分两份仓库算子:

  • jit_kernel/目录下算子:cuda_like 兼容路线,直接复用原版 CUDA 代码,靠昆仑 xPU 兼容层 xpytorch+xmlir 转换执行
  • sgl-kernel/csrc/klx/:专门为昆仑芯手写的算子,使用昆仑硬件特有原语。 qknorm_rope 属于前者:代码不改,依赖平台兼容层。 ⚠️ 关键提醒:源码存在 ≠ 在 P800 上一定跑通! 兼容性取决于:xmlir 能不能正确翻译 warp shuffle、向量加载等 CUDA 原语,需要上板实测验证正确性与性能。

#第三部分:从零复写这个算子的完整实操指南 仓库文档规定:轻量 kernel(无 CUTLASS)选择JIT 方案。 目录放置位置:

复制代码
python/sglang/jit_kernel/csrc/diffusion/qknorm_rope.cuh     # CUDA kernel源码
python/sglang/jit_kernel/diffusion/qknorm_rope.py           # Python wrapper
python/sglang/jit_kernel/tests/diffusion/test_qknorm_rope.py #单元测试
python/sglang/jit_kernel/benchmark/diffusion/bench_qknorm_rope.py #性能压测

Step1 编写 CUDA kernel

推荐开发顺序:

  1. 先写数据通路:加载向量 → fp32 转换 → 平方求和归约 → RMSNorm → RoPE 旋转 → 写回显存
  2. 设计线程映射:warp-per-head
  3. 增加 static_assert 编译期约束护栏
  4. host 端 run 函数:TensorMatcher 校验、k_ptr 负偏移、block 数量计算、LaunchKernel 启动

Step2 Python wrapper(模板代码)

复制代码
@cache_once
def _jit_qknorm_rope_module(head_dim, rope_dim, is_neox, dtype) -> Module:
    args = make_cpp_args(head_dim, rope_dim, is_neox, is_arch_support_pdl(), dtype)
    return load_jit("qknorm_rope", *args,
cuda_files=["diffusion/qknorm_rope.cuh"],
cuda_wrappers=[("qknorm_rope", f"QKNormRopeKernel<{args}>::run")])

@register_custom_op(mutates_args=["q", "k"])
def fused_inplace_qknorm_rope(q, k, q_weight, k_weight, cos_sin_cache,
positions, *, is_neox, eps=1e-6, head_dim=0, rope_dim=0):
    head_dim = head_dim or q.size(-1)
    rope_dim = rope_dim or cos_sin_cache.size(-1)
    module = _jit_qknorm_rope_module(head_dim, rope_dim, is_neox, q.dtype)
    module.qknorm_rope(q, k, q_weight, k_weight, cos_sin_cache, positions, eps)

要点:

  • cache_once 装饰器,不用 lru_cache
  • mutates_args 必须标记原地修改张量
  • build marker 只包含编译期模板参数

Step3 编译 flags

可选传递 nvcc 编译参数;硬件版本判断放在 python 层,提前报错。

Step4 单元测试(必须写)

基准参考:分步实现 RMSNorm + FlashInfer RoPE,用来核对结果正确性。 容差:bf16 浮点误差 atol=8e-2, rtol=1e-2 测试网格:head_dim (64,128,256) × rope_dim × is_neox (True/False) × int32/int64 positions,奇数 batch 大小 1/9/129 等,验证 grid-stride 循环边界。 本地执行命令:

复制代码
pytest python/sglang/jit_kernel/tests/diffusion/test_qknorm_rope.py -v

Step5 Benchmark 性能测试

⚠️重点:原地算子不能用 CUDAGraph 计时 ,原地修改张量,graph 多次回放会累积错误。使用run_benchmark_no_cudagraph 测试 case 使用真实模型配置:FLUX / Qwen-Image / Z-Image;同时跑分步版本、融合版本,对比耗时,计算加速比。注册到 CI 性能套件。

Step6 收尾

NCU(NVIDIA 性能分析工具)profile,查看显存带宽利用率、SM 占用率;把算子接入模型 runtime 四层降级逻辑。

开发踩坑清单汇总

表格

坑 规避方案
使用 lru_cache 保存 JIT 模块 统一使用 cache_once
运行时参数写入 JIT build marker marker 只放编译期模板参数
原地算子忘记 mutates_args 声明 @register_custom_op 标记 mutates_args="q","k"
平方求和在 fp16/bf16 低精度累加 加载之后立刻转 fp32 做平方累加
NeoX rotary_lanes 不是 2 的幂 Python 门禁 + C++ static_assert 双层拦截
QK 寻址两套分支 host 端 k 指针负偏移技巧统一寻址公式
CI 注册 est_time 填变量 必须字面量数字,CI 靠 AST 静态解析
in-place 算子 bench 使用 cudagraph 使用 no_cudagraph 版本
门控忘记检查权重 dtype 和输入张量一致 门控增加 weight.dtype 校验

一页极简总结(复习用)

  1. 算子:DiT 的 Attention 前置融合,RMSNorm+RoPE 合并单 kernel,原地更新 Q/K;带宽瓶颈场景,减少显存读写实现加速。
  2. 核心 CUDA 工程:warp-per-head,warp-shuffle 归约求和;NeoX 模式用 shuffle_xor 跨 lane 交换向量对;128bit 向量对齐访存提升带宽。host 侧负偏移统一 QK 寻址。模板 + JIT 编译缓存。
  3. 百度昆仑移植:kernel 代码原样复制,仅调整目录;上层 runtime 门控快照老旧,GQA 模型无法进入融合路径,但仓库现有 MHA 模型不受影响。移植路线是 cuda_like 兼容,不是原生 KLX 硬件定制算子。
  4. 开发规范:双层参数校验(Python 门控 + C++ static_assert),四层降级兜底;单元测试 + bench 必须配套,原地算子注意 torch.compile 和 cudagraph 陷阱。

📖 给初学者的学习路线建议(你可以按顺序学)

  1. 先吃透基础:Transformer、DiT、RMSNorm 数学公式、RoPE 两种配对方式;弄懂 Q/K/V、MHA/GQA
  2. GPU 基础:GPU 硬件层次(Grid/Block/Warp/Lane)、寄存器 / 共享内存 / 显存区别、shuffle 指令含义、带宽受限 vs 计算受限
  3. SGLang JIT 体系:什么是 JIT、模板实例化、缓存机制、custom op、mutates_args 原地语义
  4. 阅读简化版 kernel,跑通单元测试;尝试修改 head_dim 参数,观察 static_assert 报错
  5. 学习性能分析 NCU,看算子带宽占用
  6. 理解昆仑 cuda_like 兼容栈:xpytorch+xmlir 如何翻译 CUDA 代码到 XPU 指令
相关推荐
萧瑟余晖4 小时前
Dubbo 集群容错与负载均衡详解
架构·负载均衡·dubbo
微学AI18 天前
百度、115、夸克太分散?用 LitePan 把多网盘和影音入口收拢到一起
dubbo
能源革命18 天前
AI 日报 2026-09-10
人工智能·dubbo
Meta3919 天前
PostgreSQL报SELECT rule‘s target entry X has different type from column “XXX“
数据库·postgresql·dubbo
vHelios19 天前
【电商项目】生成二维码测试报错Dubbo超时以及data无返回问题
dubbo
小马架构21 天前
Dubbo从入门到实践
dubbo
2601_9652354823 天前
快速了解docker
docker·容器·dubbo
互联网中的一颗神经元1 个月前
09. mheap:全局堆管理器
java·spring·dubbo
学编程的小程1 个月前
影音库不想反复扫网盘:LitePan 聚合多网盘,用 WebDAV 和 STRM 整理播放链路
dubbo