diff --git a/cmake/onnxruntime_unittests.cmake b/cmake/onnxruntime_unittests.cmake index dfede7b813e6e..dbd917136f161 100644 --- a/cmake/onnxruntime_unittests.cmake +++ b/cmake/onnxruntime_unittests.cmake @@ -792,7 +792,9 @@ if(onnxruntime_USE_JSEP) endif() if(onnxruntime_USE_WEBGPU AND NOT onnxruntime_USE_EP_API_ADAPTERS) - list(APPEND onnxruntime_test_framework_src_patterns ${TEST_SRC_DIR}/providers/webgpu/*) + list(APPEND onnxruntime_test_framework_src_patterns + ${TEST_SRC_DIR}/providers/webgpu/* + ${TEST_SRC_DIR}/providers/webgpu/math/*) list(APPEND onnxruntime_test_providers_dependencies onnxruntime_providers_webgpu) list(APPEND onnxruntime_test_providers_libs onnxruntime_providers_webgpu) endif() diff --git a/docs/ContribOperators.md b/docs/ContribOperators.md index ff9e5e4024e3b..f8eef96e1b995 100644 --- a/docs/ContribOperators.md +++ b/docs/ContribOperators.md @@ -219,7 +219,6 @@ This version of the operator has been available since version 1 of the 'com.micr
Constrain mask index to integer types
- ### **com.microsoft.AttnLSTM** Computes an one-layer RNN where its RNN Cell is an AttentionWrapper wrapped a LSTM Cell. The RNN layer @@ -5661,8 +5660,13 @@ This version of the operator has been available since version 1 of the 'com.micr If block_size is provided, both hidden_size and inter_size must be divisible by the block size, and the dequantization is performed per block of size block_size along the K (input feature) dimension. - If block_size and zero_point are provided, both hidden_size and inter_size must be divisible by block_size * pack_size, - where pack_size = 8 / expert_weight_bits. + Packed byte dimensions are computed as logical_element_count * effective_expert_weight_bits / 8. + Weight rows must be byte-aligned. Zero-point rows are padded to a whole byte when necessary. + + fc1_expert_weight_bits, fc2_expert_weight_bits, and fc3_expert_weight_bits optionally override + expert_weight_bits for the corresponding projection. An omitted override inherits expert_weight_bits. + When SwiGLU is fused, FC3 is stored in FC1 and fc3_expert_weight_bits must be omitted or equal to + fc1_expert_weight_bits after inheritance. The SwiGLU (Swish-Gated Linear Unit) activation function is like: g = xW + b @@ -5695,6 +5699,12 @@ This version of the operator has been available since version 1 of the 'com.micr
Size of each quantization block along the K (input feature) dimension. Must be power of two and ≥ 16 (e.g., 16, 32, 64, 128). Both hidden_size and inter_size must be divisible by the block size. The FP4 modes always use blocking: MXFP4 ('fp4'/'wfp4afp8') is normalized to block_size 32 and NVFP4 ('nvfp4') to block_size 16, even when block_size is omitted. For integer quantization ('int'), omitting block_size means there is no blocking and a whole column shares one scaling factor.
expert_weight_bits : int
Number of bits used in quantized weights. Supported values are 2, 4, and 8. Default is 4 bits
+
fc1_expert_weight_bits : int
+
Optional FC1 override for expert_weight_bits. Inherits expert_weight_bits when omitted.
+
fc2_expert_weight_bits : int
+
Optional FC2 override for expert_weight_bits. Inherits expert_weight_bits when omitted.
+
fc3_expert_weight_bits : int
+
Optional FC3 override for expert_weight_bits. Inherits expert_weight_bits when omitted. For fused SwiGLU, the effective FC3 width must equal the effective FC1 width.
k : int
Number of top experts to select from expert pool
normalize_routing_weights : int
@@ -5719,29 +5729,29 @@ This version of the operator has been available since version 1 of the 'com.micr
router_probs : T
2D tensor with shape (num_tokens, num_experts)
fc1_experts_weights : T1
-
3D tensor with shape (num_experts, fusion_size * inter_size, hidden_size / pack_size), The fusion_size is 2 for fused swiglu, or 1 otherwise. The pack_size is 8 / expert_weight_bits.
+
3D tensor with shape (num_experts, fusion_size * inter_size, hidden_size * effective_fc1_bits / 8). The last dimension must be byte-aligned. The fusion_size is 2 for fused swiglu, or 1 otherwise. effective_fc1_bits is fc1_expert_weight_bits when provided, otherwise expert_weight_bits.
fc1_scales (optional) : T2
Optional weight scales. For quant_type='int', this is a 2D tensor with shape (num_experts, fusion_size * inter_size), or a 3D tensor with shape (num_experts, fusion_size * inter_size, hidden_size / block_size) when block_size is provided. For quant_type='fp4' or 'wfp4afp8', this is a float8e8m0 MXFP block-scale tensor with shape (num_experts, fusion_size * inter_size, hidden_size / 32). For quant_type='nvfp4', this is a float8e4m3fn NVFP4 block-scale tensor with shape (num_experts, fusion_size * inter_size, hidden_size / 16). Not used for quant_type='fp8'.
fc1_experts_bias (optional) : T
2D optional tensor with shape (num_experts, fusion_size * inter_size)
fc2_experts_weights : T1
-
3D tensor with shape (num_experts, hidden_size, inter_size / pack_size)
+
3D tensor with shape (num_experts, hidden_size, inter_size * effective_fc2_bits / 8). The last dimension must be byte-aligned. effective_fc2_bits is fc2_expert_weight_bits when provided, otherwise expert_weight_bits.
fc2_scales (optional) : T2
Optional weight scales. For quant_type='int', this is a 2D tensor with shape (num_experts, hidden_size), or a 3D tensor with shape (num_experts, hidden_size, inter_size / block_size) when block_size is provided. For quant_type='fp4' or 'wfp4afp8', this is a float8e8m0 MXFP block-scale tensor with shape (num_experts, hidden_size, inter_size / 32). For quant_type='nvfp4', this is a float8e4m3fn NVFP4 block-scale tensor with shape (num_experts, hidden_size, inter_size / 16). Not used for quant_type='fp8'.
fc2_experts_bias (optional) : T
2D optional tensor with shape (num_experts, hidden_size)
fc3_experts_weights (optional) : T1
-
3D optional tensor with shape (num_experts, inter_size, hidden_size / pack_size)
+
3D optional tensor with shape (num_experts, inter_size, hidden_size * effective_fc3_bits / 8). The last dimension must be byte-aligned. effective_fc3_bits is fc3_expert_weight_bits when provided, otherwise expert_weight_bits.
fc3_scales (optional) : T2
Optional weight scales. For quant_type='int', this is a 2D tensor with shape (num_experts, inter_size), or a 3D tensor with shape (num_experts, inter_size, hidden_size / block_size) when block_size is provided. For quant_type='fp4' or 'wfp4afp8', this is a float8e8m0 MXFP block-scale tensor with shape (num_experts, inter_size, hidden_size / 32). Not used for quant_type='fp8'.
fc3_experts_bias (optional) : T
2D optional tensor with shape (num_experts, inter_size)
fc1_zero_points (optional) : T1
-
2D tensor with shape (num_experts, fusion_size * inter_size / pack_size), or 3D tensor with shape (num_experts, fusion_size * inter_size, hidden_size / block_size / pack_size) when block_size is provided.
+
2D tensor with shape (num_experts, ceil(fusion_size * inter_size * effective_fc1_bits / 8)), or 3D tensor with shape (num_experts, fusion_size * inter_size, ceil((hidden_size / block_size) * effective_fc1_bits / 8)) when block_size is provided.
fc2_zero_points (optional) : T1
-
2D tensor with shape (num_experts, hidden_size / pack_size), or 3D tensor with shape (num_experts, hidden_size, inter_size / block_size / pack_size) when block_size is provided.
+
2D tensor with shape (num_experts, ceil(hidden_size * effective_fc2_bits / 8)), or 3D tensor with shape (num_experts, hidden_size, ceil((inter_size / block_size) * effective_fc2_bits / 8)) when block_size is provided.
fc3_zero_points (optional) : T1
-
2D optional tensor with shape (num_experts, inter_size / pack_size), or 3D optional tensor with shape (num_experts, inter_size, hidden_size / block_size / pack_size) when block_size is provided.
+
2D optional tensor with shape (num_experts, ceil(inter_size * effective_fc3_bits / 8)), or 3D optional tensor with shape (num_experts, inter_size, ceil((hidden_size / block_size) * effective_fc3_bits / 8)) when block_size is provided.
router_weights (optional) : T
2D optional tensor with shape (num_tokens, num_experts). When provided, router_probs is used only for Top-K expert selection, and router_weights is used for aggregating expert outputs (the values at the selected expert indices are gathered and used as mixing weights). This enables DeepSeek-style noaux_tc routing where different tensors are used for selection and aggregation. When not provided, router_probs is used for both selection and aggregation (backward compatible).
fc1_global_scale (optional) : T4
@@ -7749,3 +7759,4 @@ No versioning maintained for experimental ops.
Constrain input and output types to float32 tensors.
+ diff --git a/docs/contrib_ops/cuda/paged_attention.md b/docs/contrib_ops/cuda/paged_attention.md index bad080525ef0e..2b0fa4111cdac 100644 --- a/docs/contrib_ops/cuda/paged_attention.md +++ b/docs/contrib_ops/cuda/paged_attention.md @@ -235,6 +235,7 @@ ops without translation. | `kv_num_heads` | INT | required | existing | | `scale` | FLOAT | `1/sqrt(head_size)` | existing — mandatory in `LATENT` (§12.6) | | `softcap` | FLOAT | `0.0` | existing | +| `is_causal` | INT | `1` | `0` removes the right-hand causal bound on all CUDA backends | | `local_window_size` | INT | `-1` | existing — §9 | | `do_rotary` | INT | `0` | existing | | `rotary_interleaved` | INT | `0` | existing | @@ -1630,6 +1631,16 @@ These block the feature work and should land ahead of it. > falls back to the memory-efficient backend, which gathers pages into a dense buffer first and > therefore accepts any block size. The op only errors when neither backend is eligible. > Lifting this properly requires teaching the Flash paged loader to split a tile across pages. + > + > **Non-causal attention.** `is_causal=0` works with all CUDA backends: FlashAttention, + > memory-efficient attention, paged decode (including XQA), and latent attention. In particular, + > a native cache with `head_size=128` and 16-, 32-, or 64-token pages uses MEA for prefill/multi-token + > drafting and paged decode for decode-shaped batches, without requiring Flash-compatible pages. + > Each query can attend through the sequence's full live KV length (`past_seqlens + query_length`). + > A positive `local_window_size` still bounds the left side at `query_position - window_size + 1`; + > the right side remains unbounded. XQA's speculative mask admits every live draft + > token when non-causal; its single-token kernel needs no different mask. Other backend eligibility + > constraints, including XQA's page alignment and MEA's lack of attention-sink support, are unchanged. 2. **Out-of-bounds binary search.** The binary search over `cumulative_seqlens_q` in `ReshapeAndCache` and `GatherAndExpandPagedKVCache` can yield `batch_id == batch_size` when `token_id >= cumulative_seqlens_q[batch_size]`, producing OOB reads of `past_seqlens` and diff --git a/onnxruntime/contrib_ops/cpu/moe/moe_helper.h b/onnxruntime/contrib_ops/cpu/moe/moe_helper.h index d769ef94bca17..7ee07439cfa24 100644 --- a/onnxruntime/contrib_ops/cpu/moe/moe_helper.h +++ b/onnxruntime/contrib_ops/cpu/moe/moe_helper.h @@ -36,6 +36,30 @@ struct MoEParameters { }; namespace moe_helper { +struct MoEWeightBits { + int64_t fc1; + int64_t fc2; + int64_t fc3; +}; + +inline int64_t PackedByteCount(int64_t logical_element_count, int64_t bits) { + const int64_t bit_count = SafeInt(logical_element_count) * bits; + ORT_ENFORCE(bit_count % 8 == 0, "Packed weights must be byte-aligned."); + return bit_count / 8; +} + +inline int64_t PackedByteCountWithPadding(int64_t logical_element_count, int64_t bits) { + const int64_t bit_count = SafeInt(logical_element_count) * bits; + return (SafeInt(bit_count) + 7) / 8; +} + +inline MoEWeightBits WeightBitsFromPackSize(int64_t pack_size) { + ORT_ENFORCE(pack_size > 0 && 8 % pack_size == 0, + "pack_size must be a positive divisor of 8, got ", pack_size, "."); + const int64_t bits = 8 / pack_size; + return MoEWeightBits{bits, bits, bits}; +} + // Helper to check shape dimensions #define ASSERT_SHAPE_DIMENSION(shape_ptr, dim, name) \ if (shape_ptr != nullptr) { \ @@ -74,10 +98,14 @@ Status CheckInputs(MoEParameters& parameters, const Tensor* fc3_experts_bias, // optional const Tensor* fc3_experts_scales, // required for qMoE; NULL for MOE const Tensor* fc3_zero_points, // optional, for qMoE - const int64_t pack_size, // number of weights packed together (like 2 for uint4 packed to uint8) + const MoEWeightBits& weight_bits, const bool is_fused_swiglu, const int64_t block_size = 0) { // block size for block-wise quantization - ORT_RETURN_IF(pack_size <= 0, "pack_size must be positive, got ", pack_size); + ORT_RETURN_IF(weight_bits.fc1 <= 0 || weight_bits.fc1 > 8 || + weight_bits.fc2 <= 0 || weight_bits.fc2 > 8 || + weight_bits.fc3 <= 0 || weight_bits.fc3 > 8, + "FC weight bits must be between 1 and 8, got FC1=", weight_bits.fc1, + ", FC2=", weight_bits.fc2, ", FC3=", weight_bits.fc3, "."); // Required inputs if (input == nullptr) { @@ -107,20 +135,23 @@ Status CheckInputs(MoEParameters& parameters, int64_t hidden_size = input_dims[input_dims.size() - 1]; int64_t num_experts = router_probs_dims[1]; - ORT_RETURN_IF(hidden_size % pack_size != 0, - "hidden_size (", hidden_size, ") must be divisible by pack_size (", pack_size, ")."); + const bool has_uniform_weight_bits = weight_bits.fc1 == weight_bits.fc2 && weight_bits.fc2 == weight_bits.fc3; + const bool has_legacy_pack_size = has_uniform_weight_bits && 8 % weight_bits.fc1 == 0; + const int64_t legacy_pack_size = has_legacy_pack_size ? 8 / weight_bits.fc1 : 0; + ORT_RETURN_IF(has_legacy_pack_size && hidden_size % legacy_pack_size != 0, + "hidden_size (", hidden_size, ") must be divisible by pack_size (", legacy_pack_size, ")."); int64_t local_num_experts = fc1_experts_weights_shape->GetDims()[0]; const int64_t inter_size_numerator = SafeInt(fc2_experts_weights_shape->GetDims()[1]) * - fc2_experts_weights_shape->GetDims()[2] * pack_size; - ORT_RETURN_IF(inter_size_numerator % hidden_size != 0, + fc2_experts_weights_shape->GetDims()[2] * 8; + const int64_t inter_size_denominator = SafeInt(hidden_size) * weight_bits.fc2; + ORT_RETURN_IF(inter_size_numerator % inter_size_denominator != 0, "Unable to infer inter_size from fc2_experts_weights shape ", *fc2_experts_weights_shape, " and hidden_size ", hidden_size, "."); - int64_t inter_size = inter_size_numerator / hidden_size; - ORT_RETURN_IF(inter_size % pack_size != 0, - "inter_size (", inter_size, ") must be divisible by pack_size (", pack_size, ")."); - + int64_t inter_size = inter_size_numerator / inter_size_denominator; + ORT_RETURN_IF(has_legacy_pack_size && inter_size % legacy_pack_size != 0, + "inter_size (", inter_size, ") must be divisible by pack_size (", legacy_pack_size, ")."); bool legacy_shape = false; const auto& fc2_experts_weights_dims = fc2_experts_weights_shape->GetDims(); const auto& fc1_experts_weights_dims = fc1_experts_weights_shape->GetDims(); @@ -129,17 +160,30 @@ Status CheckInputs(MoEParameters& parameters, // Fused swiglu doubles the output dimension of FC1 since it fused two GEMMs into one. const int64_t fc1_inter_size = is_fused_swiglu ? (inter_size + inter_size) : inter_size; - const int64_t zp_pack_size = pack_size; // Zero points packing (1 for 8-bit, 2 for 4-bit) if (legacy_shape) { // legacy shape does not match column major memory layout. This is for backward compatibility. - CHECK_SHAPE(fc1_experts_weights_shape, "fc1_experts_weights", num_experts, hidden_size, fc1_inter_size / pack_size); - CHECK_SHAPE(fc2_experts_weights_shape, "fc2_experts_weights", num_experts, inter_size, hidden_size / pack_size); - CHECK_SHAPE(fc3_experts_weights_shape, "fc3_experts_weights", num_experts, hidden_size, inter_size / pack_size); + ORT_RETURN_IF((SafeInt(fc1_inter_size) * weight_bits.fc1) % 8 != 0 || + (SafeInt(hidden_size) * weight_bits.fc2) % 8 != 0 || + (SafeInt(inter_size) * weight_bits.fc3) % 8 != 0, + "Expert dimensions and FC weight bits must produce byte-aligned weights for the legacy layout."); + CHECK_SHAPE(fc1_experts_weights_shape, "fc1_experts_weights", num_experts, hidden_size, + PackedByteCount(fc1_inter_size, weight_bits.fc1)); + CHECK_SHAPE(fc2_experts_weights_shape, "fc2_experts_weights", num_experts, inter_size, + PackedByteCount(hidden_size, weight_bits.fc2)); + CHECK_SHAPE(fc3_experts_weights_shape, "fc3_experts_weights", num_experts, hidden_size, + PackedByteCount(inter_size, weight_bits.fc3)); } else { - CHECK_SHAPE(fc1_experts_weights_shape, "fc1_experts_weights", num_experts, fc1_inter_size, hidden_size / pack_size); - CHECK_SHAPE(fc2_experts_weights_shape, "fc2_experts_weights", num_experts, hidden_size, inter_size / pack_size); - CHECK_SHAPE(fc3_experts_weights_shape, "fc3_experts_weights", num_experts, inter_size, hidden_size / pack_size); + ORT_RETURN_IF((SafeInt(hidden_size) * weight_bits.fc1) % 8 != 0 || + (SafeInt(inter_size) * weight_bits.fc2) % 8 != 0 || + (SafeInt(hidden_size) * weight_bits.fc3) % 8 != 0, + "Expert dimensions and FC weight bits must produce byte-aligned weights."); + CHECK_SHAPE(fc1_experts_weights_shape, "fc1_experts_weights", num_experts, fc1_inter_size, + PackedByteCount(hidden_size, weight_bits.fc1)); + CHECK_SHAPE(fc2_experts_weights_shape, "fc2_experts_weights", num_experts, hidden_size, + PackedByteCount(inter_size, weight_bits.fc2)); + CHECK_SHAPE(fc3_experts_weights_shape, "fc3_experts_weights", num_experts, inter_size, + PackedByteCount(hidden_size, weight_bits.fc3)); } CHECK_TENSOR_SHAPE(router_probs, num_rows, num_experts); @@ -171,9 +215,9 @@ Status CheckInputs(MoEParameters& parameters, CHECK_TENSOR_SHAPE(fc3_experts_scales, num_experts, inter_size, fc3_blocks_per_row); // Validate zero-point tensors (block-wise) - const int64_t fc1_zp_blocks = (fc1_blocks_per_row + zp_pack_size - 1) / zp_pack_size; - const int64_t fc2_zp_blocks = (fc2_blocks_per_row + zp_pack_size - 1) / zp_pack_size; - const int64_t fc3_zp_blocks = (fc3_blocks_per_row + zp_pack_size - 1) / zp_pack_size; + const int64_t fc1_zp_blocks = PackedByteCountWithPadding(fc1_blocks_per_row, weight_bits.fc1); + const int64_t fc2_zp_blocks = PackedByteCountWithPadding(fc2_blocks_per_row, weight_bits.fc2); + const int64_t fc3_zp_blocks = PackedByteCountWithPadding(fc3_blocks_per_row, weight_bits.fc3); CHECK_TENSOR_SHAPE(fc1_zero_points, num_experts, fc1_inter_size, fc1_zp_blocks); CHECK_TENSOR_SHAPE(fc2_zero_points, num_experts, hidden_size, fc2_zp_blocks); @@ -185,10 +229,12 @@ Status CheckInputs(MoEParameters& parameters, const auto& fc1_scales_dims = fc1_experts_scales->Shape().GetDims(); if (fc1_scales_dims.size() == 2) { CHECK_TENSOR_SHAPE(fc1_experts_scales, num_experts, fc1_inter_size); - CHECK_TENSOR_SHAPE(fc1_zero_points, num_experts, (fc1_inter_size + zp_pack_size - 1) / zp_pack_size); + CHECK_TENSOR_SHAPE(fc1_zero_points, num_experts, + PackedByteCountWithPadding(fc1_inter_size, weight_bits.fc1)); } else if (fc1_scales_dims.size() == 3) { CHECK_TENSOR_SHAPE(fc1_experts_scales, num_experts, fc1_inter_size, 1); - CHECK_TENSOR_SHAPE(fc1_zero_points, num_experts, (fc1_inter_size + zp_pack_size - 1) / zp_pack_size); + CHECK_TENSOR_SHAPE(fc1_zero_points, num_experts, + PackedByteCountWithPadding(fc1_inter_size, weight_bits.fc1)); } else { ORT_THROW("fc1_experts_scales must be 2D or 3D tensor"); } @@ -198,10 +244,12 @@ Status CheckInputs(MoEParameters& parameters, const auto& fc2_scales_dims = fc2_experts_scales->Shape().GetDims(); if (fc2_scales_dims.size() == 2) { CHECK_TENSOR_SHAPE(fc2_experts_scales, num_experts, hidden_size); - CHECK_TENSOR_SHAPE(fc2_zero_points, num_experts, (hidden_size + zp_pack_size - 1) / zp_pack_size); + CHECK_TENSOR_SHAPE(fc2_zero_points, num_experts, + PackedByteCountWithPadding(hidden_size, weight_bits.fc2)); } else if (fc2_scales_dims.size() == 3) { CHECK_TENSOR_SHAPE(fc2_experts_scales, num_experts, hidden_size, 1); - CHECK_TENSOR_SHAPE(fc2_zero_points, num_experts, (hidden_size + zp_pack_size - 1) / zp_pack_size); + CHECK_TENSOR_SHAPE(fc2_zero_points, num_experts, + PackedByteCountWithPadding(hidden_size, weight_bits.fc2)); } else { ORT_THROW("fc2_experts_scales must be 2D or 3D tensor"); } @@ -211,10 +259,12 @@ Status CheckInputs(MoEParameters& parameters, const auto& fc3_scales_dims = fc3_experts_scales->Shape().GetDims(); if (fc3_scales_dims.size() == 2) { CHECK_TENSOR_SHAPE(fc3_experts_scales, num_experts, inter_size); - CHECK_TENSOR_SHAPE(fc3_zero_points, num_experts, (inter_size + zp_pack_size - 1) / zp_pack_size); + CHECK_TENSOR_SHAPE(fc3_zero_points, num_experts, + PackedByteCountWithPadding(inter_size, weight_bits.fc3)); } else if (fc3_scales_dims.size() == 3) { CHECK_TENSOR_SHAPE(fc3_experts_scales, num_experts, inter_size, 1); - CHECK_TENSOR_SHAPE(fc3_zero_points, num_experts, (inter_size + zp_pack_size - 1) / zp_pack_size); + CHECK_TENSOR_SHAPE(fc3_zero_points, num_experts, + PackedByteCountWithPadding(inter_size, weight_bits.fc3)); } else { ORT_THROW("fc3_experts_scales must be 2D or 3D tensor"); } @@ -255,6 +305,23 @@ Status CheckInputs(MoEParameters& parameters, return Status::OK(); } +template +Status CheckInputs(MoEParameters& parameters, + const Tensor* input, const Tensor* router_probs, + const TensorShape* fc1_experts_weights_shape, const Tensor* fc1_experts_bias, + const Tensor* fc1_experts_scales, const Tensor* fc1_zero_points, + const TensorShape* fc2_experts_weights_shape, const Tensor* fc2_experts_bias, + const Tensor* fc2_experts_scales, const Tensor* fc2_zero_points, + const TensorShape* fc3_experts_weights_shape, const Tensor* fc3_experts_bias, + const Tensor* fc3_experts_scales, const Tensor* fc3_zero_points, + const int64_t pack_size, const bool is_fused_swiglu, const int64_t block_size = 0) { + return CheckInputs(parameters, input, router_probs, + fc1_experts_weights_shape, fc1_experts_bias, fc1_experts_scales, fc1_zero_points, + fc2_experts_weights_shape, fc2_experts_bias, fc2_experts_scales, fc2_zero_points, + fc3_experts_weights_shape, fc3_experts_bias, fc3_experts_scales, fc3_zero_points, + WeightBitsFromPackSize(pack_size), is_fused_swiglu, block_size); +} + template Status CheckInputs(MoEParameters& parameters, const Tensor* input, // required @@ -271,7 +338,7 @@ Status CheckInputs(MoEParameters& parameters, const Tensor* fc3_experts_bias, // optional const Tensor* fc3_experts_scales, // required for qMoE; NULL for MOE const Tensor* fc3_zero_points, // optional, for qMoE - const int64_t pack_size, // number of weights packed together (like 2 for uint4 packed to uint8) + const MoEWeightBits& weight_bits, const bool is_fused_swiglu, const int64_t block_size = 0) { // block size for block-wise quantization @@ -282,7 +349,24 @@ Status CheckInputs(MoEParameters& parameters, return CheckInputs(parameters, input, router_probs, fc1_shape, fc1_experts_bias, fc1_experts_scales, fc1_zero_points, fc2_shape, fc2_experts_bias, fc2_experts_scales, fc2_zero_points, fc3_shape, fc3_experts_bias, fc3_experts_scales, fc3_zero_points, - pack_size, is_fused_swiglu, block_size); + weight_bits, is_fused_swiglu, block_size); +} + +template +Status CheckInputs(MoEParameters& parameters, + const Tensor* input, const Tensor* router_probs, + const Tensor* fc1_experts_weights, const Tensor* fc1_experts_bias, + const Tensor* fc1_experts_scales, const Tensor* fc1_zero_points, + const Tensor* fc2_experts_weights, const Tensor* fc2_experts_bias, + const Tensor* fc2_experts_scales, const Tensor* fc2_zero_points, + const Tensor* fc3_experts_weights, const Tensor* fc3_experts_bias, + const Tensor* fc3_experts_scales, const Tensor* fc3_zero_points, + const int64_t pack_size, const bool is_fused_swiglu, const int64_t block_size = 0) { + return CheckInputs(parameters, input, router_probs, + fc1_experts_weights, fc1_experts_bias, fc1_experts_scales, fc1_zero_points, + fc2_experts_weights, fc2_experts_bias, fc2_experts_scales, fc2_zero_points, + fc3_experts_weights, fc3_experts_bias, fc3_experts_scales, fc3_zero_points, + WeightBitsFromPackSize(pack_size), is_fused_swiglu, block_size); } } // namespace moe_helper diff --git a/onnxruntime/contrib_ops/cpu/moe/moe_quantization_cpu.cc b/onnxruntime/contrib_ops/cpu/moe/moe_quantization_cpu.cc index 4d1002a0c5bee..3a57d2e151ab7 100644 --- a/onnxruntime/contrib_ops/cpu/moe/moe_quantization_cpu.cc +++ b/onnxruntime/contrib_ops/cpu/moe/moe_quantization_cpu.cc @@ -633,6 +633,12 @@ Status QMoECPU::PrePack(const Tensor& tensor, int input_idx, AllocatorPtr all /*out*/ PrePackedWeights* prepacked_weights) { is_packed = false; + if (fc1_expert_weight_bits_ != expert_weight_bits_ || + fc2_expert_weight_bits_ != expert_weight_bits_ || + fc3_expert_weight_bits_ != expert_weight_bits_) { + return Status::OK(); + } + // If scales are prepacked, they are constant initializers. if (input_idx == 3) { return Status::OK(); @@ -924,6 +930,15 @@ QMoECPU::QMoECPU(const OpKernelInfo& op_kernel_info) ORT_ENFORCE(op_kernel_info.GetAttr("expert_weight_bits", &expert_weight_bits_).IsOK()); ORT_ENFORCE(expert_weight_bits_ == 2 || expert_weight_bits_ == 4 || expert_weight_bits_ == 8, "Attribute 'expert_weight_bits' must be 2, 4, or 8."); + fc1_expert_weight_bits_ = op_kernel_info.GetAttrOrDefault("fc1_expert_weight_bits", expert_weight_bits_); + fc2_expert_weight_bits_ = op_kernel_info.GetAttrOrDefault("fc2_expert_weight_bits", expert_weight_bits_); + fc3_expert_weight_bits_ = op_kernel_info.GetAttrOrDefault("fc3_expert_weight_bits", expert_weight_bits_); + ORT_ENFORCE((fc1_expert_weight_bits_ == 2 || fc1_expert_weight_bits_ == 4 || fc1_expert_weight_bits_ == 8) && + (fc2_expert_weight_bits_ == 2 || fc2_expert_weight_bits_ == 4 || fc2_expert_weight_bits_ == 8) && + (fc3_expert_weight_bits_ == 2 || fc3_expert_weight_bits_ == 4 || fc3_expert_weight_bits_ == 8), + "FC-specific expert weight bits must be 2, 4, or 8."); + ORT_ENFORCE(swiglu_fusion_ == 0 || fc3_expert_weight_bits_ == fc1_expert_weight_bits_, + "Fused SwiGLU requires FC1 and FC3 expert weight bits to match."); block_size_ = op_kernel_info.GetAttrOrDefault("block_size", 0); ORT_ENFORCE(block_size_ >= 0); @@ -1211,10 +1226,19 @@ Status QMoECPU::Compute(OpKernelContext* context) const { fc1_shape_ptr, inputs.fc1_experts_bias, inputs.fc1_scales, inputs.fc1_zero_points, fc2_shape_ptr, inputs.fc2_experts_bias, inputs.fc2_scales, inputs.fc2_zero_points, fc3_shape_ptr, inputs.fc3_experts_bias, inputs.fc3_scales, inputs.fc3_zero_points, - 8 / expert_weight_bits_, + moe_helper::MoEWeightBits{fc1_expert_weight_bits_, + fc2_expert_weight_bits_, + fc3_expert_weight_bits_}, activation_type_ == ActivationType::SwiGLU, block_size_)); + if (fc1_expert_weight_bits_ != expert_weight_bits_ || + fc2_expert_weight_bits_ != expert_weight_bits_ || + fc3_expert_weight_bits_ != expert_weight_bits_) { + return ORT_MAKE_STATUS(ONNXRUNTIME, NOT_IMPLEMENTED, + "Mixed-width QMoE execution is not yet implemented on CPU."); + } + if (fc3_shape_ptr || inputs.fc3_experts_bias || inputs.fc3_scales || inputs.fc3_zero_points) { return ORT_MAKE_STATUS(ONNXRUNTIME, NOT_IMPLEMENTED, "FC3 gating is not yet implemented on CPU for QMoE"); } diff --git a/onnxruntime/contrib_ops/cpu/moe/moe_quantization_cpu.h b/onnxruntime/contrib_ops/cpu/moe/moe_quantization_cpu.h index 69f1e8fa2dbfe..5dfac3e04f39f 100644 --- a/onnxruntime/contrib_ops/cpu/moe/moe_quantization_cpu.h +++ b/onnxruntime/contrib_ops/cpu/moe/moe_quantization_cpu.h @@ -86,6 +86,9 @@ class QMoECPU final : public OpKernel, public MoEBaseCPU { /*out*/ bool& is_packed, /*out*/ PrePackedWeights* prepacked_weights); int64_t expert_weight_bits_; + int64_t fc1_expert_weight_bits_; + int64_t fc2_expert_weight_bits_; + int64_t fc3_expert_weight_bits_; int64_t block_size_; bool use_mlas_q4_gemm_{false}; bool use_mlas_q4_gemm_overridden_{false}; diff --git a/onnxruntime/contrib_ops/cuda/bert/cutlass_fmha/fmha_launch_template.h b/onnxruntime/contrib_ops/cuda/bert/cutlass_fmha/fmha_launch_template.h index 6eafc9dc7ccf4..52f9222d7c21e 100644 --- a/onnxruntime/contrib_ops/cuda/bert/cutlass_fmha/fmha_launch_template.h +++ b/onnxruntime/contrib_ops/cuda/bert/cutlass_fmha/fmha_launch_template.h @@ -74,7 +74,9 @@ struct RightPaddingBatchHook { batch_id * lse_dim * p.num_heads + head_id * lse_dim + query_start; } - if (p.custom_mask_type == AttentionKernel::CausalFromBottomRight) { + if (p.custom_mask_type == AttentionKernel::CausalFromBottomRight || + (p.custom_mask_type == AttentionKernel::NoCustomMask && p.window_size > 0)) { + // Keep a non-causal left window anchored to the bottom-right query position. // May be negative when num_keys < num_queries (nonpad external KV cache, onnx#8068 / ORT #28904). // causal_diagonal_offset is int32_t so the negative value is preserved (no unsigned wrap). p.causal_diagonal_offset = p.num_keys - p.num_queries; @@ -97,7 +99,8 @@ struct RightPaddingBatchHook { // 15/16th of tensor core compute In that case : // - we only launch kernels for head_id % kQueriesPerBlock == 0 // - we iterate over heads instead of queries (strideM = strideH) - if (p.num_queries == 1 && p.k_strideH == 0 && p.v_strideH == 0) { + // A local window must not treat these head rows as different query positions. + if (p.num_queries == 1 && p.k_strideH == 0 && p.v_strideH == 0 && p.window_size <= 0) { if (head_id % kQueriesPerBlock != 0) return false; p.q_strideM = p.q_strideH; diff --git a/onnxruntime/contrib_ops/cuda/bert/cutlass_fmha/kernel_forward.h b/onnxruntime/contrib_ops/cuda/bert/cutlass_fmha/kernel_forward.h index 5a02ef65933c0..fc85570203739 100644 --- a/onnxruntime/contrib_ops/cuda/bert/cutlass_fmha/kernel_forward.h +++ b/onnxruntime/contrib_ops/cuda/bert/cutlass_fmha/kernel_forward.h @@ -348,7 +348,9 @@ struct AttentionKernel { } // Custom masking - if (custom_mask_type == CausalFromBottomRight) { + if (custom_mask_type == CausalFromBottomRight || + (custom_mask_type == NoCustomMask && window_size > 0)) { + // A non-causal left window uses the same bottom-right query positions, without an upper mask. // May be negative when num_keys < num_queries (nonpad external KV cache, onnx#8068 / ORT #28904). // causal_diagonal_offset is int32_t so the negative value is preserved (no unsigned wrap). causal_diagonal_offset = num_keys - num_queries; @@ -374,7 +376,8 @@ struct AttentionKernel { // 15/16th of tensor core compute In that case : // - we only launch kernels for head_id % kQueriesPerBlock == 0 // - we iterate over heads instead of queries (strideM = strideH) - if (num_queries == 1 && k_strideH == 0 && v_strideH == 0) { + // A local window must not treat these head rows as different query positions. + if (num_queries == 1 && k_strideH == 0 && v_strideH == 0 && window_size <= 0) { if (head_id % kQueriesPerBlock != 0) return false; q_strideM = q_strideH; @@ -684,12 +687,6 @@ struct AttentionKernel { XFORMERS_CHECK( p.custom_mask_type < NumCustomMaskTypes, "invalid value for `custom_mask_type`"); - if (p.window_size > 0) { - XFORMERS_CHECK( - p.custom_mask_type == CausalFromTopLeft || - p.custom_mask_type == CausalFromBottomRight, - "invalid value for custom_mask_type"); - } return true; } diff --git a/onnxruntime/contrib_ops/cuda/bert/paged_attention.cc b/onnxruntime/contrib_ops/cuda/bert/paged_attention.cc index 9a9cca3a91770..b65fc98c3b3d7 100644 --- a/onnxruntime/contrib_ops/cuda/bert/paged_attention.cc +++ b/onnxruntime/contrib_ops/cuda/bert/paged_attention.cc @@ -496,7 +496,7 @@ Status PagedAttention::ComputeInternal(OpKernelContext* context) cons is_supported_quant_type(k_quant_type_) && is_supported_quant_type(v_quant_type_) && (!is_fp8_cache || device_prop.major >= 9 || (device_prop.major == 8 && device_prop.minor == 9))); // Speculative verification steps (2..8 new tokens per sequence) run on the paged XQA kernel with - // a packed lower-triangular mask built by PagedXqaSpecDecCausalMaskKernel. The gate is the + // a packed mask built by PagedXqaSpecDecMaskKernel. The gate is the // metadata query bound, not the aggregate token count: a zero-heavy ragged step can have // token_count <= batch_size while still carrying a multi-token sequence. Local windows and // attention sinks stay eligible: the kernel's rows are flattened (query token, query head) pairs, @@ -509,7 +509,6 @@ Status PagedAttention::ComputeInternal(OpKernelContext* context) cons const bool portable_spec_dec_candidate = has_metadata_bounds && max_query_len_bound > 1 && max_query_len_bound <= 8 && (std::is_same_v || (kIsQuantizedCache && per_channel_k && !enable_per_channel_xqa_)); - // cuDNN paged SDPA (decode-only, unquantized cache). Preferred over FlashAttention when eligible; // XQA still wins its target case (fp16, group_size 6, head_size 256, native page size). The // eligibility is intentionally metadata-gated so the selection never triggers a new D->H readback @@ -574,23 +573,16 @@ Status PagedAttention::ComputeInternal(OpKernelContext* context) cons } } - // Only the FlashAttention backend takes a causality flag; the paged decode and CUTLASS kernels - // both hard-code a bottom-right causal mask. + // cuDNN paged SDPA is causal-only. FlashAttention, paged decode, and CUTLASS all consume the + // causality flag directly. bool use_paged_decode = !use_cudnn_paged && - decode_eligible && parameters.is_causal && + decode_eligible && ((decode_shaped && (kIsQuantizedCache || fp16_xqa_eligible || !flash_eligible)) || xqa_spec_dec_candidate || portable_spec_dec_candidate); bool use_flash_attention = flash_eligible && !use_paged_decode && !use_cudnn_paged; const bool use_memory_efficient_attention = - mea_eligible && !use_paged_decode && !use_cudnn_paged && parameters.is_causal; - - if (!parameters.is_causal && !use_flash_attention) { - return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, - "PagedAttention: is_causal=0 requires the FlashAttention backend (sm>=80, fp16/bf16, " - "head_size ", - parameters.head_size, ", block_size ", parameters.block_size, ")."); - } + mea_eligible && !use_paged_decode && !use_cudnn_paged; // Both gather-based backends need a dense KV staging buffer when the cache is quantized // (FlashAttention cannot read a quantized page, and the CUTLASS kernel is not paged at all). diff --git a/onnxruntime/contrib_ops/cuda/bert/paged_attention_impl.cu b/onnxruntime/contrib_ops/cuda/bert/paged_attention_impl.cu index edd5d6f7a585c..3ffd1a441a58b 100644 --- a/onnxruntime/contrib_ops/cuda/bert/paged_attention_impl.cu +++ b/onnxruntime/contrib_ops/cuda/bert/paged_attention_impl.cu @@ -828,6 +828,7 @@ __global__ void PagedDecodeSplitKV(const T* __restrict__ query, const float scale, const float softcap, const int local_window_size, + const bool is_causal, const bool k_per_channel) { extern __shared__ float paged_decode_smem[]; const int channel_groups = PagedDecodeChannelGroups(head_size); @@ -869,22 +870,20 @@ __global__ void PagedDecodeSplitKV(const T* __restrict__ query, const int64_t partial_head_index = (static_cast(split_id) * token_count + token_id) * num_heads + head_id; - // Causality is resolved per query token instead of assuming one new token per sequence: - // kv_len - q_len is past_seqlens[batch_id], so the token at offset q_index inside its sequence - // attends to cached positions [0, past + q_index]. For the decode case (q_len == 1) this reduces - // to the whole live context, as before. + // Resolve the query position even for non-causal attention: the left window stays anchored + // there, while the right bound expands to the full live context. const int q_index = token_id - cumulative_seqlens_q[batch_id]; const int q_len = cumulative_seqlens_q[batch_id + 1] - cumulative_seqlens_q[batch_id]; const int seq_kv_len = cumulative_seqlens_kv[batch_id + 1] - cumulative_seqlens_kv[batch_id]; - const int kv_len = seq_kv_len - q_len + q_index + 1; + const int query_end = seq_kv_len - q_len + q_index + 1; + const int kv_len = is_causal ? query_end : seq_kv_len; - // Sliding window matches FlashAttention's window_size_left = local_window_size - 1 convention at - // query position kv_len - 1, i.e. positions in [kv_len - local_window_size, kv_len). + // Sliding window matches FlashAttention's window_size_left = local_window_size - 1. const int tokens_per_split = (kv_len + num_splits - 1) / num_splits; int kv_begin = split_id * tokens_per_split; const int kv_end = min(kv_len, kv_begin + tokens_per_split); if (local_window_size > 0) { - kv_begin = max(kv_begin, kv_len - local_window_size); + kv_begin = max(kv_begin, query_end - local_window_size); } if (kv_begin >= kv_end) { @@ -1154,14 +1153,14 @@ Status LaunchPagedDecodeAttention(const T* query, const TCACHE* key_cache, const const int batch_size, const int num_heads, const int kv_num_heads, const int head_size, const int block_size, const int max_num_blocks_per_seq, const int token_count, const int num_splits, const float scale, - const float softcap, const int local_window_size, + const float softcap, const int local_window_size, const bool is_causal, const bool use_smooth_softmax, cudaStream_t stream) { const size_t smem_bytes = GetPagedDecodeSharedMemoryBytes(head_size); const dim3 grid(num_heads, token_count, num_splits); PagedDecodeSplitKV<<>>( query, key_cache, value_cache, k_scale, cumulative_seqlens_q, cumulative_seqlens_kv, block_table, partial_out, partial_max, partial_sum, batch_size, num_heads, kv_num_heads, head_size, block_size, - max_num_blocks_per_seq, token_count, num_splits, scale, softcap, local_window_size, k_per_channel); + max_num_blocks_per_seq, token_count, num_splits, scale, softcap, local_window_size, is_causal, k_per_channel); CUDA_RETURN_IF_ERROR(cudaGetLastError()); const dim3 reduce_grid(num_heads, token_count); @@ -1221,6 +1220,7 @@ __global__ void PagedLatentAttentionKernel(const T* __restrict__ query, const float scale, const float softcap, const int local_window_size, + const bool is_causal, const bool k_per_channel, const bool v_per_channel) { extern __shared__ float paged_latent_smem[]; @@ -1250,14 +1250,13 @@ __global__ void PagedLatentAttentionKernel(const T* __restrict__ query, const int batch_id = left; const int s = token_id - cumulative_seqlens_q[batch_id]; - // Causality: this token's logical position is past_seqlens[b] + s, and it attends every cached - // position up to and including its own. That is exactly FlashAttention's bottom-right-aligned - // causal convention for seqlen_k = past + seqlen_q. - const int kv_end = past_seqlens[batch_id] + s + 1; + const int query_end = past_seqlens[batch_id] + s + 1; + const int q_len = cumulative_seqlens_q[batch_id + 1] - cumulative_seqlens_q[batch_id]; + const int kv_end = is_causal ? query_end : past_seqlens[batch_id] + q_len; int kv_begin = 0; if (local_window_size > 0) { // local_window_size counts the current token, matching mha_varlen_fwd's window_size_left = W-1. - kv_begin = max(0, kv_end - local_window_size); + kv_begin = max(0, query_end - local_window_size); } const int kv_head_id = head_id / (num_heads / kv_num_heads); @@ -1396,13 +1395,14 @@ Status LaunchPagedLatentAttention(const T* query, const TCACHE* key_cache, const const int batch_size, const int num_heads, const int kv_num_heads, const int head_size, const int v_head_size, const int block_size, const int max_num_blocks_per_seq, const int token_count, const float scale, - const float softcap, const int local_window_size, cudaStream_t stream) { + const float softcap, const int local_window_size, const bool is_causal, + cudaStream_t stream) { const size_t smem_bytes = GetPagedLatentSharedMemoryBytes(head_size, v_head_size); const dim3 grid(token_count, num_heads); PagedLatentAttentionKernel<<>>( query, key_cache, value_cache, k_scale, v_scale, cumulative_seqlens_q, past_seqlens, block_table, output, batch_size, num_heads, kv_num_heads, head_size, v_head_size, block_size, max_num_blocks_per_seq, scale, - softcap, local_window_size, k_per_channel, v_per_channel); + softcap, local_window_size, is_causal, k_per_channel, v_per_channel); return CUDA_CALL(cudaGetLastError()); } @@ -1516,7 +1516,7 @@ Status LatentAttention( data.cumulative_seqlens_q, data.past_seqlens, data.block_table, data.output, parameters.batch_size, parameters.num_heads, parameters.kv_num_heads, parameters.head_size, parameters.v_head_size, parameters.block_size, parameters.max_num_blocks_per_seq, - parameters.token_count, scale, parameters.softcap, parameters.local_window_size, stream))); + parameters.token_count, scale, parameters.softcap, parameters.local_window_size, parameters.is_causal, stream))); DUMP_TENSOR_INIT(); DUMP_TENSOR("latent (MLA) paged attention output", data.output, parameters.token_count, parameters.num_heads, @@ -1544,7 +1544,8 @@ Status PagedDecodeAttention( data.decode_partial_out, data.decode_partial_max, data.decode_partial_sum, parameters.batch_size, parameters.num_heads, parameters.kv_num_heads, parameters.head_size, parameters.block_size, parameters.max_num_blocks_per_seq, parameters.token_count, data.num_splits, - scale, parameters.softcap, parameters.local_window_size, parameters.use_smooth_softmax, stream))); + scale, parameters.softcap, parameters.local_window_size, parameters.is_causal, + parameters.use_smooth_softmax, stream))); DUMP_TENSOR_INIT(); DUMP_TENSOR("paged decode attention output", data.output, parameters.token_count, parameters.num_heads, @@ -1679,13 +1680,14 @@ __global__ void PagedConvertHeadSinkToFloatKernel(float* __restrict__ dst, const } } -// Lower-triangular packed mask for the speculative XQA kernel: row = query token (global, packed -// token-major), bit p of word w = "may attend to draft token w*32+p". -__global__ void PagedXqaSpecDecCausalMaskKernel(uint32_t* __restrict__ mask, - const int* __restrict__ cumulative_seqlens_q, - const int batch_size, - const int token_count, - const int words_per_row) { +// Packed mask for the speculative XQA kernel: row = query token (global, packed token-major), +// bit p of word w = "may attend to draft token w*32+p". +__global__ void PagedXqaSpecDecMaskKernel(uint32_t* __restrict__ mask, + const int* __restrict__ cumulative_seqlens_q, + const int batch_size, + const int token_count, + const int words_per_row, + const bool is_causal) { const int i = blockIdx.x * blockDim.x + threadIdx.x; if (i >= token_count * words_per_row) { return; @@ -1705,7 +1707,8 @@ __global__ void PagedXqaSpecDecCausalMaskKernel(uint32_t* __restrict__ mask, } const int local_row = token_id - cumulative_seqlens_q[lo]; - const int allowed_bits = local_row + 1 - word * 32; + const int q_len = cumulative_seqlens_q[lo + 1] - cumulative_seqlens_q[lo]; + const int allowed_bits = (is_causal ? local_row + 1 : q_len) - word * 32; // Do NOT write this as `clamped == 32`. ptxas (CUDA 13.0.48) folds min/max into VIMNMX.RELU and // then reuses that instruction's clamp predicate for an equality test against the clamp bound // with inverted polarity, so `max(0, min(32, x)) == 32` is true for every x. That silently made @@ -1802,9 +1805,9 @@ Status PagedXqaDecodeAttention( const int words_per_row = (data.max_query_len + 31) / 32; const int mask_words = parameters.token_count * words_per_row; const int blocks = (mask_words + max_threads_per_block - 1) / max_threads_per_block; - PagedXqaSpecDecCausalMaskKernel<<>>( + PagedXqaSpecDecMaskKernel<<>>( data.xqa_spec_dec_mask, data.cumulative_seqlens_q, batch_size, - parameters.token_count, words_per_row); + parameters.token_count, words_per_row, parameters.is_causal); CUDA_RETURN_IF_ERROR(cudaGetLastError()); ORT_RETURN_IF_ERROR(LaunchXQAPagedSpecDecKernel( @@ -2057,7 +2060,7 @@ Status EfficientAttention( p.max_sequence_length = total_kv_tokens; p.qk_head_size = head_size; p.v_head_size = head_size; - p.causal = true; + p.causal = parameters.is_causal; p.scale = scale; p.softcap = parameters.softcap; p.local_window_size = local_window_size; diff --git a/onnxruntime/contrib_ops/cuda/moe/moe_quantization.cc b/onnxruntime/contrib_ops/cuda/moe/moe_quantization.cc index ad87f2ddebd16..810b1893a966c 100644 --- a/onnxruntime/contrib_ops/cuda/moe/moe_quantization.cc +++ b/onnxruntime/contrib_ops/cuda/moe/moe_quantization.cc @@ -276,6 +276,15 @@ QMoE::QMoE(const OpKernelInfo& op_kernel_info) : CudaKernel(op_kernel_info), MoE ORT_ENFORCE(op_kernel_info.GetAttr("expert_weight_bits", &expert_weight_bits_).IsOK()); ORT_ENFORCE(expert_weight_bits_ == 8 || expert_weight_bits_ == 4, "expert_weight_bits must be 4 or 8, but got ", expert_weight_bits_); + fc1_expert_weight_bits_ = op_kernel_info.GetAttrOrDefault("fc1_expert_weight_bits", expert_weight_bits_); + fc2_expert_weight_bits_ = op_kernel_info.GetAttrOrDefault("fc2_expert_weight_bits", expert_weight_bits_); + fc3_expert_weight_bits_ = op_kernel_info.GetAttrOrDefault("fc3_expert_weight_bits", expert_weight_bits_); + ORT_ENFORCE((fc1_expert_weight_bits_ == 2 || fc1_expert_weight_bits_ == 4 || fc1_expert_weight_bits_ == 8) && + (fc2_expert_weight_bits_ == 2 || fc2_expert_weight_bits_ == 4 || fc2_expert_weight_bits_ == 8) && + (fc3_expert_weight_bits_ == 2 || fc3_expert_weight_bits_ == 4 || fc3_expert_weight_bits_ == 8), + "FC-specific expert weight bits must be 2, 4, or 8."); + ORT_ENFORCE(swiglu_fusion_ == 0 || fc3_expert_weight_bits_ == fc1_expert_weight_bits_, + "Fused SwiGLU requires FC1 and FC3 expert weight bits to match."); block_size_ = op_kernel_info.GetAttrOrDefault("block_size", -1); this->quant_type_ = op_kernel_info.GetAttrOrDefault("quant_type", "int"); @@ -629,6 +638,9 @@ Status QMoE::ComputeInternal(OpKernelContext* context) const { const bool is_fp8 = (quant_type_ == "fp8"); const bool is_wfp4afp8 = (quant_type_ == "wfp4afp8"); const bool is_int = (quant_type_ == "int"); + const bool is_mixed_width = fc1_expert_weight_bits_ != expert_weight_bits_ || + fc2_expert_weight_bits_ != expert_weight_bits_ || + fc3_expert_weight_bits_ != expert_weight_bits_; // Modes that consume FP4 weight block scales (inputs 3/6) and per-expert global weight scales. const bool uses_fp4_weight_scales = is_fp4_family || is_wfp4afp8; // Modes that consume per-expert FP-format global weight scales (inputs 15/16). @@ -665,7 +677,7 @@ Status QMoE::ComputeInternal(OpKernelContext* context) const { // them. If PrePack never ran (e.g. ``session.disable_prepacking`` is set), the prepack // buffers stay null and falling through to the raw initializer pointers would feed // non-CUTLASS bytes to the runner, producing silently wrong output. Fail loudly instead. - if (is_int && !weights_prepacked_ && + if (is_int && !is_mixed_width && !weights_prepacked_ && (packed_fc1_weights_ == nullptr || packed_fc2_weights_ == nullptr)) { return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "QMoE weights_prepacked=0 requires PrePack to run, but the int weight " @@ -746,7 +758,6 @@ Status QMoE::ComputeInternal(OpKernelContext* context) const { "Use block_size >= 32 or remove fc*_zero_points."); } - int64_t pack_size = expert_weight_bits_ == 4 ? 2 : 1; bool is_fused_swiglu = activation_type_ == onnxruntime::llm::kernels::cutlass_kernels::ActivationType::Swiglu; MoEParameters moe_params; // Prefer the cached shapes when PrePack consumed the source initializer. @@ -757,7 +768,12 @@ Status QMoE::ComputeInternal(OpKernelContext* context) const { fc1_experts_bias_optional, fc1_scales, fc1_zeros, &fc2_shape, fc2_experts_bias_optional, fc2_scales, fc2_zeros, nullptr, nullptr, nullptr, nullptr, - pack_size, is_fused_swiglu, block_size_)); + moe_helper::MoEWeightBits{fc1_expert_weight_bits_, + fc2_expert_weight_bits_, + fc3_expert_weight_bits_}, + is_fused_swiglu, block_size_)); + ORT_RETURN_IF(is_mixed_width, + "Mixed-width QMoE execution is not yet implemented on CUDA."); ORT_RETURN_IF_NOT(k_ > 0 && k_ <= moe_params.num_experts, "QMoE requires 0 < k <= num_experts, got k=", k_, " and num_experts=", moe_params.num_experts); @@ -2070,6 +2086,12 @@ Status QMoE::PrePack(const Tensor& tensor, int input_idx, AllocatorPtr alloc, ORT_UNUSED_PARAMETER(prepacked_weights); is_packed = false; + if (fc1_expert_weight_bits_ != expert_weight_bits_ || + fc2_expert_weight_bits_ != expert_weight_bits_ || + fc3_expert_weight_bits_ != expert_weight_bits_) { + return Status::OK(); + } + cudaStream_t stream = 0; // Use default stream for PrePack operations DUMP_TENSOR_INIT(); diff --git a/onnxruntime/contrib_ops/cuda/moe/moe_quantization.h b/onnxruntime/contrib_ops/cuda/moe/moe_quantization.h index 173c1e429bca4..5b46b61eb7e45 100644 --- a/onnxruntime/contrib_ops/cuda/moe/moe_quantization.h +++ b/onnxruntime/contrib_ops/cuda/moe/moe_quantization.h @@ -67,6 +67,9 @@ class QMoE final : public CudaKernel, public MoEBase { void PrePackIntExpertWeights(const Tensor& tensor, cudaStream_t stream, AllocatorPtr alloc, IAllocatorUniquePtr& packed_buf, bool& is_packed); int64_t expert_weight_bits_; + int64_t fc1_expert_weight_bits_; + int64_t fc2_expert_weight_bits_; + int64_t fc3_expert_weight_bits_; bool is_fp16_; // When true, the int4/int8 fc1/fc2 weight initializers are already in a // CUTLASS fpA_intB layout — produced offline e.g. via diff --git a/onnxruntime/contrib_ops/webgpu/moe/qmoe.cc b/onnxruntime/contrib_ops/webgpu/moe/qmoe.cc index b2797dbb93f52..fdc665e5ba0cb 100755 --- a/onnxruntime/contrib_ops/webgpu/moe/qmoe.cc +++ b/onnxruntime/contrib_ops/webgpu/moe/qmoe.cc @@ -200,9 +200,18 @@ Status QMoE::ComputeInternal(ComputeContext& context) const { fc1_experts_weights, fc1_experts_bias_optional, fc1_scales, fc1_zero_points, fc2_experts_weights, fc2_experts_bias_optional, fc2_scales, fc2_zero_points, fc3_experts_weights_optional, fc3_experts_bias_optional, fc3_scales_optional, fc3_zero_points, - expert_weight_bits_ == 4 ? 2 : 1, + moe_helper::MoEWeightBits{fc1_expert_weight_bits_, + fc2_expert_weight_bits_, + fc3_expert_weight_bits_}, activation_type_ == MoEActivationType::SwiGLU, block_size_)); + if (fc1_expert_weight_bits_ != expert_weight_bits_ || + fc2_expert_weight_bits_ != expert_weight_bits_ || + fc3_expert_weight_bits_ != expert_weight_bits_) { + return ORT_MAKE_STATUS(ONNXRUNTIME, NOT_IMPLEMENTED, + "Mixed-width QMoE execution is not yet implemented on WebGPU."); + } + const auto& input_shape = hidden_state->Shape(); // SwiGLU validation diff --git a/onnxruntime/contrib_ops/webgpu/moe/qmoe.h b/onnxruntime/contrib_ops/webgpu/moe/qmoe.h index 2b398e514c44d..3a192aecb0406 100755 --- a/onnxruntime/contrib_ops/webgpu/moe/qmoe.h +++ b/onnxruntime/contrib_ops/webgpu/moe/qmoe.h @@ -22,6 +22,15 @@ class QMoE final : public MoE { ORT_ENFORCE(info.GetAttr("expert_weight_bits", &expert_weight_bits_).IsOK()); ORT_ENFORCE(expert_weight_bits_ == 8 || expert_weight_bits_ == 4, "expert_weight_bits must be 4 or 8, but got ", expert_weight_bits_); + fc1_expert_weight_bits_ = info.GetAttrOrDefault("fc1_expert_weight_bits", expert_weight_bits_); + fc2_expert_weight_bits_ = info.GetAttrOrDefault("fc2_expert_weight_bits", expert_weight_bits_); + fc3_expert_weight_bits_ = info.GetAttrOrDefault("fc3_expert_weight_bits", expert_weight_bits_); + ORT_ENFORCE((fc1_expert_weight_bits_ == 2 || fc1_expert_weight_bits_ == 4 || fc1_expert_weight_bits_ == 8) && + (fc2_expert_weight_bits_ == 2 || fc2_expert_weight_bits_ == 4 || fc2_expert_weight_bits_ == 8) && + (fc3_expert_weight_bits_ == 2 || fc3_expert_weight_bits_ == 4 || fc3_expert_weight_bits_ == 8), + "FC-specific expert weight bits must be 2, 4, or 8."); + ORT_ENFORCE(swiglu_fusion_ == 0 || fc3_expert_weight_bits_ == fc1_expert_weight_bits_, + "Fused SwiGLU requires FC1 and FC3 expert weight bits to match."); block_size_ = static_cast(info.GetAttrOrDefault("block_size", 0)); } @@ -29,6 +38,9 @@ class QMoE final : public MoE { private: int64_t expert_weight_bits_; + int64_t fc1_expert_weight_bits_; + int64_t fc2_expert_weight_bits_; + int64_t fc3_expert_weight_bits_; int64_t block_size_; }; diff --git a/onnxruntime/contrib_ops/webgpu/quantization/subgroup_matrix_matmul_nbits.cc b/onnxruntime/contrib_ops/webgpu/quantization/subgroup_matrix_matmul_nbits.cc index ec87ebbb9a410..92ec6ad1f0b32 100644 --- a/onnxruntime/contrib_ops/webgpu/quantization/subgroup_matrix_matmul_nbits.cc +++ b/onnxruntime/contrib_ops/webgpu/quantization/subgroup_matrix_matmul_nbits.cc @@ -14,7 +14,7 @@ namespace webgpu { // The subgroup matrix config table, support check, and component-type validation live in the // shared core header (core/providers/webgpu/math/subgroup_matrix_config.h) so both this contrib // kernel and the core subgroup-matrix MatMul share them. -using onnxruntime::webgpu::IsSubgroupMatrixConfigSupported; +using onnxruntime::webgpu::SelectSubgroupMatrixConfig; using onnxruntime::webgpu::supported_subgroup_matrix_configs; // This program optimizes the layout of input matrix A(MxK) for SubgroupMatrixLoad, so that all elements of each @@ -347,26 +347,24 @@ bool CanApplySubgroupMatrixMatMulNBits(onnxruntime::webgpu::ComputeContext& cont return false; } - bool has_subgroup_matrix = context.HasFeature(wgpu::FeatureName::ChromiumExperimentalSubgroupMatrix); - if (has_subgroup_matrix) { - // Check if the adapter reports a subgroup matrix config we support. - has_subgroup_matrix = IsSubgroupMatrixConfigSupported(context, is_fp16, config_index); - if (has_subgroup_matrix) { - if (context.AdapterInfo().vendor == std::string_view{"apple"}) { - // For now SubgroupMatrixMatMulNBits is only supported for accuracy level 4, because with Fp16 there are - // some precision issues with subgroupMatrixMultiplyAccumulate. It is possible to support higher accuracy - // by setting compute_precision to Fp32, but that will be slower. For 1K token prefill FP16 Phi 3.5 is around 5s, - // FP32 is around 7s. - has_subgroup_matrix = accuracy_level == 4; - } - } + // On Apple, this kernel is only validated for accuracy level 4. Higher accuracy requires + // a slower f32-compute variant. + if (context.AdapterInfo().vendor == std::string_view{"apple"} && accuracy_level != 4) { + return false; + } + + if (block_size != 32 || batch_count != 1 || K % 32 != 0 || N % 64 != 0) { + return false; + } + + const auto selected_config = SelectSubgroupMatrixConfig( + context, is_fp16, {{16, 16, 16, 32}, {8, 16, 16, 32}, {8, 8, 8, 32}}); + if (!selected_config) { + return false; } - return has_subgroup_matrix && - block_size == 32 && - batch_count == 1 && - K % 32 == 0 && - N % 64 == 0; + config_index = *selected_config; + return true; } } // namespace webgpu } // namespace contrib diff --git a/onnxruntime/core/graph/contrib_ops/bert_defs.cc b/onnxruntime/core/graph/contrib_ops/bert_defs.cc index 556c167d95fe3..01b9e94856a21 100644 --- a/onnxruntime/core/graph/contrib_ops/bert_defs.cc +++ b/onnxruntime/core/graph/contrib_ops/bert_defs.cc @@ -239,7 +239,8 @@ void MultiHeadAttentionTypeAndShapeInference(ONNX_NAMESPACE::InferenceContext& c void BaseGroupQueryAttentionTypeAndShapeInference(ONNX_NAMESPACE::InferenceContext& ctx, int past_key_index = -1, int use_max_past_present_buffer = -1, - int output_qk_index = -1) { + int output_qk_index = -1, + int total_sequence_length_index = -1) { // Type inference for outputs ONNX_NAMESPACE::propagateElemTypeFromInputToOutput(ctx, 0, 0); // output @@ -301,9 +302,13 @@ void BaseGroupQueryAttentionTypeAndShapeInference(ONNX_NAMESPACE::InferenceConte if (ctx.getNumOutputs() >= 3) { // has present output int64_t total_sequence_length_value = 0; - const auto* total_sequence_length_data = ctx.getInputData(6); + const auto* total_sequence_length_data = + total_sequence_length_index >= 0 ? ctx.getInputData(total_sequence_length_index) : nullptr; if (total_sequence_length_data != nullptr) { const auto& data = ParseData(total_sequence_length_data); + if (data.size() != 1) { + fail_shape_inference("total_sequence_length input must contain a single element"); + } total_sequence_length_value = static_cast(data[0]); } @@ -456,14 +461,18 @@ void GroupQueryAttentionTypeAndShapeInference(ONNX_NAMESPACE::InferenceContext& // capacity C, which is deliberately smaller than total_sequence_length. present therefore keeps // the past buffer's own sequence dimension instead of growing with the total sequence length. const int64_t sliding_window_cache = getAttribute(ctx, "sliding_window_cache", 0); + constexpr int total_sequence_length_index = 6; BaseGroupQueryAttentionTypeAndShapeInference( - ctx, past_key_index, sliding_window_cache == 1 ? 1 : use_max_past_present_buffer, qk_output_index); + ctx, past_key_index, sliding_window_cache == 1 ? 1 : use_max_past_present_buffer, qk_output_index, + total_sequence_length_index); } void SparseAttentionTypeAndShapeInference(ONNX_NAMESPACE::InferenceContext& ctx, int past_key_index) { constexpr int use_max_past_present_buffer = 1; constexpr int qk_output_index = -1; - BaseGroupQueryAttentionTypeAndShapeInference(ctx, past_key_index, use_max_past_present_buffer, qk_output_index); + constexpr int total_sequence_length_index = 7; + BaseGroupQueryAttentionTypeAndShapeInference(ctx, past_key_index, use_max_past_present_buffer, qk_output_index, + total_sequence_length_index); } constexpr const char* Attention_ver1_doc = R"DOC( @@ -3026,7 +3035,12 @@ ONNX_MS_OPERATOR_SET_SCHEMA( fail_shape_inference("CausalConvWithState: channels_last must be 0 or 1, got ", channels_last); } - if (channels_last == 1 && getAttribute(ctx, "ndim", 1) != 1) { + + const int64_t ndim = getAttribute(ctx, "ndim", 1); + if (ndim < 1 || ndim > 3) { + fail_shape_inference("CausalConvWithState: ndim must be 1, 2, or 3, got ", ndim); + } + if (channels_last == 1 && ndim != 1) { fail_shape_inference("CausalConvWithState: channels_last requires ndim = 1"); } @@ -3040,13 +3054,20 @@ ONNX_MS_OPERATOR_SET_SCHEMA( if (hasInputShape(ctx, 0) && hasInputShape(ctx, 1)) { auto& input_shape = getInputShape(ctx, 0); auto& weight_shape = getInputShape(ctx, 1); - if (input_shape.dim_size() < 2) { - fail_shape_inference("CausalConvWithState: input must have rank >= 2"); + // weight is always channels-first: (channels, 1, k_1, ..., k_ndim), rank == ndim + 2. + if (weight_shape.dim_size() != ndim + 2) { + fail_shape_inference("CausalConvWithState: weight must have rank ndim + 2 (", + ndim + 2, "), got rank ", weight_shape.dim_size()); } - if (weight_shape.dim_size() < 2) { - fail_shape_inference("CausalConvWithState: weight must have rank >= 2"); + if (channels_last == 1) { + // (batch_size, sequence_length, d_1, ..., d_n). Check the lower bound. + if (input_shape.dim_size() < 3) { + fail_shape_inference("CausalConvWithState: channels_last input must have rank >= 3"); + } + } else if (input_shape.dim_size() != ndim + 2) { + fail_shape_inference("CausalConvWithState: input must have rank ndim + 2 (", + ndim + 2, "), got rank ", input_shape.dim_size()); } - int64_t ndim = getAttribute(ctx, "ndim", 1); // (kernel_size - 1) * dilation, or an unset dim when kernel_size is symbolic. const int last_kernel_dim = weight_shape.dim_size() - 1; TensorShapeProto::Dimension state_length; @@ -3702,9 +3723,31 @@ ONNX_MS_OPERATOR_SET_SCHEMA( const int64_t head_size_qk = query_shape.dim(token_dims + 1).dim_value(); const int64_t num_heads_v = value_shape.dim(token_dims).dim_value(); const int64_t head_size_v = value_shape.dim(token_dims + 1).dim_value(); - width->set_dim_value(state_update_capacity * - (num_heads_v + num_heads_k * head_size_qk + - num_heads_v * head_size_v)); + if (num_heads_k <= 0 || head_size_qk <= 0 || num_heads_v <= 0 || head_size_v <= 0) { + fail_shape_inference( + "GatedDeltaNet: head counts and head sizes must be positive"); + } + + constexpr int64_t max_dimension = std::numeric_limits::max(); + if (num_heads_k > max_dimension / head_size_qk || + num_heads_v > max_dimension / head_size_v) { + fail_shape_inference("GatedDeltaNet: state_update width overflows int64"); + } + const int64_t key_width = num_heads_k * head_size_qk; + const int64_t value_width = num_heads_v * head_size_v; + if (num_heads_v > max_dimension - key_width) { + fail_shape_inference("GatedDeltaNet: state_update width overflows int64"); + } + const int64_t width_without_capacity = num_heads_v + key_width; + if (value_width > max_dimension - width_without_capacity) { + fail_shape_inference("GatedDeltaNet: state_update width overflows int64"); + } + const int64_t per_token_width = width_without_capacity + value_width; + if (state_update_capacity > 0 && + per_token_width > max_dimension / state_update_capacity) { + fail_shape_inference("GatedDeltaNet: state_update width overflows int64"); + } + width->set_dim_value(state_update_capacity * per_token_width); } updateOutputShape(ctx, 2, capsule_shape); } diff --git a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc index 820c7ee579fdc..98208bd505ba5 100644 --- a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc +++ b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc @@ -1452,8 +1452,13 @@ constexpr const char* qMoE_ver1_doc = R"DOC( If block_size is provided, both hidden_size and inter_size must be divisible by the block size, and the dequantization is performed per block of size block_size along the K (input feature) dimension. - If block_size and zero_point are provided, both hidden_size and inter_size must be divisible by block_size * pack_size, - where pack_size = 8 / expert_weight_bits. + Packed byte dimensions are computed as logical_element_count * effective_expert_weight_bits / 8. + Weight rows must be byte-aligned. Zero-point rows are padded to a whole byte when necessary. + + fc1_expert_weight_bits, fc2_expert_weight_bits, and fc3_expert_weight_bits optionally override + expert_weight_bits for the corresponding projection. An omitted override inherits expert_weight_bits. + When SwiGLU is fused, FC3 is stored in FC1 and fc3_expert_weight_bits must be omitted or equal to + fc1_expert_weight_bits after inheritance. The SwiGLU (Swish-Gated Linear Unit) activation function is like: g = xW + b @@ -1491,6 +1496,19 @@ ONNX_MS_OPERATOR_SET_SCHEMA( "Number of bits used in quantized weights. Supported values are 2, 4, and 8. Default is 4 bits", AttributeProto::INT, static_cast(4)) + .Attr("fc1_expert_weight_bits", + "Optional FC1 override for expert_weight_bits. Inherits expert_weight_bits when omitted.", + AttributeProto::INT, + OPTIONAL_VALUE) + .Attr("fc2_expert_weight_bits", + "Optional FC2 override for expert_weight_bits. Inherits expert_weight_bits when omitted.", + AttributeProto::INT, + OPTIONAL_VALUE) + .Attr("fc3_expert_weight_bits", + "Optional FC3 override for expert_weight_bits. Inherits expert_weight_bits when omitted. " + "For fused SwiGLU, the effective FC3 width must equal the effective FC1 width.", + AttributeProto::INT, + OPTIONAL_VALUE) .Attr("swiglu_fusion", "0: not fused, 1: fused and interleaved. 2: fused and not interleaved.", AttributeProto::INT, @@ -1549,8 +1567,10 @@ ONNX_MS_OPERATOR_SET_SCHEMA( "T") .Input(2, "fc1_experts_weights", - "3D tensor with shape (num_experts, fusion_size * inter_size, hidden_size / pack_size), " - "The fusion_size is 2 for fused swiglu, or 1 otherwise. The pack_size is 8 / expert_weight_bits.", + "3D tensor with shape (num_experts, fusion_size * inter_size, " + "hidden_size * effective_fc1_bits / 8). The last dimension must be byte-aligned. " + "The fusion_size is 2 for fused swiglu, or 1 otherwise. effective_fc1_bits is " + "fc1_expert_weight_bits when provided, otherwise expert_weight_bits.", "T1") .Input(3, "fc1_scales", @@ -1568,7 +1588,9 @@ ONNX_MS_OPERATOR_SET_SCHEMA( "2D optional tensor with shape (num_experts, fusion_size * inter_size)", "T", OpSchema::Optional) .Input(5, "fc2_experts_weights", - "3D tensor with shape (num_experts, hidden_size, inter_size / pack_size)", + "3D tensor with shape (num_experts, hidden_size, inter_size * effective_fc2_bits / 8). " + "The last dimension must be byte-aligned. effective_fc2_bits is fc2_expert_weight_bits " + "when provided, otherwise expert_weight_bits.", "T1") .Input(6, "fc2_scales", @@ -1588,7 +1610,9 @@ ONNX_MS_OPERATOR_SET_SCHEMA( OpSchema::Optional) .Input(8, "fc3_experts_weights", - "3D optional tensor with shape (num_experts, inter_size, hidden_size / pack_size)", + "3D optional tensor with shape (num_experts, inter_size, hidden_size * effective_fc3_bits / 8). " + "The last dimension must be byte-aligned. effective_fc3_bits is fc3_expert_weight_bits " + "when provided, otherwise expert_weight_bits.", "T1", OpSchema::Optional) .Input(9, @@ -1607,20 +1631,23 @@ ONNX_MS_OPERATOR_SET_SCHEMA( OpSchema::Optional) .Input(11, "fc1_zero_points", - "2D tensor with shape (num_experts, fusion_size * inter_size / pack_size), or " - "3D tensor with shape (num_experts, fusion_size * inter_size, hidden_size / block_size / pack_size) when block_size is provided.", + "2D tensor with shape (num_experts, ceil(fusion_size * inter_size * effective_fc1_bits / 8)), or " + "3D tensor with shape (num_experts, fusion_size * inter_size, " + "ceil((hidden_size / block_size) * effective_fc1_bits / 8)) when block_size is provided.", "T1", OpSchema::Optional) .Input(12, "fc2_zero_points", - "2D tensor with shape (num_experts, hidden_size / pack_size), or " - "3D tensor with shape (num_experts, hidden_size, inter_size / block_size / pack_size) when block_size is provided.", + "2D tensor with shape (num_experts, ceil(hidden_size * effective_fc2_bits / 8)), or " + "3D tensor with shape (num_experts, hidden_size, " + "ceil((inter_size / block_size) * effective_fc2_bits / 8)) when block_size is provided.", "T1", OpSchema::Optional) .Input(13, "fc3_zero_points", - "2D optional tensor with shape (num_experts, inter_size / pack_size), or " - "3D optional tensor with shape (num_experts, inter_size, hidden_size / block_size / pack_size) when block_size is provided.", + "2D optional tensor with shape (num_experts, ceil(inter_size * effective_fc3_bits / 8)), or " + "3D optional tensor with shape (num_experts, inter_size, " + "ceil((hidden_size / block_size) * effective_fc3_bits / 8)) when block_size is provided.", "T1", OpSchema::Optional) .Input(14, diff --git a/onnxruntime/core/providers/webgpu/math/subgroup_matrix_config.cc b/onnxruntime/core/providers/webgpu/math/subgroup_matrix_config.cc index f3db479b3ec14..c27f04042ae68 100644 --- a/onnxruntime/core/providers/webgpu/math/subgroup_matrix_config.cc +++ b/onnxruntime/core/providers/webgpu/math/subgroup_matrix_config.cc @@ -5,14 +5,21 @@ #include +#include "core/common/inlined_containers.h" #include "core/providers/webgpu/compute_context.h" namespace onnxruntime { namespace webgpu { +namespace { -bool IsSubgroupMatrixConfigSupported(const ComputeContextBase& context, bool is_fp16, int32_t& config_index) { +using SubgroupMatrixConfigIndices = InlinedVector; + +SubgroupMatrixConfigIndices GetSupportedSubgroupMatrixConfigIndices( + const ComputeContextBase& context, + bool is_fp16) { + SubgroupMatrixConfigIndices candidates; if (!context.HasFeature(wgpu::FeatureName::ChromiumExperimentalSubgroupMatrix)) { - return false; + return {}; } const wgpu::AdapterInfo& adapter_info = context.AdapterInfo(); const wgpu::AdapterPropertiesSubgroupMatrixConfigs& subgroup_matrix_configs = context.SubgroupMatrixConfigs(); @@ -35,13 +42,45 @@ bool IsSubgroupMatrixConfigSupported(const ComputeContextBase& context, bool is_ IsSubgroupSizeSupported(adapter_info.subgroupMinSize, adapter_info.subgroupMaxSize, supported_config.subgroupSize, context.HasFeature(wgpu::FeatureName::SubgroupSizeControl))) { - config_index = index; - return true; + candidates.push_back(index); + break; } } index++; } - return false; + return candidates; +} + +} // namespace + +namespace detail { + +std::optional SelectSubgroupMatrixConfigFromCandidates( + gsl::span candidate_indices, + std::initializer_list preferences) { + for (const auto& preference : preferences) { + for (const int32_t index : candidate_indices) { + if (index < 0 || static_cast(index) >= supported_subgroup_matrix_configs.size()) { + continue; + } + const auto& config = supported_subgroup_matrix_configs[index]; + if (config.Is(preference.M, preference.N, preference.K) && + config.subgroupSize == preference.subgroupSize) { + return index; + } + } + } + return std::nullopt; +} + +} // namespace detail + +std::optional SelectSubgroupMatrixConfig( + const ComputeContextBase& context, + bool is_fp16, + std::initializer_list preferences) { + const auto candidates = GetSupportedSubgroupMatrixConfigIndices(context, is_fp16); + return detail::SelectSubgroupMatrixConfigFromCandidates(candidates, preferences); } } // namespace webgpu diff --git a/onnxruntime/core/providers/webgpu/math/subgroup_matrix_config.h b/onnxruntime/core/providers/webgpu/math/subgroup_matrix_config.h index c6a3929db2bce..080238019733b 100644 --- a/onnxruntime/core/providers/webgpu/math/subgroup_matrix_config.h +++ b/onnxruntime/core/providers/webgpu/math/subgroup_matrix_config.h @@ -6,8 +6,12 @@ #include #include #include +#include +#include #include +#include + #include "core/providers/webgpu/webgpu_external_header.h" namespace onnxruntime { @@ -68,6 +72,13 @@ struct SupportedSubgroupMatrixConfig { } }; +struct SubgroupMatrixConfigPreference { + uint32_t M; + uint32_t N; + uint32_t K; + uint32_t subgroupSize; +}; + // A fixed-size adapter already guarantees the required size. An adapter exposing a range needs // subgroup-size control so the kernel can select its required size instead of relying on the // implementation's choice. @@ -88,10 +99,29 @@ inline constexpr std::array supported_subgroup {wgpu::SubgroupMatrixComponentType::F32, wgpu::SubgroupMatrixComponentType::F32, 8, 8, 8, 32, false}, }}; -// Returns true and sets config_index (into supported_subgroup_matrix_configs) when the device -// reports one of the supported configs matching the requested output precision. is_fp16 selects -// F16-output configs; otherwise F32-output configs. -bool IsSubgroupMatrixConfigSupported(const ComputeContextBase& context, bool is_fp16, int32_t& config_index); +// Selects a subgroup-matrix configuration supported by both the operation and the device. +// +// `is_fp16` restricts candidates to F16 configs when true and F32 configs when false. +// `preferences` lists the matrix shape and subgroup size combinations implemented by the +// operation, in performance-preference order. A candidate must also be reported by the adapter +// and have a usable subgroup size. Fixed-size adapters need no subgroup-size-control feature; +// adapters reporting a size range must support subgroup-size control. +// +// Returns the selected index in `supported_subgroup_matrix_configs`, or `std::nullopt` when no +// configuration satisfies all requirements. +std::optional SelectSubgroupMatrixConfig( + const ComputeContextBase& context, + bool is_fp16, + std::initializer_list preferences); + +namespace detail { + +// Separated from device capability discovery for focused preference-order testing. +std::optional SelectSubgroupMatrixConfigFromCandidates( + gsl::span candidate_indices, + std::initializer_list preferences); + +} // namespace detail } // namespace webgpu } // namespace onnxruntime diff --git a/onnxruntime/core/providers/webgpu/math/subgroup_matrix_gemm.cc b/onnxruntime/core/providers/webgpu/math/subgroup_matrix_gemm.cc index 1206f9773bdd7..8ade9a875e83c 100644 --- a/onnxruntime/core/providers/webgpu/math/subgroup_matrix_gemm.cc +++ b/onnxruntime/core/providers/webgpu/math/subgroup_matrix_gemm.cc @@ -136,7 +136,9 @@ class SubgroupMatrixGemmImpl final : public Gemm::GemmOptImpl { SubgroupMatrixGemmProgram program{has_c, trans_a, trans_b, config_index_, sg_mat_count_m, sg_mat_count_n, split_k}; program.SetWorkgroupSize(config.subgroupSize * split_k); - program.SetSubgroupSize(config.subgroupSize); + if (context.HasFeature(wgpu::FeatureName::SubgroupSizeControl)) { + program.SetSubgroupSize(config.subgroupSize); + } program.SetDispatchGroupSize(dispatch_x, dispatch_y, 1); program.CacheHint(has_c, trans_a, trans_b, config_index_, sg_mat_count_m, sg_mat_count_n, split_k) .AddInputs({{a, ProgramTensorMetadataDependency::TypeAndRank, 1}, @@ -203,14 +205,10 @@ Status SubgroupMatrixGemmProgram::GenerateShaderCode(ShaderHelper& shader) const std::unique_ptr CreateSubgroupMatrixGemmImpl( const Gemm& parent, const ComputeContextBase& context) { - // Only run on devices that report the fixed 8x16x16 F16 subgroup-matrix config - // this kernel is implemented for. That config's adapters expose a 16-32 subgroup - // size range, so the kernel's fixed 32 lanes per subgroup must be pinned with - // subgroup-size control. - int32_t config_index = 0; - if (!IsSubgroupMatrixConfigSupported(context, /*is_fp16=*/true, config_index) || - !supported_subgroup_matrix_configs[config_index].Is(8, 16, 16) || - !context.HasFeature(wgpu::FeatureName::SubgroupSizeControl)) { + // Only run on devices that report the 8x16x16 F16 subgroup-matrix config this + // kernel is implemented for and can provide its required subgroup size. + const auto config_index = SelectSubgroupMatrixConfig(context, /*is_fp16=*/true, {{8, 16, 16, 32}}); + if (!config_index) { return nullptr; } // Intel GPUs use a tuned/heuristic tiling policy; every other vendor falls back @@ -221,7 +219,7 @@ std::unique_ptr CreateSubgroupMatrixGemmImpl( if (!tiling_selector) { return nullptr; } - return std::make_unique(parent, config_index, std::move(tiling_selector)); + return std::make_unique(parent, *config_index, std::move(tiling_selector)); } } // namespace webgpu diff --git a/onnxruntime/core/providers/webgpu/math/subgroup_matrix_matmul.cc b/onnxruntime/core/providers/webgpu/math/subgroup_matrix_matmul.cc index 4888b5e981b37..562b395a9cbfa 100644 --- a/onnxruntime/core/providers/webgpu/math/subgroup_matrix_matmul.cc +++ b/onnxruntime/core/providers/webgpu/math/subgroup_matrix_matmul.cc @@ -188,7 +188,9 @@ class SubgroupMatrixMatMulImpl final : public MatMulOptImpl { SubgroupMatrixMatMulProgram program{activation, has_bias, config_index_, sg_mat_count_m, sg_mat_count_n, split_k}; program.SetWorkgroupSize(config.subgroupSize * split_k); - program.SetSubgroupSize(config.subgroupSize); + if (context.HasFeature(wgpu::FeatureName::SubgroupSizeControl)) { + program.SetSubgroupSize(config.subgroupSize); + } program.SetDispatchGroupSize(dispatch_x, dispatch_y, batch); program.CacheHint(activation.CacheKey(), has_bias, config_index_, sg_mat_count_m, sg_mat_count_n, split_k) @@ -317,14 +319,10 @@ Status SubgroupMatrixMatMulProgram::GenerateShaderCode(ShaderHelper& shader) con } std::unique_ptr CreateSubgroupMatrixMatMulImpl(const ComputeContextBase& context) { - // Only run on devices that report the fixed 8x16x16 F16 subgroup-matrix config - // this kernel is implemented for. That config's adapters expose a 16-32 subgroup - // size range, so the kernel's fixed 32 lanes per subgroup must be pinned with - // subgroup-size control. - int32_t config_index = 0; - if (!IsSubgroupMatrixConfigSupported(context, /*is_fp16=*/true, config_index) || - !supported_subgroup_matrix_configs[config_index].Is(8, 16, 16) || - !context.HasFeature(wgpu::FeatureName::SubgroupSizeControl)) { + // Only run on devices that report the 8x16x16 F16 subgroup-matrix config this + // kernel is implemented for and can provide its required subgroup size. + const auto config_index = SelectSubgroupMatrixConfig(context, /*is_fp16=*/true, {{8, 16, 16, 32}}); + if (!config_index) { return nullptr; } // Intel GPUs use a tuned/heuristic tiling policy; every other vendor falls back @@ -335,7 +333,7 @@ std::unique_ptr CreateSubgroupMatrixMatMulImpl(const ComputeConte if (!tiling_selector) { return nullptr; } - return std::make_unique(config_index, std::move(tiling_selector)); + return std::make_unique(*config_index, std::move(tiling_selector)); } } // namespace webgpu diff --git a/onnxruntime/test/contrib_ops/causal_conv_with_state_op_test.cc b/onnxruntime/test/contrib_ops/causal_conv_with_state_op_test.cc index d14add051807b..759488ce7e005 100644 --- a/onnxruntime/test/contrib_ops/causal_conv_with_state_op_test.cc +++ b/onnxruntime/test/contrib_ops/causal_conv_with_state_op_test.cc @@ -1100,6 +1100,67 @@ TEST(CausalConvWithStateTest, StateWindowAboveMaxIsRejected) { test.Run(OpTester::ExpectResult::kExpectFailure, "state_window must be in [0, 8]"); } +// These tests exercise shape inference which uses fail_shape_inference (throws InferenceError). +// In no-exception builds, fail_shape_inference calls abort(), so these tests must be skipped. +#ifndef ORT_NO_EXCEPTIONS +// ndim outside [1, 3] range is rejected. +TEST(CausalConvWithStateTest, NdimOutOfRangeIsRejected) { + OpTester test("CausalConvWithState", 1, onnxruntime::kMSDomain); + test.AddAttribute("activation", "none"); + test.AddAttribute("ndim", 4); + test.AddInput("input", {1, 1, 2}, {1.0f, 2.0f}); + test.AddInput("weight", {1, 1, 2}, {0.5f, 0.25f}); + test.AddOptionalInputEdge(); // bias + test.AddOptionalInputEdge(); // past_state + test.AddOutput("output", {1, 1, 2}, {0.0f, 0.0f}); + test.AddOutput("present_state", {1, 1, 1}, {0.0f}); + test.Run(OpTester::ExpectResult::kExpectFailure, "ndim must be 1, 2, or 3"); +} + +// With ndim=2 the channels-first input must have rank ndim + 2 == 4 +TEST(CausalConvWithStateTest, InputRankMismatchIsRejected) { + OpTester test("CausalConvWithState", 1, onnxruntime::kMSDomain); + test.AddAttribute("activation", "none"); + test.AddAttribute("ndim", 2); + test.AddInput("input", {1, 1, 2}, {1.0f, 2.0f}); + test.AddInput("weight", {1, 1, 2, 2}, {0.5f, 0.25f, 0.5f, 0.25f}); + test.AddOptionalInputEdge(); // bias + test.AddOptionalInputEdge(); // past_state + test.AddOutput("output", {1, 1, 2}, {0.0f, 0.0f}); + test.AddOutput("present_state", {1, 1, 1}, {0.0f}); + test.Run(OpTester::ExpectResult::kExpectFailure, "input must have rank ndim + 2"); +} + +// The weight is always channels-first with rank ndim + 2. +TEST(CausalConvWithStateTest, WeightRankMismatchIsRejected) { + OpTester test("CausalConvWithState", 1, onnxruntime::kMSDomain); + test.AddAttribute("activation", "none"); + test.AddAttribute("ndim", 1); + test.AddInput("input", {1, 1, 2}, {1.0f, 2.0f}); + test.AddInput("weight", {1, 1}, {0.5f}); + test.AddOptionalInputEdge(); // bias + test.AddOptionalInputEdge(); // past_state + test.AddOutput("output", {1, 1, 2}, {0.0f, 0.0f}); + test.AddOutput("present_state", {1, 1, 1}, {0.0f}); + test.Run(OpTester::ExpectResult::kExpectFailure, "weight must have rank ndim + 2"); +} + +// channels_last requires ndim == 1 and lays the input out as (batch, length, ...channels). +TEST(CausalConvWithStateTest, ChannelsLastInputRankBelowThreeIsRejected) { + OpTester test("CausalConvWithState", 1, onnxruntime::kMSDomain); + test.AddAttribute("activation", "none"); + test.AddAttribute("ndim", 1); + test.AddAttribute("channels_last", 1); + test.AddInput("input", {1, 2}, {1.0f, 2.0f}); + test.AddInput("weight", {1, 1, 2}, {0.5f, 0.25f}); + test.AddOptionalInputEdge(); // bias + test.AddOptionalInputEdge(); // past_state + test.AddOutput("output", {1, 2}, {0.0f, 0.0f}); + test.AddOutput("present_state", {1, 1, 1}, {0.0f}); + test.Run(OpTester::ExpectResult::kExpectFailure, "channels_last input must have rank >= 3"); +} +#endif // !ORT_NO_EXCEPTIONS + #ifdef USE_CUDA TEST(CausalConvWithStateTest, StateWindowRejectsEmptySequence) { auto ep = DefaultCudaExecutionProvider(); diff --git a/onnxruntime/test/contrib_ops/gated_delta_net_op_test.cc b/onnxruntime/test/contrib_ops/gated_delta_net_op_test.cc index b67cb5847481f..ac760fd31caed 100644 --- a/onnxruntime/test/contrib_ops/gated_delta_net_op_test.cc +++ b/onnxruntime/test/contrib_ops/gated_delta_net_op_test.cc @@ -2,7 +2,9 @@ // Licensed under the MIT License. #include +#include #include +#include #include #include #include @@ -1266,6 +1268,76 @@ TEST(GatedDeltaNetTest, RejectsPerKeyDtBias) { {}, nullptr, &eps); } +// This test exercises shape inference which uses fail_shape_inference (throws InferenceError). +// In no-exception builds, fail_shape_inference calls abort(), so this test must be skipped. +#ifndef ORT_NO_EXCEPTIONS +TEST(GatedDeltaNetTest, RejectsStateUpdateWidthOverflow) { + struct Case { + std::array query_dims; + std::array value_dims; + int64_t state_update_capacity; + }; + + const int64_t max_dimension = std::numeric_limits::max(); + const int64_t large_head_size = 4000000000LL; + const std::vector cases = { + // num_heads_k * head_size_qk. + {{1, large_head_size, large_head_size}, {1, 1, 1}, 1}, + // num_heads_v * head_size_v. + {{1, 1, 1}, {1, large_head_size, large_head_size}, 1}, + // num_heads_v + (num_heads_k * head_size_qk). + {{1, 1, max_dimension}, {1, 1, 1}, 1}, + // value_width + (num_heads_v + key_width). + {{1, 1, 1}, {1, 1, max_dimension}, 1}, + // state_update_capacity * per_token_width. + {{1, 1, 1}, {1, 1, max_dimension / 8}, 8}, + }; + + for (size_t case_index = 0; case_index < cases.size(); ++case_index) { + SCOPED_TRACE(case_index); + const auto& test_case = cases[case_index]; + std::unordered_map domain_to_version = {{kMSDomain, 1}}; + std::vector functions; + auto model = std::make_unique( + "gated_delta_net_overflow", true, ModelMetaData(), PathString(), + IOnnxRuntimeOpSchemaRegistryList(), domain_to_version, functions, + DefaultLoggingManager().DefaultLogger(), ModelOptions(true, true)); + auto& graph = model->MainGraph(); + + std::vector types; + types.reserve(6); + auto tensor_type = [&](int elem_type, const auto& dims) { + types.emplace_back(); + auto* type = &types.back(); + type->mutable_tensor_type()->set_elem_type(elem_type); + for (int64_t dim : dims) { + type->mutable_tensor_type()->mutable_shape()->add_dim()->set_dim_value(dim); + } + return type; + }; + + auto& query_arg = graph.GetOrCreateNodeArg( + "query", tensor_type(ONNX_NAMESPACE::TensorProto_DataType_FLOAT, test_case.query_dims)); + auto& key_arg = graph.GetOrCreateNodeArg( + "key", tensor_type(ONNX_NAMESPACE::TensorProto_DataType_FLOAT, std::array{1, 1, 1})); + auto& value_arg = graph.GetOrCreateNodeArg( + "value", tensor_type(ONNX_NAMESPACE::TensorProto_DataType_FLOAT, test_case.value_dims)); + auto& output_arg = graph.GetOrCreateNodeArg( + "output", tensor_type(ONNX_NAMESPACE::TensorProto_DataType_FLOAT, std::array{})); + auto& final_state_arg = graph.GetOrCreateNodeArg( + "final_state", tensor_type(ONNX_NAMESPACE::TensorProto_DataType_FLOAT, std::array{})); + auto& state_update_arg = graph.GetOrCreateNodeArg( + "state_update", tensor_type(ONNX_NAMESPACE::TensorProto_DataType_FLOAT, std::array{})); + + auto& node = graph.AddNode("node", "GatedDeltaNet", "", {&query_arg, &key_arg, &value_arg}, + {&output_arg, &final_state_arg, &state_update_arg}, nullptr, kMSDomain); + node.AddAttribute("state_update_capacity", test_case.state_update_capacity); + + ASSERT_STATUS_NOT_OK_AND_HAS_SUBSTR(graph.Resolve(), "overflows int64"); + } +} +#endif // !ORT_NO_EXCEPTIONS + TEST(GatedDeltaNetTest, RequiresCaptureCountExactlyWhenCapacityIsPositive) { if (NeedSkipGatedDeltaNetTest()) return; Geometry g{4, 1, 1, 2, 64, 32}; diff --git a/onnxruntime/test/contrib_ops/group_query_attention_op_test.cc b/onnxruntime/test/contrib_ops/group_query_attention_op_test.cc index a42452d0e538a..299930fe7807f 100644 --- a/onnxruntime/test/contrib_ops/group_query_attention_op_test.cc +++ b/onnxruntime/test/contrib_ops/group_query_attention_op_test.cc @@ -74,7 +74,10 @@ static void RunGQASeqlensKTest( const std::string& expected_message, bool provide_past = false, int past_seq_len = 0, - const std::optional>& seqlens_k_shape = std::nullopt) { + const std::optional>& seqlens_k_shape = std::nullopt, + const std::optional>& total_seq_len_shape = std::nullopt, + const std::optional>& total_seq_len_data = std::nullopt, + bool total_seq_len_is_initializer = false) { constexpr int num_heads = 1; constexpr int kv_num_heads = 1; constexpr int head_size = 8; @@ -108,7 +111,9 @@ static void RunGQASeqlensKTest( ? *seqlens_k_shape : std::vector{batch_size}; tester.AddInput("seqlens_k", shape, seqlens_k_data); - tester.AddInput("total_sequence_length", {1}, {total_seq_len}); + const std::vector ts_shape = total_seq_len_shape.value_or(std::vector{1}); + const std::vector ts_data = total_seq_len_data.value_or(std::vector{total_seq_len}); + tester.AddInput("total_sequence_length", ts_shape, ts_data, total_seq_len_is_initializer); tester.AddOptionalInputEdge(); // cos_cache tester.AddOptionalInputEdge(); // sin_cache @@ -1155,6 +1160,27 @@ TEST(GroupQueryAttentionTest, SeqlensKScalarRejected) { /*seqlens_k_shape=*/std::vector{}); } +// This test exercises shape inference which uses fail_shape_inference (throws InferenceError). +// In no-exception builds, fail_shape_inference calls abort(), so this test must be skipped. +#ifndef ORT_NO_EXCEPTIONS +// total_sequence_length constant must have a single element. +TEST(GroupQueryAttentionTest, EmptyTotalSequenceLengthInitializerRejected) { + RunGQASeqlensKTest( + /*seqlens_k_data=*/{0}, + /*total_seq_len=*/1, + /*batch_size=*/1, + /*sequence_length=*/1, + OpTester::ExpectResult::kExpectFailure, + "total_sequence_length input must contain a single element", + /*provide_past=*/false, + /*past_seq_len=*/0, + /*seqlens_k_shape=*/std::nullopt, + /*total_seq_len_shape=*/std::vector{0}, + /*total_seq_len_data=*/std::vector{}, + /*total_seq_len_is_initializer=*/true); +} +#endif // !ORT_NO_EXCEPTIONS + // Helper to compare two output vectors (non-zero check + element-wise tolerance). static void ExpectOutputsMatch(const std::vector& a, const std::vector& b, float tolerance, const char* label) { diff --git a/onnxruntime/test/contrib_ops/moe_test.cc b/onnxruntime/test/contrib_ops/moe_test.cc index e27a8130bbd82..fb2e0865f33aa 100644 --- a/onnxruntime/test/contrib_ops/moe_test.cc +++ b/onnxruntime/test/contrib_ops/moe_test.cc @@ -14,6 +14,7 @@ #include #include "nlohmann/json.hpp" +#include "contrib_ops/cpu/moe/moe_helper.h" #include "core/mlas/inc/mlas_qnbit.h" #include "core/session/onnxruntime_session_options_config_keys.h" #include "test/util/include/scoped_env_vars.h" @@ -1879,7 +1880,6 @@ TEST(MoETest, QMoETest_CPU_Int2_BlockWiseLutIdentity) { } TEST(MoETest, QMoETest_CPU_Int2_InvalidHiddenSize) { -#ifdef USE_MLAS auto cpu_ep = DefaultCpuExecutionProvider(); if (!cpu_ep) { GTEST_SKIP() << "CPU execution provider not available"; @@ -1934,9 +1934,220 @@ TEST(MoETest, QMoETest_CPU_Int2_InvalidHiddenSize) { {}, nullptr, &cpu_execution_providers); -#else - GTEST_SKIP() << "Skipping CPU QMoE test"; +} + +static void RunQMoEMixedWidthContractTest(bool invalid_fc1_shape, + std::unique_ptr execution_provider, + const char* provider_name, + bool use_raw_weights = false, + bool use_float16_scales = false, + int64_t hidden_size = 8, + int64_t inter_size = 8, + bool legacy_layout = false) { + constexpr int64_t num_rows = 1; + constexpr int64_t num_experts = 1; + constexpr int64_t fc1_bits = 2; + constexpr int64_t fc2_bits = 4; + constexpr int64_t fc1_pack_size = 8 / fc1_bits; + constexpr int64_t fc2_pack_size = 8 / fc2_bits; + + const std::vector input_dims = {num_rows, hidden_size}; + const std::vector router_probs_dims = {num_rows, num_experts}; + std::vector fc1_weights_dims = legacy_layout + ? std::vector{num_experts, hidden_size, + inter_size / fc1_pack_size} + : std::vector{num_experts, inter_size, + hidden_size / fc1_pack_size}; + if (invalid_fc1_shape) { + ++fc1_weights_dims[2]; + } + const std::vector fc2_weights_dims = legacy_layout + ? std::vector{num_experts, inter_size, + hidden_size / fc2_pack_size} + : std::vector{num_experts, hidden_size, + inter_size / fc2_pack_size}; + const std::vector fc1_scales_dims = {num_experts, inter_size}; + const std::vector fc2_scales_dims = {num_experts, hidden_size}; + + OpTester tester("QMoE", 1, onnxruntime::kMSDomain); + tester.AddAttribute("k", 1); + tester.AddAttribute("activation_type", "identity"); + tester.AddAttribute("expert_weight_bits", fc2_bits); + tester.AddAttribute("fc1_expert_weight_bits", fc1_bits); + if (use_raw_weights) { + tester.AddAttribute("weights_prepacked", 0); + } + tester.AddInput("input", input_dims, std::vector(num_rows * hidden_size)); + tester.AddInput("router_probs", router_probs_dims, std::vector(num_rows * num_experts)); + tester.AddInput("fc1_experts_weights", fc1_weights_dims, + std::vector(static_cast(fc1_weights_dims[1] * fc1_weights_dims[2]))); + if (use_float16_scales) { + tester.AddInput("fc1_scales", fc1_scales_dims, + std::vector(num_experts * inter_size, MLFloat16(1.0f))); + } else { + tester.AddInput("fc1_scales", fc1_scales_dims, std::vector(num_experts * inter_size, 1.0f)); + } + tester.AddOptionalInputEdge(); + tester.AddInput("fc2_experts_weights", fc2_weights_dims, + std::vector(static_cast(fc2_weights_dims[1] * fc2_weights_dims[2]))); + if (use_float16_scales) { + tester.AddInput("fc2_scales", fc2_scales_dims, + std::vector(num_experts * hidden_size, MLFloat16(1.0f))); + } else { + tester.AddInput("fc2_scales", fc2_scales_dims, std::vector(num_experts * hidden_size, 1.0f)); + } + tester.AddOptionalInputEdge(); + tester.AddOptionalInputEdge(); + tester.AddOptionalInputEdge(); + tester.AddOptionalInputEdge(); + tester.AddOutput("output", input_dims, std::vector(num_rows * hidden_size)); + + std::vector> execution_providers; + execution_providers.push_back(std::move(execution_provider)); + const std::string expected_error = + invalid_fc1_shape ? "Input 'fc1_experts_weights' is expected to have shape" + : MakeString("Mixed-width QMoE execution is not yet implemented on ", provider_name, "."); + tester.Run(OpTester::ExpectResult::kExpectFailure, + expected_error, {}, nullptr, &execution_providers); +} + +TEST(MoETest, QMoETest_MixedWidthContract) { + RunQMoEMixedWidthContractTest(false, DefaultCpuExecutionProvider(), "CPU"); +} + +TEST(MoETest, QMoETest_MixedWidthInvalidFC1Shape) { + RunQMoEMixedWidthContractTest(true, DefaultCpuExecutionProvider(), "CPU"); +} + +TEST(MoETest, QMoETest_MixedWidthNonSquareLayouts) { + RunQMoEMixedWidthContractTest(false, DefaultCpuExecutionProvider(), "CPU", false, false, 8, 16, false); + RunQMoEMixedWidthContractTest(false, DefaultCpuExecutionProvider(), "CPU", false, false, 8, 16, true); +} + +#if defined(USE_CUDA) +TEST(MoETest, QMoETest_MixedWidthContract_CUDA) { + if (!HasCudaEnvironment(700)) { + GTEST_SKIP() << "CUDA device with compute capability 7.0 or newer is required."; + } + RunQMoEMixedWidthContractTest(false, DefaultCudaExecutionProvider(), "CUDA", true, true); +} +#endif + +#if defined(USE_WEBGPU) +TEST(MoETest, QMoETest_MixedWidthContract_WebGPU) { + auto execution_provider = DefaultWebGpuExecutionProvider(); + if (!execution_provider) { + GTEST_SKIP() << "WebGPU execution provider not available"; + } + RunQMoEMixedWidthContractTest(false, std::move(execution_provider), "WebGPU"); +} #endif + +static void RunQMoEMixedWidthZeroPointTest(bool block_wise, bool invalid_fc2_shape) { + constexpr int64_t num_experts = 1; + constexpr int64_t hidden_size = 48; + constexpr int64_t inter_size = 48; + constexpr int64_t fc1_bits = 2; + constexpr int64_t fc2_bits = 4; + constexpr int64_t block_size = 16; + constexpr int64_t blocks_per_row = hidden_size / block_size; + + const std::vector fc1_scales_dims = + block_wise ? std::vector{num_experts, inter_size, blocks_per_row} + : std::vector{num_experts, inter_size}; + const std::vector fc2_scales_dims = + block_wise ? std::vector{num_experts, hidden_size, blocks_per_row} + : std::vector{num_experts, hidden_size}; + const std::vector fc1_zero_points_dims = + block_wise ? std::vector{num_experts, inter_size, 1} + : std::vector{num_experts, 12}; + std::vector fc2_zero_points_dims = + block_wise ? std::vector{num_experts, hidden_size, 2} + : std::vector{num_experts, 24}; + if (invalid_fc2_shape) { + --fc2_zero_points_dims.back(); + } + + OpTester tester("QMoE", 1, onnxruntime::kMSDomain); + tester.AddAttribute("k", 1); + tester.AddAttribute("activation_type", "identity"); + tester.AddAttribute("expert_weight_bits", fc2_bits); + tester.AddAttribute("fc1_expert_weight_bits", fc1_bits); + if (block_wise) { + tester.AddAttribute("block_size", block_size); + } + tester.AddInput("input", {1, hidden_size}, std::vector(hidden_size)); + tester.AddInput("router_probs", {1, num_experts}, std::vector(num_experts)); + tester.AddInput("fc1_experts_weights", {num_experts, inter_size, 12}, + std::vector(num_experts * inter_size * 12)); + tester.AddInput("fc1_scales", fc1_scales_dims, + std::vector(static_cast(TensorShape(fc1_scales_dims).Size()), 1.0f)); + tester.AddOptionalInputEdge(); + tester.AddInput("fc2_experts_weights", {num_experts, hidden_size, 24}, + std::vector(num_experts * hidden_size * 24)); + tester.AddInput("fc2_scales", fc2_scales_dims, + std::vector(static_cast(TensorShape(fc2_scales_dims).Size()), 1.0f)); + tester.AddOptionalInputEdge(); + tester.AddOptionalInputEdge(); + tester.AddOptionalInputEdge(); + tester.AddOptionalInputEdge(); + tester.AddInput("fc1_zero_points", fc1_zero_points_dims, + std::vector(static_cast(TensorShape(fc1_zero_points_dims).Size()))); + tester.AddInput("fc2_zero_points", fc2_zero_points_dims, + std::vector(static_cast(TensorShape(fc2_zero_points_dims).Size()))); + tester.AddOutput("output", {1, hidden_size}, std::vector(hidden_size)); + + std::vector> execution_providers; + execution_providers.push_back(DefaultCpuExecutionProvider()); + tester.Run(OpTester::ExpectResult::kExpectFailure, + invalid_fc2_shape ? "Input 'fc2_zero_points' is expected to have shape" + : "Mixed-width QMoE execution is not yet implemented on CPU.", + {}, nullptr, &execution_providers); +} + +TEST(MoETest, QMoETest_MixedWidthRowWiseZeroPoints) { + RunQMoEMixedWidthZeroPointTest(false, false); + RunQMoEMixedWidthZeroPointTest(false, true); +} + +TEST(MoETest, QMoETest_MixedWidthBlockWiseZeroPoints) { + RunQMoEMixedWidthZeroPointTest(true, false); + RunQMoEMixedWidthZeroPointTest(true, true); +} + +TEST(MoETest, QMoETest_MixedWidthFusedSwiGLURequiresMatchingFC1AndFC3) { + OpTester tester("QMoE", 1, onnxruntime::kMSDomain); + tester.AddAttribute("activation_type", "swiglu"); + tester.AddAttribute("swiglu_fusion", 1); + tester.AddAttribute("expert_weight_bits", 4); + tester.AddAttribute("fc1_expert_weight_bits", 2); + tester.AddAttribute("fc3_expert_weight_bits", 4); + tester.AddInput("input", {1, 8}, std::vector(8)); + tester.AddInput("router_probs", {1, 1}, std::vector(1)); + tester.AddInput("fc1_experts_weights", {1, 16, 2}, std::vector(32)); + tester.AddInput("fc1_scales", {1, 16}, std::vector(16, 1.0f)); + tester.AddOptionalInputEdge(); + tester.AddInput("fc2_experts_weights", {1, 8, 4}, std::vector(32)); + tester.AddInput("fc2_scales", {1, 8}, std::vector(8, 1.0f)); + tester.AddOptionalInputEdge(); + tester.AddOptionalInputEdge(); + tester.AddOptionalInputEdge(); + tester.AddOptionalInputEdge(); + tester.AddOutput("output", {1, 8}, std::vector(8)); + + std::vector> execution_providers; + execution_providers.push_back(DefaultCpuExecutionProvider()); + tester.Run(OpTester::ExpectResult::kExpectFailure, + "Fused SwiGLU requires FC1 and FC3 expert weight bits to match.", + {}, nullptr, &execution_providers); +} + +TEST(MoETest, QMoETest_PackedByteCountSupportsArbitraryBitWidths) { + EXPECT_EQ(contrib::moe_helper::PackedByteCount(8, 3), 3); + EXPECT_EQ(contrib::moe_helper::PackedByteCount(8, 5), 5); + EXPECT_EQ(contrib::moe_helper::PackedByteCount(4, 6), 3); + EXPECT_EQ(contrib::moe_helper::PackedByteCountWithPadding(1, 3), 1); + EXPECT_EQ(contrib::moe_helper::PackedByteCountWithPadding(3, 5), 2); } // Regression test: row-wise asymmetric 2-bit with dimensions that trigger diff --git a/onnxruntime/test/contrib_ops/sparse_attention_op_test.cc b/onnxruntime/test/contrib_ops/sparse_attention_op_test.cc index d7953442d738e..ca27c86d0d027 100644 --- a/onnxruntime/test/contrib_ops/sparse_attention_op_test.cc +++ b/onnxruntime/test/contrib_ops/sparse_attention_op_test.cc @@ -258,6 +258,42 @@ TEST(SparseAttentionTest, RejectsZeroDimBlockRowIndices) { {}, nullptr, &execution_providers); } +// This test exercises shape inference which uses fail_shape_inference (throws InferenceError). +// In no-exception builds, fail_shape_inference calls abort(), so this test must be skipped. +#ifndef ORT_NO_EXCEPTIONS +TEST(SparseAttentionTest, RejectsEmptyTotalSequenceLengthInitializer) { + OpTester test("SparseAttention", 1, onnxruntime::kMSDomain); + test.AddAttribute("num_heads", 2); + test.AddAttribute("kv_num_heads", 2); + test.AddAttribute("sparse_block_size", 1); + test.AddAttribute("scale", 1.0f); + test.AddAttribute("do_rotary", 0); + test.AddAttribute("rotary_interleaved", 0); + + test.AddInput("query", {1, 1, 16}, std::vector(16, 0.0f)); + test.AddInput("key", {1, 1, 16}, std::vector(16, 0.0f)); + test.AddInput("value", {1, 1, 16}, std::vector(16, 0.0f)); + test.AddInput("past_key", {1, 2, 4, 8}, std::vector(64, 0.0f)); + test.AddInput("past_value", {1, 2, 4, 8}, std::vector(64, 0.0f)); + test.AddInput("block_row_indices", {1, 5}, {0, 1, 2, 3, 4}); + test.AddInput("block_col_indices", {1, 1}, std::vector{0}, /*is_initializer=*/true); + // Empty total_sequence_length initializer at input 7 must be rejected by shape inference. + test.AddInput("total_sequence_length", {0}, std::vector{}, /*is_initializer=*/true); + test.AddInput("key_total_sequence_lengths", {1}, {4}); + test.AddOptionalInputEdge(); + test.AddOptionalInputEdge(); + + test.AddOutput("output", {1, 1, 16}, std::vector(16, 0.0f)); + test.AddOutput("present_key", {1, 2, 4, 8}, std::vector(64, 0.0f)); + test.AddOutput("present_value", {1, 2, 4, 8}, std::vector(64, 0.0f)); + + std::vector> execution_providers; + execution_providers.push_back(DefaultCpuExecutionProvider()); + test.Run(OpTester::ExpectResult::kExpectFailure, + "total_sequence_length input must contain a single element", {}, nullptr, &execution_providers); +} +#endif // !ORT_NO_EXCEPTIONS + // Helper for CSR value-validation tests. // Uses: num_heads=2, kv_num_heads=2, sparse_block_size=16, head_size=8. // block_row_indices shape: (1, max_blocks+1), block_col_indices shape: (1, col_count). diff --git a/onnxruntime/test/providers/cpu/math/matmul_test.cc b/onnxruntime/test/providers/cpu/math/matmul_test.cc index e6fe606ad5ed9..a239e5755292e 100644 --- a/onnxruntime/test/providers/cpu/math/matmul_test.cc +++ b/onnxruntime/test/providers/cpu/math/matmul_test.cc @@ -9,10 +9,6 @@ #include "test/common/tensor_op_test_utils.h" #include "default_providers.h" -#if defined(USE_WEBGPU) -#include "core/providers/webgpu/math/subgroup_matrix_config.h" -#endif - namespace onnxruntime { namespace test { @@ -769,19 +765,6 @@ TEST(MathOpTest, MatMulBatchedSplitK) { } #if defined(USE_WEBGPU) -TEST(SubgroupMatrixConfigTest, RequiredSubgroupSizeCompatibility) { - using webgpu::IsSubgroupSizeSupported; - - EXPECT_TRUE(IsSubgroupSizeSupported(32, 32, 32, false)); // NVIDIA and Apple fixed-size adapters - EXPECT_TRUE(IsSubgroupSizeSupported(32, 64, 32, true)); // AMD variable-size adapter - EXPECT_TRUE(IsSubgroupSizeSupported(16, 32, 32, true)); // Intel variable-size adapter - - EXPECT_FALSE(IsSubgroupSizeSupported(32, 64, 32, false)); - EXPECT_FALSE(IsSubgroupSizeSupported(64, 64, 32, true)); - EXPECT_FALSE(IsSubgroupSizeSupported(16, 16, 32, true)); - EXPECT_FALSE(IsSubgroupSizeSupported(64, 32, 32, true)); -} - // f16 MatMul cases that exercise the Intel 8x16x16 subgroup-matrix impl. // The host picks the tile shape adaptively (TileM in {8,16,32,64}, TileN in // {16,32,64}); M and N may be any size and K must be a multiple of 16. When the diff --git a/onnxruntime/test/providers/webgpu/math/subgroup_matrix_config_test.cc b/onnxruntime/test/providers/webgpu/math/subgroup_matrix_config_test.cc new file mode 100644 index 0000000000000..48b7d776f5b19 --- /dev/null +++ b/onnxruntime/test/providers/webgpu/math/subgroup_matrix_config_test.cc @@ -0,0 +1,67 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include +#include +#include + +#include "gtest/gtest.h" + +#include "core/common/inlined_containers.h" +#include "core/providers/webgpu/math/subgroup_matrix_config.h" + +namespace onnxruntime { +namespace test { + +TEST(SubgroupMatrixConfigTest, RequiredSubgroupSizeCompatibility) { + using webgpu::IsSubgroupSizeSupported; + + EXPECT_TRUE(IsSubgroupSizeSupported(32, 32, 32, false)); // NVIDIA and Apple fixed-size adapters + EXPECT_TRUE(IsSubgroupSizeSupported(32, 64, 32, true)); // AMD variable-size adapter + EXPECT_TRUE(IsSubgroupSizeSupported(16, 32, 32, true)); // Intel variable-size adapter + + EXPECT_FALSE(IsSubgroupSizeSupported(32, 64, 32, false)); + EXPECT_FALSE(IsSubgroupSizeSupported(64, 64, 32, true)); + EXPECT_FALSE(IsSubgroupSizeSupported(16, 16, 32, true)); + EXPECT_FALSE(IsSubgroupSizeSupported(64, 32, 32, true)); +} + +TEST(SubgroupMatrixConfigTest, OperationPreferenceSelectsFromAllCandidates) { + using webgpu::supported_subgroup_matrix_configs; + using webgpu::detail::SelectSubgroupMatrixConfigFromCandidates; + + const auto find_index = [](uint32_t m, uint32_t n, uint32_t k, uint32_t subgroup_size) { + for (size_t i = 0; i < supported_subgroup_matrix_configs.size(); ++i) { + const auto& config = supported_subgroup_matrix_configs[i]; + if (config.Is(m, n, k) && config.subgroupSize == subgroup_size) { + return static_cast(i); + } + } + return int32_t{-1}; + }; + + const int32_t matmul_nbits = find_index(16, 16, 16, 32); + const int32_t intel = find_index(8, 16, 16, 32); + const int32_t apple = find_index(8, 8, 8, 32); + ASSERT_GE(matmul_nbits, 0); + ASSERT_GE(intel, 0); + ASSERT_GE(apple, 0); + + // Deliberately scramble candidate order. The operation preference, rather than candidate or + // global table order, must decide which valid kernel wins. + const InlinedVector candidates{matmul_nbits, apple, intel}; + const auto prefer_intel = + SelectSubgroupMatrixConfigFromCandidates(candidates, {{8, 16, 16, 32}, {16, 16, 16, 32}}); + ASSERT_TRUE(prefer_intel.has_value()); + EXPECT_EQ(*prefer_intel, intel); + + const auto prefer_matmul_nbits = + SelectSubgroupMatrixConfigFromCandidates(candidates, {{16, 16, 16, 32}, {8, 16, 16, 32}}); + ASSERT_TRUE(prefer_matmul_nbits.has_value()); + EXPECT_EQ(*prefer_matmul_nbits, matmul_nbits); + + EXPECT_EQ(SelectSubgroupMatrixConfigFromCandidates(candidates, {{8, 8, 8, 64}}), std::nullopt); +} + +} // namespace test +} // namespace onnxruntime diff --git a/onnxruntime/test/python/transformers/test_paged_attention.py b/onnxruntime/test/python/transformers/test_paged_attention.py index da3deed76d582..01595dbde8141 100644 --- a/onnxruntime/test/python/transformers/test_paged_attention.py +++ b/onnxruntime/test/python/transformers/test_paged_attention.py @@ -1553,21 +1553,44 @@ def test_non_causal_with_rotary_and_packed(self): config = self._config(is_causal=False, rotary=True, packed=True) parity_check_paged_attention(config, rtol=5e-3, atol=5e-3) + @parameterized.expand( + [ + (f"page{block_size}_{backend}_{step}_{mask}", block_size, sdpa_kernel, cached, local) + for block_size, backend, sdpa_kernel in [ + (16, "mea", SDPA_KERNEL_EFFICIENT_ATTENTION), + (32, "mea", SDPA_KERNEL_EFFICIENT_ATTENTION), + (64, "mea", SDPA_KERNEL_EFFICIENT_ATTENTION), + (64, "default", 0), + ] + for step, cached in [("prefill", False), ("cached", True)] + for mask, local in [("full", False), ("local", True)] + ] + ) @unittest.skipIf( not has_memory_efficient_attention(), reason="MemoryEfficientAttention (fp16) requires sm>=53", ) - def test_non_causal_rejected_without_flash_attention(self): - # The CUTLASS and paged-decode kernels hard-code a causal mask, so asking for is_causal=0 - # on a backend that cannot express it has to fail loudly instead of returning a causal result. - with self.assertRaises(Exception) as ctx: - parity_check_paged_attention( - self._config(is_causal=False), - rtol=5e-3, - atol=5e-3, - sdpa_kernel=SDPA_KERNEL_EFFICIENT_ATTENTION, - ) - self.assertIn("is_causal=0 requires the FlashAttention backend", str(ctx.exception)) + def test_non_causal_mea(self, _, block_size, sdpa_kernel, cached, local): + # Multiple query tiles must retain the left bound at past + query_pos - W + 1, + # while every row can still see the rest of its sequence on the right. + config = self._config( + sequence_length=137, + total_sequence_length=384, + num_heads=4, + head_size=128, + paged_kv_block_size=block_size, + is_causal=False, + local=local, + ) + parity_check_paged_attention( + config, + rtol=5e-3, + atol=5e-3, + sdpa_kernel=sdpa_kernel, + new_seqlens_override=torch.tensor([137, 67, 0, 9], dtype=torch.int32), + past_seqlens_override=torch.tensor([97, 31, 0, 15] if cached else [0, 0, 0, 0], dtype=torch.int32), + local_window_size_override=4 if local else None, + ) # ----------------------------------------------------------------------------- @@ -2188,6 +2211,44 @@ def test_decode_ragged_multi_split(self): new_seqlens_override=new_seqlens, ) + @parameterized.expand( + [ + (f"{name}_page{block_size}", new_seqlens, total_length, local, score_transforms, block_size) + for name, new_seqlens, total_length, local, score_transforms in [ + ("ragged", [3, 0, 0, 1], 128, False, False), + ("sparse_local", [3, 0, 0, 0], 128, True, False), + ("multi_split", [2, 0], 4096, False, False), + ("multi_split_local_sink_softcap", [2, 0], 4096, True, True), + ] + for block_size in [16, 32, 64] + ] + ) + def test_decode_non_causal_ragged(self, _, new_seqlens, total_length, local, score_transforms, block_size): + # token_count <= batch_size selects portable decode despite multiple queries in a sequence. + config = self._config( + batch_size=len(new_seqlens), + sequence_length=max(new_seqlens), + total_sequence_length=total_length, + num_heads=2, + kv_num_heads=1, + head_size=128, + paged_kv_block_size=block_size, + is_causal=False, + local=local, + use_head_sink=score_transforms, + softcap=2.0 if score_transforms else 0.0, + ) + past_seqlens = [total_length - config.sequence_length - 17 * b for b in range(config.batch_size)] + parity_check_paged_attention( + config, + rtol=5e-3, + atol=5e-3, + sdpa_kernel=SDPA_KERNEL_DECODER_ATTENTION, + new_seqlens_override=torch.tensor(new_seqlens, dtype=torch.int32), + past_seqlens_override=torch.tensor(past_seqlens, dtype=torch.int32), + local_window_size_override=4 if local else None, + ) + def test_decode_ragged_int8_cache(self): # Auto-selected: a quantized decode-shaped step. XQA cannot serve it (its output layout is # one row per batch index), so it must fall through to this kernel and still be correct. @@ -2309,6 +2370,7 @@ def _check_xqa( k_scale_max_override=None, expect_xqa=None, per_channel_xqa=None, + past_seqlens_override=None, **overrides, ): if kv_cache_type == "fp8": @@ -2325,7 +2387,13 @@ def _check_xqa( ) def run(): - parity_check_paged_attention(config, rtol=rtol, atol=atol, k_scale_max_override=k_scale_max_override) + parity_check_paged_attention( + config, + rtol=rtol, + atol=atol, + k_scale_max_override=k_scale_max_override, + past_seqlens_override=past_seqlens_override, + ) if expect_xqa is None: run() @@ -2509,6 +2577,30 @@ def test_xqa_context_not_page_aligned(self): # The live length is not a multiple of 128, so the last page is partially valid. self._check_xqa(total_sequence_length=1000) + @parameterized.expand( + [ + (f"{cache}_{mask}", cache, local) + for cache in ["float16", "int8"] + for mask, local in [("full", False), ("local", True)] + ] + ) + def test_xqa_non_causal(self, _, kv_cache_type, local): + with patch.dict(os.environ, {"ORT_ENABLE_XQA_NATIVE_KV": "1"}): + self._check_xqa( + kv_cache_type=kv_cache_type, + quant_type="NONE" if kv_cache_type == "float16" else "PER_TENSOR", + num_heads=6, + kv_num_heads=1, + head_size=256, + total_sequence_length=128, + is_causal=False, + local=local, + local_window_size=4, + use_attention_metadata=True, + past_seqlens_override=torch.tensor([17, 29, 7, 63], dtype=torch.int32), + expect_xqa=True, + ) + # ---- quantization granularity ------------------------------------------------------- @parameterized.expand( @@ -2615,6 +2707,7 @@ def _check( atol=5e-3, expect_xqa=None, local_window_size=None, + past_seqlens=None, **overrides, ): kwargs = { @@ -2653,6 +2746,9 @@ def run(): atol=atol, new_seqlens_override=override, local_window_size_override=local_window_size, + past_seqlens_override=torch.tensor(past_seqlens, dtype=torch.int32) + if past_seqlens is not None + else None, ) # Output parity alone would still pass if a dispatch regression routed the case to the @@ -2737,6 +2833,30 @@ def test_spec_dec_ragged_batch(self, _, new_seqlens): # contribute no token at all. self._check(new_seqlens=new_seqlens, expect_xqa=True) + @parameterized.expand( + [ + (f"{cache}_{mask}", cache, local) + for cache in ["float16", "int8"] + for mask, local in [("full", False), ("local", True)] + ] + ) + def test_spec_dec_non_causal_ragged(self, _, kv_cache_type, local): + # Eight tokens at group size six cross XQA's 32-row tile; W=4 must not mask future tokens. + quant_type = "NONE" if kv_cache_type == "float16" else "PER_TENSOR" + with patch.dict(os.environ, {"ORT_ENABLE_XQA_NATIVE_KV": "1"}): + self._check( + kv_cache_type=kv_cache_type, + k_quant_type=quant_type, + v_quant_type=quant_type, + total_sequence_length=128, + new_seqlens=[1, 8, 0, 3], + past_seqlens=[17, 29, 7, 0], + is_causal=False, + local=local, + local_window_size=4 if local else None, + expect_xqa=True, + ) + @parameterized.expand([("64", 64), ("256", 256), ("1000", 1000), ("8192", 8192)]) def test_spec_dec_context_length(self, _, total_sequence_length): # 8192 splits the sequence across CTAs and reduces through the XQA scratch; 1000 is not @@ -2953,6 +3073,7 @@ def __init__( qk_nope_head_dim=None, kv_cache_type="float16", k_quant_type="NONE", + is_causal=True, ): self.batch_size = batch_size self.num_heads = num_heads @@ -2968,6 +3089,7 @@ def __init__( self.qk_nope_head_dim = qk_nope_head_dim self.kv_cache_type = kv_cache_type self.k_quant_type = k_quant_type + self.is_causal = is_causal @property def softmax_scale(self): @@ -3015,6 +3137,8 @@ def create_mla_graph( } if scale is not None: attrs["scale"] = scale + if not mla_config.is_causal: + attrs["is_causal"] = 0 if do_rotary: attrs["do_rotary"] = 1 attrs["rotary_interleaved"] = 1 if rotary_interleaved else 0 @@ -3171,8 +3295,9 @@ def mla_reference( for b in range(mla_config.batch_size): start = int(cum_seqlens[b].item()) for j in range(int(new_seqlens[b].item())): - kv_end = int(past_seqlens[b].item()) + j + 1 - kv_begin = max(0, kv_end - local_window_size) if local_window_size > 0 else 0 + query_end = int(past_seqlens[b].item()) + j + 1 + kv_end = query_end if mla_config.is_causal else int((past_seqlens[b] + new_seqlens[b]).item()) + kv_begin = max(0, query_end - local_window_size) if local_window_size > 0 else 0 k = latent_cache[b, kv_begin:kv_end].to(torch.float32) # [L, head_size] v = k[:, :v_head_size] logits = torch.einsum("nh,lh->nl", q[start + j], k) * scale @@ -3312,6 +3437,21 @@ def test_chunked_prefill(self): config = self._config(batch_size=3) self._run_case(config, past_seqlens=[20, 9, 5], new_seqlens=[6, 1, 0]) + @parameterized.expand( + [ + (f"{step}_{mask}", cached, local) + for step, cached in [("prefill", False), ("cached", True)] + for mask, local in [("full", False), ("local", True)] + ] + ) + def test_non_causal_ragged(self, _, cached, local): + self._run_case( + self._config(batch_size=4, is_causal=False), + past_seqlens=[17, 129, 7, 63] if cached else [0, 0, 0, 0], + new_seqlens=[1, 8, 0, 3], + local_window_size=4 if local else -1, + ) + def test_local_window(self): config = self._config() self._run_case(config, past_seqlens=[24, 24], new_seqlens=[5, 5], local_window_size=8)