From af2ffdfabb1d053abaaf5bc110bd750411ef1209 Mon Sep 17 00:00:00 2001 From: BBuf <1182563586@qq.com> Date: Mon, 17 Aug 2026 15:09:11 +0800 Subject: [PATCH 1/5] [kernels] Reorganize ops/diffusion by operator domain behind a lazy facade `ops/diffusion` was the only operator group in `kernels/ops` organized by *backend* (`triton/`, `cutedsl/`, `flydsl/`) rather than by operator, so the several implementations of one logical op lived in three directories and answering "does this fusion exist on ROCm?" meant grepping the whole group. It also mixed three kinds of thing: kernels, request-scoped mount policy, and JIT C++/CUDA extensions that are not operators at all. Structure. One subpackage per operator domain -- `norm`, `modulate`, `rope`, `activation`, `attention`, `layout` -- with the backend as a filename suffix (`_triton` / `_jit` / `_cutedsl` / `_flydsl`, or `_bitexact` where that says more), matching `ops/attention` and `ops/gemm`. `common/` holds the numerics primitives and platform plumbing, `sites/` the quality-gate and mount policy, `ext/` the Hunyuan3D rasterizer and mesh processor. Kernel bodies are moved verbatim; the SASS-level bit-exact replications are not touched. Import surface. `__init__.py` is now a PEP-562 lazy facade over 117 exports plus 38 `KernelSpec` registrations (was 4). Lazy is required, not stylistic: the backends have disjoint heavy dependencies -- Triton, CUTLASS/CuTe-DSL, FlyDSL (ROCm), MLX (Apple) -- and an eager re-export makes every one of them an import-time requirement everywhere. All 141 external import sites now go through the facade, so the internal layout can move again without touching call sites; `test_import_surface.py` enforces that, checks both directions of the export table, and asserts no leaf module is imported eagerly. Protocol. `can_fuse_*` is folded into `can_use_*`, and the kernels that signalled "unsupported" by returning `None` (`group_norm_silu_4d`, `group_norm_silu_rows`, `wan_rmsnorm_silu`) now raise, with a public predicate callers check first -- a forgotten `is None` check produced a wrong image rather than an exception. The four hand-rolled import-time platform `if` chains collapse into one `common.platform.select_impl` seam, which also removes three of the four upward imports of `multimodal_gen.current_platform` (the last one is confined to `common/platform.py` and documented). `KernelBackend.FLYDSL` is added (with its `BACKEND_METHODS` entry) so the ROCm gfx950 norm kernels can be registered as their own provenance rather than mislabelled as Triton. Tests: 37 files -> 8. Five were not diffusion-kernel tests and move to their real homes (two quantization suites to `ops/quantization`, the perf-logger sync test to `profiling/`). The rest merge into one suite per operator domain plus `test_sites.py` (gate protocols) and `test_model_fast_paths.py` (per-model wiring), collapsing the duplicated per-model and per-backend copies into parametrized cases. The timestep-embedding "perf test" -- skipped by default, asserting nothing -- becomes a benchmark. Test bodies are preserved; only the genuinely mechanical duplication is parametrized away. `README.md` adds the selection matrix: nine norm implementations that look interchangeable and are not, each with its numerical contract and the shapes and layouts it supports. Co-Authored-By: Claude Opus 5 (1M context) --- .../references/fuse-overlap-catalog.md | 14 +- python/sglang/kernels/fused_op.py | 1 + python/sglang/kernels/ops/diffusion/README.md | 153 ++++ .../sglang/kernels/ops/diffusion/__init__.py | 548 +++++++++++--- .../ops/diffusion/activation/__init__.py | 1 + .../sana_conv_post_triton.py} | 2 +- .../silu_mul_bitexact.py | 2 +- .../ops/diffusion/attention/__init__.py | 1 + .../sana_wm_gdn_chunkwise_triton.py} | 2 +- .../sana_wm_gdn_triton.py} | 2 +- .../sparse_linear_attn_triton.py} | 0 .../kernels/ops/diffusion/common/__init__.py | 6 + .../fallback_mps.py} | 8 +- .../fallback_npu.py} | 0 .../fallback_torch.py} | 0 .../diffusion/{triton => common}/numerics.py | 0 .../kernels/ops/diffusion/common/platform.py | 89 +++ .../kernels/ops/diffusion/ext/__init__.py | 7 + .../hunyuan3d_rasterizer/__init__.py | 2 +- .../ext/hunyuan3d_rasterizer/rasterizer.cpp | 172 +++++ .../ext/hunyuan3d_rasterizer/rasterizer.h | 61 ++ .../hunyuan3d_rasterizer/rasterizer_gpu.cu | 173 +++++ .../{render/__init__.py => ext/loader.py} | 5 +- .../mesh_processor/__init__.py | 2 +- .../ext/mesh_processor/mesh_processor.cpp | 170 +++++ .../kernels/ops/diffusion/layout/__init__.py | 5 + .../causal_conv3d_cat_pad_jit.py} | 0 .../causal_conv3d_cat_pad_triton.py} | 0 .../ulysses_qkv_triton.py} | 0 .../usp_relayout_jit.py} | 0 .../varlen_pack_pad_triton.py} | 0 .../wan_causal_cache_triton.py} | 0 .../ops/diffusion/modulate/__init__.py | 1 + .../indexed_modulation_triton.py} | 2 +- .../ltx2_ada_values_triton.py} | 0 .../modulate_scale_shift_jit.py} | 0 .../residual_gate_add_jit.py} | 0 .../scale_shift_triton.py} | 41 +- .../timestep_embedding_jit.py} | 0 .../wan_temb_table_slices_triton.py} | 0 .../kernels/ops/diffusion/norm/__init__.py | 7 + .../diffusion/norm/cutedsl_common/__init__.py | 1 + .../cutedsl_common}/norm_fusion.py | 2 +- .../common => norm/cutedsl_common}/reduce.py | 0 .../utils.py => norm/cutedsl_utils.py} | 0 .../fused_residual_norm_flydsl.py} | 0 .../diffusion/{ => norm}/group_norm_silu.py | 2 +- .../group_norm_silu_triton.py} | 0 .../group_norm_silu_twopass_triton.py} | 68 +- .../layernorm_modulate_triton.py} | 2 +- .../native_bf16_rmsnorm_triton.py} | 0 .../norm_scale_shift_jit.py} | 0 .../{triton/norm.py => norm/norm_triton.py} | 23 +- .../rmsnorm_onepass_triton.py} | 20 +- .../rmsnorm_scale_shift_bitexact.py | 2 +- .../scale_residual_norm_cutedsl.py} | 8 +- .../wan_rmsnorm_silu_triton.py} | 11 +- .../zimage_qk_rmsnorm_triton.py} | 0 .../hunyuan3d_rasterizer/rasterizer.cpp | 140 ---- .../render/hunyuan3d_rasterizer/rasterizer.h | 56 -- .../hunyuan3d_rasterizer/rasterizer_gpu.cu | 130 ---- .../render/mesh_processor/mesh_processor.cpp | 163 ---- .../kernels/ops/diffusion/rope/__init__.py | 1 + .../hunyuan_qkv_pack_triton.py} | 0 .../ltx2_qknorm_split_rope_jit.py} | 0 .../ltx2_rotary_triton.py} | 0 .../qknorm_rope_jit.py} | 0 .../rope_rotate_half_bitexact.py | 2 +- .../rotary.py => rope/rotary_triton.py} | 22 +- .../kernels/ops/diffusion/sites/__init__.py | 14 + .../diffusion/{ => sites}/bitexact_gate.py | 0 .../fused_gate_rmsnorm_site.py} | 8 +- .../fused_linear_gelu_site.py} | 8 +- .../fused_ln_modulate_site.py} | 4 +- .../hunyuan_qknorm_site.py} | 2 +- .../ltx2_rmsnorm_modulate_site.py} | 6 +- .../ops/diffusion/{ => sites}/quality_gate.py | 0 python/sglang/kernels/spec.py | 1 + .../existing-fast-paths.md | 54 +- .../attention/backends/sparse_linear_attn.py | 5 +- .../runtime/layers/attention/layer.py | 2 +- .../runtime/layers/elementwise.py | 2 +- .../runtime/layers/fused_scale_shift_gate.py | 2 +- .../runtime/layers/layernorm.py | 20 +- .../runtime/layers/parallel_conv.py | 8 +- .../runtime/layers/rotary_embedding/utils.py | 2 +- .../multimodal_gen/runtime/layers/usp.py | 5 +- .../runtime/layers/visual_embedding.py | 4 +- .../runtime/models/dits/cosmos3video.py | 2 +- .../runtime/models/dits/ernie_image.py | 16 +- .../runtime/models/dits/flux.py | 32 +- .../runtime/models/dits/flux_2.py | 10 +- .../runtime/models/dits/glm_image.py | 20 +- .../runtime/models/dits/hunyuanvideo.py | 16 +- .../runtime/models/dits/ideogram.py | 18 +- .../runtime/models/dits/krea2.py | 12 +- .../runtime/models/dits/ltx_2.py | 36 +- .../runtime/models/dits/minimax_h3.py | 4 +- .../runtime/models/dits/qwen_image.py | 6 +- .../runtime/models/dits/sana.py | 14 +- .../runtime/models/dits/sana_wm_components.py | 6 +- .../runtime/models/dits/wanvideo.py | 6 +- .../runtime/models/dits/zimage.py | 10 +- .../models/upsampler/latent_upsampler.py | 2 +- .../runtime/models/vaes/flux2_vae_cuda_opt.py | 24 +- .../runtime/models/vaes/hunyuanvae.py | 2 +- .../vaes/minimax_h3_video_vae/base_module.py | 4 +- .../runtime/models/vaes/wan_vae_cuda_opt.py | 10 +- .../runtime/models/vaes/wanvae.py | 5 +- .../pipelines_core/stages/denoising.py | 18 +- .../runtime/utils/mesh3d_utils.py | 6 +- .../test_latent_upsampler_group_norm_silu.py | 4 +- .../unit/test_layernorm_cutedsl_dispatch.py | 2 +- .../test/unit/test_minimax_h3_dit_contract.py | 4 +- .../diffusion/bench_causal_conv3d_cat_pad.py | 8 +- .../diffusion/bench_group_norm_silu.py | 2 +- .../diffusion/bench_ltx2_qknorm_split_rope.py | 4 +- .../benchmark/diffusion/bench_norm_impls.py | 7 +- .../benchmark/diffusion/bench_qknorm_rope.py | 2 +- .../diffusion/bench_qwen_image_modulation.py | 4 +- .../diffusion/bench_residual_gate_add.py | 3 +- .../diffusion/bench_timestep_embedding.py | 70 ++ .../kernels/ops/diffusion/test_activation.py | 112 +++ .../diffusion/test_autoencoder_kl_fastpath.py | 62 -- .../ops/diffusion/test_bitexact_gate.py | 159 ---- .../diffusion/test_causal_conv3d_cat_pad.py | 89 --- .../diffusion/test_ernie_norm_scale_shift.py | 135 ---- .../ops/diffusion/test_flux2_eager_fusions.py | 101 --- .../ops/diffusion/test_flux2_vae_fastpath.py | 64 -- .../ops/diffusion/test_flux_ln_modulate.py | 75 -- .../ops/diffusion/test_flydsl_fused_norm.py | 133 ---- .../ops/diffusion/test_fused_gate_rmsnorm.py | 54 -- .../ops/diffusion/test_fused_linear_gelu.py | 86 --- .../ops/diffusion/test_fused_ln_modulate.py | 82 -- .../diffusion/test_fused_norm_scale_shift.py | 251 ------- .../diffusion/test_glm_image_ln_modulate.py | 55 -- .../ops/diffusion/test_group_norm_silu.py | 104 --- .../test_hunyuanvideo_eager_fusions.py | 94 --- .../ops/diffusion/test_import_surface.py | 210 ++++++ .../kernels/ops/diffusion/test_layout.py | 582 +++++++++++++++ .../ops/diffusion/test_ltx2_ada_values.py | 92 --- .../diffusion/test_ltx2_qknorm_split_rope.py | 221 ------ .../diffusion/test_ltx2_rms_norm_modulate.py | 88 --- .../ops/diffusion/test_model_fast_paths.py | 703 ++++++++++++++++++ .../kernels/ops/diffusion/test_modulate.py | 424 +++++++++++ .../diffusion/test_modulate_scale_shift.py | 58 -- .../ops/diffusion/test_native_bf16_rmsnorm.py | 87 --- .../kernels/ops/diffusion/test_norm.py | 527 +++++++++++++ .../ops/diffusion/test_quality_gate.py | 47 -- .../diffusion/test_qwen_image_modulation.py | 228 ------ .../ops/diffusion/test_residual_gate_add.py | 123 --- .../{test_qknorm_rope.py => test_rope.py} | 263 ++++++- .../ops/diffusion/test_sana_conv_post.py | 60 -- .../ops/diffusion/test_sana_ln_modulate.py | 55 -- .../kernels/ops/diffusion/test_scale_shift.py | 40 - .../kernels/ops/diffusion/test_sites.py | 416 +++++++++++ .../ops/diffusion/test_timestep_embedding.py | 184 ----- .../kernels/ops/diffusion/test_ulysses_qkv.py | 54 -- .../ops/diffusion/test_usp_relayout.py | 69 -- .../ops/diffusion/test_varlen_pack_pad.py | 197 ----- .../test_varlen_uspattn_equivalence.py | 157 ---- .../ops/diffusion/test_wan_causal_cache.py | 164 ---- .../ops/diffusion/test_wan_vae_fastpath.py | 79 -- .../test_diffusion_modelopt_fp8_scaled_mm.py | 0 .../test_diffusion_nvfp4_scaled_mm.py | 0 .../test_diffusion_stage_profiler_sync.py} | 0 166 files changed, 4911 insertions(+), 4539 deletions(-) create mode 100644 python/sglang/kernels/ops/diffusion/README.md create mode 100644 python/sglang/kernels/ops/diffusion/activation/__init__.py rename python/sglang/kernels/ops/diffusion/{triton/sana_conv_post.py => activation/sana_conv_post_triton.py} (98%) rename python/sglang/kernels/ops/diffusion/{triton => activation}/silu_mul_bitexact.py (98%) create mode 100644 python/sglang/kernels/ops/diffusion/attention/__init__.py rename python/sglang/kernels/ops/diffusion/{triton/sana_wm_gdn_chunkwise.py => attention/sana_wm_gdn_chunkwise_triton.py} (99%) rename python/sglang/kernels/ops/diffusion/{triton/sana_wm_gdn.py => attention/sana_wm_gdn_triton.py} (98%) rename python/sglang/kernels/ops/diffusion/{sparse_linear_attn_kernels.py => attention/sparse_linear_attn_triton.py} (100%) create mode 100644 python/sglang/kernels/ops/diffusion/common/__init__.py rename python/sglang/kernels/ops/diffusion/{triton/mps_fallback.py => common/fallback_mps.py} (96%) rename python/sglang/kernels/ops/diffusion/{triton/npu_fallback.py => common/fallback_npu.py} (100%) rename python/sglang/kernels/ops/diffusion/{triton/torch_fallback.py => common/fallback_torch.py} (100%) rename python/sglang/kernels/ops/diffusion/{triton => common}/numerics.py (100%) create mode 100644 python/sglang/kernels/ops/diffusion/common/platform.py create mode 100644 python/sglang/kernels/ops/diffusion/ext/__init__.py rename python/sglang/kernels/ops/diffusion/{render => ext}/hunyuan3d_rasterizer/__init__.py (96%) create mode 100644 python/sglang/kernels/ops/diffusion/ext/hunyuan3d_rasterizer/rasterizer.cpp create mode 100644 python/sglang/kernels/ops/diffusion/ext/hunyuan3d_rasterizer/rasterizer.h create mode 100644 python/sglang/kernels/ops/diffusion/ext/hunyuan3d_rasterizer/rasterizer_gpu.cu rename python/sglang/kernels/ops/diffusion/{render/__init__.py => ext/loader.py} (97%) rename python/sglang/kernels/ops/diffusion/{render => ext}/mesh_processor/__init__.py (95%) create mode 100644 python/sglang/kernels/ops/diffusion/ext/mesh_processor/mesh_processor.cpp create mode 100644 python/sglang/kernels/ops/diffusion/layout/__init__.py rename python/sglang/kernels/ops/diffusion/{causal_conv3d_cat_pad.py => layout/causal_conv3d_cat_pad_jit.py} (100%) rename python/sglang/kernels/ops/diffusion/{triton/causal_conv3d_pad.py => layout/causal_conv3d_cat_pad_triton.py} (100%) rename python/sglang/kernels/ops/diffusion/{triton/ulysses_qkv.py => layout/ulysses_qkv_triton.py} (100%) rename python/sglang/kernels/ops/diffusion/{usp_relayout.py => layout/usp_relayout_jit.py} (100%) rename python/sglang/kernels/ops/diffusion/{triton/varlen_pack_pad.py => layout/varlen_pack_pad_triton.py} (100%) rename python/sglang/kernels/ops/diffusion/{triton/wan_causal_cache.py => layout/wan_causal_cache_triton.py} (100%) create mode 100644 python/sglang/kernels/ops/diffusion/modulate/__init__.py rename python/sglang/kernels/ops/diffusion/{triton/indexed_modulation.py => modulate/indexed_modulation_triton.py} (98%) rename python/sglang/kernels/ops/diffusion/{triton/ltx2_ada_values.py => modulate/ltx2_ada_values_triton.py} (100%) rename python/sglang/kernels/ops/diffusion/{modulate_scale_shift.py => modulate/modulate_scale_shift_jit.py} (100%) rename python/sglang/kernels/ops/diffusion/{residual_gate_add.py => modulate/residual_gate_add_jit.py} (100%) rename python/sglang/kernels/ops/diffusion/{triton/scale_shift.py => modulate/scale_shift_triton.py} (96%) rename python/sglang/kernels/ops/diffusion/{timestep_embedding.py => modulate/timestep_embedding_jit.py} (100%) rename python/sglang/kernels/ops/diffusion/{triton/wan_temb_table_slices.py => modulate/wan_temb_table_slices_triton.py} (100%) create mode 100644 python/sglang/kernels/ops/diffusion/norm/__init__.py create mode 100644 python/sglang/kernels/ops/diffusion/norm/cutedsl_common/__init__.py rename python/sglang/kernels/ops/diffusion/{cutedsl/common => norm/cutedsl_common}/norm_fusion.py (99%) rename python/sglang/kernels/ops/diffusion/{cutedsl/common => norm/cutedsl_common}/reduce.py (100%) rename python/sglang/kernels/ops/diffusion/{cutedsl/utils.py => norm/cutedsl_utils.py} (100%) rename python/sglang/kernels/ops/diffusion/{flydsl/fused_residual_norm.py => norm/fused_residual_norm_flydsl.py} (100%) rename python/sglang/kernels/ops/diffusion/{ => norm}/group_norm_silu.py (90%) rename python/sglang/kernels/ops/diffusion/{triton/group_norm_silu.py => norm/group_norm_silu_triton.py} (100%) rename python/sglang/kernels/ops/diffusion/{triton/group_norm_silu_twopass.py => norm/group_norm_silu_twopass_triton.py} (84%) rename python/sglang/kernels/ops/diffusion/{triton/layernorm_modulate.py => norm/layernorm_modulate_triton.py} (99%) rename python/sglang/kernels/ops/diffusion/{triton/native_bf16_rmsnorm.py => norm/native_bf16_rmsnorm_triton.py} (100%) rename python/sglang/kernels/ops/diffusion/{norm_scale_shift_native.py => norm/norm_scale_shift_jit.py} (100%) rename python/sglang/kernels/ops/diffusion/{triton/norm.py => norm/norm_triton.py} (97%) rename python/sglang/kernels/ops/diffusion/{triton/rmsnorm_onepass.py => norm/rmsnorm_onepass_triton.py} (81%) rename python/sglang/kernels/ops/diffusion/{triton => norm}/rmsnorm_scale_shift_bitexact.py (99%) rename python/sglang/kernels/ops/diffusion/{cutedsl/scale_residual_norm_scale_shift.py => norm/scale_residual_norm_cutedsl.py} (98%) rename python/sglang/kernels/ops/diffusion/{triton/wan_rmsnorm_silu.py => norm/wan_rmsnorm_silu_triton.py} (95%) rename python/sglang/kernels/ops/diffusion/{triton/zimage_native_norm.py => norm/zimage_qk_rmsnorm_triton.py} (100%) delete mode 100644 python/sglang/kernels/ops/diffusion/render/hunyuan3d_rasterizer/rasterizer.cpp delete mode 100644 python/sglang/kernels/ops/diffusion/render/hunyuan3d_rasterizer/rasterizer.h delete mode 100644 python/sglang/kernels/ops/diffusion/render/hunyuan3d_rasterizer/rasterizer_gpu.cu delete mode 100644 python/sglang/kernels/ops/diffusion/render/mesh_processor/mesh_processor.cpp create mode 100644 python/sglang/kernels/ops/diffusion/rope/__init__.py rename python/sglang/kernels/ops/diffusion/{triton/hunyuan_qkv_pack.py => rope/hunyuan_qkv_pack_triton.py} (100%) rename python/sglang/kernels/ops/diffusion/{ltx2_qknorm_split_rope.py => rope/ltx2_qknorm_split_rope_jit.py} (100%) rename python/sglang/kernels/ops/diffusion/{triton/ltx2_rotary.py => rope/ltx2_rotary_triton.py} (100%) rename python/sglang/kernels/ops/diffusion/{qknorm_rope.py => rope/qknorm_rope_jit.py} (100%) rename python/sglang/kernels/ops/diffusion/{triton => rope}/rope_rotate_half_bitexact.py (98%) rename python/sglang/kernels/ops/diffusion/{triton/rotary.py => rope/rotary_triton.py} (88%) create mode 100644 python/sglang/kernels/ops/diffusion/sites/__init__.py rename python/sglang/kernels/ops/diffusion/{ => sites}/bitexact_gate.py (100%) rename python/sglang/kernels/ops/diffusion/{fused_gate_rmsnorm.py => sites/fused_gate_rmsnorm_site.py} (93%) rename python/sglang/kernels/ops/diffusion/{fused_linear_gelu.py => sites/fused_linear_gelu_site.py} (96%) rename python/sglang/kernels/ops/diffusion/{fused_ln_modulate.py => sites/fused_ln_modulate_site.py} (95%) rename python/sglang/kernels/ops/diffusion/{hunyuan_qknorm.py => sites/hunyuan_qknorm_site.py} (96%) rename python/sglang/kernels/ops/diffusion/{ltx2_rmsnorm_modulate.py => sites/ltx2_rmsnorm_modulate_site.py} (93%) rename python/sglang/kernels/ops/diffusion/{ => sites}/quality_gate.py (100%) create mode 100644 test/registered/kernels/benchmark/diffusion/bench_timestep_embedding.py create mode 100644 test/registered/kernels/ops/diffusion/test_activation.py delete mode 100644 test/registered/kernels/ops/diffusion/test_autoencoder_kl_fastpath.py delete mode 100644 test/registered/kernels/ops/diffusion/test_bitexact_gate.py delete mode 100644 test/registered/kernels/ops/diffusion/test_causal_conv3d_cat_pad.py delete mode 100644 test/registered/kernels/ops/diffusion/test_ernie_norm_scale_shift.py delete mode 100644 test/registered/kernels/ops/diffusion/test_flux2_eager_fusions.py delete mode 100644 test/registered/kernels/ops/diffusion/test_flux2_vae_fastpath.py delete mode 100644 test/registered/kernels/ops/diffusion/test_flux_ln_modulate.py delete mode 100644 test/registered/kernels/ops/diffusion/test_flydsl_fused_norm.py delete mode 100644 test/registered/kernels/ops/diffusion/test_fused_gate_rmsnorm.py delete mode 100644 test/registered/kernels/ops/diffusion/test_fused_linear_gelu.py delete mode 100644 test/registered/kernels/ops/diffusion/test_fused_ln_modulate.py delete mode 100644 test/registered/kernels/ops/diffusion/test_fused_norm_scale_shift.py delete mode 100644 test/registered/kernels/ops/diffusion/test_glm_image_ln_modulate.py delete mode 100644 test/registered/kernels/ops/diffusion/test_group_norm_silu.py delete mode 100644 test/registered/kernels/ops/diffusion/test_hunyuanvideo_eager_fusions.py create mode 100644 test/registered/kernels/ops/diffusion/test_import_surface.py create mode 100644 test/registered/kernels/ops/diffusion/test_layout.py delete mode 100644 test/registered/kernels/ops/diffusion/test_ltx2_ada_values.py delete mode 100644 test/registered/kernels/ops/diffusion/test_ltx2_qknorm_split_rope.py delete mode 100644 test/registered/kernels/ops/diffusion/test_ltx2_rms_norm_modulate.py create mode 100644 test/registered/kernels/ops/diffusion/test_model_fast_paths.py create mode 100644 test/registered/kernels/ops/diffusion/test_modulate.py delete mode 100644 test/registered/kernels/ops/diffusion/test_modulate_scale_shift.py delete mode 100644 test/registered/kernels/ops/diffusion/test_native_bf16_rmsnorm.py create mode 100644 test/registered/kernels/ops/diffusion/test_norm.py delete mode 100644 test/registered/kernels/ops/diffusion/test_quality_gate.py delete mode 100644 test/registered/kernels/ops/diffusion/test_qwen_image_modulation.py delete mode 100644 test/registered/kernels/ops/diffusion/test_residual_gate_add.py rename test/registered/kernels/ops/diffusion/{test_qknorm_rope.py => test_rope.py} (66%) delete mode 100644 test/registered/kernels/ops/diffusion/test_sana_conv_post.py delete mode 100644 test/registered/kernels/ops/diffusion/test_sana_ln_modulate.py delete mode 100644 test/registered/kernels/ops/diffusion/test_scale_shift.py create mode 100644 test/registered/kernels/ops/diffusion/test_sites.py delete mode 100644 test/registered/kernels/ops/diffusion/test_timestep_embedding.py delete mode 100644 test/registered/kernels/ops/diffusion/test_ulysses_qkv.py delete mode 100644 test/registered/kernels/ops/diffusion/test_usp_relayout.py delete mode 100644 test/registered/kernels/ops/diffusion/test_varlen_pack_pad.py delete mode 100644 test/registered/kernels/ops/diffusion/test_varlen_uspattn_equivalence.py delete mode 100644 test/registered/kernels/ops/diffusion/test_wan_causal_cache.py delete mode 100644 test/registered/kernels/ops/diffusion/test_wan_vae_fastpath.py rename test/registered/kernels/ops/{diffusion => quantization}/test_diffusion_modelopt_fp8_scaled_mm.py (100%) rename test/registered/kernels/ops/{diffusion => quantization}/test_diffusion_nvfp4_scaled_mm.py (100%) rename test/registered/{kernels/ops/diffusion/test_stage_profiler_sync.py => profiling/test_diffusion_stage_profiler_sync.py} (100%) diff --git a/.claude/skills/llm-torch-profiler-analysis/references/fuse-overlap-catalog.md b/.claude/skills/llm-torch-profiler-analysis/references/fuse-overlap-catalog.md index 6127101805c5..ef953174d8ef 100644 --- a/.claude/skills/llm-torch-profiler-analysis/references/fuse-overlap-catalog.md +++ b/.claude/skills/llm-torch-profiler-analysis/references/fuse-overlap-catalog.md @@ -111,15 +111,15 @@ in-flight row as shipped. | Pattern | Trace keywords | Primary code | Existing path | Skill should conclude | | --- | --- | --- | --- | --- | -| Fused residual + norm + scale + shift | residual add, norm, scale, shift, gate around DiT blocks | `python/sglang/kernels/ops/diffusion/cutedsl/scale_residual_norm_scale_shift.py`
`python/sglang/multimodal_gen/runtime/layers/layernorm.py` | `fused_scale_residual_norm_scale_shift(...)` | Treat split residual + norm + modulation as a missing existing diffusion fusion first. | -| Fused norm + scale + shift | norm followed by scale / shift elementwise kernels | `python/sglang/kernels/ops/diffusion/cutedsl/scale_residual_norm_scale_shift.py`
`python/sglang/multimodal_gen/runtime/layers/layernorm.py` | `fused_norm_scale_shift(...)` | Existing modulation fusion already covers this family. | -| Triton scale / shift and gate-select kernels | tiny scale / shift or gate-select kernels dominate modulation blocks | `python/sglang/kernels/ops/diffusion/triton/scale_shift.py`
`python/sglang/multimodal_gen/runtime/layers/elementwise.py` | `fuse_scale_shift_kernel(...)` and `fuse_layernorm_scale_shift_gate_select01_kernel(...)` | Check whether the runtime is missing these existing Triton fusions. | -| Fused add-RMSNorm and one-pass RMSNorm | residual add plus RMSNorm still split on short hidden sizes | `python/sglang/multimodal_gen/runtime/layers/layernorm.py`
`python/sglang/kernels/ops/diffusion/triton/rmsnorm_onepass.py` | `fused_add_rmsnorm(...)` and `triton_one_pass_rms_norm(...)` | For short hidden-size diffusion blocks, this is already an established fusion family. | -| Fused diffusion QK norm + RoPE | split QK norm and RoPE in diffusion attention blocks | `python/sglang/kernels/ops/diffusion/qknorm_rope.py`
`python/sglang/multimodal_gen/runtime/layers/layernorm.py::apply_qk_norm_rope` | `fused_inplace_qknorm_rope(...)`, with fallback to QK norm plus `apply_flashinfer_rope_qk_inplace(...)` | Distinguish between missing fused qknorm + rope and the existing FlashInfer RoPE fallback. | +| Fused residual + norm + scale + shift | residual add, norm, scale, shift, gate around DiT blocks | `python/sglang/kernels/ops/diffusion/norm/scale_residual_norm_cutedsl.py`
`python/sglang/multimodal_gen/runtime/layers/layernorm.py` | `fused_scale_residual_norm_scale_shift(...)` | Treat split residual + norm + modulation as a missing existing diffusion fusion first. | +| Fused norm + scale + shift | norm followed by scale / shift elementwise kernels | `python/sglang/kernels/ops/diffusion/norm/scale_residual_norm_cutedsl.py`
`python/sglang/multimodal_gen/runtime/layers/layernorm.py` | `fused_norm_scale_shift(...)` | Existing modulation fusion already covers this family. | +| Triton scale / shift and gate-select kernels | tiny scale / shift or gate-select kernels dominate modulation blocks | `python/sglang/kernels/ops/diffusion/modulate/scale_shift_triton.py`
`python/sglang/multimodal_gen/runtime/layers/elementwise.py` | `fuse_scale_shift_kernel(...)` and `fuse_layernorm_scale_shift_gate_select01_kernel(...)` | Check whether the runtime is missing these existing Triton fusions. | +| Fused add-RMSNorm and one-pass RMSNorm | residual add plus RMSNorm still split on short hidden sizes | `python/sglang/multimodal_gen/runtime/layers/layernorm.py`
`python/sglang/kernels/ops/diffusion/norm/rmsnorm_onepass_triton.py` | `fused_add_rmsnorm(...)` and `triton_one_pass_rms_norm(...)` | For short hidden-size diffusion blocks, this is already an established fusion family. | +| Fused diffusion QK norm + RoPE | split QK norm and RoPE in diffusion attention blocks | `python/sglang/kernels/ops/diffusion/rope/qknorm_rope_jit.py`
`python/sglang/multimodal_gen/runtime/layers/layernorm.py::apply_qk_norm_rope` | `fused_inplace_qknorm_rope(...)`, with fallback to QK norm plus `apply_flashinfer_rope_qk_inplace(...)` | Distinguish between missing fused qknorm + rope and the existing FlashInfer RoPE fallback. | | Z-Image fused `norm(x) * tanh(scale) + shift` | `fused_norm_tanh_mul_add`
`tanh(gate) * rmsnorm(x)` | `python/sglang/kernels/ops/diffusion/cutedsl/norm_tanh_mul_add_norm_scale.py`
`python/sglang/multimodal_gen/runtime/layers/layernorm.py` | CuTeDSL kernel plus runtime helper for Z-Image residual-form modulation | Treat split Z-Image residual-form modulation as a missing existing diffusion fusion, not a novel idea. | | Z-Image fused residual modulation + next norm-scale | `fused_norm_tanh_mul_add_norm_scale`
`residual + tanh(gate) * rmsnorm(x)`
`ffn_norm1(x) * scale_mlp` | `python/sglang/kernels/ops/diffusion/cutedsl/norm_tanh_mul_add_norm_scale.py`
`python/sglang/multimodal_gen/runtime/models/dits/zimage.py` | One CuTeDSL kernel fuses the first residual-form modulation and the next normalization / scale stage | If you see this chain split in Z-Image traces, report it as a missing existing mainline fusion family. | -| LTX2 fused Ada values | `ltx2_ada_values9`
`get_ada_values`
`scale_shift_table + timestep.reshape` | `python/sglang/kernels/ops/diffusion/triton/ltx2_ada_values.py`
`python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py` | PR `#29390` fuses LTX-2.3 Ada value materialization for video/audio streams and reuses the 9 Ada tensors across self-attention, MLP, and prompt-cross-attention blocks | Treat repeated Ada add/reshape/slice ladders in LTX2 traces as a missing shipped SGLang fusion first. | -| LTX2 residual-gate add | `diffusion_residual_gate_add`
`residual_gate_add`
`residual + update * gate` | `python/sglang/kernels/ops/diffusion/residual_gate_add.py`
`python/sglang/kernels/jit/csrc/diffusion/residual_gate_add.cuh`
`python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py` | PR `#29361` fuses LTX2 `residual + update * gate` sites for attention, cross-attention, and feed-forward updates into one CUDA custom op when dtype, shape, device, and contiguity guards pass | Treat split add/mul gate ladders in LTX2 traces as a missing shipped SGLang fusion first. | +| LTX2 fused Ada values | `ltx2_ada_values9`
`get_ada_values`
`scale_shift_table + timestep.reshape` | `python/sglang/kernels/ops/diffusion/modulate/ltx2_ada_values_triton.py`
`python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py` | PR `#29390` fuses LTX-2.3 Ada value materialization for video/audio streams and reuses the 9 Ada tensors across self-attention, MLP, and prompt-cross-attention blocks | Treat repeated Ada add/reshape/slice ladders in LTX2 traces as a missing shipped SGLang fusion first. | +| LTX2 residual-gate add | `diffusion_residual_gate_add`
`residual_gate_add`
`residual + update * gate` | `python/sglang/kernels/ops/diffusion/modulate/residual_gate_add_jit.py`
`python/sglang/kernels/jit/csrc/diffusion/residual_gate_add.cuh`
`python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py` | PR `#29361` fuses LTX2 `residual + update * gate` sites for attention, cross-attention, and feed-forward updates into one CUDA custom op when dtype, shape, device, and contiguity guards pass | Treat split add/mul gate ladders in LTX2 traces as a missing shipped SGLang fusion first. | | Nunchaku fused GELU MLP | `_fused_gelu_mlp`
`fused_gelu_mlp` | `python/sglang/multimodal_gen/runtime/models/dits/flux.py` | Nunchaku path fuses `fc1 GEMM + GELU + shift + re-quant + fc2.lora_down` before the second GEMM | Treat split GELU-MLP on Nunchaku checkpoints as an existing fused family, not a new discovery. | ## 5. Diffusion kernel-overlap and async-communication families diff --git a/python/sglang/kernels/fused_op.py b/python/sglang/kernels/fused_op.py index 441bd3107632..b9cd0fad4af7 100644 --- a/python/sglang/kernels/fused_op.py +++ b/python/sglang/kernels/fused_op.py @@ -107,6 +107,7 @@ KernelBackend.JIT: "forward_jit", KernelBackend.AOT: "forward_aot", KernelBackend.CUTE_DSL: "forward_cute_dsl", + KernelBackend.FLYDSL: "forward_flydsl", KernelBackend.FLASHINFER: "forward_flashinfer", KernelBackend.DEEPGEMM: "forward_deepgemm", KernelBackend.AITER: "forward_aiter", diff --git a/python/sglang/kernels/ops/diffusion/README.md b/python/sglang/kernels/ops/diffusion/README.md new file mode 100644 index 000000000000..c4815be15864 --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/README.md @@ -0,0 +1,153 @@ +# `sglang.kernels.ops.diffusion` + +Fused kernels for diffusion (multimodal-generation) models — DiT transformer +blocks, VAE encoders/decoders, and the sequence-parallel plumbing around them. + +Unlike the LLM operator groups, almost nothing here is a general-purpose +operator. Each kernel replaces a **specific eager op chain in a specific +model**, and its value comes as much from *which rounding boundaries it +reproduces* as from its bandwidth. Multi-step denoising amplifies a per-step +rounding difference into visible quality loss, so "close enough" is a +different product from "bit-exact", and the two are gated differently. + +## Import surface + +```python +from sglang.kernels.ops.diffusion import fused_rmsnorm_scale_shift_bitexact +``` + +**Import from the package, never from a submodule.** The internal layout is +free to move; the facade is not. `test_import_surface.py` enforces this, with +a small allowlist for tests that deliberately exercise one backend. + +Resolution is lazy (PEP 562): the backends have disjoint heavy dependencies +(Triton, CUTLASS/CuTe-DSL, FlyDSL on ROCm, MLX on Apple), so an eager +re-export would make all of them import-time requirements everywhere. + +## Layout + +One subpackage per **operator domain**; the backend is a **filename suffix** +(`_triton`, `_jit`, `_cutedsl`, `_flydsl`, or `_bitexact` where that says more). +This matches `ops/attention` and `ops/gemm`, and it keeps every implementation +of one logical op in one directory. + +``` +norm/ RMSNorm / LayerNorm / GroupNorm and their fused epilogues +modulate/ adaLN modulate, gating, timestep conditioning +rope/ rotary embeddings and the QK-norm chains fused into them +activation/ SiLU / GLU / GELU fusions +attention/ sparse linear attention, gated delta-net +layout/ pure data movement: USP/Ulysses relayout, varlen pack, causal pad +common/ numerics primitives, platform predicates, non-Triton fallbacks +sites/ request-scoped mount policy — NOT kernels (see below) +ext/ JIT C++/CUDA extensions (Hunyuan3D raster/inpaint) — NOT kernels +``` + +## The two numerical contracts + +**Bit-exact (`torch.equal` vs the eager chain) → mounted unconditionally.** +These kernels reproduce every aten rounding boundary, sometimes down to the +reduction tree: `norm/layernorm_modulate_triton.py` replicates torch 2.11's +`vectorized_layer_norm_kernel` (128-thread Welford, `_rcp4` guarded +reciprocal, `shfl.down` fold order, `div.rn` + `MUFU.RSQ`), and +`norm/rmsnorm_scale_shift_bitexact.py` replicates flashinfer's CuTe-DSL +`RMSNormKernel` fragment order and `shfl.bfly` fold. They still verify +themselves against the live eager chain on first sight via +`sites/bitexact_gate.py` and fall back permanently on mismatch — the +dispatch they replicate can change under them. + +**Not bit-exact → quality-gated.** Mounted onto marked `nn.Module` sites only +for `quality="high"` requests, at batch boundaries, all-or-nothing per +transformer (`sites/quality_gate.py`). A plain fp32 single-pass norm fusion +looks harmless and is not: on ERNIE-Image it moved the 50-step trajectory to +PSNR 18.83 dB at `quality=high`, which is what motivated the bit-exact +rewrite. + +## Entry-point protocol + +Every public kernel is a **predicate + kernel** pair: + +```python +if can_use_(...): + out = (...) +else: + out = +``` + +The kernel raises on an unsupported input. It does not return `None` — a +silent `None` is too easy to forget to check, and the failure mode is a +wrong-looking image rather than an exception. + +## Selection matrix + +Several norms look interchangeable and are not. Start here. + +### Norm + scale/shift (adaLN) + +| Entry point | Backend | Contract | Applies to | +|---|---|---|---| +| `fused_rmsnorm_scale_shift_bitexact` | Triton | bit-exact vs flashinfer CuTe RMSNorm + aten modulate | bf16, contiguous rows, `H == 64 * threads_per_row` | +| `fused_scale_residual_rmsnorm_scale_shift_bitexact` | Triton | bit-exact, incl. the preceding residual-gate add | as above | +| `fused_layernorm_modulate` | Triton | bit-exact vs aten `vectorized_layer_norm` | bf16, `N % 4 == 0`, 16B-aligned | +| `fused_norm_scale_shift` / `fused_scale_residual_norm_scale_shift` | CuTe-DSL | fp32 statistics, close | fp16/bf16/fp32, LN or RMS, many broadcast modes | +| `flydsl_norm_scale_shift` / `flydsl_fused_residual_norm_scale_shift` | FlyDSL | close | **ROCm gfx950 only** | +| `fuse_layernorm_scale_shift_gate_select01_kernel` | Triton | close | per-token select between two modulation rows (Qwen-Image) | +| `norm_infer` / `rms_norm_fn` | Triton (+torch/NPU/MPS fallbacks) | close | the generic entry point; use when nothing above fits | + +### Norm variants + +| Entry point | Backend | Contract | Applies to | +|---|---|---|---| +| `triton_group_norm_silu` / `apply_group_norm_silu` | Triton | close | NCHW-contiguous, any channels-per-group, always applies SiLU | +| `group_norm_silu_4d` / `group_norm_silu_rows` | Triton | close | **channels_last only**; power-of-two `C <= 2048`; optional SiLU. This is what lets a VAE decoder run channels_last end-to-end with no `nchwToNhwc` | +| `wan_rmsnorm_silu` | Triton | close | `channels_last_3d` 5D, Wan VAE channel-first RMSNorm + SiLU | +| `rmsnorm_scale` / `rmsnorm_tanh_residual` | Triton | bf16-native statistics | Z-Image (matches its own reference exactly), Ideogram 4 (gated) | +| `zimage_qk_rmsnorm_native` | Triton | bit-exact | Z-Image per-head QK RMSNorm | +| `fused_qk_head_layernorm` | Triton | bit-exact | per-head LN on q/k, `dim_head % 4 == 0`, `<= 128` | +| `triton_one_pass_rms_norm` | Triton | close | standalone RMSNorm, one pass | + +### RoPE / QK-norm + +| Entry point | Backend | Contract | +|---|---|---| +| `fused_inplace_qknorm_rope` | JIT CUDA | one bf16 rounding step vs split baseline; `round_norm_before_rope=True` makes it exact | +| `fused_qknorm_rope_pack_kv` | JIT CUDA | as above, also packs prefix K/V | +| `fused_rope_rotate_half_bitexact` | Triton | bit-exact (elementwise only) | +| `ltx2_qknorm_split_rope_cuda` | JIT CUDA | close; **validated on B200** | +| `apply_rotary_embedding` | Triton (+fallbacks) | close; the generic entry point | +| `hunyuan_qkv_rope_pack` | Triton | bit-exact; packs QKV and applies RoPE in one pass | + +### Data movement (all bit-exact by construction) + +`usp_merge_heads`, `pack_qkv_destination_major`, `fused_pack_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_temb_table_slices` is worth knowing about: the eager +`(table + temb.float()).chunk(6, dim=2)` materializes ~8 GB of fp32 at +704p/121f *and* hands six strided slices downstream, whose `.contiguous()` +calls copy each one again. + +## What is not a kernel + +`sites/` rewrites `nn.Module` trees (mark / mount / unmount) and `ext/` builds +C++/CUDA extensions that have no backend dimension and no numerical contract. +They live here because they are diffusion-specific and share this package's +build machinery, but they are deliberately in their own directories: nothing +in `sites/` or `ext/` belongs in an operator domain, and `sites/` is the one +place allowed to reference `multimodal_gen` types (lazily, inside functions) — +inspecting model modules is its whole job. + +## Adding a kernel + +1. Put it in the operator domain it belongs to, with a backend suffix. +2. Export it from `__init__.py` (`_EXPORTS`) and register a `KernelSpec` + (`_SPECS`) — `test_import_surface.py` checks both resolve. +3. Give it a `can_use_*` predicate; raise, don't return `None`. +4. State the numerical contract in the module docstring, including which + shapes it was verified on. +5. If it is not bit-exact, gate it through `sites/`. Do not mount it by + default. +6. Test it in the domain suite (`test/registered/kernels/ops/diffusion/`), and + the model wiring in `test_model_fast_paths.py`. diff --git a/python/sglang/kernels/ops/diffusion/__init__.py b/python/sglang/kernels/ops/diffusion/__init__.py index 02c56db51bb7..2759c5d251e0 100644 --- a/python/sglang/kernels/ops/diffusion/__init__.py +++ b/python/sglang/kernels/ops/diffusion/__init__.py @@ -1,15 +1,34 @@ -"""Registered diffusion-model kernels and their public wrappers. +"""Fused kernels for diffusion (multimodal-generation) models. -Hot paths import concrete implementations from submodules. The package-level -wrappers remain available for backward compatibility. +This module is the **only** supported import surface for these kernels:: + + from sglang.kernels.ops.diffusion import fused_rmsnorm_scale_shift_bitexact + +Importing a submodule directly (``...diffusion.norm.norm_triton``) couples the +caller to the file layout; ``test_import_surface.py`` guards against it. The +one exception is a test that deliberately exercises a single backend. + +Layout -- one subpackage per **operator domain** (``norm``, ``modulate``, +``rope``, ``activation``, ``attention``, ``layout``) with the backend carried +as a filename suffix (``_triton`` / ``_jit`` / ``_cutedsl`` / ``_flydsl``, or +``_bitexact`` where that is the more informative label), matching how +``ops/attention`` and ``ops/gemm`` are organized. ``common`` holds shared +numerics and platform plumbing, ``sites`` the request-scoped mount policy, and +``ext`` the JIT C++/CUDA extensions that are not kernels. Start from +``README.md``: several norms look interchangeable and are not. + +Resolution is lazy (PEP 562). The backends have disjoint, heavy dependencies +-- Triton, CUTLASS/CuTe-DSL, FlyDSL (ROCm), MLX (Apple) -- so an eager +re-export would turn every one of them into a hard import-time requirement on +every platform. ``_EXPORTS`` maps a symbol to its module and the import +happens on first attribute access. """ from __future__ import annotations -from typing import TYPE_CHECKING +from typing import Any from sglang.kernels.registry import register_kernel -from sglang.kernels.selector import get_kernel from sglang.kernels.spec import ( CapabilityRequirement, FormatSignature, @@ -17,106 +36,443 @@ KernelSpec, ) -if TYPE_CHECKING: - import torch - from torch import nn - _CUDA = frozenset({CapabilityRequirement.CUDA}) +_HIP = frozenset({CapabilityRequirement.HIP}) -register_kernel( - KernelSpec( - op="diffusion.apply_group_norm_silu", - backend=KernelBackend.TRITON, - target="sglang.kernels.ops.diffusion.group_norm_silu:apply_group_norm_silu", - capabilities=_CUDA, - format_signature=FormatSignature(description="fused GroupNorm + SiLU"), - description="Fused group-norm + SiLU (Triton).", - ) -) -register_kernel( - KernelSpec( - op="diffusion.residual_gate_add", - backend=KernelBackend.JIT, - target="sglang.kernels.ops.diffusion.residual_gate_add:residual_gate_add", - capabilities=_CUDA, - format_signature=FormatSignature(description="residual + gate * update"), - description="Fused residual gate-add (sglang.kernels.jit).", - ) +# --------------------------------------------------------------------------- +# Registry inventory. Metadata only -- registering imports neither torch nor a +# backend and triggers no JIT build. Ops carrying several backends (e.g. +# ``scale_residual_norm_scale_shift`` in Triton, CuTe-DSL and FlyDSL) are +# inventory: callers name the one they want via ``select_kernel``. +# --------------------------------------------------------------------------- +_SPECS: tuple[tuple[str, KernelBackend, str, frozenset, str], ...] = ( + ( + "diffusion.apply_group_norm_silu", + KernelBackend.TRITON, + "norm.group_norm_silu:apply_group_norm_silu", + _CUDA, + "Fused GroupNorm + SiLU.", + ), + ( + "diffusion.group_norm_silu_4d", + KernelBackend.TRITON, + "norm.group_norm_silu_twopass_triton:group_norm_silu_4d", + _CUDA, + "Channels-last two-pass GroupNorm(+SiLU), 4D.", + ), + ( + "diffusion.group_norm_silu_rows", + KernelBackend.TRITON, + "norm.group_norm_silu_twopass_triton:group_norm_silu_rows", + _CUDA, + "Channels-last two-pass GroupNorm(+SiLU) over (N, L, C) rows.", + ), + ( + "diffusion.wan_rmsnorm_silu", + KernelBackend.TRITON, + "norm.wan_rmsnorm_silu_triton:wan_rmsnorm_silu", + _CUDA, + "Wan VAE channels_last_3d RMSNorm + SiLU.", + ), + ( + "diffusion.rmsnorm_scale_shift", + KernelBackend.TRITON, + "norm.rmsnorm_scale_shift_bitexact:fused_rmsnorm_scale_shift_bitexact", + _CUDA, + "Bit-exact RMSNorm + adaLN scale/shift.", + ), + ( + "diffusion.scale_residual_norm_scale_shift", + KernelBackend.TRITON, + "norm.rmsnorm_scale_shift_bitexact:fused_scale_residual_rmsnorm_scale_shift_bitexact", + _CUDA, + "Bit-exact residual-gate add + RMSNorm + scale/shift.", + ), + ( + "diffusion.scale_residual_norm_scale_shift", + KernelBackend.CUTE_DSL, + "norm.scale_residual_norm_cutedsl:fused_scale_residual_norm_scale_shift", + _CUDA, + "CuTe-DSL residual + norm + scale/shift.", + ), + ( + "diffusion.scale_residual_norm_scale_shift", + KernelBackend.FLYDSL, + "norm.fused_residual_norm_flydsl:flydsl_fused_residual_norm_scale_shift", + _HIP, + "FlyDSL (ROCm gfx950) residual + norm + scale/shift.", + ), + ( + "diffusion.norm_scale_shift", + KernelBackend.CUTE_DSL, + "norm.scale_residual_norm_cutedsl:fused_norm_scale_shift", + _CUDA, + "CuTe-DSL norm + scale/shift.", + ), + ( + "diffusion.norm_scale_shift", + KernelBackend.FLYDSL, + "norm.fused_residual_norm_flydsl:flydsl_norm_scale_shift", + _HIP, + "FlyDSL (ROCm gfx950) norm + scale/shift.", + ), + ( + "diffusion.layernorm_modulate", + KernelBackend.TRITON, + "norm.layernorm_modulate_triton:fused_layernorm_modulate", + _CUDA, + "Bit-exact LayerNorm + adaLN modulate.", + ), + ( + "diffusion.qk_head_layernorm", + KernelBackend.TRITON, + "norm.layernorm_modulate_triton:fused_qk_head_layernorm", + _CUDA, + "Bit-exact per-head LayerNorm for q/k.", + ), + ( + "diffusion.qk_rmsnorm_native", + KernelBackend.TRITON, + "norm.zimage_qk_rmsnorm_triton:zimage_qk_rmsnorm_native", + _CUDA, + "Z-Image bf16-native per-head QK RMSNorm.", + ), + ( + "diffusion.rmsnorm_scale", + KernelBackend.TRITON, + "norm.native_bf16_rmsnorm_triton:rmsnorm_scale", + _CUDA, + "BF16-native RMSNorm * scale.", + ), + ( + "diffusion.rmsnorm_tanh_residual", + KernelBackend.TRITON, + "norm.native_bf16_rmsnorm_triton:rmsnorm_tanh_residual", + _CUDA, + "BF16-native x + tanh(gate) * RMSNorm(y).", + ), + ( + "diffusion.modulate_scale_shift", + KernelBackend.JIT, + "modulate.modulate_scale_shift_jit:modulate_scale_shift", + _CUDA, + "Bit-exact adaLN modulate x * (1 + scale) + shift.", + ), + ( + "diffusion.residual_gate_add", + KernelBackend.JIT, + "modulate.residual_gate_add_jit:residual_gate_add", + _CUDA, + "Fused residual + gate * update.", + ), + ( + "diffusion.timestep_embedding", + KernelBackend.JIT, + "modulate.timestep_embedding_jit:timestep_embedding", + _CUDA, + "Sinusoidal timestep embedding.", + ), + ( + "diffusion.temb_table_slices", + KernelBackend.TRITON, + "modulate.wan_temb_table_slices_triton:fused_temb_table_slices", + _CUDA, + "Contiguous adaLN slices for Wan2.2-TI2V.", + ), + ( + "diffusion.ltx2_ada_values", + KernelBackend.TRITON, + "modulate.ltx2_ada_values_triton:ltx2_ada_values9", + _CUDA, + "LTX-2 nine-way adaLN value split.", + ), + ( + "diffusion.fused_inplace_qknorm_rope", + KernelBackend.JIT, + "rope.qknorm_rope_jit:fused_inplace_qknorm_rope", + _CUDA, + "Fused in-place QK RMS-norm + RoPE.", + ), + ( + "diffusion.ltx2_qknorm_split_rope", + KernelBackend.JIT, + "rope.ltx2_qknorm_split_rope_jit:ltx2_qknorm_split_rope_cuda", + _CUDA, + "LTX-2 QK-norm + split RoPE.", + ), + ( + "diffusion.rope_rotate_half", + KernelBackend.TRITON, + "rope.rope_rotate_half_bitexact:fused_rope_rotate_half_bitexact", + _CUDA, + "Bit-exact rotate-half RoPE.", + ), + ( + "diffusion.hunyuan_qkv_rope_pack", + KernelBackend.TRITON, + "rope.hunyuan_qkv_pack_triton:hunyuan_qkv_rope_pack", + _CUDA, + "HunyuanVideo QKV pack + RoPE.", + ), + ( + "diffusion.silu_mul", + KernelBackend.TRITON, + "activation.silu_mul_bitexact:fused_silu_mul_bitexact", + _CUDA, + "Bit-exact silu(a) * b for split-projection SwiGLU.", + ), + ( + "diffusion.bias_silu", + KernelBackend.TRITON, + "activation.sana_conv_post_triton:fused_bias_silu", + _CUDA, + "Bit-exact conv bias + SiLU (Sana GLUMB).", + ), + ( + "diffusion.bias_glu", + KernelBackend.TRITON, + "activation.sana_conv_post_triton:fused_bias_glu", + _CUDA, + "Bit-exact conv bias + GLU (Sana GLUMB).", + ), + ( + "diffusion.linear_gelu_tanh", + KernelBackend.AOT, + "sites.fused_linear_gelu_site:fused_linear_gelu_tanh", + _CUDA, + "Linear + tanh-GELU via the cublasLt epilogue.", + ), + ( + "diffusion.sparse_linear_attn_fwd", + KernelBackend.TRITON, + "attention.sparse_linear_attn_triton:_attn_fwd", + _CUDA, + "Sparse linear attention forward.", + ), + ( + "diffusion.bigdn", + KernelBackend.TRITON, + "attention.sana_wm_gdn_triton:fused_bigdn_func", + _CUDA, + "Sana-WM bidirectional gated delta-net.", + ), + ( + "diffusion.usp_merge_heads", + KernelBackend.JIT, + "layout.usp_relayout_jit:usp_merge_heads", + _CUDA, + "USP all-to-all output head merge.", + ), + ( + "diffusion.pack_qkv_destination_major", + KernelBackend.TRITON, + "layout.ulysses_qkv_triton:pack_qkv_destination_major", + _CUDA, + "Ulysses destination-major QKV pack.", + ), + ( + "diffusion.varlen_pack_qkv", + KernelBackend.TRITON, + "layout.varlen_pack_pad_triton:fused_pack_qkv", + _CUDA, + "Varlen gather of Q/K/V at valid positions.", + ), + ( + "diffusion.varlen_scatter_to_padded", + KernelBackend.TRITON, + "layout.varlen_pack_pad_triton:fused_scatter_to_padded", + _CUDA, + "Varlen scatter back to the dense layout.", + ), + ( + "diffusion.causal_conv3d_cat_pad", + KernelBackend.JIT, + "layout.causal_conv3d_cat_pad_jit:fused_causal_conv3d_cat_pad_cuda", + _CUDA, + "Causal Conv3d cat + pad.", + ), + ( + "diffusion.causal_conv3d_cat_pad", + KernelBackend.TRITON, + "layout.causal_conv3d_cat_pad_triton:fused_causal_conv3d_cat_pad", + _CUDA, + "Causal Conv3d cat + pad (Triton).", + ), + ( + "diffusion.cat_pad_channels_last_3d", + KernelBackend.TRITON, + "layout.wan_causal_cache_triton:cat_pad_channels_last_3d", + _CUDA, + "Wan causal VAE cat + pad in channels_last_3d.", + ), + ( + "diffusion.dup_up3d_add", + KernelBackend.TRITON, + "layout.wan_causal_cache_triton:dup_up3d_add", + _CUDA, + "Wan causal VAE main + DupUp3D(src).", + ), ) -register_kernel( - KernelSpec( - op="diffusion.fused_inplace_qknorm_rope", - backend=KernelBackend.JIT, - target="sglang.kernels.ops.diffusion.qknorm_rope:fused_inplace_qknorm_rope", - capabilities=_CUDA, - format_signature=FormatSignature( - in_place=True, description="fused in-place QK-norm + RoPE" - ), - description="Fused QK-norm + RoPE (sglang.kernels.jit).", - ) -) -# Migrated from multimodal_gen (RFC #29630, Phase 2.5). Hot paths import the -# Triton symbol directly; the registry entry remains for namespace discovery. -register_kernel( - KernelSpec( - op="diffusion.sparse_linear_attn_fwd", - backend=KernelBackend.TRITON, - target="sglang.kernels.ops.diffusion.sparse_linear_attn_kernels:_attn_fwd", - capabilities=_CUDA, - format_signature=FormatSignature(description="sparse linear attention fwd"), - description="Sparse linear attention forward (Triton).", + +for _op, _backend, _target, _caps, _description in _SPECS: + register_kernel( + KernelSpec( + op=_op, + backend=_backend, + target=f"sglang.kernels.ops.diffusion.{_target}", + capabilities=_caps, + format_signature=FormatSignature(description=_description), + description=_description, + ) ) -) +# --------------------------------------------------------------------------- +# Public export table: symbol -> owning submodule. Sorted by domain, module, +# then symbol; a new public kernel belongs here and nowhere else. +# --------------------------------------------------------------------------- +_EXPORTS: dict[str, str] = { + # Normalization: RMSNorm / LayerNorm / GroupNorm and their fused epilogues + "FLYDSL_NORM_MIN_ALIGNED_DIM": "norm.fused_residual_norm_flydsl", + "flydsl_fused_residual_norm_scale_shift": "norm.fused_residual_norm_flydsl", + "flydsl_norm_scale_shift": "norm.fused_residual_norm_flydsl", + "apply_group_norm_silu": "norm.group_norm_silu", + "triton_group_norm_silu": "norm.group_norm_silu_triton", + "can_use_group_norm_silu_4d": "norm.group_norm_silu_twopass_triton", + "can_use_group_norm_silu_rows": "norm.group_norm_silu_twopass_triton", + "group_norm_silu_4d": "norm.group_norm_silu_twopass_triton", + "group_norm_silu_rows": "norm.group_norm_silu_twopass_triton", + "can_use_fused_layernorm_modulate": "norm.layernorm_modulate_triton", + "can_use_fused_qk_head_layernorm": "norm.layernorm_modulate_triton", + "fused_layernorm_modulate": "norm.layernorm_modulate_triton", + "fused_layernorm_modulate_raw": "norm.layernorm_modulate_triton", + "fused_qk_head_layernorm": "norm.layernorm_modulate_triton", + "is_plain_layer_norm": "norm.layernorm_modulate_triton", + "rmsnorm_scale": "norm.native_bf16_rmsnorm_triton", + "rmsnorm_tanh_residual": "norm.native_bf16_rmsnorm_triton", + "norm_infer": "norm.norm_triton", + "rms_norm_fn": "norm.norm_triton", + "triton_one_pass_rms_norm": "norm.rmsnorm_onepass_triton", + "can_use_fused_rmsnorm_scale_shift": "norm.rmsnorm_scale_shift_bitexact", + "can_use_fused_scale_residual_rmsnorm_scale_shift": "norm.rmsnorm_scale_shift_bitexact", + "fused_rmsnorm_scale_shift_bitexact": "norm.rmsnorm_scale_shift_bitexact", + "fused_scale_residual_rmsnorm_scale_shift_bitexact": "norm.rmsnorm_scale_shift_bitexact", + "fused_norm_scale_shift": "norm.scale_residual_norm_cutedsl", + "fused_scale_residual_norm_scale_shift": "norm.scale_residual_norm_cutedsl", + "validate_scale_shift": "norm.scale_residual_norm_cutedsl", + "can_use_wan_rmsnorm_silu": "norm.wan_rmsnorm_silu_triton", + "wan_rmsnorm_silu": "norm.wan_rmsnorm_silu_triton", + "can_use_qk_rmsnorm_native": "norm.zimage_qk_rmsnorm_triton", + "zimage_qk_rmsnorm_native": "norm.zimage_qk_rmsnorm_triton", + # adaLN modulation, gating and timestep conditioning + "indexed_gate_bf16": "modulate.indexed_modulation_triton", + "indexed_gate_bf16_": "modulate.indexed_modulation_triton", + "indexed_scale_shift_bf16_": "modulate.indexed_modulation_triton", + "ltx2_ada_values9": "modulate.ltx2_ada_values_triton", + "can_use_modulate_scale_shift_cuda": "modulate.modulate_scale_shift_jit", + "modulate_scale_shift": "modulate.modulate_scale_shift_jit", + "modulate_scale_shift_cuda": "modulate.modulate_scale_shift_jit", + "can_use_residual_gate_add_cuda": "modulate.residual_gate_add_jit", + "residual_gate_add": "modulate.residual_gate_add_jit", + "residual_gate_add_cuda": "modulate.residual_gate_add_jit", + "fuse_layernorm_scale_shift_gate_select01_kernel": "modulate.scale_shift_triton", + "fuse_residual_layernorm_scale_shift_gate_select01_kernel": "modulate.scale_shift_triton", + "fuse_scale_shift_kernel": "modulate.scale_shift_triton", + "try_fused_scaled_residual_add_exact": "modulate.scale_shift_triton", + "timestep_embedding": "modulate.timestep_embedding_jit", + "can_use_fused_temb_table_slices": "modulate.wan_temb_table_slices_triton", + "fused_temb_table_slices": "modulate.wan_temb_table_slices_triton", + # Rotary embeddings and the QK-norm chains fused around them + "hunyuan_qkv_rope_pack": "rope.hunyuan_qkv_pack_triton", + "can_use_ltx2_qknorm_split_rope_cuda": "rope.ltx2_qknorm_split_rope_jit", + "ltx2_qknorm_split_rope_cuda": "rope.ltx2_qknorm_split_rope_jit", + "apply_ltx2_split_rotary_emb": "rope.ltx2_rotary_triton", + "can_use_fused_inplace_qknorm_rope": "rope.qknorm_rope_jit", + "fused_inplace_qknorm_rope": "rope.qknorm_rope_jit", + "fused_qknorm_rope_pack_kv": "rope.qknorm_rope_jit", + "can_use_fused_rope_rotate_half": "rope.rope_rotate_half_bitexact", + "fused_rope_rotate_half_bitexact": "rope.rope_rotate_half_bitexact", + "apply_rotary_embedding": "rope.rotary_triton", + # Activation-function fusions + "can_use_fused_bias_glu": "activation.sana_conv_post_triton", + "can_use_fused_bias_silu": "activation.sana_conv_post_triton", + "fused_bias_glu": "activation.sana_conv_post_triton", + "fused_bias_silu": "activation.sana_conv_post_triton", + "can_use_fused_silu_mul": "activation.silu_mul_bitexact", + "fused_packed_silu_mul_bitexact": "activation.silu_mul_bitexact", + "fused_silu_mul_bitexact": "activation.silu_mul_bitexact", + # Diffusion attention kernels + "cam_scan_bidi_chunkwise": "attention.sana_wm_gdn_chunkwise_triton", + "fused_bigdn_func": "attention.sana_wm_gdn_triton", + "fused_qk_inv_rms": "attention.sana_wm_gdn_triton", + "prepare_rope_tables": "attention.sana_wm_gdn_triton", + "_attn_fwd": "attention.sparse_linear_attn_triton", + "get_block_map": "attention.sparse_linear_attn_triton", + # Data movement: bitwise identical to the aten chains they replace + "can_use_fused_causal_conv3d_cat_pad_cuda": "layout.causal_conv3d_cat_pad_jit", + "fused_causal_conv3d_cat_pad_cuda": "layout.causal_conv3d_cat_pad_jit", + "fused_causal_conv3d_cat_pad": "layout.causal_conv3d_cat_pad_triton", + "pack_qkv_destination_major": "layout.ulysses_qkv_triton", + "can_use_usp_merge_heads": "layout.usp_relayout_jit", + "usp_merge_heads": "layout.usp_relayout_jit", + "build_inv_indices": "layout.varlen_pack_pad_triton", + "fused_pack_qkv": "layout.varlen_pack_pad_triton", + "fused_scatter_to_padded": "layout.varlen_pack_pad_triton", + "cat_pad_channels_last_3d": "layout.wan_causal_cache_triton", + "dup_up3d_add": "layout.wan_causal_cache_triton", + # Fusion-site policy: quality gate, first-sight verification, mount + "BitExactFusionGate": "sites.bitexact_gate", + "flashinfer_rmsnorm_diagnostic_hint": "sites.bitexact_gate", + "tensors_equal": "sites.bitexact_gate", + "fused_gate_rmsnorm_active": "sites.fused_gate_rmsnorm_site", + "fused_rmsnorm_scale": "sites.fused_gate_rmsnorm_site", + "fused_rmsnorm_tanh_residual": "sites.fused_gate_rmsnorm_site", + "mark_fused_gate_rmsnorm_site": "sites.fused_gate_rmsnorm_site", + "mount_fused_gate_rmsnorm": "sites.fused_gate_rmsnorm_site", + "unmount_fused_gate_rmsnorm": "sites.fused_gate_rmsnorm_site", + "can_use_linear_gelu": "sites.fused_linear_gelu_site", + "fused_gelu_active": "sites.fused_linear_gelu_site", + "fused_linear_gelu_tanh": "sites.fused_linear_gelu_site", + "mark_fused_gelu_site": "sites.fused_linear_gelu_site", + "mount_fused_linear_gelu": "sites.fused_linear_gelu_site", + "unmount_fused_linear_gelu": "sites.fused_linear_gelu_site", + "can_use_ln_modulate": "sites.fused_ln_modulate_site", + "fused_ln_modulate": "sites.fused_ln_modulate_site", + "fused_ln_modulate_active": "sites.fused_ln_modulate_site", + "mark_fused_ln_modulate_site": "sites.fused_ln_modulate_site", + "mount_fused_ln_modulate": "sites.fused_ln_modulate_site", + "unmount_fused_ln_modulate": "sites.fused_ln_modulate_site", + "mark_hunyuan_qknorm_site": "sites.hunyuan_qknorm_site", + "mount_hunyuan_qknorm": "sites.hunyuan_qknorm_site", + "try_hunyuan_qknorm": "sites.hunyuan_qknorm_site", + "unmount_hunyuan_qknorm": "sites.hunyuan_qknorm_site", + "can_use_ltx2_rms_norm_modulate": "sites.ltx2_rmsnorm_modulate_site", + "fused_ltx2_rms_norm_modulate": "sites.ltx2_rmsnorm_modulate_site", + "ltx2_rms_norm_modulate_active": "sites.ltx2_rmsnorm_modulate_site", + "mark_ltx2_rms_norm_modulate_site": "sites.ltx2_rmsnorm_modulate_site", + "mount_ltx2_rms_norm_modulate": "sites.ltx2_rmsnorm_modulate_site", + "unmount_ltx2_rms_norm_modulate": "sites.ltx2_rmsnorm_modulate_site", + "QualityGatedFusion": "sites.quality_gate", + # JIT C++/CUDA extensions (not kernels, not in the registry) + "interpolate": "ext.hunyuan3d_rasterizer", + "rasterize": "ext.hunyuan3d_rasterizer", + "meshVerticeInpaint": "ext.mesh_processor", +} -def apply_group_norm_silu( - x: torch.Tensor, norm: nn.Module, activation: nn.Module -) -> torch.Tensor: - """Fused GroupNorm + SiLU (falls back to eager when unsupported).""" - return get_kernel("diffusion.apply_group_norm_silu", KernelBackend.TRITON)( - x, norm, activation - ) +def __getattr__(name: str) -> Any: + """Resolve a public symbol to its submodule on first access (PEP 562).""" + module = _EXPORTS.get(name) + if module is None: + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + from importlib import import_module -def residual_gate_add( - residual: torch.Tensor, update: torch.Tensor, gate: torch.Tensor -) -> torch.Tensor: - """Fused ``residual + gate * update``.""" - return get_kernel("diffusion.residual_gate_add", KernelBackend.JIT)( - residual, update, gate - ) + value = getattr(import_module(f"{__name__}.{module}"), name) + globals()[name] = value # cache; later lookups skip __getattr__ entirely + return value -def fused_inplace_qknorm_rope( - q: torch.Tensor, - k: torch.Tensor, - q_weight: torch.Tensor, - k_weight: torch.Tensor, - cos_sin_cache: torch.Tensor, - positions: torch.Tensor, - *, - is_neox: bool, - eps: float = 1e-6, - head_dim: int = 0, - rope_dim: int = 0, -) -> None: - """Fused in-place QK RMS-norm + RoPE.""" - return get_kernel("diffusion.fused_inplace_qknorm_rope", KernelBackend.JIT)( - q, - k, - q_weight, - k_weight, - cos_sin_cache, - positions, - is_neox=is_neox, - eps=eps, - head_dim=head_dim, - rope_dim=rope_dim, - ) +def __dir__() -> list[str]: + return sorted(set(globals()) | set(_EXPORTS)) -__all__ = [ - "apply_group_norm_silu", - "residual_gate_add", - "fused_inplace_qknorm_rope", -] +__all__ = sorted(_EXPORTS) diff --git a/python/sglang/kernels/ops/diffusion/activation/__init__.py b/python/sglang/kernels/ops/diffusion/activation/__init__.py new file mode 100644 index 000000000000..f4010e8df7cc --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/activation/__init__.py @@ -0,0 +1 @@ +"""Activation-function fusions (SiLU/GLU/GELU epilogues).""" diff --git a/python/sglang/kernels/ops/diffusion/triton/sana_conv_post.py b/python/sglang/kernels/ops/diffusion/activation/sana_conv_post_triton.py similarity index 98% rename from python/sglang/kernels/ops/diffusion/triton/sana_conv_post.py rename to python/sglang/kernels/ops/diffusion/activation/sana_conv_post_triton.py index 2ee2b8a7b90c..69d4f4f5ad6e 100644 --- a/python/sglang/kernels/ops/diffusion/triton/sana_conv_post.py +++ b/python/sglang/kernels/ops/diffusion/activation/sana_conv_post_triton.py @@ -7,7 +7,7 @@ import triton # type: ignore import triton.language as tl # type: ignore -from sglang.kernels.ops.diffusion.triton.numerics import round_bf16_to_fp32 +from sglang.kernels.ops.diffusion.common.numerics import round_bf16_to_fp32 @triton.jit diff --git a/python/sglang/kernels/ops/diffusion/triton/silu_mul_bitexact.py b/python/sglang/kernels/ops/diffusion/activation/silu_mul_bitexact.py similarity index 98% rename from python/sglang/kernels/ops/diffusion/triton/silu_mul_bitexact.py rename to python/sglang/kernels/ops/diffusion/activation/silu_mul_bitexact.py index 9163707b334b..94860908363c 100644 --- a/python/sglang/kernels/ops/diffusion/triton/silu_mul_bitexact.py +++ b/python/sglang/kernels/ops/diffusion/activation/silu_mul_bitexact.py @@ -20,7 +20,7 @@ import triton # type: ignore import triton.language as tl # type: ignore -from sglang.kernels.ops.diffusion.triton.numerics import round_bf16_to_fp32 +from sglang.kernels.ops.diffusion.common.numerics import round_bf16_to_fp32 from sglang.srt.utils.custom_op import register_custom_op diff --git a/python/sglang/kernels/ops/diffusion/attention/__init__.py b/python/sglang/kernels/ops/diffusion/attention/__init__.py new file mode 100644 index 000000000000..29ebb5694622 --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/attention/__init__.py @@ -0,0 +1 @@ +"""Diffusion attention kernels: sparse linear attention and gated delta-net.""" diff --git a/python/sglang/kernels/ops/diffusion/triton/sana_wm_gdn_chunkwise.py b/python/sglang/kernels/ops/diffusion/attention/sana_wm_gdn_chunkwise_triton.py similarity index 99% rename from python/sglang/kernels/ops/diffusion/triton/sana_wm_gdn_chunkwise.py rename to python/sglang/kernels/ops/diffusion/attention/sana_wm_gdn_chunkwise_triton.py index 92fd1d104378..cf697d07edfd 100644 --- a/python/sglang/kernels/ops/diffusion/triton/sana_wm_gdn_chunkwise.py +++ b/python/sglang/kernels/ops/diffusion/attention/sana_wm_gdn_chunkwise_triton.py @@ -1619,7 +1619,7 @@ def fused_bigdn_bidi_chunkwise( def _default_dot_prec() -> int: try: - from sglang.kernels.ops.diffusion.triton.sana_wm_gdn import ( + from sglang.kernels.ops.diffusion.attention.sana_wm_gdn_triton import ( _resolve_launch_config, ) diff --git a/python/sglang/kernels/ops/diffusion/triton/sana_wm_gdn.py b/python/sglang/kernels/ops/diffusion/attention/sana_wm_gdn_triton.py similarity index 98% rename from python/sglang/kernels/ops/diffusion/triton/sana_wm_gdn.py rename to python/sglang/kernels/ops/diffusion/attention/sana_wm_gdn_triton.py index 8b8e025a0a5f..07760f61c691 100644 --- a/python/sglang/kernels/ops/diffusion/triton/sana_wm_gdn.py +++ b/python/sglang/kernels/ops/diffusion/attention/sana_wm_gdn_triton.py @@ -220,7 +220,7 @@ def fused_bigdn_func( Thin entry point kept for call-site stability; delegates to :func:`fused_bigdn_bidi_chunkwise` from ``sana_wm_gdn_chunkwise``. """ - from sglang.kernels.ops.diffusion.triton.sana_wm_gdn_chunkwise import ( + from sglang.kernels.ops.diffusion.attention.sana_wm_gdn_chunkwise_triton import ( fused_bigdn_bidi_chunkwise, ) diff --git a/python/sglang/kernels/ops/diffusion/sparse_linear_attn_kernels.py b/python/sglang/kernels/ops/diffusion/attention/sparse_linear_attn_triton.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/sparse_linear_attn_kernels.py rename to python/sglang/kernels/ops/diffusion/attention/sparse_linear_attn_triton.py diff --git a/python/sglang/kernels/ops/diffusion/common/__init__.py b/python/sglang/kernels/ops/diffusion/common/__init__.py new file mode 100644 index 000000000000..3a29cfaaf0ab --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/common/__init__.py @@ -0,0 +1,6 @@ +"""Shared infrastructure for the diffusion kernels -- no kernels of its own. + +- ``numerics`` : rounding/opmath primitives the bit-exact kernels are built from +- ``platform`` : device predicates and the Triton-vs-fallback selector +- ``fallback_*``: pure-torch / NPU / MPS implementations for Triton-less devices +""" diff --git a/python/sglang/kernels/ops/diffusion/triton/mps_fallback.py b/python/sglang/kernels/ops/diffusion/common/fallback_mps.py similarity index 96% rename from python/sglang/kernels/ops/diffusion/triton/mps_fallback.py rename to python/sglang/kernels/ops/diffusion/common/fallback_mps.py index 6c1770099963..31298b434eb7 100644 --- a/python/sglang/kernels/ops/diffusion/triton/mps_fallback.py +++ b/python/sglang/kernels/ops/diffusion/common/fallback_mps.py @@ -2,7 +2,7 @@ Triton is not available on macOS / Metal, so these pure-PyTorch (and optionally MLX-accelerated) implementations replace the Triton kernels -at import time when ``current_platform.is_mps()`` is True. +at import time when the live platform is MPS (see ``common.platform``). MLX acceleration (opt-in via ``SGLANG_USE_MLX=1``): Norm ops use ``mx.fast.rms_norm`` / ``mx.fast.layer_norm`` — single fused @@ -17,13 +17,13 @@ from sglang.srt.utils.tensor_bridge import mlx_to_torch, torch_to_mlx, use_mlx -from .torch_fallback import ( +from .fallback_torch import ( apply_rotary_embedding_native as apply_rotary_embedding_native, ) -from .torch_fallback import ( +from .fallback_torch import ( fuse_scale_shift_kernel_native as fuse_scale_shift_kernel_native, ) -from .torch_fallback import ( +from .fallback_torch import ( norm_infer_native, rms_norm_fn_native, triton_one_pass_rms_norm_native, diff --git a/python/sglang/kernels/ops/diffusion/triton/npu_fallback.py b/python/sglang/kernels/ops/diffusion/common/fallback_npu.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/triton/npu_fallback.py rename to python/sglang/kernels/ops/diffusion/common/fallback_npu.py diff --git a/python/sglang/kernels/ops/diffusion/triton/torch_fallback.py b/python/sglang/kernels/ops/diffusion/common/fallback_torch.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/triton/torch_fallback.py rename to python/sglang/kernels/ops/diffusion/common/fallback_torch.py diff --git a/python/sglang/kernels/ops/diffusion/triton/numerics.py b/python/sglang/kernels/ops/diffusion/common/numerics.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/triton/numerics.py rename to python/sglang/kernels/ops/diffusion/common/numerics.py diff --git a/python/sglang/kernels/ops/diffusion/common/platform.py b/python/sglang/kernels/ops/diffusion/common/platform.py new file mode 100644 index 000000000000..2ac0401208df --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/common/platform.py @@ -0,0 +1,89 @@ +"""Platform predicates and the import-time fallback selector. + +Several diffusion Triton kernels have no Triton on the live device (Ascend +NPU, Apple MPS, MUSA, CPU) and must resolve to a pure-``torch`` — or +MLX-accelerated — implementation. That choice is made once at import time, +which used to mean a hand-rolled four-branch ``if`` block repeated in every +such module, each importing ``current_platform`` directly. + +This module owns both halves of that: + +- :func:`platform_key` — the one place the diffusion kernels ask what device + they are on; +- :func:`select_impl` — the one place the Triton-vs-fallback choice is made. + +Layering note: the authority on "what platform is this" is the +``multimodal_gen`` platform plugin registry, because it is the only one that +consults out-of-tree vendor plugins (NPU/MUSA). ``kernels.spec.PlatformInfo`` +cannot replace it until it grows MPS/MUSA members and plugin support (see the +``DeviceType`` TODO in ``kernels/spec.py``). Until then the dependency is +deliberately confined to this single file and resolved lazily, so no other +kernel module imports upward. +""" + +from __future__ import annotations + +from typing import Callable, TypeVar + +F = TypeVar("F", bound=Callable) + +_CUDA_LIKE = frozenset({"cuda", "hip"}) + + +def platform_key() -> str: + """Return the live device family: ``cuda``/``hip``/``npu``/``mps``/``musa``/``cpu``.""" + from sglang.multimodal_gen.runtime.platforms import current_platform + + for name in ("cuda", "hip", "npu", "mps", "musa"): + if getattr(current_platform, f"is_{name}")(): + return name + return "cpu" + + +def is_cuda() -> bool: + return platform_key() == "cuda" + + +def is_hip() -> bool: + return platform_key() == "hip" + + +def has_triton() -> bool: + """True when the live device runs the Triton implementations.""" + return platform_key() in _CUDA_LIKE + + +def lazy_fallback(kind: str, name: str) -> Callable: + """Name a fallback without importing its module. + + ``select_impl`` is handed every candidate at once, so a plain import here + would pull in *all* fallback modules on every platform -- including MLX on + CUDA hosts. The returned shim imports ``common.fallback_`` on its + first call instead, which for the unselected candidates never happens. + """ + + def _call(*args, **kwargs): + from importlib import import_module + + impl = getattr( + import_module(f"sglang.kernels.ops.diffusion.common.fallback_{kind}"), name + ) + return impl(*args, **kwargs) + + _call.__name__ = name + _call.__qualname__ = f"{kind}_fallback.{name}" + return _call + + +def select_impl(triton_impl: F, **fallbacks: F) -> F: + """Pick ``triton_impl`` on CUDA/HIP, else the fallback for this platform. + + Callers pass the fallbacks they actually have, keyed by platform + (``npu=``, ``mps=``, ``musa=``, ``cpu=``); an unlisted platform keeps the + Triton implementation, which is what the pre-existing per-module ``if`` + chains did. Keeping the whole decision in one call means a module's + exported name is bound exactly once, so the fallback wiring stays greppable + and can later be replaced wholesale by a ``BaseFusedOp`` dispatch without + touching call sites. + """ + return fallbacks.get(platform_key(), triton_impl) diff --git a/python/sglang/kernels/ops/diffusion/ext/__init__.py b/python/sglang/kernels/ops/diffusion/ext/__init__.py new file mode 100644 index 000000000000..d177b930d96e --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/ext/__init__.py @@ -0,0 +1,7 @@ +"""JIT-built C++/CUDA extensions that are not kernels. + +Mesh rasterization and texture inpainting for Hunyuan3D: no backend dimension, +no numerical contract, not in the kernel registry. Kept beside the diffusion +kernels because they share the JIT build/recovery machinery in +:mod:`.loader`. +""" diff --git a/python/sglang/kernels/ops/diffusion/render/hunyuan3d_rasterizer/__init__.py b/python/sglang/kernels/ops/diffusion/ext/hunyuan3d_rasterizer/__init__.py similarity index 96% rename from python/sglang/kernels/ops/diffusion/render/hunyuan3d_rasterizer/__init__.py rename to python/sglang/kernels/ops/diffusion/ext/hunyuan3d_rasterizer/__init__.py index 74ae6c262ea9..3c958908ebe2 100644 --- a/python/sglang/kernels/ops/diffusion/render/hunyuan3d_rasterizer/__init__.py +++ b/python/sglang/kernels/ops/diffusion/ext/hunyuan3d_rasterizer/__init__.py @@ -13,7 +13,7 @@ import torch -from sglang.kernels.ops.diffusion.render import load_extension_with_recovery +from sglang.kernels.ops.diffusion.ext.loader import load_extension_with_recovery _abs_path = os.path.dirname(os.path.abspath(__file__)) _custom_rasterizer_kernel = None diff --git a/python/sglang/kernels/ops/diffusion/ext/hunyuan3d_rasterizer/rasterizer.cpp b/python/sglang/kernels/ops/diffusion/ext/hunyuan3d_rasterizer/rasterizer.cpp new file mode 100644 index 000000000000..47db1d26a4a2 --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/ext/hunyuan3d_rasterizer/rasterizer.cpp @@ -0,0 +1,172 @@ +// SPDX-License-Identifier: Apache-2.0 +// Adapted from Hunyuan3D-2: https://github.com/Tencent/Hunyuan3D-2 +// Original license: TENCENT HUNYUAN NON-COMMERCIAL LICENSE AGREEMENT + +#include "rasterizer.h" + +void rasterizeTriangleCPU(int idx, float *vt0, float *vt1, float *vt2, + int width, int height, INT64 *zbuffer, float *d, + float occlusion_truncation) { + float x_min = std::min(vt0[0], std::min(vt1[0], vt2[0])); + float x_max = std::max(vt0[0], std::max(vt1[0], vt2[0])); + float y_min = std::min(vt0[1], std::min(vt1[1], vt2[1])); + float y_max = std::max(vt0[1], std::max(vt1[1], vt2[1])); + + for (int px = x_min; px < x_max + 1; ++px) { + if (px < 0 || px >= width) + continue; + for (int py = y_min; py < y_max + 1; ++py) { + if (py < 0 || py >= height) + continue; + float vt[2] = {px + 0.5f, py + 0.5f}; + float baryCentricCoordinate[3]; + calculateBarycentricCoordinate(vt0, vt1, vt2, vt, baryCentricCoordinate); + if (isBarycentricCoordInBounds(baryCentricCoordinate)) { + int pixel = py * width + px; + if (zbuffer == 0) { + zbuffer[pixel] = (INT64)(idx + 1); + continue; + } + + float depth = baryCentricCoordinate[0] * vt0[2] + + baryCentricCoordinate[1] * vt1[2] + + baryCentricCoordinate[2] * vt2[2]; + float depth_thres = 0; + if (d) { + depth_thres = d[pixel] * 0.49999f + 0.5f + occlusion_truncation; + } + + int z_quantize = depth * (2 << 17); + INT64 token = (INT64)z_quantize * MAXINT + (INT64)(idx + 1); + if (depth < depth_thres) + continue; + zbuffer[pixel] = std::min(zbuffer[pixel], token); + } + } + } +} + +void barycentricFromImgcoordCPU(float *V, int *F, int *findices, INT64 *zbuffer, + int width, int height, int num_vertices, + int num_faces, float *barycentric_map, + int pix) { + INT64 f = zbuffer[pix] % MAXINT; + if (f == (MAXINT - 1)) { + findices[pix] = 0; + barycentric_map[pix * 3] = 0; + barycentric_map[pix * 3 + 1] = 0; + barycentric_map[pix * 3 + 2] = 0; + return; + } + findices[pix] = f; + f -= 1; + float barycentric[3] = {0, 0, 0}; + if (f >= 0) { + float vt[2] = {float(pix % width) + 0.5f, float(pix / width) + 0.5f}; + float *vt0_ptr = V + (F[f * 3] * 4); + float *vt1_ptr = V + (F[f * 3 + 1] * 4); + float *vt2_ptr = V + (F[f * 3 + 2] * 4); + + float vt0[2] = { + (vt0_ptr[0] / vt0_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, + (0.5f + 0.5f * vt0_ptr[1] / vt0_ptr[3]) * (height - 1) + 0.5f}; + float vt1[2] = { + (vt1_ptr[0] / vt1_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, + (0.5f + 0.5f * vt1_ptr[1] / vt1_ptr[3]) * (height - 1) + 0.5f}; + float vt2[2] = { + (vt2_ptr[0] / vt2_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, + (0.5f + 0.5f * vt2_ptr[1] / vt2_ptr[3]) * (height - 1) + 0.5f}; + + calculateBarycentricCoordinate(vt0, vt1, vt2, vt, barycentric); + + barycentric[0] = barycentric[0] / vt0_ptr[3]; + barycentric[1] = barycentric[1] / vt1_ptr[3]; + barycentric[2] = barycentric[2] / vt2_ptr[3]; + float w = 1.0f / (barycentric[0] + barycentric[1] + barycentric[2]); + barycentric[0] *= w; + barycentric[1] *= w; + barycentric[2] *= w; + } + barycentric_map[pix * 3] = barycentric[0]; + barycentric_map[pix * 3 + 1] = barycentric[1]; + barycentric_map[pix * 3 + 2] = barycentric[2]; +} + +void rasterizeImagecoordsKernelCPU(float *V, int *F, float *d, INT64 *zbuffer, + float occlusion_trunc, int width, int height, + int num_vertices, int num_faces, int f) { + float *vt0_ptr = V + (F[f * 3] * 4); + float *vt1_ptr = V + (F[f * 3 + 1] * 4); + float *vt2_ptr = V + (F[f * 3 + 2] * 4); + + float vt0[3] = {(vt0_ptr[0] / vt0_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, + (0.5f + 0.5f * vt0_ptr[1] / vt0_ptr[3]) * (height - 1) + 0.5f, + vt0_ptr[2] / vt0_ptr[3] * 0.49999f + 0.5f}; + float vt1[3] = {(vt1_ptr[0] / vt1_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, + (0.5f + 0.5f * vt1_ptr[1] / vt1_ptr[3]) * (height - 1) + 0.5f, + vt1_ptr[2] / vt1_ptr[3] * 0.49999f + 0.5f}; + float vt2[3] = {(vt2_ptr[0] / vt2_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, + (0.5f + 0.5f * vt2_ptr[1] / vt2_ptr[3]) * (height - 1) + 0.5f, + vt2_ptr[2] / vt2_ptr[3] * 0.49999f + 0.5f}; + + rasterizeTriangleCPU(f, vt0, vt1, vt2, width, height, zbuffer, d, + occlusion_trunc); +} + +std::vector rasterize_image_cpu(torch::Tensor V, torch::Tensor F, + torch::Tensor D, int width, + int height, + float occlusion_truncation, + int use_depth_prior) { + int num_faces = F.size(0); + int num_vertices = V.size(0); + auto options = + torch::TensorOptions().dtype(torch::kInt32).requires_grad(false); + auto INT64_options = + torch::TensorOptions().dtype(torch::kInt64).requires_grad(false); + auto findices = torch::zeros({height, width}, options); + INT64 maxint = (INT64)MAXINT * (INT64)MAXINT + (MAXINT - 1); + auto z_min = torch::ones({height, width}, INT64_options) * (int64_t)maxint; + + if (!use_depth_prior) { + for (int i = 0; i < num_faces; ++i) { + rasterizeImagecoordsKernelCPU(V.data_ptr(), F.data_ptr(), 0, + (INT64 *)z_min.data_ptr(), + occlusion_truncation, width, height, + num_vertices, num_faces, i); + } + } else { + for (int i = 0; i < num_faces; ++i) + rasterizeImagecoordsKernelCPU( + V.data_ptr(), F.data_ptr(), D.data_ptr(), + (INT64 *)z_min.data_ptr(), occlusion_truncation, width, + height, num_vertices, num_faces, i); + } + + auto float_options = + torch::TensorOptions().dtype(torch::kFloat32).requires_grad(false); + auto barycentric = torch::zeros({height, width, 3}, float_options); + for (int i = 0; i < width * height; ++i) + barycentricFromImgcoordCPU( + V.data_ptr(), F.data_ptr(), findices.data_ptr(), + (INT64 *)z_min.data_ptr(), width, height, num_vertices, + num_faces, barycentric.data_ptr(), i); + + return {findices, barycentric}; +} + +std::vector +rasterize_image(torch::Tensor V, torch::Tensor F, torch::Tensor D, int width, + int height, float occlusion_truncation, int use_depth_prior) { +#ifdef CUDA_ENABLED + return rasterize_image_gpu(V, F, D, width, height, occlusion_truncation, + use_depth_prior); +#else + return rasterize_image_cpu(V, F, D, width, height, occlusion_truncation, + use_depth_prior); +#endif +} + +PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { + m.def("rasterize_image", &rasterize_image, "Custom image rasterization"); +} diff --git a/python/sglang/kernels/ops/diffusion/ext/hunyuan3d_rasterizer/rasterizer.h b/python/sglang/kernels/ops/diffusion/ext/hunyuan3d_rasterizer/rasterizer.h new file mode 100644 index 000000000000..c9f6cec53863 --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/ext/hunyuan3d_rasterizer/rasterizer.h @@ -0,0 +1,61 @@ +// SPDX-License-Identifier: Apache-2.0 +// Adapted from Hunyuan3D-2: https://github.com/Tencent/Hunyuan3D-2 +// Original license: TENCENT HUNYUAN NON-COMMERCIAL LICENSE AGREEMENT + +#ifndef RASTERIZER_H_ +#define RASTERIZER_H_ + +#include +#include +#include + +#ifdef CUDA_ENABLED +#include +#else +#define __host__ +#define __device__ +#endif + +#define INT64 unsigned long long +#define MAXINT 2147483647 + +__host__ __device__ inline float calculateSignedArea2(float *a, float *b, + float *c) { + return ((c[0] - a[0]) * (b[1] - a[1]) - (b[0] - a[0]) * (c[1] - a[1])); +} + +__host__ __device__ inline void +calculateBarycentricCoordinate(float *a, float *b, float *c, float *p, + float *barycentric) { + float beta_tri = calculateSignedArea2(a, p, c); + float gamma_tri = calculateSignedArea2(a, b, p); + float area = calculateSignedArea2(a, b, c); + if (area == 0) { + barycentric[0] = -1.0; + barycentric[1] = -1.0; + barycentric[2] = -1.0; + return; + } + float tri_inv = 1.0 / area; + float beta = beta_tri * tri_inv; + float gamma = gamma_tri * tri_inv; + float alpha = 1.0 - beta - gamma; + barycentric[0] = alpha; + barycentric[1] = beta; + barycentric[2] = gamma; +} + +__host__ __device__ inline bool +isBarycentricCoordInBounds(float *barycentricCoord) { + return barycentricCoord[0] >= 0.0 && barycentricCoord[0] <= 1.0 && + barycentricCoord[1] >= 0.0 && barycentricCoord[1] <= 1.0 && + barycentricCoord[2] >= 0.0 && barycentricCoord[2] <= 1.0; +} + +std::vector rasterize_image_gpu(torch::Tensor V, torch::Tensor F, + torch::Tensor D, int width, + int height, + float occlusion_truncation, + int use_depth_prior); + +#endif diff --git a/python/sglang/kernels/ops/diffusion/ext/hunyuan3d_rasterizer/rasterizer_gpu.cu b/python/sglang/kernels/ops/diffusion/ext/hunyuan3d_rasterizer/rasterizer_gpu.cu new file mode 100644 index 000000000000..961ae630cec8 --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/ext/hunyuan3d_rasterizer/rasterizer_gpu.cu @@ -0,0 +1,173 @@ +// SPDX-License-Identifier: Apache-2.0 +// Adapted from Hunyuan3D-2: https://github.com/Tencent/Hunyuan3D-2 +// Original license: TENCENT HUNYUAN NON-COMMERCIAL LICENSE AGREEMENT + +#include "rasterizer.h" + +__device__ void rasterizeTriangleGPU(int idx, float *vt0, float *vt1, + float *vt2, int width, int height, + INT64 *zbuffer, float *d, + float occlusion_truncation) { + float x_min = std::min(vt0[0], std::min(vt1[0], vt2[0])); + float x_max = std::max(vt0[0], std::max(vt1[0], vt2[0])); + float y_min = std::min(vt0[1], std::min(vt1[1], vt2[1])); + float y_max = std::max(vt0[1], std::max(vt1[1], vt2[1])); + + for (int px = x_min; px < x_max + 1; ++px) { + if (px < 0 || px >= width) + continue; + for (int py = y_min; py < y_max + 1; ++py) { + if (py < 0 || py >= height) + continue; + float vt[2] = {px + 0.5f, py + 0.5f}; + float baryCentricCoordinate[3]; + calculateBarycentricCoordinate(vt0, vt1, vt2, vt, baryCentricCoordinate); + if (isBarycentricCoordInBounds(baryCentricCoordinate)) { + int pixel = py * width + px; + if (zbuffer == 0) { + atomicExch(&zbuffer[pixel], (INT64)(idx + 1)); + continue; + } + float depth = baryCentricCoordinate[0] * vt0[2] + + baryCentricCoordinate[1] * vt1[2] + + baryCentricCoordinate[2] * vt2[2]; + float depth_thres = 0; + if (d) { + depth_thres = d[pixel] * 0.49999f + 0.5f + occlusion_truncation; + } + + int z_quantize = depth * (2 << 17); + INT64 token = (INT64)z_quantize * MAXINT + (INT64)(idx + 1); + if (depth < depth_thres) + continue; + atomicMin(&zbuffer[pixel], token); + } + } + } +} + +__global__ void barycentricFromImgcoordGPU(float *V, int *F, int *findices, + INT64 *zbuffer, int width, + int height, int num_vertices, + int num_faces, + float *barycentric_map) { + int pix = blockIdx.x * blockDim.x + threadIdx.x; + if (pix >= width * height) + return; + INT64 f = zbuffer[pix] % MAXINT; + if (f == (MAXINT - 1)) { + findices[pix] = 0; + barycentric_map[pix * 3] = 0; + barycentric_map[pix * 3 + 1] = 0; + barycentric_map[pix * 3 + 2] = 0; + return; + } + findices[pix] = f; + f -= 1; + float barycentric[3] = {0, 0, 0}; + if (f >= 0) { + float vt[2] = {float(pix % width) + 0.5f, float(pix / width) + 0.5f}; + float *vt0_ptr = V + (F[f * 3] * 4); + float *vt1_ptr = V + (F[f * 3 + 1] * 4); + float *vt2_ptr = V + (F[f * 3 + 2] * 4); + + float vt0[2] = { + (vt0_ptr[0] / vt0_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, + (0.5f + 0.5f * vt0_ptr[1] / vt0_ptr[3]) * (height - 1) + 0.5f}; + float vt1[2] = { + (vt1_ptr[0] / vt1_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, + (0.5f + 0.5f * vt1_ptr[1] / vt1_ptr[3]) * (height - 1) + 0.5f}; + float vt2[2] = { + (vt2_ptr[0] / vt2_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, + (0.5f + 0.5f * vt2_ptr[1] / vt2_ptr[3]) * (height - 1) + 0.5f}; + + calculateBarycentricCoordinate(vt0, vt1, vt2, vt, barycentric); + + barycentric[0] = barycentric[0] / vt0_ptr[3]; + barycentric[1] = barycentric[1] / vt1_ptr[3]; + barycentric[2] = barycentric[2] / vt2_ptr[3]; + float w = 1.0f / (barycentric[0] + barycentric[1] + barycentric[2]); + barycentric[0] *= w; + barycentric[1] *= w; + barycentric[2] *= w; + } + barycentric_map[pix * 3] = barycentric[0]; + barycentric_map[pix * 3 + 1] = barycentric[1]; + barycentric_map[pix * 3 + 2] = barycentric[2]; +} + +__global__ void rasterizeImagecoordsKernelGPU(float *V, int *F, float *d, + INT64 *zbuffer, + float occlusion_trunc, int width, + int height, int num_vertices, + int num_faces) { + int f = blockIdx.x * blockDim.x + threadIdx.x; + if (f >= num_faces) + return; + + float *vt0_ptr = V + (F[f * 3] * 4); + float *vt1_ptr = V + (F[f * 3 + 1] * 4); + float *vt2_ptr = V + (F[f * 3 + 2] * 4); + + float vt0[3] = {(vt0_ptr[0] / vt0_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, + (0.5f + 0.5f * vt0_ptr[1] / vt0_ptr[3]) * (height - 1) + 0.5f, + vt0_ptr[2] / vt0_ptr[3] * 0.49999f + 0.5f}; + float vt1[3] = {(vt1_ptr[0] / vt1_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, + (0.5f + 0.5f * vt1_ptr[1] / vt1_ptr[3]) * (height - 1) + 0.5f, + vt1_ptr[2] / vt1_ptr[3] * 0.49999f + 0.5f}; + float vt2[3] = {(vt2_ptr[0] / vt2_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, + (0.5f + 0.5f * vt2_ptr[1] / vt2_ptr[3]) * (height - 1) + 0.5f, + vt2_ptr[2] / vt2_ptr[3] * 0.49999f + 0.5f}; + + rasterizeTriangleGPU(f, vt0, vt1, vt2, width, height, zbuffer, d, + occlusion_trunc); +} + +std::vector rasterize_image_gpu(torch::Tensor V, torch::Tensor F, + torch::Tensor D, int width, + int height, + float occlusion_truncation, + int use_depth_prior) { + int device_id = V.get_device(); + cudaSetDevice(device_id); + int num_faces = F.size(0); + int num_vertices = V.size(0); + auto options = torch::TensorOptions() + .dtype(torch::kInt32) + .device(torch::kCUDA, device_id) + .requires_grad(false); + auto INT64_options = torch::TensorOptions() + .dtype(torch::kInt64) + .device(torch::kCUDA, device_id) + .requires_grad(false); + auto findices = torch::zeros({height, width}, options); + INT64 maxint = (INT64)MAXINT * (INT64)MAXINT + (MAXINT - 1); + auto z_min = torch::ones({height, width}, INT64_options) * (int64_t)maxint; + + if (!use_depth_prior) { + rasterizeImagecoordsKernelGPU<<<(num_faces + 255) / 256, 256, 0, + at::cuda::getCurrentCUDAStream()>>>( + V.data_ptr(), F.data_ptr(), 0, + (INT64 *)z_min.data_ptr(), occlusion_truncation, width, height, + num_vertices, num_faces); + } else { + rasterizeImagecoordsKernelGPU<<<(num_faces + 255) / 256, 256, 0, + at::cuda::getCurrentCUDAStream()>>>( + V.data_ptr(), F.data_ptr(), D.data_ptr(), + (INT64 *)z_min.data_ptr(), occlusion_truncation, width, height, + num_vertices, num_faces); + } + + auto float_options = torch::TensorOptions() + .dtype(torch::kFloat32) + .device(torch::kCUDA, device_id) + .requires_grad(false); + auto barycentric = torch::zeros({height, width, 3}, float_options); + barycentricFromImgcoordGPU<<<(width * height + 255) / 256, 256, 0, + at::cuda::getCurrentCUDAStream()>>>( + V.data_ptr(), F.data_ptr(), findices.data_ptr(), + (INT64 *)z_min.data_ptr(), width, height, num_vertices, + num_faces, barycentric.data_ptr()); + + return {findices, barycentric}; +} diff --git a/python/sglang/kernels/ops/diffusion/render/__init__.py b/python/sglang/kernels/ops/diffusion/ext/loader.py similarity index 97% rename from python/sglang/kernels/ops/diffusion/render/__init__.py rename to python/sglang/kernels/ops/diffusion/ext/loader.py index d7277542ca3c..dd25a30d47be 100644 --- a/python/sglang/kernels/ops/diffusion/render/__init__.py +++ b/python/sglang/kernels/ops/diffusion/ext/loader.py @@ -1,5 +1,6 @@ from __future__ import annotations +import logging import os import shutil import sys @@ -8,9 +9,7 @@ import torch -from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger - -logger = init_logger(__name__) +logger = logging.getLogger(__name__) def _get_build_directory(name: str) -> Path: diff --git a/python/sglang/kernels/ops/diffusion/render/mesh_processor/__init__.py b/python/sglang/kernels/ops/diffusion/ext/mesh_processor/__init__.py similarity index 95% rename from python/sglang/kernels/ops/diffusion/render/mesh_processor/__init__.py rename to python/sglang/kernels/ops/diffusion/ext/mesh_processor/__init__.py index 9c6a9d67b637..fa399908aabf 100644 --- a/python/sglang/kernels/ops/diffusion/render/mesh_processor/__init__.py +++ b/python/sglang/kernels/ops/diffusion/ext/mesh_processor/__init__.py @@ -13,7 +13,7 @@ import numpy as np -from sglang.kernels.ops.diffusion.render import load_extension_with_recovery +from sglang.kernels.ops.diffusion.ext.loader import load_extension_with_recovery _abs_path = os.path.dirname(os.path.abspath(__file__)) _mesh_processor_kernel = None diff --git a/python/sglang/kernels/ops/diffusion/ext/mesh_processor/mesh_processor.cpp b/python/sglang/kernels/ops/diffusion/ext/mesh_processor/mesh_processor.cpp new file mode 100644 index 000000000000..efc2bd00c349 --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/ext/mesh_processor/mesh_processor.cpp @@ -0,0 +1,170 @@ +// SPDX-License-Identifier: Apache-2.0 +// Adapted from Hunyuan3D-2: https://github.com/Tencent/Hunyuan3D-2 +// Original license: TENCENT HUNYUAN NON-COMMERCIAL LICENSE AGREEMENT + +#include +#include +#include +#include +#include +#include +#include +#include + +namespace py = pybind11; +using namespace std; + +std::pair, py::array_t> +meshVerticeInpaint_smooth(py::array_t texture, py::array_t mask, + py::array_t vtx_pos, py::array_t vtx_uv, + py::array_t pos_idx, py::array_t uv_idx) { + auto texture_buf = texture.request(); + auto mask_buf = mask.request(); + auto vtx_pos_buf = vtx_pos.request(); + auto vtx_uv_buf = vtx_uv.request(); + auto pos_idx_buf = pos_idx.request(); + auto uv_idx_buf = uv_idx.request(); + + int texture_height = texture_buf.shape[0]; + int texture_width = texture_buf.shape[1]; + int texture_channel = texture_buf.shape[2]; + float *texture_ptr = static_cast(texture_buf.ptr); + uint8_t *mask_ptr = static_cast(mask_buf.ptr); + + int vtx_num = vtx_pos_buf.shape[0]; + float *vtx_pos_ptr = static_cast(vtx_pos_buf.ptr); + float *vtx_uv_ptr = static_cast(vtx_uv_buf.ptr); + int *pos_idx_ptr = static_cast(pos_idx_buf.ptr); + int *uv_idx_ptr = static_cast(uv_idx_buf.ptr); + + vector vtx_mask(vtx_num, 0.0f); + vector> vtx_color(vtx_num, + vector(texture_channel, 0.0f)); + vector uncolored_vtxs; + + vector> G(vtx_num); + + for (int i = 0; i < uv_idx_buf.shape[0]; ++i) { + for (int k = 0; k < 3; ++k) { + int vtx_uv_idx = uv_idx_ptr[i * 3 + k]; + int vtx_idx = pos_idx_ptr[i * 3 + k]; + int uv_v = round(vtx_uv_ptr[vtx_uv_idx * 2] * (texture_width - 1)); + int uv_u = + round((1.0 - vtx_uv_ptr[vtx_uv_idx * 2 + 1]) * (texture_height - 1)); + + if (mask_ptr[uv_u * texture_width + uv_v] > 0) { + vtx_mask[vtx_idx] = 1.0f; + for (int c = 0; c < texture_channel; ++c) { + vtx_color[vtx_idx][c] = + texture_ptr[(uv_u * texture_width + uv_v) * texture_channel + c]; + } + } else { + uncolored_vtxs.push_back(vtx_idx); + } + + G[pos_idx_ptr[i * 3 + k]].push_back(pos_idx_ptr[i * 3 + (k + 1) % 3]); + } + } + + int smooth_count = 2; + int last_uncolored_vtx_count = 0; + while (smooth_count > 0) { + int uncolored_vtx_count = 0; + + for (int vtx_idx : uncolored_vtxs) { + + vector sum_color(texture_channel, 0.0f); + float total_weight = 0.0f; + + array vtx_0 = {vtx_pos_ptr[vtx_idx * 3], + vtx_pos_ptr[vtx_idx * 3 + 1], + vtx_pos_ptr[vtx_idx * 3 + 2]}; + for (int connected_idx : G[vtx_idx]) { + if (vtx_mask[connected_idx] > 0) { + array vtx1 = {vtx_pos_ptr[connected_idx * 3], + vtx_pos_ptr[connected_idx * 3 + 1], + vtx_pos_ptr[connected_idx * 3 + 2]}; + float dist_weight = 1.0f / max(sqrt(pow(vtx_0[0] - vtx1[0], 2) + + pow(vtx_0[1] - vtx1[1], 2) + + pow(vtx_0[2] - vtx1[2], 2)), + 1E-4); + dist_weight = dist_weight * dist_weight; + for (int c = 0; c < texture_channel; ++c) { + sum_color[c] += vtx_color[connected_idx][c] * dist_weight; + } + total_weight += dist_weight; + } + } + + if (total_weight > 0.0f) { + for (int c = 0; c < texture_channel; ++c) { + vtx_color[vtx_idx][c] = sum_color[c] / total_weight; + } + vtx_mask[vtx_idx] = 1.0f; + } else { + uncolored_vtx_count++; + } + } + + if (last_uncolored_vtx_count == uncolored_vtx_count) { + smooth_count--; + } else { + smooth_count++; + } + last_uncolored_vtx_count = uncolored_vtx_count; + } + + py::array_t new_texture(texture_buf.size); + py::array_t new_mask(mask_buf.size); + + auto new_texture_buf = new_texture.request(); + auto new_mask_buf = new_mask.request(); + + float *new_texture_ptr = static_cast(new_texture_buf.ptr); + uint8_t *new_mask_ptr = static_cast(new_mask_buf.ptr); + std::copy(texture_ptr, texture_ptr + texture_buf.size, new_texture_ptr); + std::copy(mask_ptr, mask_ptr + mask_buf.size, new_mask_ptr); + + for (int face_idx = 0; face_idx < uv_idx_buf.shape[0]; ++face_idx) { + for (int k = 0; k < 3; ++k) { + int vtx_uv_idx = uv_idx_ptr[face_idx * 3 + k]; + int vtx_idx = pos_idx_ptr[face_idx * 3 + k]; + + if (vtx_mask[vtx_idx] == 1.0f) { + int uv_v = round(vtx_uv_ptr[vtx_uv_idx * 2] * (texture_width - 1)); + int uv_u = round((1.0 - vtx_uv_ptr[vtx_uv_idx * 2 + 1]) * + (texture_height - 1)); + + for (int c = 0; c < texture_channel; ++c) { + new_texture_ptr[(uv_u * texture_width + uv_v) * texture_channel + c] = + vtx_color[vtx_idx][c]; + } + new_mask_ptr[uv_u * texture_width + uv_v] = 255; + } + } + } + + new_texture.resize({texture_height, texture_width, 3}); + new_mask.resize({texture_height, texture_width}); + return std::make_pair(new_texture, new_mask); +} + +std::pair, py::array_t> +meshVerticeInpaint(py::array_t texture, py::array_t mask, + py::array_t vtx_pos, py::array_t vtx_uv, + py::array_t pos_idx, py::array_t uv_idx, + const std::string &method = "smooth") { + if (method == "smooth") { + return meshVerticeInpaint_smooth(texture, mask, vtx_pos, vtx_uv, pos_idx, + uv_idx); + } else { + throw std::invalid_argument("Invalid method. Use 'smooth'."); + } +} + +PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { + m.def("meshVerticeInpaint", &meshVerticeInpaint, + "Mesh-aware texture inpainting", py::arg("texture"), py::arg("mask"), + py::arg("vtx_pos"), py::arg("vtx_uv"), py::arg("pos_idx"), + py::arg("uv_idx"), py::arg("method") = "smooth"); +} diff --git a/python/sglang/kernels/ops/diffusion/layout/__init__.py b/python/sglang/kernels/ops/diffusion/layout/__init__.py new file mode 100644 index 000000000000..fe819e2cd9e5 --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/layout/__init__.py @@ -0,0 +1,5 @@ +"""Pure data-movement kernels: sequence-parallel relayout, varlen pack/scatter, causal padding. + +Every kernel here only moves values (plus zero fill), so each is bitwise +identical to the aten chain it replaces. +""" diff --git a/python/sglang/kernels/ops/diffusion/causal_conv3d_cat_pad.py b/python/sglang/kernels/ops/diffusion/layout/causal_conv3d_cat_pad_jit.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/causal_conv3d_cat_pad.py rename to python/sglang/kernels/ops/diffusion/layout/causal_conv3d_cat_pad_jit.py diff --git a/python/sglang/kernels/ops/diffusion/triton/causal_conv3d_pad.py b/python/sglang/kernels/ops/diffusion/layout/causal_conv3d_cat_pad_triton.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/triton/causal_conv3d_pad.py rename to python/sglang/kernels/ops/diffusion/layout/causal_conv3d_cat_pad_triton.py diff --git a/python/sglang/kernels/ops/diffusion/triton/ulysses_qkv.py b/python/sglang/kernels/ops/diffusion/layout/ulysses_qkv_triton.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/triton/ulysses_qkv.py rename to python/sglang/kernels/ops/diffusion/layout/ulysses_qkv_triton.py diff --git a/python/sglang/kernels/ops/diffusion/usp_relayout.py b/python/sglang/kernels/ops/diffusion/layout/usp_relayout_jit.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/usp_relayout.py rename to python/sglang/kernels/ops/diffusion/layout/usp_relayout_jit.py diff --git a/python/sglang/kernels/ops/diffusion/triton/varlen_pack_pad.py b/python/sglang/kernels/ops/diffusion/layout/varlen_pack_pad_triton.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/triton/varlen_pack_pad.py rename to python/sglang/kernels/ops/diffusion/layout/varlen_pack_pad_triton.py diff --git a/python/sglang/kernels/ops/diffusion/triton/wan_causal_cache.py b/python/sglang/kernels/ops/diffusion/layout/wan_causal_cache_triton.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/triton/wan_causal_cache.py rename to python/sglang/kernels/ops/diffusion/layout/wan_causal_cache_triton.py diff --git a/python/sglang/kernels/ops/diffusion/modulate/__init__.py b/python/sglang/kernels/ops/diffusion/modulate/__init__.py new file mode 100644 index 000000000000..9a45d65410ec --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/modulate/__init__.py @@ -0,0 +1 @@ +"""adaLN modulation: ``x * (1 + scale) + shift``, gating, and timestep conditioning.""" diff --git a/python/sglang/kernels/ops/diffusion/triton/indexed_modulation.py b/python/sglang/kernels/ops/diffusion/modulate/indexed_modulation_triton.py similarity index 98% rename from python/sglang/kernels/ops/diffusion/triton/indexed_modulation.py rename to python/sglang/kernels/ops/diffusion/modulate/indexed_modulation_triton.py index 726f273218f0..f6aff00832ae 100644 --- a/python/sglang/kernels/ops/diffusion/triton/indexed_modulation.py +++ b/python/sglang/kernels/ops/diffusion/modulate/indexed_modulation_triton.py @@ -4,7 +4,7 @@ import triton import triton.language as tl -from sglang.kernels.ops.diffusion.triton.numerics import round_bf16_to_fp32 +from sglang.kernels.ops.diffusion.common.numerics import round_bf16_to_fp32 @triton.jit diff --git a/python/sglang/kernels/ops/diffusion/triton/ltx2_ada_values.py b/python/sglang/kernels/ops/diffusion/modulate/ltx2_ada_values_triton.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/triton/ltx2_ada_values.py rename to python/sglang/kernels/ops/diffusion/modulate/ltx2_ada_values_triton.py diff --git a/python/sglang/kernels/ops/diffusion/modulate_scale_shift.py b/python/sglang/kernels/ops/diffusion/modulate/modulate_scale_shift_jit.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/modulate_scale_shift.py rename to python/sglang/kernels/ops/diffusion/modulate/modulate_scale_shift_jit.py diff --git a/python/sglang/kernels/ops/diffusion/residual_gate_add.py b/python/sglang/kernels/ops/diffusion/modulate/residual_gate_add_jit.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/residual_gate_add.py rename to python/sglang/kernels/ops/diffusion/modulate/residual_gate_add_jit.py diff --git a/python/sglang/kernels/ops/diffusion/triton/scale_shift.py b/python/sglang/kernels/ops/diffusion/modulate/scale_shift_triton.py similarity index 96% rename from python/sglang/kernels/ops/diffusion/triton/scale_shift.py rename to python/sglang/kernels/ops/diffusion/modulate/scale_shift_triton.py index 89572fd1ad09..668cc23725a2 100644 --- a/python/sglang/kernels/ops/diffusion/triton/scale_shift.py +++ b/python/sglang/kernels/ops/diffusion/modulate/scale_shift_triton.py @@ -2,8 +2,13 @@ import triton # type: ignore import triton.language as tl # type: ignore -from sglang.kernels.ops.diffusion.triton.numerics import mul_rn_f32 -from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.kernels.ops.diffusion.common.numerics import mul_rn_f32 +from sglang.kernels.ops.diffusion.common.platform import ( + is_cuda, + is_hip, + lazy_fallback, + select_impl, +) @triton.jit @@ -33,7 +38,7 @@ def try_fused_scaled_residual_add_exact( ) -> torch.Tensor | None: """Fuse ``residual + x * scale`` without changing eager FP32 rounding.""" if ( - not current_platform.is_cuda() + not is_cuda() or torch.is_grad_enabled() or torch.compiler.is_compiling() or residual.dtype != torch.float32 @@ -422,7 +427,7 @@ def grid(meta): # Compact scale [B, F, 1, C] -> [B*F, C] (per-frame) scale_reshaped = scale.squeeze(2).reshape(-1, C).contiguous() - if shift.dim() == 4 and current_platform.is_hip(): + if shift.dim() == 4 and is_hip(): # ROCm has no fused CUTLASS scale-shift kernel, so this native path # handles the causal Wan / LingBot output AdaLN, which passes a # per-frame shift [B, F, 1, C]. Broadcast it across each frame's @@ -728,24 +733,10 @@ def fuse_residual_layernorm_scale_shift_gate_select01_kernel( return output, residual_out, gate_out -if current_platform.is_npu(): - from .npu_fallback import fuse_scale_shift_native - - fuse_scale_shift_kernel = fuse_scale_shift_native - -if current_platform.is_mps(): - from .mps_fallback import fuse_scale_shift_kernel_native - - fuse_scale_shift_kernel = fuse_scale_shift_kernel_native - -if current_platform.is_musa(): - from .torch_fallback import fuse_scale_shift_kernel_native - - fuse_scale_shift_kernel = fuse_scale_shift_kernel_native - -if current_platform.is_cpu(): - from .torch_fallback import ( - fuse_scale_shift_kernel_native, - ) - - fuse_scale_shift_kernel = fuse_scale_shift_kernel_native +fuse_scale_shift_kernel = select_impl( + fuse_scale_shift_kernel, + npu=lazy_fallback("npu", "fuse_scale_shift_native"), + mps=lazy_fallback("mps", "fuse_scale_shift_kernel_native"), + musa=lazy_fallback("torch", "fuse_scale_shift_kernel_native"), + cpu=lazy_fallback("torch", "fuse_scale_shift_kernel_native"), +) diff --git a/python/sglang/kernels/ops/diffusion/timestep_embedding.py b/python/sglang/kernels/ops/diffusion/modulate/timestep_embedding_jit.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/timestep_embedding.py rename to python/sglang/kernels/ops/diffusion/modulate/timestep_embedding_jit.py diff --git a/python/sglang/kernels/ops/diffusion/triton/wan_temb_table_slices.py b/python/sglang/kernels/ops/diffusion/modulate/wan_temb_table_slices_triton.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/triton/wan_temb_table_slices.py rename to python/sglang/kernels/ops/diffusion/modulate/wan_temb_table_slices_triton.py diff --git a/python/sglang/kernels/ops/diffusion/norm/__init__.py b/python/sglang/kernels/ops/diffusion/norm/__init__.py new file mode 100644 index 000000000000..ca274b82e228 --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/norm/__init__.py @@ -0,0 +1,7 @@ +"""Normalization kernels: RMSNorm / LayerNorm / GroupNorm and their fused epilogues. + +Which implementation to pick is documented in the selection matrix in +``sglang/kernels/ops/diffusion/README.md`` -- there are several per norm type +and they differ by numerical contract (bit-exact vs close), activation layout +and backend, not by speed alone. +""" diff --git a/python/sglang/kernels/ops/diffusion/norm/cutedsl_common/__init__.py b/python/sglang/kernels/ops/diffusion/norm/cutedsl_common/__init__.py new file mode 100644 index 000000000000..6eeffeae984f --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/norm/cutedsl_common/__init__.py @@ -0,0 +1 @@ +"""CuTe-DSL building blocks shared by the CUTLASS-backed norm fusions.""" diff --git a/python/sglang/kernels/ops/diffusion/cutedsl/common/norm_fusion.py b/python/sglang/kernels/ops/diffusion/norm/cutedsl_common/norm_fusion.py similarity index 99% rename from python/sglang/kernels/ops/diffusion/cutedsl/common/norm_fusion.py rename to python/sglang/kernels/ops/diffusion/norm/cutedsl_common/norm_fusion.py index 0b6bc098c776..aef5123e9597 100644 --- a/python/sglang/kernels/ops/diffusion/cutedsl/common/norm_fusion.py +++ b/python/sglang/kernels/ops/diffusion/norm/cutedsl_common/norm_fusion.py @@ -5,7 +5,7 @@ import torch from einops import rearrange -from sglang.kernels.ops.diffusion.cutedsl.common.reduce import ( +from sglang.kernels.ops.diffusion.norm.cutedsl_common.reduce import ( cta_reduce_sum, warp_reduce_sum, ) diff --git a/python/sglang/kernels/ops/diffusion/cutedsl/common/reduce.py b/python/sglang/kernels/ops/diffusion/norm/cutedsl_common/reduce.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/cutedsl/common/reduce.py rename to python/sglang/kernels/ops/diffusion/norm/cutedsl_common/reduce.py diff --git a/python/sglang/kernels/ops/diffusion/cutedsl/utils.py b/python/sglang/kernels/ops/diffusion/norm/cutedsl_utils.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/cutedsl/utils.py rename to python/sglang/kernels/ops/diffusion/norm/cutedsl_utils.py diff --git a/python/sglang/kernels/ops/diffusion/flydsl/fused_residual_norm.py b/python/sglang/kernels/ops/diffusion/norm/fused_residual_norm_flydsl.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/flydsl/fused_residual_norm.py rename to python/sglang/kernels/ops/diffusion/norm/fused_residual_norm_flydsl.py diff --git a/python/sglang/kernels/ops/diffusion/group_norm_silu.py b/python/sglang/kernels/ops/diffusion/norm/group_norm_silu.py similarity index 90% rename from python/sglang/kernels/ops/diffusion/group_norm_silu.py rename to python/sglang/kernels/ops/diffusion/norm/group_norm_silu.py index 43be074638bc..326fe0bb897d 100644 --- a/python/sglang/kernels/ops/diffusion/group_norm_silu.py +++ b/python/sglang/kernels/ops/diffusion/norm/group_norm_silu.py @@ -18,7 +18,7 @@ def apply_group_norm_silu( and norm.weight is not None and norm.bias is not None ): - from sglang.kernels.ops.diffusion.triton.group_norm_silu import ( + from sglang.kernels.ops.diffusion.norm.group_norm_silu_triton import ( triton_group_norm_silu, ) diff --git a/python/sglang/kernels/ops/diffusion/triton/group_norm_silu.py b/python/sglang/kernels/ops/diffusion/norm/group_norm_silu_triton.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/triton/group_norm_silu.py rename to python/sglang/kernels/ops/diffusion/norm/group_norm_silu_triton.py diff --git a/python/sglang/kernels/ops/diffusion/triton/group_norm_silu_twopass.py b/python/sglang/kernels/ops/diffusion/norm/group_norm_silu_twopass_triton.py similarity index 84% rename from python/sglang/kernels/ops/diffusion/triton/group_norm_silu_twopass.py rename to python/sglang/kernels/ops/diffusion/norm/group_norm_silu_twopass_triton.py index fa4bad151a1d..62866c9d0731 100644 --- a/python/sglang/kernels/ops/diffusion/triton/group_norm_silu_twopass.py +++ b/python/sglang/kernels/ops/diffusion/norm/group_norm_silu_twopass_triton.py @@ -16,8 +16,9 @@ pure elementwise kernel; - optional SiLU epilogue (``apply_silu=False`` gives plain GroupNorm); - restricted static shapes: power-of-two ``C <= 2048`` that ``num_groups`` - divides. Callers must treat a ``None`` return as "unsupported" and fall - back to their reference path. + divides. Support is a predicate (``can_use_group_norm_silu_4d`` / + ``can_use_group_norm_silu_rows``); the kernels raise on an unsupported + input rather than silently returning ``None``. """ import torch @@ -185,6 +186,39 @@ def _twopass_supported(x, weight, bias, num_groups) -> bool: return triton.next_power_of_2(c) == c and c <= _MAX_CHANNELS +def can_use_group_norm_silu_4d( + x: torch.Tensor, + weight: torch.Tensor, + bias: torch.Tensor, + num_groups: int, +) -> bool: + """Whether :func:`group_norm_silu_4d` supports this channels_last input.""" + if x.dim() != 4 or not _twopass_supported(x, weight, bias, num_groups): + return False + _, c, h, w = x.shape + # c > 1 and a non-trivial spatial extent make the channels_last check + # unambiguous (degenerate shapes are contiguous in both formats). + return ( + c > 1 + and (h > 1 or w > 1) + and x.is_contiguous(memory_format=torch.channels_last) + ) + + +def can_use_group_norm_silu_rows( + x3: torch.Tensor, + weight: torch.Tensor, + bias: torch.Tensor, + num_groups: int, +) -> bool: + """Whether :func:`group_norm_silu_rows` supports this (N, L, C) input.""" + return ( + x3.dim() == 3 + and x3.is_contiguous() + and _twopass_supported(x3, weight, bias, num_groups) + ) + + def group_norm_silu_4d( x: torch.Tensor, weight: torch.Tensor, @@ -192,24 +226,16 @@ def group_norm_silu_4d( num_groups: int, eps: float, apply_silu: bool = True, -) -> torch.Tensor | None: +) -> torch.Tensor: """Fused GroupNorm(+SiLU) for a channels_last 4D (N, C, H, W) activation. Runs the rows kernel on the free (N, H*W, C) view (no layout copy) and - preserves the channels_last output layout. Returns ``None`` when the - input is unsupported; callers must fall back to their reference path. + preserves the channels_last output layout. Guard with + :func:`can_use_group_norm_silu_4d`. """ - if x.dim() != 4 or not _twopass_supported(x, weight, bias, num_groups): - return None + if not can_use_group_norm_silu_4d(x, weight, bias, num_groups): + raise ValueError("unsupported input for group_norm_silu_4d") n_batch, c, h, w = x.shape - # c > 1 and a non-trivial spatial extent make the channels_last check - # unambiguous (degenerate shapes are contiguous in both formats). - if not ( - c > 1 - and (h > 1 or w > 1) - and x.is_contiguous(memory_format=torch.channels_last) - ): - return None x3 = x.permute(0, 2, 3, 1).reshape(n_batch, h * w, c) y3 = _gn_silu_rows(x3, weight, bias, num_groups, eps, apply_silu) return y3.reshape(n_batch, h, w, c).permute(0, 3, 1, 2) @@ -222,19 +248,19 @@ def group_norm_silu_rows( num_groups: int, eps: float, apply_silu: bool = True, -) -> torch.Tensor | None: +) -> torch.Tensor: """Fused GroupNorm(+SiLU) over (N, L, C) rows (C = channels, innermost). - Returns ``None`` when the input is unsupported; callers must fall back. + Guard with :func:`can_use_group_norm_silu_rows`. """ - if x3.dim() != 3 or not x3.is_contiguous(): - return None - if not _twopass_supported(x3, weight, bias, num_groups): - return None + if not can_use_group_norm_silu_rows(x3, weight, bias, num_groups): + raise ValueError("unsupported input for group_norm_silu_rows") return _gn_silu_rows(x3, weight, bias, num_groups, eps, apply_silu) __all__ = [ + "can_use_group_norm_silu_4d", + "can_use_group_norm_silu_rows", "group_norm_silu_4d", "group_norm_silu_rows", ] diff --git a/python/sglang/kernels/ops/diffusion/triton/layernorm_modulate.py b/python/sglang/kernels/ops/diffusion/norm/layernorm_modulate_triton.py similarity index 99% rename from python/sglang/kernels/ops/diffusion/triton/layernorm_modulate.py rename to python/sglang/kernels/ops/diffusion/norm/layernorm_modulate_triton.py index 91314b519597..2082d4071f27 100644 --- a/python/sglang/kernels/ops/diffusion/triton/layernorm_modulate.py +++ b/python/sglang/kernels/ops/diffusion/norm/layernorm_modulate_triton.py @@ -47,7 +47,7 @@ import triton.language as tl # type: ignore from sglang.kernels.jit.utils import get_jit_cuda_arch -from sglang.kernels.ops.diffusion.triton.numerics import ( +from sglang.kernels.ops.diffusion.common.numerics import ( cuda_rsqrtf, div_rn_f32, round_bf16_to_fp32, diff --git a/python/sglang/kernels/ops/diffusion/triton/native_bf16_rmsnorm.py b/python/sglang/kernels/ops/diffusion/norm/native_bf16_rmsnorm_triton.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/triton/native_bf16_rmsnorm.py rename to python/sglang/kernels/ops/diffusion/norm/native_bf16_rmsnorm_triton.py diff --git a/python/sglang/kernels/ops/diffusion/norm_scale_shift_native.py b/python/sglang/kernels/ops/diffusion/norm/norm_scale_shift_jit.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/norm_scale_shift_native.py rename to python/sglang/kernels/ops/diffusion/norm/norm_scale_shift_jit.py diff --git a/python/sglang/kernels/ops/diffusion/triton/norm.py b/python/sglang/kernels/ops/diffusion/norm/norm_triton.py similarity index 97% rename from python/sglang/kernels/ops/diffusion/triton/norm.py rename to python/sglang/kernels/ops/diffusion/norm/norm_triton.py index bc44a2be10cd..22f6d3aa7d58 100644 --- a/python/sglang/kernels/ops/diffusion/triton/norm.py +++ b/python/sglang/kernels/ops/diffusion/norm/norm_triton.py @@ -5,7 +5,7 @@ import triton.language as tl # type: ignore from torch import Tensor -from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.kernels.ops.diffusion.common.platform import lazy_fallback, select_impl from sglang.srt.utils.custom_op import register_custom_op @@ -647,14 +647,13 @@ def norm_infer( return out -if current_platform.is_mps(): - from .mps_fallback import norm_infer_native, rms_norm_fn_native - - norm_infer = norm_infer_native - rms_norm_fn = rms_norm_fn_native - -if current_platform.is_cpu(): - from .torch_fallback import norm_infer_native, rms_norm_fn_native - - norm_infer = norm_infer_native - rms_norm_fn = rms_norm_fn_native +norm_infer = select_impl( + norm_infer, + mps=lazy_fallback("mps", "norm_infer_native"), + cpu=lazy_fallback("torch", "norm_infer_native"), +) +rms_norm_fn = select_impl( + rms_norm_fn, + mps=lazy_fallback("mps", "rms_norm_fn_native"), + cpu=lazy_fallback("torch", "rms_norm_fn_native"), +) diff --git a/python/sglang/kernels/ops/diffusion/triton/rmsnorm_onepass.py b/python/sglang/kernels/ops/diffusion/norm/rmsnorm_onepass_triton.py similarity index 81% rename from python/sglang/kernels/ops/diffusion/triton/rmsnorm_onepass.py rename to python/sglang/kernels/ops/diffusion/norm/rmsnorm_onepass_triton.py index 065205381c48..670d1fc2db82 100644 --- a/python/sglang/kernels/ops/diffusion/triton/rmsnorm_onepass.py +++ b/python/sglang/kernels/ops/diffusion/norm/rmsnorm_onepass_triton.py @@ -3,7 +3,7 @@ import triton.language as tl # type: ignore from sglang.kernel_api_logging import debug_kernel_api -from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.kernels.ops.diffusion.common.platform import lazy_fallback, select_impl from sglang.srt.utils.custom_op import register_custom_op @@ -69,15 +69,9 @@ def triton_one_pass_rms_norm(x: torch.Tensor, w: torch.Tensor, eps: float = 1e-6 return _triton_one_pass_rms_norm_cuda(x, w, eps) -if current_platform.is_mps(): - from .mps_fallback import triton_one_pass_rms_norm_native - - @debug_kernel_api - def triton_one_pass_rms_norm(x: torch.Tensor, w: torch.Tensor, eps: float = 1e-6): - return triton_one_pass_rms_norm_native(x, w, eps) - - -if current_platform.is_cpu(): - from .torch_fallback import triton_one_pass_rms_norm_native - - triton_one_pass_rms_norm = triton_one_pass_rms_norm_native +triton_one_pass_rms_norm = select_impl( + triton_one_pass_rms_norm, + # MPS keeps the api-logging wrapper the Triton entry point carries. + mps=debug_kernel_api(lazy_fallback("mps", "triton_one_pass_rms_norm_native")), + cpu=lazy_fallback("torch", "triton_one_pass_rms_norm_native"), +) diff --git a/python/sglang/kernels/ops/diffusion/triton/rmsnorm_scale_shift_bitexact.py b/python/sglang/kernels/ops/diffusion/norm/rmsnorm_scale_shift_bitexact.py similarity index 99% rename from python/sglang/kernels/ops/diffusion/triton/rmsnorm_scale_shift_bitexact.py rename to python/sglang/kernels/ops/diffusion/norm/rmsnorm_scale_shift_bitexact.py index 521d1b0a974b..6f4db351f72f 100644 --- a/python/sglang/kernels/ops/diffusion/triton/rmsnorm_scale_shift_bitexact.py +++ b/python/sglang/kernels/ops/diffusion/norm/rmsnorm_scale_shift_bitexact.py @@ -55,7 +55,7 @@ import triton # type: ignore import triton.language as tl # type: ignore -from sglang.kernels.ops.diffusion.triton.numerics import ( +from sglang.kernels.ops.diffusion.common.numerics import ( mul_rn_f32, round_bf16_to_fp32, rsqrt_approx_f32, diff --git a/python/sglang/kernels/ops/diffusion/cutedsl/scale_residual_norm_scale_shift.py b/python/sglang/kernels/ops/diffusion/norm/scale_residual_norm_cutedsl.py similarity index 98% rename from python/sglang/kernels/ops/diffusion/cutedsl/scale_residual_norm_scale_shift.py rename to python/sglang/kernels/ops/diffusion/norm/scale_residual_norm_cutedsl.py index e0c3c4ce3ba5..8e6d049501e1 100644 --- a/python/sglang/kernels/ops/diffusion/cutedsl/scale_residual_norm_scale_shift.py +++ b/python/sglang/kernels/ops/diffusion/norm/scale_residual_norm_cutedsl.py @@ -5,12 +5,12 @@ import cutlass.cute as cute import torch -from sglang.kernels.ops.diffusion.cutedsl.common.norm_fusion import ( +from sglang.kernels.ops.diffusion.norm.cutedsl_common.norm_fusion import ( apply_norm_cta, broadcast_tensor_for_bsfd, tensor_slice_for_bsfd, ) -from sglang.kernels.ops.diffusion.cutedsl.utils import ( +from sglang.kernels.ops.diffusion.norm.cutedsl_utils import ( WARP_SIZE, to_fake_cute_args, ) @@ -263,7 +263,7 @@ def fused_norm_scale_shift( D must be a multiple of 256 and <= 8192 to enable LDG.128 vectorized loads per thread and avoid predicated loads (e.g., bounds checks such as `index < D`). """ - from sglang.kernels.ops.diffusion.norm_scale_shift_native import ( + from sglang.kernels.ops.diffusion.norm.norm_scale_shift_jit import ( try_fused_norm_scale_shift as _try_qwen_native_norm_scale_shift, ) @@ -349,7 +349,7 @@ def fused_scale_residual_norm_scale_shift( D must be a multiple of 256 and <= 8192 to enable LDG.128 vectorized loads per thread and avoid predicated loads (e.g., bounds checks such as `index < D`). """ - from sglang.kernels.ops.diffusion.norm_scale_shift_native import ( + from sglang.kernels.ops.diffusion.norm.norm_scale_shift_jit import ( try_fused_scale_residual_norm_scale_shift as _try_qwen_native_residual_path, ) diff --git a/python/sglang/kernels/ops/diffusion/triton/wan_rmsnorm_silu.py b/python/sglang/kernels/ops/diffusion/norm/wan_rmsnorm_silu_triton.py similarity index 95% rename from python/sglang/kernels/ops/diffusion/triton/wan_rmsnorm_silu.py rename to python/sglang/kernels/ops/diffusion/norm/wan_rmsnorm_silu_triton.py index f3d8c4a48844..7cd1a455e9c4 100644 --- a/python/sglang/kernels/ops/diffusion/triton/wan_rmsnorm_silu.py +++ b/python/sglang/kernels/ops/diffusion/norm/wan_rmsnorm_silu_triton.py @@ -11,9 +11,8 @@ aten promotion to fp32 at ``* gamma`` for half-precision x with fp32 affine params -- the autocast case), SiLU in fp32. Bitwise equality with aten is still not guaranteed (different reduction and SiLU paths), so callers must -keep this behind an opt-in gate. ``wan_rmsnorm_silu`` returns ``None`` for -unsupported inputs (see ``can_use_wan_rmsnorm_silu``); callers must fall -back to their reference path. +keep this behind an opt-in gate. Support is a predicate +(``can_use_wan_rmsnorm_silu``); the kernel raises on an unsupported input. """ from __future__ import annotations @@ -175,13 +174,13 @@ def wan_rmsnorm_silu( bias: torch.Tensor | None = None, rms_scale: float | None = None, eps: float = 1e-12, -) -> torch.Tensor | None: +) -> torch.Tensor: """Fused ``SiLU(F.normalize(x, dim=1) * rms_scale * gamma + bias)``. - Returns ``None`` when the input is unsupported; callers must fall back. + Guard with :func:`can_use_wan_rmsnorm_silu`. """ if not can_use_wan_rmsnorm_silu(x, gamma, bias): - return None + raise ValueError("unsupported input for wan_rmsnorm_silu") channels = x.shape[1] gamma = gamma.reshape(channels).contiguous() diff --git a/python/sglang/kernels/ops/diffusion/triton/zimage_native_norm.py b/python/sglang/kernels/ops/diffusion/norm/zimage_qk_rmsnorm_triton.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/triton/zimage_native_norm.py rename to python/sglang/kernels/ops/diffusion/norm/zimage_qk_rmsnorm_triton.py diff --git a/python/sglang/kernels/ops/diffusion/render/hunyuan3d_rasterizer/rasterizer.cpp b/python/sglang/kernels/ops/diffusion/render/hunyuan3d_rasterizer/rasterizer.cpp deleted file mode 100644 index 16773e857e3d..000000000000 --- a/python/sglang/kernels/ops/diffusion/render/hunyuan3d_rasterizer/rasterizer.cpp +++ /dev/null @@ -1,140 +0,0 @@ -// SPDX-License-Identifier: Apache-2.0 -// Adapted from Hunyuan3D-2: https://github.com/Tencent/Hunyuan3D-2 -// Original license: TENCENT HUNYUAN NON-COMMERCIAL LICENSE AGREEMENT - -#include "rasterizer.h" - -void rasterizeTriangleCPU(int idx, float* vt0, float* vt1, float* vt2, int width, int height, INT64* zbuffer, float* d, float occlusion_truncation) { - float x_min = std::min(vt0[0], std::min(vt1[0],vt2[0])); - float x_max = std::max(vt0[0], std::max(vt1[0],vt2[0])); - float y_min = std::min(vt0[1], std::min(vt1[1],vt2[1])); - float y_max = std::max(vt0[1], std::max(vt1[1],vt2[1])); - - for (int px = x_min; px < x_max + 1; ++px) { - if (px < 0 || px >= width) - continue; - for (int py = y_min; py < y_max + 1; ++py) { - if (py < 0 || py >= height) - continue; - float vt[2] = {px + 0.5f, py + 0.5f}; - float baryCentricCoordinate[3]; - calculateBarycentricCoordinate(vt0, vt1, vt2, vt, baryCentricCoordinate); - if (isBarycentricCoordInBounds(baryCentricCoordinate)) { - int pixel = py * width + px; - if (zbuffer == 0) { - zbuffer[pixel] = (INT64)(idx + 1); - continue; - } - - float depth = baryCentricCoordinate[0] * vt0[2] + baryCentricCoordinate[1] * vt1[2] + baryCentricCoordinate[2] * vt2[2]; - float depth_thres = 0; - if (d) { - depth_thres = d[pixel] * 0.49999f + 0.5f + occlusion_truncation; - } - - int z_quantize = depth * (2<<17); - INT64 token = (INT64)z_quantize * MAXINT + (INT64)(idx + 1); - if (depth < depth_thres) - continue; - zbuffer[pixel] = std::min(zbuffer[pixel], token); - } - } - } -} - -void barycentricFromImgcoordCPU(float* V, int* F, int* findices, INT64* zbuffer, int width, int height, int num_vertices, int num_faces, - float* barycentric_map, int pix) -{ - INT64 f = zbuffer[pix] % MAXINT; - if (f == (MAXINT-1)) { - findices[pix] = 0; - barycentric_map[pix * 3] = 0; - barycentric_map[pix * 3 + 1] = 0; - barycentric_map[pix * 3 + 2] = 0; - return; - } - findices[pix] = f; - f -= 1; - float barycentric[3] = {0, 0, 0}; - if (f >= 0) { - float vt[2] = {float(pix % width) + 0.5f, float(pix / width) + 0.5f}; - float* vt0_ptr = V + (F[f * 3] * 4); - float* vt1_ptr = V + (F[f * 3 + 1] * 4); - float* vt2_ptr = V + (F[f * 3 + 2] * 4); - - float vt0[2] = {(vt0_ptr[0] / vt0_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, (0.5f + 0.5f * vt0_ptr[1] / vt0_ptr[3]) * (height - 1) + 0.5f}; - float vt1[2] = {(vt1_ptr[0] / vt1_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, (0.5f + 0.5f * vt1_ptr[1] / vt1_ptr[3]) * (height - 1) + 0.5f}; - float vt2[2] = {(vt2_ptr[0] / vt2_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, (0.5f + 0.5f * vt2_ptr[1] / vt2_ptr[3]) * (height - 1) + 0.5f}; - - calculateBarycentricCoordinate(vt0, vt1, vt2, vt, barycentric); - - barycentric[0] = barycentric[0] / vt0_ptr[3]; - barycentric[1] = barycentric[1] / vt1_ptr[3]; - barycentric[2] = barycentric[2] / vt2_ptr[3]; - float w = 1.0f / (barycentric[0] + barycentric[1] + barycentric[2]); - barycentric[0] *= w; - barycentric[1] *= w; - barycentric[2] *= w; - } - barycentric_map[pix * 3] = barycentric[0]; - barycentric_map[pix * 3 + 1] = barycentric[1]; - barycentric_map[pix * 3 + 2] = barycentric[2]; -} - -void rasterizeImagecoordsKernelCPU(float* V, int* F, float* d, INT64* zbuffer, float occlusion_trunc, int width, int height, int num_vertices, int num_faces, int f) -{ - float* vt0_ptr = V + (F[f * 3] * 4); - float* vt1_ptr = V + (F[f * 3 + 1] * 4); - float* vt2_ptr = V + (F[f * 3 + 2] * 4); - - float vt0[3] = {(vt0_ptr[0] / vt0_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, (0.5f + 0.5f * vt0_ptr[1] / vt0_ptr[3]) * (height - 1) + 0.5f, vt0_ptr[2] / vt0_ptr[3] * 0.49999f + 0.5f}; - float vt1[3] = {(vt1_ptr[0] / vt1_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, (0.5f + 0.5f * vt1_ptr[1] / vt1_ptr[3]) * (height - 1) + 0.5f, vt1_ptr[2] / vt1_ptr[3] * 0.49999f + 0.5f}; - float vt2[3] = {(vt2_ptr[0] / vt2_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, (0.5f + 0.5f * vt2_ptr[1] / vt2_ptr[3]) * (height - 1) + 0.5f, vt2_ptr[2] / vt2_ptr[3] * 0.49999f + 0.5f}; - - rasterizeTriangleCPU(f, vt0, vt1, vt2, width, height, zbuffer, d, occlusion_trunc); -} - -std::vector rasterize_image_cpu(torch::Tensor V, torch::Tensor F, torch::Tensor D, - int width, int height, float occlusion_truncation, int use_depth_prior) -{ - int num_faces = F.size(0); - int num_vertices = V.size(0); - auto options = torch::TensorOptions().dtype(torch::kInt32).requires_grad(false); - auto INT64_options = torch::TensorOptions().dtype(torch::kInt64).requires_grad(false); - auto findices = torch::zeros({height, width}, options); - INT64 maxint = (INT64)MAXINT * (INT64)MAXINT + (MAXINT - 1); - auto z_min = torch::ones({height, width}, INT64_options) * (int64_t)maxint; - - if (!use_depth_prior) { - for (int i = 0; i < num_faces; ++i) { - rasterizeImagecoordsKernelCPU(V.data_ptr(), F.data_ptr(), 0, - (INT64*)z_min.data_ptr(), occlusion_truncation, width, height, num_vertices, num_faces, i); - } - } else { - for (int i = 0; i < num_faces; ++i) - rasterizeImagecoordsKernelCPU(V.data_ptr(), F.data_ptr(), D.data_ptr(), - (INT64*)z_min.data_ptr(), occlusion_truncation, width, height, num_vertices, num_faces, i); - } - - auto float_options = torch::TensorOptions().dtype(torch::kFloat32).requires_grad(false); - auto barycentric = torch::zeros({height, width, 3}, float_options); - for (int i = 0; i < width * height; ++i) - barycentricFromImgcoordCPU(V.data_ptr(), F.data_ptr(), - findices.data_ptr(), (INT64*)z_min.data_ptr(), width, height, num_vertices, num_faces, barycentric.data_ptr(), i); - - return {findices, barycentric}; -} - -std::vector rasterize_image(torch::Tensor V, torch::Tensor F, torch::Tensor D, - int width, int height, float occlusion_truncation, int use_depth_prior) -{ -#ifdef CUDA_ENABLED - return rasterize_image_gpu(V, F, D, width, height, occlusion_truncation, use_depth_prior); -#else - return rasterize_image_cpu(V, F, D, width, height, occlusion_truncation, use_depth_prior); -#endif -} - -PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { - m.def("rasterize_image", &rasterize_image, "Custom image rasterization"); -} diff --git a/python/sglang/kernels/ops/diffusion/render/hunyuan3d_rasterizer/rasterizer.h b/python/sglang/kernels/ops/diffusion/render/hunyuan3d_rasterizer/rasterizer.h deleted file mode 100644 index 84e12ca71b10..000000000000 --- a/python/sglang/kernels/ops/diffusion/render/hunyuan3d_rasterizer/rasterizer.h +++ /dev/null @@ -1,56 +0,0 @@ -// SPDX-License-Identifier: Apache-2.0 -// Adapted from Hunyuan3D-2: https://github.com/Tencent/Hunyuan3D-2 -// Original license: TENCENT HUNYUAN NON-COMMERCIAL LICENSE AGREEMENT - -#ifndef RASTERIZER_H_ -#define RASTERIZER_H_ - -#include -#include -#include - -#ifdef CUDA_ENABLED -#include -#else -#define __host__ -#define __device__ -#endif - -#define INT64 unsigned long long -#define MAXINT 2147483647 - -__host__ __device__ inline float calculateSignedArea2(float* a, float* b, float* c) { - return ((c[0] - a[0]) * (b[1] - a[1]) - (b[0] - a[0]) * (c[1] - a[1])); -} - -__host__ __device__ inline void calculateBarycentricCoordinate(float* a, float* b, float* c, float* p, - float* barycentric) -{ - float beta_tri = calculateSignedArea2(a, p, c); - float gamma_tri = calculateSignedArea2(a, b, p); - float area = calculateSignedArea2(a, b, c); - if (area == 0) { - barycentric[0] = -1.0; - barycentric[1] = -1.0; - barycentric[2] = -1.0; - return; - } - float tri_inv = 1.0 / area; - float beta = beta_tri * tri_inv; - float gamma = gamma_tri * tri_inv; - float alpha = 1.0 - beta - gamma; - barycentric[0] = alpha; - barycentric[1] = beta; - barycentric[2] = gamma; -} - -__host__ __device__ inline bool isBarycentricCoordInBounds(float* barycentricCoord) { - return barycentricCoord[0] >= 0.0 && barycentricCoord[0] <= 1.0 && - barycentricCoord[1] >= 0.0 && barycentricCoord[1] <= 1.0 && - barycentricCoord[2] >= 0.0 && barycentricCoord[2] <= 1.0; -} - -std::vector rasterize_image_gpu(torch::Tensor V, torch::Tensor F, torch::Tensor D, - int width, int height, float occlusion_truncation, int use_depth_prior); - -#endif diff --git a/python/sglang/kernels/ops/diffusion/render/hunyuan3d_rasterizer/rasterizer_gpu.cu b/python/sglang/kernels/ops/diffusion/render/hunyuan3d_rasterizer/rasterizer_gpu.cu deleted file mode 100644 index f1317270d23d..000000000000 --- a/python/sglang/kernels/ops/diffusion/render/hunyuan3d_rasterizer/rasterizer_gpu.cu +++ /dev/null @@ -1,130 +0,0 @@ -// SPDX-License-Identifier: Apache-2.0 -// Adapted from Hunyuan3D-2: https://github.com/Tencent/Hunyuan3D-2 -// Original license: TENCENT HUNYUAN NON-COMMERCIAL LICENSE AGREEMENT - -#include "rasterizer.h" - -__device__ void rasterizeTriangleGPU(int idx, float* vt0, float* vt1, float* vt2, int width, int height, INT64* zbuffer, float* d, float occlusion_truncation) { - float x_min = std::min(vt0[0], std::min(vt1[0],vt2[0])); - float x_max = std::max(vt0[0], std::max(vt1[0],vt2[0])); - float y_min = std::min(vt0[1], std::min(vt1[1],vt2[1])); - float y_max = std::max(vt0[1], std::max(vt1[1],vt2[1])); - - for (int px = x_min; px < x_max + 1; ++px) { - if (px < 0 || px >= width) - continue; - for (int py = y_min; py < y_max + 1; ++py) { - if (py < 0 || py >= height) - continue; - float vt[2] = {px + 0.5f, py + 0.5f}; - float baryCentricCoordinate[3]; - calculateBarycentricCoordinate(vt0, vt1, vt2, vt, baryCentricCoordinate); - if (isBarycentricCoordInBounds(baryCentricCoordinate)) { - int pixel = py * width + px; - if (zbuffer == 0) { - atomicExch(&zbuffer[pixel], (INT64)(idx + 1)); - continue; - } - float depth = baryCentricCoordinate[0] * vt0[2] + baryCentricCoordinate[1] * vt1[2] + baryCentricCoordinate[2] * vt2[2]; - float depth_thres = 0; - if (d) { - depth_thres = d[pixel] * 0.49999f + 0.5f + occlusion_truncation; - } - - int z_quantize = depth * (2<<17); - INT64 token = (INT64)z_quantize * MAXINT + (INT64)(idx + 1); - if (depth < depth_thres) - continue; - atomicMin(&zbuffer[pixel], token); - } - } - } -} - -__global__ void barycentricFromImgcoordGPU(float* V, int* F, int* findices, INT64* zbuffer, int width, int height, int num_vertices, int num_faces, - float* barycentric_map) -{ - int pix = blockIdx.x * blockDim.x + threadIdx.x; - if (pix >= width * height) - return; - INT64 f = zbuffer[pix] % MAXINT; - if (f == (MAXINT-1)) { - findices[pix] = 0; - barycentric_map[pix * 3] = 0; - barycentric_map[pix * 3 + 1] = 0; - barycentric_map[pix * 3 + 2] = 0; - return; - } - findices[pix] = f; - f -= 1; - float barycentric[3] = {0, 0, 0}; - if (f >= 0) { - float vt[2] = {float(pix % width) + 0.5f, float(pix / width) + 0.5f}; - float* vt0_ptr = V + (F[f * 3] * 4); - float* vt1_ptr = V + (F[f * 3 + 1] * 4); - float* vt2_ptr = V + (F[f * 3 + 2] * 4); - - float vt0[2] = {(vt0_ptr[0] / vt0_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, (0.5f + 0.5f * vt0_ptr[1] / vt0_ptr[3]) * (height - 1) + 0.5f}; - float vt1[2] = {(vt1_ptr[0] / vt1_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, (0.5f + 0.5f * vt1_ptr[1] / vt1_ptr[3]) * (height - 1) + 0.5f}; - float vt2[2] = {(vt2_ptr[0] / vt2_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, (0.5f + 0.5f * vt2_ptr[1] / vt2_ptr[3]) * (height - 1) + 0.5f}; - - calculateBarycentricCoordinate(vt0, vt1, vt2, vt, barycentric); - - barycentric[0] = barycentric[0] / vt0_ptr[3]; - barycentric[1] = barycentric[1] / vt1_ptr[3]; - barycentric[2] = barycentric[2] / vt2_ptr[3]; - float w = 1.0f / (barycentric[0] + barycentric[1] + barycentric[2]); - barycentric[0] *= w; - barycentric[1] *= w; - barycentric[2] *= w; - } - barycentric_map[pix * 3] = barycentric[0]; - barycentric_map[pix * 3 + 1] = barycentric[1]; - barycentric_map[pix * 3 + 2] = barycentric[2]; -} - -__global__ void rasterizeImagecoordsKernelGPU(float* V, int* F, float* d, INT64* zbuffer, float occlusion_trunc, int width, int height, int num_vertices, int num_faces) -{ - int f = blockIdx.x * blockDim.x + threadIdx.x; - if (f >= num_faces) - return; - - float* vt0_ptr = V + (F[f * 3] * 4); - float* vt1_ptr = V + (F[f * 3 + 1] * 4); - float* vt2_ptr = V + (F[f * 3 + 2] * 4); - - float vt0[3] = {(vt0_ptr[0] / vt0_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, (0.5f + 0.5f * vt0_ptr[1] / vt0_ptr[3]) * (height - 1) + 0.5f, vt0_ptr[2] / vt0_ptr[3] * 0.49999f + 0.5f}; - float vt1[3] = {(vt1_ptr[0] / vt1_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, (0.5f + 0.5f * vt1_ptr[1] / vt1_ptr[3]) * (height - 1) + 0.5f, vt1_ptr[2] / vt1_ptr[3] * 0.49999f + 0.5f}; - float vt2[3] = {(vt2_ptr[0] / vt2_ptr[3] * 0.5f + 0.5f) * (width - 1) + 0.5f, (0.5f + 0.5f * vt2_ptr[1] / vt2_ptr[3]) * (height - 1) + 0.5f, vt2_ptr[2] / vt2_ptr[3] * 0.49999f + 0.5f}; - - rasterizeTriangleGPU(f, vt0, vt1, vt2, width, height, zbuffer, d, occlusion_trunc); -} - -std::vector rasterize_image_gpu(torch::Tensor V, torch::Tensor F, torch::Tensor D, - int width, int height, float occlusion_truncation, int use_depth_prior) -{ - int device_id = V.get_device(); - cudaSetDevice(device_id); - int num_faces = F.size(0); - int num_vertices = V.size(0); - auto options = torch::TensorOptions().dtype(torch::kInt32).device(torch::kCUDA, device_id).requires_grad(false); - auto INT64_options = torch::TensorOptions().dtype(torch::kInt64).device(torch::kCUDA, device_id).requires_grad(false); - auto findices = torch::zeros({height, width}, options); - INT64 maxint = (INT64)MAXINT * (INT64)MAXINT + (MAXINT - 1); - auto z_min = torch::ones({height, width}, INT64_options) * (int64_t)maxint; - - if (!use_depth_prior) { - rasterizeImagecoordsKernelGPU<<<(num_faces+255)/256,256,0,at::cuda::getCurrentCUDAStream()>>>(V.data_ptr(), F.data_ptr(), 0, - (INT64*)z_min.data_ptr(), occlusion_truncation, width, height, num_vertices, num_faces); - } else { - rasterizeImagecoordsKernelGPU<<<(num_faces+255)/256,256,0,at::cuda::getCurrentCUDAStream()>>>(V.data_ptr(), F.data_ptr(), D.data_ptr(), - (INT64*)z_min.data_ptr(), occlusion_truncation, width, height, num_vertices, num_faces); - } - - auto float_options = torch::TensorOptions().dtype(torch::kFloat32).device(torch::kCUDA, device_id).requires_grad(false); - auto barycentric = torch::zeros({height, width, 3}, float_options); - barycentricFromImgcoordGPU<<<(width * height + 255)/256, 256, 0, at::cuda::getCurrentCUDAStream()>>>(V.data_ptr(), F.data_ptr(), - findices.data_ptr(), (INT64*)z_min.data_ptr(), width, height, num_vertices, num_faces, barycentric.data_ptr()); - - return {findices, barycentric}; -} diff --git a/python/sglang/kernels/ops/diffusion/render/mesh_processor/mesh_processor.cpp b/python/sglang/kernels/ops/diffusion/render/mesh_processor/mesh_processor.cpp deleted file mode 100644 index 1ce0d35c2853..000000000000 --- a/python/sglang/kernels/ops/diffusion/render/mesh_processor/mesh_processor.cpp +++ /dev/null @@ -1,163 +0,0 @@ -// SPDX-License-Identifier: Apache-2.0 -// Adapted from Hunyuan3D-2: https://github.com/Tencent/Hunyuan3D-2 -// Original license: TENCENT HUNYUAN NON-COMMERCIAL LICENSE AGREEMENT - -#include -#include -#include -#include -#include -#include -#include -#include - -namespace py = pybind11; -using namespace std; - -std::pair, - py::array_t> meshVerticeInpaint_smooth(py::array_t texture, -py::array_t mask, - py::array_t vtx_pos, py::array_t vtx_uv, - py::array_t pos_idx, py::array_t uv_idx) { - auto texture_buf = texture.request(); - auto mask_buf = mask.request(); - auto vtx_pos_buf = vtx_pos.request(); - auto vtx_uv_buf = vtx_uv.request(); - auto pos_idx_buf = pos_idx.request(); - auto uv_idx_buf = uv_idx.request(); - - int texture_height = texture_buf.shape[0]; - int texture_width = texture_buf.shape[1]; - int texture_channel = texture_buf.shape[2]; - float* texture_ptr = static_cast(texture_buf.ptr); - uint8_t* mask_ptr = static_cast(mask_buf.ptr); - - int vtx_num = vtx_pos_buf.shape[0]; - float* vtx_pos_ptr = static_cast(vtx_pos_buf.ptr); - float* vtx_uv_ptr = static_cast(vtx_uv_buf.ptr); - int* pos_idx_ptr = static_cast(pos_idx_buf.ptr); - int* uv_idx_ptr = static_cast(uv_idx_buf.ptr); - - vector vtx_mask(vtx_num, 0.0f); - vector> vtx_color(vtx_num, vector(texture_channel, 0.0f)); - vector uncolored_vtxs; - - vector> G(vtx_num); - - for (int i = 0; i < uv_idx_buf.shape[0]; ++i) { - for (int k = 0; k < 3; ++k) { - int vtx_uv_idx = uv_idx_ptr[i * 3 + k]; - int vtx_idx = pos_idx_ptr[i * 3 + k]; - int uv_v = round(vtx_uv_ptr[vtx_uv_idx * 2] * (texture_width - 1)); - int uv_u = round((1.0 - vtx_uv_ptr[vtx_uv_idx * 2 + 1]) * (texture_height - 1)); - - if (mask_ptr[uv_u * texture_width + uv_v] > 0) { - vtx_mask[vtx_idx] = 1.0f; - for (int c = 0; c < texture_channel; ++c) { - vtx_color[vtx_idx][c] = texture_ptr[(uv_u * texture_width + uv_v) * texture_channel + c]; - } - }else{ - uncolored_vtxs.push_back(vtx_idx); - } - - G[pos_idx_ptr[i * 3 + k]].push_back(pos_idx_ptr[i * 3 + (k + 1) % 3]); - } - } - - int smooth_count = 2; - int last_uncolored_vtx_count = 0; - while (smooth_count>0) { - int uncolored_vtx_count = 0; - - for (int vtx_idx : uncolored_vtxs) { - - vector sum_color(texture_channel, 0.0f); - float total_weight = 0.0f; - - array vtx_0 = {vtx_pos_ptr[vtx_idx * 3], -vtx_pos_ptr[vtx_idx * 3 + 1], vtx_pos_ptr[vtx_idx * 3 + 2]}; - for (int connected_idx : G[vtx_idx]) { - if (vtx_mask[connected_idx] > 0) { - array vtx1 = {vtx_pos_ptr[connected_idx * 3], - vtx_pos_ptr[connected_idx * 3 + 1], vtx_pos_ptr[connected_idx * 3 + 2]}; - float dist_weight = 1.0f / max(sqrt(pow(vtx_0[0] - vtx1[0], 2) + pow(vtx_0[1] - vtx1[1], 2) + \ - pow(vtx_0[2] - vtx1[2], 2)), 1E-4); - dist_weight = dist_weight * dist_weight; - for (int c = 0; c < texture_channel; ++c) { - sum_color[c] += vtx_color[connected_idx][c] * dist_weight; - } - total_weight += dist_weight; - } - } - - if (total_weight > 0.0f) { - for (int c = 0; c < texture_channel; ++c) { - vtx_color[vtx_idx][c] = sum_color[c] / total_weight; - } - vtx_mask[vtx_idx] = 1.0f; - } else { - uncolored_vtx_count++; - } - - } - - if(last_uncolored_vtx_count==uncolored_vtx_count){ - smooth_count--; - }else{ - smooth_count++; - } - last_uncolored_vtx_count = uncolored_vtx_count; - } - - py::array_t new_texture(texture_buf.size); - py::array_t new_mask(mask_buf.size); - - auto new_texture_buf = new_texture.request(); - auto new_mask_buf = new_mask.request(); - - float* new_texture_ptr = static_cast(new_texture_buf.ptr); - uint8_t* new_mask_ptr = static_cast(new_mask_buf.ptr); - std::copy(texture_ptr, texture_ptr + texture_buf.size, new_texture_ptr); - std::copy(mask_ptr, mask_ptr + mask_buf.size, new_mask_ptr); - - for (int face_idx = 0; face_idx < uv_idx_buf.shape[0]; ++face_idx) { - for (int k = 0; k < 3; ++k) { - int vtx_uv_idx = uv_idx_ptr[face_idx * 3 + k]; - int vtx_idx = pos_idx_ptr[face_idx * 3 + k]; - - if (vtx_mask[vtx_idx] == 1.0f) { - int uv_v = round(vtx_uv_ptr[vtx_uv_idx * 2] * (texture_width - 1)); - int uv_u = round((1.0 - vtx_uv_ptr[vtx_uv_idx * 2 + 1]) * (texture_height - 1)); - - for (int c = 0; c < texture_channel; ++c) { - new_texture_ptr[(uv_u * texture_width + uv_v) * texture_channel + c] = vtx_color[vtx_idx][c]; - } - new_mask_ptr[uv_u * texture_width + uv_v] = 255; - } - } - } - - new_texture.resize({texture_height, texture_width, 3}); - new_mask.resize({texture_height, texture_width}); - return std::make_pair(new_texture, new_mask); -} - - -std::pair, py::array_t> meshVerticeInpaint(py::array_t texture, - py::array_t mask, - py::array_t vtx_pos, py::array_t vtx_uv, - py::array_t pos_idx, py::array_t uv_idx, const std::string& method = "smooth") { - if (method == "smooth") { - return meshVerticeInpaint_smooth(texture, mask, vtx_pos, vtx_uv, pos_idx, uv_idx); - } else { - throw std::invalid_argument("Invalid method. Use 'smooth'."); - } -} - -PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { - m.def("meshVerticeInpaint", &meshVerticeInpaint, "Mesh-aware texture inpainting", - py::arg("texture"), py::arg("mask"), - py::arg("vtx_pos"), py::arg("vtx_uv"), - py::arg("pos_idx"), py::arg("uv_idx"), - py::arg("method") = "smooth"); -} diff --git a/python/sglang/kernels/ops/diffusion/rope/__init__.py b/python/sglang/kernels/ops/diffusion/rope/__init__.py new file mode 100644 index 000000000000..72733429df3c --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/rope/__init__.py @@ -0,0 +1 @@ +"""Rotary embeddings and the QK-norm chains fused around them.""" diff --git a/python/sglang/kernels/ops/diffusion/triton/hunyuan_qkv_pack.py b/python/sglang/kernels/ops/diffusion/rope/hunyuan_qkv_pack_triton.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/triton/hunyuan_qkv_pack.py rename to python/sglang/kernels/ops/diffusion/rope/hunyuan_qkv_pack_triton.py diff --git a/python/sglang/kernels/ops/diffusion/ltx2_qknorm_split_rope.py b/python/sglang/kernels/ops/diffusion/rope/ltx2_qknorm_split_rope_jit.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/ltx2_qknorm_split_rope.py rename to python/sglang/kernels/ops/diffusion/rope/ltx2_qknorm_split_rope_jit.py diff --git a/python/sglang/kernels/ops/diffusion/triton/ltx2_rotary.py b/python/sglang/kernels/ops/diffusion/rope/ltx2_rotary_triton.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/triton/ltx2_rotary.py rename to python/sglang/kernels/ops/diffusion/rope/ltx2_rotary_triton.py diff --git a/python/sglang/kernels/ops/diffusion/qknorm_rope.py b/python/sglang/kernels/ops/diffusion/rope/qknorm_rope_jit.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/qknorm_rope.py rename to python/sglang/kernels/ops/diffusion/rope/qknorm_rope_jit.py diff --git a/python/sglang/kernels/ops/diffusion/triton/rope_rotate_half_bitexact.py b/python/sglang/kernels/ops/diffusion/rope/rope_rotate_half_bitexact.py similarity index 98% rename from python/sglang/kernels/ops/diffusion/triton/rope_rotate_half_bitexact.py rename to python/sglang/kernels/ops/diffusion/rope/rope_rotate_half_bitexact.py index b8868fe33e6d..a0d6af11c353 100644 --- a/python/sglang/kernels/ops/diffusion/triton/rope_rotate_half_bitexact.py +++ b/python/sglang/kernels/ops/diffusion/rope/rope_rotate_half_bitexact.py @@ -29,7 +29,7 @@ import triton # type: ignore import triton.language as tl # type: ignore -from sglang.kernels.ops.diffusion.triton.numerics import round_bf16_to_fp32 +from sglang.kernels.ops.diffusion.common.numerics import round_bf16_to_fp32 from sglang.srt.utils.custom_op import register_custom_op diff --git a/python/sglang/kernels/ops/diffusion/triton/rotary.py b/python/sglang/kernels/ops/diffusion/rope/rotary_triton.py similarity index 88% rename from python/sglang/kernels/ops/diffusion/triton/rotary.py rename to python/sglang/kernels/ops/diffusion/rope/rotary_triton.py index 616e31650c45..dc8830b91d54 100644 --- a/python/sglang/kernels/ops/diffusion/triton/rotary.py +++ b/python/sglang/kernels/ops/diffusion/rope/rotary_triton.py @@ -2,7 +2,7 @@ import triton # type: ignore import triton.language as tl # type: ignore -from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.kernels.ops.diffusion.common.platform import lazy_fallback, select_impl @triton.autotune( @@ -125,17 +125,9 @@ def apply_rotary_embedding( return output -if current_platform.is_npu(): - from .npu_fallback import apply_rotary_embedding_native - - apply_rotary_embedding = apply_rotary_embedding_native - -if current_platform.is_mps(): - from .mps_fallback import apply_rotary_embedding_native - - apply_rotary_embedding = apply_rotary_embedding_native - -if current_platform.is_cpu(): - from .torch_fallback import apply_rotary_embedding_native - - apply_rotary_embedding = apply_rotary_embedding_native +apply_rotary_embedding = select_impl( + apply_rotary_embedding, + npu=lazy_fallback("npu", "apply_rotary_embedding_native"), + mps=lazy_fallback("mps", "apply_rotary_embedding_native"), + cpu=lazy_fallback("torch", "apply_rotary_embedding_native"), +) diff --git a/python/sglang/kernels/ops/diffusion/sites/__init__.py b/python/sglang/kernels/ops/diffusion/sites/__init__.py new file mode 100644 index 000000000000..814af374af13 --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/sites/__init__.py @@ -0,0 +1,14 @@ +"""Request-scoped fusion *policy* -- module-tree rewriting, not kernels. + +A fusion whose result is not bit-exact vs the reference chain may not be on by +default: multi-step denoising amplifies per-step rounding differences into +visible quality loss. Such fusions are mounted onto marked ``nn.Module`` sites +only for ``quality="high"`` requests, at batch boundaries, all-or-nothing per +transformer (:mod:`.quality_gate`). Fusions that *are* bit-exact mount +unconditionally but still verify themselves against the live eager chain on +first sight and fall back permanently on mismatch (:mod:`.bitexact_gate`). + +Because these modules inspect and rewrite model modules, they are the one place +in this package allowed to reference ``multimodal_gen`` types, and they do so +lazily inside functions. +""" diff --git a/python/sglang/kernels/ops/diffusion/bitexact_gate.py b/python/sglang/kernels/ops/diffusion/sites/bitexact_gate.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/bitexact_gate.py rename to python/sglang/kernels/ops/diffusion/sites/bitexact_gate.py diff --git a/python/sglang/kernels/ops/diffusion/fused_gate_rmsnorm.py b/python/sglang/kernels/ops/diffusion/sites/fused_gate_rmsnorm_site.py similarity index 93% rename from python/sglang/kernels/ops/diffusion/fused_gate_rmsnorm.py rename to python/sglang/kernels/ops/diffusion/sites/fused_gate_rmsnorm_site.py index 8d210b85c02b..46aae1844b0d 100644 --- a/python/sglang/kernels/ops/diffusion/fused_gate_rmsnorm.py +++ b/python/sglang/kernels/ops/diffusion/sites/fused_gate_rmsnorm_site.py @@ -4,7 +4,7 @@ modulate/gate around each RMSNorm: ``RMSNorm(x) * scale`` before attention/FFN and ``x + tanh(gate) * RMSNorm(out)`` after. Shared BF16-native Triton kernels -(:mod:`sglang.kernels.ops.diffusion.triton.native_bf16_rmsnorm`) fuse each +(:mod:`sglang.kernels.ops.diffusion.norm.native_bf16_rmsnorm_triton`) fuse each chain into a single kernel (RMSNorm + tanh + mul + add in one pass). Z-Image mounts those kernels unconditionally because they reproduce its own @@ -29,7 +29,7 @@ import torch import torch.nn as nn -from sglang.kernels.ops.diffusion.quality_gate import QualityGatedFusion +from sglang.kernels.ops.diffusion.sites.quality_gate import QualityGatedFusion logger = logging.getLogger(__name__) @@ -50,7 +50,7 @@ def fused_rmsnorm_scale( x: torch.Tensor, weight: torch.Tensor, scale: torch.Tensor, eps: float ) -> torch.Tensor | None: """``RMSNorm(x, weight, eps) * scale`` in one Triton kernel (or None).""" - from sglang.kernels.ops.diffusion.triton.native_bf16_rmsnorm import ( + from sglang.kernels.ops.diffusion.norm.native_bf16_rmsnorm_triton import ( rmsnorm_scale, ) @@ -65,7 +65,7 @@ def fused_rmsnorm_tanh_residual( eps: float, ) -> torch.Tensor | None: """``residual + tanh(gate) * RMSNorm(x, weight, eps)`` fused (or None).""" - from sglang.kernels.ops.diffusion.triton.native_bf16_rmsnorm import ( + from sglang.kernels.ops.diffusion.norm.native_bf16_rmsnorm_triton import ( rmsnorm_tanh_residual, ) diff --git a/python/sglang/kernels/ops/diffusion/fused_linear_gelu.py b/python/sglang/kernels/ops/diffusion/sites/fused_linear_gelu_site.py similarity index 96% rename from python/sglang/kernels/ops/diffusion/fused_linear_gelu.py rename to python/sglang/kernels/ops/diffusion/sites/fused_linear_gelu_site.py index 59a678746226..e24be82fae01 100644 --- a/python/sglang/kernels/ops/diffusion/fused_linear_gelu.py +++ b/python/sglang/kernels/ops/diffusion/sites/fused_linear_gelu_site.py @@ -32,7 +32,7 @@ import torch.nn as nn from sglang.kernels.jit.utils import get_jit_cuda_arch -from sglang.kernels.ops.diffusion.quality_gate import QualityGatedFusion +from sglang.kernels.ops.diffusion.sites.quality_gate import QualityGatedFusion from sglang.srt.utils.custom_op import register_custom_op logger = logging.getLogger(__name__) @@ -123,12 +123,12 @@ def _static_reject_reason(linear: Any) -> str | None: return None -def can_fuse_linear_gelu_static(linear: Any) -> bool: +def can_use_linear_gelu_static(linear: Any) -> bool: """Input-independent guards: whether ``linear`` may ever use the epilogue.""" return _static_reject_reason(linear) is None -def can_fuse_linear_gelu(linear: Any, x: torch.Tensor) -> bool: +def can_use_linear_gelu(linear: Any, x: torch.Tensor) -> bool: """Whether ``gelu(linear(x))`` can use the fused cublasLt epilogue now.""" if not (x.is_cuda and x.dtype in (torch.bfloat16, torch.float16)): return False @@ -141,7 +141,7 @@ def can_fuse_linear_gelu(linear: Any, x: torch.Tensor) -> bool: return False if getattr(linear, "weight", None) is None or x.dtype != linear.weight.dtype: return False - return can_fuse_linear_gelu_static(linear) + return can_use_linear_gelu_static(linear) def mark_fused_gelu_site(module: nn.Module, linear_attr: str) -> None: diff --git a/python/sglang/kernels/ops/diffusion/fused_ln_modulate.py b/python/sglang/kernels/ops/diffusion/sites/fused_ln_modulate_site.py similarity index 95% rename from python/sglang/kernels/ops/diffusion/fused_ln_modulate.py rename to python/sglang/kernels/ops/diffusion/sites/fused_ln_modulate_site.py index 29a75699f51c..a818309785ae 100644 --- a/python/sglang/kernels/ops/diffusion/fused_ln_modulate.py +++ b/python/sglang/kernels/ops/diffusion/sites/fused_ln_modulate_site.py @@ -20,7 +20,7 @@ import torch.nn.functional as F from torch import nn -from sglang.kernels.ops.diffusion.quality_gate import QualityGatedFusion +from sglang.kernels.ops.diffusion.sites.quality_gate import QualityGatedFusion _SITE_MARKER_ATTR = "_sgl_fused_ln_modulate_site" _SITE_ENABLED_ATTR = "_sgl_fused_ln_modulate_enabled" @@ -50,7 +50,7 @@ def unmount_fused_ln_modulate(root: nn.Module) -> None: _FUSION.unmount(root) -def can_fuse_ln_modulate( +def can_use_ln_modulate( x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor ) -> bool: """Per-call guard: the folded affine is a [D] row, so batch must be 1.""" diff --git a/python/sglang/kernels/ops/diffusion/hunyuan_qknorm.py b/python/sglang/kernels/ops/diffusion/sites/hunyuan_qknorm_site.py similarity index 96% rename from python/sglang/kernels/ops/diffusion/hunyuan_qknorm.py rename to python/sglang/kernels/ops/diffusion/sites/hunyuan_qknorm_site.py index ffd8716c2f8f..12b7e75f4417 100644 --- a/python/sglang/kernels/ops/diffusion/hunyuan_qknorm.py +++ b/python/sglang/kernels/ops/diffusion/sites/hunyuan_qknorm_site.py @@ -8,7 +8,7 @@ import torch import torch.nn as nn -from sglang.kernels.ops.diffusion.quality_gate import QualityGatedFusion +from sglang.kernels.ops.diffusion.sites.quality_gate import QualityGatedFusion logger = logging.getLogger(__name__) diff --git a/python/sglang/kernels/ops/diffusion/ltx2_rmsnorm_modulate.py b/python/sglang/kernels/ops/diffusion/sites/ltx2_rmsnorm_modulate_site.py similarity index 93% rename from python/sglang/kernels/ops/diffusion/ltx2_rmsnorm_modulate.py rename to python/sglang/kernels/ops/diffusion/sites/ltx2_rmsnorm_modulate_site.py index f2d36ec7e390..1bad9f14bf1f 100644 --- a/python/sglang/kernels/ops/diffusion/ltx2_rmsnorm_modulate.py +++ b/python/sglang/kernels/ops/diffusion/sites/ltx2_rmsnorm_modulate_site.py @@ -21,11 +21,11 @@ import torch from torch import nn -from sglang.kernels.ops.diffusion.quality_gate import QualityGatedFusion -from sglang.kernels.ops.diffusion.triton.rmsnorm_scale_shift_bitexact import ( +from sglang.kernels.ops.diffusion.norm.rmsnorm_scale_shift_bitexact import ( can_use_fused_rmsnorm_scale_shift, fused_rmsnorm_scale_shift_bitexact, ) +from sglang.kernels.ops.diffusion.sites.quality_gate import QualityGatedFusion _SITE_MARKER_ATTR = "_sgl_ltx2_rms_norm_modulate_site" _SITE_ENABLED_ATTR = "_sgl_ltx2_rms_norm_modulate_enabled" @@ -65,7 +65,7 @@ def _ones_weight(x: torch.Tensor) -> torch.Tensor: return w -def can_fuse_ltx2_rms_norm_modulate( +def can_use_ltx2_rms_norm_modulate( x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor ) -> bool: if x.dtype is not torch.bfloat16 or not x.is_cuda: diff --git a/python/sglang/kernels/ops/diffusion/quality_gate.py b/python/sglang/kernels/ops/diffusion/sites/quality_gate.py similarity index 100% rename from python/sglang/kernels/ops/diffusion/quality_gate.py rename to python/sglang/kernels/ops/diffusion/sites/quality_gate.py diff --git a/python/sglang/kernels/spec.py b/python/sglang/kernels/spec.py index dba590ac1614..d1220a35ef97 100644 --- a/python/sglang/kernels/spec.py +++ b/python/sglang/kernels/spec.py @@ -43,6 +43,7 @@ class KernelBackend(str, Enum): JIT = "jit" # sglang.kernels.jit (nvcc / hipcc) AOT = "aot" # sgl_kernel wheel (CUDA / ROCm builds) CUTE_DSL = "cute_dsl" + FLYDSL = "flydsl" # FlyDSL MLIR compiler (device=HIP, gfx950) FLASHINFER = "flashinfer" DEEPGEMM = "deepgemm" AITER = "aiter" # AMD aiter library (device=HIP) diff --git a/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-benchmark-profile/existing-fast-paths.md b/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-benchmark-profile/existing-fast-paths.md index c558acf51519..6b6f5590458c 100644 --- a/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-benchmark-profile/existing-fast-paths.md +++ b/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-benchmark-profile/existing-fast-paths.md @@ -10,33 +10,33 @@ framework-specific optimization workflow. - `python/sglang/multimodal_gen/runtime/layers/elementwise.py` - `python/sglang/multimodal_gen/runtime/layers/fused_scale_shift_gate.py` - `python/sglang/multimodal_gen/runtime/layers/rotary_embedding/utils.py` -- `python/sglang/kernels/ops/diffusion/triton/scale_shift.py` -- `python/sglang/kernels/ops/diffusion/modulate_scale_shift.py` -- `python/sglang/kernels/ops/diffusion/fused_ln_modulate.py` -- `python/sglang/kernels/ops/diffusion/quality_gate.py` -- `python/sglang/kernels/ops/diffusion/bitexact_gate.py` -- `python/sglang/kernels/ops/diffusion/group_norm_silu.py` -- `python/sglang/kernels/ops/diffusion/triton/group_norm_silu.py` -- `python/sglang/kernels/ops/diffusion/triton/group_norm_silu_twopass.py` -- `python/sglang/kernels/ops/diffusion/triton/norm.py` -- `python/sglang/kernels/ops/diffusion/triton/rmsnorm_onepass.py` -- `python/sglang/kernels/ops/diffusion/triton/layernorm_modulate.py` -- `python/sglang/kernels/ops/diffusion/triton/native_bf16_rmsnorm.py` -- `python/sglang/kernels/ops/diffusion/triton/zimage_native_norm.py` -- `python/sglang/kernels/ops/diffusion/triton/rotary.py` -- `python/sglang/kernels/ops/diffusion/triton/ltx2_rotary.py` -- `python/sglang/kernels/ops/diffusion/ltx2_qknorm_split_rope.py` -- `python/sglang/kernels/ops/diffusion/ltx2_rmsnorm_modulate.py` -- `python/sglang/kernels/ops/diffusion/triton/indexed_modulation.py` -- `python/sglang/kernels/ops/diffusion/triton/ulysses_qkv.py` -- `python/sglang/kernels/ops/diffusion/usp_relayout.py` +- `python/sglang/kernels/ops/diffusion/modulate/scale_shift_triton.py` +- `python/sglang/kernels/ops/diffusion/modulate/modulate_scale_shift_jit.py` +- `python/sglang/kernels/ops/diffusion/sites/fused_ln_modulate_site.py` +- `python/sglang/kernels/ops/diffusion/sites/quality_gate.py` +- `python/sglang/kernels/ops/diffusion/sites/bitexact_gate.py` +- `python/sglang/kernels/ops/diffusion/norm/group_norm_silu.py` +- `python/sglang/kernels/ops/diffusion/norm/group_norm_silu_triton.py` +- `python/sglang/kernels/ops/diffusion/norm/group_norm_silu_twopass_triton.py` +- `python/sglang/kernels/ops/diffusion/norm/norm_triton.py` +- `python/sglang/kernels/ops/diffusion/norm/rmsnorm_onepass_triton.py` +- `python/sglang/kernels/ops/diffusion/norm/layernorm_modulate_triton.py` +- `python/sglang/kernels/ops/diffusion/norm/native_bf16_rmsnorm_triton.py` +- `python/sglang/kernels/ops/diffusion/norm/zimage_qk_rmsnorm_triton.py` +- `python/sglang/kernels/ops/diffusion/rope/rotary_triton.py` +- `python/sglang/kernels/ops/diffusion/rope/ltx2_rotary_triton.py` +- `python/sglang/kernels/ops/diffusion/rope/ltx2_qknorm_split_rope_jit.py` +- `python/sglang/kernels/ops/diffusion/sites/ltx2_rmsnorm_modulate_site.py` +- `python/sglang/kernels/ops/diffusion/modulate/indexed_modulation_triton.py` +- `python/sglang/kernels/ops/diffusion/layout/ulysses_qkv_triton.py` +- `python/sglang/kernels/ops/diffusion/layout/usp_relayout_jit.py` - `python/sglang/multimodal_gen/runtime/layers/usp.py` - `python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py` -- `python/sglang/kernels/ops/diffusion/residual_gate_add.py` +- `python/sglang/kernels/ops/diffusion/modulate/residual_gate_add_jit.py` - `python/sglang/kernels/jit/csrc/diffusion/residual_gate_add.cuh` -- `python/sglang/kernels/ops/diffusion/triton/varlen_pack_pad.py` -- `python/sglang/kernels/ops/diffusion/triton/wan_causal_cache.py` -- `python/sglang/kernels/ops/diffusion/cutedsl/scale_residual_norm_scale_shift.py` +- `python/sglang/kernels/ops/diffusion/layout/varlen_pack_pad_triton.py` +- `python/sglang/kernels/ops/diffusion/layout/wan_causal_cache_triton.py` +- `python/sglang/kernels/ops/diffusion/norm/scale_residual_norm_cutedsl.py` - `python/sglang/multimodal_gen/runtime/models/vaes/fast_path_gate.py` - `python/sglang/multimodal_gen/runtime/models/vaes/flux2_vae_cuda_opt.py` - `python/sglang/multimodal_gen/runtime/models/vaes/wan_vae_cuda_opt.py` @@ -249,7 +249,7 @@ framework-specific optimization workflow. **QK Norm + RoPE Optimization** - Entry point: `apply_qk_norm_rope` in `layernorm.py`. -- Fast path: JIT fused inplace QK norm + RoPE from `python/sglang/kernels/ops/diffusion/qknorm_rope.py` via `fused_inplace_qknorm_rope`. +- Fast path: JIT fused inplace QK norm + RoPE from `python/sglang/kernels/ops/diffusion/rope/qknorm_rope_jit.py` via `fused_inplace_qknorm_rope`. - Toggle: `SGLANG_ENABLE_FUSED_QKNORM_ROPE=1` keeps the fused path enabled by default. - Preconditions for fused path: - CUDA only. @@ -310,12 +310,12 @@ framework-specific optimization workflow. - LTX2 split RoPE: `apply_ltx2_split_rotary_emb` in `ltx_2.py`. - LTX2 RMSNorm+modulate and FFN GELU epilogue under `quality="high"`: `mark_ltx2_rms_norm_modulate_site` / `fused_ltx2_rms_norm_modulate` in - `kernels/ops/diffusion/ltx2_rmsnorm_modulate.py` (mount-based + `kernels/ops/diffusion/sites/ltx2_rmsnorm_modulate_site.py` (mount-based `QualityGatedFusion`, not a first-sight `BitExactFusionGate` — the fused kernel is <=1 ULP off aten, so it is request-gated instead of verified), wired at the six `LTX2TransformerBlock` adaLN sites in `ltx_2.py`. - LTX2 residual-gate add: `ltx_2.py` calls `residual_gate_add` from - `kernels/ops/diffusion/residual_gate_add.py` directly for attention, + `kernels/ops/diffusion/modulate/residual_gate_add_jit.py` directly for attention, cross-attention, and MLP residual updates. - Wan causal VAE: `cat_pad_channels_last_3d` and `dup_up3d_add` in `wanvae.py`, backed by `triton/wan_causal_cache.py`. diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/backends/sparse_linear_attn.py b/python/sglang/multimodal_gen/runtime/layers/attention/backends/sparse_linear_attn.py index a5536390862f..3acae28e650b 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/backends/sparse_linear_attn.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/backends/sparse_linear_attn.py @@ -36,10 +36,7 @@ # ==================================SLA Functions=================================== -from sglang.kernels.ops.diffusion.sparse_linear_attn_kernels import ( - _attn_fwd, - get_block_map, -) +from sglang.kernels.ops.diffusion import _attn_fwd, get_block_map def _get_cuda_arch(device_index: int) -> str: diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py index a46cc6aebe56..0aaadfc4cfac 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py @@ -12,7 +12,7 @@ from torch.nn.attention import SDPBackend, sdpa_kernel from sglang.kernels.ops.attention.flash_attention import flash_attn_varlen_func -from sglang.kernels.ops.diffusion.triton.varlen_pack_pad import ( +from sglang.kernels.ops.diffusion import ( build_inv_indices, fused_pack_qkv, fused_scatter_to_padded, diff --git a/python/sglang/multimodal_gen/runtime/layers/elementwise.py b/python/sglang/multimodal_gen/runtime/layers/elementwise.py index d0c5c0eec4b5..3acc516f48a3 100644 --- a/python/sglang/multimodal_gen/runtime/layers/elementwise.py +++ b/python/sglang/multimodal_gen/runtime/layers/elementwise.py @@ -1,6 +1,6 @@ import torch -from sglang.kernels.ops.diffusion.triton.scale_shift import fuse_scale_shift_kernel +from sglang.kernels.ops.diffusion import fuse_scale_shift_kernel from sglang.multimodal_gen.runtime.layers.custom_op import CustomOp diff --git a/python/sglang/multimodal_gen/runtime/layers/fused_scale_shift_gate.py b/python/sglang/multimodal_gen/runtime/layers/fused_scale_shift_gate.py index 6fa1703a06b9..95bf35742ea4 100644 --- a/python/sglang/multimodal_gen/runtime/layers/fused_scale_shift_gate.py +++ b/python/sglang/multimodal_gen/runtime/layers/fused_scale_shift_gate.py @@ -10,7 +10,7 @@ _is_cuda = current_platform.is_cuda() if _is_cuda: - from sglang.kernels.ops.diffusion.triton.scale_shift import ( + from sglang.kernels.ops.diffusion import ( fuse_layernorm_scale_shift_gate_select01_kernel, fuse_residual_layernorm_scale_shift_gate_select01_kernel, ) diff --git a/python/sglang/multimodal_gen/runtime/layers/layernorm.py b/python/sglang/multimodal_gen/runtime/layers/layernorm.py index c89a9ad3bd56..fe47a9fab942 100755 --- a/python/sglang/multimodal_gen/runtime/layers/layernorm.py +++ b/python/sglang/multimodal_gen/runtime/layers/layernorm.py @@ -11,12 +11,12 @@ import torch.nn as nn import torch.nn.functional as F -from sglang.kernels.ops.diffusion.qknorm_rope import ( +from sglang.kernels.ops.diffusion import ( can_use_fused_inplace_qknorm_rope, + fuse_scale_shift_kernel, fused_inplace_qknorm_rope, + triton_one_pass_rms_norm, ) -from sglang.kernels.ops.diffusion.triton.rmsnorm_onepass import triton_one_pass_rms_norm -from sglang.kernels.ops.diffusion.triton.scale_shift import fuse_scale_shift_kernel from sglang.kernels.ops.layernorm.norm import ( can_use_fused_inplace_qknorm, fused_inplace_qknorm, @@ -58,7 +58,7 @@ from sgl_kernel import fused_inplace_qknorm_rope if not _is_cpu: - from sglang.kernels.ops.diffusion.triton.norm import norm_infer, rms_norm_fn + from sglang.kernels.ops.diffusion import norm_infer, rms_norm_fn # Copied and adapted from sglang @@ -614,9 +614,7 @@ def forward_cuda( ) return self.forward_native(residual, x, gate, shift, scale) - from sglang.kernels.ops.diffusion.cutedsl.scale_residual_norm_scale_shift import ( - fused_scale_residual_norm_scale_shift, - ) + from sglang.kernels.ops.diffusion import fused_scale_residual_norm_scale_shift if isinstance(gate, int) and gate != 1: raise ValueError( @@ -647,7 +645,7 @@ def forward_hip( return self.forward_native(residual, x, gate, shift, scale) try: - from sglang.kernels.ops.diffusion.flydsl.fused_residual_norm import ( + from sglang.kernels.ops.diffusion import ( FLYDSL_NORM_MIN_ALIGNED_DIM, flydsl_fused_residual_norm_scale_shift, ) @@ -792,9 +790,7 @@ def forward_cuda( ) return self.forward_native(x, shift, scale) - from sglang.kernels.ops.diffusion.cutedsl.scale_residual_norm_scale_shift import ( - fused_norm_scale_shift, - ) + from sglang.kernels.ops.diffusion import fused_norm_scale_shift return fused_norm_scale_shift( x.contiguous(), @@ -816,7 +812,7 @@ def forward_hip( return self.forward_native(x, shift, scale) try: - from sglang.kernels.ops.diffusion.flydsl.fused_residual_norm import ( + from sglang.kernels.ops.diffusion import ( FLYDSL_NORM_MIN_ALIGNED_DIM, flydsl_norm_scale_shift, ) diff --git a/python/sglang/multimodal_gen/runtime/layers/parallel_conv.py b/python/sglang/multimodal_gen/runtime/layers/parallel_conv.py index 2c009487371b..baf49b964e60 100644 --- a/python/sglang/multimodal_gen/runtime/layers/parallel_conv.py +++ b/python/sglang/multimodal_gen/runtime/layers/parallel_conv.py @@ -18,13 +18,15 @@ logger = init_logger(__name__) if current_platform.is_cuda(): - from sglang.kernels.ops.diffusion.causal_conv3d_cat_pad import ( + from sglang.kernels.ops.diffusion import ( can_use_fused_causal_conv3d_cat_pad_cuda, - fused_causal_conv3d_cat_pad_cuda, ) - from sglang.kernels.ops.diffusion.triton.causal_conv3d_pad import ( + from sglang.kernels.ops.diffusion import ( fused_causal_conv3d_cat_pad as fused_causal_conv3d_cat_pad_triton, ) + from sglang.kernels.ops.diffusion import ( + fused_causal_conv3d_cat_pad_cuda, + ) else: can_use_fused_causal_conv3d_cat_pad_cuda = None fused_causal_conv3d_cat_pad_cuda = None diff --git a/python/sglang/multimodal_gen/runtime/layers/rotary_embedding/utils.py b/python/sglang/multimodal_gen/runtime/layers/rotary_embedding/utils.py index 9fa7d8de1155..89462c286a23 100644 --- a/python/sglang/multimodal_gen/runtime/layers/rotary_embedding/utils.py +++ b/python/sglang/multimodal_gen/runtime/layers/rotary_embedding/utils.py @@ -5,7 +5,7 @@ import torch from sglang.kernel_api_logging import debug_kernel_api -from sglang.kernels.ops.diffusion.triton.rotary import apply_rotary_embedding +from sglang.kernels.ops.diffusion import apply_rotary_embedding from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.srt.utils.custom_op import register_custom_op_from_extern diff --git a/python/sglang/multimodal_gen/runtime/layers/usp.py b/python/sglang/multimodal_gen/runtime/layers/usp.py index 1ca6619bc1e5..b98550edb5b3 100644 --- a/python/sglang/multimodal_gen/runtime/layers/usp.py +++ b/python/sglang/multimodal_gen/runtime/layers/usp.py @@ -9,10 +9,7 @@ from torch.distributed.tensor.experimental._attention import _cp_options from sglang.kernels.ops.attention.flash_attention import flash_attn_varlen_func -from sglang.kernels.ops.diffusion.triton.ulysses_qkv import ( - pack_qkv_destination_major, -) -from sglang.kernels.ops.diffusion.usp_relayout import usp_merge_heads +from sglang.kernels.ops.diffusion import pack_qkv_destination_major, usp_merge_heads from sglang.multimodal_gen.runtime.distributed.parallel_state import ( get_ring_ctx, get_sp_group, diff --git a/python/sglang/multimodal_gen/runtime/layers/visual_embedding.py b/python/sglang/multimodal_gen/runtime/layers/visual_embedding.py index 1a06bca78895..64efd15cbf9b 100644 --- a/python/sglang/multimodal_gen/runtime/layers/visual_embedding.py +++ b/python/sglang/multimodal_gen/runtime/layers/visual_embedding.py @@ -22,9 +22,7 @@ get_timestep_embedding as timestep_embedding_diffusers, ) -from sglang.kernels.ops.diffusion.timestep_embedding import ( - timestep_embedding as timestep_embedding_cuda, -) +from sglang.kernels.ops.diffusion import timestep_embedding as timestep_embedding_cuda from sglang.multimodal_gen.runtime.layers.activation import get_act_fn from sglang.multimodal_gen.runtime.layers.linear import ColumnParallelLinear from sglang.multimodal_gen.runtime.layers.mlp import MLP diff --git a/python/sglang/multimodal_gen/runtime/models/dits/cosmos3video.py b/python/sglang/multimodal_gen/runtime/models/dits/cosmos3video.py index 9f76363ea8ff..1971d1e26846 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/cosmos3video.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/cosmos3video.py @@ -14,7 +14,7 @@ import torch.nn as nn import torch.nn.functional as F -from sglang.kernels.ops.diffusion.qknorm_rope import ( +from sglang.kernels.ops.diffusion import ( can_use_fused_inplace_qknorm_rope, fused_qknorm_rope_pack_kv, ) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/ernie_image.py b/python/sglang/multimodal_gen/runtime/models/dits/ernie_image.py index 721e1ffe3d0d..3b6ebd69361d 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/ernie_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/ernie_image.py @@ -22,21 +22,17 @@ from sglang.kernels.ops.activation.activation import ( gelu_and_mul_with_activation_rounding, ) -from sglang.kernels.ops.diffusion.bitexact_gate import ( +from sglang.kernels.ops.diffusion import ( BitExactFusionGate, - flashinfer_rmsnorm_diagnostic_hint, - tensors_equal, -) -from sglang.kernels.ops.diffusion.residual_gate_add import residual_gate_add -from sglang.kernels.ops.diffusion.triton.rmsnorm_scale_shift_bitexact import ( can_use_fused_rmsnorm_scale_shift, + can_use_fused_rope_rotate_half, can_use_fused_scale_residual_rmsnorm_scale_shift, + flashinfer_rmsnorm_diagnostic_hint, fused_rmsnorm_scale_shift_bitexact, - fused_scale_residual_rmsnorm_scale_shift_bitexact, -) -from sglang.kernels.ops.diffusion.triton.rope_rotate_half_bitexact import ( - can_use_fused_rope_rotate_half, fused_rope_rotate_half_bitexact, + fused_scale_residual_rmsnorm_scale_shift_bitexact, + residual_gate_add, + tensors_equal, ) from sglang.multimodal_gen.configs.models.dits.ernie_image import ( ErnieImageDitConfig, diff --git a/python/sglang/multimodal_gen/runtime/models/dits/flux.py b/python/sglang/multimodal_gen/runtime/models/dits/flux.py index 0035ca821a79..1e328fe91490 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/flux.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/flux.py @@ -28,25 +28,21 @@ ) from torch.nn import LayerNorm as LayerNorm -from sglang.kernels.ops.diffusion.bitexact_gate import BitExactFusionGate -from sglang.kernels.ops.diffusion.fused_linear_gelu import ( - can_fuse_linear_gelu, +from sglang.kernels.ops.diffusion import ( + BitExactFusionGate, + can_use_fused_layernorm_modulate, + can_use_linear_gelu, + can_use_ln_modulate, fused_gelu_active, + fused_layernorm_modulate, fused_linear_gelu_tanh, - mark_fused_gelu_site, -) -from sglang.kernels.ops.diffusion.fused_ln_modulate import ( - can_fuse_ln_modulate, fused_ln_modulate, fused_ln_modulate_active, - mark_fused_ln_modulate_site, -) -from sglang.kernels.ops.diffusion.modulate_scale_shift import modulate_scale_shift -from sglang.kernels.ops.diffusion.residual_gate_add import residual_gate_add -from sglang.kernels.ops.diffusion.triton.layernorm_modulate import ( - can_use_fused_layernorm_modulate, - fused_layernorm_modulate, is_plain_layer_norm, + mark_fused_gelu_site, + mark_fused_ln_modulate_site, + modulate_scale_shift, + residual_gate_add, ) from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig from sglang.multimodal_gen.runtime.distributed import ( @@ -182,7 +178,7 @@ def _flux_norm_modulate( out = _flux_fused_ln_modulate(norm, x, scale, shift) if out is not None: return out - if fused_ln_modulate_active(site) and can_fuse_ln_modulate(x, scale, shift): + if fused_ln_modulate_active(site) and can_use_ln_modulate(x, scale, shift): return fused_ln_modulate(x, scale, shift, norm.eps) return modulate_scale_shift(norm(x), scale, shift) @@ -396,7 +392,7 @@ def __init__( mark_fused_gelu_site(self, "proj") def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if fused_gelu_active(self) and can_fuse_linear_gelu(self.proj, hidden_states): + if fused_gelu_active(self) and can_use_linear_gelu(self.proj, hidden_states): return fused_linear_gelu_tanh( hidden_states, self.proj.weight, self.proj.bias ) @@ -420,7 +416,7 @@ def __init__(self, proj: nn.Linear): mark_fused_gelu_site(self, "proj") def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if fused_gelu_active(self) and can_fuse_linear_gelu(self.proj, hidden_states): + if fused_gelu_active(self) and can_use_linear_gelu(self.proj, hidden_states): return fused_linear_gelu_tanh( hidden_states, self.proj.weight, self.proj.bias ) @@ -896,7 +892,7 @@ def forward( hidden_states = gate * hidden_states hidden_states = residual + hidden_states else: - if fused_gelu_active(self) and can_fuse_linear_gelu( + if fused_gelu_active(self) and can_use_linear_gelu( self.proj_mlp, norm_hidden_states ): mlp_hidden_states = fused_linear_gelu_tanh( diff --git a/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py b/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py index ee64cdfe6d25..50d0561e4bb6 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py @@ -21,15 +21,13 @@ from diffusers.models.embeddings import TimestepEmbedding, Timesteps from diffusers.models.normalization import AdaLayerNormContinuous -from sglang.kernels.ops.diffusion.bitexact_gate import BitExactFusionGate -from sglang.kernels.ops.diffusion.residual_gate_add import residual_gate_add -from sglang.kernels.ops.diffusion.triton.layernorm_modulate import ( +from sglang.kernels.ops.diffusion import ( + BitExactFusionGate, can_use_fused_layernorm_modulate, fused_layernorm_modulate_raw, - is_plain_layer_norm, -) -from sglang.kernels.ops.diffusion.triton.silu_mul_bitexact import ( fused_packed_silu_mul_bitexact, + is_plain_layer_norm, + residual_gate_add, ) from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig from sglang.multimodal_gen.runtime.distributed import ( diff --git a/python/sglang/multimodal_gen/runtime/models/dits/glm_image.py b/python/sglang/multimodal_gen/runtime/models/dits/glm_image.py index 65b3f7263343..45647eb7877f 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/glm_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/glm_image.py @@ -18,23 +18,19 @@ import torch.nn as nn import torch.nn.functional as F -from sglang.kernels.ops.diffusion.bitexact_gate import ( +from sglang.kernels.ops.diffusion import ( BitExactFusionGate, - tensors_equal, -) -from sglang.kernels.ops.diffusion.fused_linear_gelu import ( - can_fuse_linear_gelu, - fused_gelu_active, - fused_linear_gelu_tanh, - mark_fused_gelu_site, -) -from sglang.kernels.ops.diffusion.residual_gate_add import residual_gate_add -from sglang.kernels.ops.diffusion.triton.layernorm_modulate import ( can_use_fused_layernorm_modulate, can_use_fused_qk_head_layernorm, + can_use_linear_gelu, + fused_gelu_active, fused_layernorm_modulate, + fused_linear_gelu_tanh, fused_qk_head_layernorm, is_plain_layer_norm, + mark_fused_gelu_site, + residual_gate_add, + tensors_equal, ) from sglang.multimodal_gen.configs.models.dits.glmimage import GlmImageDitConfig from sglang.multimodal_gen.runtime.distributed.parallel_state import ( @@ -456,7 +452,7 @@ def __init__( mark_fused_gelu_site(self, "proj") def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if fused_gelu_active(self) and can_fuse_linear_gelu(self.proj, hidden_states): + if fused_gelu_active(self) and can_use_linear_gelu(self.proj, hidden_states): return fused_linear_gelu_tanh( hidden_states, self.proj.weight, self.proj.bias ) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/hunyuanvideo.py b/python/sglang/multimodal_gen/runtime/models/dits/hunyuanvideo.py index 5f3bf4317594..acdb6c7f919e 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/hunyuanvideo.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/hunyuanvideo.py @@ -8,23 +8,17 @@ import torch import torch.nn as nn -from sglang.kernels.ops.diffusion.bitexact_gate import ( +from sglang.kernels.ops.diffusion import ( BitExactFusionGate, - tensors_equal, -) -from sglang.kernels.ops.diffusion.fused_linear_gelu import ( - can_fuse_linear_gelu, + can_use_linear_gelu, fused_gelu_active, fused_linear_gelu_tanh, + hunyuan_qkv_rope_pack, mark_fused_gelu_site, -) -from sglang.kernels.ops.diffusion.hunyuan_qknorm import ( mark_hunyuan_qknorm_site, + tensors_equal, try_hunyuan_qknorm, ) -from sglang.kernels.ops.diffusion.triton.hunyuan_qkv_pack import ( - hunyuan_qkv_rope_pack, -) from sglang.multimodal_gen.configs.models.dits import HunyuanVideoConfig from sglang.multimodal_gen.configs.models.fsdp import ( is_double_block, @@ -115,7 +109,7 @@ def __init__(self, *args, **kwargs) -> None: mark_fused_gelu_site(self, "fc_in") def forward(self, x: torch.Tensor) -> torch.Tensor: - if fused_gelu_active(self) and can_fuse_linear_gelu(self.fc_in, x): + if fused_gelu_active(self) and can_use_linear_gelu(self.fc_in, x): x = fused_linear_gelu_tanh(x, self.fc_in.weight, self.fc_in.bias) else: x, _ = self.fc_in(x) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/ideogram.py b/python/sglang/multimodal_gen/runtime/models/dits/ideogram.py index 1380212ea955..30cecf1cbc21 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/ideogram.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/ideogram.py @@ -7,24 +7,18 @@ import torch.nn as nn import torch.nn.functional as F -from sglang.kernels.ops.diffusion.bitexact_gate import ( +from sglang.kernels.ops.diffusion import ( BitExactFusionGate, - tensors_equal, -) -from sglang.kernels.ops.diffusion.fused_gate_rmsnorm import ( + can_use_fused_silu_mul, fused_gate_rmsnorm_active, fused_rmsnorm_scale, fused_rmsnorm_tanh_residual, - mark_fused_gate_rmsnorm_site, -) -from sglang.kernels.ops.diffusion.modulate_scale_shift import modulate_scale_shift -from sglang.kernels.ops.diffusion.residual_gate_add import residual_gate_add -from sglang.kernels.ops.diffusion.triton.rope_rotate_half_bitexact import ( fused_rope_rotate_half_bitexact, -) -from sglang.kernels.ops.diffusion.triton.silu_mul_bitexact import ( - can_use_fused_silu_mul, fused_silu_mul_bitexact, + mark_fused_gate_rmsnorm_site, + modulate_scale_shift, + residual_gate_add, + tensors_equal, ) from sglang.multimodal_gen.configs.models.dits.ideogram import Ideogram4DiTConfig from sglang.multimodal_gen.configs.models.fsdp import is_layer diff --git a/python/sglang/multimodal_gen/runtime/models/dits/krea2.py b/python/sglang/multimodal_gen/runtime/models/dits/krea2.py index 3d88d2a0c804..eeff31c4ed75 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/krea2.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/krea2.py @@ -73,9 +73,7 @@ def _fused_qknorm_rope_enabled() -> bool: def _can_use_fused_qknorm_rope(head_dim: int, dtype: torch.dtype) -> bool: - from sglang.kernels.ops.diffusion.qknorm_rope import ( - can_use_fused_inplace_qknorm_rope, - ) + from sglang.kernels.ops.diffusion import can_use_fused_inplace_qknorm_rope return can_use_fused_inplace_qknorm_rope(head_dim, head_dim, False, dtype) @@ -119,9 +117,7 @@ def norm_scale_shift( pass ``scale + 1``), kept off the checkpoint so the identity load is unaffected. """ if x.is_cuda and x.shape[-1] % 256 == 0 and x.shape[-1] <= 8192: - from sglang.kernels.ops.diffusion.cutedsl.scale_residual_norm_scale_shift import ( - fused_norm_scale_shift, - ) + from sglang.kernels.ops.diffusion import fused_norm_scale_shift return fused_norm_scale_shift( x.contiguous(), @@ -358,9 +354,7 @@ def forward( and _fused_qknorm_rope_enabled() and _can_use_fused_qknorm_rope(hd, q.dtype) ): - from sglang.kernels.ops.diffusion.qknorm_rope import ( - fused_inplace_qknorm_rope, - ) + from sglang.kernels.ops.diffusion import fused_inplace_qknorm_rope b, s = qkv.shape[0], qkv.shape[1] q = q.view(b, s, self.local_heads, hd) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py b/python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py index 1efab045ca32..5850cc0ffc02 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py @@ -10,28 +10,22 @@ import torch.nn as nn import torch.nn.functional as F -from sglang.kernels.ops.diffusion.bitexact_gate import BitExactFusionGate -from sglang.kernels.ops.diffusion.fused_linear_gelu import ( - can_fuse_linear_gelu, +from sglang.kernels.ops.diffusion import ( + BitExactFusionGate, + can_use_linear_gelu, + can_use_ltx2_qknorm_split_rope_cuda, + can_use_ltx2_rms_norm_modulate, + can_use_modulate_scale_shift_cuda, fused_gelu_active, fused_linear_gelu_tanh, - mark_fused_gelu_site, -) -from sglang.kernels.ops.diffusion.ltx2_qknorm_split_rope import ( - can_use_ltx2_qknorm_split_rope_cuda, - ltx2_qknorm_split_rope_cuda, -) -from sglang.kernels.ops.diffusion.ltx2_rmsnorm_modulate import ( - can_fuse_ltx2_rms_norm_modulate, fused_ltx2_rms_norm_modulate, + ltx2_qknorm_split_rope_cuda, ltx2_rms_norm_modulate_active, + mark_fused_gelu_site, mark_ltx2_rms_norm_modulate_site, -) -from sglang.kernels.ops.diffusion.modulate_scale_shift import ( - can_use_modulate_scale_shift_cuda, modulate_scale_shift_cuda, + residual_gate_add, ) -from sglang.kernels.ops.diffusion.residual_gate_add import residual_gate_add from sglang.multimodal_gen.configs.models.dits.ltx_2 import LTX2ArchConfig, LTX2Config from sglang.multimodal_gen.configs.models.fsdp import ( is_blocks_or_transformer_blocks, @@ -209,7 +203,7 @@ def _ltx2_rms_norm_modulate( default). The fused kernel is not bit-exact (<=1 bf16 ULP) so it is gated on the request-scoped mount rather than a runtime self-check. """ - if ltx2_rms_norm_modulate_active(block) and can_fuse_ltx2_rms_norm_modulate( + if ltx2_rms_norm_modulate_active(block) and can_use_ltx2_rms_norm_modulate( x, scale, shift ): return fused_ltx2_rms_norm_modulate(x, scale, shift, eps) @@ -256,9 +250,7 @@ def _ltx2_try_fused_ada_values9( return None try: - from sglang.kernels.ops.diffusion.triton.ltx2_ada_values import ( - ltx2_ada_values9, - ) + from sglang.kernels.ops.diffusion import ltx2_ada_values9 return ltx2_ada_values9(scale_shift_table, timestep) except Exception as exc: @@ -338,9 +330,7 @@ def apply_split_rotary_emb( and cos.is_cuda and sin.is_cuda ): - from sglang.kernels.ops.diffusion.triton.ltx2_rotary import ( - apply_ltx2_split_rotary_emb, - ) + from sglang.kernels.ops.diffusion import apply_ltx2_split_rotary_emb return apply_ltx2_split_rotary_emb(x, cos, sin) @@ -1080,7 +1070,7 @@ def __init__( mark_fused_gelu_site(self, "proj_in") def forward(self, x: torch.Tensor) -> torch.Tensor: - if fused_gelu_active(self) and can_fuse_linear_gelu(self.proj_in, x): + if fused_gelu_active(self) and can_use_linear_gelu(self.proj_in, x): x = fused_linear_gelu_tanh(x, self.proj_in.weight, self.proj_in.bias) else: x, _ = self.proj_in(x) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py b/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py index 1de3f02871c2..22f78948d4e0 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py @@ -20,11 +20,9 @@ from sglang.kernels.ops.activation.activation import ( silu_and_mul_with_activation_rounding_, ) -from sglang.kernels.ops.diffusion.qknorm_rope import ( +from sglang.kernels.ops.diffusion import ( can_use_fused_inplace_qknorm_rope, fused_inplace_qknorm_rope, -) -from sglang.kernels.ops.diffusion.triton.indexed_modulation import ( indexed_gate_bf16, indexed_gate_bf16_, indexed_scale_shift_bf16_, diff --git a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py index a01986bcfefb..593e6fa3adac 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py @@ -14,8 +14,8 @@ from diffusers.models.modeling_outputs import Transformer2DModelOutput from diffusers.models.normalization import AdaLayerNormContinuous -from sglang.kernels.ops.diffusion.fused_linear_gelu import ( - can_fuse_linear_gelu, +from sglang.kernels.ops.diffusion import ( + can_use_linear_gelu, fused_gelu_active, fused_linear_gelu_tanh, mark_fused_gelu_site, @@ -842,7 +842,7 @@ def __init__( mark_fused_gelu_site(self, "proj") def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - if fused_gelu_active(self) and can_fuse_linear_gelu(self.proj, hidden_states): + if fused_gelu_active(self) and can_use_linear_gelu(self.proj, hidden_states): return fused_linear_gelu_tanh( hidden_states, self.proj.weight, self.proj.bias ) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/sana.py b/python/sglang/multimodal_gen/runtime/models/dits/sana.py index cc12993d7c59..470fbdcae9cd 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/sana.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/sana.py @@ -5,18 +5,16 @@ import torch.nn.functional as F from diffusers.models.embeddings import PixArtAlphaTextProjection, TimestepEmbedding -from sglang.kernels.ops.diffusion.bitexact_gate import BitExactFusionGate -from sglang.kernels.ops.diffusion.residual_gate_add import residual_gate_add -from sglang.kernels.ops.diffusion.triton.layernorm_modulate import ( - can_use_fused_layernorm_modulate, - fused_layernorm_modulate_raw, - is_plain_layer_norm, -) -from sglang.kernels.ops.diffusion.triton.sana_conv_post import ( +from sglang.kernels.ops.diffusion import ( + BitExactFusionGate, can_use_fused_bias_glu, can_use_fused_bias_silu, + can_use_fused_layernorm_modulate, fused_bias_glu, fused_bias_silu, + fused_layernorm_modulate_raw, + is_plain_layer_norm, + residual_gate_add, ) from sglang.multimodal_gen.configs.models.dits.sana import SanaConfig from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm diff --git a/python/sglang/multimodal_gen/runtime/models/dits/sana_wm_components.py b/python/sglang/multimodal_gen/runtime/models/dits/sana_wm_components.py index 1b0ec4b858e2..92526ba0b69c 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/sana_wm_components.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/sana_wm_components.py @@ -2137,7 +2137,7 @@ def _maybe_main_branch_triton_gdn( return None try: - from sglang.kernels.ops.diffusion.triton.sana_wm_gdn import ( + from sglang.kernels.ops.diffusion import ( fused_bigdn_func, fused_qk_inv_rms, prepare_rope_tables, @@ -2269,9 +2269,7 @@ def _maybe_cam_branch_triton_scan( return None try: - from sglang.kernels.ops.diffusion.triton.sana_wm_gdn_chunkwise import ( - cam_scan_bidi_chunkwise, - ) + from sglang.kernels.ops.diffusion import cam_scan_bidi_chunkwise B, heads, _, _ = q.shape T, H_sp, W_sp = HW diff --git a/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py b/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py index f62e0a7365ef..088732303fe0 100755 --- a/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py @@ -9,13 +9,11 @@ import torch import torch.nn as nn -from sglang.kernels.ops.diffusion.bitexact_gate import ( +from sglang.kernels.ops.diffusion import ( BitExactFusionGate, - tensors_equal, -) -from sglang.kernels.ops.diffusion.triton.wan_temb_table_slices import ( can_use_fused_temb_table_slices, fused_temb_table_slices, + tensors_equal, ) from sglang.multimodal_gen.configs.models.dits import WanVideoConfig from sglang.multimodal_gen.configs.models.fsdp import is_block diff --git a/python/sglang/multimodal_gen/runtime/models/dits/zimage.py b/python/sglang/multimodal_gen/runtime/models/dits/zimage.py index 2394a8d8a9d8..a3fd410d981b 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/zimage.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/zimage.py @@ -87,9 +87,7 @@ def zimage_rmsnorm_tanh_mul_add( enable_fused: bool = True, ) -> torch.Tensor: if enable_fused: - from sglang.kernels.ops.diffusion.triton.native_bf16_rmsnorm import ( - rmsnorm_tanh_residual, - ) + from sglang.kernels.ops.diffusion import rmsnorm_tanh_residual y = rmsnorm_tanh_residual( x, @@ -110,9 +108,7 @@ def zimage_rmsnorm_scale( enable_fused: bool = True, ) -> torch.Tensor: if enable_fused: - from sglang.kernels.ops.diffusion.triton.native_bf16_rmsnorm import ( - rmsnorm_scale, - ) + from sglang.kernels.ops.diffusion import rmsnorm_scale y = rmsnorm_scale( x, @@ -138,7 +134,7 @@ def zimage_native_qk_rmsnorm( with one Triton launch per tensor that reads the strided fused-qkv slices directly. Returns contiguous (q, k) or None when unsupported. """ - from sglang.kernels.ops.diffusion.triton.zimage_native_norm import ( + from sglang.kernels.ops.diffusion import ( can_use_qk_rmsnorm_native, zimage_qk_rmsnorm_native, ) diff --git a/python/sglang/multimodal_gen/runtime/models/upsampler/latent_upsampler.py b/python/sglang/multimodal_gen/runtime/models/upsampler/latent_upsampler.py index 24fa3ae6bd87..b926b612cb99 100644 --- a/python/sglang/multimodal_gen/runtime/models/upsampler/latent_upsampler.py +++ b/python/sglang/multimodal_gen/runtime/models/upsampler/latent_upsampler.py @@ -8,7 +8,7 @@ import torch.nn.functional as F from einops import rearrange -from sglang.kernels.ops.diffusion.group_norm_silu import apply_group_norm_silu +from sglang.kernels.ops.diffusion import apply_group_norm_silu from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( LayerwiseOffloadableModuleMixin, ) diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/flux2_vae_cuda_opt.py b/python/sglang/multimodal_gen/runtime/models/vaes/flux2_vae_cuda_opt.py index 9f8d27fe3f15..73fb7905c42c 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/flux2_vae_cuda_opt.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/flux2_vae_cuda_opt.py @@ -36,7 +36,9 @@ logger = init_logger(__name__) try: - from sglang.kernels.ops.diffusion.triton.group_norm_silu_twopass import ( + from sglang.kernels.ops.diffusion import ( + can_use_group_norm_silu_4d, + can_use_group_norm_silu_rows, group_norm_silu_4d, group_norm_silu_rows, ) @@ -70,8 +72,12 @@ def __init__(self, norm: nn.GroupNorm, gate: VaeFastPathGate) -> None: self._sgl_gate = gate def forward(self, x: torch.Tensor) -> torch.Tensor: - if self._sgl_gate.enabled and x.dim() == 4: - y = group_norm_silu_4d( + if ( + self._sgl_gate.enabled + and x.dim() == 4 + and can_use_group_norm_silu_4d(x, self.weight, self.bias, self.num_groups) + ): + return group_norm_silu_4d( x, self.weight, self.bias, @@ -79,8 +85,6 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: self.eps, apply_silu=True, ) - if y is not None: - return y return F.silu( F.group_norm(x, self.num_groups, self.weight, self.bias, self.eps) ) @@ -257,10 +261,12 @@ def _attn_fast_forward( if self.group_norm is not None: gn = self.group_norm - y = group_norm_silu_rows( - hs, gn.weight, gn.bias, gn.num_groups, gn.eps, apply_silu=False - ) - hs = y if y is not None else gn(hs.transpose(1, 2)).transpose(1, 2) + if can_use_group_norm_silu_rows(hs, gn.weight, gn.bias, gn.num_groups): + hs = group_norm_silu_rows( + hs, gn.weight, gn.bias, gn.num_groups, gn.eps, apply_silu=False + ) + else: + hs = gn(hs.transpose(1, 2)).transpose(1, 2) query = self.to_q(hs) key = self.to_k(hs) diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/hunyuanvae.py b/python/sglang/multimodal_gen/runtime/models/vaes/hunyuanvae.py index 438cf62a4744..5067a890e6e6 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/hunyuanvae.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/hunyuanvae.py @@ -22,7 +22,7 @@ import torch.nn as nn import torch.nn.functional as F -from sglang.kernels.ops.diffusion.group_norm_silu import apply_group_norm_silu +from sglang.kernels.ops.diffusion import apply_group_norm_silu from sglang.multimodal_gen.configs.models.vaes import HunyuanVAEConfig from sglang.multimodal_gen.runtime.distributed.parallel_state import ( get_decode_parallel_rank, diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/base_module.py b/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/base_module.py index 695f7613b5aa..80026a7e62a9 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/base_module.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/minimax_h3_video_vae/base_module.py @@ -11,9 +11,7 @@ from sglang.kernels.ops.activation.activation import ( silu_and_mul_with_activation_rounding, ) -from sglang.kernels.ops.diffusion.triton.scale_shift import ( - try_fused_scaled_residual_add_exact, -) +from sglang.kernels.ops.diffusion import try_fused_scaled_residual_add_exact from .attention import Attention from .vit_utils import _env_flag, _vit_torch_compile_kwargs diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/wan_vae_cuda_opt.py b/python/sglang/multimodal_gen/runtime/models/vaes/wan_vae_cuda_opt.py index d9f5857f1735..3044762bd857 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/wan_vae_cuda_opt.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/wan_vae_cuda_opt.py @@ -22,7 +22,10 @@ logger = init_logger(__name__) try: - from sglang.kernels.ops.diffusion.triton.wan_rmsnorm_silu import wan_rmsnorm_silu + from sglang.kernels.ops.diffusion import ( + can_use_wan_rmsnorm_silu, + wan_rmsnorm_silu, + ) _HAS_TRITON = True except ImportError: # pragma: no cover @@ -48,9 +51,8 @@ def __init__(self, norm: nn.Module, gate: VaeFastPathGate) -> None: def forward(self, x: torch.Tensor) -> torch.Tensor: if self._sgl_gate.enabled and not torch.compiler.is_compiling(): bias = self.bias if isinstance(self.bias, torch.Tensor) else None - y = wan_rmsnorm_silu(x, self.gamma, bias, rms_scale=self.scale) - if y is not None: - return y + if can_use_wan_rmsnorm_silu(x, self.gamma, bias): + return wan_rmsnorm_silu(x, self.gamma, bias, rms_scale=self.scale) # WanRMS_norm.forward (channel-first) + SiLU, same ops in the same # order, so the off-path stays bit-identical. return F.silu(F.normalize(x, dim=1) * self.scale * self.gamma + self.bias) diff --git a/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py b/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py index cd2119b8383b..f631bdccd54e 100644 --- a/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py +++ b/python/sglang/multimodal_gen/runtime/models/vaes/wanvae.py @@ -56,10 +56,7 @@ if current_platform.is_cuda(): try: - from sglang.kernels.ops.diffusion.triton.wan_causal_cache import ( - cat_pad_channels_last_3d, - dup_up3d_add, - ) + from sglang.kernels.ops.diffusion import cat_pad_channels_last_3d, dup_up3d_add except ImportError: # pragma: no cover cat_pad_channels_last_3d = None dup_up3d_add = None diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py index 76c66125b85c..5bd497015159 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -20,24 +20,16 @@ import torch import torch.nn as nn -from sglang.kernels.ops.diffusion.fused_gate_rmsnorm import ( +from sglang.kernels.ops.diffusion import ( mount_fused_gate_rmsnorm, - unmount_fused_gate_rmsnorm, -) -from sglang.kernels.ops.diffusion.fused_linear_gelu import ( mount_fused_linear_gelu, - unmount_fused_linear_gelu, -) -from sglang.kernels.ops.diffusion.fused_ln_modulate import ( mount_fused_ln_modulate, - unmount_fused_ln_modulate, -) -from sglang.kernels.ops.diffusion.hunyuan_qknorm import ( mount_hunyuan_qknorm, - unmount_hunyuan_qknorm, -) -from sglang.kernels.ops.diffusion.ltx2_rmsnorm_modulate import ( mount_ltx2_rms_norm_modulate, + unmount_fused_gate_rmsnorm, + unmount_fused_linear_gelu, + unmount_fused_ln_modulate, + unmount_hunyuan_qknorm, unmount_ltx2_rms_norm_modulate, ) from sglang.multimodal_gen import envs diff --git a/python/sglang/multimodal_gen/runtime/utils/mesh3d_utils.py b/python/sglang/multimodal_gen/runtime/utils/mesh3d_utils.py index 3ee1e49c73b6..28200028614f 100644 --- a/python/sglang/multimodal_gen/runtime/utils/mesh3d_utils.py +++ b/python/sglang/multimodal_gen/runtime/utils/mesh3d_utils.py @@ -18,7 +18,7 @@ logger = init_logger(__name__) # Import C++ mesh processor extension -from sglang.kernels.ops.diffusion.render.mesh_processor import meshVerticeInpaint +from sglang.kernels.ops.diffusion import meshVerticeInpaint def transform_pos( @@ -363,7 +363,7 @@ def _rasterize( resolution: Tuple[int, int], ) -> torch.Tensor: """Rasterize using CUDA rasterizer.""" - from sglang.kernels.ops.diffusion.render.hunyuan3d_rasterizer import rasterize + from sglang.kernels.ops.diffusion import rasterize if pos_clip.dim() == 2: pos_clip = pos_clip.unsqueeze(0) @@ -380,7 +380,7 @@ def _interpolate( tri: torch.Tensor, ) -> torch.Tensor: """Interpolate vertex attributes.""" - from sglang.kernels.ops.diffusion.render.hunyuan3d_rasterizer import interpolate + from sglang.kernels.ops.diffusion import interpolate barycentric = rast_out[0, ..., :-1] findices = rast_out[0, ..., -1].int() diff --git a/python/sglang/multimodal_gen/test/unit/test_latent_upsampler_group_norm_silu.py b/python/sglang/multimodal_gen/test/unit/test_latent_upsampler_group_norm_silu.py index e84f2154d2be..88615fbc4c62 100644 --- a/python/sglang/multimodal_gen/test/unit/test_latent_upsampler_group_norm_silu.py +++ b/python/sglang/multimodal_gen/test/unit/test_latent_upsampler_group_norm_silu.py @@ -6,7 +6,7 @@ import torch import sglang.multimodal_gen.runtime.models.upsampler.latent_upsampler as lu_mod -from sglang.kernels.ops.diffusion.group_norm_silu import apply_group_norm_silu +from sglang.kernels.ops.diffusion import apply_group_norm_silu from sglang.multimodal_gen.runtime.models.upsampler.latent_upsampler import ( LatentUpsampler, ResBlock, @@ -237,7 +237,7 @@ def test_latent_upsampler_forward_parity_cuda( @requires_cuda def test_resblock_actually_uses_triton_kernel_cuda(): - from sglang.kernels.ops.diffusion.triton import group_norm_silu as triton_mod + from sglang.kernels.ops.diffusion.norm import group_norm_silu_triton as triton_mod torch.manual_seed(0) device = torch.device("cuda") diff --git a/python/sglang/multimodal_gen/test/unit/test_layernorm_cutedsl_dispatch.py b/python/sglang/multimodal_gen/test/unit/test_layernorm_cutedsl_dispatch.py index 2b5ed1907434..948f630d5eff 100644 --- a/python/sglang/multimodal_gen/test/unit/test_layernorm_cutedsl_dispatch.py +++ b/python/sglang/multimodal_gen/test/unit/test_layernorm_cutedsl_dispatch.py @@ -10,7 +10,7 @@ ScaleResidualRMSNormScaleShift, ) -_CUTEDSL_MODULE = "sglang.kernels.ops.diffusion.cutedsl.scale_residual_norm_scale_shift" +_CUTEDSL_MODULE = "sglang.kernels.ops.diffusion.norm.scale_residual_norm_cutedsl" @pytest.mark.parametrize("hidden_size", [257, 8448]) diff --git a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_dit_contract.py b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_dit_contract.py index 4592a4e8d96b..84a40d81679f 100644 --- a/python/sglang/multimodal_gen/test/unit/test_minimax_h3_dit_contract.py +++ b/python/sglang/multimodal_gen/test/unit/test_minimax_h3_dit_contract.py @@ -400,9 +400,7 @@ def fake_all_to_all(actual, role=None): @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") def test_cuda_ulysses_qkv_pack_is_bit_exact(): - from sglang.kernels.ops.diffusion.triton.ulysses_qkv import ( - pack_qkv_destination_major, - ) + from sglang.kernels.ops.diffusion import pack_qkv_destination_major torch.manual_seed(23) rows, world_size, heads, head_size = 65, 8, 56, 128 diff --git a/test/registered/kernels/benchmark/diffusion/bench_causal_conv3d_cat_pad.py b/test/registered/kernels/benchmark/diffusion/bench_causal_conv3d_cat_pad.py index 44a864d59868..15c351c03578 100644 --- a/test/registered/kernels/benchmark/diffusion/bench_causal_conv3d_cat_pad.py +++ b/test/registered/kernels/benchmark/diffusion/bench_causal_conv3d_cat_pad.py @@ -3,12 +3,12 @@ import torch from sglang.kernels.jit.benchmark import marker -from sglang.kernels.ops.diffusion.causal_conv3d_cat_pad import ( - fused_causal_conv3d_cat_pad_cuda, -) -from sglang.kernels.ops.diffusion.triton.causal_conv3d_pad import ( +from sglang.kernels.ops.diffusion import ( fused_causal_conv3d_cat_pad as fused_causal_conv3d_cat_pad_triton, ) +from sglang.kernels.ops.diffusion import ( + fused_causal_conv3d_cat_pad_cuda, +) from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci register_cuda_ci( diff --git a/test/registered/kernels/benchmark/diffusion/bench_group_norm_silu.py b/test/registered/kernels/benchmark/diffusion/bench_group_norm_silu.py index f3af11577585..b60240e9bc0a 100644 --- a/test/registered/kernels/benchmark/diffusion/bench_group_norm_silu.py +++ b/test/registered/kernels/benchmark/diffusion/bench_group_norm_silu.py @@ -10,7 +10,7 @@ import torch.nn.functional as F import triton.testing -from sglang.kernels.ops.diffusion.triton.group_norm_silu import triton_group_norm_silu +from sglang.kernels.ops.diffusion import triton_group_norm_silu from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci from sglang.utils import is_in_ci diff --git a/test/registered/kernels/benchmark/diffusion/bench_ltx2_qknorm_split_rope.py b/test/registered/kernels/benchmark/diffusion/bench_ltx2_qknorm_split_rope.py index 0ab89ecd5443..f3fcc1799c95 100644 --- a/test/registered/kernels/benchmark/diffusion/bench_ltx2_qknorm_split_rope.py +++ b/test/registered/kernels/benchmark/diffusion/bench_ltx2_qknorm_split_rope.py @@ -4,9 +4,7 @@ import torch -from sglang.kernels.ops.diffusion.ltx2_qknorm_split_rope import ( - ltx2_qknorm_split_rope_cuda, -) +from sglang.kernels.ops.diffusion import ltx2_qknorm_split_rope_cuda from sglang.test.ci.ci_register import register_cuda_ci from sglang.utils import is_in_ci diff --git a/test/registered/kernels/benchmark/diffusion/bench_norm_impls.py b/test/registered/kernels/benchmark/diffusion/bench_norm_impls.py index fff0ddee6f1b..309c533c86cb 100644 --- a/test/registered/kernels/benchmark/diffusion/bench_norm_impls.py +++ b/test/registered/kernels/benchmark/diffusion/bench_norm_impls.py @@ -15,8 +15,11 @@ from sglang.kernels.jit.benchmark.utils import DEFAULT_DEVICE from sglang.kernels.jit.utils import KERNEL_PATH -from sglang.kernels.ops.diffusion.triton.norm import norm_infer, rms_norm_fn -from sglang.kernels.ops.diffusion.triton.rmsnorm_onepass import triton_one_pass_rms_norm +from sglang.kernels.ops.diffusion import ( + norm_infer, + rms_norm_fn, + triton_one_pass_rms_norm, +) from sglang.kernels.ops.layernorm.norm import fused_add_rmsnorm as jit_fused_add_rmsnorm from sglang.kernels.ops.layernorm.norm import rmsnorm as jit_rmsnorm from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci diff --git a/test/registered/kernels/benchmark/diffusion/bench_qknorm_rope.py b/test/registered/kernels/benchmark/diffusion/bench_qknorm_rope.py index 7189ebf0c0bb..86aadcf79623 100644 --- a/test/registered/kernels/benchmark/diffusion/bench_qknorm_rope.py +++ b/test/registered/kernels/benchmark/diffusion/bench_qknorm_rope.py @@ -153,7 +153,7 @@ def split_qknorm_rope(inputs: dict[str, torch.Tensor | bool]) -> None: def fused_qknorm_rope(inputs: dict[str, torch.Tensor | bool]) -> None: - from sglang.kernels.ops.diffusion.qknorm_rope import fused_inplace_qknorm_rope + from sglang.kernels.ops.diffusion import fused_inplace_qknorm_rope fused_inplace_qknorm_rope( inputs["q"], diff --git a/test/registered/kernels/benchmark/diffusion/bench_qwen_image_modulation.py b/test/registered/kernels/benchmark/diffusion/bench_qwen_image_modulation.py index 5b956a2862e6..5d7ce04b25b2 100644 --- a/test/registered/kernels/benchmark/diffusion/bench_qwen_image_modulation.py +++ b/test/registered/kernels/benchmark/diffusion/bench_qwen_image_modulation.py @@ -4,10 +4,10 @@ import triton.testing from sglang.kernels.jit.benchmark.utils import run_benchmark_no_cudagraph -from sglang.kernels.ops.diffusion.triton.norm import norm_infer -from sglang.kernels.ops.diffusion.triton.scale_shift import ( +from sglang.kernels.ops.diffusion import ( fuse_layernorm_scale_shift_gate_select01_kernel, fuse_residual_layernorm_scale_shift_gate_select01_kernel, + norm_infer, ) from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci from sglang.utils import is_in_ci diff --git a/test/registered/kernels/benchmark/diffusion/bench_residual_gate_add.py b/test/registered/kernels/benchmark/diffusion/bench_residual_gate_add.py index 6bedeb36b7b5..eb9d796de2f2 100644 --- a/test/registered/kernels/benchmark/diffusion/bench_residual_gate_add.py +++ b/test/registered/kernels/benchmark/diffusion/bench_residual_gate_add.py @@ -4,8 +4,7 @@ import torch -from sglang.kernels.ops.diffusion.residual_gate_add import residual_gate_add_cuda -from sglang.kernels.ops.diffusion.triton.scale_shift import fuse_scale_shift_kernel +from sglang.kernels.ops.diffusion import fuse_scale_shift_kernel, residual_gate_add_cuda from sglang.test.ci.ci_register import register_cuda_ci from sglang.utils import is_in_ci diff --git a/test/registered/kernels/benchmark/diffusion/bench_timestep_embedding.py b/test/registered/kernels/benchmark/diffusion/bench_timestep_embedding.py new file mode 100644 index 000000000000..140c43622cd6 --- /dev/null +++ b/test/registered/kernels/benchmark/diffusion/bench_timestep_embedding.py @@ -0,0 +1,70 @@ +"""CUDA vs eager-diffusers timing for the sinusoidal timestep embedding. + +Moved out of ``ops/diffusion/test_timestep_embedding.py``: it asserted nothing +and was skipped unless ``SGLANG_RUN_JIT_KERNEL_PERF_TESTS=1``, so it belonged +with the other benchmarks rather than in the correctness suite. +""" + +import sys + +import torch + +from sglang.kernels.ops.diffusion import timestep_embedding +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci( + est_time=20, stage="base-b-kernel-benchmark", runner_config="1-gpu-large" +) + +BATCHES = [1, 2, 8, 63, 256, 512, 613, 1024, 1536] +DIMS = [32, 64, 128, 256, 512, 1024, 2048, 4096] + + +def _reference(timesteps, dim, max_period=10000): + half_dim = dim // 2 + exponent = -torch.log( + torch.tensor(max_period, dtype=torch.float32, device=timesteps.device) + ) * torch.arange(0, half_dim, dtype=torch.float32, device=timesteps.device) + emb = torch.exp(exponent / (half_dim - 1)) + emb = timesteps[:, None].float() * emb[None, :] + return torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1) + + +def _time_ms(fn, *args, warmup=4, repeat=20): + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + for _ in range(warmup): + fn(*args) + torch.cuda.synchronize() + start.record() + for _ in range(repeat): + fn(*args) + end.record() + end.synchronize() + return start.elapsed_time(end) / repeat + + +def main() -> int: + if not torch.cuda.is_available(): + print("CUDA required") + return 0 + print(f"{'batch':>7} {'dim':>6} {'torch(ms)':>11} {'cuda(ms)':>10} {'speedup':>8}") + speedups = [] + for batch in BATCHES: + for dim in DIMS: + t = torch.linspace( + 0, max(100000, batch), steps=batch, device="cuda", dtype=torch.float32 + ) + torch_ms = _time_ms(_reference, t, dim) + cuda_ms = _time_ms(timestep_embedding, t, dim) + speedups.append(torch_ms / cuda_ms) + print( + f"{batch:>7} {dim:>6} {torch_ms:>11.6f} {cuda_ms:>10.6f} " + f"{speedups[-1]:>8.3f}" + ) + print(f"average speedup: {sum(speedups) / len(speedups):.4f}") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/test/registered/kernels/ops/diffusion/test_activation.py b/test/registered/kernels/ops/diffusion/test_activation.py new file mode 100644 index 000000000000..4a2802ad2719 --- /dev/null +++ b/test/registered/kernels/ops/diffusion/test_activation.py @@ -0,0 +1,112 @@ +"""``diffusion.activation``: activation-function fusions. + +All of these are bit-exact by construction -- they are elementwise chains with +no reduction, so reproducing aten's per-op fp32-opmath / round-to-bf16 +boundaries is enough and ``torch.equal`` is the assertion. + +The cublasLt linear+tanh-GELU epilogue is *not* here: it is not bit-exact and +is therefore quality-gated, so it is tested through its mount protocol in +``test_sites.py``. +""" + +import sys + +import pytest +import torch +import torch.nn.functional as F + +from sglang.kernels.ops.diffusion import ( + can_use_fused_bias_glu, + can_use_fused_bias_silu, + can_use_fused_silu_mul, + fused_bias_glu, + fused_bias_silu, + fused_packed_silu_mul_bitexact, + fused_silu_mul_bitexact, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=6, stage="base-b-kernel-unit", runner_config="1-gpu-large") +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") + + +@pytest.mark.parametrize("channels", [2240, 11200]) +def test_sana_bias_silu_is_bit_exact(channels): + torch.manual_seed(0) + x = torch.randn( + (1, channels, 7, 5), + device="cuda", + dtype=torch.bfloat16, + ).to(memory_format=torch.channels_last) + bias = torch.randn(channels, device="cuda", dtype=torch.bfloat16) + + assert can_use_fused_bias_silu(x, bias) + actual = fused_bias_silu(x, bias) + expected = F.silu(x + bias[None, :, None, None]) + + assert actual.is_contiguous(memory_format=torch.channels_last) + assert torch.equal(actual, expected) + + +@pytest.mark.parametrize("channels", [2240, 5600]) +def test_sana_bias_glu_is_bit_exact(channels): + torch.manual_seed(1) + x = torch.randn( + (1, 2 * channels, 7, 5), + device="cuda", + dtype=torch.bfloat16, + ).to(memory_format=torch.channels_last) + bias = torch.randn(2 * channels, device="cuda", dtype=torch.bfloat16) + + assert can_use_fused_bias_glu(x, bias) + actual = fused_bias_glu(x, bias) + biased = x + bias[None, :, None, None] + hidden, gate = torch.chunk(biased, 2, dim=1) + expected = hidden * F.silu(gate) + + assert actual.is_contiguous(memory_format=torch.channels_last) + assert torch.equal(actual, expected) + + +# --------------------------------------------------------------------------- +# silu(a) * b for split-projection SwiGLU +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("shape", [(1, 4096, 3072), (2, 17, 512)]) +def test_silu_mul_is_bit_exact(shape): + # Separate gate/up GEMMs, so the concatenated ``silu_and_mul`` kernels do + # not apply without an extra full-width cat -- this kernel replaces the + # eager ``F.silu(a) * b`` pair instead. + torch.manual_seed(0) + a = torch.randn(shape, device="cuda", dtype=torch.bfloat16) + b = torch.randn_like(a) + + assert can_use_fused_silu_mul(a, b) + assert torch.equal(fused_silu_mul_bitexact(a, b), F.silu(a) * b) + + +@pytest.mark.parametrize("hidden", [384, 3072]) +@pytest.mark.parametrize("strided", [False, True]) +def test_packed_silu_mul_is_bit_exact(hidden, strided): + # The packed form splits one [.., 2 * hidden] projection in-kernel; it must + # accept the strided view a wider projection slice produces. + torch.manual_seed(1) + if strided: + x = torch.randn(1, 19, 3 * hidden, device="cuda", dtype=torch.bfloat16) + x = x[..., : 2 * hidden] + else: + x = torch.randn(1, 19, 2 * hidden, device="cuda", dtype=torch.bfloat16) + + expected = F.silu(x[..., :hidden]) * x[..., hidden:] + assert torch.equal(fused_packed_silu_mul_bitexact(x), expected) + + +def test_silu_mul_rejects_mismatched_operands(): + a = torch.randn(1, 8, 64, device="cuda", dtype=torch.bfloat16) + assert not can_use_fused_silu_mul(a, a.float()) # mixed dtypes + assert not can_use_fused_silu_mul(a, a[:, :-1]) # mismatched shapes + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernels/ops/diffusion/test_autoencoder_kl_fastpath.py b/test/registered/kernels/ops/diffusion/test_autoencoder_kl_fastpath.py deleted file mode 100644 index 22f190db302e..000000000000 --- a/test/registered/kernels/ops/diffusion/test_autoencoder_kl_fastpath.py +++ /dev/null @@ -1,62 +0,0 @@ -"""Install-path checks for the generic AutoencoderKL CUDA fast path.""" - -import sys - -import pytest -import torch - -from sglang.multimodal_gen.configs.models.vaes.stablediffusion3 import ( - StableDiffusion3VAEConfig, -) -from sglang.multimodal_gen.runtime.models.vaes import flux2_vae_cuda_opt as vae_opt -from sglang.multimodal_gen.runtime.models.vaes.autoencoder import AutoencoderKL -from sglang.multimodal_gen.runtime.models.vaes.fast_path_gate import ( - use_vae_fast_path, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=40, stage="base-b-kernel-unit", runner_config="1-gpu-large") -pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") - - -def _small_config(): - config = StableDiffusion3VAEConfig() - config.arch_config.latent_channels = 2 - config.arch_config.block_out_channels = (4, 4) - config.arch_config.down_block_types = ("DownEncoderBlock2D",) * 2 - config.arch_config.up_block_types = ("UpDecoderBlock2D",) * 2 - config.arch_config.layers_per_block = 1 - config.arch_config.norm_num_groups = 1 - config.arch_config.sample_size = 8 - return config - - -@torch.no_grad() -def test_autoencoder_kl_fastpath_install(): - torch.manual_seed(0) - vae = AutoencoderKL(_small_config()).to("cuda", torch.bfloat16).eval() - ref_names = {n for n, _ in vae.named_parameters()} - ref_sd = {k: v.clone() for k, v in vae.state_dict().items()} - z = torch.randn(1, 2, 8, 8, device="cuda", dtype=torch.bfloat16) - ref = vae.decode(z) - - opt = vae_opt.maybe_optimize_autoencoder_kl(vae) - # Wrappers must not change parameter FQNs; strict load must round-trip. - assert {n for n, _ in opt.named_parameters()} == ref_names - opt.load_state_dict(ref_sd, strict=True) - # Gate off: bit-for-bit the original path. - assert torch.equal(opt.decode(z), ref) - # use_vae_fast_path() is a no-op when nothing registered a gate, so check - # the wrappers went in before relying on it to switch paths. - assert any( - isinstance(m, (vae_opt.FusedGroupNormSiLU, vae_opt.FusedUpsample2xConv2d)) - for m in opt.modules() - ) - # Gate on: fast path runs and stays close; leaving the scope restores exact. - with use_vae_fast_path(opt, True): - torch.testing.assert_close(opt.decode(z).float(), ref.float(), atol=0.1, rtol=0) - assert torch.equal(opt.decode(z), ref) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernels/ops/diffusion/test_bitexact_gate.py b/test/registered/kernels/ops/diffusion/test_bitexact_gate.py deleted file mode 100644 index 8b98c653bde4..000000000000 --- a/test/registered/kernels/ops/diffusion/test_bitexact_gate.py +++ /dev/null @@ -1,159 +0,0 @@ -import sys -from types import ModuleType -from unittest.mock import MagicMock, patch - -import pytest -import torch - -from sglang.kernels.ops.diffusion.bitexact_gate import ( - BitExactFusionGate, - flashinfer_rmsnorm_diagnostic_hint, - tensors_equal, -) -from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.test_utils import CustomTestCase - -register_cpu_ci(est_time=2, suite="base-a-test-cpu") - - -def test_bitexact_gate_once_mode_verifies_then_reuses(): - gate = BitExactFusionGate("once") - calls = {"fused": 0, "ref": 0} - - def fused(): - calls["fused"] += 1 - return torch.tensor([1.0]) - - def ref(): - calls["ref"] += 1 - return torch.tensor([1.0]) - - assert torch.equal(gate.accept_or_fallback(fused(), ref()), torch.tensor([1.0])) - assert gate.verified and not gate.disabled and calls == {"fused": 1, "ref": 1} - assert torch.equal(fused(), torch.tensor([1.0])) - assert calls == {"fused": 2, "ref": 1} - - -def test_bitexact_gate_mismatch_disables_permanently(): - gate = BitExactFusionGate("mismatch") - - out = gate.accept_or_fallback( - torch.tensor([1.0]), - torch.tensor([2.0]), - mismatch_msg="mismatch", - ) - assert torch.equal(out, torch.tensor([2.0])) - assert gate.disabled and not gate.verified - - -def test_bitexact_gate_per_signature_tracks_each_sig(): - gate = BitExactFusionGate("sig", per_signature=True) - a = torch.tensor([1.0]) - assert torch.equal(gate.accept_or_fallback(a, a, sig=("a",)), a) - assert gate.is_verified(("a",)) - assert not gate.is_verified(("b",)) - assert torch.equal(gate.accept_or_fallback(a, a, sig=("b",)), a) - assert gate.verified_sigs == {("a",), ("b",)} - - -def test_bitexact_gate_skips_first_sight_during_graph_capture(monkeypatch): - # Negative-branch contract: an unverified gate must not attempt first-sight - # verification inside CUDA graph capture — the eager-reference host sync - # would abort the capture (and BCG would permanently block the signature). - gate = BitExactFusionGate("capture") - monkeypatch.setattr(torch.cuda, "is_available", lambda: True) - monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: True) - assert not gate.can_attempt_once() - # A verified gate replays the fused kernel alone, which is capture-safe. - gate.mark_verified() - assert gate.can_attempt_once() - - -def test_tensors_equal_supports_sequences(): - assert tensors_equal( - (torch.tensor([1.0]), torch.tensor([2.0])), - (torch.tensor([1.0]), torch.tensor([2.0])), - ) - assert not tensors_equal( - (torch.tensor([1.0]), torch.tensor([2.0])), - (torch.tensor([1.0]), torch.tensor([3.0])), - ) - - -class TestBitExactFallbackDiagnostics(CustomTestCase): - def test_mismatch_warning_is_actionable_and_diagnostic_is_lazy(self): - logger = MagicMock() - diagnostic = MagicMock(return_value="backend=CuTe DSL") - gate = BitExactFusionGate("diagnostic") - - matched = gate.accept_or_fallback( - torch.tensor([1.0]), - torch.tensor([1.0]), - logger=logger, - diagnostic_hint=diagnostic, - ) - self.assertTrue(torch.equal(matched, torch.tensor([1.0]))) - diagnostic.assert_not_called() - logger.warning_once.assert_not_called() - - gate = BitExactFusionGate("diagnostic") - fallback = gate.accept_or_fallback( - torch.tensor([1.0]), - torch.tensor([2.0]), - logger=logger, - diagnostic_hint=diagnostic, - ) - - self.assertTrue(torch.equal(fallback, torch.tensor([2.0]))) - diagnostic.assert_called_once_with() - warning = logger.warning_once.call_args.args[0] - self.assertIn("Correctness is preserved", warning) - self.assertIn("reference kernel or reduction-order change", warning) - self.assertIn("backend=CuTe DSL", warning) - - def test_diagnostic_failure_cannot_break_the_eager_fallback(self): - logger = MagicMock() - - def broken_diagnostic(): - raise RuntimeError("diagnostics unavailable") - - gate = BitExactFusionGate("diagnostic") - fallback = gate.accept_or_fallback( - torch.tensor([1.0]), - torch.tensor([2.0]), - logger=logger, - diagnostic_hint=broken_diagnostic, - ) - - self.assertTrue(torch.equal(fallback, torch.tensor([2.0]))) - self.assertTrue(gate.disabled) - self.assertIn("Correctness is preserved", logger.warning_once.call_args.args[0]) - - def test_flashinfer_rmsnorm_hint_reports_backend_and_versions(self): - flashinfer = ModuleType("flashinfer") - flashinfer_norm = ModuleType("flashinfer.norm") - flashinfer_norm._USE_CUDA_NORM = False - versions = { - "flashinfer-python": "0.6.12", - "flashinfer-cubin": "0.6.12", - "flashinfer-jit-cache": "0.6.12+cu130", - } - - with ( - patch.dict( - sys.modules, - {"flashinfer": flashinfer, "flashinfer.norm": flashinfer_norm}, - ), - patch("importlib.metadata.version", side_effect=versions.__getitem__), - patch.dict("os.environ", {"FLASHINFER_USE_CUDA_NORM": "0"}), - ): - hint = flashinfer_rmsnorm_diagnostic_hint() - - self.assertIn("backend=CuTe DSL", hint) - self.assertIn("FLASHINFER_USE_CUDA_NORM=0", hint) - for package, version in versions.items(): - self.assertIn(f"{package}={version}", hint) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernels/ops/diffusion/test_causal_conv3d_cat_pad.py b/test/registered/kernels/ops/diffusion/test_causal_conv3d_cat_pad.py deleted file mode 100644 index f920d6478711..000000000000 --- a/test/registered/kernels/ops/diffusion/test_causal_conv3d_cat_pad.py +++ /dev/null @@ -1,89 +0,0 @@ -import sys - -import pytest -import torch - -from sglang.kernels.jit.utils import get_ci_test_range -from sglang.kernels.ops.diffusion.causal_conv3d_cat_pad import ( - fused_causal_conv3d_cat_pad_cuda, -) -from sglang.kernels.ops.diffusion.triton.causal_conv3d_pad import ( - fused_causal_conv3d_cat_pad as fused_causal_conv3d_cat_pad_triton, -) -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci - -register_cuda_ci(est_time=45, stage="base-b-kernel-unit", runner_config="1-gpu-large") -register_cuda_ci(est_time=45, stage="base-b-kernel-unit", runner_config="4-gpu-b200") -register_amd_ci(est_time=10, stage="jit-kernel-unit", runner_config="amd") - -DEVICE = "cuda" -DTYPE = torch.bfloat16 - -COSMOS3_CASES = get_ci_test_range( - [ - (1024, 1, 30, 52, 1), - (1024, 1, 30, 52, 2), - (1024, 2, 60, 104, 1), - (1024, 2, 60, 104, 2), - (512, 4, 120, 208, 1), - (512, 4, 120, 208, 2), - (256, 4, 240, 416, 1), - (256, 4, 240, 416, 2), - ], - [(1024, 1, 30, 52, 1), (512, 4, 120, 208, 2)], -) - - -def _make_inputs( - channels: int, - t_size: int, - h_size: int, - w_size: int, - cache_t: int, -) -> tuple[torch.Tensor, torch.Tensor, tuple[int, ...]]: - generator = torch.Generator(device=DEVICE) - generator.manual_seed(channels * 1009 + t_size * 251 + h_size + cache_t) - x = torch.randn( - (1, channels, t_size, h_size, w_size), - device=DEVICE, - dtype=DTYPE, - generator=generator, - ) - cache_x = torch.randn( - (1, channels, cache_t, h_size, w_size), - device=DEVICE, - dtype=DTYPE, - generator=generator, - ) - padding = (1, 1, 1, 1, cache_t, 0) - return x, cache_x, padding - - -@pytest.mark.parametrize("channels,t_size,h_size,w_size,cache_t", COSMOS3_CASES) -def test_causal_conv3d_cat_pad( - channels: int, - t_size: int, - h_size: int, - w_size: int, - cache_t: int, -) -> None: - x, cache_x, padding = _make_inputs(channels, t_size, h_size, w_size, cache_t) - actual = fused_causal_conv3d_cat_pad_cuda(x, cache_x, padding) - expected = fused_causal_conv3d_cat_pad_triton(x, cache_x, padding) - torch.testing.assert_close(actual, expected, atol=0, rtol=0) - - -def test_causal_conv3d_cat_pad_torch_compile() -> None: - x, cache_x, padding = _make_inputs(1024, 1, 30, 52, 1) - - @torch.compile(fullgraph=True) - def fn(x: torch.Tensor, cache_x: torch.Tensor) -> torch.Tensor: - return fused_causal_conv3d_cat_pad_cuda(x, cache_x, padding) - - actual = fn(x, cache_x) - expected = fused_causal_conv3d_cat_pad_triton(x, cache_x, padding) - torch.testing.assert_close(actual, expected, atol=0, rtol=0) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__])) diff --git a/test/registered/kernels/ops/diffusion/test_ernie_norm_scale_shift.py b/test/registered/kernels/ops/diffusion/test_ernie_norm_scale_shift.py deleted file mode 100644 index 747a4d807301..000000000000 --- a/test/registered/kernels/ops/diffusion/test_ernie_norm_scale_shift.py +++ /dev/null @@ -1,135 +0,0 @@ -"""ERNIE fused norm/scale/shift fast paths must stay bit-exact vs eager.""" - -import sys -from unittest.mock import patch - -import pytest -import torch - -import sglang.multimodal_gen.runtime.models.dits.ernie_image as ernie_image -from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm -from sglang.multimodal_gen.runtime.models.dits.ernie_image import ( - _ernie_gated_norm_scale_shift, - _ernie_norm_scale_shift, - _ernie_qknorm_rope, - _ernie_qknorm_rope_reference, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=4, stage="base-b-kernel-unit", runner_config="1-gpu-large") -pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") - - -@pytest.mark.parametrize("shape", [(1, 4216, 4096), (2, 1140, 4096), (1, 128, 2048)]) -def test_fused_norm_scale_shift_is_bit_exact(shape): - # (1, 4216, 4096) is the real ERNIE-Image shape (1024^2 image + text - # tokens, hidden 4096); 2048 covers the threads_per_row=32 regime. - torch.manual_seed(0) - batch, seq, hidden = shape - norm = RMSNorm(hidden, eps=1e-6).to(device="cuda", dtype=torch.bfloat16) - with torch.no_grad(): - norm.weight.copy_(torch.randn(hidden)) - x = torch.randn(batch, seq, hidden, device="cuda", dtype=torch.bfloat16) - residual = torch.randn_like(x) - update = torch.randn_like(x) - scale = torch.randn(batch, 1, hidden, device="cuda", dtype=torch.bfloat16) * 0.1 - shift = torch.randn(batch, 1, hidden, device="cuda", dtype=torch.bfloat16) * 0.1 - gate = torch.randn(batch, 1, hidden, device="cuda", dtype=torch.bfloat16) - - with torch.no_grad(): - out = _ernie_norm_scale_shift(norm, x, scale, shift) - ref = norm(x) * (1 + scale) + shift - assert torch.equal(out, ref) - - out2, res = _ernie_gated_norm_scale_shift( - norm, residual, update, gate, scale, shift - ) - res_ref = residual + gate * update - ref2 = norm(res_ref) * (1 + scale) + shift - assert torch.equal(res, res_ref) - assert torch.equal(out2, ref2) - - # the fast paths must actually be in use (not silently disabled) - assert ernie_image._ERNIE_NORM.verified - assert ernie_image._ERNIE_GATED_NORM.verified - assert not ernie_image._ERNIE_NORM.disabled - assert not ernie_image._ERNIE_GATED_NORM.disabled - - -def test_fused_qknorm_rope_is_bit_exact(): - torch.manual_seed(1) - ernie_image._ERNIE_QKNORM_ROPE.disabled = False - ernie_image._ERNIE_QKNORM_ROPE.verified = False - batch, seq, heads, head_dim = 1, 257, 32, 128 - q = torch.randn(batch, seq, heads, head_dim, device="cuda", dtype=torch.bfloat16) - k = torch.randn_like(q) - q_norm = RMSNorm(head_dim, eps=1e-6).to(device="cuda", dtype=torch.bfloat16) - k_norm = RMSNorm(head_dim, eps=1e-6).to(device="cuda", dtype=torch.bfloat16) - cos = torch.randn(seq, head_dim, device="cuda", dtype=torch.bfloat16) - sin = torch.randn_like(cos) - cache = torch.cat((cos, sin), dim=-1).contiguous() - positions = torch.arange(seq, device="cuda", dtype=torch.long) - - q_ref, k_ref = _ernie_qknorm_rope_reference( - q.clone(), k.clone(), q_norm, k_norm, head_dim, cos, sin - ) - q_out, k_out = _ernie_qknorm_rope( - q, - k, - q_norm, - k_norm, - head_dim, - cos, - sin, - cache, - positions, - ) - - assert torch.equal(q_out, q_ref) - assert torch.equal(k_out, k_ref) - assert ernie_image._ERNIE_QKNORM_ROPE.verified - assert not ernie_image._ERNIE_QKNORM_ROPE.disabled - - -def test_qknorm_rope_first_attempt_exception_uses_pristine_inputs(): - torch.manual_seed(2) - ernie_image._ERNIE_QKNORM_ROPE.disabled = False - ernie_image._ERNIE_QKNORM_ROPE.verified = False - batch, seq, heads, head_dim = 1, 17, 4, 128 - q = torch.randn(batch, seq, heads, head_dim, device="cuda", dtype=torch.bfloat16) - k = torch.randn_like(q) - q_norm = RMSNorm(head_dim, eps=1e-6).to(device="cuda", dtype=torch.bfloat16) - k_norm = RMSNorm(head_dim, eps=1e-6).to(device="cuda", dtype=torch.bfloat16) - cos = torch.randn(seq, head_dim, device="cuda", dtype=torch.bfloat16) - sin = torch.randn_like(cos) - cache = torch.cat((cos, sin), dim=-1).contiguous() - positions = torch.arange(seq, device="cuda", dtype=torch.long) - q_ref, k_ref = _ernie_qknorm_rope_reference( - q.clone(), k.clone(), q_norm, k_norm, head_dim, cos, sin - ) - - def mutate_then_raise(**kwargs): - kwargs["q"].zero_() - kwargs["k"].zero_() - raise RuntimeError("synthetic kernel failure") - - with patch.object(ernie_image, "apply_qk_norm_rope", mutate_then_raise): - q_out, k_out = _ernie_qknorm_rope( - q, - k, - q_norm, - k_norm, - head_dim, - cos, - sin, - cache, - positions, - ) - - assert torch.equal(q_out, q_ref) - assert torch.equal(k_out, k_ref) - assert ernie_image._ERNIE_QKNORM_ROPE.disabled - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__])) diff --git a/test/registered/kernels/ops/diffusion/test_flux2_eager_fusions.py b/test/registered/kernels/ops/diffusion/test_flux2_eager_fusions.py deleted file mode 100644 index 5301c9ddd855..000000000000 --- a/test/registered/kernels/ops/diffusion/test_flux2_eager_fusions.py +++ /dev/null @@ -1,101 +0,0 @@ -"""FLUX.2 eager fusions must be bit-exact for real packed/view layouts.""" - -import unittest -from unittest.mock import patch - -import torch -import torch.nn.functional as F - -import sglang.multimodal_gen.runtime.models.dits.flux_2 as flux2 -from sglang.multimodal_gen.runtime.models.dits.flux_2 import ( - _flux2_norm_modulate, - _flux2_swiglu, -) -from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.test_utils import CustomTestCase - -register_cuda_ci(est_time=12, stage="base-b-kernel-unit", runner_config="1-gpu-large") - - -@unittest.skipUnless(torch.cuda.is_available(), "CUDA required") -class TestFlux2EagerFusions(CustomTestCase): - def setUp(self): - flux2._FLUX2_LN_MOD.disabled = False - flux2._FLUX2_LN_MOD.verified = False - flux2._FLUX2_LN_MOD_SIGS.clear() - flux2._FLUX2_SWIGLU.disabled = False - flux2._FLUX2_SWIGLU.verified = False - flux2._FLUX2_SWIGLU_SIGS.clear() - - def test_norm_modulate_is_bit_exact_across_sequence_lengths(self): - torch.manual_seed(0) - hidden = 256 - norm = torch.nn.LayerNorm( - hidden, eps=1e-6, elementwise_affine=False, device="cuda" - ) - # FLUX.2 modulation values are views of one packed projection. - params = torch.randn(1, 1, 6 * hidden, device="cuda").bfloat16() - shift, scale = params.chunk(6, dim=-1)[:2] - - for seq in (17, 65): - x = torch.randn(1, seq, hidden, device="cuda").bfloat16() - expected = norm(x) * (1 + scale) + shift - actual = _flux2_norm_modulate(norm, x, scale, shift) - self.assertTrue(torch.equal(actual, expected)) - - self.assertFalse(flux2._FLUX2_LN_MOD.disabled) - self.assertEqual(len(flux2._FLUX2_LN_MOD_SIGS), 1) - - def test_packed_swiglu_is_bit_exact_for_contiguous_and_strided_views(self): - torch.manual_seed(1) - hidden = 384 - inputs = [ - torch.randn(1, 19, 2 * hidden, device="cuda").bfloat16(), - torch.randn(1, 19, 3 * hidden, device="cuda").bfloat16()[..., : 2 * hidden], - ] - for x in inputs: - expected = F.silu(x[..., :hidden]) * x[..., hidden:] - actual = _flux2_swiglu(x) - self.assertTrue(torch.equal(actual, expected)) - - self.assertFalse(flux2._FLUX2_SWIGLU.disabled) - self.assertEqual(len(flux2._FLUX2_SWIGLU_SIGS), 2) - - def test_fp16_preserves_reference_path(self): - x = torch.randn(1, 17, 512, device="cuda", dtype=torch.float16) - expected = F.silu(x[..., :256]) * x[..., 256:] - actual = _flux2_swiglu(x) - self.assertTrue(torch.equal(actual, expected)) - self.assertFalse(flux2._FLUX2_SWIGLU.disabled) - - def test_packed_swiglu_rejects_non_dense_outer_stride(self): - base = torch.randn(2, 23, 512, device="cuda", dtype=torch.bfloat16) - x = base[:, :19] - self.assertNotEqual(x.stride(0), x.shape[1] * x.stride(1)) - - expected = F.silu(x[..., :256]) * x[..., 256:] - actual = _flux2_swiglu(x) - self.assertTrue(torch.equal(actual, expected)) - self.assertEqual(len(flux2._FLUX2_SWIGLU_SIGS), 0) - - def test_new_swiglu_signature_is_not_verified_during_graph_capture(self): - first = torch.randn(1, 17, 512, device="cuda", dtype=torch.bfloat16) - self.assertTrue( - torch.equal( - _flux2_swiglu(first), - F.silu(first[..., :256]) * first[..., 256:], - ) - ) - self.assertEqual(len(flux2._FLUX2_SWIGLU_SIGS), 1) - - second = torch.randn(1, 19, 768, device="cuda", dtype=torch.bfloat16) - with patch("torch.cuda.is_current_stream_capturing", return_value=True): - actual = _flux2_swiglu(second) - - expected = F.silu(second[..., :384]) * second[..., 384:] - self.assertTrue(torch.equal(actual, expected)) - self.assertEqual(len(flux2._FLUX2_SWIGLU_SIGS), 1) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/kernels/ops/diffusion/test_flux2_vae_fastpath.py b/test/registered/kernels/ops/diffusion/test_flux2_vae_fastpath.py deleted file mode 100644 index 9796f213e8dd..000000000000 --- a/test/registered/kernels/ops/diffusion/test_flux2_vae_fastpath.py +++ /dev/null @@ -1,64 +0,0 @@ -"""Focused correctness checks for the FLUX.2 VAE CUDA fast path.""" - -import sys - -import pytest -import torch -import torch.nn as nn -import torch.nn.functional as F -from diffusers.models.upsampling import Upsample2D - -from sglang.kernels.ops.diffusion.triton import group_norm_silu_twopass as gn_kernel -from sglang.multimodal_gen.runtime.models.vaes import flux2_vae_cuda_opt as vae_opt -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=40, stage="base-b-kernel-unit", runner_config="1-gpu-large") -pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") - - -@torch.no_grad() -def test_flux2_vae_fastpath(): - torch.manual_seed(0) - gate = vae_opt.VaeFastPathGate() - gn = nn.GroupNorm(32, 128, eps=1e-6).to("cuda", torch.bfloat16) - x = torch.randn(1, 128, 64, 64, device="cuda", dtype=torch.bfloat16).to( - memory_format=torch.channels_last - ) - ref = F.silu(gn(x)) - fused_gn = vae_opt.FusedGroupNormSiLU(gn, gate) - assert set(fused_gn.state_dict()) == {"weight", "bias"} - assert torch.equal(fused_gn(x), ref) - assert ( - gn_kernel.group_norm_silu_4d(x.contiguous(), gn.weight, gn.bias, 32, 1e-6) - is None - ) - assert gn_kernel.group_norm_silu_4d(x, gn.weight.cpu(), gn.bias, 32, 1e-6) is None - assert ( - gn_kernel.group_norm_silu_4d(x[..., :0, :], gn.weight, gn.bias, 32, 1e-6) - is None - ) - - gate.enabled = True - fast = fused_gn(x) - assert fast.is_contiguous(memory_format=torch.channels_last) - torch.testing.assert_close(fast.float(), ref.float(), atol=0.06, rtol=0) - - gate.enabled = False - up = Upsample2D(channels=32, use_conv=True).to("cuda", torch.bfloat16) - fused_up = vae_opt.FusedUpsample2xConv2d(up, gate) - assert set(fused_up.state_dict()) == {"conv.weight", "conv.bias"} - x = torch.randn(2, 32, 33, 29, device="cuda", dtype=torch.bfloat16) - ref = up(x) - assert torch.equal(fused_up(x), ref) - assert fused_up._fused_weight is None - - gate.enabled = True - fast = fused_up(x) - assert fused_up._fused_weight is not None - ref_range = ref.float().max() - ref.float().min() - relative_mse = F.mse_loss(fast.float(), ref.float()) / ref_range.square() - assert relative_mse < 3.2e-5 - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernels/ops/diffusion/test_flux_ln_modulate.py b/test/registered/kernels/ops/diffusion/test_flux_ln_modulate.py deleted file mode 100644 index 7e42c7e314de..000000000000 --- a/test/registered/kernels/ops/diffusion/test_flux_ln_modulate.py +++ /dev/null @@ -1,75 +0,0 @@ -"""FLUX.1 fused LN+modulate fast path must stay bit-exact vs eager.""" - -import pytest -import torch - -import sglang.multimodal_gen.runtime.models.dits.flux as flux -from sglang.kernels.ops.diffusion.fused_ln_modulate import ( - mark_fused_ln_modulate_site, - mount_fused_ln_modulate, -) -from sglang.multimodal_gen.runtime.models.dits.flux import ( - _flux_fused_ln_modulate, - _flux_norm_modulate, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=4, stage="base-b-kernel-unit", runner_config="1-gpu-large") -pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") - - -def _eager(norm, x, scale, shift): - return norm(x) * (1 + scale[:, None]) + shift[:, None] - - -def _make_site_inputs(shape, chunks, seed): - torch.manual_seed(seed) - batch, seq, hidden = shape - norm = torch.nn.LayerNorm(hidden, eps=1e-6, elementwise_affine=False).cuda() - x = (torch.randn(batch, seq, hidden, device="cuda") * 8).bfloat16() - emb = torch.randn(batch, chunks * hidden, device="cuda").bfloat16() - parts = emb.chunk(chunks, dim=1) # strided adaLN projection views - return norm, x, parts[0], parts[1] - - -@pytest.mark.parametrize( - "shape,chunks", - [ - ((1, 4096, 3072), 6), # dual-stream image tokens (1024^2), chunk(6) - ((1, 512, 3072), 6), # dual-stream text tokens - ((1, 4608, 3072), 3), # single-stream concat, chunk(3) - ((2, 300, 3072), 6), # CFG batch, odd seq - ], -) -def test_flux_fused_ln_modulate_is_bit_exact(shape, chunks): - # Every distinct (shape, stride, eps) signature the FLUX.1 sites emit - # must verify torch.equal on first sight and stay enabled. - norm, x, shift, scale = _make_site_inputs(shape, chunks, seed=0) - out = _flux_fused_ln_modulate(norm, x, scale, shift) - assert out is not None - assert torch.equal(out, _eager(norm, x, scale, shift)) - assert not flux._FLUX_LN_MOD.disabled - assert flux._FLUX_LN_MOD.verified - - -def test_flux_norm_modulate_bitexact_supersedes_high_fold(): - # With the quality="high" affine fold mounted, the bit-exact kernel - # still takes priority, so the site output stays lossless. - site = torch.nn.Module() - mark_fused_ln_modulate_site(site) - assert mount_fused_ln_modulate(site) - norm, x, shift, scale = _make_site_inputs((1, 128, 3072), 6, seed=1) - out = _flux_norm_modulate(site, norm, x, scale, shift) - assert torch.equal(out, _eager(norm, x, scale, shift)) - - -def test_flux_fused_ln_modulate_rejects_unsupported_hidden(): - # hidden % 4 != 0 is outside the kernel contract and must bail out. - norm, x, shift, scale = _make_site_inputs((1, 64, 3070), 6, seed=2) - assert _flux_fused_ln_modulate(norm, x, scale, shift) is None - - -if __name__ == "__main__": - import sys - - sys.exit(pytest.main([__file__])) diff --git a/test/registered/kernels/ops/diffusion/test_flydsl_fused_norm.py b/test/registered/kernels/ops/diffusion/test_flydsl_fused_norm.py deleted file mode 100644 index 2571e2d69bd3..000000000000 --- a/test/registered/kernels/ops/diffusion/test_flydsl_fused_norm.py +++ /dev/null @@ -1,133 +0,0 @@ -import sys - -import pytest -import torch -import torch.nn.functional as F - -from sglang.test.ci.ci_register import register_amd_ci - -register_amd_ci(est_time=30, stage="jit-kernel-unit", runner_config="amd") - -DEVICE = "cuda" -D = 5120 -EPS = 1e-6 - - -def _ref_rms_norm(x_f32, weight, eps): - var = x_f32.pow(2).mean(-1, keepdim=True) - return x_f32 * torch.rsqrt(var + eps) - - -def _ref_fused_residual_norm_ss( - residual, x, gate, weight, bias, scale, shift, norm_type, eps -): - ref_res = residual.float() + x.float() * (gate.float() if gate is not None else 1) - ref_res_bf16 = ref_res.to(torch.bfloat16) - if norm_type == "layer": - normed = F.layer_norm(ref_res_bf16.float(), (D,), weight, bias, eps) - else: - normed = _ref_rms_norm(ref_res_bf16.float(), weight, eps) * weight.float() - y = (normed * (1.0 + scale.float()) + shift.float()).to(torch.bfloat16) - return y, ref_res_bf16 - - -def _ref_norm_ss(x, weight, bias, scale, shift, norm_type, eps): - if norm_type == "layer": - normed = F.layer_norm(x.float(), (D,), weight, bias, eps) - else: - normed = _ref_rms_norm(x.float(), weight, eps) * weight.float() - return (normed * (1.0 + scale.float()) + shift.float()).to(torch.bfloat16) - - -@pytest.fixture(autouse=True) -def cuda_setup(): - if not torch.cuda.is_available(): - pytest.skip("CUDA required") - if not hasattr(torch.version, "hip") or not torch.version.hip: - pytest.skip("ROCm/HIP required for FlyDSL kernels") - torch.manual_seed(42) - - -FUSED_CASES = [ - ("rms", 1, 16), - ("rms", 2, 16), - ("layer", 2, 16), - ("rms", 1, 90000), -] - - -@pytest.mark.parametrize("norm_type,B,L", FUSED_CASES) -def test_fused_residual_norm_scale_shift(norm_type, B, L): - from sglang.kernels.ops.diffusion.flydsl.fused_residual_norm import ( - flydsl_fused_residual_norm_scale_shift, - ) - - residual = torch.randn(B, L, D, device=DEVICE, dtype=torch.bfloat16) - x = torch.randn(B, L, D, device=DEVICE, dtype=torch.bfloat16) - gate = torch.randn(B, 1, D, device=DEVICE, dtype=torch.bfloat16) - weight = torch.randn(D, device=DEVICE, dtype=torch.float32) - bias = ( - torch.randn(D, device=DEVICE, dtype=torch.float32) - if norm_type == "layer" - else None - ) - scale = torch.randn(B, 1, D, device=DEVICE, dtype=torch.bfloat16) - shift = torch.randn(B, 1, D, device=DEVICE, dtype=torch.bfloat16) - - y, res_out = flydsl_fused_residual_norm_scale_shift( - residual, - x, - gate, - weight, - bias, - scale, - shift, - norm_type, - EPS, - ) - y_ref, res_ref = _ref_fused_residual_norm_ss( - residual, - x, - gate, - weight, - bias, - scale, - shift, - norm_type, - EPS, - ) - torch.testing.assert_close(res_out, res_ref, atol=5e-2, rtol=5e-2) - torch.testing.assert_close(y, y_ref, atol=1.0, rtol=5e-2) - - -NSS_CASES = [ - ("rms", 2, 16), - ("layer", 2, 16), - ("rms", 1, 90000), - ("layer", 1, 90000), -] - - -@pytest.mark.parametrize("norm_type,B,L", NSS_CASES) -def test_norm_scale_shift(norm_type, B, L): - from sglang.kernels.ops.diffusion.flydsl.fused_residual_norm import ( - flydsl_norm_scale_shift, - ) - - x = torch.randn(B, L, D, device=DEVICE, dtype=torch.bfloat16) - weight = torch.randn(D, device=DEVICE, dtype=torch.float32) - bias = ( - torch.randn(D, device=DEVICE, dtype=torch.float32) - if norm_type == "layer" - else None - ) - scale = torch.randn(B, 1, D, device=DEVICE, dtype=torch.bfloat16) - shift = torch.randn(B, 1, D, device=DEVICE, dtype=torch.bfloat16) - - y = flydsl_norm_scale_shift(x, weight, bias, scale, shift, norm_type, EPS) - y_ref = _ref_norm_ss(x, weight, bias, scale, shift, norm_type, EPS) - torch.testing.assert_close(y, y_ref, atol=1.0, rtol=5e-2) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/kernels/ops/diffusion/test_fused_gate_rmsnorm.py b/test/registered/kernels/ops/diffusion/test_fused_gate_rmsnorm.py deleted file mode 100644 index 57d72ecd17d1..000000000000 --- a/test/registered/kernels/ops/diffusion/test_fused_gate_rmsnorm.py +++ /dev/null @@ -1,54 +0,0 @@ -"""Core checks for the quality-gated fused gate-RMSNorm path.""" - -import sys - -import pytest -import torch -import torch.nn as nn -import torch.nn.functional as F - -from sglang.kernels.ops.diffusion import fused_gate_rmsnorm as fgn -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=4, stage="base-b-kernel-unit", runner_config="1-gpu-large") -pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") -DIM, EPS = 4608, 1e-5 # Ideogram 4 hidden size / norm_eps - - -class _Site(nn.Module): - def __init__(self, dtype=torch.bfloat16): - super().__init__() - self.norm = nn.RMSNorm(DIM, eps=EPS, device="cuda", dtype=dtype) - fgn.mark_fused_gate_rmsnorm_site(self, ("norm",)) - - -def test_fused_matches_ideogram_reference(): - torch.manual_seed(0) - site = _Site() - w = site.norm.weight.data - x = torch.randn(1, 64, DIM, device="cuda", dtype=torch.bfloat16) - residual = torch.randn_like(x) - # adaln-style strided chunks, as produced by Ideogram's modulation .chunk() - mods = torch.randn(1, 1, 2 * DIM, device="cuda", dtype=torch.bfloat16) - scale, gate = mods.chunk(2, dim=-1) - assert fgn.mount_fused_gate_rmsnorm(site) - got_scale = fgn.fused_rmsnorm_scale(x, w, 1.0 + scale, EPS) - got_gate = fgn.fused_rmsnorm_tanh_residual(x, gate, residual, w, EPS) - ref_scale = F.rms_norm(x, (DIM,), w, EPS) * (1.0 + scale) - ref_gate = residual + torch.tanh(gate) * F.rms_norm(x, (DIM,), w, EPS) - # fused path uses bf16-native norm statistics: close, not bit-exact - torch.testing.assert_close(got_scale, ref_scale, atol=8e-2, rtol=4e-2) - torch.testing.assert_close(got_gate, ref_gate, atol=8e-2, rtol=4e-2) - - -def test_mount_guards_all_or_nothing(): - good, bad = _Site(), _Site(torch.float32) - assert not fgn.mount_fused_gate_rmsnorm(nn.ModuleList([good, bad])) - assert not fgn.fused_gate_rmsnorm_active(good) - assert fgn.mount_fused_gate_rmsnorm(good) - fgn.unmount_fused_gate_rmsnorm(good) - assert not fgn.fused_gate_rmsnorm_active(good) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__])) diff --git a/test/registered/kernels/ops/diffusion/test_fused_linear_gelu.py b/test/registered/kernels/ops/diffusion/test_fused_linear_gelu.py deleted file mode 100644 index e30a9f8c9d14..000000000000 --- a/test/registered/kernels/ops/diffusion/test_fused_linear_gelu.py +++ /dev/null @@ -1,86 +0,0 @@ -"""Core checks for the quality-gated linear + tanh-GELU fusion.""" - -import sys - -import pytest -import torch -import torch.nn as nn -import torch.nn.functional as F - -from sglang.kernels.ops.diffusion import fused_linear_gelu as gelu -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=4, stage="base-b-kernel-unit", runner_config="1-gpu-large") -pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") - - -class _Site(nn.Module): - def __init__(self, dtype=torch.bfloat16, bias=True): - super().__init__() - self.proj = nn.Linear(64, 256, bias=bias, device="cuda", dtype=dtype) - gelu.mark_fused_gelu_site(self, "proj") - - def forward(self, x): - if gelu.fused_gelu_active(self) and gelu.can_fuse_linear_gelu(self.proj, x): - return gelu.fused_linear_gelu_tanh(x, self.proj.weight, self.proj.bias) - return F.gelu(self.proj(x), approximate="tanh") - - -@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) -def test_fused_matches_reference(dtype): - torch.manual_seed(0) - site = _Site(dtype) - x = torch.randn(512, 64, device="cuda", dtype=dtype) - ref = site(x) - assert gelu.mount_fused_linear_gelu(site) - atol = 2e-2 if dtype == torch.bfloat16 else 4e-3 - torch.testing.assert_close(site(x), ref, atol=atol, rtol=2e-2) - - -def test_flux_gelu_proj_site(): - """FLUX.1 shared-FF site: gate off is bit-exact, gate on is close.""" - from sglang.multimodal_gen.runtime.models.dits.flux import FluxFusedGELUProj - - torch.manual_seed(0) - proj = nn.Linear(3072, 12288, device="cuda", dtype=torch.bfloat16) - site = FluxFusedGELUProj(proj) - x = torch.randn(1, 512, 3072, device="cuda", dtype=torch.bfloat16) - ref = F.gelu(proj(x), approximate="tanh") - - assert torch.equal(site(x), ref) # unmounted default: bit-exact reference - assert gelu.mount_fused_linear_gelu(site) - torch.testing.assert_close(site(x), ref, atol=2e-2, rtol=2e-2) - gelu.unmount_fused_linear_gelu(site) - assert torch.equal(site(x), ref) - - -def test_mount_guards_and_lossless_path(): - torch.manual_seed(0) - good, bad = _Site(), _Site(torch.float32) - model = nn.ModuleList([good, bad]) - assert not gelu.mount_fused_linear_gelu(model) - assert not gelu.fused_gelu_active(good) - - x = torch.randn(16, 64, device="cuda", dtype=torch.bfloat16) - ref = good(x) - assert gelu.mount_fused_linear_gelu(good) - gelu.unmount_fused_linear_gelu(good) - assert torch.equal(good(x), ref) - - no_bias = nn.Linear(8, 8, bias=False, device="cuda", dtype=torch.bfloat16) - assert not gelu.can_fuse_linear_gelu_static(no_bias) - assert not gelu.can_fuse_linear_gelu(good.proj, x.float()) - - -@torch.no_grad() -def test_mounted_site_torch_compile_fullgraph(): - site = _Site() - x = torch.randn(16, 64, device="cuda", dtype=torch.bfloat16) - assert gelu.mount_fused_linear_gelu(site) - expected = site(x) - actual = torch.compile(site, fullgraph=True)(x) - torch.testing.assert_close(actual, expected, atol=2e-2, rtol=2e-2) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__])) diff --git a/test/registered/kernels/ops/diffusion/test_fused_ln_modulate.py b/test/registered/kernels/ops/diffusion/test_fused_ln_modulate.py deleted file mode 100644 index ebc6ed72ca03..000000000000 --- a/test/registered/kernels/ops/diffusion/test_fused_ln_modulate.py +++ /dev/null @@ -1,82 +0,0 @@ -import pytest -import torch -import torch.nn as nn - -from sglang.kernels.ops.diffusion.fused_ln_modulate import ( - can_fuse_ln_modulate, - fused_ln_modulate, - fused_ln_modulate_active, - mark_fused_ln_modulate_site, - mount_fused_ln_modulate, - unmount_fused_ln_modulate, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large") - - -@pytest.fixture(autouse=True) -def cuda_setup(): - if not torch.cuda.is_available(): - pytest.skip("CUDA required") - torch.cuda.manual_seed(0) - - -@pytest.mark.parametrize("seq_len", [4096, 512]) -def test_fused_ln_modulate_matches_reference(seq_len): - x = torch.randn((1, seq_len, 3072), device="cuda", dtype=torch.bfloat16) - scale = torch.randn((1, 3072), device="cuda", dtype=torch.bfloat16) - shift = torch.randn_like(scale) - assert can_fuse_ln_modulate(x, scale, shift) - out = fused_ln_modulate(x, scale, shift, eps=1e-6) - norm = nn.LayerNorm(3072, eps=1e-6, elementwise_affine=False).cuda() - ref = norm(x) * (1 + scale[:, None]) + shift[:, None] - # Contract: bf16 rounding-order-level difference only, not bit-exact. - torch.testing.assert_close(out, ref, atol=0.0625, rtol=0.05) - - -def test_fused_ln_modulate_guards_and_mount_protocol(): - x = torch.randn((2, 64, 3072), device="cuda", dtype=torch.bfloat16) - row = torch.randn((2, 3072), device="cuda", dtype=torch.bfloat16) - assert not can_fuse_ln_modulate(x, row, row) # folded affine needs B == 1 - root = nn.Module() - root.child = nn.Module() - mark_fused_ln_modulate_site(root.child) - assert not fused_ln_modulate_active(root.child) - assert mount_fused_ln_modulate(root) - assert fused_ln_modulate_active(root.child) - unmount_fused_ln_modulate(root) - assert not fused_ln_modulate_active(root.child) - assert not mount_fused_ln_modulate(nn.Module()) # no marked sites - - -@torch.no_grad() -def test_mounted_ln_modulate_site_torch_compile_fullgraph(): - class Site(nn.Module): - def __init__(self): - super().__init__() - mark_fused_ln_modulate_site(self) - - def forward(self, x, scale, shift): - if fused_ln_modulate_active(self) and can_fuse_ln_modulate(x, scale, shift): - return fused_ln_modulate(x, scale, shift, eps=1e-6) - return ( - nn.functional.layer_norm(x, (x.shape[-1],), eps=1e-6) - * (1 + scale[:, None]) - + shift[:, None] - ) - - site = Site() - assert mount_fused_ln_modulate(site) - x = torch.randn(1, 64, 128, device="cuda", dtype=torch.bfloat16) - scale = torch.randn(1, 128, device="cuda", dtype=torch.bfloat16) - shift = torch.randn_like(scale) - expected = site(x, scale, shift) - actual = torch.compile(site, fullgraph=True)(x, scale, shift) - torch.testing.assert_close(actual, expected, atol=0.0625, rtol=0.05) - - -if __name__ == "__main__": - import sys - - sys.exit(pytest.main([__file__])) diff --git a/test/registered/kernels/ops/diffusion/test_fused_norm_scale_shift.py b/test/registered/kernels/ops/diffusion/test_fused_norm_scale_shift.py deleted file mode 100644 index 3219aa6878b8..000000000000 --- a/test/registered/kernels/ops/diffusion/test_fused_norm_scale_shift.py +++ /dev/null @@ -1,251 +0,0 @@ -import sys -from typing import Optional, Tuple - -import pytest -import torch -from einops import rearrange -from torch import Tensor - -from sglang.kernels.ops.diffusion.cutedsl.scale_residual_norm_scale_shift import ( - fused_norm_scale_shift, - fused_scale_residual_norm_scale_shift, - validate_scale_shift, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=28, stage="base-b-kernel-unit", runner_config="1-gpu-large") - -DEVICE = "cuda" -SHAPE_MAP = { - "1": lambda B, S, F, D: (1,), - "D": lambda B, S, F, D: (D,), - "1D": lambda B, S, F, D: (1, D), - "BD": lambda B, S, F, D: (B, D), - "11D": lambda B, S, F, D: (1, 1, D), - "B1D": lambda B, S, F, D: (B, 1, D), - "1SD": lambda B, S, F, D: (1, S, D), - "BSD": lambda B, S, F, D: (B, S, D), - "BF1D": lambda B, S, F, D: (B, F, 1, D), -} -SHAPES = [ - # (B, S, F, D) - (1, 115200, 1, 3072), # Hunyuan - (1, 32760, 1, 1536), # Wan - (1, 6, 1, 3072), # Qwen - (1, 1024, 8, 3072), - (4, 512, 16, 3072), -] -DTYPES = [torch.float16, torch.bfloat16, torch.float32] -NORM_TYPES = ["layer", "rms"] -AFFINE_MODES = ["D", "NAT"] -INDEX_MODES = ["BSD", "1", "1SD", "BD", "B1D", "D", "1D", "11D", "BF1D"] - - -def _tol(dtype: torch.dtype): - return 1e-5 if dtype == torch.float32 else 5e-2 - - -@pytest.fixture(autouse=True) -def cuda_setup(): - if not torch.cuda.is_available(): - pytest.skip("CUDA required") - torch.cuda.manual_seed(0) - - -def _apply_scale_shift(y: Tensor, scale: Tensor, shift: Tensor) -> Tensor: - if scale.ndim == 4: - num_frame = scale.shape[1] - return rearrange( - rearrange(y, "b (f l) d -> b f l d", f=num_frame) * (1 + scale) + shift, - "b f l d -> b (f l) d", - ) - else: - scale = rearrange(scale, "b d -> b 1 d") if scale.ndim == 2 else scale - shift = rearrange(shift, "b d -> b 1 d") if shift.ndim == 2 else shift - return y * (1 + scale) + shift - - -def fused_norm_scale_shift_ref( - x: Tensor, - weight: Optional[Tensor], - bias: Optional[Tensor], - scale: Tensor, - shift: Tensor, - norm_type: str, - eps: float, -) -> Tensor: - original_dtype = x.dtype - x, weight, bias, scale, shift = ( - v.float() if v is not None else v for v in [x, weight, bias, scale, shift] - ) - if norm_type == "layer": - norm = torch.layer_norm(x, x.shape[-1:], eps=eps, weight=weight, bias=bias) - else: - norm = torch.rms_norm(x, x.shape[-1:], eps=eps, weight=weight) - return _apply_scale_shift(norm, scale, shift).to(original_dtype) - - -def fused_scale_residual_norm_scale_shift_ref( - residual: Tensor, - x: Tensor, - gate: Optional[Tensor] | int, - weight: Optional[Tensor], - bias: Optional[Tensor], - scale: Tensor, - shift: Tensor, - norm_type: str, - eps: float, -): - original_dtype = x.dtype - residual, x, gate, weight, bias, scale, shift = ( - v.float() if isinstance(v, Tensor) else v - for v in [residual, x, gate, weight, bias, scale, shift] - ) - if isinstance(gate, int): - x = residual + gate * x - else: - if gate.ndim == 4: - num_frame = gate.shape[1] - x_fld = rearrange(x, "b (f l) d -> b f l d", f=num_frame) - x = residual + rearrange(x_fld * gate, "b f l d -> b (f l) d") - else: - gate = rearrange(gate, "b d -> b 1 d") if gate.ndim == 2 else gate - x = residual + gate * x - if norm_type == "layer": - norm = torch.layer_norm(x, x.shape[-1:], eps=eps, weight=weight, bias=bias) - else: - norm = torch.rms_norm(x, x.shape[-1:], eps=eps, weight=weight) - y_ref = _apply_scale_shift(norm, scale, shift) - return y_ref.to(original_dtype), x.to(original_dtype) - - -def _make_tensor(index_mode: str, shape: Tuple, dtype: torch.dtype): - if index_mode == "NAT": - return None - return torch.randn(*SHAPE_MAP[index_mode](*shape), device=DEVICE, dtype=dtype) - - -def test_validate_scale_shift_rejects_non_divisible_frames(): - with pytest.raises(ValueError, match=r"S\(10\) must be divisible by F\(4\)"): - validate_scale_shift( - torch.empty((1, 4, 1, 256), device=DEVICE, dtype=torch.float16), - 1, - 10, - 256, - ) - - -@torch.no_grad() -def run_norm_scale_shift( - shape=SHAPES[0], - dtype=DTYPES[0], - affine_dtype=DTYPES[0], - scale_dtype=DTYPES[0], - shift_dtype=DTYPES[0], - norm_type=NORM_TYPES[0], - affine_mode=AFFINE_MODES[0], - scale_mode="BSD", - shift_mode="BSD", - eps=1e-5, -): - x = _make_tensor("BSD", shape, dtype) - weight = _make_tensor(affine_mode, shape, affine_dtype) - bias = _make_tensor(affine_mode, shape, affine_dtype) - scale = _make_tensor(scale_mode, shape, scale_dtype) - shift = _make_tensor(shift_mode, shape, shift_dtype) - y_dev = fused_norm_scale_shift(x, weight, bias, scale, shift, norm_type, eps) - y_ref = fused_norm_scale_shift_ref(x, weight, bias, scale, shift, norm_type, eps) - torch.testing.assert_close(y_dev, y_ref, atol=_tol(dtype), rtol=_tol(dtype)) - - -@torch.no_grad() -def run_scale_resi_norm_scale_shift( - shape=SHAPES[0], - dtype=DTYPES[0], - affine_dtype=DTYPES[0], - scale_dtype=DTYPES[0], - shift_dtype=DTYPES[0], - norm_type=NORM_TYPES[0], - affine_mode=AFFINE_MODES[0], - gate_mode="B1D", - scale_mode="BSD", - shift_mode="BSD", - eps=1e-5, -): - residual = _make_tensor("BSD", shape, dtype) - x = _make_tensor("BSD", shape, dtype) - gate = _make_tensor(gate_mode, shape, dtype) - weight = _make_tensor(affine_mode, shape, affine_dtype) - bias = _make_tensor(affine_mode, shape, affine_dtype) - scale = _make_tensor(scale_mode, shape, scale_dtype) - shift = _make_tensor(shift_mode, shape, shift_dtype) - y_dev, res_dev = fused_scale_residual_norm_scale_shift( - residual, x, gate, weight, bias, scale, shift, norm_type, eps - ) - y_ref, res_ref = fused_scale_residual_norm_scale_shift_ref( - residual, x, gate, weight, bias, scale, shift, norm_type, eps - ) - torch.testing.assert_close(y_dev, y_ref, atol=_tol(dtype), rtol=_tol(dtype)) - torch.testing.assert_close(res_dev, res_ref, atol=_tol(dtype), rtol=_tol(dtype)) - - -@pytest.mark.parametrize("norm_type", NORM_TYPES) -class TestFusedNormScaleShift: - @pytest.mark.parametrize("shape", SHAPES) - @pytest.mark.parametrize("dtype", DTYPES) - def test_shape_dtype(self, shape, dtype, norm_type): - run_norm_scale_shift(shape=shape, dtype=dtype, norm_type=norm_type) - - @pytest.mark.parametrize("dtype", DTYPES) - def test_dtype_0(self, dtype, norm_type): - run_norm_scale_shift(affine_dtype=dtype, norm_type=norm_type) - - @pytest.mark.parametrize("dtype", DTYPES) - def test_dtype_1(self, dtype, norm_type): - run_norm_scale_shift(scale_dtype=dtype, shift_dtype=dtype, norm_type=norm_type) - - @pytest.mark.parametrize("affine_mode", AFFINE_MODES) - def test_normtype_affine(self, affine_mode, norm_type): - run_norm_scale_shift(affine_mode=affine_mode, norm_type=norm_type) - - @pytest.mark.parametrize("index_mode", INDEX_MODES) - def test_index_mode(self, index_mode, norm_type): - run_norm_scale_shift( - scale_mode=index_mode, shift_mode=index_mode, norm_type=norm_type - ) - - -@pytest.mark.parametrize("norm_type", NORM_TYPES) -class TestFusedScaleResidualNormScaleShift: - @pytest.mark.parametrize("shape", SHAPES) - @pytest.mark.parametrize("dtype", DTYPES) - def test_shape_dtype(self, shape, dtype, norm_type): - run_scale_resi_norm_scale_shift(shape=shape, dtype=dtype, norm_type=norm_type) - - @pytest.mark.parametrize("dtype", DTYPES) - def test_dtype_0(self, dtype, norm_type): - run_scale_resi_norm_scale_shift(affine_dtype=dtype, norm_type=norm_type) - - @pytest.mark.parametrize("dtype", DTYPES) - def test_dtype_1(self, dtype, norm_type): - run_scale_resi_norm_scale_shift( - scale_dtype=dtype, shift_dtype=dtype, norm_type=norm_type - ) - - @pytest.mark.parametrize("affine_mode", AFFINE_MODES) - def test_normtype_affine(self, affine_mode, norm_type): - run_scale_resi_norm_scale_shift(affine_mode=affine_mode, norm_type=norm_type) - - @pytest.mark.parametrize("index_mode", INDEX_MODES) - def test_scale_shift_index_mode(self, index_mode, norm_type): - run_scale_resi_norm_scale_shift( - scale_mode=index_mode, shift_mode=index_mode, norm_type=norm_type - ) - - @pytest.mark.parametrize("index_mode", INDEX_MODES) - def test_gate_index_mode(self, index_mode, norm_type): - run_scale_resi_norm_scale_shift(gate_mode=index_mode, norm_type=norm_type) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/kernels/ops/diffusion/test_glm_image_ln_modulate.py b/test/registered/kernels/ops/diffusion/test_glm_image_ln_modulate.py deleted file mode 100644 index a3d177218b1d..000000000000 --- a/test/registered/kernels/ops/diffusion/test_glm_image_ln_modulate.py +++ /dev/null @@ -1,55 +0,0 @@ -"""GLM-Image fused LN+modulate / qk-LN fast paths must stay bit-exact vs eager.""" - -import pytest -import torch - -import sglang.multimodal_gen.runtime.models.dits.glm_image as glm_image -from sglang.multimodal_gen.runtime.models.dits.glm_image import ( - _eager_ln_modulate, - _glm_ln_modulate, - _glm_qk_layernorm, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=4, stage="base-b-kernel-unit", runner_config="1-gpu-large") -pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") - - -@pytest.mark.parametrize("shape", [(1, 4096, 4096), (2, 301, 4096), (1, 1, 2560)]) -def test_fused_ln_modulate_is_bit_exact(shape): - # (1, 4096, 4096) is the real GLM-Image image-stream shape (1024^2, - # hidden 4096); the others cover the text stream and another hidden. - torch.manual_seed(0) - batch, seq, hidden = shape - norm = torch.nn.LayerNorm(hidden, eps=1e-5, elementwise_affine=False).cuda() - x = (torch.randn(batch, seq, hidden, device="cuda") * 8).bfloat16() - emb = torch.randn(batch, 12 * hidden, device="cuda").bfloat16() - chunks = emb.chunk(12, dim=1) # strided adaLN projection views - shift, scale = chunks[0], chunks[2] - out = _glm_ln_modulate(norm, x, scale, shift, x.dtype) - assert torch.equal(out, _eager_ln_modulate(norm, x, scale, shift, x.dtype)) - assert glm_image._GLM_LN_MOD.verified - assert not glm_image._GLM_LN_MOD.disabled - - -@pytest.mark.parametrize("shape", [(1, 4360, 32, 128), (2, 37, 3, 40), (1, 129, 5, 64)]) -def test_fused_qk_head_layernorm_is_bit_exact(shape): - # (1, 4360, 32, 128) is the real GLM-Image q/k shape (text + image - # tokens, 32 heads of dim 128); the others cover partially-filled warps. - torch.manual_seed(1) - batch, seq, heads, head_dim = shape - norm_q = torch.nn.LayerNorm(head_dim, eps=1e-5, elementwise_affine=False).cuda() - norm_k = torch.nn.LayerNorm(head_dim, eps=1e-5, elementwise_affine=False).cuda() - q = (torch.randn(batch, seq, heads, head_dim, device="cuda") * 5).bfloat16() - k = (torch.randn(batch, seq, heads, head_dim, device="cuda") * 5).bfloat16() - q_out, k_out = _glm_qk_layernorm(norm_q, norm_k, q, k, q.dtype) - assert torch.equal(q_out, norm_q(q).to(q.dtype)) - assert torch.equal(k_out, norm_k(k).to(k.dtype)) - assert glm_image._GLM_QK_LN.verified - assert not glm_image._GLM_QK_LN.disabled - - -if __name__ == "__main__": - import sys - - sys.exit(pytest.main([__file__])) diff --git a/test/registered/kernels/ops/diffusion/test_group_norm_silu.py b/test/registered/kernels/ops/diffusion/test_group_norm_silu.py deleted file mode 100644 index d04aaf340a50..000000000000 --- a/test/registered/kernels/ops/diffusion/test_group_norm_silu.py +++ /dev/null @@ -1,104 +0,0 @@ -import sys - -import pytest -import torch -import torch.nn as nn -import torch.nn.functional as F - -from sglang.kernels.ops.diffusion.group_norm_silu import apply_group_norm_silu -from sglang.kernels.ops.diffusion.triton.group_norm_silu import triton_group_norm_silu -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci - -register_cuda_ci(est_time=8, stage="base-b-kernel-unit", runner_config="1-gpu-large") -register_amd_ci(est_time=15, suite="nightly-amd-kernel-1-gpu", nightly=True) - -DEVICE = "cuda" -DTYPES = [torch.float16, torch.bfloat16, torch.float32] -TEST_CASES = [ - pytest.param((2, 64, 32, 32), 32, id="image_2d"), - pytest.param((1, 64, 4, 16, 16), 32, id="video_3d"), - pytest.param((4, 128), 32, id="token_2d"), -] -LARGE_TILE_CASE = ((1, 128, 20, 256, 256), 32) - - -def _tol(dtype: torch.dtype) -> tuple[float, float]: - if dtype == torch.float32: - return 1e-5, 1e-5 - if dtype == torch.bfloat16: - return 7e-2, 2e-2 - return 3e-3, 3e-3 - - -@pytest.fixture(autouse=True) -def cuda_setup(): - if not torch.cuda.is_available(): - pytest.skip("CUDA required") - torch.cuda.manual_seed(0) - - -def _reference( - x: torch.Tensor, - weight: torch.Tensor, - bias: torch.Tensor, - num_groups: int, - eps: float = 1e-5, -) -> torch.Tensor: - return F.silu(F.group_norm(x, num_groups, weight=weight, bias=bias, eps=eps)) - - -@torch.no_grad() -@pytest.mark.parametrize("shape,num_groups", TEST_CASES) -@pytest.mark.parametrize("dtype", DTYPES) -def test_triton_group_norm_silu( - shape: tuple[int, ...], num_groups: int, dtype: torch.dtype -) -> None: - channels = shape[1] - x = torch.randn(shape, device=DEVICE, dtype=dtype) - weight = torch.randn(channels, device=DEVICE, dtype=dtype) - bias = torch.randn(channels, device=DEVICE, dtype=dtype) - - actual = triton_group_norm_silu(x, weight, bias, num_groups=num_groups) - expected = _reference(x, weight, bias, num_groups) - - atol, rtol = _tol(dtype) - torch.testing.assert_close(actual, expected, atol=atol, rtol=rtol) - - -@torch.no_grad() -@pytest.mark.parametrize("shape,num_groups", TEST_CASES[:2]) -@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) -def test_apply_group_norm_silu( - shape: tuple[int, ...], - num_groups: int, - dtype: torch.dtype, -) -> None: - norm = nn.GroupNorm(num_groups, shape[1], eps=1e-5, affine=True).to( - device=DEVICE, dtype=dtype - ) - activation = nn.SiLU() - hidden_states = torch.randn(shape, device=DEVICE, dtype=dtype) - - actual = apply_group_norm_silu(hidden_states, norm, activation) - expected = activation(norm(hidden_states)) - - atol, rtol = _tol(dtype) - torch.testing.assert_close(actual, expected, atol=atol, rtol=rtol) - - -@torch.no_grad() -def test_triton_group_norm_silu_large_tile_bf16() -> None: - shape, num_groups = LARGE_TILE_CASE - x = torch.randn(shape, device=DEVICE, dtype=torch.bfloat16) - weight = torch.randn(shape[1], device=DEVICE, dtype=torch.bfloat16) - bias = torch.randn(shape[1], device=DEVICE, dtype=torch.bfloat16) - - actual = triton_group_norm_silu(x, weight, bias, num_groups=num_groups) - expected = _reference(x, weight, bias, num_groups) - - atol, rtol = _tol(torch.bfloat16) - torch.testing.assert_close(actual, expected, atol=atol, rtol=rtol) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/kernels/ops/diffusion/test_hunyuanvideo_eager_fusions.py b/test/registered/kernels/ops/diffusion/test_hunyuanvideo_eager_fusions.py deleted file mode 100644 index 8f438beaf66d..000000000000 --- a/test/registered/kernels/ops/diffusion/test_hunyuanvideo_eager_fusions.py +++ /dev/null @@ -1,94 +0,0 @@ -"""HunyuanVideo eager QKV/RoPE and quality-gated QKNorm tests.""" - -import sys -from unittest.mock import patch - -import pytest -import torch - -import sglang.kernels.ops.diffusion.hunyuan_qknorm as hunyuan_qknorm -from sglang.kernels.ops.diffusion.hunyuan_qknorm import ( - mark_hunyuan_qknorm_site, - mount_hunyuan_qknorm, - unmount_hunyuan_qknorm, -) -from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm -from sglang.multimodal_gen.runtime.layers.rotary_embedding.utils import ( - _apply_rotary_emb, -) -from sglang.multimodal_gen.runtime.models.dits.hunyuanvideo import ( - _hunyuan_pack_qkv, - _hunyuan_qknorm, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=8, stage="base-b-kernel-unit", runner_config="1-gpu-large") -pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") - - -@pytest.mark.parametrize("img_tokens,txt_tokens", [(257, 31), (4096, 256)]) -def test_hunyuan_qkv_rope_pack_is_bit_exact(img_tokens, txt_tokens): - torch.manual_seed(0) - shape_img = (1, img_tokens, 24, 128) - shape_txt = (1, txt_tokens, 24, 128) - img_q, img_k, img_v = ( - torch.randn(shape_img, device="cuda", dtype=torch.bfloat16) for _ in range(3) - ) - txt_q, txt_k, txt_v = ( - torch.randn(shape_txt, device="cuda", dtype=torch.bfloat16) for _ in range(3) - ) - cos = torch.randn(img_tokens, 64, device="cuda") - sin = torch.randn_like(cos) - - q, k, v = _hunyuan_pack_qkv(img_q, img_k, img_v, txt_q, txt_k, txt_v, cos, sin) - q_ref = torch.cat( - (_apply_rotary_emb(img_q, cos, sin, is_neox_style=False), txt_q), dim=1 - ) - k_ref = torch.cat( - (_apply_rotary_emb(img_k, cos, sin, is_neox_style=False), txt_k), dim=1 - ) - v_ref = torch.cat((img_v, txt_v), dim=1) - - assert torch.equal(q, q_ref) - assert torch.equal(k, k_ref) - assert torch.equal(v, v_ref) - - -def test_hunyuan_quality_qknorm_matches_rmsnorm(): - torch.manual_seed(1) - site = torch.nn.Module() - mark_hunyuan_qknorm_site(site) - q_norm = RMSNorm(128, eps=1e-6).to(device="cuda", dtype=torch.bfloat16) - k_norm = RMSNorm(128, eps=1e-6).to(device="cuda", dtype=torch.bfloat16) - packed = torch.randn(1, 257, 3, 24, 128, device="cuda", dtype=torch.bfloat16) - q, k = packed[:, :, 0], packed[:, :, 1] - q_ref = q_norm(q.contiguous()).to(q) - k_ref = k_norm(k.contiguous()).to(k) - - q_unmounted, k_unmounted = _hunyuan_qknorm(site, q, k, q_norm, k_norm) - assert torch.equal(q_unmounted, q_ref) - assert torch.equal(k_unmounted, k_ref) - - assert mount_hunyuan_qknorm(site) - q_out, k_out = _hunyuan_qknorm(site, q, k, q_norm, k_norm) - torch.testing.assert_close(q_out, q_ref, atol=2e-2, rtol=2e-2) - torch.testing.assert_close(k_out, k_ref, atol=2e-2, rtol=2e-2) - - unmount_hunyuan_qknorm(site) - q_unmounted, k_unmounted = _hunyuan_qknorm(site, q, k, q_norm, k_norm) - assert torch.equal(q_unmounted, q_ref) - assert torch.equal(k_unmounted, k_ref) - - -def test_hunyuan_quality_qknorm_stays_unmounted_without_cute_kernel(): - site = torch.nn.Module() - mark_hunyuan_qknorm_site(site) - - with patch.object(hunyuan_qknorm, "_get_qk_rmsnorm_cute", return_value=None): - assert not mount_hunyuan_qknorm(site) - - assert not hunyuan_qknorm._FUSION.is_enabled(site) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernels/ops/diffusion/test_import_surface.py b/test/registered/kernels/ops/diffusion/test_import_surface.py new file mode 100644 index 000000000000..482524581396 --- /dev/null +++ b/test/registered/kernels/ops/diffusion/test_import_surface.py @@ -0,0 +1,210 @@ +"""Guards that keep the ``diffusion`` package's import surface from eroding. + +The reorganization only stays useful if two invariants hold: + +1. runtime code imports from ``sglang.kernels.ops.diffusion`` and not from a + submodule, so the internal layout can move without touching call sites; +2. the facade's ``_EXPORTS`` table and the registry's ``_SPECS`` table both + point at symbols that actually exist. + +Neither is checkable by the type system, and both fail silently -- a stale +``_EXPORTS`` entry only raises when some model happens to call that kernel, on +a GPU, at serving time. These are pure-CPU tests: they read the tables and +resolve them with ``importlib``/``ast`` without importing torch backends. +""" + +import ast +import importlib +import pathlib +import subprocess +import sys + +import pytest + +from sglang.kernels.ops.diffusion import _EXPORTS, _SPECS +from sglang.kernels.registry import registry +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=4, suite="base-a-test-cpu") + +PACKAGE = "sglang.kernels.ops.diffusion" +_PACKAGE_DIR = pathlib.Path(importlib.import_module(PACKAGE).__file__ or "").parent +_REPO_ROOT = _PACKAGE_DIR.parents[4] # /python/sglang/kernels/ops/diffusion + +# Backend-specific test files may name a leaf module on purpose; everything +# else -- all runtime code -- must go through the facade. +_DEEP_IMPORT_ALLOWLIST = { + "python/sglang/multimodal_gen/test/unit/test_latent_upsampler_group_norm_silu.py", + "test/registered/kernels/ops/diffusion/test_model_fast_paths.py", + "test/registered/kernels/ops/diffusion/test_sites.py", +} + + +def _module_defines(module_path: str) -> set[str]: + """Top-level names bound by a submodule, without importing it. + + Importing would pull in Triton / CuTe-DSL / FlyDSL, none of which are + installed on the CPU CI lane -- so this reads the source instead. + """ + path = _PACKAGE_DIR / (module_path.replace(".", "/") + ".py") + if not path.exists(): + path = _PACKAGE_DIR / module_path.replace(".", "/") / "__init__.py" + assert path.exists(), f"{PACKAGE}.{module_path} does not exist" + + names: set[str] = set() + for node in ast.parse(path.read_text(encoding="utf-8")).body: + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)): + names.add(node.name) + elif isinstance(node, ast.Assign): + names.update(t.id for t in node.targets if isinstance(t, ast.Name)) + elif isinstance(node, ast.AnnAssign) and isinstance(node.target, ast.Name): + names.add(node.target.id) + elif isinstance(node, (ast.Import, ast.ImportFrom)): + names.update((a.asname or a.name).split(".")[0] for a in node.names) + elif isinstance(node, (ast.If, ast.Try)): + # Platform-conditional rebinds (``x = select_impl(...)``) and + # guarded defs still bind a public name. + for inner in ast.walk(node): + if isinstance(inner, (ast.FunctionDef, ast.ClassDef)): + names.add(inner.name) + elif isinstance(inner, ast.Assign): + names.update(t.id for t in inner.targets if isinstance(t, ast.Name)) + return names + + +def test_every_export_resolves_to_a_real_symbol(): + missing = [ + f"{symbol} -> {module}" + for symbol, module in sorted(_EXPORTS.items()) + if symbol not in _module_defines(module) + ] + assert not missing, f"stale _EXPORTS entries: {missing}" + + +def test_every_symbol_imported_from_the_facade_is_exported(): + """The reverse of the check above, and the one that actually bites. + + A missing ``_EXPORTS`` entry raises ``ImportError`` at module import, so a + module-level ``from ...diffusion import x`` fails loudly. A *function-local* + one -- the pattern used for optional backends -- fails only when that test + or code path runs, on the platform that has the backend. Enumerating the + call sites catches it here instead. + """ + unexported = set() + for root in ("python/sglang", "test", "benchmark"): + root_dir = _REPO_ROOT / root + if not root_dir.exists(): + continue + for path in root_dir.rglob("*.py"): + rel = path.relative_to(_REPO_ROOT).as_posix() + if rel.startswith("python/sglang/kernels/ops/diffusion/"): + continue + try: + tree = ast.parse(path.read_text(encoding="utf-8")) + except (SyntaxError, UnicodeDecodeError): + continue + for node in ast.walk(tree): + if isinstance(node, ast.ImportFrom) and node.module == PACKAGE: + unexported.update( + a.name + for a in node.names + if a.name not in _EXPORTS and not a.name.startswith("_") + ) + assert not unexported, f"imported but not in _EXPORTS: {sorted(unexported)}" + + +def test_every_registered_spec_target_resolves(): + missing = [] + for _op, _backend, target, _caps, _description in _SPECS: + module, _, attr = target.partition(":") + if attr not in _module_defines(module): + missing.append(target) + assert not missing, f"stale _SPECS targets: {missing}" + + +def test_registry_holds_the_diffusion_ops(): + # Registration happens at package import, is metadata-only, and is what + # ``select_kernel`` / the tracing tools read. + registered = {op for op in registry.ops() if op.startswith("diffusion.")} + assert {op for op, *_ in _SPECS} <= registered + + +def test_facade_rejects_unknown_attributes(): + module = sys.modules[PACKAGE] + with pytest.raises(AttributeError): + module.definitely_not_a_kernel + assert set(module.__all__) == set(_EXPORTS) + assert set(_EXPORTS) <= set(dir(module)) + + +def test_importing_the_package_does_not_import_any_leaf_module(): + """The reason ``__getattr__`` is lazy rather than a block of re-exports. + + The backends have disjoint, heavy, mutually-exclusive dependencies -- + Triton (CUDA/ROCm), CUTLASS/CuTe-DSL, FlyDSL (gfx950), MLX (Apple). If + ``_EXPORTS`` ever degrades into eager ``from .norm.x import y`` lines, all + of them become import-time requirements on every platform, which is how a + CPU-only or Apple install starts failing at ``import sglang``. + + Asserted on this package's own leaf modules rather than on ``triton`` in + ``sys.modules``: sibling operator groups import Triton for their own + reasons, so a global check would not isolate this package's behavior. + Run in a fresh interpreter because this process has already resolved + exports through the facade. + """ + code = ( + "import importlib, sys\n" + f"importlib.import_module('{PACKAGE}')\n" + f"prefix = '{PACKAGE}.'\n" + "leaves = [m for m in sys.modules if m.startswith(prefix)" + " and not m.endswith('__init__')]\n" + "print(','.join(sorted(m for m in leaves if '.' in m[len(prefix):]" + " or sys.modules[m].__file__ and not sys.modules[m].__file__" + ".endswith('__init__.py'))))\n" + ) + result = subprocess.run( + [sys.executable, "-c", code], capture_output=True, text=True, timeout=600 + ) + assert result.returncode == 0, result.stderr + leaked = [m for m in result.stdout.strip().split(",") if m] + assert not leaked, f"importing {PACKAGE} eagerly imported: {leaked}" + + +@pytest.mark.parametrize("root", ["python/sglang", "test", "benchmark"]) +def test_runtime_code_imports_only_through_the_facade(root): + root_dir = _REPO_ROOT / root + if not root_dir.exists(): # source checkouts only + pytest.skip(f"{root} not present in this install") + + offenders = [] + for path in root_dir.rglob("*.py"): + rel = path.relative_to(_REPO_ROOT).as_posix() + if rel.startswith("python/sglang/kernels/ops/diffusion/"): + continue # intra-package imports are the point of the subpackages + if rel in _DEEP_IMPORT_ALLOWLIST: + continue + try: + tree = ast.parse(path.read_text(encoding="utf-8")) + except (SyntaxError, UnicodeDecodeError): + continue + for node in ast.walk(tree): + if ( + isinstance(node, ast.ImportFrom) + and node.module + and node.module.startswith(f"{PACKAGE}.") + ): + offenders.append(f"{rel}:{node.lineno} imports {node.module}") + elif isinstance(node, ast.Import): + offenders.extend( + f"{rel}:{node.lineno} imports {a.name}" + for a in node.names + if a.name.startswith(f"{PACKAGE}.") + ) + assert not offenders, ( + "import from sglang.kernels.ops.diffusion instead of a submodule:\n " + + "\n ".join(offenders) + ) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernels/ops/diffusion/test_layout.py b/test/registered/kernels/ops/diffusion/test_layout.py new file mode 100644 index 000000000000..93c37d281250 --- /dev/null +++ b/test/registered/kernels/ops/diffusion/test_layout.py @@ -0,0 +1,582 @@ +"""``diffusion.layout``: data-movement kernels. + +Every kernel in this domain only moves values (plus zero fill, plus at most +one same-order add), so each one is *bitwise* identical to the aten chain it +replaces. That makes ``torch.equal`` -- not ``assert_close`` -- the right +assertion throughout this file; a tolerance here would hide a real bug. + +Covered: USP output head merge, Ulysses destination-major QKV pack, varlen +pack/scatter, causal Conv3d cat+pad (CUDA and Triton), and the Wan causal-VAE +cache kernels. +""" + +import sys +from unittest.mock import patch + +import pytest +import torch +import torch.nn.functional as F + +from sglang.kernels.jit.utils import get_ci_test_range +from sglang.kernels.ops.attention.flash_attention import flash_attn_varlen_func +from sglang.kernels.ops.diffusion import ( + build_inv_indices, + can_use_usp_merge_heads, + cat_pad_channels_last_3d, + dup_up3d_add, +) +from sglang.kernels.ops.diffusion import ( + fused_causal_conv3d_cat_pad as fused_causal_conv3d_cat_pad_triton, +) +from sglang.kernels.ops.diffusion import ( + fused_causal_conv3d_cat_pad_cuda, + fused_pack_qkv, + fused_scatter_to_padded, + pack_qkv_destination_major, + usp_merge_heads, +) +from sglang.multimodal_gen.runtime.layers.attention.backends import ( + flash_attn as _fa_backend, +) +from sglang.multimodal_gen.runtime.layers.attention.layer import build_varlen_mask_meta +from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci + +register_cuda_ci(est_time=110, stage="base-b-kernel-unit", runner_config="1-gpu-large") +register_cuda_ci(est_time=45, stage="base-b-kernel-unit", runner_config="4-gpu-b200") +# Nightly is not redundant: it sets SGLANG_JIT_KERNEL_RUN_FULL_TESTS=1, which +# expands the get_ci_test_range sweeps below. +register_cuda_ci(est_time=20, stage="nightly", runner_config="1-gpu-large") +register_amd_ci(est_time=10, stage="jit-kernel-unit", runner_config="amd") +register_amd_ci(est_time=15, suite="nightly-amd-kernel-1-gpu", nightly=True) + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") + +DEVICE = "cuda" + + +def _cl3d(shape, dtype): + return torch.randn(shape, device=DEVICE, dtype=dtype).contiguous( + memory_format=torch.channels_last_3d + ) + + +# --------------------------------------------------------------------------- +# USP output head merge +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize( + "world,seq,batch,h_local,head_dim", + [ + (4, 7936, 1, 14, 128), # H3 768p production shape (Ulysses 4) + (2, 64, 3, 4, 64), # batched + (4, 33, 2, 4, 100), # scalar fallback inside the CUDA kernel + ], +) +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16, torch.float32]) +def test_usp_merge_heads_bitwise(dtype, world, seq, batch, h_local, head_dim): + generator = torch.Generator(device=DEVICE).manual_seed(4321) + x = torch.randn( + world, + seq, + batch, + h_local, + head_dim, + dtype=dtype, + device=DEVICE, + generator=generator, + ) + assert can_use_usp_merge_heads(x) + out = usp_merge_heads(x) + ref = x.permute(2, 1, 0, 3, 4).contiguous() + assert out.shape == ref.shape + assert torch.equal(out, ref) + + +def test_usp_merge_heads_unsupported_inputs_use_exact_fallback(): + # The wrapper degrades to the aten permute for anything the fast path + # rejects -- a wrong rank, a transposed view, an empty leading dim, or a + # ROCm build -- so callers never need their own guard. + x = torch.randn(2, 4, 1, 4, 64, dtype=torch.bfloat16, device=DEVICE) + for value in (x.transpose(0, 1), x[:0], x[0]): + assert not can_use_usp_merge_heads(value) + if value.dim() == 5: + assert torch.equal( + usp_merge_heads(value), value.permute(2, 1, 0, 3, 4).contiguous() + ) + + with patch.object(torch.version, "hip", "6.3"): + assert not can_use_usp_merge_heads(x) + assert torch.equal(usp_merge_heads(x), x.permute(2, 1, 0, 3, 4).contiguous()) + + +# --------------------------------------------------------------------------- +# Ulysses destination-major QKV pack +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) +def test_pack_qkv_destination_major_is_bit_exact(dtype): + torch.manual_seed(0) + rows, world_size, global_heads, head_size = 17, 4, 12, 64 + q, k, v = ( + torch.randn(rows, global_heads, head_size, device=DEVICE, dtype=dtype) + for _ in range(3) + ) + + local_heads = global_heads // world_size + expected = torch.empty( + world_size, rows, local_heads, 3 * head_size, device=DEVICE, dtype=dtype + ) + for index, tensor in enumerate((q, k, v)): + shards = tensor.view(rows, world_size, local_heads, head_size).permute( + 1, 0, 2, 3 + ) + expected[..., index * head_size : (index + 1) * head_size].copy_(shards) + + assert torch.equal(pack_qkv_destination_major(q, k, v, world_size), expected) + + +def test_pack_qkv_destination_major_validates_inputs(): + q = torch.empty(2, 4, 8, device=DEVICE, dtype=torch.bfloat16) + with pytest.raises(ValueError, match="same 3D shape"): + pack_qkv_destination_major(q, q[:, :-1], q, 2) + with pytest.raises(ValueError, match="divide global_heads"): + pack_qkv_destination_major(q, q, q, 3) + with pytest.raises(ValueError, match="expected shape"): + pack_qkv_destination_major(q, q, q, 2, out=torch.empty_like(q)) + + +# --------------------------------------------------------------------------- +# Varlen pack / scatter +# --------------------------------------------------------------------------- + +VARLEN_DTYPES = get_ci_test_range([torch.bfloat16, torch.float16], [torch.bfloat16]) +# (name, bs, s_txt, s_img, num_heads, head_dim, valid_txt_lens) +VARLEN_SHAPES = get_ci_test_range( + [ + ("small_c2", 2, 64, 128, 4, 64, [32, 48]), + ("prod_c2", 2, 256, 1024, 24, 128, [128, 200]), + ("all_valid_b1", 1, 64, 128, 4, 64, [64]), + ("all_valid_b4", 4, 64, 128, 4, 64, [64, 64, 64, 64]), + ("c8_prod", 8, 256, 4096, 24, 128, [128, 200, 256, 100, 50, 256, 256, 50]), + # one batch with zero valid text tokens (image side still valid) + ("zero_txt_one_batch", 2, 64, 128, 4, 64, [0, 32]), + # bs=1 with no text validity (only image rows packed) + ("bs1_zero_txt", 1, 64, 128, 4, 64, [0]), + ], + [ + ("small_c2", 2, 64, 128, 4, 64, [32, 48]), + ("prod_c2", 2, 256, 1024, 24, 128, [128, 200]), + ("all_valid_b4", 4, 64, 128, 4, 64, [64, 64, 64, 64]), + ], +) + + +def _build_mask(bs, s_txt, s_img, valid_txt_lens): + mask = torch.zeros(bs, s_txt + s_img, dtype=torch.bool, device=DEVICE) + for b, vt in enumerate(valid_txt_lens): + mask[b, :vt] = True + mask[b, s_txt:] = True + return mask + + +def _build_meta(mask): + bs, seq = mask.shape + indices = mask.reshape(-1).nonzero(as_tuple=False).flatten() + return indices, build_inv_indices(indices, bs * seq) + + +@pytest.mark.parametrize("dtype", VARLEN_DTYPES) +@pytest.mark.parametrize("shape", VARLEN_SHAPES, ids=lambda s: s[0]) +def test_varlen_pack_matches_index_select(dtype, shape): + _, bs, s_txt, s_img, num_heads, head_dim, valid_txt_lens = shape + torch.manual_seed(0) + s = s_txt + s_img + indices, _ = _build_meta(_build_mask(bs, s_txt, s_img, valid_txt_lens)) + + q, k, v = ( + torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) + for _ in range(3) + ) + fused = fused_pack_qkv(q, k, v, indices) + for got, src in zip(fused, (q, k, v), strict=True): + want = src.reshape(bs * s, num_heads, head_dim).index_select(0, indices) + assert torch.equal(got, want) + + +@pytest.mark.parametrize("dtype", VARLEN_DTYPES) +@pytest.mark.parametrize("shape", VARLEN_SHAPES, ids=lambda s: s[0]) +def test_varlen_scatter_matches_index_copy(dtype, shape): + _, bs, s_txt, s_img, num_heads, head_dim, valid_txt_lens = shape + torch.manual_seed(1) + s = s_txt + s_img + mask = _build_mask(bs, s_txt, s_img, valid_txt_lens) + indices, inv_indices = _build_meta(mask) + + out_unpad = torch.randn( + indices.shape[0], num_heads, head_dim, dtype=dtype, device=DEVICE + ) + flat = torch.zeros(bs * s, num_heads, head_dim, dtype=dtype, device=DEVICE) + flat.index_copy_(0, indices, out_unpad) + out_ref = flat.view(bs, s, num_heads, head_dim) + + out_fused = fused_scatter_to_padded(out_unpad, inv_indices, bs, s) + assert torch.equal(out_ref, out_fused) + invalid = ~mask + if invalid.any(): + # Padding rows must be exactly zero, not merely small. + assert out_fused[invalid].abs().max().item() == 0.0 + + +def test_varlen_pack_handles_non_contiguous_input(): + # Q/K/V arrive as (B, H, S, D) permutes from attention; the helper must + # make them contiguous itself rather than reading the wrong strides. + torch.manual_seed(2) + bs, s_txt, s_img, num_heads, head_dim = 2, 64, 128, 4, 64 + indices, _ = _build_meta(_build_mask(bs, s_txt, s_img, [32, 48])) + + pre = torch.randn( + bs, num_heads, s_txt + s_img, head_dim, dtype=torch.bfloat16, device=DEVICE + ) + q, k, v = (torch.randn_like(pre).permute(0, 2, 1, 3) for _ in range(3)) + assert not q.is_contiguous() + + fused = fused_pack_qkv(q, k, v, indices) + for got, src in zip(fused, (q, k, v), strict=True): + want = src.contiguous().flatten(0, 1).index_select(0, indices) + assert torch.equal(got, want) + + +def test_build_inv_indices_matches_manual(): + torch.manual_seed(3) + bs, s = 2, 32 + mask = torch.bernoulli(torch.full((bs, s), 0.6, device=DEVICE)).to(torch.bool) + indices = mask.reshape(-1).nonzero(as_tuple=False).flatten() + + manual = torch.full((bs * s,), -1, dtype=torch.int32, device=DEVICE) + if indices.numel(): + manual[indices.long()] = torch.arange( + indices.numel(), dtype=torch.int32, device=DEVICE + ) + assert torch.equal(build_inv_indices(indices, bs * s), manual) + + +def test_varlen_empty_valid_set_handled(): + # An all-False mask is reachable (a request whose text side is fully + # masked): pack must return empty tensors and scatter an all-zero dense + # output rather than launching a degenerate grid. + bs, s, num_heads, head_dim = 2, 16, 4, 64 + indices = torch.zeros(bs, s, dtype=torch.bool, device=DEVICE).reshape(-1).nonzero() + indices = indices.flatten() + inv_indices = build_inv_indices(indices, bs * s) + assert indices.numel() == 0 + + q = torch.randn(bs, s, num_heads, head_dim, dtype=torch.bfloat16, device=DEVICE) + unpad = fused_pack_qkv(q, q.clone(), q.clone(), indices) + assert all(t.shape == (0, num_heads, head_dim) for t in unpad) + + out_padded = fused_scatter_to_padded(unpad[0], inv_indices, bs, s) + assert out_padded.shape == (bs, s, num_heads, head_dim) + assert out_padded.abs().max().item() == 0.0 + + +# The kernels above are unit-tested against index_select/index_copy_; this +# section drives them through the production USPAttention masked branch, where +# a wrong index layout would produce plausible-looking attention output rather +# than an obvious mismatch. + + +def _sdpa_with_key_mask(q, k, v, key_mask, softmax_scale): + """Reference: SDPA with a ``[B, S]`` key mask broadcast to ``[B, 1, 1, S]``.""" + q_ = q.transpose(1, 2) + k_ = k.transpose(1, 2) + v_ = v.transpose(1, 2) + mask = key_mask.to(dtype=q.dtype)[:, None, None, :] + mask = (mask - 1.0) * torch.finfo(q.dtype).max + out = F.scaled_dot_product_attention( + q_, + k_, + v_, + attn_mask=mask, + dropout_p=0.0, + is_causal=False, + scale=softmax_scale, + ) + return out.transpose(1, 2) + + +def _varlen_path(q, k, v, key_mask, softmax_scale): + """Production varlen path matching USPAttention.forward masked branch.""" + bs, seq = q.shape[0], q.shape[1] + meta = build_varlen_mask_meta(key_mask) + indices = meta["indices"] + if indices.shape[0] == 0: + return torch.zeros_like(q) + q_unpad, k_unpad, v_unpad = fused_pack_qkv(q, k, v, indices) + out_unpad = flash_attn_varlen_func( + q=q_unpad, + k=k_unpad, + v=v_unpad, + cu_seqlens_q=meta["cu_seqlens"], + cu_seqlens_k=meta["cu_seqlens"], + max_seqlen_q=meta["max_seqlen"], + max_seqlen_k=meta["max_seqlen"], + softmax_scale=softmax_scale, + causal=False, + ver=_fa_backend.fa_ver, + ) + return fused_scatter_to_padded(out_unpad, meta["inv_indices"], bs, seq) + + +@pytest.mark.parametrize("dtype", VARLEN_DTYPES) +@pytest.mark.parametrize("shape", VARLEN_SHAPES, ids=lambda s: s[0]) +def test_varlen_path_matches_sdpa_on_valid_rows(dtype, shape): + """Valid rows: varlen output ≈ SDPA output within FA tolerance.""" + _, bs, s_txt, s_img, num_heads, head_dim, valid_txt_lens = shape + torch.manual_seed(0) + s = s_txt + s_img + softmax_scale = head_dim**-0.5 + mask = _build_mask(bs, s_txt, s_img, valid_txt_lens) + q = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) + k = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) + v = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) + + out_sdpa = _sdpa_with_key_mask(q, k, v, mask, softmax_scale) + out_varlen = _varlen_path(q, k, v, mask, softmax_scale) + + valid = mask[..., None, None].expand_as(out_sdpa) + rtol = 1e-2 if dtype == torch.bfloat16 else 5e-3 + atol = 5e-2 if dtype == torch.bfloat16 else 1e-2 + torch.testing.assert_close( + out_sdpa[valid], + out_varlen[valid], + rtol=rtol, + atol=atol, + ) + + +@pytest.mark.parametrize("dtype", VARLEN_DTYPES) +@pytest.mark.parametrize("shape", VARLEN_SHAPES, ids=lambda s: s[0]) +def test_varlen_path_zeros_masked_rows(dtype, shape): + """Masked rows: varlen path produces exact zeros (documented contract).""" + _, bs, s_txt, s_img, num_heads, head_dim, valid_txt_lens = shape + torch.manual_seed(1) + s = s_txt + s_img + softmax_scale = head_dim**-0.5 + mask = _build_mask(bs, s_txt, s_img, valid_txt_lens) + q = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) + k = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) + v = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) + + out_varlen = _varlen_path(q, k, v, mask, softmax_scale) + + invalid = ~mask + if invalid.any(): + assert (out_varlen[invalid] == 0).all(), "masked rows must be zero-filled" + + +# --------------------------------------------------------------------------- +# Causal Conv3d cat + pad (CUDA JIT vs Triton) +# --------------------------------------------------------------------------- + +CONV3D_CASES = get_ci_test_range( + [ + (1024, 1, 30, 52, 1), + (1024, 1, 30, 52, 2), + (1024, 2, 60, 104, 1), + (1024, 2, 60, 104, 2), + (512, 4, 120, 208, 1), + (512, 4, 120, 208, 2), + (256, 4, 240, 416, 1), + (256, 4, 240, 416, 2), + ], + [(1024, 1, 30, 52, 1), (512, 4, 120, 208, 2)], +) + + +def _conv3d_inputs(channels, t_size, h_size, w_size, cache_t): + generator = torch.Generator(device=DEVICE) + generator.manual_seed(channels * 1009 + t_size * 251 + h_size + cache_t) + x = torch.randn( + (1, channels, t_size, h_size, w_size), + device=DEVICE, + dtype=torch.bfloat16, + generator=generator, + ) + cache_x = torch.randn( + (1, channels, cache_t, h_size, w_size), + device=DEVICE, + dtype=torch.bfloat16, + generator=generator, + ) + return x, cache_x, (1, 1, 1, 1, cache_t, 0) + + +@pytest.mark.parametrize("channels,t_size,h_size,w_size,cache_t", CONV3D_CASES) +def test_causal_conv3d_cat_pad_cuda_matches_triton( + channels, t_size, h_size, w_size, cache_t +): + x, cache_x, padding = _conv3d_inputs(channels, t_size, h_size, w_size, cache_t) + actual = fused_causal_conv3d_cat_pad_cuda(x, cache_x, padding) + expected = fused_causal_conv3d_cat_pad_triton(x, cache_x, padding) + assert torch.equal(actual, expected) + + +def test_causal_conv3d_cat_pad_torch_compile(): + # The CUDA path is a registered custom op, so a fullgraph compile must not + # graph-break on it. + x, cache_x, padding = _conv3d_inputs(1024, 1, 30, 52, 1) + + @torch.compile(fullgraph=True) + def fn(x, cache_x): + return fused_causal_conv3d_cat_pad_cuda(x, cache_x, padding) + + assert torch.equal( + fn(x, cache_x), fused_causal_conv3d_cat_pad_triton(x, cache_x, padding) + ) + + +# --------------------------------------------------------------------------- +# Wan causal VAE cache kernels +# --------------------------------------------------------------------------- + + +def _ref_cat_pad(x, cache, padding): + p = list(padding) + if cache is not None: + x = torch.cat([cache, x], dim=2) + p[4] -= cache.shape[2] + if any(p): + x = F.pad(x, p) + return x.contiguous(memory_format=torch.channels_last_3d) + + +@torch.no_grad() +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float32]) +@pytest.mark.parametrize( + "c,t,h,w,cache_t,pads", + [ + (96, 1, 10, 14, 0, (1, 1, 1, 1, 2, 0)), # first chunk, zero-fill front + (96, 1, 10, 14, 1, (1, 1, 1, 1, 2, 0)), # legacy 1-frame cache + (96, 1, 10, 14, 2, (1, 1, 1, 1, 2, 0)), # steady state k3 conv + (64, 1, 10, 14, 2, (0, 0, 0, 0, 2, 0)), # time_conv (temporal only) + (48, 4, 10, 14, 2, (1, 1, 1, 1, 2, 0)), # encoder-style T=4 chunk + ], +) +def test_cat_pad_channels_last_3d_bitwise(dtype, c, t, h, w, cache_t, pads): + torch.cuda.manual_seed(0) + x = _cl3d((1, c, t, h, w), dtype) + cache = None + if cache_t: + # Strided interior view: caches may arrive as non-contiguous slices. + ph, pw = pads[2], pads[0] + buf = _cl3d((1, c, cache_t, h + 2 * ph, w + 2 * pw), dtype) + cache = buf[:, :, :, ph : ph + h, pw : pw + w] + ref = _ref_cat_pad(x, cache, pads) + + out = cat_pad_channels_last_3d(x, cache, pads) + assert out is not None and out.shape == ref.shape + assert out.is_contiguous(memory_format=torch.channels_last_3d) + assert torch.equal(out, ref) + + # Dual-output mode: the same pass also emits the compact feature cache + # (unpadded interior of the last frames), bitwise equal to the slice. + pair = cat_pad_channels_last_3d(x, cache, pads, keep_cache_t=2) + assert pair is not None + out2, keep = pair + assert torch.equal(out2, ref) + ph, pw = pads[2], pads[0] + keep_t = min(2, ref.shape[2]) + want = ref[:, :, ref.shape[2] - keep_t :, ph : ph + h, pw : pw + w] + assert keep.shape == want.shape + assert keep.is_contiguous(memory_format=torch.channels_last_3d) + assert torch.equal(keep, want) + + +@torch.no_grad() +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float32]) +@pytest.mark.parametrize( + "c_in,c_out,t,h,w,ft,fs,drop", + [ + (128, 64, 1, 10, 14, 2, 2, False), + (128, 64, 1, 10, 14, 2, 2, True), # first_chunk slicing + (64, 32, 2, 10, 14, 1, 2, False), + ], +) +def test_dup_up3d_add_bitwise(dtype, c_in, c_out, t, h, w, ft, fs, drop): + torch.cuda.manual_seed(0) + repeats = c_out * ft * fs * fs // c_in + src = _cl3d((1, c_in, t, h, w), dtype) + t_out = t * ft - (ft - 1 if drop else 0) + # Main arm as a permuted view, like the WanResample 2D output. + main = torch.randn( + (1, t_out, c_out, h * fs, w * fs), device=DEVICE, dtype=dtype + ).permute(0, 2, 1, 3, 4) + + dup = src.repeat_interleave(repeats, dim=1) + dup = dup.view(1, c_out, ft, fs, fs, t, h, w) + dup = dup.permute(0, 1, 5, 2, 6, 3, 7, 4).contiguous() + dup = dup.view(1, c_out, t * ft, h * fs, w * fs) + if drop: + dup = dup[:, :, ft - 1 :, :, :] + ref = main + dup + + out = dup_up3d_add(main, src, ft, fs, repeats, drop) + assert out is not None and out.shape == ref.shape + # Layout must match the aten add output exactly (downstream reductions + # are layout-sensitive), and every value must be bitwise identical. + assert out.stride() == ref.stride() + assert torch.equal(out, ref) + + +@torch.no_grad() +@pytest.mark.parametrize("pads_temporal_only", [False, True]) +def test_wan_cached_conv_chunk_loop_bitwise(pads_temporal_only): + """The fused conv-input/compact-cache scheme must reproduce the original + clone/cat bookkeeping bitwise across a chunked decode, including the + first-chunk zero fill and the "Rep" marker start used by WanResample.""" + from sglang.multimodal_gen.runtime.models.vaes import wanvae + from sglang.multimodal_gen.runtime.models.vaes.wanvae import ( + CACHE_T, + WanCausalConv3d, + _cache_payload, + _run_cached_causal_conv, + ) + + torch.cuda.manual_seed(0) + c = 64 + if pads_temporal_only: + conv = WanCausalConv3d(c, 2 * c, (3, 1, 1), padding=(1, 0, 0)) + else: + conv = WanCausalConv3d(c, c, 3, padding=1) + conv = conv.to(device=DEVICE, dtype=torch.float32) + conv.weight.data = conv.weight.data.contiguous(memory_format=torch.channels_last_3d) + chunks = [_cl3d((1, c, 1, 10, 14), torch.float32) for _ in range(4)] + + def run(force_fallback, start): + cache = [start] + outs = [] + orig = wanvae.cat_pad_channels_last_3d + if force_fallback: + wanvae.cat_pad_channels_last_3d = None + try: + for x in chunks: + outs.append(_run_cached_causal_conv(conv, x, cache, 0)) + finally: + wanvae.cat_pad_channels_last_3d = orig + return outs, cache[0] + + for start in (None, "Rep"): + fused_outs, fused_cache = run(False, start) + ref_outs, ref_cache = run(True, start) + for got, want in zip(fused_outs, ref_outs, strict=True): + assert torch.equal(got, want) + got_payload = _cache_payload(fused_cache) + assert got_payload is not None and got_payload.shape[2] == CACHE_T + # Reference cache holds the last CACHE_T unpadded frames. + assert torch.equal(got_payload, ref_cache[:, :, -CACHE_T:]) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernels/ops/diffusion/test_ltx2_ada_values.py b/test/registered/kernels/ops/diffusion/test_ltx2_ada_values.py deleted file mode 100644 index 3f780b2dfb29..000000000000 --- a/test/registered/kernels/ops/diffusion/test_ltx2_ada_values.py +++ /dev/null @@ -1,92 +0,0 @@ -import sys - -import pytest -import torch - -from sglang.kernels.ops.diffusion.triton.ltx2_ada_values import ltx2_ada_values9 -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci - -register_cuda_ci(est_time=8, stage="base-b-kernel-unit", runner_config="1-gpu-large") -register_amd_ci(est_time=8, suite="nightly-amd-kernel-1-gpu", nightly=True) - -DEVICE = "cuda" - - -@pytest.fixture(autouse=True) -def cuda_setup(): - if not torch.cuda.is_available(): - pytest.skip("CUDA required") - torch.cuda.manual_seed(0) - - -def _reference( - scale_shift_table: torch.Tensor, - timestep: torch.Tensor, -) -> tuple[torch.Tensor, ...]: - batch, seq, _ = timestep.shape - hidden = scale_shift_table.shape[1] - return ( - scale_shift_table.to(device=timestep.device, dtype=timestep.dtype) - .view(1, 1, 9, hidden) - .add(timestep.reshape(batch, seq, 9, hidden)) - .unbind(dim=2) - ) - - -@torch.no_grad() -@pytest.mark.parametrize("batch,seq,hidden", [(1, 1, 4096), (2, 3, 2048)]) -@pytest.mark.parametrize("table_dtype", [torch.bfloat16, torch.float32]) -def test_ltx2_ada_values9( - batch: int, - seq: int, - hidden: int, - table_dtype: torch.dtype, -) -> None: - scale_shift_table = torch.randn( - 9, hidden, device=DEVICE, dtype=table_dtype - ).contiguous() - timestep = torch.randn( - batch, seq, 9 * hidden, device=DEVICE, dtype=torch.bfloat16 - ).contiguous() - - actual = ltx2_ada_values9(scale_shift_table, timestep) - expected = _reference(scale_shift_table, timestep) - - assert len(actual) == 9 - for actual_value, expected_value in zip(actual, expected): - assert actual_value.is_contiguous() - torch.testing.assert_close(actual_value, expected_value, atol=0, rtol=0) - - -@torch.no_grad() -def test_ltx2_ada_values9_torch_compile_fullgraph() -> None: - hidden = 4096 - scale_shift_table = torch.randn( - 9, hidden, device=DEVICE, dtype=torch.bfloat16 - ).contiguous() - timestep = torch.randn( - 1, 1, 9 * hidden, device=DEVICE, dtype=torch.bfloat16 - ).contiguous() - - actual = torch.compile(ltx2_ada_values9, fullgraph=True)( - scale_shift_table, timestep - ) - expected = _reference(scale_shift_table, timestep) - - assert len(actual) == 9 - for actual_value, expected_value in zip(actual, expected): - assert actual_value.is_contiguous() - torch.testing.assert_close(actual_value, expected_value, atol=0, rtol=0) - - -@torch.no_grad() -def test_ltx2_ada_values9_rejects_unsupported_shape() -> None: - scale_shift_table = torch.randn(8, 4096, device=DEVICE, dtype=torch.bfloat16) - timestep = torch.randn(1, 1, 9 * 4096, device=DEVICE, dtype=torch.bfloat16) - - with pytest.raises(ValueError, match="scale_shift_table"): - ltx2_ada_values9(scale_shift_table, timestep) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/kernels/ops/diffusion/test_ltx2_qknorm_split_rope.py b/test/registered/kernels/ops/diffusion/test_ltx2_qknorm_split_rope.py deleted file mode 100644 index 5aa73636f1a8..000000000000 --- a/test/registered/kernels/ops/diffusion/test_ltx2_qknorm_split_rope.py +++ /dev/null @@ -1,221 +0,0 @@ -import sys - -import pytest -import torch -import torch.nn.functional as F - -from sglang.kernels.ops.diffusion.ltx2_qknorm_split_rope import ( - can_use_ltx2_qknorm_split_rope_cuda, - ltx2_qknorm_split_rope_cuda, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=45, stage="base-b-kernel-unit", runner_config="4-gpu-b200") - -BF16_FUSED_ATOL = 1.6e-1 - - -def _require_cuda_b200() -> None: - if not torch.cuda.is_available(): - pytest.skip("CUDA required") - if torch.cuda.get_device_capability()[0] < 10: - pytest.skip("LTX2 QKNorm split-RoPE CUDA path is validated on B200") - - -@pytest.fixture(autouse=True) -def cuda_setup(): - _require_cuda_b200() - torch.cuda.manual_seed(20260630) - - -def _make_cos_sin( - batch: int, seq_len: int, num_heads: int, head_dim: int -) -> tuple[torch.Tensor, torch.Tensor]: - half_dim = head_dim // 2 - cos = torch.randn( - batch, seq_len, num_heads, half_dim, device="cuda", dtype=torch.bfloat16 - ).transpose(1, 2) - sin = torch.randn( - batch, seq_len, num_heads, half_dim, device="cuda", dtype=torch.bfloat16 - ).transpose(1, 2) - return cos, sin - - -def _apply_split_rotary_ref( - x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor -) -> torch.Tensor: - x_dtype = x.dtype - batch = x.shape[0] - _, num_heads, seq_len, _ = cos.shape - x = x.reshape(batch, seq_len, num_heads, -1).swapaxes(1, 2) - last = x.shape[-1] - half = last // 2 - - split_x = x.reshape(*x.shape[:-1], 2, half) - first_x = split_x[..., :1, :] - second_x = split_x[..., 1:, :] - cos_u = cos.unsqueeze(-2) - sin_u = sin.unsqueeze(-2) - - out = split_x * cos_u - out[..., :1, :].addcmul_(-sin_u, second_x) - out[..., 1:, :].addcmul_(sin_u, first_x) - out = out.reshape(*out.shape[:-2], last) - return out.swapaxes(1, 2).reshape(batch, seq_len, -1).to(dtype=x_dtype) - - -def _reference( - q: torch.Tensor, - k: torch.Tensor, - q_cos: torch.Tensor, - q_sin: torch.Tensor, - k_cos: torch.Tensor, - k_sin: torch.Tensor, - q_weight: torch.Tensor, - k_weight: torch.Tensor, - eps: float, -) -> tuple[torch.Tensor, torch.Tensor]: - # rms_norm isn't autocast fp32-preserving, so feed fp32 inputs directly - # to keep the normalized value unrounded until the final RoPE output. - q_norm = F.rms_norm(q.float(), (q.shape[-1],), q_weight.float(), eps) - k_norm = F.rms_norm(k.float(), (k.shape[-1],), k_weight.float(), eps) - q_ref = _apply_split_rotary_ref(q_norm, q_cos, q_sin) - k_ref = _apply_split_rotary_ref(k_norm, k_cos, k_sin) - return q_ref.to(dtype=torch.bfloat16), k_ref.to(dtype=torch.bfloat16) - - -@pytest.mark.parametrize( - "batch,q_seq,k_seq,num_heads,head_dim", - [ - (1, 3, 3, 32, 128), - (1, 5, 2, 32, 64), - (2, 4, 3, 32, 64), - ], -) -def test_ltx2_qknorm_split_rope_matches_torch_exactly( - batch: int, q_seq: int, k_seq: int, num_heads: int, head_dim: int -) -> None: - hidden = num_heads * head_dim - eps = 1e-6 - q = torch.randn(batch, q_seq, hidden, device="cuda", dtype=torch.bfloat16) - k = torch.randn(batch, k_seq, hidden, device="cuda", dtype=torch.bfloat16) - q_cos, q_sin = _make_cos_sin(batch, q_seq, num_heads, head_dim) - k_cos, k_sin = _make_cos_sin(batch, k_seq, num_heads, head_dim) - q_weight = torch.randn(hidden, device="cuda", dtype=torch.bfloat16) - k_weight = torch.randn(hidden, device="cuda", dtype=torch.bfloat16) - - assert can_use_ltx2_qknorm_split_rope_cuda( - q, - q_cos, - q_sin, - q_weight, - k, - k_cos, - k_sin, - k_weight, - num_heads=num_heads, - head_dim=head_dim, - ) - - q_ref, k_ref = _reference(q, k, q_cos, q_sin, k_cos, k_sin, q_weight, k_weight, eps) - q_out, k_out = ltx2_qknorm_split_rope_cuda( - q, - q_cos, - q_sin, - q_weight, - k, - k_cos, - k_sin, - k_weight, - eps=eps, - num_heads=num_heads, - head_dim=head_dim, - ) - torch.cuda.synchronize() - - torch.testing.assert_close(q_out, q_ref, rtol=0, atol=BF16_FUSED_ATOL) - torch.testing.assert_close(k_out, k_ref, rtol=0, atol=BF16_FUSED_ATOL) - - -def test_ltx2_qknorm_split_rope_rejects_unsupported_inputs() -> None: - q = torch.randn((1, 3, 4096), device="cuda", dtype=torch.bfloat16) - k = torch.randn_like(q) - q_cos, q_sin = _make_cos_sin(1, 3, 32, 128) - q_weight = torch.randn(4096, device="cuda", dtype=torch.bfloat16) - k_weight = torch.randn(4096, device="cuda", dtype=torch.bfloat16) - - assert can_use_ltx2_qknorm_split_rope_cuda( - q, - q_cos, - q_sin, - q_weight, - k, - q_cos, - q_sin, - k_weight, - num_heads=32, - head_dim=128, - ) - assert not can_use_ltx2_qknorm_split_rope_cuda( - q.float(), - q_cos, - q_sin, - q_weight, - k, - q_cos, - q_sin, - k_weight, - num_heads=32, - head_dim=128, - ) - assert not can_use_ltx2_qknorm_split_rope_cuda( - q, - q_cos, - q_sin, - q_weight, - k, - q_cos.transpose(-1, -2), - q_sin, - k_weight, - num_heads=32, - head_dim=128, - ) - - -def test_ltx2_qknorm_split_rope_custom_op_torch_compile_fullgraph() -> None: - batch, q_seq, k_seq, num_heads, head_dim = 1, 3, 2, 32, 64 - hidden = num_heads * head_dim - q = torch.randn(batch, q_seq, hidden, device="cuda", dtype=torch.bfloat16) - k = torch.randn(batch, k_seq, hidden, device="cuda", dtype=torch.bfloat16) - q_cos, q_sin = _make_cos_sin(batch, q_seq, num_heads, head_dim) - k_cos, k_sin = _make_cos_sin(batch, k_seq, num_heads, head_dim) - q_weight = torch.randn(hidden, device="cuda", dtype=torch.bfloat16) - k_weight = torch.randn(hidden, device="cuda", dtype=torch.bfloat16) - - def fn(q, k, q_cos, q_sin, k_cos, k_sin, q_weight, k_weight): - return ltx2_qknorm_split_rope_cuda( - q, - q_cos, - q_sin, - q_weight, - k, - k_cos, - k_sin, - k_weight, - eps=1e-6, - num_heads=num_heads, - head_dim=head_dim, - ) - - compiled = torch.compile(fn, fullgraph=True) - q_out, k_out = compiled(q, k, q_cos, q_sin, k_cos, k_sin, q_weight, k_weight) - q_ref, k_ref = _reference( - q, k, q_cos, q_sin, k_cos, k_sin, q_weight, k_weight, 1e-6 - ) - torch.cuda.synchronize() - torch.testing.assert_close(q_out, q_ref, rtol=0, atol=BF16_FUSED_ATOL) - torch.testing.assert_close(k_out, k_ref, rtol=0, atol=BF16_FUSED_ATOL) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/kernels/ops/diffusion/test_ltx2_rms_norm_modulate.py b/test/registered/kernels/ops/diffusion/test_ltx2_rms_norm_modulate.py deleted file mode 100644 index 37b6ea37bb28..000000000000 --- a/test/registered/kernels/ops/diffusion/test_ltx2_rms_norm_modulate.py +++ /dev/null @@ -1,88 +0,0 @@ -"""LTX-2 quality=high RMSNorm+modulate fusion: gated, close to eager.""" - -import sys - -import pytest -import torch -from torch import nn - -import sglang.multimodal_gen.runtime.models.dits.ltx_2 as ltx2_module -from sglang.kernels.ops.diffusion.ltx2_rmsnorm_modulate import ( - fused_ltx2_rms_norm_modulate, - mark_ltx2_rms_norm_modulate_site, - mount_ltx2_rms_norm_modulate, - unmount_ltx2_rms_norm_modulate, -) -from sglang.multimodal_gen.runtime.layers.layernorm import RMSNormNoWeight -from sglang.multimodal_gen.runtime.models.dits.ltx_2 import _ltx2_rms_norm_modulate -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci - -register_cuda_ci(est_time=8, stage="base-b-kernel-unit", runner_config="1-gpu-large") -register_amd_ci(est_time=8, suite="nightly-amd-kernel-1-gpu", nightly=True) - - -@pytest.fixture(autouse=True) -def _setup(): - if not torch.cuda.is_available(): - pytest.skip("CUDA required") - torch.cuda.manual_seed(0) - - -def _eager(rms, x, scale, shift, eps): - return rms(x, eps) * (1 + scale) + shift - - -def _inputs(hidden, batch=1, seq=4096): - rms = RMSNormNoWeight() - x = torch.randn(batch, seq, hidden, device="cuda", dtype=torch.bfloat16) - scale = torch.randn(batch, 1, hidden, device="cuda", dtype=torch.bfloat16) * 0.1 - shift = torch.randn(batch, 1, hidden, device="cuda", dtype=torch.bfloat16) * 0.1 - return rms, x, scale, shift - - -# hidden 4096 = LTX-2 video stream, 2048 = audio stream. -@pytest.mark.parametrize("hidden", [4096, 2048]) -def test_lossless_default_is_bitexact(hidden): - # A marked-but-unmounted site uses only the self-verified bit-exact - # modulate fast path after the reference aten RMSNorm. - block = nn.Module() - mark_ltx2_rms_norm_modulate_site(block) - rms, x, scale, shift = _inputs(hidden) - out = _ltx2_rms_norm_modulate(block, rms, x, scale, shift, 1e-6) - assert torch.equal(out, _eager(rms, x, scale, shift, 1e-6)) - - -def test_lossless_compile_keeps_expression_visible_to_inductor(monkeypatch): - block = nn.Module() - mark_ltx2_rms_norm_modulate_site(block) - rms, x, scale, shift = _inputs(2048, seq=126) - monkeypatch.setattr(torch.compiler, "is_compiling", lambda: True) - monkeypatch.setattr( - ltx2_module, - "_ltx2_modulate", - lambda *_args: pytest.fail("compiled path must not call the opaque custom op"), - ) - - out = _ltx2_rms_norm_modulate(block, rms, x, scale, shift, 1e-6) - assert torch.equal(out, _eager(rms, x, scale, shift, 1e-6)) - - -@pytest.mark.parametrize("hidden", [4096, 2048]) -def test_mounted_high_uses_fused_kernel(hidden): - block = nn.Module() - mark_ltx2_rms_norm_modulate_site(block) - assert mount_ltx2_rms_norm_modulate(block) - try: - rms, x, scale, shift = _inputs(hidden) - out = _ltx2_rms_norm_modulate(block, rms, x, scale, shift, 1e-6) - # The mounted path routes through the fused kernel exactly. - assert torch.equal(out, fused_ltx2_rms_norm_modulate(x, scale, shift, 1e-6)) - # And stays within half-precision rounding of the eager reference. - ref = _eager(rms, x, scale, shift, 1e-6) - assert torch.allclose(out.float(), ref.float(), atol=3e-2, rtol=1e-2) - finally: - unmount_ltx2_rms_norm_modulate(block) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/kernels/ops/diffusion/test_model_fast_paths.py b/test/registered/kernels/ops/diffusion/test_model_fast_paths.py new file mode 100644 index 000000000000..c78de2ca99b2 --- /dev/null +++ b/test/registered/kernels/ops/diffusion/test_model_fast_paths.py @@ -0,0 +1,703 @@ +"""Per-model fast paths: each model wrapper must reproduce its own reference. + +A diffusion kernel is only as good as the wrapper that decides when to use it, +and that decision is model-specific: FLUX.1 and GLM-Image feed different adaLN +layouts to the same LayerNorm+modulate kernel, Sana only engages on non-default +streams, ERNIE runs a bit-exact chain unconditionally. Kernel-level numerics +live in ``test_norm.py`` / ``test_modulate.py`` / ``test_rope.py``; this file +covers the wiring: right kernel, right reference, gate ends in the right state. + +Two assertion styles appear, and the difference is the point: + +- ``torch.equal`` for the bit-exact default-on paths. These self-verify at + runtime, so a tolerance here would let a real regression through -- the gate + would silently fall back to eager and the fusion would simply stop running. +- a tolerance for the quality-gated paths, which are *documented* as differing + from eager at half-precision rounding-order level. + +Each section also asserts the gate ended up ``verified`` / not ``disabled``: +without it a test still passes when the fast path never engaged at all. +""" + +import sys +import unittest +from unittest.mock import patch + +import pytest +import torch +import torch.nn as nn +import torch.nn.functional as F +from diffusers.models.upsampling import Upsample2D + +import sglang.kernels.ops.diffusion.sites.hunyuan_qknorm_site as hunyuan_qknorm +import sglang.multimodal_gen.runtime.models.dits.ernie_image as ernie_image +import sglang.multimodal_gen.runtime.models.dits.flux as flux +import sglang.multimodal_gen.runtime.models.dits.flux_2 as flux2 +import sglang.multimodal_gen.runtime.models.dits.glm_image as glm_image +import sglang.multimodal_gen.runtime.models.dits.ltx_2 as ltx2_module +import sglang.multimodal_gen.runtime.models.dits.sana as sana +from sglang.kernels.ops.diffusion import ( + can_use_wan_rmsnorm_silu, + fused_ltx2_rms_norm_modulate, + mark_fused_ln_modulate_site, + mark_hunyuan_qknorm_site, + mark_ltx2_rms_norm_modulate_site, + mount_fused_ln_modulate, + mount_hunyuan_qknorm, + mount_ltx2_rms_norm_modulate, + unmount_hunyuan_qknorm, + unmount_ltx2_rms_norm_modulate, + wan_rmsnorm_silu, +) +from sglang.multimodal_gen.configs.models.vaes.stablediffusion3 import ( + StableDiffusion3VAEConfig, +) +from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm, RMSNormNoWeight +from sglang.multimodal_gen.runtime.layers.rotary_embedding.utils import ( + _apply_rotary_emb, +) +from sglang.multimodal_gen.runtime.models.dits.ernie_image import ( + _ernie_gated_norm_scale_shift, + _ernie_norm_scale_shift, + _ernie_qknorm_rope, + _ernie_qknorm_rope_reference, +) +from sglang.multimodal_gen.runtime.models.dits.flux import ( + _flux_fused_ln_modulate, + _flux_norm_modulate, +) +from sglang.multimodal_gen.runtime.models.dits.flux_2 import ( + _flux2_norm_modulate, + _flux2_swiglu, +) +from sglang.multimodal_gen.runtime.models.dits.glm_image import ( + _eager_ln_modulate as _glm_eager_ln_modulate, +) +from sglang.multimodal_gen.runtime.models.dits.glm_image import ( + _glm_ln_modulate, + _glm_qk_layernorm, +) +from sglang.multimodal_gen.runtime.models.dits.hunyuanvideo import ( + _hunyuan_pack_qkv, + _hunyuan_qknorm, +) +from sglang.multimodal_gen.runtime.models.dits.ltx_2 import _ltx2_rms_norm_modulate +from sglang.multimodal_gen.runtime.models.dits.sana import ( + _eager_ln_modulate as _sana_eager_ln_modulate, +) +from sglang.multimodal_gen.runtime.models.dits.sana import ( + _sana_ln_modulate, +) +from sglang.multimodal_gen.runtime.models.vaes import flux2_vae_cuda_opt as vae_opt +from sglang.multimodal_gen.runtime.models.vaes.autoencoder import AutoencoderKL +from sglang.multimodal_gen.runtime.models.vaes.fast_path_gate import use_vae_fast_path +from sglang.multimodal_gen.runtime.models.vaes.wan_vae_cuda_opt import ( + FusedWanRMSNormSiLU, + VaeFastPathGate, +) +from sglang.multimodal_gen.runtime.models.vaes.wanvae import WanRMS_norm +from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=95, stage="base-b-kernel-unit", runner_config="1-gpu-large") +register_amd_ci(est_time=8, suite="nightly-amd-kernel-1-gpu", nightly=True) + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") + + +@pytest.fixture(autouse=True) +def _seed_cuda(): + """Every wrapper below asserts against a reference computed from the same + random draw, so the seed must be fixed per test, not per module.""" + torch.cuda.manual_seed(0) + + +# ------------------------------------------------------------------------- +# FLUX.1 -- LayerNorm + adaLN modulate, and the shared-FF GELU site +# ------------------------------------------------------------------------- + + +def _flux_eager(norm, x, scale, shift): + return norm(x) * (1 + scale[:, None]) + shift[:, None] + + +def _flux_site_inputs(shape, chunks, seed): + torch.manual_seed(seed) + batch, seq, hidden = shape + norm = torch.nn.LayerNorm(hidden, eps=1e-6, elementwise_affine=False).cuda() + x = (torch.randn(batch, seq, hidden, device="cuda") * 8).bfloat16() + emb = torch.randn(batch, chunks * hidden, device="cuda").bfloat16() + parts = emb.chunk(chunks, dim=1) # strided adaLN projection views + return norm, x, parts[0], parts[1] + + +@pytest.mark.parametrize( + "shape,chunks", + [ + ((1, 4096, 3072), 6), # dual-stream image tokens (1024^2), chunk(6) + ((1, 512, 3072), 6), # dual-stream text tokens + ((1, 4608, 3072), 3), # single-stream concat, chunk(3) + ((2, 300, 3072), 6), # CFG batch, odd seq + ], +) +def test_flux_fused_ln_modulate_is_bit_exact(shape, chunks): + # Every distinct (shape, stride, eps) signature the FLUX.1 sites emit + # must verify torch.equal on first sight and stay enabled. + norm, x, shift, scale = _flux_site_inputs(shape, chunks, seed=0) + out = _flux_fused_ln_modulate(norm, x, scale, shift) + assert out is not None + assert torch.equal(out, _flux_eager(norm, x, scale, shift)) + assert not flux._FLUX_LN_MOD.disabled + assert flux._FLUX_LN_MOD.verified + + +def test_flux_norm_modulate_bitexact_supersedes_high_fold(): + # With the quality="high" affine fold mounted, the bit-exact kernel + # still takes priority, so the site output stays lossless. + site = torch.nn.Module() + mark_fused_ln_modulate_site(site) + assert mount_fused_ln_modulate(site) + norm, x, shift, scale = _flux_site_inputs((1, 128, 3072), 6, seed=1) + out = _flux_norm_modulate(site, norm, x, scale, shift) + assert torch.equal(out, _flux_eager(norm, x, scale, shift)) + + +def test_flux_fused_ln_modulate_rejects_unsupported_hidden(): + # hidden % 4 != 0 is outside the kernel contract and must bail out. + norm, x, shift, scale = _flux_site_inputs((1, 64, 3070), 6, seed=2) + assert _flux_fused_ln_modulate(norm, x, scale, shift) is None + + +# ------------------------------------------------------------------------- +# FLUX.2 -- packed norm+modulate and packed SwiGLU views +# ------------------------------------------------------------------------- + + +@unittest.skipUnless(torch.cuda.is_available(), "CUDA required") +class TestFlux2EagerFusions(CustomTestCase): + def setUp(self): + flux2._FLUX2_LN_MOD.disabled = False + flux2._FLUX2_LN_MOD.verified = False + flux2._FLUX2_LN_MOD_SIGS.clear() + flux2._FLUX2_SWIGLU.disabled = False + flux2._FLUX2_SWIGLU.verified = False + flux2._FLUX2_SWIGLU_SIGS.clear() + + def test_norm_modulate_is_bit_exact_across_sequence_lengths(self): + torch.manual_seed(0) + hidden = 256 + norm = torch.nn.LayerNorm( + hidden, eps=1e-6, elementwise_affine=False, device="cuda" + ) + # FLUX.2 modulation values are views of one packed projection. + params = torch.randn(1, 1, 6 * hidden, device="cuda").bfloat16() + shift, scale = params.chunk(6, dim=-1)[:2] + + for seq in (17, 65): + x = torch.randn(1, seq, hidden, device="cuda").bfloat16() + expected = norm(x) * (1 + scale) + shift + actual = _flux2_norm_modulate(norm, x, scale, shift) + self.assertTrue(torch.equal(actual, expected)) + + self.assertFalse(flux2._FLUX2_LN_MOD.disabled) + self.assertEqual(len(flux2._FLUX2_LN_MOD_SIGS), 1) + + def test_packed_swiglu_is_bit_exact_for_contiguous_and_strided_views(self): + torch.manual_seed(1) + hidden = 384 + inputs = [ + torch.randn(1, 19, 2 * hidden, device="cuda").bfloat16(), + torch.randn(1, 19, 3 * hidden, device="cuda").bfloat16()[..., : 2 * hidden], + ] + for x in inputs: + expected = F.silu(x[..., :hidden]) * x[..., hidden:] + actual = _flux2_swiglu(x) + self.assertTrue(torch.equal(actual, expected)) + + self.assertFalse(flux2._FLUX2_SWIGLU.disabled) + self.assertEqual(len(flux2._FLUX2_SWIGLU_SIGS), 2) + + def test_fp16_preserves_reference_path(self): + x = torch.randn(1, 17, 512, device="cuda", dtype=torch.float16) + expected = F.silu(x[..., :256]) * x[..., 256:] + actual = _flux2_swiglu(x) + self.assertTrue(torch.equal(actual, expected)) + self.assertFalse(flux2._FLUX2_SWIGLU.disabled) + + def test_packed_swiglu_rejects_non_dense_outer_stride(self): + base = torch.randn(2, 23, 512, device="cuda", dtype=torch.bfloat16) + x = base[:, :19] + self.assertNotEqual(x.stride(0), x.shape[1] * x.stride(1)) + + expected = F.silu(x[..., :256]) * x[..., 256:] + actual = _flux2_swiglu(x) + self.assertTrue(torch.equal(actual, expected)) + self.assertEqual(len(flux2._FLUX2_SWIGLU_SIGS), 0) + + def test_new_swiglu_signature_is_not_verified_during_graph_capture(self): + first = torch.randn(1, 17, 512, device="cuda", dtype=torch.bfloat16) + self.assertTrue( + torch.equal( + _flux2_swiglu(first), + F.silu(first[..., :256]) * first[..., 256:], + ) + ) + self.assertEqual(len(flux2._FLUX2_SWIGLU_SIGS), 1) + + second = torch.randn(1, 19, 768, device="cuda", dtype=torch.bfloat16) + with patch("torch.cuda.is_current_stream_capturing", return_value=True): + actual = _flux2_swiglu(second) + + expected = F.silu(second[..., :384]) * second[..., 384:] + self.assertTrue(torch.equal(actual, expected)) + self.assertEqual(len(flux2._FLUX2_SWIGLU_SIGS), 1) + + +# ------------------------------------------------------------------------- +# GLM-Image -- LayerNorm + modulate and per-head qk LayerNorm +# ------------------------------------------------------------------------- + + +@pytest.mark.parametrize("shape", [(1, 4096, 4096), (2, 301, 4096), (1, 1, 2560)]) +def test_glm_ln_modulate_is_bit_exact(shape): + # (1, 4096, 4096) is the real GLM-Image image-stream shape (1024^2, + # hidden 4096); the others cover the text stream and another hidden. + torch.manual_seed(0) + batch, seq, hidden = shape + norm = torch.nn.LayerNorm(hidden, eps=1e-5, elementwise_affine=False).cuda() + x = (torch.randn(batch, seq, hidden, device="cuda") * 8).bfloat16() + emb = torch.randn(batch, 12 * hidden, device="cuda").bfloat16() + chunks = emb.chunk(12, dim=1) # strided adaLN projection views + shift, scale = chunks[0], chunks[2] + out = _glm_ln_modulate(norm, x, scale, shift, x.dtype) + assert torch.equal(out, _glm_eager_ln_modulate(norm, x, scale, shift, x.dtype)) + assert glm_image._GLM_LN_MOD.verified + assert not glm_image._GLM_LN_MOD.disabled + + +@pytest.mark.parametrize("shape", [(1, 4360, 32, 128), (2, 37, 3, 40), (1, 129, 5, 64)]) +def test_glm_qk_head_layernorm_is_bit_exact(shape): + # (1, 4360, 32, 128) is the real GLM-Image q/k shape (text + image + # tokens, 32 heads of dim 128); the others cover partially-filled warps. + torch.manual_seed(1) + batch, seq, heads, head_dim = shape + norm_q = torch.nn.LayerNorm(head_dim, eps=1e-5, elementwise_affine=False).cuda() + norm_k = torch.nn.LayerNorm(head_dim, eps=1e-5, elementwise_affine=False).cuda() + q = (torch.randn(batch, seq, heads, head_dim, device="cuda") * 5).bfloat16() + k = (torch.randn(batch, seq, heads, head_dim, device="cuda") * 5).bfloat16() + q_out, k_out = _glm_qk_layernorm(norm_q, norm_k, q, k, q.dtype) + assert torch.equal(q_out, norm_q(q).to(q.dtype)) + assert torch.equal(k_out, norm_k(k).to(k.dtype)) + assert glm_image._GLM_QK_LN.verified + assert not glm_image._GLM_QK_LN.disabled + + +# ------------------------------------------------------------------------- +# Sana -- stream-conditional LayerNorm + modulate +# ------------------------------------------------------------------------- + + +@pytest.mark.parametrize( + "shape,nmod,transposed", + [ + ((2, 1024, 2240), 6, False), + ((2, 1024, 2240), 2, False), + ((1, 1024, 2240), 6, True), + ((1, 37, 2240), 6, False), + ], +) +def test_sana_fused_ln_modulate_is_bit_exact(shape, nmod, transposed): + # (., 1024, 2240) is the real Sana 1024px shape; hidden 2240 % 512 != 0 + # exercises the kernel's partial tail chunk. nmod mirrors the two adaLN + # chunk layouts, transposed the permuted layout the Sana DiT serves. + torch.manual_seed(0) + batch, seq, hidden = shape + norm = torch.nn.LayerNorm(hidden, eps=1e-6, elementwise_affine=False).cuda() + x = (torch.randn(batch, seq, hidden, device="cuda") * 4).bfloat16() + if transposed: + x = x.permute(0, 2, 1).contiguous().permute(0, 2, 1) + emb = torch.randn(batch, nmod, hidden, device="cuda").bfloat16() + shift, scale = emb.chunk(nmod, dim=1)[0], emb.chunk(nmod, dim=1)[-1] + # default-stream eager serving must stay on the untouched eager chain + n_sigs = len(sana._SANA_LN_MOD.verified_sigs) + _sana_ln_modulate(norm, x, scale, shift) + assert len(sana._SANA_LN_MOD.verified_sigs) == n_sigs + # the fusion engages on non-default streams (the BCG warmup/capture path) + with torch.cuda.stream(torch.cuda.Stream()): + out = _sana_ln_modulate(norm, x, scale, shift) + assert len(sana._SANA_LN_MOD.verified_sigs) == n_sigs + 1 # verified + out2 = _sana_ln_modulate(norm, x, scale, shift) # verified-sig lane + torch.cuda.synchronize() + assert torch.equal(out, _sana_eager_ln_modulate(norm, x, scale, shift)) + assert torch.equal(out2, out) and not sana._SANA_LN_MOD.disabled + + +# ------------------------------------------------------------------------- +# ERNIE-Image -- bit-exact RMSNorm scale/shift and rotate-half RoPE +# ------------------------------------------------------------------------- + + +@pytest.mark.parametrize("shape", [(1, 4216, 4096), (2, 1140, 4096), (1, 128, 2048)]) +def test_ernie_norm_scale_shift_is_bit_exact(shape): + # (1, 4216, 4096) is the real ERNIE-Image shape (1024^2 image + text + # tokens, hidden 4096); 2048 covers the threads_per_row=32 regime. + torch.manual_seed(0) + batch, seq, hidden = shape + norm = RMSNorm(hidden, eps=1e-6).to(device="cuda", dtype=torch.bfloat16) + with torch.no_grad(): + norm.weight.copy_(torch.randn(hidden)) + x = torch.randn(batch, seq, hidden, device="cuda", dtype=torch.bfloat16) + residual = torch.randn_like(x) + update = torch.randn_like(x) + scale = torch.randn(batch, 1, hidden, device="cuda", dtype=torch.bfloat16) * 0.1 + shift = torch.randn(batch, 1, hidden, device="cuda", dtype=torch.bfloat16) * 0.1 + gate = torch.randn(batch, 1, hidden, device="cuda", dtype=torch.bfloat16) + + with torch.no_grad(): + out = _ernie_norm_scale_shift(norm, x, scale, shift) + ref = norm(x) * (1 + scale) + shift + assert torch.equal(out, ref) + + out2, res = _ernie_gated_norm_scale_shift( + norm, residual, update, gate, scale, shift + ) + res_ref = residual + gate * update + ref2 = norm(res_ref) * (1 + scale) + shift + assert torch.equal(res, res_ref) + assert torch.equal(out2, ref2) + + # the fast paths must actually be in use (not silently disabled) + assert ernie_image._ERNIE_NORM.verified + assert ernie_image._ERNIE_GATED_NORM.verified + assert not ernie_image._ERNIE_NORM.disabled + assert not ernie_image._ERNIE_GATED_NORM.disabled + + +def test_ernie_qknorm_rope_is_bit_exact(): + torch.manual_seed(1) + ernie_image._ERNIE_QKNORM_ROPE.disabled = False + ernie_image._ERNIE_QKNORM_ROPE.verified = False + batch, seq, heads, head_dim = 1, 257, 32, 128 + q = torch.randn(batch, seq, heads, head_dim, device="cuda", dtype=torch.bfloat16) + k = torch.randn_like(q) + q_norm = RMSNorm(head_dim, eps=1e-6).to(device="cuda", dtype=torch.bfloat16) + k_norm = RMSNorm(head_dim, eps=1e-6).to(device="cuda", dtype=torch.bfloat16) + cos = torch.randn(seq, head_dim, device="cuda", dtype=torch.bfloat16) + sin = torch.randn_like(cos) + cache = torch.cat((cos, sin), dim=-1).contiguous() + positions = torch.arange(seq, device="cuda", dtype=torch.long) + + q_ref, k_ref = _ernie_qknorm_rope_reference( + q.clone(), k.clone(), q_norm, k_norm, head_dim, cos, sin + ) + q_out, k_out = _ernie_qknorm_rope( + q, + k, + q_norm, + k_norm, + head_dim, + cos, + sin, + cache, + positions, + ) + + assert torch.equal(q_out, q_ref) + assert torch.equal(k_out, k_ref) + assert ernie_image._ERNIE_QKNORM_ROPE.verified + assert not ernie_image._ERNIE_QKNORM_ROPE.disabled + + +def test_ernie_qknorm_rope_first_attempt_exception_uses_pristine_inputs(): + torch.manual_seed(2) + ernie_image._ERNIE_QKNORM_ROPE.disabled = False + ernie_image._ERNIE_QKNORM_ROPE.verified = False + batch, seq, heads, head_dim = 1, 17, 4, 128 + q = torch.randn(batch, seq, heads, head_dim, device="cuda", dtype=torch.bfloat16) + k = torch.randn_like(q) + q_norm = RMSNorm(head_dim, eps=1e-6).to(device="cuda", dtype=torch.bfloat16) + k_norm = RMSNorm(head_dim, eps=1e-6).to(device="cuda", dtype=torch.bfloat16) + cos = torch.randn(seq, head_dim, device="cuda", dtype=torch.bfloat16) + sin = torch.randn_like(cos) + cache = torch.cat((cos, sin), dim=-1).contiguous() + positions = torch.arange(seq, device="cuda", dtype=torch.long) + q_ref, k_ref = _ernie_qknorm_rope_reference( + q.clone(), k.clone(), q_norm, k_norm, head_dim, cos, sin + ) + + def mutate_then_raise(**kwargs): + kwargs["q"].zero_() + kwargs["k"].zero_() + raise RuntimeError("synthetic kernel failure") + + with patch.object(ernie_image, "apply_qk_norm_rope", mutate_then_raise): + q_out, k_out = _ernie_qknorm_rope( + q, + k, + q_norm, + k_norm, + head_dim, + cos, + sin, + cache, + positions, + ) + + assert torch.equal(q_out, q_ref) + assert torch.equal(k_out, k_ref) + assert ernie_image._ERNIE_QKNORM_ROPE.disabled + + +# ------------------------------------------------------------------------- +# LTX-2 -- weightless RMSNorm + modulate (quality-gated) +# ------------------------------------------------------------------------- + + +def _ltx2_eager(rms, x, scale, shift, eps): + return rms(x, eps) * (1 + scale) + shift + + +def _ltx2_inputs(hidden, batch=1, seq=4096): + rms = RMSNormNoWeight() + x = torch.randn(batch, seq, hidden, device="cuda", dtype=torch.bfloat16) + scale = torch.randn(batch, 1, hidden, device="cuda", dtype=torch.bfloat16) * 0.1 + shift = torch.randn(batch, 1, hidden, device="cuda", dtype=torch.bfloat16) * 0.1 + return rms, x, scale, shift + + +# hidden 4096 = LTX-2 video stream, 2048 = audio stream. +@pytest.mark.parametrize("hidden", [4096, 2048]) +def test_ltx2_lossless_default_is_bitexact(hidden): + # A marked-but-unmounted site uses only the self-verified bit-exact + # modulate fast path after the reference aten RMSNorm. + block = nn.Module() + mark_ltx2_rms_norm_modulate_site(block) + rms, x, scale, shift = _ltx2_inputs(hidden) + out = _ltx2_rms_norm_modulate(block, rms, x, scale, shift, 1e-6) + assert torch.equal(out, _ltx2_eager(rms, x, scale, shift, 1e-6)) + + +def test_ltx2_lossless_compile_keeps_expression_visible_to_inductor(monkeypatch): + block = nn.Module() + mark_ltx2_rms_norm_modulate_site(block) + rms, x, scale, shift = _ltx2_inputs(2048, seq=126) + monkeypatch.setattr(torch.compiler, "is_compiling", lambda: True) + monkeypatch.setattr( + ltx2_module, + "_ltx2_modulate", + lambda *_args: pytest.fail("compiled path must not call the opaque custom op"), + ) + + out = _ltx2_rms_norm_modulate(block, rms, x, scale, shift, 1e-6) + assert torch.equal(out, _ltx2_eager(rms, x, scale, shift, 1e-6)) + + +@pytest.mark.parametrize("hidden", [4096, 2048]) +def test_ltx2_mounted_high_uses_fused_kernel(hidden): + block = nn.Module() + mark_ltx2_rms_norm_modulate_site(block) + assert mount_ltx2_rms_norm_modulate(block) + try: + rms, x, scale, shift = _ltx2_inputs(hidden) + out = _ltx2_rms_norm_modulate(block, rms, x, scale, shift, 1e-6) + # The mounted path routes through the fused kernel exactly. + assert torch.equal(out, fused_ltx2_rms_norm_modulate(x, scale, shift, 1e-6)) + # And stays within half-precision rounding of the eager reference. + ref = _ltx2_eager(rms, x, scale, shift, 1e-6) + assert torch.allclose(out.float(), ref.float(), atol=3e-2, rtol=1e-2) + finally: + unmount_ltx2_rms_norm_modulate(block) + + +# ------------------------------------------------------------------------- +# HunyuanVideo -- QKV/RoPE pack and quality-gated QK RMSNorm +# ------------------------------------------------------------------------- + + +@pytest.mark.parametrize("img_tokens,txt_tokens", [(257, 31), (4096, 256)]) +def test_hunyuan_qkv_rope_pack_is_bit_exact(img_tokens, txt_tokens): + torch.manual_seed(0) + shape_img = (1, img_tokens, 24, 128) + shape_txt = (1, txt_tokens, 24, 128) + img_q, img_k, img_v = ( + torch.randn(shape_img, device="cuda", dtype=torch.bfloat16) for _ in range(3) + ) + txt_q, txt_k, txt_v = ( + torch.randn(shape_txt, device="cuda", dtype=torch.bfloat16) for _ in range(3) + ) + cos = torch.randn(img_tokens, 64, device="cuda") + sin = torch.randn_like(cos) + + q, k, v = _hunyuan_pack_qkv(img_q, img_k, img_v, txt_q, txt_k, txt_v, cos, sin) + q_ref = torch.cat( + (_apply_rotary_emb(img_q, cos, sin, is_neox_style=False), txt_q), dim=1 + ) + k_ref = torch.cat( + (_apply_rotary_emb(img_k, cos, sin, is_neox_style=False), txt_k), dim=1 + ) + v_ref = torch.cat((img_v, txt_v), dim=1) + + assert torch.equal(q, q_ref) + assert torch.equal(k, k_ref) + assert torch.equal(v, v_ref) + + +def test_hunyuan_quality_qknorm_matches_rmsnorm(): + torch.manual_seed(1) + site = torch.nn.Module() + mark_hunyuan_qknorm_site(site) + q_norm = RMSNorm(128, eps=1e-6).to(device="cuda", dtype=torch.bfloat16) + k_norm = RMSNorm(128, eps=1e-6).to(device="cuda", dtype=torch.bfloat16) + packed = torch.randn(1, 257, 3, 24, 128, device="cuda", dtype=torch.bfloat16) + q, k = packed[:, :, 0], packed[:, :, 1] + q_ref = q_norm(q.contiguous()).to(q) + k_ref = k_norm(k.contiguous()).to(k) + + q_unmounted, k_unmounted = _hunyuan_qknorm(site, q, k, q_norm, k_norm) + assert torch.equal(q_unmounted, q_ref) + assert torch.equal(k_unmounted, k_ref) + + assert mount_hunyuan_qknorm(site) + q_out, k_out = _hunyuan_qknorm(site, q, k, q_norm, k_norm) + torch.testing.assert_close(q_out, q_ref, atol=2e-2, rtol=2e-2) + torch.testing.assert_close(k_out, k_ref, atol=2e-2, rtol=2e-2) + + unmount_hunyuan_qknorm(site) + q_unmounted, k_unmounted = _hunyuan_qknorm(site, q, k, q_norm, k_norm) + assert torch.equal(q_unmounted, q_ref) + assert torch.equal(k_unmounted, k_ref) + + +def test_hunyuan_quality_qknorm_stays_unmounted_without_cute_kernel(): + site = torch.nn.Module() + mark_hunyuan_qknorm_site(site) + + with patch.object(hunyuan_qknorm, "_get_qk_rmsnorm_cute", return_value=None): + assert not mount_hunyuan_qknorm(site) + + assert not hunyuan_qknorm._FUSION.is_enabled(site) + + +# ------------------------------------------------------------------------- +# Wan VAE -- fused RMSNorm+SiLU module gate +# ------------------------------------------------------------------------- + + +def _wan_cl3d(shape, dtype): + return torch.randn(shape, device="cuda", dtype=dtype).contiguous( + memory_format=torch.channels_last_3d + ) + + +@torch.no_grad() +def test_wan_vae_gate_dispatch() -> None: + # Gate off must stay bit-exact; gate on must route to the fused kernel. + torch.cuda.manual_seed(0) + norm = WanRMS_norm(96, images=False).to(device="cuda", dtype=torch.bfloat16) + norm.gamma.add_(torch.randn_like(norm.gamma)) + gate = VaeFastPathGate() + fused = FusedWanRMSNormSiLU(norm, gate) + # Parameter names must not change (weight transfer matches by name). + assert [n for n, _ in fused.named_parameters()] == ["gamma"] + x = _wan_cl3d((1, 96, 3, 10, 14), torch.bfloat16) + assert torch.equal(fused(x), nn.SiLU()(norm(x))) + gate.enabled = True + expected = wan_rmsnorm_silu(x, norm.gamma, rms_scale=float(norm.scale)) + assert torch.equal(fused(x), expected) + + +@torch.no_grad() +def test_wan_vae_rejects_empty_input() -> None: + x = torch.empty(1, 96, 0, 2, 2, device="cuda", dtype=torch.bfloat16).to( + memory_format=torch.channels_last_3d + ) + gamma = torch.ones(96, 1, 1, 1, device="cuda", dtype=torch.bfloat16) + assert not can_use_wan_rmsnorm_silu(x, gamma, None) + + +# ------------------------------------------------------------------------- +# FLUX.2 VAE -- fused GroupNorm+SiLU and folded 2x upsample conv +# ------------------------------------------------------------------------- + + +@torch.no_grad() +def test_flux2_vae_fast_path(): + torch.manual_seed(0) + gate = vae_opt.VaeFastPathGate() + gn = nn.GroupNorm(32, 128, eps=1e-6).to("cuda", torch.bfloat16) + x = torch.randn(1, 128, 64, 64, device="cuda", dtype=torch.bfloat16).to( + memory_format=torch.channels_last + ) + ref = F.silu(gn(x)) + fused_gn = vae_opt.FusedGroupNormSiLU(gn, gate) + assert set(fused_gn.state_dict()) == {"weight", "bias"} + assert torch.equal(fused_gn(x), ref) # gate off: bit-exact reference + + gate.enabled = True + fast = fused_gn(x) + assert fast.is_contiguous(memory_format=torch.channels_last) + torch.testing.assert_close(fast.float(), ref.float(), atol=0.06, rtol=0) + + gate.enabled = False + up = Upsample2D(channels=32, use_conv=True).to("cuda", torch.bfloat16) + fused_up = vae_opt.FusedUpsample2xConv2d(up, gate) + assert set(fused_up.state_dict()) == {"conv.weight", "conv.bias"} + x = torch.randn(2, 32, 33, 29, device="cuda", dtype=torch.bfloat16) + ref = up(x) + assert torch.equal(fused_up(x), ref) + assert fused_up._fused_weight is None + + gate.enabled = True + fast = fused_up(x) + assert fused_up._fused_weight is not None + ref_range = ref.float().max() - ref.float().min() + relative_mse = F.mse_loss(fast.float(), ref.float()) / ref_range.square() + assert relative_mse < 3.2e-5 + + +# --------------------------------------------------------------------------- +# AutoencoderKL (generic) -- fast-path install must not disturb the checkpoint +# --------------------------------------------------------------------------- + + +def _small_config(): + config = StableDiffusion3VAEConfig() + config.arch_config.latent_channels = 2 + config.arch_config.block_out_channels = (4, 4) + config.arch_config.down_block_types = ("DownEncoderBlock2D",) * 2 + config.arch_config.up_block_types = ("UpDecoderBlock2D",) * 2 + config.arch_config.layers_per_block = 1 + config.arch_config.norm_num_groups = 1 + config.arch_config.sample_size = 8 + return config + + +@torch.no_grad() +def test_autoencoder_kl_fastpath_install(): + torch.manual_seed(0) + vae = AutoencoderKL(_small_config()).to("cuda", torch.bfloat16).eval() + ref_names = {n for n, _ in vae.named_parameters()} + ref_sd = {k: v.clone() for k, v in vae.state_dict().items()} + z = torch.randn(1, 2, 8, 8, device="cuda", dtype=torch.bfloat16) + ref = vae.decode(z) + + opt = vae_opt.maybe_optimize_autoencoder_kl(vae) + # Wrappers must not change parameter FQNs; strict load must round-trip. + assert {n for n, _ in opt.named_parameters()} == ref_names + opt.load_state_dict(ref_sd, strict=True) + # Gate off: bit-for-bit the original path. + assert torch.equal(opt.decode(z), ref) + # use_vae_fast_path() is a no-op when nothing registered a gate, so check + # the wrappers went in before relying on it to switch paths. + assert any( + isinstance(m, (vae_opt.FusedGroupNormSiLU, vae_opt.FusedUpsample2xConv2d)) + for m in opt.modules() + ) + # Gate on: fast path runs and stays close; leaving the scope restores exact. + with use_vae_fast_path(opt, True): + torch.testing.assert_close(opt.decode(z).float(), ref.float(), atol=0.1, rtol=0) + assert torch.equal(opt.decode(z), ref) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernels/ops/diffusion/test_modulate.py b/test/registered/kernels/ops/diffusion/test_modulate.py new file mode 100644 index 000000000000..3ed13deebc13 --- /dev/null +++ b/test/registered/kernels/ops/diffusion/test_modulate.py @@ -0,0 +1,424 @@ +"""``diffusion.modulate``: adaLN modulation, gating and timestep conditioning. + +The bit-exact kernels here (``modulate_scale_shift``, ``residual_gate_add``, +``ltx2_ada_values9``, ``try_fused_scaled_residual_add_exact``) reproduce every +aten rounding boundary, so they are asserted with ``torch.equal``. The +select-0/1 LayerNorm fusions compute their statistics differently from the +reference chain and are asserted with a tolerance. +""" + +import sys + +import pytest +import torch + +from sglang.kernels.jit.utils import get_ci_test_range +from sglang.kernels.ops.diffusion import ( + can_use_modulate_scale_shift_cuda, + can_use_residual_gate_add_cuda, + fuse_layernorm_scale_shift_gate_select01_kernel, + fuse_residual_layernorm_scale_shift_gate_select01_kernel, + ltx2_ada_values9, + modulate_scale_shift, + modulate_scale_shift_cuda, + norm_infer, + residual_gate_add, + residual_gate_add_cuda, + timestep_embedding, + try_fused_scaled_residual_add_exact, +) +from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci + +register_cuda_ci(est_time=75, stage="base-b-kernel-unit", runner_config="1-gpu-large") +register_cuda_ci(est_time=60, stage="base-b-kernel-unit", runner_config="4-gpu-b200") +# Nightly is not redundant: it sets SGLANG_JIT_KERNEL_RUN_FULL_TESTS=1, which +# expands the get_ci_test_range sweeps below. +register_cuda_ci(est_time=50, stage="nightly", runner_config="1-gpu-large") +register_amd_ci(est_time=38, suite="nightly-amd-kernel-1-gpu", nightly=True) + +DEVICE = "cuda" + + +@pytest.fixture(autouse=True) +def cuda_setup(): + if not torch.cuda.is_available(): + pytest.skip("CUDA required") + torch.cuda.manual_seed(0) + + +# --------------------------------------------------------------------------- +# modulate: x * (1 + scale) + shift +# --------------------------------------------------------------------------- + +# FLUX.1 1024^2 adaLN shapes (D=3072) plus batched and odd-length coverage. +MODULATE_CASES = [ + (1, 4096, 3072), + (1, 512, 3072), + (1, 4608, 3072), + (2, 1024, 3072), + (1, 17, 64), +] + + +def _eager_modulate(x, scale, shift): + return x * (1 + scale[:, None]) + shift[:, None] + + +@pytest.mark.parametrize("shape", MODULATE_CASES) +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +def test_modulate_scale_shift_matches_eager(shape, dtype): + x = torch.randn(shape, device=DEVICE, dtype=dtype) + scale = torch.randn((shape[0], shape[-1]), device=DEVICE, dtype=dtype) + shift = torch.randn_like(scale) + assert torch.equal( + modulate_scale_shift_cuda(x, scale, shift), _eager_modulate(x, scale, shift) + ) + + +def test_modulate_scale_shift_accepts_adaln_chunk_views(): + # Production feeds strided ``emb.chunk(6)`` views, not fresh tensors. + x = torch.randn((1, 4096, 3072), device=DEVICE, dtype=torch.bfloat16) + emb = torch.randn((1, 6 * 3072), device=DEVICE, dtype=torch.bfloat16) + shift, scale = emb.chunk(6, dim=1)[:2] + assert can_use_modulate_scale_shift_cuda(x, scale, shift) + assert torch.equal( + modulate_scale_shift_cuda(x, scale, shift), _eager_modulate(x, scale, shift) + ) + + +def test_modulate_scale_shift_guards_reject_fp32(): + x = torch.randn((1, 64, 64), device=DEVICE, dtype=torch.float32) + row = torch.randn((1, 64), device=DEVICE, dtype=torch.float32) + assert not can_use_modulate_scale_shift_cuda(x, row, row) + # The public wrapper still returns the eager result on a rejected input. + assert torch.equal(modulate_scale_shift(x, row, row), _eager_modulate(x, row, row)) + + +# --------------------------------------------------------------------------- +# residual + gate * update +# --------------------------------------------------------------------------- + +GATE_CASES = [ + ((1, 1024, 4096), (1, 1, 4096)), + ((1, 512, 4096), (1, 512, 4096)), + ((1, 17, 65), (1, 1, 65)), + ((1, 17, 65), (1, 17, 65)), + # FLUX.1 / FLUX.2-klein 1024^2 shapes (D=3072): dual-stream image/text + # and single-stream/joint concat; gates are [1, 1, D] modulation rows. + ((1, 4096, 3072), (1, 1, 3072)), + ((1, 512, 3072), (1, 1, 3072)), + ((1, 4608, 3072), (1, 1, 3072)), + # FLUX.2-dev (D=6144) joint sequence. + ((1, 4608, 6144), (1, 1, 6144)), + # ERNIE-4.5-VL 1024^2 image tokens plus text tokens. + ((1, 4216, 4096), (1, 1, 4096)), +] + + +def _assert_gate_add(out, ref): + if ref.dtype == torch.float32: + # fp32 has no rounding boundary to reproduce; the kernel keeps the + # accumulation in fp32 and only order may differ. + torch.testing.assert_close(out, ref, atol=1e-5, rtol=1e-5) + else: + assert torch.equal(out, ref) + + +@pytest.mark.parametrize("residual_shape,gate_shape", GATE_CASES) +def test_residual_gate_add_matches_torch(residual_shape, gate_shape): + residual = torch.randn(residual_shape, device=DEVICE, dtype=torch.bfloat16) + update = torch.randn_like(residual) + gate = torch.randn(gate_shape, device=DEVICE, dtype=torch.bfloat16) + + ref = residual + update * gate + _assert_gate_add(residual_gate_add_cuda(residual, update, gate), ref) + assert torch.equal(residual_gate_add(residual, update, gate), ref) + + +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16, torch.float32]) +@pytest.mark.parametrize("gate_shape", [(1, 1, 64), (1, 9, 64)]) +def test_residual_gate_add_dtypes(dtype, gate_shape): + residual = torch.randn((1, 9, 64), device=DEVICE, dtype=dtype) + update = torch.randn_like(residual) + gate = torch.randn(gate_shape, device=DEVICE, dtype=dtype) + _assert_gate_add( + residual_gate_add_cuda(residual, update, gate), residual + update * gate + ) + + +def test_residual_gate_add_guards_and_eager_fallback(): + residual = torch.randn((1, 8, 64), device=DEVICE, dtype=torch.bfloat16) + update = torch.randn_like(residual) + gate = torch.randn((1, 1, 64), device=DEVICE, dtype=torch.bfloat16) + assert can_use_residual_gate_add_cuda(residual, update, gate) + + rejected = [ + (residual.cpu(), update, gate), # not on device + (residual, update.float(), gate), # mixed dtypes + (residual, update[:, ::2], gate), # strided update + (residual, update, gate[:, :, ::2]), # strided gate + (residual[:, :0], update[:, :0], gate), # empty token dim + ] + for args in rejected: + assert not can_use_residual_gate_add_cuda(*args) + + # Only [1, ..., 1, D] row-broadcast gates are supported; a batched + # [B>1, 1, D] gate is not row-broadcast here and must fall back. + batched = torch.randn((2, 8, 64), device=DEVICE, dtype=torch.bfloat16) + batched_update = torch.randn_like(batched) + batched_gate = torch.randn((2, 1, 64), device=DEVICE, dtype=torch.bfloat16) + assert not can_use_residual_gate_add_cuda(batched, batched_update, batched_gate) + assert torch.equal( + residual_gate_add(batched, batched_update, batched_gate), + batched + batched_update * batched_gate, + ) + + +def test_residual_gate_add_torch_compile_fullgraph(): + residual = torch.randn((1, 32, 128), device=DEVICE, dtype=torch.bfloat16) + update = torch.randn_like(residual) + gate = torch.randn((1, 1, 128), device=DEVICE, dtype=torch.bfloat16) + compiled = torch.compile(residual_gate_add, fullgraph=True) + assert torch.equal(compiled(residual, update, gate), residual + update * gate) + + +@torch.no_grad() +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) +def test_scaled_residual_add_is_bit_exact(dtype): + # fp32 residual accumulator + half-precision update, as the DiT blocks + # that keep their residual stream in fp32 emit it. + residual = torch.randn(2, 17, 64, device=DEVICE, dtype=torch.float32) + x = torch.randn(2, 17, 64, device=DEVICE, dtype=dtype) + scale = torch.randn(64, device=DEVICE, dtype=torch.float32) + + actual = try_fused_scaled_residual_add_exact(residual, x, scale) + assert actual is not None + assert torch.equal(actual, residual + x * scale) + + +@torch.no_grad() +def test_scaled_residual_add_rejects_unsupported_inputs(): + residual = torch.empty(2, 3, 8, device=DEVICE, dtype=torch.float32) + x = torch.empty_like(residual) + scale = torch.empty(8, device=DEVICE, dtype=torch.float32) + # A too-small hidden dim and a mismatched scale length both bail out; + # ``try_`` returning None is this helper's documented contract. + assert try_fused_scaled_residual_add_exact(residual, x, scale) is None + assert try_fused_scaled_residual_add_exact(residual, x.half(), scale[:-1]) is None + + +# --------------------------------------------------------------------------- +# LTX-2 nine-way adaLN value split +# --------------------------------------------------------------------------- + + +def _ltx2_reference(scale_shift_table, timestep): + batch, seq, _ = timestep.shape + hidden = scale_shift_table.shape[1] + return ( + scale_shift_table.to(device=timestep.device, dtype=timestep.dtype) + .view(1, 1, 9, hidden) + .add(timestep.reshape(batch, seq, 9, hidden)) + .unbind(dim=2) + ) + + +@torch.no_grad() +@pytest.mark.parametrize("batch,seq,hidden", [(1, 1, 4096), (2, 3, 2048)]) +@pytest.mark.parametrize("table_dtype", [torch.bfloat16, torch.float32]) +@pytest.mark.parametrize("compiled", [False, True]) +def test_ltx2_ada_values9(batch, seq, hidden, table_dtype, compiled): + scale_shift_table = torch.randn(9, hidden, device=DEVICE, dtype=table_dtype) + timestep = torch.randn(batch, seq, 9 * hidden, device=DEVICE, dtype=torch.bfloat16) + + fn = ( + torch.compile(ltx2_ada_values9, fullgraph=True) + if compiled + else ltx2_ada_values9 + ) + actual = fn(scale_shift_table, timestep) + expected = _ltx2_reference(scale_shift_table, timestep) + + assert len(actual) == 9 + for got, want in zip(actual, expected, strict=True): + # Each slice must come out naturally contiguous -- that is the point + # of the kernel; a strided slice would re-add the downstream copy. + assert got.is_contiguous() + assert torch.equal(got, want) + + +@torch.no_grad() +def test_ltx2_ada_values9_rejects_unsupported_shape(): + scale_shift_table = torch.randn(8, 4096, device=DEVICE, dtype=torch.bfloat16) + timestep = torch.randn(1, 1, 9 * 4096, device=DEVICE, dtype=torch.bfloat16) + with pytest.raises(ValueError, match="scale_shift_table"): + ltx2_ada_values9(scale_shift_table, timestep) + + +# --------------------------------------------------------------------------- +# select-0/1 LayerNorm modulation (Qwen-Image) +# --------------------------------------------------------------------------- + +SELECT01_DTYPES = get_ci_test_range( + [torch.float16, torch.bfloat16, torch.float32], [torch.float16, torch.bfloat16] +) +SELECT01_SHAPES = get_ci_test_range( + [(b, s, h) for b in (1, 2, 4) for s in (6, 33, 128, 257) for h in (512, 3072)], + [(1, 6, 512), (2, 128, 3072)], +) +EPS = 1e-6 + + +def _select01_reference(normalized, mods, index): + scale0, shift0, gate0, scale1, shift1, gate1 = mods + idx = index.bool().unsqueeze(-1) + scale = torch.where(idx, scale1.unsqueeze(1), scale0.unsqueeze(1)) + shift = torch.where(idx, shift1.unsqueeze(1), shift0.unsqueeze(1)) + gate = torch.where(idx, gate1.unsqueeze(1), gate0.unsqueeze(1)) + return normalized * (1 + scale) + shift, gate + + +@pytest.mark.parametrize("dtype", SELECT01_DTYPES) +@pytest.mark.parametrize("shape", SELECT01_SHAPES) +@pytest.mark.parametrize("with_residual", [False, True]) +def test_layernorm_scale_shift_gate_select01(dtype, shape, with_residual): + batch_size, seq_len, hidden_size = shape + x = torch.randn(batch_size, seq_len, hidden_size, device=DEVICE, dtype=dtype) + weight = torch.randn(hidden_size, device=DEVICE, dtype=dtype) + bias = torch.randn(hidden_size, device=DEVICE, dtype=dtype) + index = torch.randint(0, 2, (batch_size, seq_len), device=DEVICE, dtype=torch.int32) + mods = tuple( + torch.randn(batch_size, hidden_size, device=DEVICE, dtype=dtype) + for _ in range(6) + ) + scale0, shift0, gate0, scale1, shift1, gate1 = mods + + if with_residual: + residual = torch.randn_like(x) + residual_gate = torch.randn_like(x) + residual_ref = residual + residual_gate * x + normalized = norm_infer( + residual_ref.flatten(0, 1), weight, bias, eps=EPS, is_rms_norm=False + ).view_as(residual_ref) + out_ref, gate_ref = _select01_reference(normalized, mods, index) + out, residual_out, gate = ( + fuse_residual_layernorm_scale_shift_gate_select01_kernel( + x.contiguous(), + residual=residual.contiguous(), + residual_gate=residual_gate.contiguous(), + weight=weight, + bias=bias, + scale0=scale0, + shift0=shift0, + gate0=gate0, + scale1=scale1, + shift1=shift1, + gate1=gate1, + index=index, + eps=EPS, + ) + ) + else: + normalized = norm_infer( + x.flatten(0, 1), weight, bias, eps=EPS, is_rms_norm=False + ).view_as(x) + out_ref, gate_ref = _select01_reference(normalized, mods, index) + residual_ref = residual_out = None + out, gate = fuse_layernorm_scale_shift_gate_select01_kernel( + x.contiguous(), + weight=weight, + bias=bias, + scale0=scale0, + shift0=shift0, + gate0=gate0, + scale1=scale1, + shift1=shift1, + gate1=gate1, + index=index, + eps=EPS, + ) + + tol = 1e-5 if dtype == torch.float32 else 5e-2 + torch.testing.assert_close(out, out_ref, atol=tol, rtol=tol) + torch.testing.assert_close(gate, gate_ref, atol=tol, rtol=tol) + if with_residual: + torch.testing.assert_close(residual_out, residual_ref, atol=tol, rtol=tol) + + +# --------------------------------------------------------------------------- +# Sinusoidal timestep embedding +# --------------------------------------------------------------------------- + +TIMESTEP_BATCHES = get_ci_test_range( + [1, 2, 8, 128, 256, 512, 1536, 2048, 4096, 11008, 16384], [1, 128, 2048, 16384] +) +TIMESTEP_DIMS = get_ci_test_range( + [32, 128, 256, 512, 1536, 2048, 4096, 8192], [32, 512, 8192] +) +TIMESTEP_DTYPES = get_ci_test_range( + [torch.float16, torch.bfloat16, torch.float32], [torch.float16, torch.bfloat16] +) + + +def timestep_embedding_reference( + timesteps, + dim, + *, + flip_sin_to_cos=False, + downscale_freq_shift=1, + scale=1, + max_period=10000, +): + """diffusers' ``get_timestep_embedding``, kept verbatim as the oracle.""" + assert len(timesteps.shape) == 1, "Timesteps should be a 1d-array" + timesteps = timesteps.to(torch.float32) + half_dim = dim // 2 + exponent = -torch.log( + torch.tensor(max_period, dtype=torch.float32, device=timesteps.device) + ) * torch.arange( + start=0, end=half_dim, dtype=torch.float32, device=timesteps.device + ) + exponent = exponent / (half_dim - downscale_freq_shift) + + emb = torch.exp(exponent) + emb = timesteps[:, None].float() * emb[None, :] + emb = scale * emb + emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1) + if flip_sin_to_cos: + emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1) + if dim % 2 == 1: + emb = torch.nn.functional.pad(emb, (0, 1, 0, 0)) + return emb + + +@pytest.mark.parametrize("batch_size", TIMESTEP_BATCHES) +@pytest.mark.parametrize("dim", TIMESTEP_DIMS) +@pytest.mark.parametrize("dtype", TIMESTEP_DTYPES) +@pytest.mark.parametrize( + "flip_sin_to_cos,downscale_freq_shift,scale", + [ + (True, 0, 1), # the sgl-diffusion default + (False, 1, 1), # the diffusers default + (True, 1, 0.01), # scaled variant used by the SD-style embedders + ], +) +def test_timestep_embedding_matches_diffusers( + batch_size, dim, dtype, flip_sin_to_cos, downscale_freq_shift, scale +): + t = torch.randint(low=0, high=1000, size=(batch_size,), device=DEVICE).to(dtype) + kwargs = dict( + flip_sin_to_cos=flip_sin_to_cos, + downscale_freq_shift=downscale_freq_shift, + scale=scale, + max_period=10000, + ) + torch.testing.assert_close( + timestep_embedding(t, dim, **kwargs), + timestep_embedding_reference(t, dim, **kwargs), + atol=1e-3, + rtol=1e-3, + ) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernels/ops/diffusion/test_modulate_scale_shift.py b/test/registered/kernels/ops/diffusion/test_modulate_scale_shift.py deleted file mode 100644 index 68d841d8dccd..000000000000 --- a/test/registered/kernels/ops/diffusion/test_modulate_scale_shift.py +++ /dev/null @@ -1,58 +0,0 @@ -import pytest -import torch - -from sglang.kernels.ops.diffusion.modulate_scale_shift import ( - can_use_modulate_scale_shift_cuda, - modulate_scale_shift, - modulate_scale_shift_cuda, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large") -register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="4-gpu-b200") - -# FLUX.1 1024^2 adaLN shapes (D=3072) plus batched and odd-length coverage. -CASES = [(1, 4096, 3072), (1, 512, 3072), (1, 4608, 3072), (2, 1024, 3072), (1, 17, 64)] - - -@pytest.fixture(autouse=True) -def cuda_setup(): - if not torch.cuda.is_available(): - pytest.skip("CUDA required") - torch.cuda.manual_seed(0) - - -def _eager(x, scale, shift): - return x * (1 + scale[:, None]) + shift[:, None] - - -@pytest.mark.parametrize("shape", CASES) -@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) -def test_modulate_scale_shift_matches_eager(shape, dtype): - x = torch.randn(shape, device="cuda", dtype=dtype) - scale = torch.randn((shape[0], shape[-1]), device="cuda", dtype=dtype) - shift = torch.randn_like(scale) - out = modulate_scale_shift_cuda(x, scale, shift) - assert torch.equal(out, _eager(x, scale, shift)) # bitwise contract - - -def test_modulate_scale_shift_adaln_chunk_views(): - x = torch.randn((1, 4096, 3072), device="cuda", dtype=torch.bfloat16) - emb = torch.randn((1, 6 * 3072), device="cuda", dtype=torch.bfloat16) - shift, scale = emb.chunk(6, dim=1)[:2] - assert can_use_modulate_scale_shift_cuda(x, scale, shift) - out = modulate_scale_shift_cuda(x, scale, shift) - assert torch.equal(out, _eager(x, scale, shift)) - - -def test_modulate_scale_shift_guards_reject_fp32(): - x = torch.randn((1, 64, 64), device="cuda", dtype=torch.float32) - row = torch.randn((1, 64), device="cuda", dtype=torch.float32) - assert not can_use_modulate_scale_shift_cuda(x, row, row) - assert torch.equal(modulate_scale_shift(x, row, row), _eager(x, row, row)) - - -if __name__ == "__main__": - import sys - - sys.exit(pytest.main([__file__])) diff --git a/test/registered/kernels/ops/diffusion/test_native_bf16_rmsnorm.py b/test/registered/kernels/ops/diffusion/test_native_bf16_rmsnorm.py deleted file mode 100644 index 6c50d1c32a5a..000000000000 --- a/test/registered/kernels/ops/diffusion/test_native_bf16_rmsnorm.py +++ /dev/null @@ -1,87 +0,0 @@ -import pytest -import torch - -from sglang.kernels.ops.diffusion.triton.native_bf16_rmsnorm import ( - rmsnorm_scale, - rmsnorm_tanh_residual, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=8, stage="base-b-kernel-unit", runner_config="1-gpu-large") - -EPS = 1e-5 - - -def _native_bf16_rmsnorm(x: torch.Tensor, weight: torch.Tensor) -> torch.Tensor: - square = (x * x).to(torch.bfloat16) - mean_square = square.mean(dim=-1, keepdim=True).to(torch.bfloat16) - rstd = torch.rsqrt((mean_square + EPS).to(torch.bfloat16).float()).to( - torch.bfloat16 - ) - return ((x * rstd).to(torch.bfloat16) * weight).to(torch.bfloat16) - - -def test_native_bf16_rmsnorm_rejects_unsupported_inputs(): - x = torch.randn(2, 3, 16, dtype=torch.bfloat16) - weight = torch.randn(16, dtype=torch.bfloat16) - modulation = torch.randn(2, 1, 16, dtype=torch.bfloat16) - residual = torch.randn_like(x) - - assert rmsnorm_scale(x, weight, modulation, EPS) is None - assert rmsnorm_tanh_residual(x, modulation, residual, weight, EPS) is None - assert rmsnorm_scale(x, weight[:-1], modulation, EPS) is None - assert rmsnorm_tanh_residual(x, modulation, residual[..., :-1], weight, EPS) is None - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") -@pytest.mark.parametrize("shape", [(1, 32, 2560), (2, 17, 256)]) -def test_rmsnorm_scale_matches_native_bf16(shape): - torch.manual_seed(0) - batch, _, dim = shape - x = torch.randn(shape, device="cuda", dtype=torch.bfloat16) - weight = torch.randn(dim, device="cuda", dtype=torch.bfloat16) - scale = torch.randn(batch, 1, dim, device="cuda", dtype=torch.bfloat16) - - actual = rmsnorm_scale(x, weight, scale, EPS) - expected = (_native_bf16_rmsnorm(x, weight) * scale).to(torch.bfloat16) - - assert actual is not None - torch.testing.assert_close(actual, expected, atol=2e-2, rtol=2e-2) - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") -@pytest.mark.parametrize("shape", [(1, 32, 2560), (2, 17, 256)]) -def test_rmsnorm_tanh_residual_matches_native_bf16(shape): - torch.manual_seed(0) - batch, _, dim = shape - x = torch.randn(shape, device="cuda", dtype=torch.bfloat16) - gate = torch.randn(batch, 1, dim, device="cuda", dtype=torch.bfloat16) - residual = torch.randn(shape, device="cuda", dtype=torch.bfloat16) - weight = torch.randn(dim, device="cuda", dtype=torch.bfloat16) - - actual = rmsnorm_tanh_residual(x, gate, residual, weight, EPS) - norm = _native_bf16_rmsnorm(x, weight) - gated = (torch.tanh(gate.float()).to(torch.bfloat16) * norm).to(torch.bfloat16) - expected = (residual + gated).to(torch.bfloat16) - - assert actual is not None - # Triton's exp-based tanh can differ slightly from torch.tanh in BF16. - torch.testing.assert_close(actual, expected, atol=4e-2, rtol=2e-2) - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") -def test_native_bf16_rmsnorm_rejects_hidden_size_above_limit(): - dim = 8448 - x = torch.empty(1, 1, dim, device="cuda", dtype=torch.bfloat16) - weight = torch.empty(dim, device="cuda", dtype=torch.bfloat16) - modulation = torch.empty(1, 1, dim, device="cuda", dtype=torch.bfloat16) - residual = torch.empty_like(x) - - assert rmsnorm_scale(x, weight, modulation, EPS) is None - assert rmsnorm_tanh_residual(x, modulation, residual, weight, EPS) is None - - -if __name__ == "__main__": - import sys - - sys.exit(pytest.main([__file__])) diff --git a/test/registered/kernels/ops/diffusion/test_norm.py b/test/registered/kernels/ops/diffusion/test_norm.py new file mode 100644 index 000000000000..a3f91881a153 --- /dev/null +++ b/test/registered/kernels/ops/diffusion/test_norm.py @@ -0,0 +1,527 @@ +"""``diffusion.norm``: GroupNorm / RMSNorm / LayerNorm and their fused epilogues. + +This domain has the most implementations of any in the package (see the +selection matrix in ``sglang/kernels/ops/diffusion/README.md``), so the suite +is organized by *kernel*, and each section states which oracle it is held to: + +- ``triton_group_norm_silu`` / ``apply_group_norm_silu`` -> ``F.group_norm`` + + ``F.silu`` with a per-dtype tolerance (fp32 statistics, different reduction). +- the two-pass channels-last GroupNorm -> same oracle, plus its support + predicates (the kernels raise on an unsupported input rather than returning + ``None``). +- ``rmsnorm_scale`` / ``rmsnorm_tanh_residual`` -> a bf16-native reference that + reproduces Z-Image's own norm, with a tolerance for Triton's exp-based tanh. +- the CuTe-DSL and FlyDSL fused norm+scale/shift -> an fp32 reference chain. + +The *bit-exact* norms (``fused_rmsnorm_scale_shift_bitexact``, +``fused_layernorm_modulate``, ``zimage_qk_rmsnorm_native``) are exercised +through their model wrappers in ``test_model_fast_paths.py``, where the live +eager chain they must reproduce is available. +""" + +import sys + +import pytest +import torch +import torch.nn as nn +import torch.nn.functional as F +from einops import rearrange + +from sglang.kernels.ops.diffusion import ( + apply_group_norm_silu, + can_use_group_norm_silu_4d, + can_use_wan_rmsnorm_silu, + group_norm_silu_4d, + rmsnorm_scale, + rmsnorm_tanh_residual, + triton_group_norm_silu, + wan_rmsnorm_silu, +) +from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci + +register_cuda_ci(est_time=85, stage="base-b-kernel-unit", runner_config="1-gpu-large") +register_amd_ci(est_time=30, stage="jit-kernel-unit", runner_config="amd") +register_amd_ci(est_time=15, suite="nightly-amd-kernel-1-gpu", nightly=True) + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") + +DEVICE = "cuda" +DTYPES = [torch.float16, torch.bfloat16, torch.float32] +EPS = 1e-5 + + +def _tol(dtype: torch.dtype) -> tuple[float, float]: + if dtype == torch.float32: + return 1e-5, 1e-5 + if dtype == torch.bfloat16: + return 7e-2, 2e-2 + return 3e-3, 3e-3 + + +@pytest.fixture(autouse=True) +def cuda_setup(): + torch.cuda.manual_seed(0) + + +def _cl3d(shape, dtype): + return torch.randn(shape, device=DEVICE, dtype=dtype).contiguous( + memory_format=torch.channels_last_3d + ) + + +# --------------------------------------------------------------------------- +# GroupNorm + SiLU +# --------------------------------------------------------------------------- + +GN_CASES = [ + pytest.param((2, 64, 32, 32), 32, id="image_2d"), + pytest.param((1, 64, 4, 16, 16), 32, id="video_3d"), + pytest.param((4, 128), 32, id="token_2d"), +] + + +def _gn_silu_reference(x, weight, bias, num_groups, eps=EPS): + return F.silu(F.group_norm(x, num_groups, weight=weight, bias=bias, eps=eps)) + + +@torch.no_grad() +@pytest.mark.parametrize("shape,num_groups", GN_CASES) +@pytest.mark.parametrize("dtype", DTYPES) +def test_triton_group_norm_silu(shape, num_groups, dtype): + channels = shape[1] + x = torch.randn(shape, device=DEVICE, dtype=dtype) + weight = torch.randn(channels, device=DEVICE, dtype=dtype) + bias = torch.randn(channels, device=DEVICE, dtype=dtype) + + atol, rtol = _tol(dtype) + torch.testing.assert_close( + triton_group_norm_silu(x, weight, bias, num_groups=num_groups), + _gn_silu_reference(x, weight, bias, num_groups), + atol=atol, + rtol=rtol, + ) + + +@torch.no_grad() +def test_triton_group_norm_silu_large_tile_bf16(): + # A tile large enough to force the chunked launch path (128 channels over + # 20x256x256), which the small cases above never reach. + shape, num_groups = (1, 128, 20, 256, 256), 32 + x = torch.randn(shape, device=DEVICE, dtype=torch.bfloat16) + weight = torch.randn(shape[1], device=DEVICE, dtype=torch.bfloat16) + bias = torch.randn(shape[1], device=DEVICE, dtype=torch.bfloat16) + + atol, rtol = _tol(torch.bfloat16) + torch.testing.assert_close( + triton_group_norm_silu(x, weight, bias, num_groups=num_groups), + _gn_silu_reference(x, weight, bias, num_groups), + atol=atol, + rtol=rtol, + ) + + +@torch.no_grad() +@pytest.mark.parametrize("shape,num_groups", GN_CASES[:2]) +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) +def test_apply_group_norm_silu_module_wrapper(shape, num_groups, dtype): + # The nn.Module-taking wrapper must match the eager module pair it stands + # in for, including its own guard set (affine=True, non-inplace SiLU). + norm = nn.GroupNorm(num_groups, shape[1], eps=EPS, affine=True).to( + device=DEVICE, dtype=dtype + ) + activation = nn.SiLU() + x = torch.randn(shape, device=DEVICE, dtype=dtype) + + atol, rtol = _tol(dtype) + torch.testing.assert_close( + apply_group_norm_silu(x, norm, activation), + activation(norm(x)), + atol=atol, + rtol=rtol, + ) + + +@torch.no_grad() +def test_group_norm_silu_4d_channels_last_and_guards(): + gn = nn.GroupNorm(32, 128, eps=1e-6).to(DEVICE, torch.bfloat16) + x = torch.randn(1, 128, 64, 64, device=DEVICE, dtype=torch.bfloat16).to( + memory_format=torch.channels_last + ) + assert can_use_group_norm_silu_4d(x, gn.weight, gn.bias, 32) + out = group_norm_silu_4d(x, gn.weight, gn.bias, 32, 1e-6) + assert out.is_contiguous(memory_format=torch.channels_last) + torch.testing.assert_close(out.float(), F.silu(gn(x)).float(), atol=0.06, rtol=0) + + # Guards: the kernel exists only for channels_last inputs with device-side + # affine params and a non-empty spatial extent. Each rejected case must + # fail the predicate *and* raise if called anyway -- silently returning + # ``None`` is what this protocol replaced. + rejected = [ + (x.contiguous(), gn.weight, gn.bias), # contiguous (NCHW) layout + (x, gn.weight.cpu(), gn.bias), # host-side affine + (x[..., :0, :], gn.weight, gn.bias), # empty spatial extent + ] + for args in rejected: + assert not can_use_group_norm_silu_4d(*args, 32) + with pytest.raises(ValueError): + group_norm_silu_4d(*args, 32, 1e-6) + + +# --------------------------------------------------------------------------- +# BF16-native RMSNorm fusions (Z-Image / Ideogram) +# --------------------------------------------------------------------------- + + +def _native_bf16_rmsnorm(x, weight): + """Z-Image's own norm: every step materialized in bf16, no fp32 carry.""" + square = (x * x).to(torch.bfloat16) + mean_square = square.mean(dim=-1, keepdim=True).to(torch.bfloat16) + rstd = torch.rsqrt((mean_square + EPS).to(torch.bfloat16).float()).to( + torch.bfloat16 + ) + return ((x * rstd).to(torch.bfloat16) * weight).to(torch.bfloat16) + + +@pytest.mark.parametrize("shape", [(1, 32, 2560), (2, 17, 256)]) +def test_rmsnorm_scale_matches_native_bf16(shape): + batch, _, dim = shape + x = torch.randn(shape, device=DEVICE, dtype=torch.bfloat16) + weight = torch.randn(dim, device=DEVICE, dtype=torch.bfloat16) + scale = torch.randn(batch, 1, dim, device=DEVICE, dtype=torch.bfloat16) + + actual = rmsnorm_scale(x, weight, scale, EPS) + assert actual is not None + torch.testing.assert_close( + actual, + (_native_bf16_rmsnorm(x, weight) * scale).to(torch.bfloat16), + atol=2e-2, + rtol=2e-2, + ) + + +@pytest.mark.parametrize("shape", [(1, 32, 2560), (2, 17, 256)]) +def test_rmsnorm_tanh_residual_matches_native_bf16(shape): + batch, _, dim = shape + x = torch.randn(shape, device=DEVICE, dtype=torch.bfloat16) + gate = torch.randn(batch, 1, dim, device=DEVICE, dtype=torch.bfloat16) + residual = torch.randn(shape, device=DEVICE, dtype=torch.bfloat16) + weight = torch.randn(dim, device=DEVICE, dtype=torch.bfloat16) + + actual = rmsnorm_tanh_residual(x, gate, residual, weight, EPS) + norm = _native_bf16_rmsnorm(x, weight) + gated = (torch.tanh(gate.float()).to(torch.bfloat16) * norm).to(torch.bfloat16) + + assert actual is not None + # Triton's exp-based tanh can differ slightly from torch.tanh in bf16. + torch.testing.assert_close( + actual, (residual + gated).to(torch.bfloat16), atol=4e-2, rtol=2e-2 + ) + + +@pytest.mark.parametrize("on_host", [True, False]) +def test_native_bf16_rmsnorm_rejects_unsupported_inputs(on_host): + # Host tensors and a hidden size past the kernel limit are both outside + # the contract; these entry points signal that by returning None (they + # are internal fast-path probes, not public predicate+kernel pairs). + device = "cpu" if on_host else DEVICE + dim = 16 if on_host else 8448 + x = torch.randn(2, 3, dim, dtype=torch.bfloat16, device=device) + weight = torch.randn(dim, dtype=torch.bfloat16, device=device) + modulation = torch.randn(2, 1, dim, dtype=torch.bfloat16, device=device) + residual = torch.randn_like(x) + + assert rmsnorm_scale(x, weight, modulation, EPS) is None + assert rmsnorm_tanh_residual(x, modulation, residual, weight, EPS) is None + if on_host: + # Mismatched trailing dims are rejected too. + assert rmsnorm_scale(x, weight[:-1], modulation, EPS) is None + assert ( + rmsnorm_tanh_residual(x, modulation, residual[..., :-1], weight, EPS) + is None + ) + + +# --------------------------------------------------------------------------- +# Wan VAE channels_last_3d RMSNorm + SiLU +# --------------------------------------------------------------------------- + + +@torch.no_grad() +@pytest.mark.parametrize( + "x_dtype,affine_dtype,atol,rtol", + [ + (torch.float32, torch.float32, 1e-5, 1e-5), # FastWan2.2 fp32 decode + (torch.bfloat16, torch.float32, 1.5e-1, 3e-2), # Wan2.1 bf16 autocast + ], +) +def test_wan_rmsnorm_silu_numerics(x_dtype, affine_dtype, atol, rtol): + x = _cl3d((1, 96, 3, 10, 14), x_dtype) + gamma = torch.randn((96, 1, 1, 1), device=DEVICE, dtype=affine_dtype) + for bias in (None, torch.randn_like(gamma)): + expected = F.silu( + F.normalize(x, dim=1) * 96**0.5 * gamma + (0 if bias is None else bias) + ) + actual = wan_rmsnorm_silu(x, gamma, bias) + assert actual.dtype == expected.dtype + # The kernel must preserve the channels_last_3d layout; a relayout + # here would undo the reason the decoder runs in that format. + assert actual.stride() == x.stride() + torch.testing.assert_close(actual, expected, atol=atol, rtol=rtol) + + +@torch.no_grad() +def test_wan_rmsnorm_silu_rejects_empty_input(): + x = torch.empty(1, 96, 0, 2, 2, device=DEVICE, dtype=torch.bfloat16).to( + memory_format=torch.channels_last_3d + ) + gamma = torch.ones(96, 1, 1, 1, device=DEVICE, dtype=torch.bfloat16) + assert not can_use_wan_rmsnorm_silu(x, gamma, None) + with pytest.raises(ValueError): + wan_rmsnorm_silu(x, gamma) + + +# --------------------------------------------------------------------------- +# CuTe-DSL fused (residual +) norm + scale/shift +# --------------------------------------------------------------------------- + +SHAPE_MAP = { + "1": lambda B, S, F_, D: (1,), + "D": lambda B, S, F_, D: (D,), + "1D": lambda B, S, F_, D: (1, D), + "BD": lambda B, S, F_, D: (B, D), + "11D": lambda B, S, F_, D: (1, 1, D), + "B1D": lambda B, S, F_, D: (B, 1, D), + "1SD": lambda B, S, F_, D: (1, S, D), + "BSD": lambda B, S, F_, D: (B, S, D), + "BF1D": lambda B, S, F_, D: (B, F_, 1, D), +} +# (B, S, F, D) +CUTE_SHAPES = [ + (1, 115200, 1, 3072), # HunyuanVideo + (1, 32760, 1, 1536), # Wan + (1, 6, 1, 3072), # Qwen-Image + (1, 1024, 8, 3072), + (4, 512, 16, 3072), +] +NORM_TYPES = ["layer", "rms"] +AFFINE_MODES = ["D", "NAT"] +INDEX_MODES = ["BSD", "1", "1SD", "BD", "B1D", "D", "1D", "11D", "BF1D"] + + +def _make_tensor(index_mode, shape, dtype): + if index_mode == "NAT": + return None + return torch.randn(*SHAPE_MAP[index_mode](*shape), device=DEVICE, dtype=dtype) + + +def _apply_scale_shift(y, scale, shift): + if scale.ndim == 4: + num_frame = scale.shape[1] + return rearrange( + rearrange(y, "b (f l) d -> b f l d", f=num_frame) * (1 + scale) + shift, + "b f l d -> b (f l) d", + ) + scale = rearrange(scale, "b d -> b 1 d") if scale.ndim == 2 else scale + shift = rearrange(shift, "b d -> b 1 d") if shift.ndim == 2 else shift + return y * (1 + scale) + shift + + +def _cute_reference(residual, x, gate, weight, bias, scale, shift, norm_type, eps): + """fp32 oracle for both variants; ``residual is None`` = no-residual form.""" + original_dtype = x.dtype + residual, x, gate, weight, bias, scale, shift = ( + v.float() if isinstance(v, torch.Tensor) else v + for v in (residual, x, gate, weight, bias, scale, shift) + ) + residual_out = None + if residual is not None: + if isinstance(gate, int): + x = residual + gate * x + elif gate.ndim == 4: + folded = rearrange(x, "b (f l) d -> b f l d", f=gate.shape[1]) + x = residual + rearrange(folded * gate, "b f l d -> b (f l) d") + else: + g = rearrange(gate, "b d -> b 1 d") if gate.ndim == 2 else gate + x = residual + g * x + residual_out = x.to(original_dtype) + if norm_type == "layer": + norm = torch.layer_norm(x, x.shape[-1:], eps=eps, weight=weight, bias=bias) + else: + norm = torch.rms_norm(x, x.shape[-1:], eps=eps, weight=weight) + return _apply_scale_shift(norm, scale, shift).to(original_dtype), residual_out + + +@torch.no_grad() +def _run_cute( + with_residual, + shape=CUTE_SHAPES[0], + dtype=DTYPES[0], + affine_dtype=DTYPES[0], + mod_dtype=DTYPES[0], + norm_type=NORM_TYPES[0], + affine_mode=AFFINE_MODES[0], + gate_mode="B1D", + index_mode="BSD", + eps=EPS, +): + # Imported inside the helper for the same reason as FlyDSL below: the + # CuTe-DSL path needs cuda-python + CUTLASS, and naming a facade export + # resolves it immediately -- at module level that would take the Triton + # sections of this file down with it wherever CUTLASS is absent. + from sglang.kernels.ops.diffusion import ( + fused_norm_scale_shift, + fused_scale_residual_norm_scale_shift, + ) + + x = _make_tensor("BSD", shape, dtype) + weight = _make_tensor(affine_mode, shape, affine_dtype) + bias = _make_tensor(affine_mode, shape, affine_dtype) + scale = _make_tensor(index_mode, shape, mod_dtype) + shift = _make_tensor(index_mode, shape, mod_dtype) + tol = 1e-5 if dtype == torch.float32 else 5e-2 + + if with_residual: + residual = _make_tensor("BSD", shape, dtype) + gate = _make_tensor(gate_mode, shape, dtype) + y, res = fused_scale_residual_norm_scale_shift( + residual, x, gate, weight, bias, scale, shift, norm_type, eps + ) + y_ref, res_ref = _cute_reference( + residual, x, gate, weight, bias, scale, shift, norm_type, eps + ) + torch.testing.assert_close(res, res_ref, atol=tol, rtol=tol) + else: + y = fused_norm_scale_shift(x, weight, bias, scale, shift, norm_type, eps) + y_ref, _ = _cute_reference( + None, x, None, weight, bias, scale, shift, norm_type, eps + ) + torch.testing.assert_close(y, y_ref, atol=tol, rtol=tol) + + +@pytest.mark.parametrize("with_residual", [False, True]) +@pytest.mark.parametrize("norm_type", NORM_TYPES) +@pytest.mark.parametrize("shape", CUTE_SHAPES) +@pytest.mark.parametrize("dtype", DTYPES) +def test_cutedsl_norm_scale_shift_shapes(with_residual, norm_type, shape, dtype): + _run_cute(with_residual, shape=shape, dtype=dtype, norm_type=norm_type) + + +@pytest.mark.parametrize("with_residual", [False, True]) +@pytest.mark.parametrize("norm_type", NORM_TYPES) +@pytest.mark.parametrize("dtype", DTYPES) +@pytest.mark.parametrize("operand", ["affine", "modulation"]) +def test_cutedsl_norm_scale_shift_mixed_operand_dtypes( + with_residual, norm_type, dtype, operand +): + # The affine params and the modulation rows may each arrive in a dtype + # other than the activation's; both combinations must dispatch. + kwargs = {"affine_dtype" if operand == "affine" else "mod_dtype": dtype} + _run_cute(with_residual, norm_type=norm_type, **kwargs) + + +@pytest.mark.parametrize("with_residual", [False, True]) +@pytest.mark.parametrize("norm_type", NORM_TYPES) +@pytest.mark.parametrize("affine_mode", AFFINE_MODES) +def test_cutedsl_norm_scale_shift_affine_modes(with_residual, norm_type, affine_mode): + _run_cute(with_residual, norm_type=norm_type, affine_mode=affine_mode) + + +@pytest.mark.parametrize("with_residual", [False, True]) +@pytest.mark.parametrize("norm_type", NORM_TYPES) +@pytest.mark.parametrize("index_mode", INDEX_MODES) +def test_cutedsl_norm_scale_shift_index_modes(with_residual, norm_type, index_mode): + _run_cute(with_residual, norm_type=norm_type, index_mode=index_mode) + + +@pytest.mark.parametrize("norm_type", NORM_TYPES) +@pytest.mark.parametrize("index_mode", INDEX_MODES) +def test_cutedsl_scale_residual_gate_index_modes(norm_type, index_mode): + _run_cute(True, norm_type=norm_type, gate_mode=index_mode) + + +def test_validate_scale_shift_rejects_non_divisible_frames(): + from sglang.kernels.ops.diffusion import validate_scale_shift + + with pytest.raises(ValueError, match=r"S\(10\) must be divisible by F\(4\)"): + validate_scale_shift( + torch.empty((1, 4, 1, 256), device=DEVICE, dtype=torch.float16), 1, 10, 256 + ) + + +# --------------------------------------------------------------------------- +# FlyDSL fused norm + scale/shift (ROCm gfx950) +# --------------------------------------------------------------------------- + +FLYDSL_D = 5120 +FLYDSL_EPS = 1e-6 + + +def _require_rocm(): + if not torch.version.hip: + pytest.skip("ROCm/HIP required for FlyDSL kernels") + + +def _flydsl_reference(residual, x, gate, weight, bias, scale, shift, norm_type, eps): + if residual is not None: + x = (residual.float() + x.float() * gate.float()).to(torch.bfloat16) + residual_out = x + else: + residual_out = None + if norm_type == "layer": + normed = F.layer_norm(x.float(), (FLYDSL_D,), weight, bias, eps) + else: + var = x.float().pow(2).mean(-1, keepdim=True) + normed = x.float() * torch.rsqrt(var + eps) * weight.float() + y = (normed * (1.0 + scale.float()) + shift.float()).to(torch.bfloat16) + return y, residual_out + + +@pytest.mark.parametrize("with_residual", [False, True]) +@pytest.mark.parametrize( + "norm_type,B,L", + [("rms", 1, 16), ("rms", 2, 16), ("layer", 2, 16), ("rms", 1, 90000)], +) +def test_flydsl_norm_scale_shift(with_residual, norm_type, B, L): + _require_rocm() + # Imported inside the test: the FlyDSL compiler only exists on ROCm, and + # the facade resolves an export the moment it is named -- a module-level + # import here would fail collection of this whole file on CUDA. + from sglang.kernels.ops.diffusion import ( + flydsl_fused_residual_norm_scale_shift, + flydsl_norm_scale_shift, + ) + + torch.manual_seed(42) + shape = (B, L, FLYDSL_D) + x = torch.randn(shape, device=DEVICE, dtype=torch.bfloat16) + weight = torch.randn(FLYDSL_D, device=DEVICE, dtype=torch.float32) + bias = ( + torch.randn(FLYDSL_D, device=DEVICE, dtype=torch.float32) + if norm_type == "layer" + else None + ) + scale = torch.randn(B, 1, FLYDSL_D, device=DEVICE, dtype=torch.bfloat16) + shift = torch.randn(B, 1, FLYDSL_D, device=DEVICE, dtype=torch.bfloat16) + + if with_residual: + residual = torch.randn(shape, device=DEVICE, dtype=torch.bfloat16) + gate = torch.randn(B, 1, FLYDSL_D, device=DEVICE, dtype=torch.bfloat16) + y, res = flydsl_fused_residual_norm_scale_shift( + residual, x, gate, weight, bias, scale, shift, norm_type, FLYDSL_EPS + ) + y_ref, res_ref = _flydsl_reference( + residual, x, gate, weight, bias, scale, shift, norm_type, FLYDSL_EPS + ) + torch.testing.assert_close(res, res_ref, atol=5e-2, rtol=5e-2) + else: + y = flydsl_norm_scale_shift( + x, weight, bias, scale, shift, norm_type, FLYDSL_EPS + ) + y_ref, _ = _flydsl_reference( + None, x, None, weight, bias, scale, shift, norm_type, FLYDSL_EPS + ) + torch.testing.assert_close(y, y_ref, atol=1.0, rtol=5e-2) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernels/ops/diffusion/test_quality_gate.py b/test/registered/kernels/ops/diffusion/test_quality_gate.py deleted file mode 100644 index ba90c52ebade..000000000000 --- a/test/registered/kernels/ops/diffusion/test_quality_gate.py +++ /dev/null @@ -1,47 +0,0 @@ -import sys - -import pytest -import torch.nn as nn - -from sglang.kernels.ops.diffusion.quality_gate import QualityGatedFusion -from sglang.test.ci.ci_register import register_cpu_ci - -register_cpu_ci(est_time=2, suite="base-a-test-cpu") - - -def test_quality_gate_mounts_and_unmounts_all_sites(): - fusion = QualityGatedFusion( - name="test fusion", - marker_attr="_test_fusion_site", - enabled_attr="_test_fusion_enabled", - ) - root = nn.ModuleList([nn.Module(), nn.Module()]) - for index, site in enumerate(root): - fusion.mark(site, index) - - assert [fusion.metadata(site) for site in root] == [0, 1] - assert fusion.mount(root) - assert all(fusion.is_enabled(site) for site in root) - fusion.unmount(root) - assert not any(fusion.is_enabled(site) for site in root) - - -def test_quality_gate_rejection_is_all_or_nothing(): - fusion = QualityGatedFusion( - name="test fusion", - marker_attr="_test_fusion_site", - enabled_attr="_test_fusion_enabled", - ) - root = nn.ModuleList([nn.Module(), nn.Module()]) - for index, site in enumerate(root): - fusion.mark(site, index) - - assert not fusion.mount( - root, reject_reason=lambda site: "rejected" if fusion.metadata(site) else None - ) - assert not any(fusion.is_enabled(site) for site in root) - assert not fusion.mount(nn.Module()) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernels/ops/diffusion/test_qwen_image_modulation.py b/test/registered/kernels/ops/diffusion/test_qwen_image_modulation.py deleted file mode 100644 index f0a5baa17231..000000000000 --- a/test/registered/kernels/ops/diffusion/test_qwen_image_modulation.py +++ /dev/null @@ -1,228 +0,0 @@ -import sys - -import pytest -import torch -import triton - -from sglang.kernels.jit.utils import get_ci_test_range -from sglang.kernels.ops.diffusion.triton.norm import norm_infer -from sglang.kernels.ops.diffusion.triton.scale_shift import ( - fuse_layernorm_scale_shift_gate_select01_kernel, - fuse_residual_layernorm_scale_shift_gate_select01_kernel, -) -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci - -register_cuda_ci(est_time=15, stage="base-b-kernel-unit", runner_config="1-gpu-large") -# Nightly is not redundant here: it sets SGLANG_JIT_KERNEL_RUN_FULL_TESTS=1 to expand get_ci_test_range sweeps. -register_cuda_ci(est_time=30, stage="nightly", runner_config="1-gpu-large") -register_amd_ci(est_time=30, suite="nightly-amd-kernel-1-gpu", nightly=True) - -DEVICE = "cuda" -DTYPES = get_ci_test_range( - [torch.float16, torch.bfloat16, torch.float32], [torch.float16, torch.bfloat16] -) -BATCH_SIZES = get_ci_test_range([1, 2, 4], [1, 2]) -SEQ_LENS = get_ci_test_range([6, 33, 128, 257], [6, 128]) -HIDDEN_SIZES = get_ci_test_range([512, 1024, 1536, 3072], [512, 3072]) -EPS = 1e-6 - - -def _tol(dtype: torch.dtype) -> tuple[float, float]: - if dtype == torch.float32: - return 1e-5, 1e-5 - return 5e-2, 5e-2 - - -def _make_modulation_tensors(batch_size: int, hidden_size: int, dtype: torch.dtype): - scale0 = torch.randn(batch_size, hidden_size, device=DEVICE, dtype=dtype) - shift0 = torch.randn(batch_size, hidden_size, device=DEVICE, dtype=dtype) - gate0 = torch.randn(batch_size, hidden_size, device=DEVICE, dtype=dtype) - scale1 = torch.randn(batch_size, hidden_size, device=DEVICE, dtype=dtype) - shift1 = torch.randn(batch_size, hidden_size, device=DEVICE, dtype=dtype) - gate1 = torch.randn(batch_size, hidden_size, device=DEVICE, dtype=dtype) - return scale0, shift0, gate0, scale1, shift1, gate1 - - -def _baseline_select01_modulation( - x: torch.Tensor, - weight: torch.Tensor | None, - bias: torch.Tensor | None, - scale0: torch.Tensor, - shift0: torch.Tensor, - gate0: torch.Tensor, - scale1: torch.Tensor, - shift1: torch.Tensor, - gate1: torch.Tensor, - index: torch.Tensor, - eps: float, -): - normalized = norm_infer( - x.view(-1, x.shape[-1]), - weight, - bias, - eps=eps, - is_rms_norm=False, - ).view_as(x) - return _apply_select01_modulation( - normalized, scale0, shift0, gate0, scale1, shift1, gate1, index - ) - - -def _baseline_residual_select01_modulation( - x: torch.Tensor, - residual: torch.Tensor, - residual_gate: torch.Tensor, - weight: torch.Tensor | None, - bias: torch.Tensor | None, - scale0: torch.Tensor, - shift0: torch.Tensor, - gate0: torch.Tensor, - scale1: torch.Tensor, - shift1: torch.Tensor, - gate1: torch.Tensor, - index: torch.Tensor, - eps: float, -): - residual_out = residual + residual_gate * x - normalized = norm_infer( - residual_out.view(-1, residual_out.shape[-1]), - weight, - bias, - eps=eps, - is_rms_norm=False, - ).view_as(residual_out) - output, gate_out = _apply_select01_modulation( - normalized, scale0, shift0, gate0, scale1, shift1, gate1, index - ) - return output, residual_out, gate_out - - -def _apply_select01_modulation( - x: torch.Tensor, - scale0: torch.Tensor, - shift0: torch.Tensor, - gate0: torch.Tensor, - scale1: torch.Tensor, - shift1: torch.Tensor, - gate1: torch.Tensor, - index: torch.Tensor, -): - idx = index.bool().unsqueeze(-1) - scale = torch.where(idx, scale1.unsqueeze(1), scale0.unsqueeze(1)) - shift = torch.where(idx, shift1.unsqueeze(1), shift0.unsqueeze(1)) - gate = torch.where(idx, gate1.unsqueeze(1), gate0.unsqueeze(1)) - return x * (1 + scale) + shift, gate - - -@pytest.fixture(autouse=True) -def cuda_setup(): - if not torch.cuda.is_available(): - pytest.skip("CUDA required") - torch.cuda.manual_seed(0) - - -@pytest.mark.parametrize("dtype", DTYPES) -@pytest.mark.parametrize("batch_size", BATCH_SIZES) -@pytest.mark.parametrize("seq_len", SEQ_LENS) -@pytest.mark.parametrize("hidden_size", HIDDEN_SIZES) -def test_fused_layernorm_scale_shift_gate_select01( - dtype, batch_size, seq_len, hidden_size -): - x = torch.randn(batch_size, seq_len, hidden_size, device=DEVICE, dtype=dtype) - weight = torch.randn(hidden_size, device=DEVICE, dtype=dtype) - bias = torch.randn(hidden_size, device=DEVICE, dtype=dtype) - index = torch.randint(0, 2, (batch_size, seq_len), device=DEVICE, dtype=torch.int32) - scale0, shift0, gate0, scale1, shift1, gate1 = _make_modulation_tensors( - batch_size, hidden_size, dtype - ) - - out_ref, gate_ref = _baseline_select01_modulation( - x, - weight, - bias, - scale0, - shift0, - gate0, - scale1, - shift1, - gate1, - index, - EPS, - ) - out_fused, gate_fused = fuse_layernorm_scale_shift_gate_select01_kernel( - x.contiguous(), - weight=weight, - bias=bias, - scale0=scale0, - shift0=shift0, - gate0=gate0, - scale1=scale1, - shift1=shift1, - gate1=gate1, - index=index, - eps=EPS, - ) - - atol, rtol = _tol(dtype) - triton.testing.assert_close(out_ref, out_fused, atol=atol, rtol=rtol) - triton.testing.assert_close(gate_ref, gate_fused, atol=atol, rtol=rtol) - - -@pytest.mark.parametrize("dtype", DTYPES) -@pytest.mark.parametrize("batch_size", BATCH_SIZES) -@pytest.mark.parametrize("seq_len", SEQ_LENS) -@pytest.mark.parametrize("hidden_size", HIDDEN_SIZES) -def test_fused_residual_layernorm_scale_shift_gate_select01( - dtype, batch_size, seq_len, hidden_size -): - x = torch.randn(batch_size, seq_len, hidden_size, device=DEVICE, dtype=dtype) - residual = torch.randn_like(x) - residual_gate = torch.randn_like(x) - weight = torch.randn(hidden_size, device=DEVICE, dtype=dtype) - bias = torch.randn(hidden_size, device=DEVICE, dtype=dtype) - index = torch.randint(0, 2, (batch_size, seq_len), device=DEVICE, dtype=torch.int32) - scale0, shift0, gate0, scale1, shift1, gate1 = _make_modulation_tensors( - batch_size, hidden_size, dtype - ) - - out_ref, residual_ref, gate_ref = _baseline_residual_select01_modulation( - x, - residual, - residual_gate, - weight, - bias, - scale0, - shift0, - gate0, - scale1, - shift1, - gate1, - index, - EPS, - ) - out_fused, residual_fused, gate_fused = ( - fuse_residual_layernorm_scale_shift_gate_select01_kernel( - x.contiguous(), - residual=residual.contiguous(), - residual_gate=residual_gate.contiguous(), - weight=weight, - bias=bias, - scale0=scale0, - shift0=shift0, - gate0=gate0, - scale1=scale1, - shift1=shift1, - gate1=gate1, - index=index, - eps=EPS, - ) - ) - - atol, rtol = _tol(dtype) - triton.testing.assert_close(out_ref, out_fused, atol=atol, rtol=rtol) - triton.testing.assert_close(residual_ref, residual_fused, atol=atol, rtol=rtol) - triton.testing.assert_close(gate_ref, gate_fused, atol=atol, rtol=rtol) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/kernels/ops/diffusion/test_residual_gate_add.py b/test/registered/kernels/ops/diffusion/test_residual_gate_add.py deleted file mode 100644 index f281e7c89005..000000000000 --- a/test/registered/kernels/ops/diffusion/test_residual_gate_add.py +++ /dev/null @@ -1,123 +0,0 @@ -import sys - -import pytest -import torch - -from sglang.kernels.ops.diffusion.residual_gate_add import ( - can_use_residual_gate_add_cuda, - residual_gate_add, - residual_gate_add_cuda, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="1-gpu-large") -register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="4-gpu-b200") - - -CASES = [ - ((1, 1024, 4096), (1, 1, 4096)), - ((1, 512, 4096), (1, 512, 4096)), - ((1, 17, 65), (1, 1, 65)), - ((1, 17, 65), (1, 17, 65)), - # FLUX.1 / FLUX.2-klein 1024^2 shapes (D=3072): dual-stream image/text - # and single-stream/joint concat; gates are [1, 1, D] modulation rows. - ((1, 4096, 3072), (1, 1, 3072)), - ((1, 512, 3072), (1, 1, 3072)), - ((1, 4608, 3072), (1, 1, 3072)), - # FLUX.2-dev (D=6144) joint sequence. - ((1, 4608, 6144), (1, 1, 6144)), - # ERNIE-4.5-VL 1024^2 image tokens plus text tokens. - ((1, 4216, 4096), (1, 1, 4096)), -] - - -def _tol(dtype: torch.dtype) -> float: - return 1e-5 if dtype == torch.float32 else 5e-2 - - -def _assert_matches_torch(out: torch.Tensor, ref: torch.Tensor) -> None: - if ref.dtype == torch.float32: - torch.testing.assert_close(out, ref, atol=_tol(ref.dtype), rtol=_tol(ref.dtype)) - else: - torch.testing.assert_close(out, ref, atol=0, rtol=0) - - -@pytest.fixture(autouse=True) -def cuda_setup(): - if not torch.cuda.is_available(): - pytest.skip("CUDA required") - torch.cuda.manual_seed(0) - - -@pytest.mark.parametrize("residual_shape,gate_shape", CASES) -def test_residual_gate_add_matches_torch(residual_shape, gate_shape): - residual = torch.randn(residual_shape, device="cuda", dtype=torch.bfloat16) - update = torch.randn_like(residual) - gate = torch.randn(gate_shape, device="cuda", dtype=torch.bfloat16) - - out = residual_gate_add_cuda(residual, update, gate) - ref = residual + update * gate - _assert_matches_torch(out, ref) - assert torch.equal(residual_gate_add(residual, update, gate), ref) - - -@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16, torch.float32]) -@pytest.mark.parametrize("gate_shape", [(1, 1, 64), (1, 9, 64)]) -def test_residual_gate_add_dtypes(dtype, gate_shape): - residual = torch.randn((1, 9, 64), device="cuda", dtype=dtype) - update = torch.randn_like(residual) - gate = torch.randn(gate_shape, device="cuda", dtype=dtype) - - out = residual_gate_add_cuda(residual, update, gate) - ref = residual + update * gate - _assert_matches_torch(out, ref) - - -def test_can_use_residual_gate_add_cuda_rejects_unsupported_inputs(): - residual = torch.randn((1, 8, 64), device="cuda", dtype=torch.bfloat16) - update = torch.randn_like(residual) - gate = torch.randn((1, 1, 64), device="cuda", dtype=torch.bfloat16) - - assert can_use_residual_gate_add_cuda(residual, update, gate) - assert not can_use_residual_gate_add_cuda(residual.cpu(), update, gate) - assert not can_use_residual_gate_add_cuda(residual, update.float(), gate) - assert not can_use_residual_gate_add_cuda(residual, update[:, ::2], gate) - assert not can_use_residual_gate_add_cuda(residual, update, gate[:, :, ::2]) - empty_residual = residual[:, :0] - empty_update = update[:, :0] - assert not can_use_residual_gate_add_cuda(empty_residual, empty_update, gate) - assert torch.equal( - residual_gate_add(empty_residual, empty_update, gate), - empty_residual + empty_update * gate, - ) - - # Only [1, ..., 1, D] row-broadcast gates are supported; a batched - # [B>1, 1, D] gate is not row-broadcast here and must fall back. - batched_residual = torch.randn((2, 8, 64), device="cuda", dtype=torch.bfloat16) - batched_update = torch.randn_like(batched_residual) - batched_gate = torch.randn((2, 1, 64), device="cuda", dtype=torch.bfloat16) - assert not can_use_residual_gate_add_cuda( - batched_residual, batched_update, batched_gate - ) - assert torch.equal( - residual_gate_add(batched_residual, batched_update, batched_gate), - batched_residual + batched_update * batched_gate, - ) - - -def test_residual_gate_add_custom_op_torch_compile_fullgraph(): - residual = torch.randn((1, 32, 128), device="cuda", dtype=torch.bfloat16) - update = torch.randn_like(residual) - gate = torch.randn((1, 1, 128), device="cuda", dtype=torch.bfloat16) - - def fn(residual, update, gate): - return residual_gate_add(residual, update, gate) - - compiled = torch.compile(fn, fullgraph=True) - out = compiled(residual, update, gate) - ref = residual + update * gate - _assert_matches_torch(out, ref) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/kernels/ops/diffusion/test_qknorm_rope.py b/test/registered/kernels/ops/diffusion/test_rope.py similarity index 66% rename from test/registered/kernels/ops/diffusion/test_qknorm_rope.py rename to test/registered/kernels/ops/diffusion/test_rope.py index 867d23c7847f..b3537639b018 100644 --- a/test/registered/kernels/ops/diffusion/test_qknorm_rope.py +++ b/test/registered/kernels/ops/diffusion/test_rope.py @@ -1,23 +1,50 @@ +"""``diffusion.rope``: rotary embeddings and the QK-norm chains fused into them. + +Two families with different oracles: + +- ``fused_inplace_qknorm_rope`` / ``fused_qknorm_rope_pack_kv`` are compared + against the *split* baseline (a separate qknorm kernel plus FlashInfer or + sgl_kernel RoPE). In the default mode the two differ by about one bf16 + rounding step, so those cases use a tolerance; with + ``round_norm_before_rope=True`` the fused kernel reproduces the split + rounding exactly and ``torch.equal`` applies. +- The LTX-2 split-RoPE CUDA kernel is validated on B200 only (guarded per + test, not per module, so the rest of this file still runs elsewhere). +""" + import itertools import sys import pytest import torch +import torch.nn.functional as F import triton from sglang.kernels.jit.utils import get_ci_test_range +from sglang.kernels.ops.diffusion import ( + can_use_fused_inplace_qknorm_rope, + can_use_ltx2_qknorm_split_rope_cuda, + fused_inplace_qknorm_rope, + fused_qknorm_rope_pack_kv, + ltx2_qknorm_split_rope_cuda, +) from sglang.test.ci.ci_register import register_cuda_ci -register_cuda_ci(est_time=44, stage="base-b-kernel-unit", runner_config="1-gpu-large") -# Nightly is not redundant here: it sets SGLANG_JIT_KERNEL_RUN_FULL_TESTS=1 to expand get_ci_test_range sweeps. +register_cuda_ci(est_time=52, stage="base-b-kernel-unit", runner_config="1-gpu-large") +register_cuda_ci(est_time=45, stage="base-b-kernel-unit", runner_config="4-gpu-b200") +# Nightly is not redundant: it sets SGLANG_JIT_KERNEL_RUN_FULL_TESTS=1, which +# expands the get_ci_test_range sweeps below. register_cuda_ci(est_time=220, stage="nightly", runner_config="1-gpu-large") +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") + DEVICE = "cuda" DTYPE = torch.bfloat16 MAX_SEQ_LEN = 131072 ROPE_BASE = 10000.0 ATOL = 8e-2 RTOL = 1e-2 +BF16_FUSED_ATOL = 1.6e-1 def create_cos_sin_cache( @@ -70,8 +97,6 @@ def fused_qknorm_rope( positions: torch.Tensor, is_neox: bool, ) -> None: - from sglang.kernels.ops.diffusion.qknorm_rope import fused_inplace_qknorm_rope - fused_inplace_qknorm_rope( q, k, @@ -85,10 +110,6 @@ def fused_qknorm_rope( def test_qknorm_rope_rejects_unsupported_dtypes() -> None: - from sglang.kernels.ops.diffusion.qknorm_rope import ( - can_use_fused_inplace_qknorm_rope, - ) - assert not can_use_fused_inplace_qknorm_rope(128, 128, False, torch.float32) assert not can_use_fused_inplace_qknorm_rope( 128, 128, False, torch.bfloat16, torch.float64 @@ -164,9 +185,6 @@ def test_qknorm_rope( def test_qknorm_rope_preserves_split_bf16_rounding() -> None: from sgl_kernel import rotary_embedding - from sglang.kernels.ops.diffusion.qknorm_rope import ( - fused_inplace_qknorm_rope, - ) from sglang.kernels.ops.layernorm.norm import fused_inplace_qknorm num_tokens, num_heads, head_dim, rope_dim = 257, 28, 128, 96 @@ -217,7 +235,6 @@ def test_qknorm_rope_preserves_split_bf16_rounding() -> None: def test_qknorm_rope_preserves_full_width_neox_cache() -> None: - from sglang.kernels.ops.diffusion.qknorm_rope import fused_inplace_qknorm_rope from sglang.kernels.ops.layernorm.norm import fused_inplace_qknorm num_tokens, num_heads, head_dim = 257, 32, 128 @@ -256,9 +273,6 @@ def test_qknorm_rope_preserves_full_width_neox_cache() -> None: def test_qknorm_rope_requires_opt_in_for_strided_packed_gqa() -> None: - from sglang.kernels.ops.diffusion.qknorm_rope import ( - fused_inplace_qknorm_rope, - ) from sglang.multimodal_gen.runtime.layers.layernorm import ( RMSNorm, apply_qk_norm_rope, @@ -329,10 +343,6 @@ def test_qknorm_rope_requires_opt_in_for_strided_packed_gqa() -> None: def test_qknorm_rope_pack_kv_matches_separate_ops() -> None: - from sglang.kernels.ops.diffusion.qknorm_rope import ( - fused_inplace_qknorm_rope, - fused_qknorm_rope_pack_kv, - ) batch_size = 2 prefix_tokens, suffix_tokens = 17, 257 @@ -416,9 +426,6 @@ def test_qknorm_rope_pack_kv_matches_separate_ops() -> None: def test_qknorm_rope_pack_kv_preserves_split_bf16_rounding() -> None: from sgl_kernel import rotary_embedding - from sglang.kernels.ops.diffusion.qknorm_rope import ( - fused_qknorm_rope_pack_kv, - ) from sglang.kernels.ops.layernorm.norm import fused_inplace_qknorm batch_size = 1 @@ -507,7 +514,6 @@ def test_qknorm_rope_pack_kv_preserves_split_bf16_rounding() -> None: def test_qknorm_rope_accepts_empty_token_dimension() -> None: - from sglang.kernels.ops.diffusion.qknorm_rope import fused_inplace_qknorm_rope num_heads, head_dim = 8, 128 q = torch.empty(0, num_heads, head_dim, device=DEVICE, dtype=DTYPE) @@ -529,5 +535,214 @@ def test_qknorm_rope_accepts_empty_token_dimension() -> None: assert q.numel() == k.numel() == 0 +# ------------------------------------------------------------------------- +# LTX-2 split RoPE (B200) +# ------------------------------------------------------------------------- + + +def _require_b200() -> None: + if not torch.cuda.is_available(): + pytest.skip("CUDA required") + if torch.cuda.get_device_capability()[0] < 10: + pytest.skip("LTX2 QKNorm split-RoPE CUDA path is validated on B200") + + +def _ltx2_make_cos_sin( + batch: int, seq_len: int, num_heads: int, head_dim: int +) -> tuple[torch.Tensor, torch.Tensor]: + half_dim = head_dim // 2 + cos = torch.randn( + batch, seq_len, num_heads, half_dim, device="cuda", dtype=torch.bfloat16 + ).transpose(1, 2) + sin = torch.randn( + batch, seq_len, num_heads, half_dim, device="cuda", dtype=torch.bfloat16 + ).transpose(1, 2) + return cos, sin + + +def _apply_split_rotary_ref( + x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor +) -> torch.Tensor: + x_dtype = x.dtype + batch = x.shape[0] + _, num_heads, seq_len, _ = cos.shape + x = x.reshape(batch, seq_len, num_heads, -1).swapaxes(1, 2) + last = x.shape[-1] + half = last // 2 + + split_x = x.reshape(*x.shape[:-1], 2, half) + first_x = split_x[..., :1, :] + second_x = split_x[..., 1:, :] + cos_u = cos.unsqueeze(-2) + sin_u = sin.unsqueeze(-2) + + out = split_x * cos_u + out[..., :1, :].addcmul_(-sin_u, second_x) + out[..., 1:, :].addcmul_(sin_u, first_x) + out = out.reshape(*out.shape[:-2], last) + return out.swapaxes(1, 2).reshape(batch, seq_len, -1).to(dtype=x_dtype) + + +def _ltx2_reference( + q: torch.Tensor, + k: torch.Tensor, + q_cos: torch.Tensor, + q_sin: torch.Tensor, + k_cos: torch.Tensor, + k_sin: torch.Tensor, + q_weight: torch.Tensor, + k_weight: torch.Tensor, + eps: float, +) -> tuple[torch.Tensor, torch.Tensor]: + # rms_norm isn't autocast fp32-preserving, so feed fp32 inputs directly + # to keep the normalized value unrounded until the final RoPE output. + q_norm = F.rms_norm(q.float(), (q.shape[-1],), q_weight.float(), eps) + k_norm = F.rms_norm(k.float(), (k.shape[-1],), k_weight.float(), eps) + q_ref = _apply_split_rotary_ref(q_norm, q_cos, q_sin) + k_ref = _apply_split_rotary_ref(k_norm, k_cos, k_sin) + return q_ref.to(dtype=torch.bfloat16), k_ref.to(dtype=torch.bfloat16) + + +@pytest.mark.parametrize( + "batch,q_seq,k_seq,num_heads,head_dim", + [ + (1, 3, 3, 32, 128), + (1, 5, 2, 32, 64), + (2, 4, 3, 32, 64), + ], +) +def test_ltx2_qknorm_split_rope_matches_torch_exactly( + batch: int, q_seq: int, k_seq: int, num_heads: int, head_dim: int +) -> None: + _require_b200() + torch.cuda.manual_seed(20260630) + hidden = num_heads * head_dim + eps = 1e-6 + q = torch.randn(batch, q_seq, hidden, device="cuda", dtype=torch.bfloat16) + k = torch.randn(batch, k_seq, hidden, device="cuda", dtype=torch.bfloat16) + q_cos, q_sin = _ltx2_make_cos_sin(batch, q_seq, num_heads, head_dim) + k_cos, k_sin = _ltx2_make_cos_sin(batch, k_seq, num_heads, head_dim) + q_weight = torch.randn(hidden, device="cuda", dtype=torch.bfloat16) + k_weight = torch.randn(hidden, device="cuda", dtype=torch.bfloat16) + + assert can_use_ltx2_qknorm_split_rope_cuda( + q, + q_cos, + q_sin, + q_weight, + k, + k_cos, + k_sin, + k_weight, + num_heads=num_heads, + head_dim=head_dim, + ) + + q_ref, k_ref = _ltx2_reference( + q, k, q_cos, q_sin, k_cos, k_sin, q_weight, k_weight, eps + ) + q_out, k_out = ltx2_qknorm_split_rope_cuda( + q, + q_cos, + q_sin, + q_weight, + k, + k_cos, + k_sin, + k_weight, + eps=eps, + num_heads=num_heads, + head_dim=head_dim, + ) + torch.cuda.synchronize() + + torch.testing.assert_close(q_out, q_ref, rtol=0, atol=BF16_FUSED_ATOL) + torch.testing.assert_close(k_out, k_ref, rtol=0, atol=BF16_FUSED_ATOL) + + +def test_ltx2_qknorm_split_rope_rejects_unsupported_inputs() -> None: + _require_b200() + torch.cuda.manual_seed(20260630) + q = torch.randn((1, 3, 4096), device="cuda", dtype=torch.bfloat16) + k = torch.randn_like(q) + q_cos, q_sin = _ltx2_make_cos_sin(1, 3, 32, 128) + q_weight = torch.randn(4096, device="cuda", dtype=torch.bfloat16) + k_weight = torch.randn(4096, device="cuda", dtype=torch.bfloat16) + + assert can_use_ltx2_qknorm_split_rope_cuda( + q, + q_cos, + q_sin, + q_weight, + k, + q_cos, + q_sin, + k_weight, + num_heads=32, + head_dim=128, + ) + assert not can_use_ltx2_qknorm_split_rope_cuda( + q.float(), + q_cos, + q_sin, + q_weight, + k, + q_cos, + q_sin, + k_weight, + num_heads=32, + head_dim=128, + ) + assert not can_use_ltx2_qknorm_split_rope_cuda( + q, + q_cos, + q_sin, + q_weight, + k, + q_cos.transpose(-1, -2), + q_sin, + k_weight, + num_heads=32, + head_dim=128, + ) + + +def test_ltx2_qknorm_split_rope_custom_op_torch_compile_fullgraph() -> None: + _require_b200() + torch.cuda.manual_seed(20260630) + batch, q_seq, k_seq, num_heads, head_dim = 1, 3, 2, 32, 64 + hidden = num_heads * head_dim + q = torch.randn(batch, q_seq, hidden, device="cuda", dtype=torch.bfloat16) + k = torch.randn(batch, k_seq, hidden, device="cuda", dtype=torch.bfloat16) + q_cos, q_sin = _ltx2_make_cos_sin(batch, q_seq, num_heads, head_dim) + k_cos, k_sin = _ltx2_make_cos_sin(batch, k_seq, num_heads, head_dim) + q_weight = torch.randn(hidden, device="cuda", dtype=torch.bfloat16) + k_weight = torch.randn(hidden, device="cuda", dtype=torch.bfloat16) + + def fn(q, k, q_cos, q_sin, k_cos, k_sin, q_weight, k_weight): + return ltx2_qknorm_split_rope_cuda( + q, + q_cos, + q_sin, + q_weight, + k, + k_cos, + k_sin, + k_weight, + eps=1e-6, + num_heads=num_heads, + head_dim=head_dim, + ) + + compiled = torch.compile(fn, fullgraph=True) + q_out, k_out = compiled(q, k, q_cos, q_sin, k_cos, k_sin, q_weight, k_weight) + q_ref, k_ref = _ltx2_reference( + q, k, q_cos, q_sin, k_cos, k_sin, q_weight, k_weight, 1e-6 + ) + torch.cuda.synchronize() + torch.testing.assert_close(q_out, q_ref, rtol=0, atol=BF16_FUSED_ATOL) + torch.testing.assert_close(k_out, k_ref, rtol=0, atol=BF16_FUSED_ATOL) + + if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v", "-s"])) + sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernels/ops/diffusion/test_sana_conv_post.py b/test/registered/kernels/ops/diffusion/test_sana_conv_post.py deleted file mode 100644 index 89503bf21f98..000000000000 --- a/test/registered/kernels/ops/diffusion/test_sana_conv_post.py +++ /dev/null @@ -1,60 +0,0 @@ -"""Sana GLUMB post-processing fusions must match the eager bf16 chain.""" - -import pytest -import torch -import torch.nn.functional as F - -from sglang.kernels.ops.diffusion.triton.sana_conv_post import ( - can_use_fused_bias_glu, - can_use_fused_bias_silu, - fused_bias_glu, - fused_bias_silu, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=3, stage="base-b-kernel-unit", runner_config="1-gpu-large") -pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") - - -@pytest.mark.parametrize("channels", [2240, 11200]) -def test_sana_bias_silu_is_bit_exact(channels): - torch.manual_seed(0) - x = torch.randn( - (1, channels, 7, 5), - device="cuda", - dtype=torch.bfloat16, - ).to(memory_format=torch.channels_last) - bias = torch.randn(channels, device="cuda", dtype=torch.bfloat16) - - assert can_use_fused_bias_silu(x, bias) - actual = fused_bias_silu(x, bias) - expected = F.silu(x + bias[None, :, None, None]) - - assert actual.is_contiguous(memory_format=torch.channels_last) - assert torch.equal(actual, expected) - - -@pytest.mark.parametrize("channels", [2240, 5600]) -def test_sana_bias_glu_is_bit_exact(channels): - torch.manual_seed(1) - x = torch.randn( - (1, 2 * channels, 7, 5), - device="cuda", - dtype=torch.bfloat16, - ).to(memory_format=torch.channels_last) - bias = torch.randn(2 * channels, device="cuda", dtype=torch.bfloat16) - - assert can_use_fused_bias_glu(x, bias) - actual = fused_bias_glu(x, bias) - biased = x + bias[None, :, None, None] - hidden, gate = torch.chunk(biased, 2, dim=1) - expected = hidden * F.silu(gate) - - assert actual.is_contiguous(memory_format=torch.channels_last) - assert torch.equal(actual, expected) - - -if __name__ == "__main__": - import sys - - sys.exit(pytest.main([__file__])) diff --git a/test/registered/kernels/ops/diffusion/test_sana_ln_modulate.py b/test/registered/kernels/ops/diffusion/test_sana_ln_modulate.py deleted file mode 100644 index 36405ac8cbf2..000000000000 --- a/test/registered/kernels/ops/diffusion/test_sana_ln_modulate.py +++ /dev/null @@ -1,55 +0,0 @@ -"""Sana fused LN+modulate fast path must stay bit-exact vs eager.""" - -import pytest -import torch - -import sglang.multimodal_gen.runtime.models.dits.sana as sana -from sglang.multimodal_gen.runtime.models.dits.sana import ( - _eager_ln_modulate, - _sana_ln_modulate, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=3, stage="base-b-kernel-unit", runner_config="1-gpu-large") -pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") - - -@pytest.mark.parametrize( - "shape,nmod,transposed", - [ - ((2, 1024, 2240), 6, False), - ((2, 1024, 2240), 2, False), - ((1, 1024, 2240), 6, True), - ((1, 37, 2240), 6, False), - ], -) -def test_sana_fused_ln_modulate_is_bit_exact(shape, nmod, transposed): - # (., 1024, 2240) is the real Sana 1024px shape; hidden 2240 % 512 != 0 - # exercises the kernel's partial tail chunk. nmod mirrors the two adaLN - # chunk layouts, transposed the permuted layout the Sana DiT serves. - torch.manual_seed(0) - batch, seq, hidden = shape - norm = torch.nn.LayerNorm(hidden, eps=1e-6, elementwise_affine=False).cuda() - x = (torch.randn(batch, seq, hidden, device="cuda") * 4).bfloat16() - if transposed: - x = x.permute(0, 2, 1).contiguous().permute(0, 2, 1) - emb = torch.randn(batch, nmod, hidden, device="cuda").bfloat16() - shift, scale = emb.chunk(nmod, dim=1)[0], emb.chunk(nmod, dim=1)[-1] - # default-stream eager serving must stay on the untouched eager chain - n_sigs = len(sana._SANA_LN_MOD.verified_sigs) - _sana_ln_modulate(norm, x, scale, shift) - assert len(sana._SANA_LN_MOD.verified_sigs) == n_sigs - # the fusion engages on non-default streams (the BCG warmup/capture path) - with torch.cuda.stream(torch.cuda.Stream()): - out = _sana_ln_modulate(norm, x, scale, shift) - assert len(sana._SANA_LN_MOD.verified_sigs) == n_sigs + 1 # verified - out2 = _sana_ln_modulate(norm, x, scale, shift) # verified-sig lane - torch.cuda.synchronize() - assert torch.equal(out, _eager_ln_modulate(norm, x, scale, shift)) - assert torch.equal(out2, out) and not sana._SANA_LN_MOD.disabled - - -if __name__ == "__main__": - import sys - - sys.exit(pytest.main([__file__])) diff --git a/test/registered/kernels/ops/diffusion/test_scale_shift.py b/test/registered/kernels/ops/diffusion/test_scale_shift.py deleted file mode 100644 index aff10d18138c..000000000000 --- a/test/registered/kernels/ops/diffusion/test_scale_shift.py +++ /dev/null @@ -1,40 +0,0 @@ -import sys - -import pytest -import torch - -from sglang.kernels.ops.diffusion.triton.scale_shift import ( - try_fused_scaled_residual_add_exact, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=5, stage="base-b-kernel-unit", runner_config="1-gpu-large") -pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") - - -@torch.no_grad() -@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) -def test_scaled_residual_add_is_bit_exact(dtype): - torch.manual_seed(0) - residual = torch.randn(2, 17, 64, device="cuda", dtype=torch.float32) - x = torch.randn(2, 17, 64, device="cuda", dtype=dtype) - scale = torch.randn(64, device="cuda", dtype=torch.float32) - - actual = try_fused_scaled_residual_add_exact(residual, x, scale) - expected = residual + x * scale - assert actual is not None - assert torch.equal(actual, expected) - - -@torch.no_grad() -def test_scaled_residual_add_rejects_unsupported_inputs(): - residual = torch.empty(2, 3, 8, device="cuda", dtype=torch.float32) - x = torch.empty_like(residual) - scale = torch.empty(8, device="cuda", dtype=torch.float32) - - assert try_fused_scaled_residual_add_exact(residual, x, scale) is None - assert try_fused_scaled_residual_add_exact(residual, x.half(), scale[:-1]) is None - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/kernels/ops/diffusion/test_sites.py b/test/registered/kernels/ops/diffusion/test_sites.py new file mode 100644 index 000000000000..75973a0fe2b7 --- /dev/null +++ b/test/registered/kernels/ops/diffusion/test_sites.py @@ -0,0 +1,416 @@ +"""``diffusion.sites``: the two gate protocols and the mount lifecycle. + +Nothing here is a kernel. ``sites`` decides *whether* a fusion is allowed to +run for a given request and model: + +- :class:`QualityGatedFusion` -- for fusions that are **not** bit-exact. They + mount onto marked ``nn.Module`` sites only for ``quality="high"``, at batch + boundaries, all-or-nothing per transformer. +- :class:`BitExactFusionGate` -- for fusions that **are** bit-exact. They run + by default but verify themselves against the live eager chain on first + sight and disable permanently on any mismatch. + +The protocol tests are pure-CPU. The mount-lifecycle tests below use +synthetic sites (a bare ``nn.Module`` with the marker attribute) so they test +the protocol rather than any one model; the real model wrappers live in +``test_model_fast_paths.py``. +""" + +import sys +from types import ModuleType +from unittest.mock import MagicMock, patch + +import pytest +import torch +import torch.nn as nn +import torch.nn.functional as F + +import sglang.kernels.ops.diffusion.sites.fused_gate_rmsnorm_site as gate_rmsnorm +import sglang.kernels.ops.diffusion.sites.fused_linear_gelu_site as linear_gelu +from sglang.kernels.ops.diffusion import ( + BitExactFusionGate, + QualityGatedFusion, + can_use_ln_modulate, + flashinfer_rmsnorm_diagnostic_hint, + fused_ln_modulate, + fused_ln_modulate_active, + mark_fused_ln_modulate_site, + mount_fused_ln_modulate, + tensors_equal, + unmount_fused_ln_modulate, +) +from sglang.test.ci.ci_register import register_cpu_ci, register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=3, suite="base-a-test-cpu") +register_cuda_ci(est_time=38, stage="base-b-kernel-unit", runner_config="1-gpu-large") + +requires_cuda = pytest.mark.skipif( + not torch.cuda.is_available(), reason="CUDA required" +) + + +# --------------------------------------------------------------------------- +# QualityGatedFusion protocol (CPU) +# --------------------------------------------------------------------------- + + +def test_quality_gate_mounts_and_unmounts_all_sites(): + fusion = QualityGatedFusion( + name="test fusion", + marker_attr="_test_fusion_site", + enabled_attr="_test_fusion_enabled", + ) + root = nn.ModuleList([nn.Module(), nn.Module()]) + for index, site in enumerate(root): + fusion.mark(site, index) + + assert [fusion.metadata(site) for site in root] == [0, 1] + assert fusion.mount(root) + assert all(fusion.is_enabled(site) for site in root) + fusion.unmount(root) + assert not any(fusion.is_enabled(site) for site in root) + + +def test_quality_gate_rejection_is_all_or_nothing(): + fusion = QualityGatedFusion( + name="test fusion", + marker_attr="_test_fusion_site", + enabled_attr="_test_fusion_enabled", + ) + root = nn.ModuleList([nn.Module(), nn.Module()]) + for index, site in enumerate(root): + fusion.mark(site, index) + + assert not fusion.mount( + root, reject_reason=lambda site: "rejected" if fusion.metadata(site) else None + ) + assert not any(fusion.is_enabled(site) for site in root) + assert not fusion.mount(nn.Module()) + + +# ------------------------------------------------------------------------- +# BitExactFusionGate protocol (CPU) +# ------------------------------------------------------------------------- + + +def test_bitexact_gate_once_mode_verifies_then_reuses(): + gate = BitExactFusionGate("once") + calls = {"fused": 0, "ref": 0} + + def fused(): + calls["fused"] += 1 + return torch.tensor([1.0]) + + def ref(): + calls["ref"] += 1 + return torch.tensor([1.0]) + + assert torch.equal(gate.accept_or_fallback(fused(), ref()), torch.tensor([1.0])) + assert gate.verified and not gate.disabled and calls == {"fused": 1, "ref": 1} + assert torch.equal(fused(), torch.tensor([1.0])) + assert calls == {"fused": 2, "ref": 1} + + +def test_bitexact_gate_mismatch_disables_permanently(): + gate = BitExactFusionGate("mismatch") + + out = gate.accept_or_fallback( + torch.tensor([1.0]), + torch.tensor([2.0]), + mismatch_msg="mismatch", + ) + assert torch.equal(out, torch.tensor([2.0])) + assert gate.disabled and not gate.verified + + +def test_bitexact_gate_per_signature_tracks_each_sig(): + gate = BitExactFusionGate("sig", per_signature=True) + a = torch.tensor([1.0]) + assert torch.equal(gate.accept_or_fallback(a, a, sig=("a",)), a) + assert gate.is_verified(("a",)) + assert not gate.is_verified(("b",)) + assert torch.equal(gate.accept_or_fallback(a, a, sig=("b",)), a) + assert gate.verified_sigs == {("a",), ("b",)} + + +def test_bitexact_gate_skips_first_sight_during_graph_capture(monkeypatch): + # Negative-branch contract: an unverified gate must not attempt first-sight + # verification inside CUDA graph capture — the eager-reference host sync + # would abort the capture (and BCG would permanently block the signature). + gate = BitExactFusionGate("capture") + monkeypatch.setattr(torch.cuda, "is_available", lambda: True) + monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: True) + assert not gate.can_attempt_once() + # A verified gate replays the fused kernel alone, which is capture-safe. + gate.mark_verified() + assert gate.can_attempt_once() + + +def test_tensors_equal_supports_sequences(): + assert tensors_equal( + (torch.tensor([1.0]), torch.tensor([2.0])), + (torch.tensor([1.0]), torch.tensor([2.0])), + ) + assert not tensors_equal( + (torch.tensor([1.0]), torch.tensor([2.0])), + (torch.tensor([1.0]), torch.tensor([3.0])), + ) + + +class TestBitExactFallbackDiagnostics(CustomTestCase): + def test_mismatch_warning_is_actionable_and_diagnostic_is_lazy(self): + logger = MagicMock() + diagnostic = MagicMock(return_value="backend=CuTe DSL") + gate = BitExactFusionGate("diagnostic") + + matched = gate.accept_or_fallback( + torch.tensor([1.0]), + torch.tensor([1.0]), + logger=logger, + diagnostic_hint=diagnostic, + ) + self.assertTrue(torch.equal(matched, torch.tensor([1.0]))) + diagnostic.assert_not_called() + logger.warning_once.assert_not_called() + + gate = BitExactFusionGate("diagnostic") + fallback = gate.accept_or_fallback( + torch.tensor([1.0]), + torch.tensor([2.0]), + logger=logger, + diagnostic_hint=diagnostic, + ) + + self.assertTrue(torch.equal(fallback, torch.tensor([2.0]))) + diagnostic.assert_called_once_with() + warning = logger.warning_once.call_args.args[0] + self.assertIn("Correctness is preserved", warning) + self.assertIn("reference kernel or reduction-order change", warning) + self.assertIn("backend=CuTe DSL", warning) + + def test_diagnostic_failure_cannot_break_the_eager_fallback(self): + logger = MagicMock() + + def broken_diagnostic(): + raise RuntimeError("diagnostics unavailable") + + gate = BitExactFusionGate("diagnostic") + fallback = gate.accept_or_fallback( + torch.tensor([1.0]), + torch.tensor([2.0]), + logger=logger, + diagnostic_hint=broken_diagnostic, + ) + + self.assertTrue(torch.equal(fallback, torch.tensor([2.0]))) + self.assertTrue(gate.disabled) + self.assertIn("Correctness is preserved", logger.warning_once.call_args.args[0]) + + def test_flashinfer_rmsnorm_hint_reports_backend_and_versions(self): + flashinfer = ModuleType("flashinfer") + flashinfer_norm = ModuleType("flashinfer.norm") + flashinfer_norm._USE_CUDA_NORM = False + versions = { + "flashinfer-python": "0.6.12", + "flashinfer-cubin": "0.6.12", + "flashinfer-jit-cache": "0.6.12+cu130", + } + + with ( + patch.dict( + sys.modules, + {"flashinfer": flashinfer, "flashinfer.norm": flashinfer_norm}, + ), + patch("importlib.metadata.version", side_effect=versions.__getitem__), + patch.dict("os.environ", {"FLASHINFER_USE_CUDA_NORM": "0"}), + ): + hint = flashinfer_rmsnorm_diagnostic_hint() + + self.assertIn("backend=CuTe DSL", hint) + self.assertIn("FLASHINFER_USE_CUDA_NORM=0", hint) + for package, version in versions.items(): + self.assertIn(f"{package}={version}", hint) + + +# --------------------------------------------------------------------------- +# Mount lifecycle on synthetic sites (CUDA) +# --------------------------------------------------------------------------- + +LN_DIM = 3072 # FLUX.1 hidden size + + +@requires_cuda +@pytest.mark.parametrize("seq_len", [4096, 512]) +def test_fused_ln_modulate_matches_reference(seq_len): + torch.cuda.manual_seed(0) + x = torch.randn((1, seq_len, LN_DIM), device="cuda", dtype=torch.bfloat16) + scale = torch.randn((1, LN_DIM), device="cuda", dtype=torch.bfloat16) + shift = torch.randn_like(scale) + assert can_use_ln_modulate(x, scale, shift) + + norm = nn.LayerNorm(LN_DIM, eps=1e-6, elementwise_affine=False).cuda() + ref = norm(x) * (1 + scale[:, None]) + shift[:, None] + # Contract: bf16 rounding-order-level difference only, not bit-exact -- + # which is exactly why this fusion is quality-gated. + torch.testing.assert_close( + fused_ln_modulate(x, scale, shift, eps=1e-6), ref, atol=0.0625, rtol=0.05 + ) + + +@requires_cuda +def test_fused_ln_modulate_guards_and_mount_protocol(): + x = torch.randn((2, 64, LN_DIM), device="cuda", dtype=torch.bfloat16) + row = torch.randn((2, LN_DIM), device="cuda", dtype=torch.bfloat16) + assert not can_use_ln_modulate(x, row, row) # folded affine needs B == 1 + + root = nn.Module() + root.child = nn.Module() + mark_fused_ln_modulate_site(root.child) + assert not fused_ln_modulate_active(root.child) + assert mount_fused_ln_modulate(root) + assert fused_ln_modulate_active(root.child) + unmount_fused_ln_modulate(root) + assert not fused_ln_modulate_active(root.child) + assert not mount_fused_ln_modulate(nn.Module()) # no marked sites + + +@requires_cuda +@torch.no_grad() +def test_mounted_ln_modulate_site_compiles_fullgraph(): + class Site(nn.Module): + def __init__(self): + super().__init__() + mark_fused_ln_modulate_site(self) + + def forward(self, x, scale, shift): + if fused_ln_modulate_active(self) and can_use_ln_modulate(x, scale, shift): + return fused_ln_modulate(x, scale, shift, eps=1e-6) + return ( + F.layer_norm(x, (x.shape[-1],), eps=1e-6) * (1 + scale[:, None]) + + shift[:, None] + ) + + site = Site() + assert mount_fused_ln_modulate(site) + x = torch.randn(1, 64, 128, device="cuda", dtype=torch.bfloat16) + scale = torch.randn(1, 128, device="cuda", dtype=torch.bfloat16) + shift = torch.randn_like(scale) + expected = site(x, scale, shift) + torch.testing.assert_close( + torch.compile(site, fullgraph=True)(x, scale, shift), + expected, + atol=0.0625, + rtol=0.05, + ) + + +GATE_RMSNORM_DIM, GATE_RMSNORM_EPS = 4608, 1e-5 # Ideogram 4 hidden / norm_eps + + +class _GateRMSNormSite(nn.Module): + def __init__(self, dtype=torch.bfloat16): + super().__init__() + self.norm = nn.RMSNorm( + GATE_RMSNORM_DIM, eps=GATE_RMSNORM_EPS, device="cuda", dtype=dtype + ) + gate_rmsnorm.mark_fused_gate_rmsnorm_site(self, ("norm",)) + + +@requires_cuda +def test_fused_gate_rmsnorm_matches_ideogram_reference(): + torch.manual_seed(0) + site = _GateRMSNormSite() + w = site.norm.weight.data + dim = GATE_RMSNORM_DIM + x = torch.randn(1, 64, dim, device="cuda", dtype=torch.bfloat16) + residual = torch.randn_like(x) + # adaln-style strided chunks, as produced by Ideogram's modulation .chunk() + mods = torch.randn(1, 1, 2 * dim, device="cuda", dtype=torch.bfloat16) + scale, gate = mods.chunk(2, dim=-1) + + assert gate_rmsnorm.mount_fused_gate_rmsnorm(site) + got_scale = gate_rmsnorm.fused_rmsnorm_scale(x, w, 1.0 + scale, GATE_RMSNORM_EPS) + got_gate = gate_rmsnorm.fused_rmsnorm_tanh_residual( + x, gate, residual, w, GATE_RMSNORM_EPS + ) + norm = F.rms_norm(x, (dim,), w, GATE_RMSNORM_EPS) + # The fused path uses bf16-native norm statistics: close, not bit-exact. + torch.testing.assert_close(got_scale, norm * (1.0 + scale), atol=8e-2, rtol=4e-2) + torch.testing.assert_close( + got_gate, residual + torch.tanh(gate) * norm, atol=8e-2, rtol=4e-2 + ) + + +@requires_cuda +def test_fused_gate_rmsnorm_mount_is_all_or_nothing(): + good, bad = _GateRMSNormSite(), _GateRMSNormSite(torch.float32) + # One fp32 norm anywhere in the tree keeps *every* site on the reference. + assert not gate_rmsnorm.mount_fused_gate_rmsnorm(nn.ModuleList([good, bad])) + assert not gate_rmsnorm.fused_gate_rmsnorm_active(good) + assert gate_rmsnorm.mount_fused_gate_rmsnorm(good) + gate_rmsnorm.unmount_fused_gate_rmsnorm(good) + assert not gate_rmsnorm.fused_gate_rmsnorm_active(good) + + +class _GeluSite(nn.Module): + def __init__(self, dtype=torch.bfloat16, bias=True): + super().__init__() + self.proj = nn.Linear(64, 256, bias=bias, device="cuda", dtype=dtype) + linear_gelu.mark_fused_gelu_site(self, "proj") + + def forward(self, x): + if linear_gelu.fused_gelu_active(self) and linear_gelu.can_use_linear_gelu( + self.proj, x + ): + return linear_gelu.fused_linear_gelu_tanh( + x, self.proj.weight, self.proj.bias + ) + return F.gelu(self.proj(x), approximate="tanh") + + +@requires_cuda +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +def test_fused_linear_gelu_matches_reference(dtype): + torch.manual_seed(0) + site = _GeluSite(dtype) + x = torch.randn(512, 64, device="cuda", dtype=dtype) + ref = site(x) # unmounted: the reference chain + assert linear_gelu.mount_fused_linear_gelu(site) + atol = 2e-2 if dtype == torch.bfloat16 else 4e-3 + torch.testing.assert_close(site(x), ref, atol=atol, rtol=2e-2) + + +@requires_cuda +def test_fused_linear_gelu_guards_and_lossless_path(): + torch.manual_seed(0) + good, bad = _GeluSite(), _GeluSite(torch.float32) + assert not linear_gelu.mount_fused_linear_gelu(nn.ModuleList([good, bad])) + assert not linear_gelu.fused_gelu_active(good) + + x = torch.randn(16, 64, device="cuda", dtype=torch.bfloat16) + ref = good(x) + assert linear_gelu.mount_fused_linear_gelu(good) + linear_gelu.unmount_fused_linear_gelu(good) + # Unmounting must restore the reference path bit-for-bit. + assert torch.equal(good(x), ref) + + no_bias = nn.Linear(8, 8, bias=False, device="cuda", dtype=torch.bfloat16) + assert not linear_gelu.can_use_linear_gelu_static(no_bias) + assert not linear_gelu.can_use_linear_gelu(good.proj, x.float()) + + +@requires_cuda +@torch.no_grad() +def test_mounted_gelu_site_compiles_fullgraph(): + site = _GeluSite() + x = torch.randn(16, 64, device="cuda", dtype=torch.bfloat16) + assert linear_gelu.mount_fused_linear_gelu(site) + expected = site(x) + torch.testing.assert_close( + torch.compile(site, fullgraph=True)(x), expected, atol=2e-2, rtol=2e-2 + ) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernels/ops/diffusion/test_timestep_embedding.py b/test/registered/kernels/ops/diffusion/test_timestep_embedding.py deleted file mode 100644 index d803b336f6b4..000000000000 --- a/test/registered/kernels/ops/diffusion/test_timestep_embedding.py +++ /dev/null @@ -1,184 +0,0 @@ -import os -import sys - -import numpy as np -import pytest -import torch - -try: - import tabulate -except Exception: - tabulate = None - -from sglang.kernels.jit.utils import get_ci_test_range -from sglang.kernels.ops.diffusion.timestep_embedding import ( - timestep_embedding as timestep_embedding_cuda, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=16, stage="base-b-kernel-unit", runner_config="1-gpu-large") -# Nightly is not redundant here: it sets SGLANG_JIT_KERNEL_RUN_FULL_TESTS=1 to expand get_ci_test_range sweeps. -register_cuda_ci(est_time=20, stage="nightly", runner_config="1-gpu-large") - -CORRECTNESS_BATCH_SIZES = get_ci_test_range( - [1, 2, 8, 128, 256, 512, 1536, 2048, 4096, 11008, 16384], - [1, 128, 2048, 16384], -) -CORRECTNESS_DIMS = get_ci_test_range( - [32, 128, 256, 512, 1536, 2048, 4096, 8192], - [32, 512, 8192], -) -DIFFUSERS_BATCH_SIZES = get_ci_test_range( - [1, 2, 8, 128, 256, 512, 1536, 2048, 16384], - [1, 512, 16384], -) -DIFFUSERS_DIMS = get_ci_test_range([32, 256, 512, 1536, 8192], [32, 512, 8192]) -DTYPES = get_ci_test_range( - [torch.float16, torch.bfloat16, torch.float32], - [torch.float16, torch.bfloat16], -) -SCALES = get_ci_test_range([1, 0.01], [1, 0.01]) - - -def get_timestep_embedding_reference( - timesteps: torch.Tensor, - dim: int, - *, - flip_sin_to_cos: bool = False, - downscale_freq_shift: float = 1, - scale: float = 1, - max_period: int = 10000, -): - assert len(timesteps.shape) == 1, "Timesteps should be a 1d-array" - - timesteps = timesteps.to(torch.float32) - half_dim = dim // 2 - exponent = -torch.log( - torch.tensor(max_period, dtype=torch.float32, device=timesteps.device) - ) * torch.arange( - start=0, end=half_dim, dtype=torch.float32, device=timesteps.device - ) - exponent = exponent / (half_dim - downscale_freq_shift) - - emb = torch.exp(exponent) - emb = timesteps[:, None].float() * emb[None, :] - - emb = scale * emb - - emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1) - if flip_sin_to_cos: - emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1) - if dim % 2 == 1: - emb = torch.nn.functional.pad(emb, (0, 1, 0, 0)) - return emb - - -@pytest.mark.parametrize("batch_size", CORRECTNESS_BATCH_SIZES) -@pytest.mark.parametrize("dim", CORRECTNESS_DIMS) -@pytest.mark.parametrize("dtype", DTYPES) -def test_timestep_embedding_correctness_with_sgld(batch_size, dim, dtype): - device = "cuda" - t = torch.randint(low=0, high=1000, size=(batch_size,), device=device).to(dtype) - torch_output = get_timestep_embedding_reference( - t, dim, flip_sin_to_cos=True, downscale_freq_shift=0 - ) - cuda_output = timestep_embedding_cuda( - t, dim, flip_sin_to_cos=True, downscale_freq_shift=0 - ) - torch.testing.assert_close(torch_output, cuda_output, atol=1e-3, rtol=1e-3) - - -@pytest.mark.parametrize("batch_size", DIFFUSERS_BATCH_SIZES) -@pytest.mark.parametrize("dim", DIFFUSERS_DIMS) -@pytest.mark.parametrize("dtype", DTYPES) -@pytest.mark.parametrize("flip_sin_to_cos", [False, True]) -@pytest.mark.parametrize("downscale_freq_shift", [0, 1]) -@pytest.mark.parametrize("scale", SCALES) -def test_timestep_embedding_correctness_with_diffusers( - batch_size, dim, flip_sin_to_cos, downscale_freq_shift, scale, dtype -): - device = "cuda" - t = torch.randint(low=0, high=1000, size=(batch_size,), device=device).to(dtype) - torch_output = get_timestep_embedding_reference( - t, - dim, - flip_sin_to_cos=flip_sin_to_cos, - downscale_freq_shift=downscale_freq_shift, - scale=scale, - max_period=10000, - ) - cuda_output = timestep_embedding_cuda( - t, - dim, - flip_sin_to_cos=flip_sin_to_cos, - downscale_freq_shift=downscale_freq_shift, - scale=scale, - max_period=10000, - ) - torch.testing.assert_close(torch_output, cuda_output, atol=1e-3, rtol=1e-3) - - -def test_timestep_embedding_perf(): - if os.environ.get("SGLANG_RUN_JIT_KERNEL_PERF_TESTS") != "1": - pytest.skip("Perf test disabled by default") - if tabulate is None: - pytest.skip("Optional dependency 'tabulate' is not installed") - - NUM_BATCH = [1, 2, 8, 63, 256, 512, 613, 1024, 1536] - NUM_DIM = [32, 64, 128, 256, 512, 1024, 2048, 4096] - - def perf_kernel_fn(kernel_fn: callable, *args, **kwargs): - warmup_times = 4 - repeat_times = 20 - start = torch.cuda.Event(enable_timing=True) - end = torch.cuda.Event(enable_timing=True) - - for _ in range(warmup_times): - kernel_fn(*args, **kwargs) - torch.cuda.synchronize() - - start.record() - for _ in range(repeat_times): - kernel_fn(*args, **kwargs) - end.record() - end.synchronize() - return start.elapsed_time(end) / repeat_times - - device = "cuda" - results = [] - - cuda_speedups = [] - for B in NUM_BATCH: - for dim in NUM_DIM: - t = torch.linspace(0, max(100000, B), steps=B, device=device).to( - torch.float32 - ) - time_torch = perf_kernel_fn(get_timestep_embedding_reference, t, dim) - time_cuda = perf_kernel_fn(timestep_embedding_cuda, t, dim) - speedup_cuda = time_torch / time_cuda - - results.append( - { - "Batch Size": B, - "Dimension": dim, - "Torch Time (ms)": time_torch, - "CUDA Time (ms)": time_cuda, - "Speedup (CUDA)": speedup_cuda, - } - ) - cuda_speedups.append(speedup_cuda) - - print("=== Timestep Embedding Benchmark Results ===") - print( - tabulate.tabulate( - results, - headers="keys", - tablefmt="fancy_grid", - floatfmt=(".0f", ".0f", ".6f", ".6f", ".5f"), - ) - ) - print(f"Average Speedup(cuda): {np.mean(cuda_speedups):.4f}") - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/kernels/ops/diffusion/test_ulysses_qkv.py b/test/registered/kernels/ops/diffusion/test_ulysses_qkv.py deleted file mode 100644 index 7dec40df232e..000000000000 --- a/test/registered/kernels/ops/diffusion/test_ulysses_qkv.py +++ /dev/null @@ -1,54 +0,0 @@ -import sys - -import pytest -import torch - -from sglang.kernels.ops.diffusion.triton.ulysses_qkv import ( - pack_qkv_destination_major, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=8, stage="base-b-kernel-unit", runner_config="1-gpu-large") -pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") - - -@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) -def test_pack_qkv_destination_major_is_bit_exact(dtype): - torch.manual_seed(0) - rows, world_size, global_heads, head_size = 17, 4, 12, 64 - q, k, v = ( - torch.randn(rows, global_heads, head_size, device="cuda", dtype=dtype) - for _ in range(3) - ) - - local_heads = global_heads // world_size - expected = torch.empty( - world_size, - rows, - local_heads, - 3 * head_size, - device="cuda", - dtype=dtype, - ) - for index, tensor in enumerate((q, k, v)): - shards = tensor.view(rows, world_size, local_heads, head_size).permute( - 1, 0, 2, 3 - ) - expected[..., index * head_size : (index + 1) * head_size].copy_(shards) - - actual = pack_qkv_destination_major(q, k, v, world_size) - assert torch.equal(actual, expected) - - -def test_pack_qkv_destination_major_validates_inputs(): - q = torch.empty(2, 4, 8, device="cuda", dtype=torch.bfloat16) - with pytest.raises(ValueError, match="same 3D shape"): - pack_qkv_destination_major(q, q[:, :-1], q, 2) - with pytest.raises(ValueError, match="divide global_heads"): - pack_qkv_destination_major(q, q, q, 3) - with pytest.raises(ValueError, match="expected shape"): - pack_qkv_destination_major(q, q, q, 2, out=torch.empty_like(q)) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/kernels/ops/diffusion/test_usp_relayout.py b/test/registered/kernels/ops/diffusion/test_usp_relayout.py deleted file mode 100644 index 588524865ee1..000000000000 --- a/test/registered/kernels/ops/diffusion/test_usp_relayout.py +++ /dev/null @@ -1,69 +0,0 @@ -"""Bitwise tests for the generic Ulysses output head-merge fast path.""" - -import sys -from unittest.mock import patch - -import pytest -import torch - -from sglang.kernels.ops.diffusion.usp_relayout import ( - can_use_usp_merge_heads, - usp_merge_heads, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=4, stage="base-b-kernel-unit", runner_config="1-gpu-large") - -DEVICE = "cuda" - - -@pytest.mark.parametrize( - "world,seq,batch,h_local,head_dim", - [ - (4, 7936, 1, 14, 128), # H3 768p production shape (Ulysses 4) - (2, 64, 3, 4, 64), # batched - (4, 33, 2, 4, 100), # scalar fallback inside the CUDA kernel - ], -) -@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16, torch.float32]) -def test_usp_merge_heads_bitwise(dtype, world, seq, batch, h_local, head_dim): - generator = torch.Generator(device=DEVICE).manual_seed(4321) - x = torch.randn( - world, - seq, - batch, - h_local, - head_dim, - dtype=dtype, - device=DEVICE, - generator=generator, - ) - assert can_use_usp_merge_heads(x) - out = usp_merge_heads(x) - ref = x.permute(2, 1, 0, 3, 4).contiguous() - assert out.shape == ref.shape - assert torch.equal(out, ref) - - -def test_usp_merge_heads_unsupported_inputs_use_exact_fallback(): - x = torch.randn(2, 4, 1, 4, 64, dtype=torch.bfloat16, device=DEVICE) - unsupported = [x.transpose(0, 1), x[:0]] - - for value in unsupported: - assert not can_use_usp_merge_heads(value) - assert torch.equal( - usp_merge_heads(value), value.permute(2, 1, 0, 3, 4).contiguous() - ) - - with patch.object(torch.version, "hip", "6.3"): - assert not can_use_usp_merge_heads(x) - assert torch.equal(usp_merge_heads(x), x.permute(2, 1, 0, 3, 4).contiguous()) - - -def test_usp_merge_heads_fast_path_rejects_wrong_rank(): - x = torch.randn(2, 4, 1, 4, 64, dtype=torch.bfloat16, device=DEVICE) - assert not can_use_usp_merge_heads(x[0]) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernels/ops/diffusion/test_varlen_pack_pad.py b/test/registered/kernels/ops/diffusion/test_varlen_pack_pad.py deleted file mode 100644 index dfdea7838ad3..000000000000 --- a/test/registered/kernels/ops/diffusion/test_varlen_pack_pad.py +++ /dev/null @@ -1,197 +0,0 @@ -"""Numerical correctness for fused varlen pack/scatter Triton kernels. - -Bit-exact comparison against the equivalent PyTorch ops (index_select, -zeros + index_copy_) across bf16/fp16 and several shape/mask cases. -""" - -import pytest -import torch - -from sglang.kernels.jit.utils import get_ci_test_range -from sglang.kernels.ops.diffusion.triton.varlen_pack_pad import ( - build_inv_indices, - fused_pack_qkv, - fused_scatter_to_padded, -) -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci - -register_cuda_ci(est_time=10, stage="base-b-kernel-unit", runner_config="1-gpu-large") -# Nightly is not redundant here: it sets SGLANG_JIT_KERNEL_RUN_FULL_TESTS=1 to expand get_ci_test_range sweeps. -register_cuda_ci(est_time=20, stage="nightly", runner_config="1-gpu-large") -register_amd_ci(est_time=15, suite="nightly-amd-kernel-1-gpu", nightly=True) - -DEVICE = "cuda" -DTYPES = get_ci_test_range([torch.bfloat16, torch.float16], [torch.bfloat16]) -# (bs, s_txt, s_img, num_heads, head_dim, valid_txt_lens) tuples -SHAPES = get_ci_test_range( - [ - # name, bs, s_txt, s_img, H, D, valid_txt_lens - ("small_c2", 2, 64, 128, 4, 64, [32, 48]), - ("prod_c2", 2, 256, 1024, 24, 128, [128, 200]), - ("all_valid_b1", 1, 64, 128, 4, 64, [64]), - ("all_valid_b4", 4, 64, 128, 4, 64, [64, 64, 64, 64]), - ("c8_prod", 8, 256, 4096, 24, 128, [128, 200, 256, 100, 50, 256, 256, 50]), - # one batch with zero valid text tokens (image side still valid) - ("zero_txt_one_batch", 2, 64, 128, 4, 64, [0, 32]), - # bs=1 with no text validity (only image rows packed) - ("bs1_zero_txt", 1, 64, 128, 4, 64, [0]), - ], - [ - ("small_c2", 2, 64, 128, 4, 64, [32, 48]), - ("prod_c2", 2, 256, 1024, 24, 128, [128, 200]), - ("all_valid_b4", 4, 64, 128, 4, 64, [64, 64, 64, 64]), - ], -) - - -def _build_mask(bs, s_txt, s_img, valid_txt_lens): - s = s_txt + s_img - mask = torch.zeros(bs, s, dtype=torch.bool, device=DEVICE) - for b, vt in enumerate(valid_txt_lens): - mask[b, :vt] = True - mask[b, s_txt:] = True - return mask - - -def _ref_pack(q, k, v, indices): - bs, seq = q.shape[:2] - - def flat(t): - return t.reshape(bs * seq, *t.shape[2:]) - - return ( - flat(q).index_select(0, indices), - flat(k).index_select(0, indices), - flat(v).index_select(0, indices), - ) - - -def _ref_scatter(out_unpad, indices, bs, seq): - _, num_heads, head_dim = out_unpad.shape - flat = torch.zeros( - bs * seq, num_heads, head_dim, dtype=out_unpad.dtype, device=DEVICE - ) - flat.index_copy_(0, indices, out_unpad) - return flat.view(bs, seq, num_heads, head_dim) - - -def _build_meta(mask): - bs, seq = mask.shape - indices = mask.reshape(-1).nonzero(as_tuple=False).flatten() - inv_indices = build_inv_indices(indices, bs * seq) - return indices, inv_indices - - -@pytest.mark.parametrize("dtype", DTYPES) -@pytest.mark.parametrize( - "shape", SHAPES, ids=lambda s: s[0] if isinstance(s, tuple) else str(s) -) -def test_pack_matches_index_select(dtype, shape): - _, bs, s_txt, s_img, num_heads, head_dim, valid_txt_lens = shape - torch.manual_seed(0) - s = s_txt + s_img - mask = _build_mask(bs, s_txt, s_img, valid_txt_lens) - indices, _ = _build_meta(mask) - - q = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) - k = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) - v = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) - - q_ref, k_ref, v_ref = _ref_pack(q, k, v, indices) - q_fused, k_fused, v_fused = fused_pack_qkv(q, k, v, indices) - - # bit-exact: pack is pure gather, no math - assert torch.equal(q_ref, q_fused) - assert torch.equal(k_ref, k_fused) - assert torch.equal(v_ref, v_fused) - - -@pytest.mark.parametrize("dtype", DTYPES) -@pytest.mark.parametrize( - "shape", SHAPES, ids=lambda s: s[0] if isinstance(s, tuple) else str(s) -) -def test_scatter_matches_index_copy(dtype, shape): - _, bs, s_txt, s_img, num_heads, head_dim, valid_txt_lens = shape - torch.manual_seed(1) - s = s_txt + s_img - mask = _build_mask(bs, s_txt, s_img, valid_txt_lens) - indices, inv_indices = _build_meta(mask) - n_valid = indices.shape[0] - - out_unpad = torch.randn(n_valid, num_heads, head_dim, dtype=dtype, device=DEVICE) - out_ref = _ref_scatter(out_unpad, indices, bs, s) - out_fused = fused_scatter_to_padded(out_unpad, inv_indices, bs, s) - - # bit-exact: scatter is pure copy + zero-fill - assert torch.equal(out_ref, out_fused) - # Padding rows must be exactly zero - invalid = ~mask - if invalid.any(): - assert out_fused[invalid].abs().max().item() == 0.0 - - -def test_pack_handles_non_contiguous_input(): - """Helper must accept non-contiguous Q/K/V (auto .contiguous() inside).""" - torch.manual_seed(2) - bs, s_txt, s_img, num_heads, head_dim = 2, 64, 128, 4, 64 - s = s_txt + s_img - mask = _build_mask(bs, s_txt, s_img, [32, 48]) - indices, _ = _build_meta(mask) - - # Build non-contiguous tensors via permute - qkv_pre = torch.randn( - bs, num_heads, s, head_dim, dtype=torch.bfloat16, device=DEVICE - ) - q = qkv_pre.permute(0, 2, 1, 3) - k = torch.randn_like(qkv_pre).permute(0, 2, 1, 3) - v = torch.randn_like(qkv_pre).permute(0, 2, 1, 3) - assert not q.is_contiguous() - - q_ref, k_ref, v_ref = _ref_pack( - q.contiguous(), k.contiguous(), v.contiguous(), indices - ) - q_fused, k_fused, v_fused = fused_pack_qkv(q, k, v, indices) - assert torch.equal(q_ref, q_fused) - assert torch.equal(k_ref, k_fused) - assert torch.equal(v_ref, v_fused) - - -def test_build_inv_indices_matches_manual(): - """build_inv_indices output should match the manual full+scatter form.""" - torch.manual_seed(3) - bs, s = 2, 32 - mask = torch.bernoulli(torch.full((bs, s), 0.6, device=DEVICE)).to(torch.bool) - indices = mask.reshape(-1).nonzero(as_tuple=False).flatten() - n_valid = indices.shape[0] - - manual = torch.full((bs * s,), -1, dtype=torch.int32, device=DEVICE) - if n_valid > 0: - manual[indices.long()] = torch.arange(n_valid, dtype=torch.int32, device=DEVICE) - - built = build_inv_indices(indices, bs * s) - assert torch.equal(built, manual) - - -def test_empty_valid_set_handled(): - """All-False mask: pack returns empty tensors; scatter writes all zeros.""" - bs, s, num_heads, head_dim = 2, 16, 4, 64 - mask = torch.zeros(bs, s, dtype=torch.bool, device=DEVICE) - indices = mask.reshape(-1).nonzero(as_tuple=False).flatten() - inv_indices = build_inv_indices(indices, bs * s) - assert indices.numel() == 0 - - q = torch.randn(bs, s, num_heads, head_dim, dtype=torch.bfloat16, device=DEVICE) - q_unpad, k_unpad, v_unpad = fused_pack_qkv(q, q.clone(), q.clone(), indices) - assert q_unpad.shape == (0, num_heads, head_dim) - assert k_unpad.shape == (0, num_heads, head_dim) - assert v_unpad.shape == (0, num_heads, head_dim) - - out_padded = fused_scatter_to_padded(q_unpad, inv_indices, bs, s) - assert out_padded.shape == (bs, s, num_heads, head_dim) - assert out_padded.abs().max().item() == 0.0 - - -if __name__ == "__main__": - import sys - - sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/kernels/ops/diffusion/test_varlen_uspattn_equivalence.py b/test/registered/kernels/ops/diffusion/test_varlen_uspattn_equivalence.py deleted file mode 100644 index db67a6cf1926..000000000000 --- a/test/registered/kernels/ops/diffusion/test_varlen_uspattn_equivalence.py +++ /dev/null @@ -1,157 +0,0 @@ -"""End-to-end equivalence between USPAttention varlen path and SDPA reference. - -Compares the production varlen path (``build_varlen_mask_meta`` + -``fused_pack_qkv`` + ``flash_attn_varlen_func`` + ``fused_scatter_to_padded``) -against ``torch.nn.functional.scaled_dot_product_attention`` with a broadcast -key mask, for inputs the gating in ``USPAttention.forward`` would accept. - -Verifies the documented contract: - * Valid (non-masked) query rows match SDPA within FA-vs-SDPA tolerance. - * Masked query rows are exactly zero in the varlen path (differs from - SDPA, which produces deterministic attention output at those rows). -""" - -import pytest -import torch -import torch.nn.functional as F - -from sglang.kernels.jit.utils import get_ci_test_range -from sglang.kernels.ops.attention.flash_attention import flash_attn_varlen_func -from sglang.kernels.ops.diffusion.triton.varlen_pack_pad import ( - fused_pack_qkv, - fused_scatter_to_padded, -) -from sglang.multimodal_gen.runtime.layers.attention.backends import ( - flash_attn as _fa_backend, -) -from sglang.multimodal_gen.runtime.layers.attention.layer import ( - build_varlen_mask_meta, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=15, stage="base-b-kernel-unit", runner_config="1-gpu-large") -# Nightly is not redundant here: it sets SGLANG_JIT_KERNEL_RUN_FULL_TESTS=1 to expand get_ci_test_range sweeps. -register_cuda_ci(est_time=20, stage="nightly", runner_config="1-gpu-large") - -DEVICE = "cuda" -DTYPES = get_ci_test_range([torch.bfloat16, torch.float16], [torch.bfloat16]) -# (name, bs, s_txt, s_img, num_heads, head_dim, valid_txt_lens) -SHAPES = get_ci_test_range( - [ - ("small_c2", 2, 64, 128, 4, 64, [32, 48]), - ("prod_c2", 2, 256, 1024, 24, 128, [128, 200]), - ("all_valid_b1", 1, 64, 128, 4, 64, [64]), - ("zero_txt_one_batch", 2, 64, 128, 4, 64, [0, 32]), - ], - [ - ("small_c2", 2, 64, 128, 4, 64, [32, 48]), - ], -) - - -def _build_mask(bs, s_txt, s_img, valid_txt_lens): - s = s_txt + s_img - mask = torch.zeros(bs, s, dtype=torch.bool, device=DEVICE) - for b, vt in enumerate(valid_txt_lens): - mask[b, :vt] = True - mask[b, s_txt:] = True - return mask - - -def _sdpa_with_key_mask(q, k, v, key_mask, softmax_scale): - """Reference: SDPA with a ``[B, S]`` key mask broadcast to ``[B, 1, 1, S]``.""" - q_ = q.transpose(1, 2) - k_ = k.transpose(1, 2) - v_ = v.transpose(1, 2) - mask = key_mask.to(dtype=q.dtype)[:, None, None, :] - mask = (mask - 1.0) * torch.finfo(q.dtype).max - out = F.scaled_dot_product_attention( - q_, - k_, - v_, - attn_mask=mask, - dropout_p=0.0, - is_causal=False, - scale=softmax_scale, - ) - return out.transpose(1, 2) - - -def _varlen_path(q, k, v, key_mask, softmax_scale): - """Production varlen path matching USPAttention.forward masked branch.""" - bs, seq = q.shape[0], q.shape[1] - meta = build_varlen_mask_meta(key_mask) - indices = meta["indices"] - if indices.shape[0] == 0: - return torch.zeros_like(q) - q_unpad, k_unpad, v_unpad = fused_pack_qkv(q, k, v, indices) - out_unpad = flash_attn_varlen_func( - q=q_unpad, - k=k_unpad, - v=v_unpad, - cu_seqlens_q=meta["cu_seqlens"], - cu_seqlens_k=meta["cu_seqlens"], - max_seqlen_q=meta["max_seqlen"], - max_seqlen_k=meta["max_seqlen"], - softmax_scale=softmax_scale, - causal=False, - ver=_fa_backend.fa_ver, - ) - return fused_scatter_to_padded(out_unpad, meta["inv_indices"], bs, seq) - - -@pytest.mark.parametrize("dtype", DTYPES) -@pytest.mark.parametrize( - "shape", SHAPES, ids=lambda s: s[0] if isinstance(s, tuple) else str(s) -) -def test_varlen_path_matches_sdpa_on_valid_rows(dtype, shape): - """Valid rows: varlen output ≈ SDPA output within FA tolerance.""" - _, bs, s_txt, s_img, num_heads, head_dim, valid_txt_lens = shape - torch.manual_seed(0) - s = s_txt + s_img - softmax_scale = head_dim**-0.5 - mask = _build_mask(bs, s_txt, s_img, valid_txt_lens) - q = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) - k = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) - v = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) - - out_sdpa = _sdpa_with_key_mask(q, k, v, mask, softmax_scale) - out_varlen = _varlen_path(q, k, v, mask, softmax_scale) - - valid = mask[..., None, None].expand_as(out_sdpa) - rtol = 1e-2 if dtype == torch.bfloat16 else 5e-3 - atol = 5e-2 if dtype == torch.bfloat16 else 1e-2 - torch.testing.assert_close( - out_sdpa[valid], - out_varlen[valid], - rtol=rtol, - atol=atol, - ) - - -@pytest.mark.parametrize("dtype", DTYPES) -@pytest.mark.parametrize( - "shape", SHAPES, ids=lambda s: s[0] if isinstance(s, tuple) else str(s) -) -def test_varlen_path_zeros_masked_rows(dtype, shape): - """Masked rows: varlen path produces exact zeros (documented contract).""" - _, bs, s_txt, s_img, num_heads, head_dim, valid_txt_lens = shape - torch.manual_seed(1) - s = s_txt + s_img - softmax_scale = head_dim**-0.5 - mask = _build_mask(bs, s_txt, s_img, valid_txt_lens) - q = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) - k = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) - v = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) - - out_varlen = _varlen_path(q, k, v, mask, softmax_scale) - - invalid = ~mask - if invalid.any(): - assert (out_varlen[invalid] == 0).all(), "masked rows must be zero-filled" - - -if __name__ == "__main__": - import sys - - sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/kernels/ops/diffusion/test_wan_causal_cache.py b/test/registered/kernels/ops/diffusion/test_wan_causal_cache.py deleted file mode 100644 index 530c295def83..000000000000 --- a/test/registered/kernels/ops/diffusion/test_wan_causal_cache.py +++ /dev/null @@ -1,164 +0,0 @@ -"""Wan causal VAE data-movement kernels: the fused conv-input builder and the -fused DupUp3D shortcut add must be bitwise identical to the aten op chains -they replace (they are pure data movement plus zero fill / one fp32 add).""" - -import sys - -import pytest -import torch -import torch.nn.functional as F - -from sglang.kernels.ops.diffusion.triton.wan_causal_cache import ( - cat_pad_channels_last_3d, - dup_up3d_add, -) -from sglang.multimodal_gen.runtime.models.vaes import wanvae -from sglang.multimodal_gen.runtime.models.vaes.wanvae import ( - CACHE_T, - WanCausalConv3d, - _cache_payload, - _run_cached_causal_conv, -) -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=40, stage="base-b-kernel-unit", runner_config="1-gpu-large") - -pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") - - -def _cl3d(shape, dtype): - return torch.randn(shape, device="cuda", dtype=dtype).contiguous( - memory_format=torch.channels_last_3d - ) - - -def _ref_cat_pad(x, cache, padding): - p = list(padding) - if cache is not None: - x = torch.cat([cache, x], dim=2) - p[4] -= cache.shape[2] - if any(p): - x = F.pad(x, p) - return x.contiguous(memory_format=torch.channels_last_3d) - - -@torch.no_grad() -@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float32]) -@pytest.mark.parametrize( - "c,t,h,w,cache_t,pads", - [ - (96, 1, 10, 14, 0, (1, 1, 1, 1, 2, 0)), # first chunk, zero-fill front - (96, 1, 10, 14, 1, (1, 1, 1, 1, 2, 0)), # legacy 1-frame cache - (96, 1, 10, 14, 2, (1, 1, 1, 1, 2, 0)), # steady state k3 conv - (64, 1, 10, 14, 2, (0, 0, 0, 0, 2, 0)), # time_conv (temporal only) - (48, 4, 10, 14, 2, (1, 1, 1, 1, 2, 0)), # encoder-style T=4 chunk - ], -) -def test_cat_pad_bitwise(dtype, c, t, h, w, cache_t, pads) -> None: - torch.cuda.manual_seed(0) - x = _cl3d((1, c, t, h, w), dtype) - cache = None - if cache_t: - # Strided interior view: caches may arrive as non-contiguous slices. - ph, pw = pads[2], pads[0] - buf = _cl3d((1, c, cache_t, h + 2 * ph, w + 2 * pw), dtype) - cache = buf[:, :, :, ph : ph + h, pw : pw + w] - out = cat_pad_channels_last_3d(x, cache, pads) - ref = _ref_cat_pad(x, cache, pads) - assert out is not None and out.shape == ref.shape - assert out.is_contiguous(memory_format=torch.channels_last_3d) - assert torch.equal(out, ref) - - # Dual-output mode: the same pass also emits the compact feature cache - # (unpadded interior of the last frames), bitwise equal to the slice. - pair = cat_pad_channels_last_3d(x, cache, pads, keep_cache_t=2) - assert pair is not None - out2, keep = pair - assert torch.equal(out2, ref) - ph, pw = pads[2], pads[0] - keep_t = min(2, ref.shape[2]) - want = ref[:, :, ref.shape[2] - keep_t :, ph : ph + h, pw : pw + w] - assert keep.shape == want.shape - assert keep.is_contiguous(memory_format=torch.channels_last_3d) - assert torch.equal(keep, want) - - -@torch.no_grad() -@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float32]) -@pytest.mark.parametrize( - "c_in,c_out,t,h,w,ft,fs,drop", - [ - (128, 64, 1, 10, 14, 2, 2, False), - (128, 64, 1, 10, 14, 2, 2, True), # first_chunk slicing - (64, 32, 2, 10, 14, 1, 2, False), - ], -) -def test_dup_up3d_add_bitwise(dtype, c_in, c_out, t, h, w, ft, fs, drop) -> None: - torch.cuda.manual_seed(0) - repeats = c_out * ft * fs * fs // c_in - src = _cl3d((1, c_in, t, h, w), dtype) - t_out = t * ft - (ft - 1 if drop else 0) - # Main arm as a permuted view, like the WanResample 2D output. - main = torch.randn( - (1, t_out, c_out, h * fs, w * fs), device="cuda", dtype=dtype - ).permute(0, 2, 1, 3, 4) - - dup = src.repeat_interleave(repeats, dim=1) - dup = dup.view(1, c_out, ft, fs, fs, t, h, w) - dup = dup.permute(0, 1, 5, 2, 6, 3, 7, 4).contiguous() - dup = dup.view(1, c_out, t * ft, h * fs, w * fs) - if drop: - dup = dup[:, :, ft - 1 :, :, :] - ref = main + dup - - out = dup_up3d_add(main, src, ft, fs, repeats, drop) - assert out is not None and out.shape == ref.shape - # Layout must match the aten add output exactly (downstream reductions - # are layout-sensitive), and every value must be bitwise identical. - assert out.stride() == ref.stride() - assert torch.equal(out, ref) - - -@torch.no_grad() -@pytest.mark.parametrize("pads_temporal_only", [False, True]) -def test_cached_conv_chunk_loop_bitwise(pads_temporal_only) -> None: - """The fused conv-input/compact-cache scheme must reproduce the original - clone/cat bookkeeping bitwise across a chunked decode, including the - first-chunk zero fill and the "Rep" marker start used by WanResample.""" - torch.cuda.manual_seed(0) - c = 64 - if pads_temporal_only: - conv = WanCausalConv3d(c, 2 * c, (3, 1, 1), padding=(1, 0, 0)) - else: - conv = WanCausalConv3d(c, c, 3, padding=1) - conv = conv.to(device="cuda", dtype=torch.float32) - conv.weight.data = conv.weight.data.contiguous(memory_format=torch.channels_last_3d) - chunks = [_cl3d((1, c, 1, 10, 14), torch.float32) for _ in range(4)] - - def run(force_fallback, start): - cache = [start] - outs = [] - if force_fallback: - orig = wanvae.cat_pad_channels_last_3d - wanvae.cat_pad_channels_last_3d = None - try: - for x in chunks: - outs.append(_run_cached_causal_conv(conv, x, cache, 0)) - finally: - if force_fallback: - wanvae.cat_pad_channels_last_3d = orig - return outs, cache[0] - - for start in (None, "Rep"): - fused_outs, fused_cache = run(False, start) - ref_outs, ref_cache = run(True, start) - for got, want in zip(fused_outs, ref_outs, strict=True): - assert torch.equal(got, want) - got_payload = _cache_payload(fused_cache) - assert got_payload is not None and got_payload.shape[2] == CACHE_T - # Reference cache holds the last CACHE_T unpadded frames. - assert torch.equal(got_payload, ref_cache[:, :, -CACHE_T:]) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/kernels/ops/diffusion/test_wan_vae_fastpath.py b/test/registered/kernels/ops/diffusion/test_wan_vae_fastpath.py deleted file mode 100644 index 260f391c7401..000000000000 --- a/test/registered/kernels/ops/diffusion/test_wan_vae_fastpath.py +++ /dev/null @@ -1,79 +0,0 @@ -"""Wan VAE decoder fast path: fused-kernel numerics and gate dispatch -(the lossless off-path must stay bit-exact).""" - -import sys - -import pytest -import torch -import torch.nn as nn -import torch.nn.functional as F - -from sglang.kernels.ops.diffusion.triton.wan_rmsnorm_silu import wan_rmsnorm_silu -from sglang.multimodal_gen.runtime.models.vaes.wan_vae_cuda_opt import ( - FusedWanRMSNormSiLU, - VaeFastPathGate, -) -from sglang.multimodal_gen.runtime.models.vaes.wanvae import WanRMS_norm -from sglang.test.ci.ci_register import register_cuda_ci - -register_cuda_ci(est_time=40, stage="base-b-kernel-unit", runner_config="1-gpu-large") - -pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") - - -def _cl3d(shape, dtype): - return torch.randn(shape, device="cuda", dtype=dtype).contiguous( - memory_format=torch.channels_last_3d - ) - - -@torch.no_grad() -@pytest.mark.parametrize( - "x_dtype,affine_dtype,atol,rtol", - [ - (torch.float32, torch.float32, 1e-5, 1e-5), # FastWan2.2 fp32 decode - (torch.bfloat16, torch.float32, 1.5e-1, 3e-2), # Wan2.1 bf16 autocast - ], -) -def test_kernel_numerics(x_dtype, affine_dtype, atol, rtol) -> None: - torch.cuda.manual_seed(0) - x = _cl3d((1, 96, 3, 10, 14), x_dtype) - gamma = torch.randn((96, 1, 1, 1), device="cuda", dtype=affine_dtype) - for bias in (None, torch.randn_like(gamma)): - expected = F.silu( - F.normalize(x, dim=1) * 96**0.5 * gamma + (0 if bias is None else bias) - ) - actual = wan_rmsnorm_silu(x, gamma, bias) - assert actual is not None and actual.dtype == expected.dtype - assert actual.stride() == x.stride() - torch.testing.assert_close(actual, expected, atol=atol, rtol=rtol) - - -@torch.no_grad() -def test_fused_module_gate_dispatch() -> None: - # Gate off must stay bit-exact; gate on must route to the fused kernel. - torch.cuda.manual_seed(0) - norm = WanRMS_norm(96, images=False).to(device="cuda", dtype=torch.bfloat16) - norm.gamma.add_(torch.randn_like(norm.gamma)) - gate = VaeFastPathGate() - fused = FusedWanRMSNormSiLU(norm, gate) - # Parameter names must not change (weight transfer matches by name). - assert [n for n, _ in fused.named_parameters()] == ["gamma"] - x = _cl3d((1, 96, 3, 10, 14), torch.bfloat16) - assert torch.equal(fused(x), nn.SiLU()(norm(x))) - gate.enabled = True - expected = wan_rmsnorm_silu(x, norm.gamma, rms_scale=float(norm.scale)) - assert torch.equal(fused(x), expected) - - -@torch.no_grad() -def test_kernel_rejects_empty_input() -> None: - x = torch.empty(1, 96, 0, 2, 2, device="cuda", dtype=torch.bfloat16).to( - memory_format=torch.channels_last_3d - ) - gamma = torch.ones(96, 1, 1, 1, device="cuda", dtype=torch.bfloat16) - assert wan_rmsnorm_silu(x, gamma) is None - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/kernels/ops/diffusion/test_diffusion_modelopt_fp8_scaled_mm.py b/test/registered/kernels/ops/quantization/test_diffusion_modelopt_fp8_scaled_mm.py similarity index 100% rename from test/registered/kernels/ops/diffusion/test_diffusion_modelopt_fp8_scaled_mm.py rename to test/registered/kernels/ops/quantization/test_diffusion_modelopt_fp8_scaled_mm.py diff --git a/test/registered/kernels/ops/diffusion/test_diffusion_nvfp4_scaled_mm.py b/test/registered/kernels/ops/quantization/test_diffusion_nvfp4_scaled_mm.py similarity index 100% rename from test/registered/kernels/ops/diffusion/test_diffusion_nvfp4_scaled_mm.py rename to test/registered/kernels/ops/quantization/test_diffusion_nvfp4_scaled_mm.py diff --git a/test/registered/kernels/ops/diffusion/test_stage_profiler_sync.py b/test/registered/profiling/test_diffusion_stage_profiler_sync.py similarity index 100% rename from test/registered/kernels/ops/diffusion/test_stage_profiler_sync.py rename to test/registered/profiling/test_diffusion_stage_profiler_sync.py From f4cf3d704078e0c54206e5224de04fd92593ec25 Mon Sep 17 00:00:00 2001 From: BBuf <1182563586@qq.com> Date: Mon, 17 Aug 2026 18:13:40 +0800 Subject: [PATCH 2/5] [kernels] Fix CI lane inheritance in the merged diffusion test suites Merging N test files into one operator-domain suite also merges their CI lane registrations, so each part inherits the union of the lanes its siblings ran on. Where the parts have disjoint platform requirements that puts tests on runners their dependencies do not exist on. Two lanes broke: - `jit-kernel-b200-test`: `test_layout.py` picked up the B200 lane from the causal-Conv3d section, and the varlen<->USPAttention comparison folded in from `test_varlen_uspattn_equivalence.py` needs FlashAttention, which that image does not ship. `flash_attn_varlen_func` resolves its backend lazily, so it raises `ImportError` at the call, not at import. Caught at the call site and skipped there; the guard is deliberately not a module-level `flash_attn` probe, which would also skip on H200 where the dispatch resolves to a backend that does exist. - `jit-kernel-unit-test-amd-rocm720`: `test_norm.py` picked up the AMD lane from the FlyDSL section, and the CuTe-DSL cases need cuda-python (`ModuleNotFoundError: No module named 'cuda'`, 147 failures). FlyDSL is an AMD gfx950-only compiler and was the only part of that file on the AMD lane at all, so it moves to `test_norm_flydsl.py` with exactly its original registration; `test_norm.py` returns to being CUDA-only. The CuTe-DSL imports also gain an explicit skip so the constraint is stated in code rather than implied by which lanes the file happens to be on. Audited the other five merged suites the same way. `test_layout.py` already passed on the AMD lane and 43/49 on B200, so the remaining inherited lanes are real added coverage, not latent breakage. Verified on 1x H200: 2581 passed, 13 skipped -- identical to the run before this commit, confirming the new guards cost no coverage where the backends exist. Co-Authored-By: Claude Opus 5 (1M context) --- .../kernels/ops/diffusion/test_layout.py | 33 +++-- .../kernels/ops/diffusion/test_norm.py | 115 ++++-------------- .../kernels/ops/diffusion/test_norm_flydsl.py | 98 +++++++++++++++ 3 files changed, 146 insertions(+), 100 deletions(-) create mode 100644 test/registered/kernels/ops/diffusion/test_norm_flydsl.py diff --git a/test/registered/kernels/ops/diffusion/test_layout.py b/test/registered/kernels/ops/diffusion/test_layout.py index 93c37d281250..59f33a7b2360 100644 --- a/test/registered/kernels/ops/diffusion/test_layout.py +++ b/test/registered/kernels/ops/diffusion/test_layout.py @@ -314,18 +314,27 @@ def _varlen_path(q, k, v, key_mask, softmax_scale): if indices.shape[0] == 0: return torch.zeros_like(q) q_unpad, k_unpad, v_unpad = fused_pack_qkv(q, k, v, indices) - out_unpad = flash_attn_varlen_func( - q=q_unpad, - k=k_unpad, - v=v_unpad, - cu_seqlens_q=meta["cu_seqlens"], - cu_seqlens_k=meta["cu_seqlens"], - max_seqlen_q=meta["max_seqlen"], - max_seqlen_k=meta["max_seqlen"], - softmax_scale=softmax_scale, - causal=False, - ver=_fa_backend.fa_ver, - ) + try: + out_unpad = flash_attn_varlen_func( + q=q_unpad, + k=k_unpad, + v=v_unpad, + cu_seqlens_q=meta["cu_seqlens"], + cu_seqlens_k=meta["cu_seqlens"], + max_seqlen_q=meta["max_seqlen"], + max_seqlen_k=meta["max_seqlen"], + softmax_scale=softmax_scale, + causal=False, + ver=_fa_backend.fa_ver, + ) + except ImportError as exc: # pragma: no cover - image-dependent + # ``flash_attn_varlen_func`` resolves its backend lazily, so an image + # without the selected FlashAttention build raises here rather than at + # import. This file also runs on the B200 lane (for the causal-Conv3d + # section), which ships no ``flash_attn`` -- skip only this end-to-end + # comparison there; the pack/scatter kernels themselves are covered + # unit-wise above on every lane. + pytest.skip(f"FlashAttention varlen v{_fa_backend.fa_ver} unavailable: {exc}") return fused_scatter_to_padded(out_unpad, meta["inv_indices"], bs, seq) diff --git a/test/registered/kernels/ops/diffusion/test_norm.py b/test/registered/kernels/ops/diffusion/test_norm.py index a3f91881a153..d2db4cf7dbb0 100644 --- a/test/registered/kernels/ops/diffusion/test_norm.py +++ b/test/registered/kernels/ops/diffusion/test_norm.py @@ -11,7 +11,11 @@ ``None``). - ``rmsnorm_scale`` / ``rmsnorm_tanh_residual`` -> a bf16-native reference that reproduces Z-Image's own norm, with a tolerance for Triton's exp-based tanh. -- the CuTe-DSL and FlyDSL fused norm+scale/shift -> an fp32 reference chain. +- the CuTe-DSL fused norm+scale/shift -> an fp32 reference chain. + +The FlyDSL norms live in ``test_norm_flydsl.py``: they are ROCm gfx950-only, +so they run on a CI lane this file does not, and keeping them here dragged the +CUDA-only CuTe-DSL cases onto the AMD runner. The *bit-exact* norms (``fused_rmsnorm_scale_shift_bitexact``, ``fused_layernorm_modulate``, ``zimage_qk_rmsnorm_native``) are exercised @@ -37,11 +41,9 @@ triton_group_norm_silu, wan_rmsnorm_silu, ) -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci +from sglang.test.ci.ci_register import register_cuda_ci -register_cuda_ci(est_time=85, stage="base-b-kernel-unit", runner_config="1-gpu-large") -register_amd_ci(est_time=30, stage="jit-kernel-unit", runner_config="amd") -register_amd_ci(est_time=15, suite="nightly-amd-kernel-1-gpu", nightly=True) +register_cuda_ci(est_time=70, stage="base-b-kernel-unit", runner_config="1-gpu-large") pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") @@ -308,6 +310,24 @@ def test_wan_rmsnorm_silu_rejects_empty_input(): INDEX_MODES = ["BSD", "1", "1SD", "BD", "B1D", "D", "1D", "11D", "BF1D"] +def _import_cutedsl(): + """Import the CuTe-DSL entry points, skipping when the backend is absent. + + This file is registered on the AMD lane for its FlyDSL section, but the + CuTe-DSL norms need cuda-python and CUTLASS, which the ROCm image does not + ship. Guarded per test rather than by dropping this file from the AMD + lane, so the Triton and FlyDSL sections keep running there. + """ + try: + from sglang.kernels.ops.diffusion import ( + fused_norm_scale_shift, + fused_scale_residual_norm_scale_shift, + ) + except ImportError as exc: # pragma: no cover - platform-dependent + pytest.skip(f"CuTe-DSL backend unavailable: {exc}") + return fused_norm_scale_shift, fused_scale_residual_norm_scale_shift + + def _make_tensor(index_mode, shape, dtype): if index_mode == "NAT": return None @@ -364,14 +384,7 @@ def _run_cute( index_mode="BSD", eps=EPS, ): - # Imported inside the helper for the same reason as FlyDSL below: the - # CuTe-DSL path needs cuda-python + CUTLASS, and naming a facade export - # resolves it immediately -- at module level that would take the Triton - # sections of this file down with it wherever CUTLASS is absent. - from sglang.kernels.ops.diffusion import ( - fused_norm_scale_shift, - fused_scale_residual_norm_scale_shift, - ) + fused_norm_scale_shift, fused_scale_residual_norm_scale_shift = _import_cutedsl() x = _make_tensor("BSD", shape, dtype) weight = _make_tensor(affine_mode, shape, affine_dtype) @@ -440,6 +453,7 @@ def test_cutedsl_scale_residual_gate_index_modes(norm_type, index_mode): def test_validate_scale_shift_rejects_non_divisible_frames(): + _import_cutedsl() from sglang.kernels.ops.diffusion import validate_scale_shift with pytest.raises(ValueError, match=r"S\(10\) must be divisible by F\(4\)"): @@ -448,80 +462,5 @@ def test_validate_scale_shift_rejects_non_divisible_frames(): ) -# --------------------------------------------------------------------------- -# FlyDSL fused norm + scale/shift (ROCm gfx950) -# --------------------------------------------------------------------------- - -FLYDSL_D = 5120 -FLYDSL_EPS = 1e-6 - - -def _require_rocm(): - if not torch.version.hip: - pytest.skip("ROCm/HIP required for FlyDSL kernels") - - -def _flydsl_reference(residual, x, gate, weight, bias, scale, shift, norm_type, eps): - if residual is not None: - x = (residual.float() + x.float() * gate.float()).to(torch.bfloat16) - residual_out = x - else: - residual_out = None - if norm_type == "layer": - normed = F.layer_norm(x.float(), (FLYDSL_D,), weight, bias, eps) - else: - var = x.float().pow(2).mean(-1, keepdim=True) - normed = x.float() * torch.rsqrt(var + eps) * weight.float() - y = (normed * (1.0 + scale.float()) + shift.float()).to(torch.bfloat16) - return y, residual_out - - -@pytest.mark.parametrize("with_residual", [False, True]) -@pytest.mark.parametrize( - "norm_type,B,L", - [("rms", 1, 16), ("rms", 2, 16), ("layer", 2, 16), ("rms", 1, 90000)], -) -def test_flydsl_norm_scale_shift(with_residual, norm_type, B, L): - _require_rocm() - # Imported inside the test: the FlyDSL compiler only exists on ROCm, and - # the facade resolves an export the moment it is named -- a module-level - # import here would fail collection of this whole file on CUDA. - from sglang.kernels.ops.diffusion import ( - flydsl_fused_residual_norm_scale_shift, - flydsl_norm_scale_shift, - ) - - torch.manual_seed(42) - shape = (B, L, FLYDSL_D) - x = torch.randn(shape, device=DEVICE, dtype=torch.bfloat16) - weight = torch.randn(FLYDSL_D, device=DEVICE, dtype=torch.float32) - bias = ( - torch.randn(FLYDSL_D, device=DEVICE, dtype=torch.float32) - if norm_type == "layer" - else None - ) - scale = torch.randn(B, 1, FLYDSL_D, device=DEVICE, dtype=torch.bfloat16) - shift = torch.randn(B, 1, FLYDSL_D, device=DEVICE, dtype=torch.bfloat16) - - if with_residual: - residual = torch.randn(shape, device=DEVICE, dtype=torch.bfloat16) - gate = torch.randn(B, 1, FLYDSL_D, device=DEVICE, dtype=torch.bfloat16) - y, res = flydsl_fused_residual_norm_scale_shift( - residual, x, gate, weight, bias, scale, shift, norm_type, FLYDSL_EPS - ) - y_ref, res_ref = _flydsl_reference( - residual, x, gate, weight, bias, scale, shift, norm_type, FLYDSL_EPS - ) - torch.testing.assert_close(res, res_ref, atol=5e-2, rtol=5e-2) - else: - y = flydsl_norm_scale_shift( - x, weight, bias, scale, shift, norm_type, FLYDSL_EPS - ) - y_ref, _ = _flydsl_reference( - None, x, None, weight, bias, scale, shift, norm_type, FLYDSL_EPS - ) - torch.testing.assert_close(y, y_ref, atol=1.0, rtol=5e-2) - - if __name__ == "__main__": sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernels/ops/diffusion/test_norm_flydsl.py b/test/registered/kernels/ops/diffusion/test_norm_flydsl.py new file mode 100644 index 000000000000..e399477e54b4 --- /dev/null +++ b/test/registered/kernels/ops/diffusion/test_norm_flydsl.py @@ -0,0 +1,98 @@ +"""``diffusion.norm``: the FlyDSL fused norm + scale/shift kernels (ROCm). + +Split out of ``test_norm.py`` rather than merged with the other norm backends: +FlyDSL is an AMD gfx950-only compiler, so these run on the AMD CI lane and +nothing else in that file does. Keeping them together forced the CUDA-only +CuTe-DSL cases onto the ROCm runner, where cuda-python does not exist. + +Oracle: an fp32 reference chain, with a tolerance -- the kernel keeps fp32 +statistics but reorders the reduction. +""" + +import sys + +import pytest +import torch +import torch.nn.functional as F + +from sglang.test.ci.ci_register import register_amd_ci + +register_amd_ci(est_time=30, stage="jit-kernel-unit", runner_config="amd") + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="GPU required") + +DEVICE = "cuda" + +FLYDSL_D = 5120 +FLYDSL_EPS = 1e-6 + + +def _require_rocm(): + if not torch.version.hip: + pytest.skip("ROCm/HIP required for FlyDSL kernels") + + +def _flydsl_reference(residual, x, gate, weight, bias, scale, shift, norm_type, eps): + if residual is not None: + x = (residual.float() + x.float() * gate.float()).to(torch.bfloat16) + residual_out = x + else: + residual_out = None + if norm_type == "layer": + normed = F.layer_norm(x.float(), (FLYDSL_D,), weight, bias, eps) + else: + var = x.float().pow(2).mean(-1, keepdim=True) + normed = x.float() * torch.rsqrt(var + eps) * weight.float() + y = (normed * (1.0 + scale.float()) + shift.float()).to(torch.bfloat16) + return y, residual_out + + +@pytest.mark.parametrize("with_residual", [False, True]) +@pytest.mark.parametrize( + "norm_type,B,L", + [("rms", 1, 16), ("rms", 2, 16), ("layer", 2, 16), ("rms", 1, 90000)], +) +def test_flydsl_norm_scale_shift(with_residual, norm_type, B, L): + _require_rocm() + # Imported inside the test: the FlyDSL compiler only exists on ROCm, and + # the facade resolves an export the moment it is named -- a module-level + # import here would fail collection of this whole file on CUDA. + from sglang.kernels.ops.diffusion import ( + flydsl_fused_residual_norm_scale_shift, + flydsl_norm_scale_shift, + ) + + torch.manual_seed(42) + shape = (B, L, FLYDSL_D) + x = torch.randn(shape, device=DEVICE, dtype=torch.bfloat16) + weight = torch.randn(FLYDSL_D, device=DEVICE, dtype=torch.float32) + bias = ( + torch.randn(FLYDSL_D, device=DEVICE, dtype=torch.float32) + if norm_type == "layer" + else None + ) + scale = torch.randn(B, 1, FLYDSL_D, device=DEVICE, dtype=torch.bfloat16) + shift = torch.randn(B, 1, FLYDSL_D, device=DEVICE, dtype=torch.bfloat16) + + if with_residual: + residual = torch.randn(shape, device=DEVICE, dtype=torch.bfloat16) + gate = torch.randn(B, 1, FLYDSL_D, device=DEVICE, dtype=torch.bfloat16) + y, res = flydsl_fused_residual_norm_scale_shift( + residual, x, gate, weight, bias, scale, shift, norm_type, FLYDSL_EPS + ) + y_ref, res_ref = _flydsl_reference( + residual, x, gate, weight, bias, scale, shift, norm_type, FLYDSL_EPS + ) + torch.testing.assert_close(res, res_ref, atol=5e-2, rtol=5e-2) + else: + y = flydsl_norm_scale_shift( + x, weight, bias, scale, shift, norm_type, FLYDSL_EPS + ) + y_ref, _ = _flydsl_reference( + None, x, None, weight, bias, scale, shift, norm_type, FLYDSL_EPS + ) + torch.testing.assert_close(y, y_ref, atol=1.0, rtol=5e-2) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) From e2d8060859ceb97390bf94360f37350f398abb69 Mon Sep 17 00:00:00 2001 From: BBuf <1182563586@qq.com> Date: Mon, 17 Aug 2026 20:06:29 +0800 Subject: [PATCH 3/5] [kernels] Separate the B200- and HIP-exclusive diffusion test lanes Second round of the same lane-inheritance problem, on the two lanes that only started reporting once the first round unblocked them. `test_rope.py` carried the B200 lane because the LTX-2 split-RoPE kernel is validated there and nowhere else. But the `fused_inplace_qknorm_rope` cases in that file are held to the *split* baseline -- a separate qknorm kernel plus sgl_kernel / FlashInfer RoPE -- and that dispatch differs on Blackwell, so three bit-exact assertions that hold on H200 do not hold on B200. The two groups need opposite lanes, so LTX-2 moves to `test_rope_ltx2.py` with its original B200-only registration and `test_rope.py` returns to the CUDA lanes its content was written against. Same shape as the FlyDSL split in the previous commit. `test_usp_merge_heads_bitwise` reached the AMD lane through `test_layout.py` and asserts `can_use_usp_merge_heads(x)`, which returns False under HIP *by design* -- the sibling test in the same file even patches `torch.version.hip` to check that. Skipped on HIP, where the wrapper's aten fallback is what runs and is already covered. The remaining inherited lanes in `test_layout.py` are real coverage, not latent breakage: the B200 run reported 43 passed alongside the 6 FlashAttention failures, so every other section this file gained on that lane already passes there. Verified on 1x H200: 2581 passed, 13 skipped -- unchanged across the split. Co-Authored-By: Claude Opus 5 (1M context) --- .../kernels/ops/diffusion/test_layout.py | 6 + .../kernels/ops/diffusion/test_rope.py | 223 +---------------- .../kernels/ops/diffusion/test_rope_ltx2.py | 233 ++++++++++++++++++ 3 files changed, 245 insertions(+), 217 deletions(-) create mode 100644 test/registered/kernels/ops/diffusion/test_rope_ltx2.py diff --git a/test/registered/kernels/ops/diffusion/test_layout.py b/test/registered/kernels/ops/diffusion/test_layout.py index 59f33a7b2360..9cf834bba9b5 100644 --- a/test/registered/kernels/ops/diffusion/test_layout.py +++ b/test/registered/kernels/ops/diffusion/test_layout.py @@ -74,6 +74,12 @@ def _cl3d(shape, dtype): ], ) @pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16, torch.float32]) +@pytest.mark.skipif( + bool(torch.version.hip), + reason="the USP merge-heads JIT fast path is CUDA-only by design -- " + "can_use_usp_merge_heads() returns False under HIP, and the aten fallback " + "it degrades to is covered by the unsupported-inputs test below", +) def test_usp_merge_heads_bitwise(dtype, world, seq, batch, h_local, head_dim): generator = torch.Generator(device=DEVICE).manual_seed(4321) x = torch.randn( diff --git a/test/registered/kernels/ops/diffusion/test_rope.py b/test/registered/kernels/ops/diffusion/test_rope.py index b3537639b018..1889b1ba2c37 100644 --- a/test/registered/kernels/ops/diffusion/test_rope.py +++ b/test/registered/kernels/ops/diffusion/test_rope.py @@ -8,8 +8,11 @@ rounding step, so those cases use a tolerance; with ``round_norm_before_rope=True`` the fused kernel reproduces the split rounding exactly and ``torch.equal`` applies. -- The LTX-2 split-RoPE CUDA kernel is validated on B200 only (guarded per - test, not per module, so the rest of this file still runs elsewhere). +The LTX-2 split-RoPE kernel lives in ``test_rope_ltx2.py``: it is validated on +B200 and registered on that lane alone, which the cases here cannot share -- +their oracle is the *split* baseline (a separate qknorm kernel plus sgl_kernel +or FlashInfer RoPE), whose dispatch differs on Blackwell, so the bit-exact +assertions below do not hold there. """ import itertools @@ -17,21 +20,17 @@ import pytest import torch -import torch.nn.functional as F import triton from sglang.kernels.jit.utils import get_ci_test_range from sglang.kernels.ops.diffusion import ( can_use_fused_inplace_qknorm_rope, - can_use_ltx2_qknorm_split_rope_cuda, fused_inplace_qknorm_rope, fused_qknorm_rope_pack_kv, - ltx2_qknorm_split_rope_cuda, ) from sglang.test.ci.ci_register import register_cuda_ci -register_cuda_ci(est_time=52, stage="base-b-kernel-unit", runner_config="1-gpu-large") -register_cuda_ci(est_time=45, stage="base-b-kernel-unit", runner_config="4-gpu-b200") +register_cuda_ci(est_time=44, stage="base-b-kernel-unit", runner_config="1-gpu-large") # Nightly is not redundant: it sets SGLANG_JIT_KERNEL_RUN_FULL_TESTS=1, which # expands the get_ci_test_range sweeps below. register_cuda_ci(est_time=220, stage="nightly", runner_config="1-gpu-large") @@ -44,7 +43,6 @@ ROPE_BASE = 10000.0 ATOL = 8e-2 RTOL = 1e-2 -BF16_FUSED_ATOL = 1.6e-1 def create_cos_sin_cache( @@ -535,214 +533,5 @@ def test_qknorm_rope_accepts_empty_token_dimension() -> None: assert q.numel() == k.numel() == 0 -# ------------------------------------------------------------------------- -# LTX-2 split RoPE (B200) -# ------------------------------------------------------------------------- - - -def _require_b200() -> None: - if not torch.cuda.is_available(): - pytest.skip("CUDA required") - if torch.cuda.get_device_capability()[0] < 10: - pytest.skip("LTX2 QKNorm split-RoPE CUDA path is validated on B200") - - -def _ltx2_make_cos_sin( - batch: int, seq_len: int, num_heads: int, head_dim: int -) -> tuple[torch.Tensor, torch.Tensor]: - half_dim = head_dim // 2 - cos = torch.randn( - batch, seq_len, num_heads, half_dim, device="cuda", dtype=torch.bfloat16 - ).transpose(1, 2) - sin = torch.randn( - batch, seq_len, num_heads, half_dim, device="cuda", dtype=torch.bfloat16 - ).transpose(1, 2) - return cos, sin - - -def _apply_split_rotary_ref( - x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor -) -> torch.Tensor: - x_dtype = x.dtype - batch = x.shape[0] - _, num_heads, seq_len, _ = cos.shape - x = x.reshape(batch, seq_len, num_heads, -1).swapaxes(1, 2) - last = x.shape[-1] - half = last // 2 - - split_x = x.reshape(*x.shape[:-1], 2, half) - first_x = split_x[..., :1, :] - second_x = split_x[..., 1:, :] - cos_u = cos.unsqueeze(-2) - sin_u = sin.unsqueeze(-2) - - out = split_x * cos_u - out[..., :1, :].addcmul_(-sin_u, second_x) - out[..., 1:, :].addcmul_(sin_u, first_x) - out = out.reshape(*out.shape[:-2], last) - return out.swapaxes(1, 2).reshape(batch, seq_len, -1).to(dtype=x_dtype) - - -def _ltx2_reference( - q: torch.Tensor, - k: torch.Tensor, - q_cos: torch.Tensor, - q_sin: torch.Tensor, - k_cos: torch.Tensor, - k_sin: torch.Tensor, - q_weight: torch.Tensor, - k_weight: torch.Tensor, - eps: float, -) -> tuple[torch.Tensor, torch.Tensor]: - # rms_norm isn't autocast fp32-preserving, so feed fp32 inputs directly - # to keep the normalized value unrounded until the final RoPE output. - q_norm = F.rms_norm(q.float(), (q.shape[-1],), q_weight.float(), eps) - k_norm = F.rms_norm(k.float(), (k.shape[-1],), k_weight.float(), eps) - q_ref = _apply_split_rotary_ref(q_norm, q_cos, q_sin) - k_ref = _apply_split_rotary_ref(k_norm, k_cos, k_sin) - return q_ref.to(dtype=torch.bfloat16), k_ref.to(dtype=torch.bfloat16) - - -@pytest.mark.parametrize( - "batch,q_seq,k_seq,num_heads,head_dim", - [ - (1, 3, 3, 32, 128), - (1, 5, 2, 32, 64), - (2, 4, 3, 32, 64), - ], -) -def test_ltx2_qknorm_split_rope_matches_torch_exactly( - batch: int, q_seq: int, k_seq: int, num_heads: int, head_dim: int -) -> None: - _require_b200() - torch.cuda.manual_seed(20260630) - hidden = num_heads * head_dim - eps = 1e-6 - q = torch.randn(batch, q_seq, hidden, device="cuda", dtype=torch.bfloat16) - k = torch.randn(batch, k_seq, hidden, device="cuda", dtype=torch.bfloat16) - q_cos, q_sin = _ltx2_make_cos_sin(batch, q_seq, num_heads, head_dim) - k_cos, k_sin = _ltx2_make_cos_sin(batch, k_seq, num_heads, head_dim) - q_weight = torch.randn(hidden, device="cuda", dtype=torch.bfloat16) - k_weight = torch.randn(hidden, device="cuda", dtype=torch.bfloat16) - - assert can_use_ltx2_qknorm_split_rope_cuda( - q, - q_cos, - q_sin, - q_weight, - k, - k_cos, - k_sin, - k_weight, - num_heads=num_heads, - head_dim=head_dim, - ) - - q_ref, k_ref = _ltx2_reference( - q, k, q_cos, q_sin, k_cos, k_sin, q_weight, k_weight, eps - ) - q_out, k_out = ltx2_qknorm_split_rope_cuda( - q, - q_cos, - q_sin, - q_weight, - k, - k_cos, - k_sin, - k_weight, - eps=eps, - num_heads=num_heads, - head_dim=head_dim, - ) - torch.cuda.synchronize() - - torch.testing.assert_close(q_out, q_ref, rtol=0, atol=BF16_FUSED_ATOL) - torch.testing.assert_close(k_out, k_ref, rtol=0, atol=BF16_FUSED_ATOL) - - -def test_ltx2_qknorm_split_rope_rejects_unsupported_inputs() -> None: - _require_b200() - torch.cuda.manual_seed(20260630) - q = torch.randn((1, 3, 4096), device="cuda", dtype=torch.bfloat16) - k = torch.randn_like(q) - q_cos, q_sin = _ltx2_make_cos_sin(1, 3, 32, 128) - q_weight = torch.randn(4096, device="cuda", dtype=torch.bfloat16) - k_weight = torch.randn(4096, device="cuda", dtype=torch.bfloat16) - - assert can_use_ltx2_qknorm_split_rope_cuda( - q, - q_cos, - q_sin, - q_weight, - k, - q_cos, - q_sin, - k_weight, - num_heads=32, - head_dim=128, - ) - assert not can_use_ltx2_qknorm_split_rope_cuda( - q.float(), - q_cos, - q_sin, - q_weight, - k, - q_cos, - q_sin, - k_weight, - num_heads=32, - head_dim=128, - ) - assert not can_use_ltx2_qknorm_split_rope_cuda( - q, - q_cos, - q_sin, - q_weight, - k, - q_cos.transpose(-1, -2), - q_sin, - k_weight, - num_heads=32, - head_dim=128, - ) - - -def test_ltx2_qknorm_split_rope_custom_op_torch_compile_fullgraph() -> None: - _require_b200() - torch.cuda.manual_seed(20260630) - batch, q_seq, k_seq, num_heads, head_dim = 1, 3, 2, 32, 64 - hidden = num_heads * head_dim - q = torch.randn(batch, q_seq, hidden, device="cuda", dtype=torch.bfloat16) - k = torch.randn(batch, k_seq, hidden, device="cuda", dtype=torch.bfloat16) - q_cos, q_sin = _ltx2_make_cos_sin(batch, q_seq, num_heads, head_dim) - k_cos, k_sin = _ltx2_make_cos_sin(batch, k_seq, num_heads, head_dim) - q_weight = torch.randn(hidden, device="cuda", dtype=torch.bfloat16) - k_weight = torch.randn(hidden, device="cuda", dtype=torch.bfloat16) - - def fn(q, k, q_cos, q_sin, k_cos, k_sin, q_weight, k_weight): - return ltx2_qknorm_split_rope_cuda( - q, - q_cos, - q_sin, - q_weight, - k, - k_cos, - k_sin, - k_weight, - eps=1e-6, - num_heads=num_heads, - head_dim=head_dim, - ) - - compiled = torch.compile(fn, fullgraph=True) - q_out, k_out = compiled(q, k, q_cos, q_sin, k_cos, k_sin, q_weight, k_weight) - q_ref, k_ref = _ltx2_reference( - q, k, q_cos, q_sin, k_cos, k_sin, q_weight, k_weight, 1e-6 - ) - torch.cuda.synchronize() - torch.testing.assert_close(q_out, q_ref, rtol=0, atol=BF16_FUSED_ATOL) - torch.testing.assert_close(k_out, k_ref, rtol=0, atol=BF16_FUSED_ATOL) - - if __name__ == "__main__": sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernels/ops/diffusion/test_rope_ltx2.py b/test/registered/kernels/ops/diffusion/test_rope_ltx2.py new file mode 100644 index 000000000000..e6e1e4bc7075 --- /dev/null +++ b/test/registered/kernels/ops/diffusion/test_rope_ltx2.py @@ -0,0 +1,233 @@ +"""``diffusion.rope``: the LTX-2 QK-norm + split-RoPE CUDA kernel. + +Split out of ``test_rope.py`` rather than merged with the other RoPE kernels: +this one is validated on B200 and registered on that lane alone, while the +``fused_inplace_qknorm_rope`` cases there are held to the *split* baseline, +whose sgl_kernel / FlashInfer dispatch differs on Blackwell -- their bit-exact +assertions fail on B200. One file cannot carry both lane sets. +""" + +import sys + +import pytest +import torch +import torch.nn.functional as F + +from sglang.kernels.ops.diffusion import ( + can_use_ltx2_qknorm_split_rope_cuda, + ltx2_qknorm_split_rope_cuda, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=45, stage="base-b-kernel-unit", runner_config="4-gpu-b200") + +DEVICE = "cuda" +BF16_FUSED_ATOL = 1.6e-1 + + +def _require_b200() -> None: + if not torch.cuda.is_available(): + pytest.skip("CUDA required") + if torch.cuda.get_device_capability()[0] < 10: + pytest.skip("LTX2 QKNorm split-RoPE CUDA path is validated on B200") + + +def _ltx2_make_cos_sin( + batch: int, seq_len: int, num_heads: int, head_dim: int +) -> tuple[torch.Tensor, torch.Tensor]: + half_dim = head_dim // 2 + cos = torch.randn( + batch, seq_len, num_heads, half_dim, device="cuda", dtype=torch.bfloat16 + ).transpose(1, 2) + sin = torch.randn( + batch, seq_len, num_heads, half_dim, device="cuda", dtype=torch.bfloat16 + ).transpose(1, 2) + return cos, sin + + +def _apply_split_rotary_ref( + x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor +) -> torch.Tensor: + x_dtype = x.dtype + batch = x.shape[0] + _, num_heads, seq_len, _ = cos.shape + x = x.reshape(batch, seq_len, num_heads, -1).swapaxes(1, 2) + last = x.shape[-1] + half = last // 2 + + split_x = x.reshape(*x.shape[:-1], 2, half) + first_x = split_x[..., :1, :] + second_x = split_x[..., 1:, :] + cos_u = cos.unsqueeze(-2) + sin_u = sin.unsqueeze(-2) + + out = split_x * cos_u + out[..., :1, :].addcmul_(-sin_u, second_x) + out[..., 1:, :].addcmul_(sin_u, first_x) + out = out.reshape(*out.shape[:-2], last) + return out.swapaxes(1, 2).reshape(batch, seq_len, -1).to(dtype=x_dtype) + + +def _ltx2_reference( + q: torch.Tensor, + k: torch.Tensor, + q_cos: torch.Tensor, + q_sin: torch.Tensor, + k_cos: torch.Tensor, + k_sin: torch.Tensor, + q_weight: torch.Tensor, + k_weight: torch.Tensor, + eps: float, +) -> tuple[torch.Tensor, torch.Tensor]: + # rms_norm isn't autocast fp32-preserving, so feed fp32 inputs directly + # to keep the normalized value unrounded until the final RoPE output. + q_norm = F.rms_norm(q.float(), (q.shape[-1],), q_weight.float(), eps) + k_norm = F.rms_norm(k.float(), (k.shape[-1],), k_weight.float(), eps) + q_ref = _apply_split_rotary_ref(q_norm, q_cos, q_sin) + k_ref = _apply_split_rotary_ref(k_norm, k_cos, k_sin) + return q_ref.to(dtype=torch.bfloat16), k_ref.to(dtype=torch.bfloat16) + + +@pytest.mark.parametrize( + "batch,q_seq,k_seq,num_heads,head_dim", + [ + (1, 3, 3, 32, 128), + (1, 5, 2, 32, 64), + (2, 4, 3, 32, 64), + ], +) +def test_ltx2_qknorm_split_rope_matches_torch_exactly( + batch: int, q_seq: int, k_seq: int, num_heads: int, head_dim: int +) -> None: + _require_b200() + torch.cuda.manual_seed(20260630) + hidden = num_heads * head_dim + eps = 1e-6 + q = torch.randn(batch, q_seq, hidden, device="cuda", dtype=torch.bfloat16) + k = torch.randn(batch, k_seq, hidden, device="cuda", dtype=torch.bfloat16) + q_cos, q_sin = _ltx2_make_cos_sin(batch, q_seq, num_heads, head_dim) + k_cos, k_sin = _ltx2_make_cos_sin(batch, k_seq, num_heads, head_dim) + q_weight = torch.randn(hidden, device="cuda", dtype=torch.bfloat16) + k_weight = torch.randn(hidden, device="cuda", dtype=torch.bfloat16) + + assert can_use_ltx2_qknorm_split_rope_cuda( + q, + q_cos, + q_sin, + q_weight, + k, + k_cos, + k_sin, + k_weight, + num_heads=num_heads, + head_dim=head_dim, + ) + + q_ref, k_ref = _ltx2_reference( + q, k, q_cos, q_sin, k_cos, k_sin, q_weight, k_weight, eps + ) + q_out, k_out = ltx2_qknorm_split_rope_cuda( + q, + q_cos, + q_sin, + q_weight, + k, + k_cos, + k_sin, + k_weight, + eps=eps, + num_heads=num_heads, + head_dim=head_dim, + ) + torch.cuda.synchronize() + + torch.testing.assert_close(q_out, q_ref, rtol=0, atol=BF16_FUSED_ATOL) + torch.testing.assert_close(k_out, k_ref, rtol=0, atol=BF16_FUSED_ATOL) + + +def test_ltx2_qknorm_split_rope_rejects_unsupported_inputs() -> None: + _require_b200() + torch.cuda.manual_seed(20260630) + q = torch.randn((1, 3, 4096), device="cuda", dtype=torch.bfloat16) + k = torch.randn_like(q) + q_cos, q_sin = _ltx2_make_cos_sin(1, 3, 32, 128) + q_weight = torch.randn(4096, device="cuda", dtype=torch.bfloat16) + k_weight = torch.randn(4096, device="cuda", dtype=torch.bfloat16) + + assert can_use_ltx2_qknorm_split_rope_cuda( + q, + q_cos, + q_sin, + q_weight, + k, + q_cos, + q_sin, + k_weight, + num_heads=32, + head_dim=128, + ) + assert not can_use_ltx2_qknorm_split_rope_cuda( + q.float(), + q_cos, + q_sin, + q_weight, + k, + q_cos, + q_sin, + k_weight, + num_heads=32, + head_dim=128, + ) + assert not can_use_ltx2_qknorm_split_rope_cuda( + q, + q_cos, + q_sin, + q_weight, + k, + q_cos.transpose(-1, -2), + q_sin, + k_weight, + num_heads=32, + head_dim=128, + ) + + +def test_ltx2_qknorm_split_rope_custom_op_torch_compile_fullgraph() -> None: + _require_b200() + torch.cuda.manual_seed(20260630) + batch, q_seq, k_seq, num_heads, head_dim = 1, 3, 2, 32, 64 + hidden = num_heads * head_dim + q = torch.randn(batch, q_seq, hidden, device="cuda", dtype=torch.bfloat16) + k = torch.randn(batch, k_seq, hidden, device="cuda", dtype=torch.bfloat16) + q_cos, q_sin = _ltx2_make_cos_sin(batch, q_seq, num_heads, head_dim) + k_cos, k_sin = _ltx2_make_cos_sin(batch, k_seq, num_heads, head_dim) + q_weight = torch.randn(hidden, device="cuda", dtype=torch.bfloat16) + k_weight = torch.randn(hidden, device="cuda", dtype=torch.bfloat16) + + def fn(q, k, q_cos, q_sin, k_cos, k_sin, q_weight, k_weight): + return ltx2_qknorm_split_rope_cuda( + q, + q_cos, + q_sin, + q_weight, + k, + k_cos, + k_sin, + k_weight, + eps=1e-6, + num_heads=num_heads, + head_dim=head_dim, + ) + + compiled = torch.compile(fn, fullgraph=True) + q_out, k_out = compiled(q, k, q_cos, q_sin, k_cos, k_sin, q_weight, k_weight) + q_ref, k_ref = _ltx2_reference( + q, k, q_cos, q_sin, k_cos, k_sin, q_weight, k_weight, 1e-6 + ) + torch.cuda.synchronize() + torch.testing.assert_close(q_out, q_ref, rtol=0, atol=BF16_FUSED_ATOL) + torch.testing.assert_close(k_out, k_ref, rtol=0, atol=BF16_FUSED_ATOL) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"])) From bb749c59cae029bf075887aae4eea800c27e1478 Mon Sep 17 00:00:00 2001 From: BBuf <1182563586@qq.com> Date: Tue, 18 Aug 2026 10:17:45 +0800 Subject: [PATCH 4/5] [kernels] Delegate the diffusion platform predicates to the platform's own cache `is_cuda()` / `is_hip()` in `common/platform.py` derived their answer through `platform_key()`, which imports the platform module and walks a chain of `getattr(current_platform, f"is_{name}")()` calls. Those predicates sit in per-call kernel guards -- `try_fused_scaled_residual_add_exact` checks `is_cuda()` on every invocation -- where the code they replaced called `current_platform.is_cuda()`, an `lru_cache`d method. Delegate straight to that method instead, which is what the pre-refactor guards did. `platform_key()` stays uncached on purpose: `select_impl` calls it at module import time, so memoizing it would latch the answer before the platform plugin has resolved. --- .../kernels/ops/diffusion/common/platform.py | 24 ++++++++++++++++--- 1 file changed, 21 insertions(+), 3 deletions(-) diff --git a/python/sglang/kernels/ops/diffusion/common/platform.py b/python/sglang/kernels/ops/diffusion/common/platform.py index 2ac0401208df..01a744ad665c 100644 --- a/python/sglang/kernels/ops/diffusion/common/platform.py +++ b/python/sglang/kernels/ops/diffusion/common/platform.py @@ -31,7 +31,13 @@ def platform_key() -> str: - """Return the live device family: ``cuda``/``hip``/``npu``/``mps``/``musa``/``cpu``.""" + """Return the live device family: ``cuda``/``hip``/``npu``/``mps``/``musa``/``cpu``. + + Deliberately *not* memoized: :func:`select_impl` calls it at module import + time, and latching that first answer would freeze the choice before the + platform plugin has resolved. Use :func:`is_cuda` / :func:`is_hip` on hot + paths -- they delegate straight to the platform's own cached predicates. + """ from sglang.multimodal_gen.runtime.platforms import current_platform for name in ("cuda", "hip", "npu", "mps", "musa"): @@ -41,11 +47,23 @@ def platform_key() -> str: def is_cuda() -> bool: - return platform_key() == "cuda" + """Cheap enough for a per-call kernel guard. + + Delegates to ``current_platform.is_cuda``, which is ``lru_cache``d on the + platform object -- the same call the pre-refactor guards made. Going + through :func:`platform_key` instead would add an import plus a chain of + ``getattr`` lookups to every fused-elementwise dispatch. + """ + from sglang.multimodal_gen.runtime.platforms import current_platform + + return current_platform.is_cuda() def is_hip() -> bool: - return platform_key() == "hip" + """See :func:`is_cuda`; delegates to the platform's cached predicate.""" + from sglang.multimodal_gen.runtime.platforms import current_platform + + return current_platform.is_hip() def has_triton() -> bool: From b7740b98374d1c07a29503437aff71cb7509a3ed Mon Sep 17 00:00:00 2001 From: BBuf <1182563586@qq.com> Date: Tue, 18 Aug 2026 11:16:52 +0800 Subject: [PATCH 5/5] [diffusion] Fix the Sana LN+modulate test's cross-stream race `test_sana_fused_ln_modulate_is_bit_exact` fills x/scale/shift on the default stream, then enters a fresh `torch.cuda.Stream()` and reads them there without making the side stream wait. The fused kernel can therefore read a half-written tensor; the gate's first-sight `torch.equal` check then fails, and because a mismatch disables `_SANA_LN_MOD` *permanently*, every later parametrization also stops registering its signature -- so one lost race fails three of the four cases. It only loses the race when the GPU is contended, which is why it is invisible on an idle box and intermittent on shared CI runners. Reproduced under a synthetic GEMM load on one H200: main's test, unpatched: 5 / 10 runs failed with side.wait_stream(): 0 / 10 runs failed and the merged suite 0/8 under the same load, full suite 2580 passed. The production path is unaffected: the BCG runner enters its side stream and produces its own tensors there, so only this hand-rolled test needed the `wait_stream` pair that PyTorch's side-stream contract requires. --- .../kernels/ops/diffusion/test_model_fast_paths.py | 13 +++++++++++-- 1 file changed, 11 insertions(+), 2 deletions(-) diff --git a/test/registered/kernels/ops/diffusion/test_model_fast_paths.py b/test/registered/kernels/ops/diffusion/test_model_fast_paths.py index c78de2ca99b2..31238f9d1242 100644 --- a/test/registered/kernels/ops/diffusion/test_model_fast_paths.py +++ b/test/registered/kernels/ops/diffusion/test_model_fast_paths.py @@ -322,11 +322,20 @@ def test_sana_fused_ln_modulate_is_bit_exact(shape, nmod, transposed): n_sigs = len(sana._SANA_LN_MOD.verified_sigs) _sana_ln_modulate(norm, x, scale, shift) assert len(sana._SANA_LN_MOD.verified_sigs) == n_sigs - # the fusion engages on non-default streams (the BCG warmup/capture path) - with torch.cuda.stream(torch.cuda.Stream()): + # The fusion engages on non-default streams (the BCG warmup/capture path). + # x/scale/shift were filled on the default stream, so the side stream must + # wait for that work before reading them -- without this the fused kernel + # can read a half-written tensor, the first-sight torch.equal check fails, + # and the gate disables itself *permanently*, which then breaks every later + # parametrization too. It only loses the race when the GPU is contended, + # which is why it shows up on shared CI runners and not on an idle box. + side = torch.cuda.Stream() + side.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(side): out = _sana_ln_modulate(norm, x, scale, shift) assert len(sana._SANA_LN_MOD.verified_sigs) == n_sigs + 1 # verified out2 = _sana_ln_modulate(norm, x, scale, shift) # verified-sig lane + torch.cuda.current_stream().wait_stream(side) torch.cuda.synchronize() assert torch.equal(out, _sana_eager_ln_modulate(norm, x, scale, shift)) assert torch.equal(out2, out) and not sana._SANA_LN_MOD.disabled