diff --git a/python/sglang/jit_kernel/csrc/diffusion/norm_scale_shift.cuh b/python/sglang/jit_kernel/csrc/diffusion/norm_scale_shift.cuh new file mode 100644 index 000000000..dbc9af09b --- /dev/null +++ b/python/sglang/jit_kernel/csrc/diffusion/norm_scale_shift.cuh @@ -0,0 +1,213 @@ +// Minimal native-CUDA fast path for Qwen-Image diffusion norm-scale-shift. +// +// Supported shape family: +// - bf16 activations, B == 1, hidden dim == 3072 +// - layer norm only, no affine weight/bias +// - scale/shift are bf16 row-broadcast tensors ([D], [1,D], or [1,1,D]) +// - optional residual path uses a bf16 row-broadcast gate +// +// All other public-op inputs fall back to the existing CuTe-DSL implementation +// from the Python dispatcher. + +#pragma once + +#include // For TensorMatcher, SymbolicSize, SymbolicDevice + +#include // For device::math::rsqrt +#include // For SGL_DEVICE, bf16_t, LaunchKernel +#include // For AlignedVector +#include // For warp::reduce_sum + +#include + +namespace sglang_norm_scale_shift { + +namespace { + +constexpr int kHidden = 3072; +constexpr int kVecElems = 16; // 32B/thread for bf16 on Blackwell. +constexpr int kThreads = kHidden / kVecElems; +constexpr int kWarps = kThreads / device::kWarpThreads; +constexpr float kInvHidden = 1.0f / float(kHidden); + +static_assert(kThreads == 192); +static_assert(kWarps == 6); + +struct QwenImageNormParams { + void* y; + void* res_out; + const void* x; + const void* residual; + const void* gate; + const void* scale; + const void* shift; + float eps; +}; + +SGL_DEVICE float cta_reduce_sum(float v, int warp, int lane, float* scratch) { + v = device::warp::reduce_sum(v); + if (lane == 0) { + scratch[warp] = v; + } + __syncthreads(); + + if (warp == 0) { + float a = lane < kWarps ? scratch[lane] : 0.0f; + a = device::warp::reduce_sum(a); + if (lane == 0) { + scratch[kWarps] = a; + } + } + __syncthreads(); + return scratch[kWarps]; +} + +template +__global__ void qwen_image_norm_scale_shift_kernel(const QwenImageNormParams __grid_constant__ params) { + using namespace device; + using Vec = AlignedVector; + + const int row = blockIdx.x; + const int tid = threadIdx.x; + const int lane = tid & int(kWarpThreads - 1); + const int warp = tid >> 5; + const int row_offset = row * kHidden; + const int elem_offset = tid * kVecElems; + + __shared__ float scratch_a[kWarps + 1]; + __shared__ float scratch_b[kWarps + 1]; + + Vec xv; + xv.load(static_cast(params.x) + row_offset + elem_offset); + + float v[kVecElems]; +#pragma unroll + for (int i = 0; i < kVecElems; ++i) { + v[i] = static_cast(xv[i]); + } + + if constexpr (kHasResidual) { + Vec gv; + Vec rv; + Vec ro; + gv.load(static_cast(params.gate) + elem_offset); + rv.load(static_cast(params.residual) + row_offset + elem_offset); + +#pragma unroll + for (int i = 0; i < kVecElems; ++i) { + const bf16_t rounded = static_cast(v[i] * static_cast(gv[i]) + static_cast(rv[i])); + ro[i] = rounded; + v[i] = static_cast(rounded); + } + ro.store(static_cast(params.res_out) + row_offset + elem_offset); + } + + float sum = 0.0f; +#pragma unroll + for (int i = 0; i < kVecElems; ++i) { + sum += v[i]; + } + const float mean = cta_reduce_sum(sum, warp, lane, scratch_a) * kInvHidden; + + float var_sum = 0.0f; +#pragma unroll + for (int i = 0; i < kVecElems; ++i) { + const float d = v[i] - mean; + var_sum += d * d; + } + const float var = cta_reduce_sum(var_sum, warp, lane, scratch_b) * kInvHidden; + const float factor = math::rsqrt(var + params.eps); + + Vec scv; + Vec shv; + Vec yv; + scv.load(static_cast(params.scale) + elem_offset); + shv.load(static_cast(params.shift) + elem_offset); + +#pragma unroll + for (int i = 0; i < kVecElems; ++i) { + const float norm = static_cast(static_cast((v[i] - mean) * factor)); + yv[i] = static_cast(norm * (1.0f + static_cast(scv[i])) + static_cast(shv[i])); + } + yv.store(static_cast(params.y) + row_offset + elem_offset); +} + +inline uint32_t verify_qwen_geometry(host::SymbolicSize& num_rows) { + using namespace host; + RuntimeCheck(num_rows.unwrap() > 0, "num_rows must be positive"); + RuntimeCheck(num_rows.unwrap() <= int64_t(UINT32_MAX), "num_rows out of range"); + return static_cast(num_rows.unwrap()); +} + +} // namespace + +struct QwenImageNormScaleShiftKernel { + static void + run(tvm::ffi::TensorView y, + tvm::ffi::TensorView x, + tvm::ffi::TensorView scale, + tvm::ffi::TensorView shift, + double eps) { + using namespace host; + auto N = SymbolicSize{"num_rows"}; + auto device = SymbolicDevice{}; + device.set_options(); + + TensorMatcher({N, kHidden}).with_dtype().with_device(device).verify(x).verify(y); + TensorMatcher({kHidden}).with_dtype().with_device(device).verify(scale).verify(shift); + + const uint32_t grid = verify_qwen_geometry(N); + const auto params = QwenImageNormParams{ + .y = y.data_ptr(), + .res_out = nullptr, + .x = x.data_ptr(), + .residual = nullptr, + .gate = nullptr, + .scale = scale.data_ptr(), + .shift = shift.data_ptr(), + .eps = static_cast(eps), + }; + LaunchKernel(grid, kThreads, device.unwrap())(qwen_image_norm_scale_shift_kernel, params); + } +}; + +struct QwenImageScaleResidualNormScaleShiftKernel { + static void + run(tvm::ffi::TensorView y, + tvm::ffi::TensorView res_out, + tvm::ffi::TensorView residual, + tvm::ffi::TensorView x, + tvm::ffi::TensorView gate, + tvm::ffi::TensorView scale, + tvm::ffi::TensorView shift, + double eps) { + using namespace host; + auto N = SymbolicSize{"num_rows"}; + auto device = SymbolicDevice{}; + device.set_options(); + + TensorMatcher({N, kHidden}) + .with_dtype() + .with_device(device) + .verify(x) + .verify(residual) + .verify(y) + .verify(res_out); + TensorMatcher({kHidden}).with_dtype().with_device(device).verify(gate).verify(scale).verify(shift); + + const uint32_t grid = verify_qwen_geometry(N); + const auto params = QwenImageNormParams{ + .y = y.data_ptr(), + .res_out = res_out.data_ptr(), + .x = x.data_ptr(), + .residual = residual.data_ptr(), + .gate = gate.data_ptr(), + .scale = scale.data_ptr(), + .shift = shift.data_ptr(), + .eps = static_cast(eps), + }; + LaunchKernel(grid, kThreads, device.unwrap())(qwen_image_norm_scale_shift_kernel, params); + } +}; + +} // namespace sglang_norm_scale_shift diff --git a/python/sglang/jit_kernel/diffusion/cutedsl/scale_residual_norm_scale_shift.py b/python/sglang/jit_kernel/diffusion/cutedsl/scale_residual_norm_scale_shift.py index 8f102fd73..c835fea63 100644 --- a/python/sglang/jit_kernel/diffusion/cutedsl/scale_residual_norm_scale_shift.py +++ b/python/sglang/jit_kernel/diffusion/cutedsl/scale_residual_norm_scale_shift.py @@ -299,6 +299,15 @@ 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.jit_kernel.diffusion.norm_scale_shift_native import ( + try_fused_norm_scale_shift as _try_qwen_native_norm_scale_shift, + ) + + native_y = _try_qwen_native_norm_scale_shift( + x, weight, bias, scale, shift, norm_type, eps + ) + if native_y is not None: + return native_y stream = cuda.CUstream(torch.cuda.current_stream().cuda_stream) # Tensor Validation BSD = x.shape @@ -376,6 +385,15 @@ 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.jit_kernel.diffusion.norm_scale_shift_native import ( + try_fused_scale_residual_norm_scale_shift as _try_qwen_native_residual_path, + ) + + native_out = _try_qwen_native_residual_path( + residual, x, gate, weight, bias, scale, shift, norm_type, eps + ) + if native_out is not None: + return native_out # Tensor Validation BSD = x.shape validate_x(x, *BSD) diff --git a/python/sglang/jit_kernel/diffusion/norm_scale_shift_native.py b/python/sglang/jit_kernel/diffusion/norm_scale_shift_native.py new file mode 100644 index 000000000..cbc175918 --- /dev/null +++ b/python/sglang/jit_kernel/diffusion/norm_scale_shift_native.py @@ -0,0 +1,131 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +import torch + +from sglang.jit_kernel.utils import cache_once, load_jit + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +_HIDDEN = 3072 +_ALIGN = 32 + + +def _aligned(t: torch.Tensor) -> bool: + return t.data_ptr() % _ALIGN == 0 + + +def _blackwell_or_newer(device: torch.device) -> bool: + return ( + torch.cuda.is_available() and torch.cuda.get_device_capability(device)[0] >= 10 + ) + + +def _qwen_activation(t, like=None) -> bool: + return ( + isinstance(t, torch.Tensor) + and t.is_cuda + and t.dtype == torch.bfloat16 + and t.ndim == 3 + and t.shape[0] == 1 + and t.shape[-1] == _HIDDEN + and t.numel() > 0 + and t.is_contiguous() + and _aligned(t) + and (like is None or (t.device == like.device and t.shape == like.shape)) + ) + + +def _row_bf16(t, device: torch.device): + if ( + not isinstance(t, torch.Tensor) + or t.dtype != torch.bfloat16 + or not t.is_cuda + or t.device != device + or t.ndim < 1 + or t.stride(-1) != 1 + ): + return None + if t.shape == (_HIDDEN,): + row = t + elif t.shape in ((1, _HIDDEN), (1, 1, _HIDDEN)): + row = t.reshape(_HIDDEN) + else: + return None + return row if _aligned(row) else None + + +@cache_once +def _jit_norm_scale_shift_module() -> Module: + return load_jit( + "qwen_image_norm_scale_shift_native", + cuda_files=["diffusion/norm_scale_shift.cuh"], + cuda_wrappers=[ + ( + "qwen_image_nss_bf16_row", + "sglang_norm_scale_shift::QwenImageNormScaleShiftKernel::run", + ), + ( + "qwen_image_srnss_bf16_row", + "sglang_norm_scale_shift::" + "QwenImageScaleResidualNormScaleShiftKernel::run", + ), + ], + ) + + +_module = _jit_norm_scale_shift_module + + +def try_fused_norm_scale_shift(x, weight, bias, scale, shift, norm_type, eps): + if norm_type != "layer" or weight is not None or bias is not None: + return None + if not _qwen_activation(x) or not _blackwell_or_newer(x.device): + return None + + scale = _row_bf16(scale, x.device) + shift = _row_bf16(shift, x.device) + if scale is None or shift is None: + return None + + y = torch.empty_like(x) + _module().qwen_image_nss_bf16_row( + y.view(-1, _HIDDEN), x.view(-1, _HIDDEN), scale, shift, float(eps) + ) + return y + + +def try_fused_scale_residual_norm_scale_shift( + residual, x, gate, weight, bias, scale, shift, norm_type, eps +): + if norm_type != "layer" or weight is not None or bias is not None: + return None + if not ( + _qwen_activation(x) + and _qwen_activation(residual, x) + and _blackwell_or_newer(x.device) + ): + return None + + gate = _row_bf16(gate, x.device) + scale = _row_bf16(scale, x.device) + shift = _row_bf16(shift, x.device) + if gate is None or scale is None or shift is None: + return None + + y = torch.empty_like(x) + residual_out = torch.empty_like(x) + _module().qwen_image_srnss_bf16_row( + y.view(-1, _HIDDEN), + residual_out.view(-1, _HIDDEN), + residual.view(-1, _HIDDEN), + x.view(-1, _HIDDEN), + gate, + scale, + shift, + float(eps), + ) + return y, residual_out