From d633ab7349a8cdfe939d25745c3f647c508b8be5 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <35585791+BBuf@users.noreply.github.com> Date: Fri, 27 Mar 2026 14:27:08 +0800 Subject: [PATCH] [Diffusion] Add qknorm rope fuse kernel (#21440) --- .../benchmark/diffusion/bench_qknorm_rope.py | 190 ++++++++++++++ .../jit_kernel/csrc/diffusion/qknorm_rope.cuh | 246 ++++++++++++++++++ .../jit_kernel/diffusion/qknorm_rope.py | 97 +++++++ .../tests/diffusion/test_qknorm_rope.py | 153 +++++++++++ .../runtime/layers/layernorm.py | 141 ++++++++++ .../runtime/models/dits/flux.py | 85 +++--- .../runtime/models/dits/flux_2.py | 81 +++--- .../runtime/models/dits/qwen_image.py | 45 ++-- .../runtime/models/dits/zimage.py | 51 +++- 9 files changed, 986 insertions(+), 103 deletions(-) create mode 100644 python/sglang/jit_kernel/benchmark/diffusion/bench_qknorm_rope.py create mode 100644 python/sglang/jit_kernel/csrc/diffusion/qknorm_rope.cuh create mode 100644 python/sglang/jit_kernel/diffusion/qknorm_rope.py create mode 100644 python/sglang/jit_kernel/tests/diffusion/test_qknorm_rope.py diff --git a/python/sglang/jit_kernel/benchmark/diffusion/bench_qknorm_rope.py b/python/sglang/jit_kernel/benchmark/diffusion/bench_qknorm_rope.py new file mode 100644 index 000000000..8d60097bb --- /dev/null +++ b/python/sglang/jit_kernel/benchmark/diffusion/bench_qknorm_rope.py @@ -0,0 +1,190 @@ +from dataclasses import dataclass +from typing import Tuple + +import torch +import triton +import triton.testing + +from sglang.jit_kernel.benchmark.utils import ( + DEFAULT_DEVICE, + DEFAULT_DTYPE, + get_benchmark_range, + run_benchmark_no_cudagraph, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=13, suite="stage-b-kernel-benchmark-1-gpu-large") + +MAX_SEQ_LEN = 131072 +ROPE_BASE = 10000.0 + + +@dataclass(frozen=True) +class CaseSpec: + name: str + batch_size: int + num_tokens: int + num_heads: int + head_dim: int + rope_dim: int + is_neox: bool + + +BENCH_CASES = ( + CaseSpec("flux_1024", 1, 4096, 24, 128, 128, False), + CaseSpec("qwen_image_1024", 1, 4096, 32, 128, 128, False), + CaseSpec("qwen_image_partial", 1, 4096, 32, 128, 64, False), + # Z-Image-Turbo default 1024x1024 config: dim=3840, num_heads=30 -> head_dim=128. + CaseSpec("zimage_1024", 1, 4096, 30, 128, 128, False), + CaseSpec("batch2_medium", 2, 2048, 24, 128, 128, False), +) +CASE_BY_NAME = {case.name: case for case in BENCH_CASES} +CASE_NAMES = get_benchmark_range( + full_range=[case.name for case in BENCH_CASES], + ci_range=[case.name for case in BENCH_CASES], +) +LINE_VALS = ["split", "fused"] +LINE_NAMES = ["JIT QKNorm + FlashInfer RoPE", "SGL JIT Fused QKNorm+RoPE"] +STYLES = [("red", "-"), ("blue", "--")] + + +def create_cos_sin_cache( + rotary_dim: int, + max_position: int = MAX_SEQ_LEN, + base: float = ROPE_BASE, +) -> torch.Tensor: + inv_freq = 1.0 / ( + base + ** ( + torch.arange(0, rotary_dim, 2, dtype=torch.float32, device=DEFAULT_DEVICE) + / rotary_dim + ) + ) + t = torch.arange(max_position, dtype=torch.float32, device=DEFAULT_DEVICE) + freqs = torch.einsum("i,j->ij", t, inv_freq) + return torch.cat((freqs.cos(), freqs.sin()), dim=-1) + + +def make_inputs(case: CaseSpec) -> dict[str, torch.Tensor | bool]: + seed = ( + case.batch_size * 1_000_003 + + case.num_tokens * 8191 + + case.num_heads * 127 + + case.head_dim * 17 + + case.rope_dim + ) + generator = torch.Generator(device=DEFAULT_DEVICE) + generator.manual_seed(seed) + return { + "q": torch.randn( + case.batch_size * case.num_tokens, + case.num_heads, + case.head_dim, + device=DEFAULT_DEVICE, + dtype=DEFAULT_DTYPE, + generator=generator, + ), + "k": torch.randn( + case.batch_size * case.num_tokens, + case.num_heads, + case.head_dim, + device=DEFAULT_DEVICE, + dtype=DEFAULT_DTYPE, + generator=generator, + ), + "q_weight": torch.randn( + case.head_dim, + device=DEFAULT_DEVICE, + dtype=DEFAULT_DTYPE, + generator=generator, + ), + "k_weight": torch.randn( + case.head_dim, + device=DEFAULT_DEVICE, + dtype=DEFAULT_DTYPE, + generator=generator, + ), + "positions": torch.randint( + 0, + MAX_SEQ_LEN, + (case.batch_size * case.num_tokens,), + device=DEFAULT_DEVICE, + dtype=torch.int64, + generator=generator, + ), + "cos_sin_cache": create_cos_sin_cache(case.rope_dim), + "is_neox": case.is_neox, + } + + +def clone_inputs( + inputs: dict[str, torch.Tensor | bool], +) -> dict[str, torch.Tensor | bool]: + out: dict[str, torch.Tensor | bool] = {} + for key, value in inputs.items(): + out[key] = value.clone() if isinstance(value, torch.Tensor) else value + return out + + +def split_qknorm_rope(inputs: dict[str, torch.Tensor | bool]) -> None: + from flashinfer.rope import apply_rope_with_cos_sin_cache_inplace + + from sglang.jit_kernel.norm import fused_inplace_qknorm + + q = inputs["q"] + k = inputs["k"] + q_weight = inputs["q_weight"] + k_weight = inputs["k_weight"] + positions = inputs["positions"] + cos_sin_cache = inputs["cos_sin_cache"] + is_neox = bool(inputs["is_neox"]) + + fused_inplace_qknorm(q, k, q_weight, k_weight) + apply_rope_with_cos_sin_cache_inplace( + positions=positions, + query=q.view(q.shape[0], -1), + key=k.view(k.shape[0], -1), + head_size=q.shape[-1], + cos_sin_cache=cos_sin_cache, + is_neox=is_neox, + ) + + +def fused_qknorm_rope(inputs: dict[str, torch.Tensor | bool]) -> None: + from sglang.jit_kernel.diffusion.qknorm_rope import fused_inplace_qknorm_rope + + fused_inplace_qknorm_rope( + inputs["q"], + inputs["k"], + inputs["q_weight"], + inputs["k_weight"], + inputs["cos_sin_cache"], + inputs["positions"], + is_neox=bool(inputs["is_neox"]), + rope_dim=inputs["cos_sin_cache"].shape[-1], + ) + + +@triton.testing.perf_report( + triton.testing.Benchmark( + x_names=["case_name"], + x_vals=CASE_NAMES, + line_arg="provider", + line_vals=LINE_VALS, + line_names=LINE_NAMES, + styles=STYLES, + ylabel="us", + plot_name="diffusion-qknorm-rope-performance", + args={}, + ) +) +def benchmark(case_name: str, provider: str) -> Tuple[float, float, float]: + case = CASE_BY_NAME[case_name] + inputs = make_inputs(case) + fn = split_qknorm_rope if provider == "split" else fused_qknorm_rope + return run_benchmark_no_cudagraph(lambda: fn(inputs)) + + +if __name__ == "__main__": + print("Running diffusion qknorm + rope performance benchmark...") + benchmark.run(print_data=True) diff --git a/python/sglang/jit_kernel/csrc/diffusion/qknorm_rope.cuh b/python/sglang/jit_kernel/csrc/diffusion/qknorm_rope.cuh new file mode 100644 index 000000000..ab4452945 --- /dev/null +++ b/python/sglang/jit_kernel/csrc/diffusion/qknorm_rope.cuh @@ -0,0 +1,246 @@ +#include + +#include +#include +#include +#include +#include + +#include + +#include +#include + +namespace { + +struct QKNormRopeParams { + void* __restrict__ q_ptr; + void* __restrict__ k_ptr; // pre-offset by -num_qo_heads * head_stride_bytes + const void* __restrict__ q_weight_ptr; + const void* __restrict__ k_weight_ptr; + const void* __restrict__ cos_sin_cache_ptr; + const void* __restrict__ positions; + int64_t q_stride_bytes; + int64_t k_stride_bytes; + int64_t head_stride_bytes; + uint32_t num_qo_heads; + uint32_t num_kv_heads; + uint32_t num_tokens; + float eps; +}; + +constexpr uint32_t kThreadsPerBlock = 256; +constexpr uint32_t kWarpsPerBlock = kThreadsPerBlock / device::kWarpThreads; + +template +constexpr uint32_t active_mask() { + static_assert(kLaneCount <= device::kWarpThreads, "active_mask lane count must not exceed warp size"); + if constexpr (kLaneCount == device::kWarpThreads) { + return 0xffffffffu; + } else { + return (1u << kLaneCount) - 1u; + } +} + +SGL_DEVICE float load_cache_value(const float* ptr, int64_t idx) { +#ifdef USE_ROCM + return ptr[idx]; +#else + return __ldg(ptr + idx); +#endif +} + +template +__global__ void fused_qknorm_rope_warp(const QKNormRopeParams __grid_constant__ params) { + using namespace device; + + static_assert(std::is_same_v || std::is_same_v); + static_assert(kHeadDim <= 256, "Only warp-level fused qknorm+rope is supported"); + static_assert(kHeadDim % kWarpThreads == 0, "head_dim must be divisible by warp size"); + + constexpr uint32_t kElemsPerThread = kHeadDim / kWarpThreads; + constexpr uint32_t kVecSize = kElemsPerThread / 2; + constexpr uint32_t kRotaryLanes = kRopeDim / kElemsPerThread; + constexpr uint32_t kHalfRotaryLanes = kRotaryLanes / 2; + constexpr uint32_t kActiveMask = active_mask(); + constexpr int64_t kCosSinStrideBytes = kRopeDim * sizeof(float); + + static_assert(kElemsPerThread % 2 == 0, "Each lane must own an even number of elements"); + static_assert(kRopeDim > 0 && kRopeDim <= kHeadDim, "Invalid rope dimension"); + static_assert(kRopeDim % kElemsPerThread == 0, "rope_dim must align with per-lane vector width"); + static_assert( + !kIsNeox || (kRotaryLanes >= 2 && ((kRotaryLanes & (kRotaryLanes - 1)) == 0)), + "NeoX fused qknorm+rope requires rotary lane count to be a power of 2"); + + using Packed = packed_t; + using Storage = AlignedVector; + + const auto& [q_ptr, k_ptr, q_weight_ptr, k_weight_ptr, cos_sin_cache_ptr, positions, q_stride_bytes, k_stride_bytes, head_stride_bytes, num_qo_heads, num_kv_heads, num_tokens, eps] = + params; + + const uint32_t lane_id = threadIdx.x % kWarpThreads; + const uint32_t warp_id = threadIdx.x / kWarpThreads; + const uint32_t start_worker_id = blockIdx.x * kWarpsPerBlock + warp_id; + const uint32_t num_workers = gridDim.x * kWarpsPerBlock; + const uint32_t num_qk_heads = num_qo_heads + num_kv_heads; + const uint32_t num_works = num_qk_heads * num_tokens; + + PDLWaitPrimary(); + + for (uint32_t idx = start_worker_id; idx < num_works; idx += num_workers) { + const uint32_t token_id = idx / num_qk_heads; + const uint32_t head_id = idx % num_qk_heads; + const bool load_q = head_id < num_qo_heads; + const void* input = load_q ? pointer::offset(q_ptr, token_id * q_stride_bytes, head_id * head_stride_bytes) + : pointer::offset(k_ptr, token_id * k_stride_bytes, head_id * head_stride_bytes); + const void* weight_ptr = load_q ? q_weight_ptr : k_weight_ptr; + + auto input_vec = load_as(input, lane_id); + const auto weight_vec = load_as(weight_ptr, lane_id); + + float elems[kElemsPerThread]; + float sum_of_squares = 0.0f; + +#pragma unroll + for (uint32_t j = 0; j < kVecSize; ++j) { + const auto [x0, x1] = cast(input_vec[j]); + elems[2 * j] = x0; + elems[2 * j + 1] = x1; + sum_of_squares += x0 * x0 + x1 * x1; + } + + sum_of_squares = warp::reduce_sum(sum_of_squares); + const float norm_factor = math::rsqrt(sum_of_squares / static_cast(kHeadDim) + eps); + +#pragma unroll + for (uint32_t j = 0; j < kVecSize; ++j) { + const auto [w0, w1] = cast(weight_vec[j]); + elems[2 * j] *= norm_factor * w0; + elems[2 * j + 1] *= norm_factor * w1; + } + + if constexpr (kIsNeox) { + if (lane_id < kRotaryLanes) { + const auto pos = static_cast(static_cast(positions)[token_id]); + const auto cos_ptr = static_cast(pointer::offset(cos_sin_cache_ptr, pos * kCosSinStrideBytes)); + const auto sin_ptr = cos_ptr + kRopeDim / 2; + +#pragma unroll + for (uint32_t i = 0; i < kElemsPerThread; ++i) { + float swapped = __shfl_xor_sync(kActiveMask, elems[i], kHalfRotaryLanes); + if (lane_id < kHalfRotaryLanes) { + swapped = -swapped; + } + int dim_idx = static_cast(lane_id * kElemsPerThread + i); + dim_idx = (dim_idx * 2) % kRopeDim; + const int half_idx = dim_idx / 2; + const float cos = load_cache_value(cos_ptr, half_idx); + const float sin = load_cache_value(sin_ptr, half_idx); + elems[i] = elems[i] * cos + swapped * sin; + } + } + } else { + if (lane_id < kRotaryLanes) { + const auto pos = static_cast(static_cast(positions)[token_id]); + const auto cos_ptr = static_cast(pointer::offset(cos_sin_cache_ptr, pos * kCosSinStrideBytes)); + const auto sin_ptr = cos_ptr + kRopeDim / 2; + +#pragma unroll + for (uint32_t i = 0; i < kElemsPerThread; i += 2) { + const float x = elems[i]; + const float y = elems[i + 1]; + const int half_idx = static_cast(lane_id * kElemsPerThread + i) / 2; + const float cos = load_cache_value(cos_ptr, half_idx); + const float sin = load_cache_value(sin_ptr, half_idx); + elems[i] = x * cos - y * sin; + elems[i + 1] = y * cos + x * sin; + } + } + } + +#pragma unroll + for (uint32_t j = 0; j < kVecSize; ++j) { + input_vec[j] = cast({elems[2 * j], elems[2 * j + 1]}); + } + store_as(const_cast(input), input_vec, lane_id); + } + + PDLTriggerSecondary(); +} + +template +struct QKNormRopeKernel { + static_assert(kHeadDim <= 256, "Only head_dim <= 256 is supported"); + template + static constexpr auto kernel = fused_qknorm_rope_warp; + + static void + run(const tvm::ffi::TensorView q, + const tvm::ffi::TensorView k, + const tvm::ffi::TensorView q_weight, + const tvm::ffi::TensorView k_weight, + const tvm::ffi::TensorView cos_sin_cache, + const tvm::ffi::TensorView positions, + float eps) { + using namespace host; + + auto N = SymbolicSize{"num_tokens"}; + auto Q = SymbolicSize{"num_qo_heads"}; + auto K = SymbolicSize{"num_kv_heads"}; + auto D = SymbolicSize{"head_dim"}; + auto R = SymbolicSize{"rope_dim"}; + auto Dq = SymbolicSize{"q_stride"}; + auto Dk = SymbolicSize{"k_stride"}; + auto Dd = SymbolicSize{"head_stride"}; + auto device = SymbolicDevice{}; + auto id_type = SymbolicDType{}; + D.set_value(kHeadDim); + R.set_value(kRopeDim); + device.set_options(); + + TensorMatcher({N, Q, D}).with_strides({Dq, Dd, 1}).with_dtype().with_device(device).verify(q); + TensorMatcher({N, K, D}).with_strides({Dk, Dd, 1}).with_dtype().with_device(device).verify(k); + TensorMatcher({D}).with_dtype().with_device(device).verify(q_weight).verify(k_weight); + TensorMatcher({-1, R}).with_dtype().with_device(device).verify(cos_sin_cache); + TensorMatcher({N}).with_dtype(id_type).with_device(device).verify(positions); + + const auto num_tokens = static_cast(N.unwrap()); + const auto num_qo_heads = static_cast(Q.unwrap()); + const auto num_kv_heads = static_cast(K.unwrap()); + const auto q_stride_bytes = static_cast(Dq.unwrap() * sizeof(DType)); + const auto k_stride_bytes = static_cast(Dk.unwrap() * sizeof(DType)); + const auto head_stride_bytes = static_cast(Dd.unwrap() * sizeof(DType)); + + const int64_t k_offset = static_cast(num_qo_heads) * head_stride_bytes; + const auto params = QKNormRopeParams{ + .q_ptr = q.data_ptr(), + .k_ptr = pointer::offset(k.data_ptr(), -k_offset), + .q_weight_ptr = q_weight.data_ptr(), + .k_weight_ptr = k_weight.data_ptr(), + .cos_sin_cache_ptr = cos_sin_cache.data_ptr(), + .positions = positions.data_ptr(), + .q_stride_bytes = q_stride_bytes, + .k_stride_bytes = k_stride_bytes, + .head_stride_bytes = head_stride_bytes, + .num_qo_heads = num_qo_heads, + .num_kv_heads = num_kv_heads, + .num_tokens = num_tokens, + .eps = eps, + }; + + const auto is_int32 = id_type.is_type(); + const auto selected_kernel = is_int32 ? kernel : kernel; + const uint32_t kNumSM = runtime::get_sm_count(device.unwrap().device_id); + static const uint32_t kOccupancyTable[2] = { + runtime::get_blocks_per_sm(kernel, kThreadsPerBlock), + runtime::get_blocks_per_sm(kernel, kThreadsPerBlock), + }; + const auto max_blocks = kOccupancyTable[is_int32 ? 0 : 1] * kNumSM; + const auto num_works = (num_qo_heads + num_kv_heads) * num_tokens; + const auto needed_blocks = div_ceil(num_works, kWarpsPerBlock); + const auto num_blocks = std::min(max_blocks, needed_blocks); + LaunchKernel(num_blocks, kThreadsPerBlock, device.unwrap()).enable_pdl(kUsePDL)(selected_kernel, params); + } +}; + +} // namespace diff --git a/python/sglang/jit_kernel/diffusion/qknorm_rope.py b/python/sglang/jit_kernel/diffusion/qknorm_rope.py new file mode 100644 index 000000000..8dfdf8d8d --- /dev/null +++ b/python/sglang/jit_kernel/diffusion/qknorm_rope.py @@ -0,0 +1,97 @@ +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING + +import torch + +from sglang.jit_kernel.utils import ( + cache_once, + is_arch_support_pdl, + load_jit, + make_cpp_args, +) +from sglang.srt.utils.custom_op import register_custom_op + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +logger = logging.getLogger(__name__) + + +@cache_once +def _jit_qknorm_rope_module( + head_dim: int, + rope_dim: int, + is_neox: bool, + dtype: torch.dtype, +) -> Module: + args = make_cpp_args(head_dim, rope_dim, is_neox, is_arch_support_pdl(), dtype) + return load_jit( + "qknorm_rope", + *args, + cuda_files=["diffusion/qknorm_rope.cuh"], + cuda_wrappers=[("qknorm_rope", f"QKNormRopeKernel<{args}>::run")], + ) + + +@torch.compiler.assume_constant_result +@cache_once +def can_use_fused_inplace_qknorm_rope( + head_dim: int, + rope_dim: int, + is_neox: bool, + dtype: torch.dtype, +) -> bool: + if head_dim not in (64, 128, 256): + logger.warning(f"Unsupported head_dim={head_dim} for JIT fused QKNorm+RoPE") + return False + if rope_dim <= 0 or rope_dim > head_dim: + logger.warning( + f"Unsupported rope_dim={rope_dim} for head_dim={head_dim} in fused QKNorm+RoPE" + ) + return False + elems_per_thread = head_dim // 32 + if rope_dim % elems_per_thread != 0: + logger.warning( + "rope_dim=%s must be divisible by per-thread width=%s for fused QKNorm+RoPE", + rope_dim, + elems_per_thread, + ) + return False + if is_neox: + rotary_lanes = rope_dim // elems_per_thread + if rotary_lanes < 2 or rotary_lanes & (rotary_lanes - 1): + logger.warning( + "rope_dim=%s yields invalid rotary_lanes=%s for neox fused QKNorm+RoPE; rotary lane count must be a power of 2", + rope_dim, + rotary_lanes, + ) + return False + try: + _jit_qknorm_rope_module(head_dim, rope_dim, is_neox, dtype) + return True + except Exception as e: + logger.warning(f"Failed to load JIT fused QKNorm+RoPE kernel: {e}") + return False + + +@register_custom_op(mutates_args=["q", "k"]) +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: + head_dim = head_dim or q.size(-1) + rope_dim = rope_dim or cos_sin_cache.size(-1) + module = _jit_qknorm_rope_module(head_dim, rope_dim, is_neox, q.dtype) + module.qknorm_rope(q, k, q_weight, k_weight, cos_sin_cache, positions, eps) diff --git a/python/sglang/jit_kernel/tests/diffusion/test_qknorm_rope.py b/python/sglang/jit_kernel/tests/diffusion/test_qknorm_rope.py new file mode 100644 index 000000000..f12a6ae8e --- /dev/null +++ b/python/sglang/jit_kernel/tests/diffusion/test_qknorm_rope.py @@ -0,0 +1,153 @@ +import itertools +import sys + +import pytest +import torch +import triton + +from sglang.jit_kernel.utils import get_ci_test_range +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=44, suite="stage-b-kernel-unit-1-gpu-large") +register_cuda_ci(est_time=176, suite="nightly-kernel-1-gpu", nightly=True) + +DEVICE = "cuda" +DTYPE = torch.bfloat16 +MAX_SEQ_LEN = 131072 +ROPE_BASE = 10000.0 +ATOL = 8e-2 +RTOL = 1e-2 + + +def create_cos_sin_cache( + rotary_dim: int, + max_position: int = MAX_SEQ_LEN, + base: float = ROPE_BASE, +) -> torch.Tensor: + inv_freq = 1.0 / ( + base + ** ( + torch.arange(0, rotary_dim, 2, dtype=torch.float32, device=DEVICE) + / rotary_dim + ) + ) + t = torch.arange(max_position, dtype=torch.float32, device=DEVICE) + freqs = torch.einsum("i,j->ij", t, inv_freq) + return torch.cat((freqs.cos(), freqs.sin()), dim=-1) + + +def split_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, +) -> None: + from flashinfer.rope import apply_rope_with_cos_sin_cache_inplace + + from sglang.jit_kernel.norm import fused_inplace_qknorm + + fused_inplace_qknorm(q, k, q_weight, k_weight) + apply_rope_with_cos_sin_cache_inplace( + positions=positions.long(), + query=q.view(q.shape[0], -1), + key=k.view(k.shape[0], -1), + head_size=q.shape[-1], + cos_sin_cache=cos_sin_cache, + is_neox=is_neox, + ) + + +def fused_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, +) -> None: + from sglang.jit_kernel.diffusion.qknorm_rope import fused_inplace_qknorm_rope + + fused_inplace_qknorm_rope( + q, + k, + q_weight, + k_weight, + cos_sin_cache, + positions, + is_neox=is_neox, + rope_dim=cos_sin_cache.shape[-1], + ) + + +BS_LIST = [2**n for n in range(13)] +BS_LIST += [x + 1 for x in BS_LIST] +BS_LIST = get_ci_test_range(BS_LIST, [1, 9, 129, 257, 2049, 4097]) +HEADS_LIST = get_ci_test_range([8, 16, 24, 32], [8, 24]) +HEAD_DIM_LIST = get_ci_test_range([64, 128, 256], [64, 128, 256]) +IS_NEOX_LIST = [False, True] +POSITION_DTYPES = [torch.int32, torch.int64] +ROPE_DIM_CHOICES = { + 64: [64], + 128: [64, 128], + 256: [64, 128, 256], +} + + +@pytest.mark.parametrize( + "batch_size,num_heads,head_dim,is_neox,position_dtype", + list( + itertools.product( + BS_LIST, + HEADS_LIST, + HEAD_DIM_LIST, + IS_NEOX_LIST, + POSITION_DTYPES, + ) + ), +) +def test_qknorm_rope( + batch_size: int, + num_heads: int, + head_dim: int, + is_neox: bool, + position_dtype: torch.dtype, +) -> None: + rope_dims = ROPE_DIM_CHOICES[head_dim] + for rope_dim in rope_dims: + if is_neox: + elems_per_thread = head_dim // 32 + rotary_lanes = rope_dim // elems_per_thread + if rotary_lanes < 2 or rotary_lanes & (rotary_lanes - 1): + continue + + q = torch.randn(batch_size, num_heads, head_dim, device=DEVICE, dtype=DTYPE) + k = torch.randn(batch_size, num_heads, head_dim, device=DEVICE, dtype=DTYPE) + q_weight = torch.randn(head_dim, device=DEVICE, dtype=DTYPE) + k_weight = torch.randn(head_dim, device=DEVICE, dtype=DTYPE) + positions = torch.randint( + 0, MAX_SEQ_LEN, (batch_size,), device=DEVICE, dtype=position_dtype + ) + cos_sin_cache = create_cos_sin_cache(rope_dim) + + q_ref, k_ref = q.clone(), k.clone() + q_fused, k_fused = q.clone(), k.clone() + + split_qknorm_rope( + q_ref, k_ref, q_weight, k_weight, cos_sin_cache, positions, is_neox + ) + fused_qknorm_rope( + q_fused, k_fused, q_weight, k_weight, cos_sin_cache, positions, is_neox + ) + + # The split baseline mixes a separate BF16 qknorm kernel with FlashInfer RoPE, + # which differs from the fused path by about one BF16 rounding step on H200. + triton.testing.assert_close(q_ref, q_fused, atol=ATOL, rtol=RTOL) + triton.testing.assert_close(k_ref, k_fused, atol=ATOL, rtol=RTOL) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/python/sglang/multimodal_gen/runtime/layers/layernorm.py b/python/sglang/multimodal_gen/runtime/layers/layernorm.py index 62047224f..092a0b805 100644 --- a/python/sglang/multimodal_gen/runtime/layers/layernorm.py +++ b/python/sglang/multimodal_gen/runtime/layers/layernorm.py @@ -4,12 +4,17 @@ # Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/layers/layernorm.py """Custom normalization layers.""" +import os from typing import Optional, Tuple, Union import torch import torch.nn as nn import torch.nn.functional as F +from sglang.jit_kernel.diffusion.qknorm_rope import ( + can_use_fused_inplace_qknorm_rope, + fused_inplace_qknorm_rope, +) from sglang.jit_kernel.diffusion.triton.norm import norm_infer, rms_norm_fn from sglang.jit_kernel.diffusion.triton.rmsnorm_onepass import triton_one_pass_rms_norm from sglang.jit_kernel.diffusion.triton.scale_shift import fuse_scale_shift_kernel @@ -568,6 +573,142 @@ def apply_qk_norm( return q_out, k_out +def apply_qk_norm_with_optional_rope( + q: torch.Tensor, + k: torch.Tensor, + q_norm: "RMSNorm", + k_norm: "RMSNorm", + head_dim: int, + cos_sin_cache: Optional[torch.Tensor] = None, + *, + is_neox: bool = False, + positions: Optional[torch.Tensor] = None, + position_offset: int = 0, + allow_inplace: bool = True, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Apply QK RMSNorm and optionally RoPE when a cos/sin cache is provided.""" + + if cos_sin_cache is None: + return apply_qk_norm( + q=q, + k=k, + q_norm=q_norm, + k_norm=k_norm, + head_dim=head_dim, + allow_inplace=allow_inplace, + ) + + return apply_qk_norm_rope( + q=q, + k=k, + q_norm=q_norm, + k_norm=k_norm, + head_dim=head_dim, + cos_sin_cache=cos_sin_cache, + is_neox=is_neox, + positions=positions, + position_offset=position_offset, + allow_inplace=allow_inplace, + ) + + +def apply_qk_norm_rope( + q: torch.Tensor, + k: torch.Tensor, + q_norm: "RMSNorm", + k_norm: "RMSNorm", + head_dim: int, + cos_sin_cache: torch.Tensor, + *, + is_neox: bool = False, + positions: Optional[torch.Tensor] = None, + position_offset: int = 0, + allow_inplace: bool = True, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Apply QK RMSNorm followed by RoPE, fusing both on supported CUDA shapes.""" + + from sglang.multimodal_gen.runtime.layers.rotary_embedding import ( + apply_flashinfer_rope_qk_inplace, + ) + + if q.dim() != 4 or k.dim() != 4: + raise ValueError( + f"apply_qk_norm_rope expects 4D q/k tensors, got q:{tuple(q.shape)} k:{tuple(k.shape)}" + ) + if q.shape != k.shape: + raise ValueError( + f"apply_qk_norm_rope expects q/k to have the same shape, got {q.shape} vs {k.shape}" + ) + + batch_size, seq_len, _, _ = q.shape + q_eps = q_norm.variance_epsilon + k_eps = k_norm.variance_epsilon + rope_dim = cos_sin_cache.size(-1) + fused_enabled = os.getenv("SGLANG_ENABLE_FUSED_QKNORM_ROPE", "1").lower() not in { + "0", + "false", + "off", + "no", + } + + if positions is None: + pos_1d = torch.arange( + position_offset, + position_offset + seq_len, + device=q.device, + dtype=torch.int64, + ) + positions = pos_1d if batch_size == 1 else pos_1d.repeat(batch_size) + else: + if positions.dim() != 1 or positions.numel() != batch_size * seq_len: + raise ValueError( + f"positions must be 1D of length {batch_size * seq_len}, got shape={tuple(positions.shape)}" + ) + + if ( + fused_enabled + and _is_cuda + and allow_inplace + and (q_eps == k_eps) + and q.dtype in (torch.float16, torch.bfloat16) + and q_norm.weight.dtype == q.dtype + and k_norm.weight.dtype == k.dtype + and q.is_contiguous() + and k.is_contiguous() + and can_use_fused_inplace_qknorm_rope(head_dim, rope_dim, is_neox, q.dtype) + ): + fused_inplace_qknorm_rope( + q=q.reshape(-1, q.shape[-2], head_dim), + k=k.reshape(-1, k.shape[-2], head_dim), + q_weight=q_norm.weight, + k_weight=k_norm.weight, + cos_sin_cache=cos_sin_cache, + positions=positions, + is_neox=is_neox, + eps=q_eps, + head_dim=head_dim, + rope_dim=rope_dim, + ) + return q, k + + q, k = apply_qk_norm( + q=q, + k=k, + q_norm=q_norm, + k_norm=k_norm, + head_dim=head_dim, + allow_inplace=allow_inplace, + ) + return apply_flashinfer_rope_qk_inplace( + q=q, + k=k, + cos_sin_cache=cos_sin_cache, + head_size=head_dim, + is_neox=is_neox, + positions=positions, + ) + + def tensor_parallel_rms_norm(x: torch.Tensor, norm: "RMSNorm") -> torch.Tensor: tp_rank = get_tensor_model_parallel_rank() tp_size = get_tensor_model_parallel_world_size() diff --git a/python/sglang/multimodal_gen/runtime/models/dits/flux.py b/python/sglang/multimodal_gen/runtime/models/dits/flux.py index d4da6984a..f5c35a5f0 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/flux.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/flux.py @@ -29,7 +29,10 @@ from torch.nn import LayerNorm as LayerNorm from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig from sglang.multimodal_gen.runtime.layers.attention import USPAttention -from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm, apply_qk_norm +from sglang.multimodal_gen.runtime.layers.layernorm import ( + RMSNorm, + apply_qk_norm_with_optional_rope, +) from sglang.multimodal_gen.runtime.layers.linear import ( ColumnParallelLinear, MergedColumnParallelLinear, @@ -44,7 +47,6 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config i ) from sglang.multimodal_gen.runtime.layers.rotary_embedding import ( NDRotaryEmbedding, - apply_flashinfer_rope_qk_inplace, ) from sglang.multimodal_gen.runtime.layers.visual_embedding import ( CombinedTimestepGuidanceTextProjEmbeddings, @@ -354,37 +356,7 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin): query = query.unflatten(-1, (self.heads, -1)) key = key.unflatten(-1, (self.heads, -1)) value = value.unflatten(-1, (self.heads, -1)) - query, key = apply_qk_norm( - q=query, - k=key, - q_norm=self.norm_q, - k_norm=self.norm_k, - head_dim=self.head_dim, - allow_inplace=True, - ) - - if self.added_kv_proj_dim is not None: - encoder_query = encoder_query.unflatten(-1, (self.heads, -1)) - encoder_key = encoder_key.unflatten(-1, (self.heads, -1)) - encoder_value = encoder_value.unflatten(-1, (self.heads, -1)) - - encoder_query, encoder_key = apply_qk_norm( - q=encoder_query, - k=encoder_key, - q_norm=self.norm_added_q, - k_norm=self.norm_added_k, - head_dim=self.head_dim, - allow_inplace=True, - ) - - bsz, seq_len, _, _ = query.shape - query = torch.cat([encoder_query, query], dim=1) - key = torch.cat([encoder_key, key], dim=1) - value = torch.cat([encoder_value, value], dim=1) - num_replicated_prefix = ( - num_replicated_prefix or encoder_hidden_states.shape[1] - ) - + cos_sin_cache = None if freqs_cis is not None: cos, sin = freqs_cis cos_sin_cache = torch.cat( @@ -394,8 +366,51 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin): ], dim=-1, ) - query, key = apply_flashinfer_rope_qk_inplace( - query, key, cos_sin_cache, is_neox=False + + if self.added_kv_proj_dim is not None: + encoder_query = encoder_query.unflatten(-1, (self.heads, -1)) + encoder_key = encoder_key.unflatten(-1, (self.heads, -1)) + encoder_value = encoder_value.unflatten(-1, (self.heads, -1)) + + text_seq_len = encoder_query.shape[1] + encoder_query, encoder_key = apply_qk_norm_with_optional_rope( + q=encoder_query, + k=encoder_key, + q_norm=self.norm_added_q, + k_norm=self.norm_added_k, + head_dim=self.head_dim, + cos_sin_cache=cos_sin_cache, + is_neox=False, + allow_inplace=True, + ) + query, key = apply_qk_norm_with_optional_rope( + q=query, + k=key, + q_norm=self.norm_q, + k_norm=self.norm_k, + head_dim=self.head_dim, + cos_sin_cache=cos_sin_cache, + is_neox=False, + position_offset=text_seq_len, + allow_inplace=True, + ) + + query = torch.cat([encoder_query, query], dim=1) + key = torch.cat([encoder_key, key], dim=1) + value = torch.cat([encoder_value, value], dim=1) + num_replicated_prefix = ( + num_replicated_prefix or encoder_hidden_states.shape[1] + ) + else: + query, key = apply_qk_norm_with_optional_rope( + q=query, + k=key, + q_norm=self.norm_q, + k_norm=self.norm_k, + head_dim=self.head_dim, + cos_sin_cache=cos_sin_cache, + is_neox=False, + allow_inplace=True, ) x = self.attn(query, key, value, num_replicated_prefix=num_replicated_prefix) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py b/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py index 66fb8b2df..1651cdba1 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py @@ -23,7 +23,10 @@ from diffusers.models.normalization import AdaLayerNormContinuous from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig from sglang.multimodal_gen.runtime.distributed import divide, get_tp_world_size from sglang.multimodal_gen.runtime.layers.attention import USPAttention -from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm, apply_qk_norm +from sglang.multimodal_gen.runtime.layers.layernorm import ( + RMSNorm, + apply_qk_norm_with_optional_rope, +) from sglang.multimodal_gen.runtime.layers.linear import ( ColumnParallelLinear, MergedColumnParallelLinear, @@ -291,33 +294,7 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin): key = key.unflatten(-1, (self.local_heads, -1)) value = value.unflatten(-1, (self.local_heads, -1)) - query, key = apply_qk_norm( - q=query, - k=key, - q_norm=self.norm_q, - k_norm=self.norm_k, - head_dim=self.head_dim, - allow_inplace=True, - ) - - if self.added_kv_proj_dim is not None: - encoder_query = encoder_query.unflatten(-1, (self.local_heads, -1)) - encoder_key = encoder_key.unflatten(-1, (self.local_heads, -1)) - encoder_value = encoder_value.unflatten(-1, (self.local_heads, -1)) - - encoder_query, encoder_key = apply_qk_norm( - q=encoder_query, - k=encoder_key, - q_norm=self.norm_added_q, - k_norm=self.norm_added_k, - head_dim=self.head_dim, - allow_inplace=True, - ) - - query = torch.cat([encoder_query, query], dim=1) - key = torch.cat([encoder_key, key], dim=1) - value = torch.cat([encoder_value, value], dim=1) - + cos_sin_cache = None if freqs_cis is not None: cos, sin = freqs_cis cos_sin_cache = torch.cat( @@ -327,8 +304,48 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin): ], dim=-1, ) - query, key = apply_flashinfer_rope_qk_inplace( - query, key, cos_sin_cache, is_neox=False + + if self.added_kv_proj_dim is not None: + encoder_query = encoder_query.unflatten(-1, (self.local_heads, -1)) + encoder_key = encoder_key.unflatten(-1, (self.local_heads, -1)) + encoder_value = encoder_value.unflatten(-1, (self.local_heads, -1)) + + text_seq_len = encoder_query.shape[1] + encoder_query, encoder_key = apply_qk_norm_with_optional_rope( + q=encoder_query, + k=encoder_key, + q_norm=self.norm_added_q, + k_norm=self.norm_added_k, + head_dim=self.head_dim, + cos_sin_cache=cos_sin_cache, + is_neox=False, + allow_inplace=True, + ) + query, key = apply_qk_norm_with_optional_rope( + q=query, + k=key, + q_norm=self.norm_q, + k_norm=self.norm_k, + head_dim=self.head_dim, + cos_sin_cache=cos_sin_cache, + is_neox=False, + position_offset=text_seq_len, + allow_inplace=True, + ) + + query = torch.cat([encoder_query, query], dim=1) + key = torch.cat([encoder_key, key], dim=1) + value = torch.cat([encoder_value, value], dim=1) + else: + query, key = apply_qk_norm_with_optional_rope( + q=query, + k=key, + q_norm=self.norm_q, + k_norm=self.norm_k, + head_dim=self.head_dim, + cos_sin_cache=cos_sin_cache, + is_neox=False, + allow_inplace=True, ) num_rep = ( @@ -963,9 +980,7 @@ class Flux2Transformer2DModel(CachableDiT, OffloadableDiTMixin): # 0. Handle input arguments if joint_attention_kwargs is not None: joint_attention_kwargs = joint_attention_kwargs.copy() - lora_scale = joint_attention_kwargs.pop("scale", 1.0) - else: - lora_scale = 1.0 + joint_attention_kwargs.pop("scale", 1.0) num_txt_tokens = encoder_hidden_states.shape[1] diff --git a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py index 79c2db044..ebf2daa5a 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py @@ -30,7 +30,7 @@ from sglang.multimodal_gen.runtime.layers.layernorm import ( LayerNormScaleShift, RMSNorm, ScaleResidualLayerNormScaleShift, - apply_qk_norm, + apply_qk_norm_with_optional_rope, ) from sglang.multimodal_gen.runtime.layers.linear import ( MergedColumnParallelLinear, @@ -626,26 +626,7 @@ class QwenImageCrossAttention(nn.Module): txt_key = txt_key.unflatten(-1, (self.num_heads, -1)) txt_value = txt_value.unflatten(-1, (self.num_heads, -1)) - # Apply QK normalization - if self.qk_norm: - img_query, img_key = apply_qk_norm( - q=img_query, - k=img_key, - q_norm=self.norm_q, - k_norm=self.norm_k, - head_dim=img_query.shape[-1], - allow_inplace=True, - ) - txt_query, txt_key = apply_qk_norm( - q=txt_query, - k=txt_key, - q_norm=self.norm_added_q, - k_norm=self.norm_added_k, - head_dim=txt_query.shape[-1], - allow_inplace=True, - ) - - # Apply RoPE + img_cache = txt_cache = None if image_rotary_emb is not None: if not ( isinstance(image_rotary_emb[0], torch.Tensor) @@ -655,6 +636,28 @@ class QwenImageCrossAttention(nn.Module): img_cache, txt_cache = image_rotary_emb + if self.qk_norm: + img_query, img_key = apply_qk_norm_with_optional_rope( + q=img_query, + k=img_key, + q_norm=self.norm_q, + k_norm=self.norm_k, + head_dim=img_query.shape[-1], + cos_sin_cache=img_cache, + is_neox=False, + allow_inplace=True, + ) + txt_query, txt_key = apply_qk_norm_with_optional_rope( + q=txt_query, + k=txt_key, + q_norm=self.norm_added_q, + k_norm=self.norm_added_k, + head_dim=txt_query.shape[-1], + cos_sin_cache=txt_cache, + is_neox=False, + allow_inplace=True, + ) + elif img_cache is not None and txt_cache is not None: img_query, img_key = apply_flashinfer_rope_qk_inplace( img_query, img_key, img_cache, is_neox=False ) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/zimage.py b/python/sglang/multimodal_gen/runtime/models/dits/zimage.py index 3758e04e6..8a0ab6a58 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/zimage.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/zimage.py @@ -19,7 +19,10 @@ from sglang.multimodal_gen.runtime.layers.attention import ( UlyssesAttention, USPAttention, ) -from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm, apply_qk_norm +from sglang.multimodal_gen.runtime.layers.layernorm import ( + RMSNorm, + apply_qk_norm_with_optional_rope, +) from sglang.multimodal_gen.runtime.layers.linear import ( ColumnParallelLinear, MergedColumnParallelLinear, @@ -256,16 +259,6 @@ class ZImageAttention(nn.Module): k = k.view(*k.shape[:-1], self.local_num_kv_heads, self.head_dim) v = v.view(*v.shape[:-1], self.local_num_kv_heads, self.head_dim) - if self.qk_norm: - q, k = apply_qk_norm( - q=q, - k=k, - q_norm=self.norm_q, - k_norm=self.norm_k, - head_dim=self.head_dim, - allow_inplace=True, - ) - if freqs_cis is not None: cos, sin = freqs_cis if _is_cuda and q.shape == k.shape: @@ -276,12 +269,42 @@ class ZImageAttention(nn.Module): ], dim=-1, ) - q, k = apply_flashinfer_rope_qk_inplace( - q, k, cos_sin_cache, is_neox=False - ) + if self.qk_norm: + q, k = apply_qk_norm_with_optional_rope( + q=q, + k=k, + q_norm=self.norm_q, + k_norm=self.norm_k, + head_dim=self.head_dim, + cos_sin_cache=cos_sin_cache, + is_neox=False, + allow_inplace=True, + ) + else: + q, k = apply_flashinfer_rope_qk_inplace( + q, k, cos_sin_cache, is_neox=False + ) else: + if self.qk_norm: + q, k = apply_qk_norm_with_optional_rope( + q=q, + k=k, + q_norm=self.norm_q, + k_norm=self.norm_k, + head_dim=self.head_dim, + allow_inplace=True, + ) q = _apply_rotary_emb(q, cos, sin, is_neox_style=False) k = _apply_rotary_emb(k, cos, sin, is_neox_style=False) + elif self.qk_norm: + q, k = apply_qk_norm_with_optional_rope( + q=q, + k=k, + q_norm=self.norm_q, + k_norm=self.norm_k, + head_dim=self.head_dim, + allow_inplace=True, + ) if ( num_replicated_suffix > 0