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