diff --git a/python/sglang/kernels/jit/csrc/diffusion/qwen_qkv_epilogue.cuh b/python/sglang/kernels/jit/csrc/diffusion/qwen_qkv_epilogue.cuh new file mode 100644 index 000000000..35dc8eacf --- /dev/null +++ b/python/sglang/kernels/jit/csrc/diffusion/qwen_qkv_epilogue.cuh @@ -0,0 +1,265 @@ +#pragma once + +#include + +#include +#include +#include +#include +#include + +#include +#include + +namespace sglang { + +namespace qwen_qkv_epilogue { + +constexpr int kHeadDim = 128; +constexpr int kThreads = 256; +constexpr int kWarps = kThreads / device::kWarpThreads; +constexpr int kElemsPerThread = kHeadDim / device::kWarpThreads; +constexpr int kVecSize = kElemsPerThread / 2; + +struct Params { + void* joint_q; + void* joint_k; + void* joint_v; + const void* img_q; + const void* img_k; + const void* img_v; + const void* txt_q; + const void* txt_k; + const void* txt_v; + const void* img_q_weight; + const void* img_k_weight; + const void* txt_q_weight; + const void* txt_k_weight; + const void* img_cache; + const void* txt_cache; + int64_t input_token_stride_bytes; + int64_t output_token_stride_bytes; + int64_t head_stride_bytes; + uint32_t img_tokens; + uint32_t txt_tokens; + uint32_t num_heads; + float img_eps; + float txt_eps; +}; + +__global__ void qwen_qkv_epilogue_kernel(const Params __grid_constant__ params) { + using namespace device; + using Packed = packed_t; + using Storage = AlignedVector; + + const uint32_t lane = threadIdx.x % kWarpThreads; + const uint32_t warp = threadIdx.x / kWarpThreads; + const uint32_t start = blockIdx.x * kWarps + warp; + const uint32_t workers = gridDim.x * kWarps; + const uint32_t total_tokens = params.txt_tokens + params.img_tokens; + const uint32_t token_head_works = total_tokens * params.num_heads; + const uint32_t total_works = 3 * token_head_works; + + for (uint32_t work = start; work < total_works; work += workers) { + const uint32_t kind = work / token_head_works; // 0: Q, 1: K, 2: V. + const uint32_t token_head = work % token_head_works; + const uint32_t joint_token = token_head / params.num_heads; + const uint32_t head = token_head % params.num_heads; + const bool is_text = joint_token < params.txt_tokens; + const uint32_t source_token = is_text ? joint_token : joint_token - params.txt_tokens; + + const void* input_base; + void* output_base; + if (kind == 0) { + input_base = is_text ? params.txt_q : params.img_q; + output_base = params.joint_q; + } else if (kind == 1) { + input_base = is_text ? params.txt_k : params.img_k; + output_base = params.joint_k; + } else { + input_base = is_text ? params.txt_v : params.img_v; + output_base = params.joint_v; + } + + const void* input = + pointer::offset(input_base, source_token * params.input_token_stride_bytes, head * params.head_stride_bytes); + void* output = + pointer::offset(output_base, joint_token * params.output_token_stride_bytes, head * params.head_stride_bytes); + + auto input_vec = load_as(input, lane); + if (kind == 2) { + store_as(output, input_vec, lane); + continue; + } + + const void* weight_base; + if (kind == 0) { + weight_base = is_text ? params.txt_q_weight : params.img_q_weight; + } else { + weight_base = is_text ? params.txt_k_weight : params.img_k_weight; + } + const auto weight_vec = load_as(weight_base, lane); + + 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 eps = is_text ? params.txt_eps : params.img_eps; + 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; + } + + const auto* cache = static_cast(is_text ? params.txt_cache : params.img_cache); + const auto* cos_ptr = cache + source_token * kHeadDim; + const auto* sin_ptr = cos_ptr + kHeadDim / 2; +#pragma unroll + for (uint32_t i = 0; i < kElemsPerThread; i += 2) { + const float x = elems[i]; + const float y = elems[i + 1]; + const uint32_t cache_idx = (lane * kElemsPerThread + i) / 2; + const float cos = __ldg(cos_ptr + cache_idx); + const float sin = __ldg(sin_ptr + cache_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(output, input_vec, lane); + } +} + +struct QwenQKVEpilogueKernel { + static void + run(tvm::ffi::TensorView joint_q, + tvm::ffi::TensorView joint_k, + tvm::ffi::TensorView joint_v, + tvm::ffi::TensorView img_q, + tvm::ffi::TensorView img_k, + tvm::ffi::TensorView img_v, + tvm::ffi::TensorView txt_q, + tvm::ffi::TensorView txt_k, + tvm::ffi::TensorView txt_v, + tvm::ffi::TensorView img_q_weight, + tvm::ffi::TensorView img_k_weight, + tvm::ffi::TensorView txt_q_weight, + tvm::ffi::TensorView txt_k_weight, + tvm::ffi::TensorView img_cache, + tvm::ffi::TensorView txt_cache, + double img_eps, + double txt_eps) { + using namespace host; + + auto NI = SymbolicSize{"img_tokens"}; + auto NT = SymbolicSize{"txt_tokens"}; + auto N = SymbolicSize{"joint_tokens"}; + auto H = SymbolicSize{"num_heads"}; + auto D = SymbolicSize{"head_dim"}; + auto device = SymbolicDevice{}; + D.set_value(kHeadDim); + device.set_options(); + + TensorMatcher({NI, H, D}) + .with_strides({-1, D, 1}) + .with_dtype() + .with_device(device) + .verify(img_q) + .verify(img_k) + .verify(img_v); + TensorMatcher({NT, H, D}) + .with_strides({-1, D, 1}) + .with_dtype() + .with_device(device) + .verify(txt_q) + .verify(txt_k) + .verify(txt_v); + N.set_value(NI.unwrap() + NT.unwrap()); + TensorMatcher({N, H, D}).with_dtype().with_device(device).verify(joint_q).verify(joint_k).verify(joint_v); + TensorMatcher({D}) + .with_dtype() + .with_device(device) + .verify(img_q_weight) + .verify(img_k_weight) + .verify(txt_q_weight) + .verify(txt_k_weight); + TensorMatcher({-1, D}).with_dtype().with_device(device).verify(img_cache); + TensorMatcher({-1, D}).with_dtype().with_device(device).verify(txt_cache); + + RuntimeCheck( + img_q.stride(0) == img_k.stride(0) && img_q.stride(0) == img_v.stride(0), + "image QKV inputs must use the same token stride"); + RuntimeCheck( + txt_q.stride(0) == txt_k.stride(0) && txt_q.stride(0) == txt_v.stride(0), + "text QKV inputs must use the same token stride"); + RuntimeCheck(img_q.stride(0) == txt_q.stride(0), "image/text QKV token strides must match"); + RuntimeCheck( + img_q.stride(1) == kHeadDim && img_k.stride(1) == kHeadDim && img_v.stride(1) == kHeadDim, + "image QKV heads must be contiguous"); + RuntimeCheck( + txt_q.stride(1) == kHeadDim && txt_k.stride(1) == kHeadDim && txt_v.stride(1) == kHeadDim, + "text QKV heads must be contiguous"); + RuntimeCheck(joint_q.is_contiguous(), "joint QKV outputs must be contiguous"); + RuntimeCheck(joint_k.is_contiguous(), "joint QKV outputs must be contiguous"); + RuntimeCheck(joint_v.is_contiguous(), "joint QKV outputs must be contiguous"); + RuntimeCheck(img_cache.size(0) >= NI.unwrap(), "img cache is shorter than img tokens"); + RuntimeCheck(txt_cache.size(0) >= NT.unwrap(), "txt cache is shorter than txt tokens"); + + const uint32_t img_tokens = static_cast(NI.unwrap()); + const uint32_t txt_tokens = static_cast(NT.unwrap()); + const uint32_t num_heads = static_cast(H.unwrap()); + const uint32_t total_works = 3 * (img_tokens + txt_tokens) * num_heads; + if (total_works == 0) return; + + const int64_t head_stride_bytes = kHeadDim * sizeof(bf16_t); + const int64_t input_token_stride_bytes = img_q.stride(0) * sizeof(bf16_t); + const int64_t output_token_stride_bytes = num_heads * head_stride_bytes; + const auto params = Params{ + .joint_q = joint_q.data_ptr(), + .joint_k = joint_k.data_ptr(), + .joint_v = joint_v.data_ptr(), + .img_q = img_q.data_ptr(), + .img_k = img_k.data_ptr(), + .img_v = img_v.data_ptr(), + .txt_q = txt_q.data_ptr(), + .txt_k = txt_k.data_ptr(), + .txt_v = txt_v.data_ptr(), + .img_q_weight = img_q_weight.data_ptr(), + .img_k_weight = img_k_weight.data_ptr(), + .txt_q_weight = txt_q_weight.data_ptr(), + .txt_k_weight = txt_k_weight.data_ptr(), + .img_cache = img_cache.data_ptr(), + .txt_cache = txt_cache.data_ptr(), + .input_token_stride_bytes = input_token_stride_bytes, + .output_token_stride_bytes = output_token_stride_bytes, + .head_stride_bytes = head_stride_bytes, + .img_tokens = img_tokens, + .txt_tokens = txt_tokens, + .num_heads = num_heads, + .img_eps = static_cast(img_eps), + .txt_eps = static_cast(txt_eps), + }; + + const uint32_t sm_count = runtime::get_sm_count(device.unwrap().device_id); + static const uint32_t blocks_per_sm = runtime::get_blocks_per_sm(qwen_qkv_epilogue_kernel, kThreads); + const uint32_t needed_blocks = div_ceil(total_works, uint32_t(kWarps)); + const uint32_t blocks = std::min(blocks_per_sm * sm_count, needed_blocks); + LaunchKernel(blocks, kThreads, device.unwrap())(qwen_qkv_epilogue_kernel, params); + } +}; + +} // namespace qwen_qkv_epilogue + +} // namespace sglang diff --git a/python/sglang/kernels/ops/diffusion/__init__.py b/python/sglang/kernels/ops/diffusion/__init__.py index 892129a2f..5284a16af 100644 --- a/python/sglang/kernels/ops/diffusion/__init__.py +++ b/python/sglang/kernels/ops/diffusion/__init__.py @@ -37,6 +37,7 @@ from sglang.kernels.spec import ( ) _CUDA = frozenset({CapabilityRequirement.CUDA}) +_CUDA_SM100_PLUS = frozenset({CapabilityRequirement.cuda(min_sm=(10, 0))}) _HIP = frozenset({CapabilityRequirement.HIP}) # --------------------------------------------------------------------------- @@ -193,6 +194,13 @@ _SPECS: tuple[tuple[str, KernelBackend, str, frozenset, str], ...] = ( _CUDA, "Fused in-place QK RMS-norm + RoPE.", ), + ( + "diffusion.qwen_qkv_epilogue", + KernelBackend.JIT, + "rope.qwen_qkv_epilogue_jit:try_fused_qwen_qkv_epilogue", + _CUDA_SM100_PLUS, + "Qwen-Image QK RMS-norm, RoPE, and joint QKV writes.", + ), ( "diffusion.ltx2_qknorm_split_rope", KernelBackend.JIT, @@ -426,6 +434,7 @@ _EXPORTS: dict[str, str] = { "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", + "try_fused_qwen_qkv_epilogue": "rope.qwen_qkv_epilogue_jit", "can_use_fused_rope_rotate_half": "rope.rope_rotate_half_bitexact", "fused_rope_rotate_half_bitexact": "rope.rope_rotate_half_bitexact", "can_use_interleaved_rope_fp64": "rope.interleaved_rope_fp64_jit", diff --git a/python/sglang/kernels/ops/diffusion/rope/qwen_qkv_epilogue_jit.py b/python/sglang/kernels/ops/diffusion/rope/qwen_qkv_epilogue_jit.py new file mode 100644 index 000000000..5c1c48596 --- /dev/null +++ b/python/sglang/kernels/ops/diffusion/rope/qwen_qkv_epilogue_jit.py @@ -0,0 +1,148 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +import torch + +from sglang.kernels.jit.utils import cache_once, load_jit + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + +_HEAD_DIM = 128 +_ALIGN = 32 + + +@cache_once +def qwen_qkv_epilogue_module() -> Module: + return load_jit( + "qwen_qkv_epilogue_bf16", + cuda_files=["diffusion/qwen_qkv_epilogue.cuh"], + cuda_wrappers=[ + ( + "qwen_qkv_epilogue", + "qwen_qkv_epilogue::QwenQKVEpilogueKernel::run", + ) + ], + ) + + +def _qkv_tensor(tensor: torch.Tensor, like: torch.Tensor | None = None) -> bool: + return ( + isinstance(tensor, torch.Tensor) + and tensor.is_cuda + and tensor.dtype == torch.bfloat16 + and tensor.ndim == 4 + and tensor.shape[0] == 1 + and tensor.shape[-1] == _HEAD_DIM + and tensor.numel() > 0 + and tensor.stride(-1) == 1 + and tensor.stride(-2) == _HEAD_DIM + and tensor.data_ptr() % _ALIGN == 0 + and ( + like is None + or ( + tensor.device == like.device + and tensor.shape[0] == like.shape[0] + and tensor.shape[1] == like.shape[1] + and tensor.shape[2:] == like.shape[2:] + and tensor.stride(1) == like.stride(1) + ) + ) + ) + + +def try_fused_qwen_qkv_epilogue( + img_q: torch.Tensor, + img_k: torch.Tensor, + img_v: torch.Tensor, + txt_q: torch.Tensor, + txt_k: torch.Tensor, + txt_v: torch.Tensor, + img_q_weight: torch.Tensor, + img_k_weight: torch.Tensor, + txt_q_weight: torch.Tensor, + txt_k_weight: torch.Tensor, + img_cache: torch.Tensor, + txt_cache: torch.Tensor, + img_eps: float, + txt_eps: float, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None: + """Fuse Q/K normalization, RoPE, and QKV joint-buffer writes. + + The caller retains an explicit unfused path for every unsupported shape, + layout, architecture, or compilation mode. + """ + if torch.compiler.is_compiling(): + return None + if not ( + _qkv_tensor(img_q) + and _qkv_tensor(img_k, img_q) + and _qkv_tensor(img_v, img_q) + and _qkv_tensor(txt_q) + and _qkv_tensor(txt_k, txt_q) + and _qkv_tensor(txt_v, txt_q) + and img_q.shape[2] == txt_q.shape[2] + and torch.version.cuda is not None + and torch.cuda.get_device_capability(img_q.device)[0] >= 10 + ): + return None + + heads = img_q.shape[2] + weights = [] + # RMSNorm weights are shared across heads, unlike projection biases. + for tensor in (img_q_weight, img_k_weight, txt_q_weight, txt_k_weight): + if not ( + isinstance(tensor, torch.Tensor) + and tensor.is_cuda + and tensor.device == img_q.device + and tensor.dtype == torch.bfloat16 + and tensor.shape == (_HEAD_DIM,) + and tensor.is_contiguous() + and tensor.data_ptr() % _ALIGN == 0 + ): + return None + weights.append(tensor) + + if not ( + isinstance(img_cache, torch.Tensor) + and isinstance(txt_cache, torch.Tensor) + and img_cache.is_cuda + and txt_cache.is_cuda + and img_cache.device == img_q.device + and txt_cache.device == img_q.device + and img_cache.dtype == torch.float32 + and txt_cache.dtype == torch.float32 + and img_cache.ndim == 2 + and txt_cache.ndim == 2 + and img_cache.shape[1] == _HEAD_DIM + and txt_cache.shape[1] == _HEAD_DIM + and img_cache.shape[0] >= img_q.shape[1] + and txt_cache.shape[0] >= txt_q.shape[1] + and img_cache.is_contiguous() + and txt_cache.is_contiguous() + ): + return None + + joint_shape = (1, txt_q.shape[1] + img_q.shape[1], heads, _HEAD_DIM) + joint_q = torch.empty(joint_shape, dtype=img_q.dtype, device=img_q.device) + joint_k = torch.empty_like(joint_q) + joint_v = torch.empty_like(joint_q) + qwen_qkv_epilogue_module().qwen_qkv_epilogue( + joint_q.view(-1, heads, _HEAD_DIM), + joint_k.view(-1, heads, _HEAD_DIM), + joint_v.view(-1, heads, _HEAD_DIM), + img_q.view(-1, heads, _HEAD_DIM), + img_k.view(-1, heads, _HEAD_DIM), + img_v.view(-1, heads, _HEAD_DIM), + txt_q.view(-1, heads, _HEAD_DIM), + txt_k.view(-1, heads, _HEAD_DIM), + txt_v.view(-1, heads, _HEAD_DIM), + *weights, + img_cache, + txt_cache, + float(img_eps), + float(txt_eps), + ) + return joint_q, joint_k, joint_v diff --git a/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py b/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py index b11684405..8b31410a6 100755 --- a/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py +++ b/python/sglang/multimodal_gen/runtime/layers/quantization/modelopt_quant.py @@ -39,7 +39,6 @@ from sglang.srt.layers.quantization.modelopt_quant import ( ) from sglang.srt.layers.quantization.utils import ( convert_to_channelwise, - is_layer_skipped, requantize_with_max_scale, ) from sglang.srt.layers.utils.common import copy_or_rebind_param @@ -123,10 +122,7 @@ class ModelOptQuantConfig(QuantizationConfig): from sglang.multimodal_gen.runtime.layers.linear import LinearBase if isinstance(layer, LinearBase): - if self.is_layer_excluded(prefix) or ( - self.packed_modules_mapping - and is_layer_skipped(prefix, [], self.packed_modules_mapping) - ): + if self.is_layer_excluded(prefix) or self._is_packed_layer_excluded(prefix): return UnquantizedLinearMethod() return Linear(self) return None @@ -163,6 +159,24 @@ class ModelOptQuantConfig(QuantizationConfig): return True return False + def _is_packed_layer_excluded(self, prefix: str) -> bool: + proj_name = prefix.rsplit(".", 1)[-1] + shard_names = self.packed_modules_mapping.get(proj_name) + if shard_names is None: + return False + + base_prefix = prefix[: -len(proj_name)] + shard_exclusions = [ + self.is_layer_excluded(base_prefix + shard_name) + for shard_name in shard_names + ] + if any(shard_exclusions) and not all(shard_exclusions): + raise ValueError( + f"Detected some but not all shards of {prefix} are quantized. " + "All shards of fused layers must have the same precision." + ) + return all(shard_exclusions) + class ModelOptFp8Config(ModelOptQuantConfig): """Config class for ModelOpt FP8 diffusion checkpoints.""" @@ -475,16 +489,36 @@ class ModelOptFp8LinearMethod(LinearMethodBase): if input_scale is not None: copy_or_rebind_param(layer, "input_scale", input_scale) - max_w_scale, quantized_weight = requantize_with_max_scale( - weight, layer.weight_scale, layer.logical_widths + complete_shard_scales = ( + self.cutlass_fp8_supported + and len(layer.logical_widths) > 1 + and bool( + torch.all( + layer.weight_scale > torch.finfo(torch.float8_e4m3fn).min + ).item() + ) ) + if complete_shard_scales: + # CUTLASS accepts a scale per output channel. Preserve each + # checkpoint shard's original FP8 values and scale instead of + # requantizing all packed shards to the largest scale. + quantized_weight = weight + processed_weight_scale = convert_to_channelwise( + layer.weight_scale, layer.logical_widths + ) + else: + processed_weight_scale, quantized_weight = requantize_with_max_scale( + weight, layer.weight_scale, layer.logical_widths + ) + if self.cutlass_fp8_supported: + processed_weight_scale = convert_to_channelwise( + processed_weight_scale, layer.logical_widths + ) # Preserve the parameter subclass metadata while rebinding to the # transposed FP8 view expected by the runtime. layer.weight.data = quantized_weight.t().detach() layer.weight.requires_grad_(False) - if self.cutlass_fp8_supported: - max_w_scale = convert_to_channelwise(max_w_scale, layer.logical_widths) - copy_or_rebind_param(layer, "weight_scale", max_w_scale) + copy_or_rebind_param(layer, "weight_scale", processed_weight_scale) copy_or_rebind_param(layer, "input_scale", layer.input_scale.max()) def apply( diff --git a/python/sglang/multimodal_gen/runtime/models/dits/common.py b/python/sglang/multimodal_gen/runtime/models/dits/common.py index 7c05fe98b..70cba347a 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/common.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/common.py @@ -26,6 +26,8 @@ def get_qkv_projections( attn: Any, hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor | None = None, + *, + make_contiguous: bool = True, ) -> tuple[ torch.Tensor, torch.Tensor, @@ -41,10 +43,15 @@ def get_qkv_projections( set by those blocks' constructors, and ``use_fused_added_qkv`` whenever ``added_kv_proj_dim`` is not ``None`` — direct attribute access so a renamed flag fails loudly instead of silently unfusing. + + ``make_contiguous=False`` preserves zero-copy views for a caller that can + consume packed projection output strides directly. """ if attn.use_fused_qkv: qkv, _ = attn.to_qkv(hidden_states) - query, key, value = [t.contiguous() for t in qkv.chunk(3, dim=-1)] + query, key, value = qkv.chunk(3, dim=-1) + if make_contiguous: + query, key, value = [t.contiguous() for t in (query, key, value)] else: query, _ = attn.to_q(hidden_states) key, _ = attn.to_k(hidden_states) @@ -54,9 +61,11 @@ def get_qkv_projections( if encoder_hidden_states is not None and attn.added_kv_proj_dim is not None: if attn.use_fused_added_qkv: added_qkv, _ = attn.to_added_qkv(encoder_hidden_states) - encoder_query, encoder_key, encoder_value = [ - t.contiguous() for t in added_qkv.chunk(3, dim=-1) - ] + encoder_query, encoder_key, encoder_value = added_qkv.chunk(3, dim=-1) + if make_contiguous: + encoder_query, encoder_key, encoder_value = [ + t.contiguous() for t in (encoder_query, encoder_key, encoder_value) + ] else: encoder_query, _ = attn.add_q_proj(encoder_hidden_states) encoder_key, _ = attn.add_k_proj(encoder_hidden_states) 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 2b6fa6345..eefec49fc 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py @@ -27,6 +27,7 @@ from sglang.kernels.ops.diffusion import ( try_fused_bias_mul_add, try_fused_bias_scale_residual_norm_scale_shift, try_fused_norm_scale_shift_fp8, + try_fused_qwen_qkv_epilogue, try_fused_scale_residual_norm_scale_shift_fp8, ) from sglang.multimodal_gen.configs.models.dits.qwenimage import QwenImageDitConfig @@ -265,6 +266,48 @@ def _qwen_modulation_cache_key( ) +def _modelopt_quant_name( + quant_config: Optional[QuantizationConfig], +) -> str | None: + return None if quant_config is None else quant_config.get_name() + + +_MODEL_OPT_FP8_QKV_PARAM_NAMES_MAPPING = { + # ModelOpt FP8 uses one QKV GEMM per stream. Merge the three Diffusers + # projections and their static scales into MergedColumnParallelLinear. + r"^(transformer_blocks\.\d+\.attn)\.to_q\.(weight|bias|weight_scale|input_scale)$": ( + r"\1.to_qkv.\2", + 0, + 3, + ), + r"^(transformer_blocks\.\d+\.attn)\.to_k\.(weight|bias|weight_scale|input_scale)$": ( + r"\1.to_qkv.\2", + 1, + 3, + ), + r"^(transformer_blocks\.\d+\.attn)\.to_v\.(weight|bias|weight_scale|input_scale)$": ( + r"\1.to_qkv.\2", + 2, + 3, + ), + r"^(transformer_blocks\.\d+\.attn)\.add_q_proj\.(weight|bias|weight_scale|input_scale)$": ( + r"\1.to_added_qkv.\2", + 0, + 3, + ), + r"^(transformer_blocks\.\d+\.attn)\.add_k_proj\.(weight|bias|weight_scale|input_scale)$": ( + r"\1.to_added_qkv.\2", + 1, + 3, + ), + r"^(transformer_blocks\.\d+\.attn)\.add_v_proj\.(weight|bias|weight_scale|input_scale)$": ( + r"\1.to_added_qkv.\2", + 2, + 3, + ), +} + + class QwenTimestepProjEmbeddings(nn.Module): def __init__(self, embedding_dim, use_additional_t_cond=False): super().__init__() @@ -681,8 +724,14 @@ class QwenImageCrossAttention(nn.Module): self.added_kv_proj_dim = added_kv_proj_dim self.prefix = prefix self.defer_output_bias = _defer_modelopt_output_bias(quant_config) - - self.use_fused_qkv = isinstance(quant_config, NunchakuConfig) + quant_name = _modelopt_quant_name(quant_config) + self.use_fused_qkv_epilogue = quant_name in { + "modelopt_fp4", + "modelopt_fp8", + } + self.use_fused_qkv = ( + isinstance(quant_config, NunchakuConfig) or quant_name == "modelopt_fp8" + ) self.inner_dim = out_dim if out_dim is not None else head_dim * num_heads self.inner_kv_dim = self.inner_dim @@ -733,7 +782,9 @@ class QwenImageCrossAttention(nn.Module): self.norm_k = RMSNorm(head_dim, eps=eps) if qk_norm else nn.Identity() if added_kv_proj_dim is not None: - self.use_fused_added_qkv = isinstance(quant_config, NunchakuConfig) + self.use_fused_added_qkv = ( + isinstance(quant_config, NunchakuConfig) or quant_name == "modelopt_fp8" + ) if self.use_fused_added_qkv: self.to_added_qkv = MergedColumnParallelLinear( added_kv_proj_dim, @@ -854,7 +905,12 @@ class QwenImageCrossAttention(nn.Module): txt_query, txt_key, txt_value, - ) = _get_qkv_projections(self, hidden_states, encoder_hidden_states) + ) = _get_qkv_projections( + self, + hidden_states, + encoder_hidden_states, + make_contiguous=not self.use_fused_qkv_epilogue, + ) # Reshape for multi-head attention img_query = img_query.unflatten(-1, (self.local_num_heads, self.head_dim)) @@ -875,35 +931,71 @@ 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=self.head_dim, - 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=self.head_dim, - 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 - ) - txt_query, txt_key = apply_flashinfer_rope_qk_inplace( - txt_query, txt_key, txt_cache, is_neox=False + joint_qkv = None + if ( + self.use_fused_qkv_epilogue + and self.qk_norm + and img_cache is not None + and txt_cache is not None + and not sp_text_sharded + and sp_txt_pad == 0 + ): + joint_qkv = try_fused_qwen_qkv_epilogue( + img_query, + img_key, + img_value, + txt_query, + txt_key, + txt_value, + self.norm_q.weight, + self.norm_k.weight, + self.norm_added_q.weight, + self.norm_added_k.weight, + img_cache, + txt_cache, + self.norm_q.variance_epsilon, + self.norm_added_q.variance_epsilon, ) + if joint_qkv is None: + # Fused ModelOpt FP8 projections expose zero-copy Q/K/V views into + # one packed GEMM output. Unsupported epilogue cases keep the old + # contiguous contract before entering the generic QKNorm/RoPE path. + img_query, img_key, img_value = [ + tensor.contiguous() for tensor in (img_query, img_key, img_value) + ] + txt_query, txt_key, txt_value = [ + tensor.contiguous() for tensor in (txt_query, txt_key, txt_value) + ] + 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=self.head_dim, + 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=self.head_dim, + 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 + ) + txt_query, txt_key = apply_flashinfer_rope_qk_inplace( + txt_query, txt_key, txt_cache, is_neox=False + ) + # Joint order [text, image]; join_seqs relocates any SP text tail-pad # behind the image (see sp_shard.join_seqs for why). seg_qkv = None @@ -923,7 +1015,9 @@ class QwenImageCrossAttention(nn.Module): img_value, sp_txt_pad, ) - if seg_qkv is not None: + if joint_qkv is not None: + joint_query, joint_key, joint_value = joint_qkv + elif seg_qkv is not None: joint_query, joint_key, joint_value = seg_qkv else: joint_query = join_seqs(txt_query, img_query, sp_txt_pad) @@ -1807,8 +1901,21 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): _repeated_blocks = ["QwenImageTransformerBlock"] param_names_mapping = QwenImageDitConfig().arch_config.param_names_mapping + packed_modules_mapping = { + "to_qkv": ["to_q", "to_k", "to_v"], + "to_added_qkv": ["add_q_proj", "add_k_proj", "add_v_proj"], + } _fsdp_shard_conditions = [is_transformer_block] + @classmethod + def get_param_names_mapping_for_quant_config( + cls, quant_config: Optional[QuantizationConfig] + ) -> dict: + mapping = dict(cls.param_names_mapping) + if _modelopt_quant_name(quant_config) == "modelopt_fp8": + mapping.update(_MODEL_OPT_FP8_QKV_PARAM_NAMES_MAPPING) + return mapping + @classmethod def get_nunchaku_quant_rules(cls) -> dict[str, list[str]]: return { @@ -1839,6 +1946,12 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): quant_config: Optional[QuantizationConfig] = None, ): super().__init__(config=config, hf_config=hf_config) + # Only ModelOpt FP8 constructs packed QKV modules for checkpoints with + # Diffusers-style split Q/K/V names. Keep the mapping instance-local so + # eager and NVFP4 checkpoints still target their split projections. + self.param_names_mapping = self.get_param_names_mapping_for_quant_config( + quant_config + ) arch = self.config patch_size = arch.patch_size in_channels = arch.in_channels diff --git a/python/sglang/multimodal_gen/test/unit/test_modelopt_fp8_layerwise_offload_load.py b/python/sglang/multimodal_gen/test/unit/test_modelopt_fp8_layerwise_offload_load.py index 3f42243fa..45a72c656 100644 --- a/python/sglang/multimodal_gen/test/unit/test_modelopt_fp8_layerwise_offload_load.py +++ b/python/sglang/multimodal_gen/test/unit/test_modelopt_fp8_layerwise_offload_load.py @@ -4,6 +4,7 @@ layerwise offload: requantize_with_max_scale() runs scaled_fp8_quant(), a CUDA-only kernel, so a CPU-resident postprocess must never come back.""" import unittest +from unittest.mock import patch import torch from torch import nn @@ -82,9 +83,6 @@ class TestModelOptFp8LayerwiseOffloadLoad(unittest.TestCase): ensure_distributed_env_defaults() maybe_init_distributed_environment_and_model_parallel(tp_size=1, sp_size=1) - state_dict, weight_ref = _make_serialized_fp8_checkpoint() - expected_max_scale = state_dict["qkv.weight_scale"].max() - # The plan a layerwise-offload component gets when # _needs_device_weight_postprocess() returns True: load and postprocess # on GPU, then defer the CPU placement. @@ -95,50 +93,79 @@ class TestModelOptFp8LayerwiseOffloadLoad(unittest.TestCase): ) self.assertTrue(load_plan.defer_cpu_placement) - model = fsdp_load.maybe_load_fsdp_model( - model_cls=_FusedFp8Model, - init_params={ - "quant_config": ModelOptFp8Config(is_checkpoint_fp8_serialized=True) - }, - weight_dir_list=[], - device=torch.device("cuda"), - hsdp_replicate_dim=1, - hsdp_shard_dim=1, - param_dtype=torch.bfloat16, - reduce_dtype=torch.float32, - component_starts_on_cpu=True, - weight_load_plan=load_plan, - weights_iterator=iter(state_dict.items()), - ) + for cutlass_supported in (False, True): + with self.subTest(cutlass_supported=cutlass_supported): + state_dict, weight_ref = _make_serialized_fp8_checkpoint() + checkpoint_weight = state_dict["qkv.weight"].clone() + checkpoint_scales = state_dict["qkv.weight_scale"].clone() + expected_max_scale = checkpoint_scales.max() - # Postprocess ran: the weight was requantized to the shared max scale - # and rebound transposed. - weight = model.qkv.weight - self.assertEqual(weight.dtype, torch.float8_e4m3fn) - self.assertEqual(tuple(weight.shape), (_IN_FEATURES, 2 * _SHARD_OUT)) - weight_scale = model.qkv.weight_scale - torch.testing.assert_close( - weight_scale.flatten(), - expected_max_scale.expand(weight_scale.numel()), - check_device=False, - ) - torch.testing.assert_close( - model.qkv.input_scale.flatten().max(), torch.tensor(0.5), check_device=False - ) + with patch( + "sglang.multimodal_gen.runtime.layers.quantization." + "modelopt_quant.cutlass_fp8_supported", + return_value=cutlass_supported, + ): + model = fsdp_load.maybe_load_fsdp_model( + model_cls=_FusedFp8Model, + init_params={ + "quant_config": ModelOptFp8Config( + is_checkpoint_fp8_serialized=True + ) + }, + weight_dir_list=[], + device=torch.device("cuda"), + hsdp_replicate_dim=1, + hsdp_shard_dim=1, + param_dtype=torch.bfloat16, + reduce_dtype=torch.float32, + component_starts_on_cpu=True, + weight_load_plan=load_plan, + weights_iterator=iter(state_dict.items()), + ) - # The round trip through both quantizations stays close to the source. - # Loose on purpose: this guards against garbage (wrong scale, wrong - # shard order), not fp8 precision. - dequant = weight.t().float().cpu() * expected_max_scale - torch.testing.assert_close( - dequant, weight_ref, rtol=0.5, atol=float(expected_max_scale) * 8 - ) + # Both paths rebind the runtime weight transposed. CUTLASS can + # consume a channelwise scale, so it preserves the checkpoint's + # FP8 shards; the fallback requantizes them to one max scale. + weight = model.qkv.weight + self.assertEqual(weight.dtype, torch.float8_e4m3fn) + self.assertEqual(tuple(weight.shape), (_IN_FEATURES, 2 * _SHARD_OUT)) + weight_scale = model.qkv.weight_scale.flatten() + if cutlass_supported: + expected_scales = torch.repeat_interleave( + checkpoint_scales, _SHARD_OUT + ) + self.assertTrue(torch.equal(weight.t(), checkpoint_weight)) + else: + expected_scales = expected_max_scale.expand(weight_scale.numel()) + torch.testing.assert_close( + weight_scale, + expected_scales, + check_device=False, + ) + torch.testing.assert_close( + model.qkv.input_scale.flatten().max(), + torch.tensor(0.5), + check_device=False, + ) - # Layerwise offload contract: the component lands on CPU afterwards, - # with the non-checkpoint buffer rebuilt. - self.assertEqual(weight.device.type, "cpu") - self.assertFalse(model.inv_freq.is_meta) - self.assertEqual(model.inv_freq.device.type, "cpu") + # The round trip stays close to the source. Loose on purpose: + # this guards against garbage (wrong scale or shard order), not + # FP8 precision. + dequant = weight.t().float().cpu() * weight_scale.float().cpu().view( + -1, 1 + ) + torch.testing.assert_close( + dequant, + weight_ref, + rtol=0.5, + atol=float(expected_max_scale) * 8, + ) + + # Layerwise offload contract: the component lands on CPU + # afterwards, with the non-checkpoint buffer rebuilt. + self.assertEqual(weight.device.type, "cpu") + self.assertFalse(model.inv_freq.is_meta) + self.assertEqual(model.inv_freq.device.type, "cpu") if __name__ == "__main__": diff --git a/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py b/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py index 23f557c0f..98e63b8f7 100644 --- a/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py +++ b/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py @@ -79,6 +79,7 @@ from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import ( ModelOptFp4Config, ModelOptFp4LinearMethod, ModelOptFp8Config, + ModelOptFp8LinearMethod, _prepare_nvfp4_weight_bytes, ) from sglang.multimodal_gen.runtime.layers.quantization.mxfp8 import MXFP8Config @@ -103,9 +104,16 @@ from sglang.multimodal_gen.runtime.loader.transformer_load_utils import ( resolve_transformer_checkpoint_files, resolve_transformer_quant_load_spec, ) +from sglang.multimodal_gen.runtime.loader.utils import ( + get_param_names_mapping, + hf_to_custom_state_dict, +) from sglang.multimodal_gen.runtime.loader.weight_load_plan import WeightLoadPlan from sglang.multimodal_gen.runtime.models.dits.flux import FluxSingleTransformerBlock from sglang.multimodal_gen.runtime.models.dits.minimax_h3 import MiniMaxH3DiTModel +from sglang.multimodal_gen.runtime.models.dits.qwen_image import ( + QwenImageTransformer2DModel, +) from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum from sglang.multimodal_gen.runtime.platforms.interface import DeviceCapability from sglang.multimodal_gen.runtime.utils.quantization_utils import ( @@ -149,6 +157,169 @@ def _make_quant_config(name: str, **attrs): class TestTransformerQuantHelpers(unittest.TestCase): + def test_modelopt_fp8_packed_cutlass_preserves_checkpoint_shard_scales(self): + method = ModelOptFp8LinearMethod( + ModelOptFp8Config(is_checkpoint_fp8_serialized=True) + ) + method.cutlass_fp8_supported = True + layer = torch.nn.Module() + layer.logical_widths = [2, 2, 2] + weight = ( + torch.arange(24, dtype=torch.float32).reshape(6, 4).to(torch.float8_e4m3fn) + ) + layer.register_parameter( + "weight", torch.nn.Parameter(weight.clone(), requires_grad=False) + ) + layer.register_parameter( + "weight_scale", + torch.nn.Parameter( + torch.tensor([0.1, 0.2, 0.3], dtype=torch.float32), + requires_grad=False, + ), + ) + layer.register_parameter( + "input_scale", + torch.nn.Parameter(torch.ones(3, dtype=torch.float32), requires_grad=False), + ) + + method.process_weights_after_loading(layer) + + torch.testing.assert_close(layer.weight, weight.t(), rtol=0, atol=0) + torch.testing.assert_close( + layer.weight_scale, + torch.tensor([[0.1], [0.1], [0.2], [0.2], [0.3], [0.3]]), + ) + torch.testing.assert_close(layer.input_scale, torch.tensor(1.0)) + + def test_modelopt_fp8_packed_cutlass_requantizes_incomplete_shard_scales(self): + method = ModelOptFp8LinearMethod( + ModelOptFp8Config(is_checkpoint_fp8_serialized=True) + ) + method.cutlass_fp8_supported = True + layer = torch.nn.Module() + layer.logical_widths = [2, 2, 2] + weight = ( + torch.arange(24, dtype=torch.float32).reshape(6, 4).to(torch.float8_e4m3fn) + ) + layer.register_parameter( + "weight", torch.nn.Parameter(weight.clone(), requires_grad=False) + ) + layer.register_parameter( + "weight_scale", + torch.nn.Parameter( + torch.tensor( + [0.1, torch.finfo(torch.float32).min, 0.3], + dtype=torch.float32, + ), + requires_grad=False, + ), + ) + layer.register_parameter( + "input_scale", + torch.nn.Parameter(torch.ones(3, dtype=torch.float32), requires_grad=False), + ) + + with patch( + "sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant." + "requantize_with_max_scale", + return_value=(torch.tensor(0.3), weight.clone()), + ) as requantize: + method.process_weights_after_loading(layer) + + requantize.assert_called_once() + torch.testing.assert_close(layer.weight, weight.t(), rtol=0, atol=0) + torch.testing.assert_close( + layer.weight_scale, torch.full((6, 1), 0.3), rtol=0, atol=0 + ) + + def test_modelopt_packed_layer_requires_consistent_shard_precision(self): + prefix = "blocks.0.attn.to_qkv" + mapping = {"to_qkv": ["to_q", "to_k", "to_v"]} + layer = LinearBase(input_size=16, output_size=48) + + quantized = ModelOptFp8Config( + is_checkpoint_fp8_serialized=True, + packed_modules_mapping=mapping, + ) + self.assertIsInstance( + quantized.get_quant_method(layer, prefix), ModelOptFp8LinearMethod + ) + + excluded = ModelOptFp8Config( + is_checkpoint_fp8_serialized=True, + exclude_modules=[ + "blocks.0.attn.to_q", + "blocks.0.attn.to_k", + "blocks.0.attn.to_v", + ], + packed_modules_mapping=mapping, + ) + self.assertIsInstance( + excluded.get_quant_method(layer, prefix), UnquantizedLinearMethod + ) + + partial = ModelOptFp8Config( + is_checkpoint_fp8_serialized=True, + exclude_modules=["blocks.0.attn.to_q"], + packed_modules_mapping=mapping, + ) + with self.assertRaisesRegex(ValueError, "some but not all shards"): + partial.get_quant_method(layer, prefix) + + def test_qwen_modelopt_fp8_qkv_checkpoint_tensors_are_merged(self): + mapping = get_param_names_mapping( + QwenImageTransformer2DModel.get_param_names_mapping_for_quant_config( + ModelOptFp8Config(is_checkpoint_fp8_serialized=True) + ) + ) + prefix = "transformer_blocks.0.attn" + source = {} + for shard_id, shard_name in enumerate(("q", "k", "v")): + source[f"{prefix}.to_{shard_name}.weight"] = torch.full( + (2, 3), shard_id + 1, dtype=torch.float8_e4m3fn + ) + source[f"{prefix}.to_{shard_name}.bias"] = torch.full( + (2,), shard_id + 1, dtype=torch.bfloat16 + ) + source[f"{prefix}.to_{shard_name}.weight_scale"] = torch.tensor( + [0.1 * (shard_id + 1)], dtype=torch.float32 + ) + source[f"{prefix}.to_{shard_name}.input_scale"] = torch.tensor( + [0.2 * (shard_id + 1)], dtype=torch.float32 + ) + + merged, _ = hf_to_custom_state_dict(source, mapping) + + self.assertEqual(merged[f"{prefix}.to_qkv.weight"].shape, (6, 3)) + self.assertEqual(merged[f"{prefix}.to_qkv.bias"].shape, (6,)) + torch.testing.assert_close( + merged[f"{prefix}.to_qkv.weight_scale"], + torch.tensor([0.1, 0.2, 0.3], dtype=torch.float32), + ) + torch.testing.assert_close( + merged[f"{prefix}.to_qkv.input_scale"], + torch.tensor([0.2, 0.4, 0.6], dtype=torch.float32), + ) + self.assertEqual( + QwenImageTransformer2DModel.packed_modules_mapping["to_qkv"], + ["to_q", "to_k", "to_v"], + ) + + def test_qwen_non_fp8_qkv_checkpoint_tensors_are_not_merged(self): + prefix = "transformer_blocks.0.attn" + source_name = f"{prefix}.to_q.weight" + for quant_config in (None, ModelOptFp4Config()): + with self.subTest(quant_config=quant_config): + mapping = get_param_names_mapping( + QwenImageTransformer2DModel.get_param_names_mapping_for_quant_config( + quant_config + ) + ) + target_name, merge_index, total_shards = mapping(source_name) + self.assertEqual(target_name, source_name) + self.assertIsNone(merge_index) + self.assertIsNone(total_shards) + @patch( "sglang.multimodal_gen.runtime.loader.transformer_load_utils.build_nvfp4_config_from_safetensors_list", return_value=None, @@ -1021,8 +1192,8 @@ class TestTransformerQuantHelpers(unittest.TestCase): warning.assert_called_once() def test_modelopt_fp8_always_needs_device_weight_postprocess(self): - # Even a serialized checkpoint requantizes fused shards through - # scaled_fp8_quant(), which cannot process CPU tensors. + # Serialized checkpoints still transpose weights and may requantize + # packed shards through scaled_fp8_quant() on the runtime device. self.assertTrue( _needs_device_weight_postprocess( ModelOptFp8Config(is_checkpoint_fp8_serialized=True) diff --git a/test/registered/kernels/ops/diffusion/test_qwen_qkv_epilogue.py b/test/registered/kernels/ops/diffusion/test_qwen_qkv_epilogue.py new file mode 100644 index 000000000..ca1e356f9 --- /dev/null +++ b/test/registered/kernels/ops/diffusion/test_qwen_qkv_epilogue.py @@ -0,0 +1,194 @@ +import sys +from unittest.mock import patch + +import pytest +import torch + +from sglang.kernels.ops.diffusion import try_fused_qwen_qkv_epilogue +from sglang.multimodal_gen.runtime.layers.layernorm import ( + RMSNorm, + apply_qk_norm_with_optional_rope, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=15, stage="base-b-kernel-unit", runner_config="4-gpu-b200") + +pytestmark = pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] < 10, + reason="Qwen-Image QKV epilogue requires SM100+", +) + + +@pytest.fixture(autouse=True) +def _seed_cuda(): + torch.cuda.manual_seed(0) + + +def test_qwen_qkv_epilogue_is_bit_exact(): + heads = 4 + head_dim = 128 + img_tokens = 17 + txt_tokens = 7 + img_qkv = [ + torch.randn( + 1, + img_tokens, + heads, + head_dim, + device="cuda", + dtype=torch.bfloat16, + ) + for _ in range(3) + ] + txt_qkv = [ + torch.randn( + 1, + txt_tokens, + heads, + head_dim, + device="cuda", + dtype=torch.bfloat16, + ) + for _ in range(3) + ] + norms = [ + RMSNorm(head_dim, eps=1e-6).to(device="cuda", dtype=torch.bfloat16) + for _ in range(4) + ] + + def cache(tokens): + angles = torch.randn(tokens, head_dim // 2, device="cuda") + return torch.cat([angles.cos(), angles.sin()], dim=-1).contiguous() + + img_cache = cache(img_tokens) + txt_cache = cache(txt_tokens) + + img_reference = [tensor.clone() for tensor in img_qkv] + txt_reference = [tensor.clone() for tensor in txt_qkv] + img_reference[0], img_reference[1] = apply_qk_norm_with_optional_rope( + img_reference[0], + img_reference[1], + norms[0], + norms[1], + head_dim, + img_cache, + is_neox=False, + ) + txt_reference[0], txt_reference[1] = apply_qk_norm_with_optional_rope( + txt_reference[0], + txt_reference[1], + norms[2], + norms[3], + head_dim, + txt_cache, + is_neox=False, + ) + expected = tuple( + torch.cat([txt_reference[index], img_reference[index]], dim=1) + for index in range(3) + ) + + actual = try_fused_qwen_qkv_epilogue( + *img_qkv, + *txt_qkv, + norms[0].weight, + norms[1].weight, + norms[2].weight, + norms[3].weight, + img_cache, + txt_cache, + 1e-6, + 1e-6, + ) + assert actual is not None + assert all( + torch.equal(result, reference) for result, reference in zip(actual, expected) + ) + + # ModelOpt FP8 produces one packed QKV GEMM output. Its chunked Q/K/V + # tensors are zero-copy views with a 3x token stride; the epilogue must + # consume those views directly rather than launching six contiguous copies. + img_packed = torch.cat([tensor.flatten(2) for tensor in img_qkv], dim=-1) + txt_packed = torch.cat([tensor.flatten(2) for tensor in txt_qkv], dim=-1) + img_views = [ + tensor.unflatten(-1, (heads, head_dim)) + for tensor in img_packed.chunk(3, dim=-1) + ] + txt_views = [ + tensor.unflatten(-1, (heads, head_dim)) + for tensor in txt_packed.chunk(3, dim=-1) + ] + assert all(not tensor.is_contiguous() for tensor in (*img_views, *txt_views)) + + packed_actual = try_fused_qwen_qkv_epilogue( + *img_views, + *txt_views, + norms[0].weight, + norms[1].weight, + norms[2].weight, + norms[3].weight, + img_cache, + txt_cache, + 1e-6, + 1e-6, + ) + assert packed_actual is not None + assert all( + torch.equal(result, reference) + for result, reference in zip(packed_actual, expected) + ) + + +def test_qwen_qkv_epilogue_rejects_compile(): + tensor = torch.empty(1, 1, 1, 128, device="cuda", dtype=torch.bfloat16) + row = torch.empty(128, device="cuda", dtype=torch.bfloat16) + cache = torch.empty(1, 128, device="cuda", dtype=torch.float32) + with patch("torch.compiler.is_compiling", return_value=True): + assert ( + try_fused_qwen_qkv_epilogue( + tensor, + tensor, + tensor, + tensor, + tensor, + tensor, + row, + row, + row, + row, + cache, + cache, + 1e-6, + 1e-6, + ) + is None + ) + + +def test_qwen_qkv_epilogue_rejects_unsupported_head_dim(): + tensor = torch.empty(1, 1, 1, 64, device="cuda", dtype=torch.bfloat16) + row = torch.empty(64, device="cuda", dtype=torch.bfloat16) + cache = torch.empty(1, 64, device="cuda", dtype=torch.float32) + assert ( + try_fused_qwen_qkv_epilogue( + tensor, + tensor, + tensor, + tensor, + tensor, + tensor, + row, + row, + row, + row, + cache, + cache, + 1e-6, + 1e-6, + ) + is None + ) + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v"]))