SGLang 扩散模型融合算子包sglang.kernels.ops.diffusion完全指南:DiT / VAE 位精确内核与质量门控机制
【免费下载链接】sglangSGLang is a high-performance serving framework for large language models and multimodal models.项目地址: https://gitcode.com/GitHub_Trending/sg/sglang
导读
本文档面向在 SGLang 中接入扩散模型(multimodal-generation)服务、或需要为 DiT Transformer 块、VAE 编码器/解码器编写高性能融合算子的开发者。文章以python/sglang/kernels/ops/diffusion/README.md为骨架,结合 包入口 的注册表与导出表、位精确门控、质量门控 以及各算子域的 Triton / CuTe-DSL / FlyDSL / JIT CUDA 实现,深入讲解该包的设计哲学、算子选择矩阵、数值契约与接入新算子的完整流程。读完本文,你将掌握:为什么扩散算子必须区分"位精确"与"近似"两条门控路径、如何在众多看起来可互换的归一化算子中正确选型、以及如何把一个新的融合内核接入 SGLang 的统一内核注册表并保证服务质量不劣化。
定位:一个"按模型定制"的融合算子包,而非通用算子库
与 SGLang 面向 LLM 的算子分组不同,sglang.kernels.ops.diffusion中几乎没有任何通用算子。每一个内核都对应某个具体模型中的某一条具体 eager 算子链(例如 ERNIE-Image 的 adaLN 链、Wan VAE 的 channel-first RMSNorm+SiLU、FLUX.2 的 QK RMSNorm+RoPE+QKV 打包、LingBot Video 的分组受限 MoE 路由等),其价值不仅来自带宽优化,更来自它精确复现了哪一条舍入边界。
这一点在扩散模型场景下格外关键:多步去噪(multi-step denoising)会把单步的舍入差异逐级放大为肉眼可见的质量损失。因此本文档给出一个核心判断标准——"差不多"(close)与"位精确"(bit-exact)是两个不同的产品,二者的挂载门控策略完全不同:
| 契约类型 | 判定方式 | 挂载策略 |
|---|---|---|
| 位精确(bit-exact) | torch.equal与 eager 链逐位一致 | 无条件挂载(unconditionally mounted),无需质量门控 |
| 近似(close) | 数值接近但不逐位一致 | 仅质量门控:只对quality="extra-high"与quality="high"的请求,在 batch 边界、按整个 Transformer 全有或全无地挂载 |
一个关键反例:一个看起来无害的普通 fp32 单遍 norm 融合,实际上并不安全——在 ERNIE-Image 上它把 50 步去噪轨迹推到了 PSNR 18.83 dB,这正是该包重写为位精确实现的直接动因(见 rmsnorm_scale_shift_bitexact.py 的模块文档)。
另外,模型/checkpoint 原生的(model/checkpoint-native)选择——例如与所选模型或部署路径绑定的宽松契约内核、稀疏算子、FP8/NVFP4 量化产出器——不属于请求级quality档位管制的范围:quality 档位既不会选择也不会禁用这些独立的模型级选择。
导入面:只从包门面导入,绝不从子模块导入
包的对外契约非常明确:
from sglang.kernels.ops.diffusion import fused_rmsnorm_scale_shift_bitexact必须从包(facade)导入,绝不能从子模块导入。内部文件布局(norm、rope等子包)随时可能调整,而门面接口是稳定的。唯一的例外是刻意只测单一后端的测试。
这种懒解析(lazy resolution,PEP 562)设计在 包入口 中体现得十分彻底:
- 各后端拥有互不相交的重依赖(Triton、CUTLASS/CuTe-DSL、ROCm 上的 FlyDSL),若在包导入时急切地 re-export,会把所有后端变成所有平台上的硬性 import-time 依赖;
- 因此入口模块维护一张
_EXPORTS字典(符号 → 所属子模块),通过模块级__getattr__在首次属性访问时才import_module并缓存到globals(),后续查找完全绕过__getattr__; - 同时维护
_SPECS元数据元组,通过register_kernel(KernelSpec(...))把每个算子注册进 SGLang 的统一内核注册表。注册只登记元数据——不导入 torch、不触发任何后端、不触发 JIT 构建。同一算子携带多个后端(如scale_residual_norm_scale_shift同时有 Triton、CuTe-DSL、FlyDSL 三种实现)时,调用方通过select_kernel按名字选用。
以入口文件中的注册表为例,一个算子条目包含:op名、后端类型(KernelBackend.TRITON / JIT / KDA / CUTE_DSL / FLYDSL / AOT)、目标函数路径(相对或全限定)、能力约束(CUDA / HIP / SM100+,如 Qwen-Image 的 QKV epilogue 明确要求CapabilityRequirement.cuda(min_sm=(10, 0)))、以及一句话描述。
目录布局:算子域、编译器后缀、以及"不是内核"的目录
普通实现按**算子域(operator domain)**划分子包;编译器/后端则体现在文件名的后缀(_triton、_jit、_cutedsl、_flydsl,以及表示数值语义的_bitexact)。带 Kernel Design Agents(KDA)来源的实现位于sglang.kernels.kda_kernels,门面仍是其唯一受支持的运行时导入面。
norm/ RMSNorm / LayerNorm / GroupNorm 及其融合 epilogue modulate/ adaLN modulate、gating、timestep 条件化 rope/ 旋转位置编码及融合进其中的 QK-norm 链 activation/ SiLU / GLU / GELU 融合 attention/ 稀疏线性注意力、gated delta-net routing/ 扩散模型 MoE 路由与专家选择 layout/ 纯数据搬运:USP/Ulysses relayout、varlen pack、causal pad common/ 数值原语、平台谓词、非 Triton 回退 sites/ 请求作用域的挂载策略 —— 不是内核(见下文) ext/ JIT C++/CUDA 扩展(Hunyuan3D raster/inpaint)—— 不是内核 ../../kda_kernels/ Agent 生成的实现及其 JIT CUDA 源码实际目录与 README 完全一致,例如 norm 下有rmsnorm_scale_shift_bitexact.py、scale_residual_norm_cutedsl.py、group_norm_silu_triton.py等;sites 下有bitexact_gate.py、quality_gate.py及各模型站点文件。
数值契约与质量策略:位精确、质量门控与"非内核"边界
位精确内核 → 无条件挂载
位精确内核逐位复现 eager 链的每一个 aten 舍入边界,有时精确到归约树(reduction tree)本身。例如:
sglang.kernels.kda_kernels/layernorm_modulate_triton.py复现 torch 2.11 的vectorized_layer_norm_kernel(128 线程 Welford、_rcp4保护的倒数、shfl.down折叠顺序、div.rn+MUFU.RSQ);norm/rmsnorm_scale_shift_bitexact.py复现 flashinfer 的 CuTe-DSLRMSNormKernel的 fragment 顺序与shfl.bfly折叠。
即使如此,这些内核仍会在首次见到真实输入时用sites/bitexact_gate.py对实时 eager 链做一次自校验,一旦不匹配就永久回退——因为它们复现的 dispatch 本身可能随平台变化。
BitExactFusionGate 提供两种校验模式:
- once-for-all(默认):第一次
torch.equal通过后,融合路径对之后所有调用永久生效(GLM / Ernie 使用); - per-signature:每个不同的签名
sig独立校验(FLUX / Sana 使用),以匹配按形状分派的 aten LayerNorm dispatch。
can_attempt_once()明确禁止在torch.compile追踪或 CUDA graph 捕获期间做首次校验(会执行 eager 参考链 + host 同步),而accept_or_fallback()在输出与参考不一致时打印告警、永久disable()并返回 eager 参考输出,保证正确性优先。它还提供了flashinfer_rmsnorm_diagnostic_hint()回调,用于在精确度失配时输出 FlashInfer 归一化后端诊断(检测_USE_CUDA_NORM标志、FLASHINFER_USE_CUDA_NORM环境变量及 flashinfer 相关包版本)。
非位精确内核 → 质量门控
近似内核只能通过sites/quality_gate.py中的QualityGatedFusion挂载到被标记的nn.Module站点上,且仅对quality="extra-high"和quality="high"的请求生效,在 batch 边界、按每个 Transformer 全有或全无(all-or-nothing):
extra-high只追加这些请求门控的 DiT/VAE 融合;high是累积的,还可能启用模型自有的近似路径,例如 Cache-DiT 或低精度 decode。
QualityGatedFusion通过mark/mount/unmount/is_enabled维护站点的 marker 属性与 enabled 标志(enabled 是一个普通模块属性,以便编译后的模型 forward 读取时不依赖该辅助对象);mount还支持reject_reason静态守卫回调——任何站点未通过守卫时,整个模型保持参考路径。
一个质量门控的实际例子:SANA-Video 线性注意力
SANA-Video 的质量门控线性注意力站点,第一个 GEMM 保持 BF16 输入但要求 FP32 累加/输出,第二个 GEMM 用 FP32 运行;而默认路径在两次 GEMM 之前就把 Q/K/V 提升到 FP32。这个"半精度输入 + FP32 输出"的混合契约正是通过站点文件(sites/sana_video_linear_attention_site.py)实现的。
什么是"不是内核":sites 与 ext
sites/重写nn.Module树(mark / mount / unmount),不是算子;它是唯一允许(在函数内部、懒加载地)引用multimodal_gen类型的地方,因为检查模型模块就是它的全部工作;ext/构建没有后端维度、也没有数值契约的 C++/CUDA 扩展(如 Hunyuan3D rasterizer),它们共享本包的构建机制,但刻意独立成目录。
入口点协议:predicate + kernel 配对
每个公开内核都是谓词 + 内核配对:
if can_use_<op>(...): out = <op>(...) else: out = <reference chain>关键约束:内核在遇到不支持的输入时直接 raise,绝不返回None——一个静默的None太容易被调用方漏掉检查,而失败模式将是一张看起来错误的图像,而不是一个异常。这一点在group_limited_topk的入口函数中可见:can_use_group_limited_topk失败时直接raise ValueError(...),列出全部前置条件(非空连续 CUDA float32[tokens, experts]张量、每组至少两个 2 的幂专家、1 < n_group、0 < topk_group <= n_group、top_k不超过所选组容量)。
同时,公开内核用@register_custom_op注册,并配套fake_impl(meta/fake 实现),确保在 torch.compile / 元设备模式下可追踪(见rmsnorm_scale_shift_bitexact.py与group_limited_topk_triton.py)。
选择矩阵:从"看起来可互换"的算子中正确选型
README 明确警告:"好几个归一化算子看起来可互换,实际并不是。从这里开始选型。"
Norm + scale/shift(adaLN)
| Entry point | Backend | Contract | Applies to |
|---|---|---|---|
fuse_scale_shift_kernel | Triton | close | 连续 BLC;scalar/row/token 调制加 causal-video[B, F, 1, C],使用静态封顶的 2 的幂 tile 避免请求期 autotuning |
fused_rmsnorm_scale_shift_bitexact | Triton | 对 flashinfer CuTe RMSNorm + aten modulate 位精确 | bf16、连续行、H == 64 * threads_per_row |
fused_scale_residual_rmsnorm_scale_shift_bitexact | Triton | 位精确,含前置 residual-gate add | 同上 |
fused_layernorm_modulate | Triton | 对 atenvectorized_layer_norm位精确 | bf16、N % 4 == 0、16B 对齐 |
fused_norm_scale_shift/fused_scale_residual_norm_scale_shift | CuTe-DSL | fp32 统计量、close | fp16/bf16/fp32、LN 或 RMS、多种广播模式 |
flydsl_norm_scale_shift/flydsl_fused_residual_norm_scale_shift | FlyDSL | close | 仅 ROCm gfx950 |
try_fused_scale_residual_norm_scale_shift_nvfp4 | JIT CUDA | 匹配所选 NVFP4 产出器契约 | Qwen residual LayerNorm/调制 + FC1 NVFP4 量化 |
fuse_layernorm_scale_shift_gate_select01_kernel | Triton | close | 每个 token 在两行调制之间选择(Qwen-Image) |
norm_infer/rms_norm_fn | Triton(+torch/NPU/MPS 回退) | close | 通用入口;上面都不适用时用它 |
以fused_rmsnorm_scale_shift_bitexact为例,其数值契约在源码文档中逐条写死(rmsnorm_scale_shift_bitexact.py):
RMSNorm.forward_cuda分派到 flashinfer CuTe-DSLRMSNormKernel:连续 bf16 行、H == 64 * threads_per_row(H≤3072 时 32,否则 64)时,每个线程tx拥有列{8*TPR*b + 8*tx + v},fragment 按 v 最快、再 b 排序;- 平方在 fp32 中分别舍入、不使用 FMA(用不透明
mul.rn.f32阻止编译器收缩成 FMA),然后对 64 个 fragment 值做有序的顺序 fadd 链(无 reassoc); - warp 归约用
shfl.bfly偏移 1,2,4,8,16——即相邻对折叠树;rstd = rsqrt.approx.f32(sum_sq / H + eps);输出y = (bf16)(float(x) * rstd * (w + 0.0))只做一次最终舍入; - aten modulate 链在每步之后舍入到 bf16:
round(1 + scale)、round(y * that)、round(prod + shift); - 残差变体重现 eager 对
round(gate * update)、round(residual + that),再把舍入结果送入同一忠实 norm。
作者还特别提示了一个易踩的坑:折叠阶段num_warps必须与被复现内核的 warps-per-row 一致,否则会触发病态的 Triton 布局转换(实测 25µs → 500µs)。
Norm 变体
| Entry point | Backend | Contract | Applies to |
|---|---|---|---|
triton_group_norm_silu/apply_group_norm_silu | Triton | close | NCHW 连续、任意 channels-per-group、总是施加 SiLU |
group_norm_silu_4d/group_norm_silu_rows | Triton | close | 仅 channels_last;2 的幂C <= 2048;可选 SiLU。这让 VAE decoder 可以端到端保持 channels_last,无需nchwToNhwc |
wan_rmsnorm_silu | Triton | close | 稠密channels_last_3d5D(stride(C) == 1)、Wan VAE channel-first RMSNorm + SiLU |
rmsnorm_scale/rmsnorm_tanh_residual | Triton | bf16 原生统计量 | Z-Image(与其自身参考精确一致)、Ideogram 4(门控) |
zimage_qk_rmsnorm_native | Triton | bit-exact | Z-Image 每头 QK RMSNorm |
fused_qk_head_layernorm | Triton | bit-exact | 每头 LN on q/k、dim_head % 4 == 0、<= 128 |
triton_one_pass_rms_norm | Triton | close | 独立 RMSNorm,单遍 |
残差门控(Residual gating)
| Entry point | Backend | Contract | Applies to |
|---|---|---|---|
residual_gate_add | KDA(JIT CUDA) | 对residual + update * gate位精确 | 连续张量;或转置稠密[B, tokens, hidden]残差/输出 + 连续 update + row-broadcast gate(SANA-Video) |
转置稠密路径用共享内存 tile 以逻辑行主序读取 update,同时保持残差读取与输出写在其[B, hidden, tokens]底层布局上合并(coalesced)。README 特别警告:不要仅仅为了走普通路径而插入.contiguous()——那会在每个残差站点恢复一整次张量拷贝。
RoPE / QK-norm
| Entry point | Backend | Contract |
|---|---|---|
fused_inplace_qknorm_rope | JIT CUDA | 相对拆分基线只多一步 bf16 舍入;round_norm_before_rope=True时精确;支持 compact 与全宽 NeoX/interleaved cache |
fused_qknorm_rope_pack_kv | JIT CUDA | 同上,额外打包前缀 K/V |
try_fused_flux2_qkv_epilogue | KDA(JIT CUDA) | 对所选 BF16 链位精确;FLUX.2 QK RMSNorm + RoPE + 联合 QKV 打包 |
try_fused_qwen_qkv_epilogue | JIT CUDA | 对所选 BF16 链位精确;Qwen-Image QK RMSNorm + RoPE + 联合 QKV 写入;SM100+ |
fused_rope_rotate_half_bitexact | Triton | 位精确(仅逐元素) |
fused_interleaved_rope_fp64 | JIT CUDA | 对 SANA-Video 配对的 fp64 RoPE 位精确 |
fused_inplace_helios_qk_rope | JIT CUDA | 对 Helios 转置频率布局的配对就地 RoPE 位精确 |
ltx2_qknorm_split_rope_cuda | KDA(JIT CUDA) | close;在 B200 上验证 |
fused_ltx25_decoder_rope | JIT CUDA | 由缓存的 compact 轴表配对 3D RoPE,位精确 |
apply_rotary_embedding | Triton(+回退) | close;通用入口 |
hunyuan_qkv_rope_pack | Triton | 位精确;单遍打包 QKV 并施加 RoPE |
MoE 路由
| Entry point | Backend | Contract | Applies to |
|---|---|---|---|
group_limited_topk | Triton | 所选专家 id 集合与受守卫的 CUDAtorch.topk(..., sorted=False)链一致;输出顺序未指定 | LingBot Video 默认开启的 sigmoid+bias 分组受限路由;连续 fp32[tokens, experts]、每组至少两个 2 的幂专家 |
group_limited_topk_triton.py 的实现细节值得展开:参考的 LingBot Video 路由用一个小内核链完成分组受限 top-k(每组 top-2 与求和、组 top-k、scatter_进零掩码、expand/reshape广播、masked_fill为-inf、最终专家 top-k)。在 launch-bound 的单 GPU 上,这条链纯粹是开销——每个中间张量都很小,整个计算受带宽与 launch 限制。融合内核每个 token 一个 program:一次性加载自己的 score 行、在寄存器中归约组内和、用-inf掩蔽未选组、写出 top-k 专家 id,并在 128 专家 / 4 组 / 选 2 组 / top-8 的生产配置上保持所选集合一致。注意其对重复最大值的处理:掩掉所有等于m1的值会丢掉组内第二大的项(组内出现重复最大值时),因此实现用tl.min(tl.where(g == m1, group_e, BLOCK_EPG))只移除恰一份第一个最大值。
数据搬运与量化布局产出器
以下算子要么是位精确的数据搬运,要么是保持运算次序的算术(same-order arithmetic):usp_merge_heads、pack_qkv_destination_major、fused_pack_qkv、fused_pack_segmented_qkv、fused_scatter_to_padded、fused_causal_conv3d_cat_pad_cuda、cat_pad_channels_last_3d、dup_up3d_add、fused_temb_table_slices、ltx2_ada_values9。此外:
fused_layernorm_modulate_fp8_quant_raw把 FLUX.2 的 LayerNorm、adaLN 调制与静态 FP8 量化折叠进单内核;try_flux2_token_cat_fp8与try_flux2_token_cat_nvfp4把分支拼接直接融合进 FLUX.2 checkpoint 路径所选的量化表示。
fused_temb_table_slices尤其值得关注:eager 版本(table + temb.float()).chunk(6, dim=2)在 704p/121f 时物化约 8GB 的 fp32,并把六个带步长的切片交给下游,它们的.contiguous()调用又把每个切片各复制一遍——融合后这些中间拷贝全部消失。
在真实模型中的接线方式
位精确内核在python/sglang/multimodal_gen/runtime/models/dits/下的各 DiT 模型中使用。以 ERNIE-Image 为例(ernie_image.py):模块顶部为每条融合路径各建一个BitExactFusionGate(fused-norm、fused gated-norm、fused RoPE、fused QKNorm+RoPE、fused GELU-mul),正向传播中调用fused_rmsnorm_scale_shift_bitexact等入口,通过门的accept_or_fallback在首次校验通过后启用、失败则永久回退到 eager。类似的接线也出现在flux.py、flux_2.py、glm_image.py、qwen_image.py、sana.py等文件中,与 README 中 "dispatch 可能在调用方脚下变化" 的警告相互印证。
添加一个新内核:六步流程
README 给出了官方接入流程,结合源码可完整还原每一步:
- 放对位置:普通实现放进对应算子域子包(
norm/、rope/、modulate/等);KDA 工作流生成的实现放进sglang.kernels.kda_kernels,连同其源码修订信息与任何 JIT CUDA 源文件。 - 登记:在 包入口 的
_EXPORTS(符号 → 子模块映射,供 PEP 562 懒导入)与_SPECS(KernelSpec注册元数据)中各登记一条。规范上_EXPORTS按域 → 模块 → 符号排序,新公开内核只能出现在这里。 - 给出
can_use_*谓词:不支持的输入直接 raise,不返回None。 - 在模块 docstring 中声明数值契约:包括在哪些形状上验证过(例如
rmsnorm_scale_shift_bitexact声明在(1,4216,4096)/(1,4096,4096)/(2,1140,4096)/(1,128,2048)bf16 上做过torch.equal验证)。 - 若非位精确,走
sites/门控:必须同时挂载到extra-high与high,绝不能在默认的lossless路径上生效。 - 测试:算子域测试放进
test/registered/kernels/ops/diffusion/(当前包含test_norm.py、test_rope.py、test_modulate.py、test_layout.py、test_routing.py、test_sites.py、test_model_fast_paths.py等 19 个测试文件),模型接线测试进test_model_fast_paths.py。另有test/registered/e2e/diffusion/下的端到端用例(如 test_diffusion_unit.py)以run_diffusion_suite("unit")方式组织 1-GPU / 2-GPU / B200 / BCG 等 CI 泳道。
小结
sglang.kernels.ops.diffusion是 SGLang 中一个纪律性极强的融合算子包:它用"位精确 → 无条件挂载、近似 → 质量门控"的双轨契约,把扩散模型多步去噪对舍入误差的敏感性变成可工程化的准则;用"predicate + kernel"协议和 PEP 562 懒门面隔离了 Triton / CuTe-DSL / FlyDSL / JIT CUDA 多种后端的重依赖;再用sites/把"改模型模块"这一非算子职责隔离在算子域之外。无论是为新的 DiT/VAE 模型接入融合算子,还是理解 SGLang 扩散服务的数值质量保障机制,本文的选择矩阵与接入流程都可以直接作为工作起点。
【免费下载链接】sglangSGLang is a high-performance serving framework for large language models and multimodal models.项目地址: https://gitcode.com/GitHub_Trending/sg/sglang
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考