尧图网络科技YAOTU DIGITAL 获取报价
获取报价
首页 / 资讯中心 / 文章详情

AI infra(1)

发布时间:2026/9/29 21:49:46

资讯中心
01
ARTICLE

AI infra(1)

AI infra(1)
前置说明你这份文档是SGLang Diffusion 融合算子fused_inplace_qknorm_rope深度技术分析面向大模型推理内核开发现在我把全文拆解砍掉复杂长句逐块加通俗注释拆成【基础概念 → 算子原理 → CUDA Kernel 代码解读 → 百度昆仑芯移植 → 动手复写教程】全部术语大白话AI 初学者友好先把核心名词一次性解释清楚。背景这个算子是 DiT 图像生成模型FLUX、Qwen-Image 这类文生图Attention 前面的预处理融合算子。 融合算子把多个连续的 GPU 计算步骤合并成 1 个 GPU 核函数 (kernel)减少显存读写提速。词汇预习先看懂这些词后面就轻松很多表格名词通俗解释Kernel / CUDA kernel在 GPU 上并行执行的一段 C 代码CPU 调用它GPU 大量线程同时跑In-place原地计算计算结果直接覆盖原来输入内存不额外开辟新显存存中间结果省显存RMSNorm归一化算法把向量缩放防止数值爆炸稳定模型推理RoPE旋转位置编码给向量注入位置信息让 Attention 知道 token 的先后顺序DiTDiffusion Transformer现在主流文生图模型架构FLUX/Qwen-ImageQ / K / VAttention 机制的三个向量Query 查询、Key 键、Value 值GQA / MHAMHAQ 和 K 头数量完全一样GQAQ 头多、K 头少节省显存warpGPU 最小调度单位1 个 warp 固定 32 个线程 (lane)warp 内可以用 shuffle 指令交换寄存器数据lanewarp 内部单个线程编号 0~31JIT 编译运行时动态编译 C 代码不是程序启动前编译可以根据参数生成定制 kerneltemplate 模板 (C)编译期就固定参数比如 head_dim128编译出来的代码没有 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_dimrms(x) sqrt( (1/head_dim) · Σᵢ xᵢ² eps ) x_out[i] x[i] / rms(x) · w[i]eps极小值一般 1e-6防止分母等于 0除零报错w模型 checkpoint 里面保存的可学习参数逐通道缩放✅ 人话取出单个注意力头的向量 x每个元素求平方全部相加求和除以向量长度 head_dim开平方根得到 RMS向量每个元素除以 RMS乘以权重 w作用把向量数值范围稳定住防止 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 两套分支interleavedGPT-J 风格相邻两个一组 (2j,2j1)out[2j] x[2j]·cosⱼ − x[2j1]·sinⱼ out[2j1] x[2j1]·cosⱼ x[2j]·sinⱼ 一组两个数字在同一个线程寄存器里面计算简单。FLUX/Z-Image 使用这个模式。NeoXLLaMA 风格前半段和后半段配对 (d, dhalf)halfrope_dim/2。out[d] x[d]·cos_d − x[dhalf]·sin_d d half out[dhalf] x[dhalf]·cos_d x[d]·sin_d 麻烦点一对数字不在同一个线程分散在不同 lane需要 warp_shuffle 跨线程拿数据。LLaMA 文本模型常用。部分 RoPE不是 head_dim 全部维度都旋转只旋转前rope_dim维剩下维度原样保留。例 head_dim128rope_dim64只旋转前 64 维。约束rope_dim ≤ head_dim2.3 融合算子完整计算流程单头输入q 向量、k 向量权重 w_q/w_kcos/sin 表每个 token 的位置 pos1. 读取这个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 最大瓶颈显存读写带宽不是计算。表格分步分开 normrope 两个 kernel融合单 kernelGPU 启动次数2 次 kernel launch1 次读取 Q 显存2 次读一次给 normnorm 写完rope 再读一遍1 次读写入 Q 显存2 次norm 写中间结果rope 再写最终结果1 次写中间数据必须写到显存数据全程保存在寄存器不落地显存一句话显存往返减半推理速度接近翻倍。这类算子属于带宽受限算子计算量很小大量时间浪费在读写显存。##3. 输入输出契约Tensor 参数校验表契约调用这个算子张量必须满足的形状、数据类型、内存排布要求TensorMatcher 用来自动校验不满足直接报错。表格参数形状dtypestride 要求说明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同上支持 GQAQ 头数量和 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.cuhC GPU 内核代码在 GPU 设备上执行###4.1 参数结构体 QKNormRopeParamsstruct 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; // 一个block256线程 8个warp ×32laneGPU 线程层级Grid网格→Block线程块→Thread线程 这里一个 block 固定 256 线程拆成 8 个 warp每个 warp32 线程。###4.2 线程映射逻辑重点warp-per-headconst 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 个 warp 负责处理 1 个 head单个注意力头向量总任务量 全部 Q 头 全部 K 头 × token 数量idx任务编号head_id num_qo_heads→ 当前 warp 处理 Q否则处理 K。同一个 kernel 同时处理 Q 和 Kgrid-stride 循环任务数量远超 GPUblock 数量时block 循环反复领取任务避免启动过多 block提高 SM 占用率分配规则head_dim 必须被 32 整除64/128/256每个 lane 分到 head_dim/32 个元素 例 head_dim128128/324每个 lane 负责 4 个数字 为什么 warp-per-head而不是 block-per-head 一个 warp32 线程刚好处理一个 headwarp 内 shuffle 归约不需要共享内存 shared memory不需要线程同步__syncthreads。 一个 block8 个 warp并行处理 8 个 head互相独立等待少。###4.3 RMSNorm 主体代码using Packed packed_tDType; // bf16x2 / fp16x2 打包类型一次读取两个元素 using Storage AlignedVectorPacked, kVecSize; // 128bit向量16字节对齐总线一次性读取 auto input_vec load_asStorage(input, lane_id); // lane读取对齐向量 const auto weight_vec load_asStorage(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] castfp32x2_t(input_vec[j]); // bf16/fp16转fp32高精度 elems[2*j] x0; elems[2*j1] 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] castfp32x2_t(weight_vec[j]); elems[2*j] * norm_factor * w0; elems[2*j1] * norm_factor * w1; //RMSNorm计算完成结果保存在寄存器elems数组 }四个工程要点注释128bit 对齐向量读取一次读取 16 字节充分利用 GPU 内存总线带宽比逐个读取快很多升 fp32 累加平方bf16 精度很低大量数字累加误差会越来越大加载之后立刻转 fp32 计算warp::reduce_sum基于__shfl_xor_sync蝶形求和32lane 把各自的 sum 汇总成总和。不需要 shared 内存归一化结果保存在寄存器数组 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_maskkRotaryLanes(); 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_castint(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配对 lane 编号 lane_id ^ kHalfRotaryLanes异或。前一半 lane 和后一半 lane 两两配对互相拿到对方寄存器的值符号处理前半 lane 的公式需要减去 x [dhalf]所以 swapped 取负后半 lane 不需要取负(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[i1]; elems[i] x * cos[half_idx] - y * sin[half_idx]; elems[i1] 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 头QK 头编号处理 K 如果不做偏移Q 和 K 寻址公式需要两套 if 分支判断代码复杂。✅ 技巧CPU 端预先把 K 指针向前偏移一段负偏移指向 K 内存起始地址前面。 kernel 里面 Q、K 可以复用同一套寻址公式kernel 内部消除分支简化代码。###4.6 static_assert 编译期护栏 模板参数编译期静态断言提前拦截非法参数组合编译阶段直接报错不要等到 GPU 运行才崩溃。表格static_assert含义kHeadDim % kWarpThreads 0head_dim 必须被 32 整除保证 32lane 均匀分配元素kRopeDim 0 kRopeDim kHeadDim部分 RoPE 约束旋转维度不能超过 head_dimkElemsPerThread %2 0打包向量成对适配 fp16x2/bf16x2kRopeDim % kElemsPerThread 0参与旋转的 lane 必须完整拥有元素不能拆分NeoXkRotaryLanes 是 2 的幂shuffle 异或配对逻辑成立的前提###4.7 host 侧 run () 函数CPU 端入口调用 GPU kernelstatic 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_dtypeDType() .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 ? kernelint32_t : kernelint64_t; //③ 计算合适block数量不超过GPU SM最大并发防止过度启动block static const uint32_t kOccupancyTable[2] { get_blocks_per_sm(kernelint32_t, 256), ... }; const auto num_blocks std::min(max_blocks, needed_blocks); //④ LaunchKernel启动GPU kernelRAII封装自动检查cuda error LaunchKernel(num_blocks, kThreadsPerBlock, device.unwrap()) .enable_pdl(kUsePDL)(selected_kernel, params); }PDLProgrammatic Dependent Launch。SM90H100/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, fQKNormRopeKernel{args}::run)], )逻辑cache_once自定义装饰器缓存编译后的 so不用 lru_cache因为 lru_cache 和 torch.compile 冲突。make_cpp_args 收集编译期模板参数组成唯一 key只有模板参数变化才会重新编译。❗重点token 数量、eps 这类运行时参数绝对不能放进缓存 key否则每次推理尺寸变化重复编译内存爆炸。load_jit 调用 nvcc 编译 cuh 代码生成动态库 so加载到 python。模板不同生成不同版本 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_resulttorch.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, eps1e-6, head_dim0, rope_dim0) - 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(...)四层检查顺序环境变量开关可以手动关闭融合算子平台、数据类型、张量连续性、eps 相等性检查can_use_fused_inplace_qknorm_rope 能力检查试编译全部通过才走融合 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 一行变化只是文件目录移动逻辑不变单测 benchmarkimport 路径修改测试逻辑完全复用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.shapeQ 头数量必须等于 K 头数量。后果 GQA 模型Q 头≠K 头无法进入融合算子直接降级到慢路径。 但是kernel 底层代码本身原生支持 GQA只是上层 python 判断条件卡住了。 好在百度仓库内用到这个算子的模型Z-Image/FLUX/Qwen-Image全部是 MHAQ 头 K 头数量一样现有模型不受影响新增 GQA 模型才会踩坑。上游还额外增加的保护百度快照没有torch.compile 编译期保护防止图捕获阶段误入融合算子cos_sin_cache 形状、设备校验positions 自动转换设备与 dtype百度要求调用方自己保证###8.3 CI 测试注册 API 适配 上游 CI 参数register_cuda_ci(est_time44, stagexxx, runner_configxxx)百度 CI合并成 suite 单参数register_cuda_ci(est_time44, suitexxx)只是 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 上算子能力盘点✅ 已具备能力数学逻辑完全等价RMSNormRoPE 融合、interleaved/NeoX、部分 RoPE、GQA 内核支持、fp16/bf16、int32/int64 位置、原地计算全套单元测试 性能 benchmarkZimage / FLUX / Qwen-Image 在 MHA 场景下可以成功走到融合快路径四层降级兜底算子不可用时自动切分步不会崩溃⚠️ 缺口runtime 门控不支持 GQA 模型走融合路径缺少 torch.compile 编译期保护PDL 指令PDL 是英伟达 SM90 专属特性昆仑芯编译时判定不支持 PDL编译出来不带 PDL 逻辑只是少一点启动重叠不影响正确性。没有 Krea2 模型接入代码属于模型层缺失不是算子本身重要区分两份仓库算子jit_kernel/目录下算子cuda_like 兼容路线直接复用原版 CUDA 代码靠昆仑 xPU 兼容层 xpytorchxmlir 转换执行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推荐开发顺序先写数据通路加载向量 → fp32 转换 → 平方求和归约 → RMSNorm → RoPE 旋转 → 写回显存设计线程映射warp-per-head增加 static_assert 编译期约束护栏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, fQKNormRopeKernel{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, eps1e-6, head_dim0, rope_dim0): 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_cachemutates_args 必须标记原地修改张量build marker 只包含编译期模板参数Step3 编译 flags可选传递 nvcc 编译参数硬件版本判断放在 python 层提前报错。Step4 单元测试必须写基准参考分步实现 RMSNorm FlashInfer RoPE用来核对结果正确性。 容差bf16 浮点误差atol8e-2, rtol1e-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 -vStep5 Benchmark 性能测试⚠️重点原地算子不能用 CUDAGraph 计时原地修改张量graph 多次回放会累积错误。使用run_benchmark_no_cudagraph测试 case 使用真实模型配置FLUX / Qwen-Image / Z-Image同时跑分步版本、融合版本对比耗时计算加速比。注册到 CI 性能套件。Step6 收尾NCUNVIDIA 性能分析工具profile查看显存带宽利用率、SM 占用率把算子接入模型 runtime 四层降级逻辑。开发踩坑清单汇总表格坑规避方案使用 lru_cache 保存 JIT 模块统一使用 cache_once运行时参数写入 JIT build markermarker 只放编译期模板参数原地算子忘记 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 校验一页极简总结复习用算子DiT 的 Attention 前置融合RMSNormRoPE 合并单 kernel原地更新 Q/K带宽瓶颈场景减少显存读写实现加速。核心 CUDA 工程warp-per-headwarp-shuffle 归约求和NeoX 模式用 shuffle_xor 跨 lane 交换向量对128bit 向量对齐访存提升带宽。host 侧负偏移统一 QK 寻址。模板 JIT 编译缓存。百度昆仑移植kernel 代码原样复制仅调整目录上层 runtime 门控快照老旧GQA 模型无法进入融合路径但仓库现有 MHA 模型不受影响。移植路线是 cuda_like 兼容不是原生 KLX 硬件定制算子。开发规范双层参数校验Python 门控 C static_assert四层降级兜底单元测试 bench 必须配套原地算子注意 torch.compile 和 cudagraph 陷阱。 给初学者的学习路线建议你可以按顺序学先吃透基础Transformer、DiT、RMSNorm 数学公式、RoPE 两种配对方式弄懂 Q/K/V、MHA/GQAGPU 基础GPU 硬件层次Grid/Block/Warp/Lane、寄存器 / 共享内存 / 显存区别、shuffle 指令含义、带宽受限 vs 计算受限SGLang JIT 体系什么是 JIT、模板实例化、缓存机制、custom op、mutates_args 原地语义阅读简化版 kernel跑通单元测试尝试修改 head_dim 参数观察 static_assert 报错学习性能分析 NCU看算子带宽占用理解昆仑 cuda_like 兼容栈xpytorchxmlir 如何翻译 CUDA 代码到 XPU 指令
02
RELATED NEWS

相关资讯

更多网站建设与数字化升级内容

03
WHY YAOTU

想打造同款高转化官网?

懂行业、懂生意,从建站到增长一站式陪跑

◈

场景化定制

不做模板站,围绕你的业务场景量身设计,小众不撞款。

◐

营销型架构

以转化目标组织内容与路径,让官网真正带来询盘。

▲

全周期服务

设计、开发、运营、运维一体,上线只是开始。

免费获取你的建站方案

留下需求,专属顾问 24 小时内为你输出方案建议。