[kernels] Reorganize ops/diffusion by operator domain behind a lazy facade (#35114)
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5
parent
7605529bdf
commit
ae6945e112
@@ -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`<br>`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`<br>`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`<br>`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`<br>`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`<br>`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`<br>`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`<br>`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`<br>`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`<br>`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`<br>`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`<br>`tanh(gate) * rmsnorm(x)` | `python/sglang/kernels/ops/diffusion/cutedsl/norm_tanh_mul_add_norm_scale.py`<br>`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`<br>`residual + tanh(gate) * rmsnorm(x)`<br>`ffn_norm1(x) * scale_mlp` | `python/sglang/kernels/ops/diffusion/cutedsl/norm_tanh_mul_add_norm_scale.py`<br>`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`<br>`get_ada_values`<br>`scale_shift_table + timestep.reshape` | `python/sglang/kernels/ops/diffusion/triton/ltx2_ada_values.py`<br>`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`<br>`residual_gate_add`<br>`residual + update * gate` | `python/sglang/kernels/ops/diffusion/residual_gate_add.py`<br>`python/sglang/kernels/jit/csrc/diffusion/residual_gate_add.cuh`<br>`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`<br>`get_ada_values`<br>`scale_shift_table + timestep.reshape` | `python/sglang/kernels/ops/diffusion/modulate/ltx2_ada_values_triton.py`<br>`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`<br>`residual_gate_add`<br>`residual + update * gate` | `python/sglang/kernels/ops/diffusion/modulate/residual_gate_add_jit.py`<br>`python/sglang/kernels/jit/csrc/diffusion/residual_gate_add.cuh`<br>`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`<br>`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
|
||||
|
||||
@@ -107,6 +107,7 @@ BACKEND_METHODS: Dict[KernelBackend, str] = {
|
||||
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",
|
||||
|
||||
@@ -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_<op>(...):
|
||||
out = <op>(...)
|
||||
else:
|
||||
out = <reference chain>
|
||||
```
|
||||
|
||||
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`.
|
||||
@@ -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 @@ from sglang.kernels.spec import (
|
||||
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).",
|
||||
)
|
||||
)
|
||||
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).",
|
||||
)
|
||||
# ---------------------------------------------------------------------------
|
||||
# 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).",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
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
|
||||
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,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
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
|
||||
)
|
||||
# ---------------------------------------------------------------------------
|
||||
# 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 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 __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
|
||||
|
||||
value = getattr(import_module(f"{__name__}.{module}"), name)
|
||||
globals()[name] = value # cache; later lookups skip __getattr__ entirely
|
||||
return value
|
||||
|
||||
|
||||
__all__ = [
|
||||
"apply_group_norm_silu",
|
||||
"residual_gate_add",
|
||||
"fused_inplace_qknorm_rope",
|
||||
]
|
||||
def __dir__() -> list[str]:
|
||||
return sorted(set(globals()) | set(_EXPORTS))
|
||||
|
||||
|
||||
__all__ = sorted(_EXPORTS)
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""Activation-function fusions (SiLU/GLU/GELU epilogues)."""
|
||||
+1
-1
@@ -7,7 +7,7 @@ import torch
|
||||
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
|
||||
+1
-1
@@ -20,7 +20,7 @@ import torch
|
||||
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
|
||||
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""Diffusion attention kernels: sparse linear attention and gated delta-net."""
|
||||
+1
-1
@@ -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,
|
||||
)
|
||||
|
||||
+1
-1
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
"""
|
||||
+4
-4
@@ -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 torch import Tensor
|
||||
|
||||
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,
|
||||
@@ -0,0 +1,107 @@
|
||||
"""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``.
|
||||
|
||||
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"):
|
||||
if getattr(current_platform, f"is_{name}")():
|
||||
return name
|
||||
return "cpu"
|
||||
|
||||
|
||||
def is_cuda() -> bool:
|
||||
"""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:
|
||||
"""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:
|
||||
"""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_<kind>`` 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)
|
||||
@@ -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`.
|
||||
"""
|
||||
+1
-1
@@ -13,7 +13,7 @@ from typing import Tuple
|
||||
|
||||
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
|
||||
@@ -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<torch::Tensor> 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<float>(), F.data_ptr<int>(), 0,
|
||||
(INT64 *)z_min.data_ptr<int64_t>(),
|
||||
occlusion_truncation, width, height,
|
||||
num_vertices, num_faces, i);
|
||||
}
|
||||
} else {
|
||||
for (int i = 0; i < num_faces; ++i)
|
||||
rasterizeImagecoordsKernelCPU(
|
||||
V.data_ptr<float>(), F.data_ptr<int>(), D.data_ptr<float>(),
|
||||
(INT64 *)z_min.data_ptr<int64_t>(), 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<float>(), F.data_ptr<int>(), findices.data_ptr<int>(),
|
||||
(INT64 *)z_min.data_ptr<int64_t>(), width, height, num_vertices,
|
||||
num_faces, barycentric.data_ptr<float>(), i);
|
||||
|
||||
return {findices, barycentric};
|
||||
}
|
||||
|
||||
std::vector<torch::Tensor>
|
||||
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");
|
||||
}
|
||||
@@ -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 <ATen/ATen.h>
|
||||
#include <torch/extension.h>
|
||||
#include <vector>
|
||||
|
||||
#ifdef CUDA_ENABLED
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#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<torch::Tensor> rasterize_image_gpu(torch::Tensor V, torch::Tensor F,
|
||||
torch::Tensor D, int width,
|
||||
int height,
|
||||
float occlusion_truncation,
|
||||
int use_depth_prior);
|
||||
|
||||
#endif
|
||||
@@ -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<torch::Tensor> 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<float>(), F.data_ptr<int>(), 0,
|
||||
(INT64 *)z_min.data_ptr<int64_t>(), occlusion_truncation, width, height,
|
||||
num_vertices, num_faces);
|
||||
} else {
|
||||
rasterizeImagecoordsKernelGPU<<<(num_faces + 255) / 256, 256, 0,
|
||||
at::cuda::getCurrentCUDAStream()>>>(
|
||||
V.data_ptr<float>(), F.data_ptr<int>(), D.data_ptr<float>(),
|
||||
(INT64 *)z_min.data_ptr<int64_t>(), 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<float>(), F.data_ptr<int>(), findices.data_ptr<int>(),
|
||||
(INT64 *)z_min.data_ptr<int64_t>(), width, height, num_vertices,
|
||||
num_faces, barycentric.data_ptr<float>());
|
||||
|
||||
return {findices, barycentric};
|
||||
}
|
||||
+2
-3
@@ -1,5 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import shutil
|
||||
import sys
|
||||
@@ -8,9 +9,7 @@ from typing import Any, Sequence
|
||||
|
||||
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:
|
||||
+1
-1
@@ -13,7 +13,7 @@ from typing import Tuple
|
||||
|
||||
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
|
||||
@@ -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 <algorithm>
|
||||
#include <cmath>
|
||||
#include <pybind11/numpy.h>
|
||||
#include <pybind11/pybind11.h>
|
||||
#include <pybind11/stl.h>
|
||||
#include <queue>
|
||||
#include <torch/extension.h>
|
||||
#include <vector>
|
||||
|
||||
namespace py = pybind11;
|
||||
using namespace std;
|
||||
|
||||
std::pair<py::array_t<float>, py::array_t<uint8_t>>
|
||||
meshVerticeInpaint_smooth(py::array_t<float> texture, py::array_t<uint8_t> mask,
|
||||
py::array_t<float> vtx_pos, py::array_t<float> vtx_uv,
|
||||
py::array_t<int> pos_idx, py::array_t<int> 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<float *>(texture_buf.ptr);
|
||||
uint8_t *mask_ptr = static_cast<uint8_t *>(mask_buf.ptr);
|
||||
|
||||
int vtx_num = vtx_pos_buf.shape[0];
|
||||
float *vtx_pos_ptr = static_cast<float *>(vtx_pos_buf.ptr);
|
||||
float *vtx_uv_ptr = static_cast<float *>(vtx_uv_buf.ptr);
|
||||
int *pos_idx_ptr = static_cast<int *>(pos_idx_buf.ptr);
|
||||
int *uv_idx_ptr = static_cast<int *>(uv_idx_buf.ptr);
|
||||
|
||||
vector<float> vtx_mask(vtx_num, 0.0f);
|
||||
vector<vector<float>> vtx_color(vtx_num,
|
||||
vector<float>(texture_channel, 0.0f));
|
||||
vector<int> uncolored_vtxs;
|
||||
|
||||
vector<vector<int>> 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<float> sum_color(texture_channel, 0.0f);
|
||||
float total_weight = 0.0f;
|
||||
|
||||
array<float, 3> 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<float, 3> 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<float> new_texture(texture_buf.size);
|
||||
py::array_t<uint8_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<float *>(new_texture_buf.ptr);
|
||||
uint8_t *new_mask_ptr = static_cast<uint8_t *>(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<float>, py::array_t<uint8_t>>
|
||||
meshVerticeInpaint(py::array_t<float> texture, py::array_t<uint8_t> mask,
|
||||
py::array_t<float> vtx_pos, py::array_t<float> vtx_uv,
|
||||
py::array_t<int> pos_idx, py::array_t<int> 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");
|
||||
}
|
||||
@@ -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.
|
||||
"""
|
||||
@@ -0,0 +1 @@
|
||||
"""adaLN modulation: ``x * (1 + scale) + shift``, gating, and timestep conditioning."""
|
||||
+1
-1
@@ -4,7 +4,7 @@ import torch
|
||||
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
|
||||
+16
-25
@@ -2,8 +2,13 @@ import torch
|
||||
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 fuse_scale_shift_kernel(
|
||||
|
||||
# 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"),
|
||||
)
|
||||
@@ -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.
|
||||
"""
|
||||
@@ -0,0 +1 @@
|
||||
"""CuTe-DSL building blocks shared by the CUTLASS-backed norm fusions."""
|
||||
+1
-1
@@ -5,7 +5,7 @@ import cutlass.cute as cute
|
||||
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,
|
||||
)
|
||||
+1
-1
@@ -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,
|
||||
)
|
||||
|
||||
+47
-21
@@ -16,8 +16,9 @@ kernel for the channels_last VAE decoder fast path with a different contract:
|
||||
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",
|
||||
]
|
||||
+1
-1
@@ -47,7 +47,7 @@ import triton # type: ignore
|
||||
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,
|
||||
+11
-12
@@ -5,7 +5,7 @@ import triton # type: ignore
|
||||
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"),
|
||||
)
|
||||
+7
-13
@@ -3,7 +3,7 @@ import triton # type: ignore
|
||||
import triton.language as tl # type: ignore
|
||||
|
||||
from sglang.kernels.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"),
|
||||
)
|
||||
+1
-1
@@ -55,7 +55,7 @@ import torch
|
||||
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,
|
||||
+4
-4
@@ -5,12 +5,12 @@ import cutlass
|
||||
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,
|
||||
)
|
||||
|
||||
+5
-6
@@ -11,9 +11,8 @@ the same dtype boundary as eager ``WanRMS_norm.forward`` (including the
|
||||
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()
|
||||
@@ -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<torch::Tensor> 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<float>(), F.data_ptr<int>(), 0,
|
||||
(INT64*)z_min.data_ptr<int64_t>(), occlusion_truncation, width, height, num_vertices, num_faces, i);
|
||||
}
|
||||
} else {
|
||||
for (int i = 0; i < num_faces; ++i)
|
||||
rasterizeImagecoordsKernelCPU(V.data_ptr<float>(), F.data_ptr<int>(), D.data_ptr<float>(),
|
||||
(INT64*)z_min.data_ptr<int64_t>(), 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<float>(), F.data_ptr<int>(),
|
||||
findices.data_ptr<int>(), (INT64*)z_min.data_ptr<int64_t>(), width, height, num_vertices, num_faces, barycentric.data_ptr<float>(), i);
|
||||
|
||||
return {findices, barycentric};
|
||||
}
|
||||
|
||||
std::vector<torch::Tensor> 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");
|
||||
}
|
||||
@@ -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 <torch/extension.h>
|
||||
#include <vector>
|
||||
#include <ATen/ATen.h>
|
||||
|
||||
#ifdef CUDA_ENABLED
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#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<torch::Tensor> rasterize_image_gpu(torch::Tensor V, torch::Tensor F, torch::Tensor D,
|
||||
int width, int height, float occlusion_truncation, int use_depth_prior);
|
||||
|
||||
#endif
|
||||
@@ -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<torch::Tensor> 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<float>(), F.data_ptr<int>(), 0,
|
||||
(INT64*)z_min.data_ptr<int64_t>(), occlusion_truncation, width, height, num_vertices, num_faces);
|
||||
} else {
|
||||
rasterizeImagecoordsKernelGPU<<<(num_faces+255)/256,256,0,at::cuda::getCurrentCUDAStream()>>>(V.data_ptr<float>(), F.data_ptr<int>(), D.data_ptr<float>(),
|
||||
(INT64*)z_min.data_ptr<int64_t>(), 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<float>(), F.data_ptr<int>(),
|
||||
findices.data_ptr<int>(), (INT64*)z_min.data_ptr<int64_t>(), width, height, num_vertices, num_faces, barycentric.data_ptr<float>());
|
||||
|
||||
return {findices, barycentric};
|
||||
}
|
||||
@@ -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 <vector>
|
||||
#include <queue>
|
||||
#include <cmath>
|
||||
#include <algorithm>
|
||||
#include <torch/extension.h>
|
||||
#include <pybind11/pybind11.h>
|
||||
#include <pybind11/numpy.h>
|
||||
#include <pybind11/stl.h>
|
||||
|
||||
namespace py = pybind11;
|
||||
using namespace std;
|
||||
|
||||
std::pair<py::array_t<float>,
|
||||
py::array_t<uint8_t>> meshVerticeInpaint_smooth(py::array_t<float> texture,
|
||||
py::array_t<uint8_t> mask,
|
||||
py::array_t<float> vtx_pos, py::array_t<float> vtx_uv,
|
||||
py::array_t<int> pos_idx, py::array_t<int> 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<float*>(texture_buf.ptr);
|
||||
uint8_t* mask_ptr = static_cast<uint8_t*>(mask_buf.ptr);
|
||||
|
||||
int vtx_num = vtx_pos_buf.shape[0];
|
||||
float* vtx_pos_ptr = static_cast<float*>(vtx_pos_buf.ptr);
|
||||
float* vtx_uv_ptr = static_cast<float*>(vtx_uv_buf.ptr);
|
||||
int* pos_idx_ptr = static_cast<int*>(pos_idx_buf.ptr);
|
||||
int* uv_idx_ptr = static_cast<int*>(uv_idx_buf.ptr);
|
||||
|
||||
vector<float> vtx_mask(vtx_num, 0.0f);
|
||||
vector<vector<float>> vtx_color(vtx_num, vector<float>(texture_channel, 0.0f));
|
||||
vector<int> uncolored_vtxs;
|
||||
|
||||
vector<vector<int>> 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<float> sum_color(texture_channel, 0.0f);
|
||||
float total_weight = 0.0f;
|
||||
|
||||
array<float, 3> 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<float, 3> 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<float> new_texture(texture_buf.size);
|
||||
py::array_t<uint8_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<float*>(new_texture_buf.ptr);
|
||||
uint8_t* new_mask_ptr = static_cast<uint8_t*>(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<float>, py::array_t<uint8_t>> meshVerticeInpaint(py::array_t<float> texture,
|
||||
py::array_t<uint8_t> mask,
|
||||
py::array_t<float> vtx_pos, py::array_t<float> vtx_uv,
|
||||
py::array_t<int> pos_idx, py::array_t<int> 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");
|
||||
}
|
||||
@@ -0,0 +1 @@
|
||||
"""Rotary embeddings and the QK-norm chains fused around them."""
|
||||
+1
-1
@@ -29,7 +29,7 @@ import torch
|
||||
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
|
||||
|
||||
|
||||
+7
-15
@@ -2,7 +2,7 @@ import torch
|
||||
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"),
|
||||
)
|
||||
@@ -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.
|
||||
"""
|
||||
+4
-4
@@ -4,7 +4,7 @@ Adaln-style DiT blocks (Ideogram 4) spend four elementwise chains per block on
|
||||
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 @@ from importlib import import_module
|
||||
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,
|
||||
)
|
||||
|
||||
+4
-4
@@ -32,7 +32,7 @@ import torch
|
||||
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:
|
||||
+2
-2
@@ -20,7 +20,7 @@ import torch
|
||||
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."""
|
||||
+1
-1
@@ -8,7 +8,7 @@ from functools import cache
|
||||
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__)
|
||||
|
||||
+3
-3
@@ -21,11 +21,11 @@ from __future__ import annotations
|
||||
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:
|
||||
@@ -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)
|
||||
|
||||
+27
-27
@@ -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`.
|
||||
|
||||
+1
-4
@@ -36,10 +36,7 @@ logger = init_logger(__name__)
|
||||
|
||||
|
||||
# ==================================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:
|
||||
|
||||
@@ -12,7 +12,7 @@ import torch.nn as nn
|
||||
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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -10,7 +10,7 @@ from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
|
||||
_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,
|
||||
)
|
||||
|
||||
@@ -11,12 +11,12 @@ import torch
|
||||
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 @@ if _is_xpu:
|
||||
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 @@ class _ScaleResidualNormScaleShift(CustomOp):
|
||||
)
|
||||
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 @@ class _ScaleResidualNormScaleShift(CustomOp):
|
||||
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 @@ class _NormScaleShift(CustomOp):
|
||||
)
|
||||
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 @@ class _NormScaleShift(CustomOp):
|
||||
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,
|
||||
)
|
||||
|
||||
@@ -18,13 +18,15 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
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
|
||||
|
||||
@@ -5,7 +5,7 @@ from typing import Optional, Tuple
|
||||
import torch
|
||||
|
||||
from sglang.kernels.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
|
||||
|
||||
@@ -9,10 +9,7 @@ import torch.distributed._functional_collectives as ft_c
|
||||
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,
|
||||
|
||||
@@ -22,9 +22,7 @@ from diffusers.models.embeddings import (
|
||||
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
|
||||
|
||||
@@ -14,7 +14,7 @@ import torch
|
||||
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,
|
||||
)
|
||||
|
||||
@@ -22,21 +22,17 @@ from diffusers.models.embeddings import TimestepEmbedding, Timesteps
|
||||
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_scale_residual_rmsnorm_scale_shift,
|
||||
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,
|
||||
can_use_fused_scale_residual_rmsnorm_scale_shift,
|
||||
flashinfer_rmsnorm_diagnostic_hint,
|
||||
fused_rmsnorm_scale_shift_bitexact,
|
||||
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,
|
||||
|
||||
@@ -28,25 +28,21 @@ from diffusers.models.normalization import (
|
||||
)
|
||||
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 @@ class FluxGELU(nn.Module):
|
||||
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 @@ class FluxFusedGELUProj(nn.Module):
|
||||
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 @@ class FluxSingleTransformerBlock(nn.Module):
|
||||
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(
|
||||
|
||||
@@ -21,15 +21,13 @@ from diffusers.models.attention import AttentionModuleMixin
|
||||
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 (
|
||||
|
||||
@@ -18,23 +18,19 @@ import torch
|
||||
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 @@ class GlmImageGELU(nn.Module):
|
||||
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
|
||||
)
|
||||
|
||||
@@ -8,22 +8,16 @@ import numpy as np
|
||||
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,
|
||||
mark_fused_gelu_site,
|
||||
)
|
||||
from sglang.kernels.ops.diffusion.hunyuan_qknorm import (
|
||||
mark_hunyuan_qknorm_site,
|
||||
try_hunyuan_qknorm,
|
||||
)
|
||||
from sglang.kernels.ops.diffusion.triton.hunyuan_qkv_pack import (
|
||||
hunyuan_qkv_rope_pack,
|
||||
mark_fused_gelu_site,
|
||||
mark_hunyuan_qknorm_site,
|
||||
tensors_equal,
|
||||
try_hunyuan_qknorm,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.models.dits import HunyuanVideoConfig
|
||||
from sglang.multimodal_gen.configs.models.fsdp import (
|
||||
@@ -115,7 +109,7 @@ class HunyuanMLP(MLP):
|
||||
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)
|
||||
|
||||
@@ -7,24 +7,18 @@ import torch
|
||||
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
|
||||
|
||||
@@ -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 @@ class Attention(nn.Module):
|
||||
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)
|
||||
|
||||
@@ -10,28 +10,22 @@ import torch
|
||||
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 @@ class LTX2FeedForward(nn.Module):
|
||||
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)
|
||||
|
||||
@@ -20,11 +20,9 @@ from safetensors.torch import safe_open
|
||||
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_,
|
||||
|
||||
@@ -14,8 +14,8 @@ from diffusers.models.embeddings import TimestepEmbedding, Timesteps
|
||||
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 @@ class QwenImageGELU(nn.Module):
|
||||
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
|
||||
)
|
||||
|
||||
@@ -5,18 +5,16 @@ import torch.nn as nn
|
||||
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
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user