diff --git a/python/sglang/jit_kernel/csrc/diffusion/norm_scale_shift.cuh b/python/sglang/jit_kernel/csrc/diffusion/norm_scale_shift.cuh index dbc9af09b..c18e32606 100644 --- a/python/sglang/jit_kernel/csrc/diffusion/norm_scale_shift.cuh +++ b/python/sglang/jit_kernel/csrc/diffusion/norm_scale_shift.cuh @@ -8,6 +8,9 @@ // // All other public-op inputs fall back to the existing CuTe-DSL implementation // from the Python dispatcher. +// +// Developed with MIT HAN Lab Kernel Design Agents: +// https://github.com/mit-han-lab/kernel-design-agents #pragma once diff --git a/python/sglang/jit_kernel/csrc/diffusion/residual_gate_add.cuh b/python/sglang/jit_kernel/csrc/diffusion/residual_gate_add.cuh new file mode 100644 index 000000000..693641880 --- /dev/null +++ b/python/sglang/jit_kernel/csrc/diffusion/residual_gate_add.cuh @@ -0,0 +1,322 @@ +// CUDA fast path for diffusion residual-gate elementwise updates. +// +// Implements: +// out = residual + update * gate +// +// The production shapes come from LTX-2.3 HQ residual/gate updates. This is +// intentionally narrow: contiguous residual/update/out tensors, with either a +// full contiguous gate or a row-broadcast [1, 1, D] gate. +// +// Developed with MIT HAN Lab Kernel Design Agents: +// https://github.com/mit-han-lab/kernel-design-agents + +#pragma once + +#include // For host dtype helpers and TensorView metadata +#include // For RuntimeCheck and div_ceil + +#include // For dtype_trait conversions +#include // For LaunchKernel and CUDA dtype aliases + +#include + +namespace sglang_residual_gate_add { + +namespace { + +constexpr int kBlockSize = 256; +constexpr int kBcastRowsPerBlock = 4; +constexpr int kBcastColsVecPerBlock = 256; +constexpr int64_t kMaxGrid = 65535; + +enum class GateMode : int { kFull = 0, kBcastRow = 1 }; + +inline const char* data_ptr(const tvm::ffi::TensorView& t) { + return static_cast(t.data_ptr()) + t.byte_offset(); +} + +inline char* mutable_data_ptr(const tvm::ffi::TensorView& t) { + return static_cast(t.data_ptr()) + t.byte_offset(); +} + +inline bool aligned16(const void* p) { + return (reinterpret_cast(p) & 0xF) == 0; +} + +inline int64_t numel(const tvm::ffi::TensorView& t) { + int64_t n = 1; + for (int i = 0; i < t.ndim(); ++i) { + n *= t.size(i); + } + return n; +} + +inline int64_t grid_for(int64_t total) { + int64_t grid = host::div_ceil(total, static_cast(kBlockSize)); + if (grid < 1) { + grid = 1; + } + if (grid > kMaxGrid) { + grid = kMaxGrid; + } + return grid; +} + +inline bool is_dense_contiguous(const tvm::ffi::TensorView& t) { + int64_t expected = 1; + for (int i = t.ndim() - 1; i >= 0; --i) { + if (t.size(i) == 1) { + continue; + } + if (t.stride(i) != expected) { + return false; + } + expected *= t.size(i); + } + return true; +} + +template +inline void check_dtype(const tvm::ffi::TensorView& t) { + host::RuntimeCheck(host::is_type(t.dtype()), "unexpected dtype for residual_gate_add"); +} + +template +__device__ __forceinline__ float to_float(T v) { + return static_cast(v); +} + +template <> +__device__ __forceinline__ float to_float(fp16_t v) { + return __half2float(v); +} + +template <> +__device__ __forceinline__ float to_float(bf16_t v) { + return __bfloat162float(v); +} + +template +__device__ __forceinline__ T residual_gate_value(T residual, T update, T gate) { + const T product = dtype_trait::from(to_float(update) * to_float(gate)); + return dtype_trait::from(to_float(residual) + to_float(product)); +} + +template +union Vec16 { + static constexpr int kElems = 16 / sizeof(T); + uint4 raw; + T elems[kElems]; +}; + +template +__global__ void residual_gate_add_vec_kernel( + const T* __restrict__ residual, + const T* __restrict__ update, + const T* __restrict__ gate, + T* __restrict__ out, + int64_t n_vec) { + const int64_t stride = static_cast(gridDim.x) * blockDim.x; + for (int64_t v = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; v < n_vec; v += stride) { + const Vec16 r{.raw = reinterpret_cast(residual)[v]}; + const Vec16 u{.raw = reinterpret_cast(update)[v]}; + const Vec16 g{.raw = reinterpret_cast(gate)[v]}; + + Vec16 o; +#pragma unroll + for (int i = 0; i < kVec; ++i) { + o.elems[i] = residual_gate_value(r.elems[i], u.elems[i], g.elems[i]); + } + reinterpret_cast(out)[v] = o.raw; + } +} + +template +__global__ void residual_gate_add_bcast_row_tile_kernel( + const T* __restrict__ residual, + const T* __restrict__ update, + const T* __restrict__ gate, + T* __restrict__ out, + int64_t rows, + int64_t row_vec) { + const int64_t col_vec = static_cast(blockIdx.x) * kBcastColsVecPerBlock + threadIdx.x; + if (col_vec >= row_vec) { + return; + } + + const Vec16 g{.raw = SGLANG_LDG(reinterpret_cast(gate) + col_vec)}; + + // Grid-stride over row tiles so the launch stays valid even when the number + // of row tiles exceeds the gridDim.y hardware limit. + const int64_t row_tile_stride = static_cast(gridDim.y) * kBcastRowsPerBlock; + for (int64_t row_base = static_cast(blockIdx.y) * kBcastRowsPerBlock; row_base < rows; + row_base += row_tile_stride) { +#pragma unroll + for (int row_offset = 0; row_offset < kBcastRowsPerBlock; ++row_offset) { + const int64_t row = row_base + row_offset; + if (row < rows) { + const int64_t v = row * row_vec + col_vec; + const Vec16 r{.raw = reinterpret_cast(residual)[v]}; + const Vec16 u{.raw = reinterpret_cast(update)[v]}; + + Vec16 o; +#pragma unroll + for (int i = 0; i < kVec; ++i) { + o.elems[i] = residual_gate_value(r.elems[i], u.elems[i], g.elems[i]); + } + reinterpret_cast(out)[v] = o.raw; + } + } + } +} + +template +__global__ void residual_gate_add_scalar_kernel( + const T* __restrict__ residual, + const T* __restrict__ update, + const T* __restrict__ gate, + T* __restrict__ out, + int64_t begin, + int64_t total, + int64_t D) { + const int64_t stride = static_cast(gridDim.x) * blockDim.x; + for (int64_t i = begin + static_cast(blockIdx.x) * blockDim.x + threadIdx.x; i < total; i += stride) { + const T gate_value = kGate == GateMode::kFull ? gate[i] : SGLANG_LDG(gate + (i % D)); + out[i] = residual_gate_value(residual[i], update[i], gate_value); + } +} + +template +inline void launch_residual_gate_add( + const tvm::ffi::TensorView& out, + const tvm::ffi::TensorView& residual, + const tvm::ffi::TensorView& update, + const tvm::ffi::TensorView& gate, + GateMode mode) { + const int64_t total = numel(residual); + if (total == 0) { + return; + } + + const int64_t D = residual.size(residual.ndim() - 1); + const T* residual_ptr = reinterpret_cast(data_ptr(residual)); + const T* update_ptr = reinterpret_cast(data_ptr(update)); + const T* gate_ptr = reinterpret_cast(data_ptr(gate)); + T* out_ptr = reinterpret_cast(mutable_data_ptr(out)); + constexpr int kVec = 16 / sizeof(T); + + const bool vec_ok = aligned16(residual_ptr) && aligned16(update_ptr) && aligned16(gate_ptr) && aligned16(out_ptr) && + (D % kVec == 0) && (mode == GateMode::kBcastRow || total % kVec == 0); + + int64_t done = 0; + if (vec_ok) { + const int64_t n_vec = total / kVec; + const int64_t row_vec = D / kVec; + if (mode == GateMode::kFull) { + host::LaunchKernel(static_cast(grid_for(n_vec)), kBlockSize, out.device())( + residual_gate_add_vec_kernel, residual_ptr, update_ptr, gate_ptr, out_ptr, n_vec); + } else { + const int64_t rows = total / D; + const int64_t col_blocks = host::div_ceil(row_vec, static_cast(kBcastColsVecPerBlock)); + const int64_t row_tiles = host::div_ceil(rows, static_cast(kBcastRowsPerBlock)); + const int64_t row_blocks = row_tiles > kMaxGrid ? kMaxGrid : row_tiles; + host::LaunchKernel( + dim3(static_cast(col_blocks), static_cast(row_blocks)), + dim3(kBcastColsVecPerBlock), + out.device())( + residual_gate_add_bcast_row_tile_kernel, residual_ptr, update_ptr, gate_ptr, out_ptr, rows, row_vec); + } + done = n_vec * kVec; + } + + if (done < total) { + if (mode == GateMode::kFull) { + host::LaunchKernel(static_cast(grid_for(total - done)), kBlockSize, out.device())( + residual_gate_add_scalar_kernel, + residual_ptr, + update_ptr, + gate_ptr, + out_ptr, + done, + total, + D); + } else { + host::LaunchKernel(static_cast(grid_for(total - done)), kBlockSize, out.device())( + residual_gate_add_scalar_kernel, + residual_ptr, + update_ptr, + gate_ptr, + out_ptr, + done, + total, + D); + } + } +} + +template +inline GateMode validate_residual_gate_add( + const tvm::ffi::TensorView& out, + const tvm::ffi::TensorView& residual, + const tvm::ffi::TensorView& update, + const tvm::ffi::TensorView& gate) { + check_dtype(out); + check_dtype(residual); + check_dtype(update); + check_dtype(gate); + host::RuntimeCheck(residual.device().device_type == kDLCUDA, "residual must be CUDA"); + host::RuntimeCheck(update.device().device_type == kDLCUDA, "update must be CUDA"); + host::RuntimeCheck(gate.device().device_type == kDLCUDA, "gate must be CUDA"); + host::RuntimeCheck(out.device().device_type == kDLCUDA, "out must be CUDA"); + host::RuntimeCheck( + residual.device().device_id == update.device().device_id && + residual.device().device_id == gate.device().device_id && + residual.device().device_id == out.device().device_id, + "residual/update/gate/out must be on the same CUDA device"); + host::RuntimeCheck(residual.ndim() >= 2, "residual must be at least 2D"); + host::RuntimeCheck(update.ndim() == residual.ndim(), "update rank must match residual"); + host::RuntimeCheck(out.ndim() == residual.ndim(), "out rank must match residual"); + for (int i = 0; i < residual.ndim(); ++i) { + host::RuntimeCheck(update.size(i) == residual.size(i), "update shape must match residual"); + host::RuntimeCheck(out.size(i) == residual.size(i), "out shape must match residual"); + } + host::RuntimeCheck(is_dense_contiguous(residual), "residual must be contiguous"); + host::RuntimeCheck(is_dense_contiguous(update), "update must be contiguous"); + host::RuntimeCheck(is_dense_contiguous(out), "out must be contiguous"); + host::RuntimeCheck(is_dense_contiguous(gate), "gate must be contiguous"); + host::RuntimeCheck(data_ptr(out) != data_ptr(residual), "out must not alias residual"); + host::RuntimeCheck(data_ptr(out) != data_ptr(update), "out must not alias update"); + host::RuntimeCheck(data_ptr(out) != data_ptr(gate), "out must not alias gate"); + + const int D_dim = residual.ndim() - 1; + const int row_dim = residual.ndim() - 2; + host::RuntimeCheck(gate.ndim() == residual.ndim(), "gate rank must match residual"); + host::RuntimeCheck(gate.size(D_dim) == residual.size(D_dim), "gate last dim must match residual"); + + bool full_gate = true; + for (int i = 0; i < residual.ndim(); ++i) { + full_gate = full_gate && gate.size(i) == residual.size(i); + } + if (full_gate) { + return GateMode::kFull; + } + + host::RuntimeCheck(gate.size(row_dim) == 1, "broadcast gate row dim must be 1"); + for (int i = 0; i < D_dim; ++i) { + host::RuntimeCheck(gate.size(i) == 1, "broadcast gate leading dims must be 1"); + } + return GateMode::kBcastRow; +} + +} // namespace + +template +struct ResidualGateAddKernel { + static void + run(tvm::ffi::TensorView out, tvm::ffi::TensorView residual, tvm::ffi::TensorView update, tvm::ffi::TensorView gate) { + const GateMode mode = validate_residual_gate_add(out, residual, update, gate); + launch_residual_gate_add(out, residual, update, gate, mode); + } +}; + +} // namespace sglang_residual_gate_add diff --git a/python/sglang/jit_kernel/diffusion/residual_gate_add.py b/python/sglang/jit_kernel/diffusion/residual_gate_add.py new file mode 100644 index 000000000..9933e3d2a --- /dev/null +++ b/python/sglang/jit_kernel/diffusion/residual_gate_add.py @@ -0,0 +1,85 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +import torch + +from sglang.jit_kernel.utils import cache_once, load_jit, make_cpp_args +from sglang.srt.utils.custom_op import register_custom_op + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +_SUPPORTED_DTYPES = (torch.float16, torch.bfloat16, torch.float32) + + +@cache_once +def _jit_residual_gate_add_module(dtype: torch.dtype) -> Module: + args = make_cpp_args(dtype) + return load_jit( + "diffusion_residual_gate_add", + *args, + cuda_files=["diffusion/residual_gate_add.cuh"], + cuda_wrappers=[ + ( + "residual_gate_add", + "sglang_residual_gate_add::" f"ResidualGateAddKernel<{args}>::run", + ), + ], + ) + + +def _fake_impl( + residual: torch.Tensor, update: torch.Tensor, gate: torch.Tensor +) -> torch.Tensor: + return torch.empty_like(residual) + + +@register_custom_op( + op_name="diffusion_residual_gate_add", + mutates_args=[], + fake_impl=_fake_impl, +) +def _residual_gate_add_custom_op( + residual: torch.Tensor, update: torch.Tensor, gate: torch.Tensor +) -> torch.Tensor: + out = torch.empty_like(residual) + module = _jit_residual_gate_add_module(residual.dtype) + module.residual_gate_add(out, residual, update, gate) + return out + + +def _is_row_broadcast_gate(residual: torch.Tensor, gate: torch.Tensor) -> bool: + if gate.dim() != residual.dim() or gate.shape[-1] != residual.shape[-1]: + return False + row_dim = gate.dim() - 2 + return gate.shape[row_dim] == 1 and all(size == 1 for size in gate.shape[:-1]) + + +def can_use_residual_gate_add_cuda( + residual: torch.Tensor, update: torch.Tensor, gate: torch.Tensor +) -> bool: + return ( + residual.dtype in _SUPPORTED_DTYPES + and residual.dtype == update.dtype + and residual.dtype == gate.dtype + and residual.is_cuda + and update.is_cuda + and gate.is_cuda + and residual.device == update.device == gate.device + and residual.dim() >= 2 + and update.shape == residual.shape + and (gate.shape == residual.shape or _is_row_broadcast_gate(residual, gate)) + and residual.is_contiguous() + and update.is_contiguous() + and gate.is_contiguous() + ) + + +def residual_gate_add_cuda( + residual: torch.Tensor, update: torch.Tensor, gate: torch.Tensor +) -> torch.Tensor: + if not can_use_residual_gate_add_cuda(residual, update, gate): + raise RuntimeError("unsupported input for residual_gate_add CUDA") + return _residual_gate_add_custom_op(residual, update, gate) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py b/python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py index 09f647da1..fa1fbf2ac 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/ltx_2.py @@ -10,6 +10,10 @@ import torch import torch.nn as nn import torch.nn.functional as F +from sglang.jit_kernel.diffusion.residual_gate_add import ( + can_use_residual_gate_add_cuda, + residual_gate_add_cuda, +) from sglang.multimodal_gen.configs.models.dits.ltx_2 import LTX2ArchConfig, LTX2Config from sglang.multimodal_gen.runtime.distributed import ( get_sp_parallel_rank, @@ -48,6 +52,29 @@ logger = init_logger(__name__) ADALN_NUM_BASE_PARAMS = 6 ADALN_NUM_CROSS_ATTN_PARAMS = 3 +_LTX2_RESIDUAL_GATE_CUDA_DISABLED = False + + +def _ltx2_residual_gate_add( + residual: torch.Tensor, + update: torch.Tensor, + gate: torch.Tensor, +) -> torch.Tensor: + global _LTX2_RESIDUAL_GATE_CUDA_DISABLED + + if not _LTX2_RESIDUAL_GATE_CUDA_DISABLED and can_use_residual_gate_add_cuda( + residual, update, gate + ): + try: + return residual_gate_add_cuda(residual, update, gate) + except Exception as exc: + if torch.compiler.is_compiling(): + raise + logger.warning_once(f"Disabling LTX2 residual-gate CUDA fast path: {exc}") + _LTX2_RESIDUAL_GATE_CUDA_DISABLED = True + + return residual + update * gate + _LTX2_FUSED_ADA_VALUES_RUNTIME_DISABLED = False @@ -1108,7 +1135,9 @@ class LTX2TransformerBlock(nn.Module): gather_context_kv_for_sp=audio_replicated_for_sp, context_replicated_prefix_len=video_memory_prefix_len, ) - hidden_states = hidden_states + attn_hidden_states * vgate_msa + hidden_states = _ltx2_residual_gate_add( + hidden_states, attn_hidden_states, vgate_msa + ) if audio_ada_values is None: ashift_msa, ascale_msa, agate_msa = self.get_ada_values( @@ -1128,7 +1157,9 @@ class LTX2TransformerBlock(nn.Module): all_perturbed=skip_audio_self_attn, skip_sequence_parallel_override=audio_replicated_for_sp, ) - audio_hidden_states = audio_hidden_states + attn_audio_hidden_states * agate_msa + audio_hidden_states = _ltx2_residual_gate_add( + audio_hidden_states, attn_audio_hidden_states, agate_msa + ) # 2. Prompt Cross-Attention if self.cross_attention_adaln: # LTX2.3 @@ -1156,7 +1187,9 @@ class LTX2TransformerBlock(nn.Module): context=mod_encoder_hidden_states, mask=encoder_attention_mask, ) - hidden_states = hidden_states + attn_hidden_states * vgate_q + hidden_states = _ltx2_residual_gate_add( + hidden_states, attn_hidden_states, vgate_q + ) if audio_ada_values is None: ashift_q, ascale_q, agate_q = self.get_ada_values( @@ -1182,8 +1215,8 @@ class LTX2TransformerBlock(nn.Module): context=mod_audio_encoder_hidden_states, mask=audio_encoder_attention_mask, ) - audio_hidden_states = ( - audio_hidden_states + attn_audio_hidden_states * agate_q + audio_hidden_states = _ltx2_residual_gate_add( + audio_hidden_states, attn_audio_hidden_states, agate_q ) else: norm_hidden_states = self.rms_norm(hidden_states, self.norm_eps) @@ -1284,7 +1317,9 @@ class LTX2TransformerBlock(nn.Module): a2v_attn_hidden_states = ( a2v_attn_hidden_states * a2v_cross_attn_perturbation_mask ) - hidden_states = hidden_states + a2v_gate * a2v_attn_hidden_states + hidden_states = _ltx2_residual_gate_add( + hidden_states, a2v_attn_hidden_states, a2v_gate + ) # V2A mod_norm_hidden_states = ( @@ -1308,8 +1343,8 @@ class LTX2TransformerBlock(nn.Module): v2a_attn_hidden_states = ( v2a_attn_hidden_states * v2a_cross_attn_perturbation_mask ) - audio_hidden_states = ( - audio_hidden_states + v2a_gate * v2a_attn_hidden_states + audio_hidden_states = _ltx2_residual_gate_add( + audio_hidden_states, v2a_attn_hidden_states, v2a_gate ) # 4. Feedforward if video_ada_values is None: @@ -1322,7 +1357,7 @@ class LTX2TransformerBlock(nn.Module): self.rms_norm(hidden_states, self.norm_eps) * (1 + vscale_mlp) + vshift_mlp ) ff_output = self.ff(norm_hidden_states) - hidden_states = hidden_states + ff_output * vgate_mlp + hidden_states = _ltx2_residual_gate_add(hidden_states, ff_output, vgate_mlp) if audio_ada_values is None: ashift_mlp, ascale_mlp, agate_mlp = self.get_ada_values( @@ -1335,7 +1370,9 @@ class LTX2TransformerBlock(nn.Module): + ashift_mlp ) audio_ff_output = self.audio_ff(norm_audio_hidden_states) - audio_hidden_states = audio_hidden_states + audio_ff_output * agate_mlp + audio_hidden_states = _ltx2_residual_gate_add( + audio_hidden_states, audio_ff_output, agate_mlp + ) return hidden_states, audio_hidden_states diff --git a/test/registered/jit/benchmark/diffusion/bench_residual_gate_add.py b/test/registered/jit/benchmark/diffusion/bench_residual_gate_add.py new file mode 100644 index 000000000..d8e3e0e45 --- /dev/null +++ b/test/registered/jit/benchmark/diffusion/bench_residual_gate_add.py @@ -0,0 +1,114 @@ +import random +import sys +from dataclasses import dataclass + +import torch + +from sglang.jit_kernel.diffusion.residual_gate_add import residual_gate_add_cuda +from sglang.jit_kernel.diffusion.triton.scale_shift import fuse_scale_shift_kernel +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.utils import is_in_ci + +register_cuda_ci(est_time=30, suite="base-b-kernel-benchmark-1-gpu-large") + + +@dataclass(frozen=True) +class Workload: + name: str + residual_shape: tuple[int, ...] + gate_shape: tuple[int, ...] + + +FULL_WORKLOADS = [ + Workload("ltx2_bcast_s32640_c4096", (1, 32640, 4096), (1, 1, 4096)), + Workload("ltx2_full_s8160_c4096", (1, 8160, 4096), (1, 8160, 4096)), + Workload("ideogram4_bcast_s4096_c4608", (1, 4096, 4608), (1, 1, 4608)), + Workload("flux2_bcast_s4608_c3072", (1, 4608, 3072), (1, 1, 3072)), + Workload("flux2_bcast_s4096_c3072", (1, 4096, 3072), (1, 1, 3072)), + Workload("flux2_bcast_s512_c3072", (1, 512, 3072), (1, 1, 3072)), + Workload("ltx2_full_s126_c2048", (1, 126, 2048), (1, 126, 2048)), +] +CI_WORKLOADS = [ + Workload("ltx2_bcast_s1024_c4096", (1, 1024, 4096), (1, 1, 4096)), + Workload("ltx2_full_s512_c4096", (1, 512, 4096), (1, 512, 4096)), +] + + +def cuda_event_us(fn, warmups: int, repeats: int, rounds: int) -> float: + for _ in range(warmups): + fn() + torch.cuda.synchronize() + + samples = [] + for _ in range(rounds): + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + for _ in range(repeats): + fn() + end.record() + end.synchronize() + samples.append(start.elapsed_time(end) * 1000.0 / repeats) + samples.sort() + return samples[len(samples) // 2] + + +def benchmark() -> None: + if not torch.cuda.is_available(): + print("CUDA required") + return + + torch.manual_seed(20260625) + random.seed(20260625) + torch.cuda.set_device(0) + + workloads = CI_WORKLOADS if is_in_ci() else FULL_WORKLOADS + warmups = 5 if is_in_ci() else 20 + repeats = 5 if is_in_ci() else 20 + rounds = 5 if is_in_ci() else 13 + + print("| workload | gate | torch us | triton us | cuda us | cuda/triton |") + print("|---|---|---:|---:|---:|---:|") + + for workload in workloads: + residual = torch.randn( + workload.residual_shape, device="cuda", dtype=torch.bfloat16 + ) + update = torch.randn_like(residual) + gate = torch.randn(workload.gate_shape, device="cuda", dtype=torch.bfloat16) + + ref = residual + update * gate + triton_out = fuse_scale_shift_kernel(update, gate, residual, scale_constant=0) + cuda_out = residual_gate_add_cuda(residual, update, gate) + torch.cuda.synchronize() + torch.testing.assert_close(triton_out, ref, atol=5e-2, rtol=5e-2) + torch.testing.assert_close(cuda_out, ref, atol=5e-2, rtol=5e-2) + + fns = { + "torch": lambda: residual + update * gate, + "triton": lambda: fuse_scale_shift_kernel( + update, gate, residual, scale_constant=0 + ), + "cuda": lambda: residual_gate_add_cuda(residual, update, gate), + } + order = ["torch", "triton", "cuda"] + random.shuffle(order) + times = { + name: cuda_event_us(fns[name], warmups, repeats, rounds) for name in order + } + + gate_kind = ( + "bcast" if workload.gate_shape != workload.residual_shape else "full" + ) + print( + f"| {workload.name} | {gate_kind} | {times['torch']:.2f} | " + f"{times['triton']:.2f} | {times['cuda']:.2f} | " + f"{times['triton'] / times['cuda']:.3f}x |" + ) + + torch.cuda.empty_cache() + + +if __name__ == "__main__": + benchmark() + sys.exit(0) diff --git a/test/registered/jit/diffusion/test_residual_gate_add.py b/test/registered/jit/diffusion/test_residual_gate_add.py new file mode 100644 index 000000000..973fe8618 --- /dev/null +++ b/test/registered/jit/diffusion/test_residual_gate_add.py @@ -0,0 +1,101 @@ +import sys + +import pytest +import torch + +from sglang.jit_kernel.diffusion.residual_gate_add import ( + can_use_residual_gate_add_cuda, + residual_gate_add_cuda, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-large") +register_cuda_ci(est_time=30, suite="base-b-kernel-unit-1-gpu-b200") + + +CASES = [ + ((1, 1024, 4096), (1, 1, 4096)), + ((1, 512, 4096), (1, 512, 4096)), + ((1, 17, 65), (1, 1, 65)), + ((1, 17, 65), (1, 17, 65)), +] + + +def _tol(dtype: torch.dtype) -> float: + return 1e-5 if dtype == torch.float32 else 5e-2 + + +def _assert_matches_torch(out: torch.Tensor, ref: torch.Tensor) -> None: + if ref.dtype == torch.float32: + torch.testing.assert_close(out, ref, atol=_tol(ref.dtype), rtol=_tol(ref.dtype)) + else: + torch.testing.assert_close(out, ref, atol=0, rtol=0) + + +@pytest.fixture(autouse=True) +def cuda_setup(): + if not torch.cuda.is_available(): + pytest.skip("CUDA required") + torch.cuda.manual_seed(0) + + +@pytest.mark.parametrize("residual_shape,gate_shape", CASES) +def test_residual_gate_add_matches_torch(residual_shape, gate_shape): + residual = torch.randn(residual_shape, device="cuda", dtype=torch.bfloat16) + update = torch.randn_like(residual) + gate = torch.randn(gate_shape, device="cuda", dtype=torch.bfloat16) + + out = residual_gate_add_cuda(residual, update, gate) + ref = residual + update * gate + _assert_matches_torch(out, ref) + + +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16, torch.float32]) +@pytest.mark.parametrize("gate_shape", [(1, 1, 64), (1, 9, 64)]) +def test_residual_gate_add_dtypes(dtype, gate_shape): + residual = torch.randn((1, 9, 64), device="cuda", dtype=dtype) + update = torch.randn_like(residual) + gate = torch.randn(gate_shape, device="cuda", dtype=dtype) + + out = residual_gate_add_cuda(residual, update, gate) + ref = residual + update * gate + _assert_matches_torch(out, ref) + + +def test_can_use_residual_gate_add_cuda_rejects_unsupported_inputs(): + residual = torch.randn((1, 8, 64), device="cuda", dtype=torch.bfloat16) + update = torch.randn_like(residual) + gate = torch.randn((1, 1, 64), device="cuda", dtype=torch.bfloat16) + + assert can_use_residual_gate_add_cuda(residual, update, gate) + assert not can_use_residual_gate_add_cuda(residual.cpu(), update, gate) + assert not can_use_residual_gate_add_cuda(residual, update.float(), gate) + assert not can_use_residual_gate_add_cuda(residual, update[:, ::2], gate) + assert not can_use_residual_gate_add_cuda(residual, update, gate[:, :, ::2]) + + # Only [1, ..., 1, D] row-broadcast gates are supported; a batched + # [B>1, 1, D] gate is not row-broadcast here and must fall back. + batched_residual = torch.randn((2, 8, 64), device="cuda", dtype=torch.bfloat16) + batched_update = torch.randn_like(batched_residual) + batched_gate = torch.randn((2, 1, 64), device="cuda", dtype=torch.bfloat16) + assert not can_use_residual_gate_add_cuda( + batched_residual, batched_update, batched_gate + ) + + +def test_residual_gate_add_custom_op_torch_compile_fullgraph(): + residual = torch.randn((1, 32, 128), device="cuda", dtype=torch.bfloat16) + update = torch.randn_like(residual) + gate = torch.randn((1, 1, 128), device="cuda", dtype=torch.bfloat16) + + def fn(residual, update, gate): + return residual_gate_add_cuda(residual, update, gate) + + compiled = torch.compile(fn, fullgraph=True) + out = compiled(residual, update, gate) + ref = residual + update * gate + _assert_matches_torch(out, ref) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v", "-s"]))