diff --git a/docs/design/webgpu_matmul_algorithm_scheduler.md b/docs/design/webgpu_matmul_algorithm_scheduler.md new file mode 100644 index 0000000000000..15269933e4a86 --- /dev/null +++ b/docs/design/webgpu_matmul_algorithm_scheduler.md @@ -0,0 +1,87 @@ +# WebGPU MatMul Algorithm Scheduler Design + +## Goal + +Make every `MatMulComputeDispatcher` implementation path explicit and independently testable. Selection policy must be separated from execution, retain the existing common rules as a fallback, and allow vendor-specific policy to override any performance heuristic. + +## Algorithms + +Introduce `MatMulAlgorithm` with five concrete values: + +- `SubgroupMatrix`: the common subgroup-matrix implementation in `subgroup_matrix_matmul.cc`. +- `Naive`: `MatMulNaiveProgram`. +- `Subgroup`: the common subgroup implementation currently located in `vendor/intel/math/matmul.cc`. +- `Packed`: the generic packed `MatMulProgram` without Split-K. +- `PackedSplitK`: the same packed program with Split-K initialization and atomic accumulation. + +There is no `Auto` algorithm value. Automatic versus forced selection is represented by an optional forced algorithm, which prevents an unset policy state from being confused with an executable implementation. + +## Selection Architecture + +Add a `MatMulAlgorithmScheduler` base class. Automatic selection uses this order: + +1. Select `Naive` when `K == 0`. This is a correctness rule and cannot be overridden by vendor policy. +2. Ask the virtual `SelectVendorAlgorithm` policy hook for any algorithm. A vendor may use its own thresholds for every algorithm, not only vendor-specific implementations. +3. If the vendor returns no selection, call the private, non-virtual `SelectCommonAlgorithm` fallback. It selects `SubgroupMatrix` when supported, `Naive` when `N < 8 && K < 8`, `PackedSplitK` when the existing `SplitKConfig::UseSplitK` rule succeeds, and otherwise `Packed`. + +An Intel-derived scheduler implements the vendor hook. It preserves the original policy by selecting `SubgroupMatrix` first when applicable, then selecting the common `Subgroup` implementation under Intel-specific shape thresholds. Other vendors use the base scheduler unchanged. Future vendor policies can derive from the base scheduler, override as many performance ranges as needed, and return no selection to delegate the remaining ranges to the common fallback without adding vendor conditionals to `MatMulComputeDispatcher`. + +The subgroup implementation is not inherently Intel-specific: forcing it requires the WebGPU `Subgroups` feature and a supported subgroup size. Its current source location under `vendor/intel` is historical and should move to the common math directory in a follow-up, while Intel-specific automatic-selection thresholds and tuning remain under `vendor/intel`. + +The subgroup shader implements subgroup-size branches for 8, 16, and 32 lanes. A fixed-size adapter may use any of those sizes directly. An adapter reporting a size range must expose subgroup-size control so the execution plan can select a supported size explicitly; otherwise the subgroup algorithm is unavailable. + +The scheduler accepts immutable problem facts rather than mutable policy decisions: logical and packed dimensions, batch sizes, adapter architecture, input types, packing and layout facts, deterministic-compute state, activation and bias facts, and device capabilities. It also owns a copy of the immutable `SplitKConfig` selected for the adapter. The private common fallback evaluates `SplitKConfig::UseSplitK` directly from the problem facts instead of receiving a precomputed decision from `MatMulComputeDispatcher`. A vendor may ignore that common recommendation and apply independent thresholds from the raw facts. This keeps the scheduler deterministic and unit-testable without a WebGPU device. + +`SplitKConfig` contains only generic Split-K eligibility evaluation and data. Adapter routing is handled by a small generic factory, while Intel architecture profiles and their measured threshold tables live under `vendor/intel`. `WebGpuContext` owns the selected configuration so GEMM and MatMul use the same profile; the MatMul scheduler receives that configuration when it is created. A future vendor can add its own profile builder and factory route without adding conditions to `MatMulComputeDispatcher` or changing the generic evaluator. + +## Execution Configuration + +Algorithm selection and execution tuning are separate decisions. After selecting one enum, the scheduler creates a `MatMulExecutionPlan` containing that enum and a typed algorithm configuration. It first asks the protected `SelectVendorConfiguration` hook for tuning, then uses private common defaults when the vendor declines. The tuning hook runs for both automatic and forced algorithms, so the test-only forcing option controls the implementation path without disabling real device tuning. + +The packed configuration initially contains workgroup size, elements per thread, inner tile size, and Split-K size. `ApplyMatMulPacked` consumes those values directly and includes shader-affecting values in its cache key. The common configuration preserves the existing `8x8x1` workgroup, `4x1x1` or `4x4x1` elements-per-thread rule, inner tile size 32, and adapter Split-K size. A vendor may replace any of these values without changing `MatMulComputeDispatcher` or the packed implementation. + +Packed tuning is validated at the execution boundary. The current shader requires both Z workgroup dimensions to remain one because Z identifies a batch or Split-K slice rather than a tiled output axis. Split-K sizes must be greater than one and aligned to the inner tile so adjacent workgroups cannot overlap their K ranges. Dispatch counts use widened, overflow-safe arithmetic and are range-checked before conversion to WebGPU's 32-bit dimensions. + +Configuration is represented by an algorithm-specific variant rather than a bag of unrelated optional fields. The dispatcher rejects an algorithm/configuration type mismatch before execution. Empty configuration types reserve the same typed boundary for algorithms whose tuning remains inside their existing implementation. In particular, subgroup-matrix MatMul already receives a vendor-specific `SubgroupMatrixTilingSelector`; that existing selector continues to choose tile M, tile N, and split K without coupling the scheduler to device or shader classes. + +## Forced Test Selection + +Add the internal WebGPU session configuration key `ep.webgpuexecutionprovider.forceMatmulAlgorithm`. Accepted values are `subgroup_matrix`, `naive`, `subgroup`, `packed`, and `packed_split_k`. The option is parsed when the WebGPU EP is created, stored as `std::optional`, and exposed read-only through `ComputeContextBase`. + +When set, the scheduler returns the requested enum before applying heuristic rules. The dispatcher then validates the algorithm's hard prerequisites. Unsupported device features, data types, layouts, deterministic-compute settings, or other correctness constraints produce a descriptive failure naming the forced algorithm; forced mode never silently falls back. + +Heuristic thresholds are not hard prerequisites. For example, forcing subgroup bypasses Intel's current `M/N/K` performance thresholds while still requiring subgroup support. Forcing Split-K bypasses performance thresholds while still requiring a usable Split-K configuration, non-deterministic compute, compatible packing/activation, and supported bias layout. + +Invalid option strings fail during WebGPU provider creation and list accepted values. + +## Dispatch and Implementation Boundaries + +Introduce `MatMulComputeDispatcher` as the single compute entry point below the MatMul, pointwise Conv, and contrib Attention kernels. Each kernel owns one dispatcher for its lifetime. The dispatcher lazily creates adapter-dependent state from the first compute context and then owns: + +- one `MatMulAlgorithmScheduler`, which contains selection policy and immutable adapter tuning data; and +- an optional `SubgroupMatrixMatMulImpl`, which contains only the subgroup-matrix implementation's persistent device state, including cached padded constant weights. + +The dispatcher does not store a current or previously selected algorithm. Selection occurs for every invocation because shapes, input properties, activation, and bias can differ between calls. The session's forced-test configuration, when present, participates in each selection without becoming mutable dispatcher state. The resulting `MatMulExecutionPlan` is local to that invocation. + +For each call, `MatMulComputeDispatcher::Compute` computes shared shape and capability facts once, asks the scheduler for one execution plan, validates the selected algorithm's hard prerequisites and configuration type, and switches directly to the matching implementation. MatMul, pointwise Conv, and contrib Attention delegate through this same entry point rather than coordinating the scheduler and implementations themselves. + +The scheduler remains pure policy: it creates plans but does not create device programs, cache tensor data, or execute kernels. `SubgroupMatrixMatMulImpl` is named for the implementation it owns and is reached only when the plan selects `SubgroupMatrix`; its applicability query is non-mutating, and its execution method does not communicate selection through a `handled` output. This removes trial execution as a dispatch mechanism. + +The other implementations remain focused stateless functions or program builders unless they acquire persistent state in the future. The naive, subgroup, packed, and packed Split-K paths therefore do not receive empty polymorphic wrapper classes. The generic packed helper takes an explicit Split-K mode and packed configuration; it does not re-run selection or tuning heuristics. WebGPU's existing program cache continues to own reusable compiled programs. + +This replaces `MatMulOptImplCache` and the generic `MatMulOptImpl` interface. If another algorithm later needs persistent implementation state, the dispatcher can own a separately named backend for that algorithm without changing scheduler policy or pretending that one object represents whichever algorithm happened to be selected most recently. + +## Compatibility + +With no forcing option or vendor override, the common scheduler remains equivalent to the previous selection order. The MatMul, pointwise Conv, and contrib Attention operator APIs and model semantics do not change; only their internal MatMul compute ownership moves behind the dispatcher. + +The option is intentionally internal and test-only: it is declared with WebGPU provider options for configuration plumbing but is not added to public user documentation. + +## Testing + +- Add device-independent scheduler unit tests covering every common branch, forced-over-vendor precedence, the zero-K correctness guard, vendor-over-common precedence, Intel override, default fallback, independent vendor thresholds, common packed defaults, vendor tuning of a forced algorithm, Split-K profile routing, and current Intel architecture boundaries. +- Add parser/configuration tests for every accepted value and invalid input. +- Add WebGPU MatMul tests that choose shapes which normally select a different path, force a compatible algorithm, and verify numerical output. Hardware-specific forced algorithms are tested only when their hard capabilities are present; strict-failure tests cover unsupported forced choices. +- Run hardware-backed tests on the build's default WebGPU backend. Hardware-specific tests inspect the selected adapter's capabilities and skip unsupported algorithms. +- Verify the selected adapter exposes subgroup size control, f16, and the cooperative/subgroup-matrix configuration required by the 8x16x16 kernel before claiming subgroup-matrix execution coverage. +- Run scheduler/parser tests and hardware-backed MatMul tests on the local Intel Arc Vulkan adapter. macOS-arm64 Metal CI remains additional cross-backend coverage; lavapipe is not used because it cannot execute MatMul reliably. diff --git a/onnxruntime/contrib_ops/webgpu/bert/attention.cc b/onnxruntime/contrib_ops/webgpu/bert/attention.cc index 9a6f6062abaa5..7ba3075c6c9fb 100644 --- a/onnxruntime/contrib_ops/webgpu/bert/attention.cc +++ b/onnxruntime/contrib_ops/webgpu/bert/attention.cc @@ -9,7 +9,6 @@ #include "contrib_ops/webgpu/webgpu_contrib_kernels.h" #include "core/providers/webgpu/webgpu_supported_types.h" #include "core/providers/webgpu/webgpu_utils.h" -#include "core/providers/webgpu/math/matmul.h" using namespace onnxruntime::webgpu; using namespace ::onnxruntime::common; using namespace ONNX_NAMESPACE; @@ -659,7 +658,7 @@ Attention::Attention(const OpKernelInfo& info) Status PrepareQKV(onnxruntime::webgpu::ComputeContext& context, const WebgpuAttentionParameters& parameters, const Tensor* input, const Tensor* weights, const Tensor* bias, Tensor* q, Tensor* k, Tensor* v, - MatMulOptImplCache& matmul_compute_cache, bool weights_are_constant) { + MatMulComputeDispatcher& matmul_compute_dispatcher, bool weights_are_constant) { // Use MatMul to compute packed QKV output: input * weights + bias // Then use SplitPackedQKV to split into Q, K, V in BSD format // Returns Q, K, V in BSD format @@ -673,9 +672,9 @@ Status PrepareQKV(onnxruntime::webgpu::ComputeContext& context, const WebgpuAtte std::vector matmul_inputs = {input, weights, bias}; // Call MatMul: packed_qkv = input * weights + bias - ORT_RETURN_IF_ERROR(onnxruntime::webgpu::ComputeMatMul( - &context, Activation(), matmul_inputs, &packed_qkv, /*is_channels_last=*/true, - matmul_compute_cache, weights_are_constant)); + ORT_RETURN_IF_ERROR(matmul_compute_dispatcher.Compute( + context, Activation(), matmul_inputs, &packed_qkv, /*is_channels_last=*/true, + weights_are_constant)); // Output Q, K, V in BSD format return SplitPackedQKV(context, parameters, &packed_qkv, q, k, v, parameters.hidden_size_); @@ -743,7 +742,7 @@ Status Attention::ComputeInternal(onnxruntime::webgpu::ComputeContext& context) // Compute Q, K, V from input, weights, and bias (returns BSD format) ORT_RETURN_IF_ERROR(PrepareQKV(context, parameters, input, weights, bias, &Q_bsd, &K_bsd, &V_bsd, - matmul_compute_cache_, weights_are_constant_)); + matmul_compute_dispatcher_, weights_are_constant_)); parameters.qkv_format_ = Q_K_V_BSNH; // Check if we can use flash attention diff --git a/onnxruntime/contrib_ops/webgpu/bert/attention.h b/onnxruntime/contrib_ops/webgpu/bert/attention.h index 30864c17e74ef..eb08876a96fdd 100644 --- a/onnxruntime/contrib_ops/webgpu/bert/attention.h +++ b/onnxruntime/contrib_ops/webgpu/bert/attention.h @@ -4,7 +4,7 @@ #pragma once #include "core/providers/webgpu/compute_context.h" -#include "core/providers/webgpu/math/matmul.h" +#include "core/providers/webgpu/math/matmul_compute_dispatcher.h" #include "core/providers/webgpu/program.h" #include "core/providers/webgpu/shader_helper.h" #include "core/providers/webgpu/webgpu_kernel.h" @@ -150,7 +150,7 @@ class Attention final : public WebGpuKernel, public onnxruntime::contrib::Attent Status ComputeInternal(onnxruntime::webgpu::ComputeContext& context) const override; private: - mutable MatMulOptImplCache matmul_compute_cache_; + mutable MatMulComputeDispatcher matmul_compute_dispatcher_; bool weights_are_constant_ = false; }; diff --git a/onnxruntime/core/providers/webgpu/compute_context.h b/onnxruntime/core/providers/webgpu/compute_context.h index e151258a63a5b..157d9b467d3c7 100644 --- a/onnxruntime/core/providers/webgpu/compute_context.h +++ b/onnxruntime/core/providers/webgpu/compute_context.h @@ -123,6 +123,10 @@ class ComputeContextBase { return ep_.EnableMatmulFp32Accumulation(); } + inline std::optional ForcedMatMulAlgorithm() const { + return ep_.ForcedMatMulAlgorithm(); + } + // // Get the logger. // diff --git a/onnxruntime/core/providers/webgpu/math/matmul.cc b/onnxruntime/core/providers/webgpu/math/matmul.cc index 65157575fbb3b..803078b1cf125 100644 --- a/onnxruntime/core/providers/webgpu/math/matmul.cc +++ b/onnxruntime/core/providers/webgpu/math/matmul.cc @@ -11,19 +11,25 @@ #include "core/providers/webgpu/webgpu_supported_types.h" #include "core/providers/webgpu/nn/fuse_utils.h" #include "core/providers/webgpu/data_transfer.h" +#include "core/providers/webgpu/vendor/intel/math/matmul_algorithm_scheduler.h" #include "core/providers/webgpu/vendor/intel/math/matmul.h" #include "core/providers/webgpu/webgpu_utils.h" namespace onnxruntime { namespace webgpu { -std::unique_ptr CreateSubgroupMatrixMatMulImpl(const ComputeContextBase& context); +std::unique_ptr CreateSubgroupMatrixMatMulImpl(const ComputeContextBase& context); -MatMulOptImpl* MatMulOptImplCache::GetOrCreate(const ComputeContextBase& context) { - std::call_once(subgroup_impl_init_flag_, [&]() { - subgroup_impl_ = CreateSubgroupMatrixMatMulImpl(context); +void MatMulComputeDispatcher::Initialize(const ComputeContextBase& context) { + std::call_once(init_flag_, [&]() { + subgroup_matrix_impl_ = CreateSubgroupMatrixMatMulImpl(context); + if (context.AdapterInfo().vendor == std::string_view{"intel"}) { + scheduler_ = std::make_unique( + context.GetSplitKConfig()); + } else { + scheduler_ = std::make_unique(context.GetSplitKConfig()); + } }); - return subgroup_impl_.get(); } ONNX_OPERATOR_VERSIONED_KERNEL_EX( @@ -61,13 +67,17 @@ static std::string CalcResult(int64_t components, int64_t a_components, int64_t } Status MatMulNaiveProgram::GenerateShaderCode(ShaderHelper& shader) const { - const auto& a = shader.AddInput("a", ShaderUsage::UseUniform | ShaderUsage::UseIndicesTypeAlias | - ShaderUsage::UseValueTypeAlias | ShaderUsage::UseElementTypeAlias); - const auto& b = shader.AddInput("b", ShaderUsage::UseUniform | ShaderUsage::UseIndicesTypeAlias | - ShaderUsage::UseValueTypeAlias | ShaderUsage::UseElementTypeAlias); - - const int a_components = a.NumComponents(); - const int components = b.NumComponents(); // components of N + const ShaderVariableHelper* a = nullptr; + const ShaderVariableHelper* b = nullptr; + int a_components = 1; + if (!is_zero_k_) { + a = &shader.AddInput("a", ShaderUsage::UseUniform | ShaderUsage::UseIndicesTypeAlias | + ShaderUsage::UseValueTypeAlias | ShaderUsage::UseElementTypeAlias); + b = &shader.AddInput("b", ShaderUsage::UseUniform | ShaderUsage::UseIndicesTypeAlias | + ShaderUsage::UseValueTypeAlias | ShaderUsage::UseElementTypeAlias); + a_components = a->NumComponents(); + } + const int components = NumberOfComponents(Outputs()[0].var_type); // components of N std::string process_bias; if (has_bias_) { @@ -81,32 +91,34 @@ Status MatMulNaiveProgram::GenerateShaderCode(ShaderHelper& shader) const { const auto& output = shader.AddOutput("output", ShaderUsage::UseUniform | ShaderUsage::UseIndicesTypeAlias | ShaderUsage::UseValueTypeAlias | ShaderUsage::UseElementTypeAlias); shader.AdditionalImplementation() << GetActivationDeclaration(activation_, "output_value_t", "output_element_t"); - const auto& batch_dims = shader.AddIndices("batch_dims"); shader.MainFunctionBody() << shader.GuardAgainstOutOfBoundsWorkgroupSizes("uniforms.output_size") << "let col = (global_idx % (uniforms.N / " << components << ")) * " << components << ";\n" << "var index1 = global_idx / (uniforms.N / " << components << ");\n" << "let stride1 = uniforms.M / " << output_number_ << ";\n" << "let row = (index1 % stride1) * " << output_number_ << ";\n" - << "let batch = index1 / stride1;\n"; - if (output_rank_ != 2) { - shader.MainFunctionBody() << "let batch_indices = " << batch_dims.OffsetToIndices("batch") << ";\n"; + << "let batch = index1 / stride1;\n" + << "var values: array;\n"; + if (!is_zero_k_) { + const auto& batch_dims = shader.AddIndices("batch_dims"); + if (output_rank_ != 2) { + shader.MainFunctionBody() << "let batch_indices = " << batch_dims.OffsetToIndices("batch") << ";\n"; + } + shader.MainFunctionBody() << "var a_indices: a_indices_t;\n" + << ConvertOutputBatchIndicesToInputBatchIndices("a", *a, a->Rank() - 2, batch_dims.Rank(), "batch_indices") + << a->IndicesSet("a_indices", a->Rank() - 2, 0) << "\n" + << a->IndicesSet("a_indices", a->Rank() - 1, 0) << "\n" + << "let a_offset = " << a->IndicesToOffset("a_indices") << "*" << a_components << ";\n" + << "var b_indices: b_indices_t;\n" + << ConvertOutputBatchIndicesToInputBatchIndices("b", *b, b->Rank() - 2, batch_dims.Rank(), "batch_indices") + << b->IndicesSet("b_indices", b->Rank() - 2, 0) << "\n" + << b->IndicesSet("b_indices", b->Rank() - 1, 0) << "\n" + << "let b_offset = " << b->IndicesToOffset("b_indices") << " * " << components << ";\n" + << "for (var k: u32 = 0u; k < uniforms.K; k = k + " << a_components << ") {\n" + << CalcResult(components, a_components, output_number_) << "\n" + << "}\n"; } - shader.MainFunctionBody() << "var a_indices: a_indices_t;\n" - << ConvertOutputBatchIndicesToInputBatchIndices("a", a, a.Rank() - 2, batch_dims.Rank(), "batch_indices") - << a.IndicesSet("a_indices", a.Rank() - 2, 0) << "\n" - << a.IndicesSet("a_indices", a.Rank() - 1, 0) << "\n" - << "let a_offset = " << a.IndicesToOffset("a_indices") << "*" << a_components << ";\n" - << "var b_indices: b_indices_t;\n" - << ConvertOutputBatchIndicesToInputBatchIndices("b", b, b.Rank() - 2, batch_dims.Rank(), "batch_indices") - << b.IndicesSet("b_indices", b.Rank() - 2, 0) << "\n" - << b.IndicesSet("b_indices", b.Rank() - 1, 0) << "\n" - << "let b_offset = " << b.IndicesToOffset("b_indices") << " * " << components << ";\n" - << "var values: array;\n" - << "for (var k: u32 = 0u; k < uniforms.K; k = k + " << a_components << ") {\n" - << CalcResult(components, a_components, output_number_) << "\n" - << "}\n" - << "for (var i = 0u; i < " << output_number_ << "u; i++) {\n" + shader.MainFunctionBody() << "for (var i = 0u; i < " << output_number_ << "u; i++) {\n" << " var value = values[i];\n" << process_bias << "\n" << apply_activation << "\n" @@ -152,195 +164,302 @@ Status MatMul::ComputeInternal(ComputeContext& context) const { inputs[1] = &promoted_b; } - return ComputeMatMul(&context, Activation(), inputs, output_tensor, - /*is_channels_last=*/true, compute_cache_, b_is_constant_); + return compute_dispatcher_.Compute(context, Activation(), inputs, output_tensor, + /*is_channels_last=*/true, b_is_constant_); } -Status ComputeMatMul(ComputeContext* context, - const Activation& activation, std::vector& inputs, Tensor* output_tensor, bool is_channels_last, - MatMulOptImplCache& cache, - bool b_is_constant) { +static Status ApplyMatMulNaive(ComputeContext& context, + const Activation& activation, + const std::vector& inputs, + Tensor* output_tensor, + bool is_channels_last, + const MatMulComputeHelper& helper) { const auto* a = inputs[0]; const auto* b = inputs[1]; - bool has_bias = inputs.size() > 2; - const TensorShape& logical_a_shape = a->Shape(); - const TensorShape& logical_b_shape = b->Shape(); - ORT_RETURN_IF_NOT(logical_a_shape.NumDimensions() >= 2 && logical_b_shape.NumDimensions() >= 2, - "ComputeMatMul expects matrix or batched-matrix inputs."); - - MatMulComputeHelper helper; - ORT_THROW_IF_ERROR(helper.Compute(logical_a_shape, logical_b_shape)); - - MatMulOptImpl* subgroup_impl = cache.GetOrCreate(*context); - if (subgroup_impl != nullptr) { - bool handled = false; - ORT_RETURN_IF_ERROR(subgroup_impl->Compute( - *context, inputs, output_tensor, - activation, is_channels_last, b_is_constant, handled)); - if (handled) { - return Status::OK(); - } - } - if (helper.N() < 8 && helper.K() < 8) { - const uint32_t m = narrow(helper.M()); - const uint32_t n = narrow(helper.N()); - const uint32_t k = narrow(helper.K()); - const int components = GetMaxComponents(n); - const int a_components = GetMaxComponents(k); - const int64_t output_number = GetMaxComponents(m); - const TensorShape& logical_output_shape = helper.OutputShape(); - const size_t output_rank = logical_output_shape.NumDimensions(); - const TensorShape outer_dims = - output_rank > 2 ? logical_output_shape.Slice(0, output_rank - 2) : TensorShape({}); - const int64_t output_rows = logical_a_shape[logical_a_shape.NumDimensions() - 2]; - const TensorShape output_program_shape{ - outer_dims.Size(), output_rows, n / components}; - const uint32_t output_size = - narrow(logical_output_shape.Size() / components / output_number); - - MatMulNaiveProgram program{activation, output_rank, output_number, has_bias, is_channels_last}; + const bool has_bias = inputs.size() > 2; + const uint32_t m = narrow(helper.M()); + const uint32_t n = narrow(helper.N()); + const uint32_t k = narrow(helper.K()); + const bool is_zero_k = k == 0; + const int components = GetMaxComponents(n); + const int a_components = is_zero_k ? 1 : GetMaxComponents(k); + const int64_t output_number = GetMaxComponents(m); + const TensorShape& logical_output_shape = helper.OutputShape(); + const size_t output_rank = logical_output_shape.NumDimensions(); + const TensorShape outer_dims = + output_rank > 2 ? logical_output_shape.Slice(0, output_rank - 2) : TensorShape({}); + const int64_t output_rows = a->Shape()[a->Shape().NumDimensions() - 2]; + const TensorShape output_program_shape{ + outer_dims.Size(), output_rows, n / components}; + const uint32_t output_size = + narrow(logical_output_shape.Size() / components / output_number); + + MatMulNaiveProgram program{activation, output_rank, output_number, has_bias, + is_channels_last, is_zero_k}; + program + .CacheHint(activation.CacheKey(), std::to_string(components), + std::to_string(a_components), std::to_string(output_number), + std::to_string(is_channels_last), std::to_string(is_zero_k)); + if (!is_zero_k) { program - .CacheHint(activation.CacheKey(), std::to_string(components), - std::to_string(a_components), std::to_string(output_number), - std::to_string(is_channels_last)) .AddInputs({{a, ProgramTensorMetadataDependency::TypeAndRank, a_components}, - {b, ProgramTensorMetadataDependency::TypeAndRank, components}}); - if (has_bias) { - const int bias_components = is_channels_last ? components : 1; - program.AddInput({inputs[2], ProgramTensorMetadataDependency::Rank, bias_components}); - } - program - .AddOutputs({{output_tensor, ProgramTensorMetadataDependency::None, - output_program_shape, components}}) - .SetDispatchGroupSize(CeilDiv(output_size, 64u)) - .AddIndices(outer_dims) - .AddUniformVariables({{output_size}, {m}, {n}, {k}}); - AppendActivationUniformsData(activation, program); - return context->RunProgram(program); + {b, ProgramTensorMetadataDependency::TypeAndRank, components}}) + .AddIndices(outer_dims); } - - if (intel::CanApplyMatMulIntel(*context, helper.M(), helper.N(), helper.K())) { - return intel::ApplyMatMulIntel(*context, activation, inputs, output_tensor, is_channels_last); + if (has_bias) { + const int bias_components = is_channels_last ? components : 1; + program.AddInput({inputs[2], ProgramTensorMetadataDependency::Rank, bias_components}); } + program + .AddOutputs({{output_tensor, ProgramTensorMetadataDependency::None, + output_program_shape, components}}) + .SetDispatchGroupSize(CeilDiv(output_size, 64u)) + .AddUniformVariables({{output_size}, {m}, {n}, {k}}); + AppendActivationUniformsData(activation, program); + return context.RunProgram(program); +} - TensorShape a_shape = logical_a_shape; - TensorShape b_shape = logical_b_shape; +static Status ApplyMatMulPacked(ComputeContext& context, + const Activation& activation, + const std::vector& inputs, + Tensor* output_tensor, + bool is_channels_last, + const MatMulComputeHelper& helper, + const MatMulPackedConfiguration& configuration, + bool use_split_k) { + const auto* a = inputs[0]; + const auto* b = inputs[1]; + const bool has_bias = inputs.size() > 2; + TensorShape a_shape = a->Shape(); + TensorShape b_shape = b->Shape(); TensorShape output_shape = helper.OutputShape(); - const int64_t batchA = + const int64_t batch_a = a_shape.NumDimensions() > 2 ? a_shape.SizeToDimension(a_shape.NumDimensions() - 2) : 1; - const int64_t batchB = + const int64_t batch_b = b_shape.NumDimensions() > 2 ? b_shape.SizeToDimension(b_shape.NumDimensions() - 2) : 1; - // The generic path benefits from folding A's batch dimensions into M when B - // is shared. The subgroup and Intel paths derive their own dispatch shapes - // directly from the tensor views and have already declined above. - if (batchA != 1 && batchB == 1) { - const int64_t batchAndM = a_shape.SizeToDimension(a_shape.NumDimensions() - 1); - a_shape = TensorShape({batchAndM, helper.K()}); + if (batch_a != 1 && batch_b == 1) { + const int64_t batch_and_m = a_shape.SizeToDimension(a_shape.NumDimensions() - 1); + a_shape = TensorShape({batch_and_m, helper.K()}); b_shape = TensorShape({helper.K(), helper.N()}); - output_shape = TensorShape({batchAndM, helper.N()}); + output_shape = TensorShape({batch_and_m, helper.N()}); } - // helpful dimension variables - TensorShape outer_dims_a = a_shape.NumDimensions() > 2 - ? a_shape.Slice(0, a_shape.NumDimensions() - 2) - : TensorShape({}); - - TensorShape outer_dims_b = b_shape.NumDimensions() > 2 - ? b_shape.Slice(0, b_shape.NumDimensions() - 2) - : TensorShape({}); - - TensorShape outer_dims = output_shape.NumDimensions() > 2 - ? output_shape.Slice(0, output_shape.NumDimensions() - 2) - : TensorShape({}); - + const TensorShape outer_dims_a = a_shape.NumDimensions() > 2 + ? a_shape.Slice(0, a_shape.NumDimensions() - 2) + : TensorShape({}); + const TensorShape outer_dims_b = b_shape.NumDimensions() > 2 + ? b_shape.Slice(0, b_shape.NumDimensions() - 2) + : TensorShape({}); + const TensorShape outer_dims = output_shape.NumDimensions() > 2 + ? output_shape.Slice(0, output_shape.NumDimensions() - 2) + : TensorShape({}); const int64_t batch_size = outer_dims.Size(); - - // Get dimensions for matrix multiplication from TensorShape - const uint32_t dim_a_outer = narrow(a_shape[a_shape.NumDimensions() - 2]); // left matrix second dimension - const uint32_t dim_inner = narrow(a_shape[a_shape.NumDimensions() - 1]); // left matrix first dimension - const uint32_t dim_b_outer = narrow(b_shape[b_shape.NumDimensions() - 1]); // right matrix first dimension - + const uint32_t dim_a_outer = narrow(a_shape[a_shape.NumDimensions() - 2]); + const uint32_t dim_inner = narrow(a_shape[a_shape.NumDimensions() - 1]); + const uint32_t dim_b_outer = narrow(b_shape[b_shape.NumDimensions() - 1]); const bool is_vec4 = dim_inner % 4 == 0 && dim_b_outer % 4 == 0; - InlinedVector elements_per_thread = dim_a_outer <= 8 - ? InlinedVector({4, 1, 1}) - : InlinedVector({4, 4, 1}); - - const uint32_t dispatch_x = narrow((dim_b_outer + MatMul::MATMUL_PACKED_WORKGROUP_SIZE_X * elements_per_thread[0] - 1) / - (MatMul::MATMUL_PACKED_WORKGROUP_SIZE_X * elements_per_thread[0])); - const uint32_t dispatch_y = narrow((dim_a_outer + MatMul::MATMUL_PACKED_WORKGROUP_SIZE_Y * elements_per_thread[1] - 1) / - (MatMul::MATMUL_PACKED_WORKGROUP_SIZE_Y * elements_per_thread[1])); - uint32_t dispatch_z = narrow((static_cast(batch_size) + MatMul::MATMUL_PACKED_WORKGROUP_SIZE_Z * elements_per_thread[2] - 1) / - (MatMul::MATMUL_PACKED_WORKGROUP_SIZE_Z * elements_per_thread[2])); + ORT_RETURN_IF_NOT(IsMatMulPackedConfigurationValid(configuration, use_split_k), + "MatMul packed configuration is invalid for ", + use_split_k ? "Split-K." : "the packed algorithm."); + InlinedVector elements_per_thread{ + configuration.elements_per_thread[0], + configuration.elements_per_thread[1], + configuration.elements_per_thread[2]}; + const auto dispatch_x_value = TryGetMatMulPackedDispatchGroupCount( + dim_b_outer, configuration.workgroup_size[0], configuration.elements_per_thread[0]); + const auto dispatch_y_value = TryGetMatMulPackedDispatchGroupCount( + dim_a_outer, configuration.workgroup_size[1], configuration.elements_per_thread[1]); + const auto dispatch_z_value = TryGetMatMulPackedDispatchGroupCount( + batch_size, configuration.workgroup_size[2], configuration.elements_per_thread[2]); + ORT_RETURN_IF_NOT(dispatch_x_value.has_value() && + dispatch_y_value.has_value() && + dispatch_z_value.has_value(), + "MatMul packed dispatch dimensions exceed uint32 limits."); + const uint32_t dispatch_x = *dispatch_x_value; + const uint32_t dispatch_y = *dispatch_y_value; + uint32_t dispatch_z = *dispatch_z_value; const int components = is_vec4 ? 4 : 1; - const TensorShape a_shape_temp = CreateMatMulIntermediateShape(outer_dims_a, dim_a_outer, dim_inner, components); - const TensorShape b_shape_temp = CreateMatMulIntermediateShape(outer_dims_b, dim_inner, dim_b_outer, components); - const TensorShape output_shape_temp = TensorShape({batch_size, dim_a_outer, dim_b_outer / components}); - + const TensorShape a_shape_temp = + CreateMatMulIntermediateShape(outer_dims_a, dim_a_outer, dim_inner, components); + const TensorShape b_shape_temp = + CreateMatMulIntermediateShape(outer_dims_b, dim_inner, dim_b_outer, components); + const TensorShape output_shape_temp{batch_size, dim_a_outer, dim_b_outer / components}; ProgramOutput output(output_tensor, ProgramTensorMetadataDependency::Rank, output_shape_temp, components); const Tensor* bias = has_bias ? inputs[2] : nullptr; bool use_bias_in_matmul = has_bias; uint32_t split_dim_inner = 1; uint32_t splits_per_batch = 1; - // Current Split-K implementation relies on atomic operations, which are not deterministic. - if (!context->KernelContext().GetUseDeterministicCompute()) { - const SplitKConfig& split_k_config = context->GetSplitKConfig(); - const bool need_split_k = split_k_config.UseSplitK( - is_vec4, activation.activation_kind_, batch_size, dim_a_outer, dim_b_outer, - dim_inner, is_channels_last); - if (need_split_k) { - ORT_ENFORCE(is_vec4, "Split-K MatMul requires vec4 packing."); - - if (has_bias) { - ORT_ENFORCE(is_channels_last, "Split-K MatMul only supports channels-last format."); - } - - // Initialize `output_tensor` with 0 or bias before MatMulProgram with Split-K enabled. - const auto fill_bias_program = CreateMatMulFillBiasOrZeroBeforeSplitKProgram(bias, output_tensor, /*is_gemm*/ false, /*beta*/ 1.0f, /*bias_components*/ 4, output_shape_temp, narrow(batch_size)); - ORT_RETURN_IF_ERROR(context->RunProgram(fill_bias_program)); - - // `bias` has been handled in the execution of `fill_bias_program` so we don't need to set - // `bias` again in `MatMulProgram`. - use_bias_in_matmul = false; - - // With Split-K, `dim_inner` will be split into multiple parts. `dispatch_z` encodes - // both the split-k index and the batch index: dispatch_z = splits_per_batch * batch_size. - split_dim_inner = split_k_config.GetSplitDimInner(); - splits_per_batch = (dim_inner + split_dim_inner - 1) / split_dim_inner; - const uint64_t dispatch_z_u64 = static_cast(batch_size) * static_cast(splits_per_batch); - ORT_ENFORCE(dispatch_z_u64 <= static_cast(std::numeric_limits::max()), - "dispatch_z exceeds uint32_t range: ", dispatch_z_u64); - dispatch_z = narrow(dispatch_z_u64); - - // The output should be declared in atomic types in `MatMulProgram` for the use of atomic - // built-in functions. - output.is_atomic = true; - } + if (use_split_k) { + ORT_RETURN_IF(context.KernelContext().GetUseDeterministicCompute(), + "MatMul algorithm packed_split_k does not support deterministic compute."); + ORT_RETURN_IF_NOT(configuration.split_dim_inner > 1, + "MatMul algorithm packed_split_k is not configured for this adapter."); + ORT_RETURN_IF_NOT(is_vec4, + "MatMul algorithm packed_split_k requires vec4 packing."); + ORT_RETURN_IF_NOT(activation.activation_kind_ == ActivationKind::None, + "MatMul algorithm packed_split_k does not support a fused activation."); + ORT_RETURN_IF_NOT(!has_bias || is_channels_last, + "MatMul algorithm packed_split_k requires channels-last bias layout."); + + const auto fill_bias_program = CreateMatMulFillBiasOrZeroBeforeSplitKProgram( + bias, output_tensor, /*is_gemm=*/false, /*beta=*/1.0f, + /*output_components=*/4, output_shape_temp, narrow(batch_size)); + ORT_RETURN_IF_ERROR(context.RunProgram(fill_bias_program)); + use_bias_in_matmul = false; + split_dim_inner = configuration.split_dim_inner; + const auto splits_per_batch_value = TryGetMatMulPackedDispatchGroupCount( + dim_inner, split_dim_inner, 1); + ORT_RETURN_IF_NOT(splits_per_batch_value.has_value(), + "MatMul packed Split-K count exceeds uint32 limits."); + splits_per_batch = *splits_per_batch_value; + const uint64_t dispatch_z_u64 = + static_cast(batch_size) * static_cast(splits_per_batch); + ORT_RETURN_IF_NOT(dispatch_z_u64 <= static_cast(std::numeric_limits::max()), + "MatMul algorithm packed_split_k dispatch_z exceeds uint32_t range: ", dispatch_z_u64); + dispatch_z = narrow(dispatch_z_u64); + output.is_atomic = true; } - MatMulProgram matmul_program{activation, use_bias_in_matmul, is_vec4, elements_per_thread, is_channels_last, split_dim_inner}; - matmul_program - .CacheHint(activation.CacheKey(), absl::StrJoin(elements_per_thread, "-"), std::to_string(is_vec4), components, is_channels_last, split_dim_inner) + MatMulProgram program{activation, use_bias_in_matmul, is_vec4, elements_per_thread, + is_channels_last, split_dim_inner, configuration.tile_inner}; + program + .CacheHint(activation.CacheKey(), absl::StrJoin(elements_per_thread, "-"), + absl::StrJoin(configuration.workgroup_size, "-"), + std::to_string(is_vec4), components, is_channels_last, + split_dim_inner, configuration.tile_inner) .AddInputs({{a, ProgramTensorMetadataDependency::TypeAndRank, a_shape_temp, components}, {b, ProgramTensorMetadataDependency::TypeAndRank, b_shape_temp, components}}) .AddUniformVariables({{dim_a_outer}, {dim_b_outer}, {dim_inner}, {dispatch_x}, {dispatch_y}, {dispatch_z}, {splits_per_batch}}) .AddIndices(outer_dims) .SetDispatchGroupSize(dispatch_x, dispatch_y, dispatch_z) - .SetWorkgroupSize(MatMul::MATMUL_PACKED_WORKGROUP_SIZE_X, MatMul::MATMUL_PACKED_WORKGROUP_SIZE_Y, MatMul::MATMUL_PACKED_WORKGROUP_SIZE_Z) + .SetWorkgroupSize(configuration.workgroup_size[0], + configuration.workgroup_size[1], + configuration.workgroup_size[2]) .AddOutput(std::move(output)); - // Activation uniforms must remain last because definitions and values are matched by index. - AppendActivationUniformsData(activation, matmul_program); + AppendActivationUniformsData(activation, program); if (use_bias_in_matmul) { - auto bias_components = is_channels_last ? components : 1; - TensorShape reduced_bias_shape = ReduceShapeByComponents(bias->Shape(), bias_components); - matmul_program.AddInput({bias, ProgramTensorMetadataDependency::Rank, reduced_bias_shape, bias_components}); + const int bias_components = is_channels_last ? components : 1; + const TensorShape reduced_bias_shape = ReduceShapeByComponents(bias->Shape(), bias_components); + program.AddInput({bias, ProgramTensorMetadataDependency::Rank, reduced_bias_shape, bias_components}); + } + + return context.RunProgram(program); +} + +Status MatMulComputeDispatcher::Compute(ComputeContext& context, + const Activation& activation, + const std::vector& inputs, + Tensor* output_tensor, + bool is_channels_last, + bool b_is_constant) { + const auto* a = inputs[0]; + const auto* b = inputs[1]; + const bool has_bias = inputs.size() > 2; + const TensorShape& logical_a_shape = a->Shape(); + const TensorShape& logical_b_shape = b->Shape(); + ORT_RETURN_IF_NOT(logical_a_shape.NumDimensions() >= 2 && logical_b_shape.NumDimensions() >= 2, + "ComputeMatMul expects matrix or batched-matrix inputs."); + + MatMulComputeHelper helper; + ORT_RETURN_IF_ERROR(helper.Compute(logical_a_shape, logical_b_shape)); + + Initialize(context); + SubgroupMatrixMatMulImpl* subgroup_impl = subgroup_matrix_impl_.get(); + const bool can_use_subgroup_matrix = + subgroup_impl != nullptr && + subgroup_impl->CanApply(context, inputs, is_channels_last, b_is_constant); + const std::optional subgroup_size = intel::SelectMatMulSubgroupSize(context); + const bool has_subgroup_capability = subgroup_size.has_value(); + const int64_t batch_a = + logical_a_shape.NumDimensions() > 2 + ? logical_a_shape.SizeToDimension(logical_a_shape.NumDimensions() - 2) + : 1; + const int64_t batch_b = + logical_b_shape.NumDimensions() > 2 + ? logical_b_shape.SizeToDimension(logical_b_shape.NumDimensions() - 2) + : 1; + const TensorShape& logical_output_shape = helper.OutputShape(); + const uint64_t batch_size = logical_output_shape.NumDimensions() > 2 + ? narrow(logical_output_shape.SizeToDimension( + logical_output_shape.NumDimensions() - 2)) + : 1; + const bool folds_batch_into_m = batch_a != 1 && batch_b == 1; + + MatMulAlgorithmSelectionParams selection_params{}; + selection_params.m = helper.M(); + selection_params.n = helper.N(); + selection_params.k = helper.K(); + selection_params.packed_m = folds_batch_into_m + ? logical_a_shape.SizeToDimension(logical_a_shape.NumDimensions() - 1) + : helper.M(); + selection_params.batch_size = batch_size; + selection_params.packed_batch_size = folds_batch_into_m ? 1 : batch_size; + selection_params.adapter_architecture = context.AdapterInfo().architecture; + selection_params.a_data_type = a->GetElementType(); + selection_params.b_data_type = b->GetElementType(); + selection_params.can_use_subgroup_matrix = can_use_subgroup_matrix; + selection_params.has_subgroup_capability = has_subgroup_capability; + selection_params.subgroup_size = subgroup_size.value_or(0); + selection_params.is_vec4 = helper.K() % 4 == 0 && helper.N() % 4 == 0; + selection_params.deterministic_compute = context.KernelContext().GetUseDeterministicCompute(); + selection_params.has_fused_activation = activation.activation_kind_ != ActivationKind::None; + selection_params.has_bias = has_bias; + selection_params.is_channels_last = is_channels_last; + + const MatMulExecutionPlan plan = scheduler_->CreateExecutionPlan( + selection_params, context.ForcedMatMulAlgorithm()); + const MatMulAlgorithm algorithm = plan.algorithm; + ORT_RETURN_IF_NOT(IsMatMulAlgorithmConfigurationCompatible(plan), + "MatMul algorithm ", MatMulAlgorithmName(algorithm), + " received an incompatible vendor configuration."); + const auto* packed_configuration = std::get_if(&plan.configuration); + const auto* subgroup_configuration = std::get_if(&plan.configuration); + + MatMulAlgorithmPrerequisites prerequisites{}; + prerequisites.can_use_subgroup_matrix = can_use_subgroup_matrix; + prerequisites.has_subgroup_capability = has_subgroup_capability; + prerequisites.has_nonzero_k = helper.K() > 0; + prerequisites.split_k_configured = + packed_configuration != nullptr && packed_configuration->split_dim_inner > 1; + prerequisites.deterministic_compute = selection_params.deterministic_compute; + prerequisites.is_vec4 = selection_params.is_vec4; + prerequisites.has_fused_activation = selection_params.has_fused_activation; + prerequisites.split_k_bias_layout_supported = !has_bias || is_channels_last; + ORT_RETURN_IF_NOT(MeetsMatMulAlgorithmPrerequisites(algorithm, prerequisites), + "MatMul algorithm ", MatMulAlgorithmName(algorithm), + " does not support these inputs or this device."); + + switch (algorithm) { + case MatMulAlgorithm::SubgroupMatrix: + ORT_RETURN_IF_NOT(subgroup_impl != nullptr, + "MatMul algorithm subgroup_matrix is unavailable."); + return subgroup_impl->Compute( + context, inputs, output_tensor, activation, is_channels_last, b_is_constant); + case MatMulAlgorithm::Naive: + return ApplyMatMulNaive( + context, activation, inputs, output_tensor, is_channels_last, helper); + case MatMulAlgorithm::Subgroup: + return intel::ApplyMatMulSubgroup( + context, activation, inputs, output_tensor, is_channels_last, + subgroup_configuration->subgroup_size); + case MatMulAlgorithm::Packed: + return ApplyMatMulPacked( + context, activation, inputs, output_tensor, is_channels_last, helper, + *packed_configuration, + /*use_split_k=*/false); + case MatMulAlgorithm::PackedSplitK: + return ApplyMatMulPacked( + context, activation, inputs, output_tensor, is_channels_last, helper, + *packed_configuration, + /*use_split_k=*/true); } - return context->RunProgram(matmul_program); + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Unknown MatMul algorithm."); } MatMulFillBiasOrZeroBeforeSplitKProgram CreateMatMulFillBiasOrZeroBeforeSplitKProgram( diff --git a/onnxruntime/core/providers/webgpu/math/matmul.h b/onnxruntime/core/providers/webgpu/math/matmul.h index 4aaca49f0edbe..86e8c3df0a915 100644 --- a/onnxruntime/core/providers/webgpu/math/matmul.h +++ b/onnxruntime/core/providers/webgpu/math/matmul.h @@ -3,12 +3,10 @@ #pragma once -#include -#include - #include "core/providers/webgpu/webgpu_kernel.h" #include "core/providers/webgpu/program.h" #include "core/providers/cpu/math/matmul_helper.h" +#include "core/providers/webgpu/math/matmul_compute_dispatcher.h" #include "core/providers/webgpu/math/matmul_utils.h" #include "core/providers/webgpu/math/matmul_packed.h" #include "core/providers/webgpu/webgpu_utils.h" @@ -17,35 +15,6 @@ namespace onnxruntime { namespace webgpu { -class MatMulOptImpl { - public: - virtual ~MatMulOptImpl() = default; - - virtual Status Compute(ComputeContext& context, - const std::vector& inputs, - Tensor* output, - const Activation& activation, - bool is_channels_last, - bool b_is_constant, - /*out*/ bool& handled) = 0; -}; - -class MatMulOptImplCache { - public: - MatMulOptImplCache() = default; - ORT_DISALLOW_COPY_ASSIGNMENT_AND_MOVE(MatMulOptImplCache); - - MatMulOptImpl* GetOrCreate(const ComputeContextBase& context); - - private: - std::once_flag subgroup_impl_init_flag_; - std::unique_ptr subgroup_impl_; -}; - -Status ComputeMatMul(ComputeContext* context, const Activation& activation, std::vector& inputs, Tensor* output, - bool is_channels_last, MatMulOptImplCache& cache, - bool b_is_constant = false); - MatMulFillBiasOrZeroBeforeSplitKProgram CreateMatMulFillBiasOrZeroBeforeSplitKProgram( const Tensor* bias, Tensor* output, @@ -72,14 +41,15 @@ class MatMul final : public WebGpuKernel { constexpr static uint32_t MATMUL_PACKED_WORKGROUP_SIZE_Z = 1; private: - mutable MatMulOptImplCache compute_cache_; + mutable MatMulComputeDispatcher compute_dispatcher_; bool b_is_constant_ = false; }; class MatMulNaiveProgram final : public Program { public: - MatMulNaiveProgram(const Activation& activation, const size_t output_rank, int64_t output_number, bool has_bias, bool is_channels_last = false) - : Program{"MatMulNaive"}, activation_(activation), output_rank_(output_rank), output_number_(output_number), has_bias_{has_bias}, is_channels_last_(is_channels_last) { + MatMulNaiveProgram(const Activation& activation, const size_t output_rank, int64_t output_number, + bool has_bias, bool is_channels_last = false, bool is_zero_k = false) + : Program{"MatMulNaive"}, activation_(activation), output_rank_(output_rank), output_number_(output_number), has_bias_{has_bias}, is_channels_last_(is_channels_last), is_zero_k_(is_zero_k) { } Status GenerateShaderCode(ShaderHelper& sh) const override; @@ -96,6 +66,7 @@ class MatMulNaiveProgram final : public Program { const int64_t output_number_; const bool has_bias_; const bool is_channels_last_; + const bool is_zero_k_; }; } // namespace webgpu diff --git a/onnxruntime/core/providers/webgpu/math/matmul_algorithm.cc b/onnxruntime/core/providers/webgpu/math/matmul_algorithm.cc new file mode 100644 index 0000000000000..459f2ec57b771 --- /dev/null +++ b/onnxruntime/core/providers/webgpu/math/matmul_algorithm.cc @@ -0,0 +1,45 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/providers/webgpu/math/matmul_algorithm.h" + +namespace onnxruntime { +namespace webgpu { + +std::optional ParseMatMulAlgorithm(std::string_view name) { + if (name == "subgroup_matrix") { + return MatMulAlgorithm::SubgroupMatrix; + } + if (name == "naive") { + return MatMulAlgorithm::Naive; + } + if (name == "subgroup") { + return MatMulAlgorithm::Subgroup; + } + if (name == "packed") { + return MatMulAlgorithm::Packed; + } + if (name == "packed_split_k") { + return MatMulAlgorithm::PackedSplitK; + } + return std::nullopt; +} + +std::string_view MatMulAlgorithmName(MatMulAlgorithm algorithm) { + switch (algorithm) { + case MatMulAlgorithm::SubgroupMatrix: + return "subgroup_matrix"; + case MatMulAlgorithm::Naive: + return "naive"; + case MatMulAlgorithm::Subgroup: + return "subgroup"; + case MatMulAlgorithm::Packed: + return "packed"; + case MatMulAlgorithm::PackedSplitK: + return "packed_split_k"; + } + return "unknown"; +} + +} // namespace webgpu +} // namespace onnxruntime \ No newline at end of file diff --git a/onnxruntime/core/providers/webgpu/math/matmul_algorithm.h b/onnxruntime/core/providers/webgpu/math/matmul_algorithm.h new file mode 100644 index 0000000000000..1118af425ab33 --- /dev/null +++ b/onnxruntime/core/providers/webgpu/math/matmul_algorithm.h @@ -0,0 +1,25 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include +#include + +namespace onnxruntime { +namespace webgpu { + +enum class MatMulAlgorithm { + SubgroupMatrix, + Naive, + Subgroup, + Packed, + PackedSplitK, +}; + +std::optional ParseMatMulAlgorithm(std::string_view name); + +std::string_view MatMulAlgorithmName(MatMulAlgorithm algorithm); + +} // namespace webgpu +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/webgpu/math/matmul_algorithm_scheduler.cc b/onnxruntime/core/providers/webgpu/math/matmul_algorithm_scheduler.cc new file mode 100644 index 0000000000000..6d02af7659f07 --- /dev/null +++ b/onnxruntime/core/providers/webgpu/math/matmul_algorithm_scheduler.cc @@ -0,0 +1,189 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/providers/webgpu/math/matmul_algorithm_scheduler.h" + +#include +#include + +namespace onnxruntime { +namespace webgpu { + +bool IsMatMulPackedConfigurationValid( + const MatMulPackedConfiguration& configuration, + bool use_split_k) { + if (configuration.workgroup_size[0] == 0 || + configuration.workgroup_size[1] == 0 || + configuration.workgroup_size[2] != 1 || + configuration.elements_per_thread[0] == 0 || + configuration.elements_per_thread[1] == 0 || + configuration.elements_per_thread[2] != 1 || + configuration.tile_inner == 0) { + return false; + } + + return !use_split_k || + (configuration.split_dim_inner > 1 && + configuration.split_dim_inner % configuration.tile_inner == 0); +} + +std::optional TryGetMatMulPackedDispatchGroupCount( + uint64_t dimension, + uint32_t workgroup_size, + uint32_t elements_per_thread) { + if (workgroup_size == 0 || elements_per_thread == 0) { + return std::nullopt; + } + + const uint64_t elements_per_group = + static_cast(workgroup_size) * elements_per_thread; + const uint64_t group_count = + dimension == 0 ? 0 : 1 + (dimension - 1) / elements_per_group; + if (group_count > std::numeric_limits::max()) { + return std::nullopt; + } + + return static_cast(group_count); +} + +bool IsMatMulAlgorithmConfigurationCompatible(const MatMulExecutionPlan& plan) { + switch (plan.algorithm) { + case MatMulAlgorithm::SubgroupMatrix: + return std::holds_alternative(plan.configuration); + case MatMulAlgorithm::Naive: + return std::holds_alternative(plan.configuration); + case MatMulAlgorithm::Subgroup: + return std::holds_alternative(plan.configuration); + case MatMulAlgorithm::Packed: + case MatMulAlgorithm::PackedSplitK: + return std::holds_alternative(plan.configuration); + } + return false; +} + +bool MeetsMatMulAlgorithmPrerequisites( + MatMulAlgorithm algorithm, + const MatMulAlgorithmPrerequisites& prerequisites) { + switch (algorithm) { + case MatMulAlgorithm::SubgroupMatrix: + return prerequisites.can_use_subgroup_matrix; + case MatMulAlgorithm::Subgroup: + return prerequisites.has_subgroup_capability && + prerequisites.has_nonzero_k; + case MatMulAlgorithm::PackedSplitK: + return prerequisites.has_nonzero_k && + prerequisites.split_k_configured && + !prerequisites.deterministic_compute && + prerequisites.is_vec4 && + !prerequisites.has_fused_activation && + prerequisites.split_k_bias_layout_supported; + case MatMulAlgorithm::Packed: + return prerequisites.has_nonzero_k; + case MatMulAlgorithm::Naive: + return true; + } + return false; +} + +MatMulAlgorithmScheduler::MatMulAlgorithmScheduler(SplitKConfig split_k_config) + : split_k_config_{std::move(split_k_config)} {} + +MatMulAlgorithmScheduler::~MatMulAlgorithmScheduler() = default; + +MatMulAlgorithm MatMulAlgorithmScheduler::Select( + const MatMulAlgorithmSelectionParams& params, + std::optional forced_algorithm) const { + if (forced_algorithm.has_value()) { + return *forced_algorithm; + } + if (params.k == 0) { + return MatMulAlgorithm::Naive; + } + if (const auto vendor_algorithm = SelectVendorAlgorithm(params); vendor_algorithm.has_value()) { + return *vendor_algorithm; + } + return SelectCommonAlgorithm(params); +} + +MatMulExecutionPlan MatMulAlgorithmScheduler::CreateExecutionPlan( + const MatMulAlgorithmSelectionParams& params, + std::optional forced_algorithm) const { + const MatMulAlgorithm algorithm = Select(params, forced_algorithm); + std::optional configuration = + SelectVendorConfiguration(algorithm, params); + if (!configuration.has_value()) { + configuration = SelectCommonConfiguration(algorithm, params); + } + return MatMulExecutionPlan{algorithm, std::move(*configuration)}; +} + +std::optional MatMulAlgorithmScheduler::SelectVendorAlgorithm( + const MatMulAlgorithmSelectionParams& /*params*/) const { + return std::nullopt; +} + +std::optional MatMulAlgorithmScheduler::SelectVendorConfiguration( + MatMulAlgorithm /*algorithm*/, + const MatMulAlgorithmSelectionParams& /*params*/) const { + return std::nullopt; +} + +MatMulAlgorithm MatMulAlgorithmScheduler::SelectCommonAlgorithm( + const MatMulAlgorithmSelectionParams& params) const { + if (params.can_use_subgroup_matrix) { + return MatMulAlgorithm::SubgroupMatrix; + } + if (params.n < 8 && params.k < 8) { + return MatMulAlgorithm::Naive; + } + if (ShouldUseSplitK(params)) { + return MatMulAlgorithm::PackedSplitK; + } + return MatMulAlgorithm::Packed; +} + +MatMulAlgorithmConfiguration MatMulAlgorithmScheduler::SelectCommonConfiguration( + MatMulAlgorithm algorithm, + const MatMulAlgorithmSelectionParams& params) const { + switch (algorithm) { + case MatMulAlgorithm::SubgroupMatrix: + return MatMulSubgroupMatrixConfiguration{}; + case MatMulAlgorithm::Naive: + return MatMulNaiveConfiguration{}; + case MatMulAlgorithm::Subgroup: + return MatMulSubgroupConfiguration{params.subgroup_size}; + case MatMulAlgorithm::Packed: + case MatMulAlgorithm::PackedSplitK: { + MatMulPackedConfiguration configuration{}; + configuration.elements_per_thread = + params.packed_m <= 8 ? std::array{4, 1, 1} + : std::array{4, 4, 1}; + configuration.split_dim_inner = + algorithm == MatMulAlgorithm::PackedSplitK + ? split_k_config_.GetSplitDimInner() + : 1; + return configuration; + } + } + return MatMulNaiveConfiguration{}; +} + +bool MatMulAlgorithmScheduler::ShouldUseSplitK( + const MatMulAlgorithmSelectionParams& params) const { + if (params.deterministic_compute || params.has_fused_activation || + params.packed_m < 0 || params.n < 0 || params.k < 0) { + return false; + } + + return split_k_config_.UseSplitK( + params.is_vec4, + ActivationKind::None, + params.packed_batch_size, + static_cast(params.packed_m), + static_cast(params.n), + static_cast(params.k), + params.is_channels_last); +} + +} // namespace webgpu +} // namespace onnxruntime \ No newline at end of file diff --git a/onnxruntime/core/providers/webgpu/math/matmul_algorithm_scheduler.h b/onnxruntime/core/providers/webgpu/math/matmul_algorithm_scheduler.h new file mode 100644 index 0000000000000..05d102c5beb6b --- /dev/null +++ b/onnxruntime/core/providers/webgpu/math/matmul_algorithm_scheduler.h @@ -0,0 +1,131 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include +#include +#include +#include +#include + +#include "core/common/common.h" +#include "core/providers/webgpu/math/matmul_algorithm.h" +#include "core/providers/webgpu/webgpu_utils.h" + +namespace onnxruntime { +namespace webgpu { + +// Immutable problem facts used by automatic selection and execution tuning. +struct MatMulAlgorithmSelectionParams { + int64_t m = 0; + int64_t n = 0; + int64_t k = 0; + int64_t packed_m = 0; + uint64_t batch_size = 1; + uint64_t packed_batch_size = 1; + std::string_view adapter_architecture; + int32_t a_data_type = 0; + int32_t b_data_type = 0; + bool can_use_subgroup_matrix = false; + bool has_subgroup_capability = false; + uint32_t subgroup_size = 0; + bool is_vec4 = false; + bool deterministic_compute = false; + bool has_fused_activation = false; + bool has_bias = false; + bool is_channels_last = true; +}; + +// Algorithm-specific tuning selected together with the implementation. +struct MatMulSubgroupMatrixConfiguration {}; +struct MatMulNaiveConfiguration {}; +struct MatMulSubgroupConfiguration { + uint32_t subgroup_size = 0; +}; + +struct MatMulPackedConfiguration { + std::array workgroup_size{8, 8, 1}; + std::array elements_per_thread{4, 4, 1}; + uint32_t tile_inner = 32; + uint32_t split_dim_inner = 1; +}; + +bool IsMatMulPackedConfigurationValid( + const MatMulPackedConfiguration& configuration, + bool use_split_k); + +std::optional TryGetMatMulPackedDispatchGroupCount( + uint64_t dimension, + uint32_t workgroup_size, + uint32_t elements_per_thread); + +using MatMulAlgorithmConfiguration = + std::variant; + +// Complete per-invocation decision consumed by the compute dispatcher. +struct MatMulExecutionPlan { + MatMulAlgorithm algorithm; + MatMulAlgorithmConfiguration configuration; +}; + +bool IsMatMulAlgorithmConfigurationCompatible(const MatMulExecutionPlan& plan); + +// Runtime correctness constraints validated immediately before dispatch. +struct MatMulAlgorithmPrerequisites { + bool can_use_subgroup_matrix = false; + bool has_subgroup_capability = false; + bool has_nonzero_k = false; + bool split_k_configured = false; + bool deterministic_compute = false; + bool is_vec4 = false; + bool has_fused_activation = false; + bool split_k_bias_layout_supported = true; +}; + +bool MeetsMatMulAlgorithmPrerequisites( + MatMulAlgorithm algorithm, + const MatMulAlgorithmPrerequisites& prerequisites); + +// Pure selection and tuning policy. Runtime validation remains in the dispatcher. +class MatMulAlgorithmScheduler { + public: + explicit MatMulAlgorithmScheduler(SplitKConfig split_k_config = {}); + virtual ~MatMulAlgorithmScheduler(); + ORT_DISALLOW_COPY_ASSIGNMENT_AND_MOVE(MatMulAlgorithmScheduler); + + MatMulAlgorithm Select( + const MatMulAlgorithmSelectionParams& params, + std::optional forced_algorithm = std::nullopt) const; + + // A forced plan may intentionally violate runtime prerequisites; the dispatcher validates it + // and reports an error naming the requested algorithm instead of silently changing selection. + MatMulExecutionPlan CreateExecutionPlan( + const MatMulAlgorithmSelectionParams& params, + std::optional forced_algorithm = std::nullopt) const; + + protected: + virtual std::optional SelectVendorAlgorithm( + const MatMulAlgorithmSelectionParams& params) const; + + virtual std::optional SelectVendorConfiguration( + MatMulAlgorithm algorithm, + const MatMulAlgorithmSelectionParams& params) const; + + private: + MatMulAlgorithm SelectCommonAlgorithm(const MatMulAlgorithmSelectionParams& params) const; + + MatMulAlgorithmConfiguration SelectCommonConfiguration( + MatMulAlgorithm algorithm, + const MatMulAlgorithmSelectionParams& params) const; + + bool ShouldUseSplitK(const MatMulAlgorithmSelectionParams& params) const; + + SplitKConfig split_k_config_; +}; + +} // namespace webgpu +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/webgpu/math/matmul_compute_dispatcher.h b/onnxruntime/core/providers/webgpu/math/matmul_compute_dispatcher.h new file mode 100644 index 0000000000000..6d574411f6523 --- /dev/null +++ b/onnxruntime/core/providers/webgpu/math/matmul_compute_dispatcher.h @@ -0,0 +1,66 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include +#include +#include + +#include "core/common/common.h" +#include "core/common/status.h" +#include "core/providers/webgpu/math/matmul_algorithm_scheduler.h" + +namespace onnxruntime { + +class Tensor; + +namespace webgpu { + +struct Activation; +class ComputeContext; +class ComputeContextBase; + +// Stateful subgroup-matrix implementation owned independently from selection policy. +class SubgroupMatrixMatMulImpl { + public: + SubgroupMatrixMatMulImpl() = default; + virtual ~SubgroupMatrixMatMulImpl() = default; + ORT_DISALLOW_COPY_ASSIGNMENT_AND_MOVE(SubgroupMatrixMatMulImpl); + + virtual bool CanApply(const ComputeContext& context, + const std::vector& inputs, + bool is_channels_last, + bool b_is_constant) const = 0; + + virtual Status Compute(ComputeContext& context, + const std::vector& inputs, + Tensor* output, + const Activation& activation, + bool is_channels_last, + bool b_is_constant) = 0; +}; + +// Selects, validates, and dispatches one concrete implementation per invocation. +class MatMulComputeDispatcher { + public: + MatMulComputeDispatcher() = default; + ORT_DISALLOW_COPY_ASSIGNMENT_AND_MOVE(MatMulComputeDispatcher); + + Status Compute(ComputeContext& context, + const Activation& activation, + const std::vector& inputs, + Tensor* output, + bool is_channels_last, + bool b_is_constant = false); + + private: + void Initialize(const ComputeContextBase& context); + + std::once_flag init_flag_; + std::unique_ptr subgroup_matrix_impl_; + std::unique_ptr scheduler_; +}; + +} // namespace webgpu +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/webgpu/math/matmul_packed.cc b/onnxruntime/core/providers/webgpu/math/matmul_packed.cc index 6cd2f692a5cd3..59a40c6c3f12a 100644 --- a/onnxruntime/core/providers/webgpu/math/matmul_packed.cc +++ b/onnxruntime/core/providers/webgpu/math/matmul_packed.cc @@ -47,9 +47,12 @@ Status MatMulProgram::GenerateShaderCode(ShaderHelper& shader) const { ORT_RETURN_IF_ERROR(MakeMatMulPackedVec4Source( shader, elements_per_thread_, WorkgroupSizeX(), WorkgroupSizeY(), data_type, &batch_dims, /*transA = */ false, /*transB = */ false, /*alpha = */ 1.f, /*need_handle_matmul = */ true, - /*output_components = */ 4, /*tile_inner = */ 32, need_split_k, split_dim_inner_)); + /*output_components = */ 4, tile_inner_, need_split_k, split_dim_inner_)); } else { - ORT_RETURN_IF_ERROR(MakeMatMulPackedSource(shader, elements_per_thread_, WorkgroupSizeX(), WorkgroupSizeY(), data_type, &batch_dims)); + ORT_RETURN_IF_ERROR(MakeMatMulPackedSource( + shader, elements_per_thread_, WorkgroupSizeX(), WorkgroupSizeY(), data_type, &batch_dims, + /*transpose_a = */ false, /*transpose_b = */ false, /*alpha = */ 1.f, + /*need_handle_matmul = */ true, tile_inner_)); } return Status::OK(); } diff --git a/onnxruntime/core/providers/webgpu/math/matmul_packed.h b/onnxruntime/core/providers/webgpu/math/matmul_packed.h index 833305e9cddc8..e1be676298fca 100644 --- a/onnxruntime/core/providers/webgpu/math/matmul_packed.h +++ b/onnxruntime/core/providers/webgpu/math/matmul_packed.h @@ -13,13 +13,21 @@ namespace onnxruntime { namespace webgpu { class MatMulProgram final : public Program { public: - MatMulProgram(const Activation& activation, bool bias, bool is_vec4, const gsl::span& elements_per_thread, bool is_channels_last = false, uint32_t split_dim_inner = 1) : Program{"MatMul"}, - activation_(activation), - has_bias_{bias}, - is_vec4_{is_vec4}, - elements_per_thread_(elements_per_thread.begin(), elements_per_thread.end()), - is_channels_last_(is_channels_last), - split_dim_inner_(split_dim_inner) {} + MatMulProgram(const Activation& activation, + bool bias, + bool is_vec4, + const gsl::span& elements_per_thread, + bool is_channels_last = false, + uint32_t split_dim_inner = 1, + uint32_t tile_inner = 32) + : Program{"MatMul"}, + activation_(activation), + has_bias_{bias}, + is_vec4_{is_vec4}, + elements_per_thread_(elements_per_thread.begin(), elements_per_thread.end()), + is_channels_last_(is_channels_last), + split_dim_inner_(split_dim_inner), + tile_inner_(tile_inner) {} Status GenerateShaderCode(ShaderHelper& sh) const override; WEBGPU_PROGRAM_DEFINE_UNIFORM_VARIABLES({"dim_a_outer", ProgramUniformVariableDataType::Uint32}, @@ -40,6 +48,7 @@ class MatMulProgram final : public Program { const InlinedVector elements_per_thread_; bool is_channels_last_ = false; uint32_t split_dim_inner_ = 1; + uint32_t tile_inner_ = 32; }; // The program to initialize the output with 0 or bias before doing MatMul with Split-K. In Split-K, diff --git a/onnxruntime/core/providers/webgpu/math/subgroup_matrix_matmul.cc b/onnxruntime/core/providers/webgpu/math/subgroup_matrix_matmul.cc index 3b2c75a9f1391..9fe0aff886aaa 100644 --- a/onnxruntime/core/providers/webgpu/math/subgroup_matrix_matmul.cc +++ b/onnxruntime/core/providers/webgpu/math/subgroup_matrix_matmul.cc @@ -76,21 +76,16 @@ class SubgroupMatrixMatMulPadBProgram final : public Program& inputs, - Tensor* output, - const Activation& activation, - bool is_channels_last, - bool b_is_constant, - /*out*/ bool& handled) override { - handled = false; - + bool CanApply(const ComputeContext& context, + const std::vector& inputs, + bool is_channels_last, + bool b_is_constant) const override { const auto* a = inputs[0]; const auto* b = inputs[1]; const auto& a_shape = a->Shape(); @@ -100,55 +95,102 @@ class SubgroupMatrixMatMulImpl final : public MatMulOptImpl { const size_t b_rank = b_shape.NumDimensions(); if ((!is_channels_last && has_bias) || a_rank < 2 || b_rank < 2 || !a->IsDataType() || !b->IsDataType()) { - return Status::OK(); + return false; } const uint32_t K = narrow(a_shape[a_rank - 1]); if (K == 0) { - return Status::OK(); + return false; } uint32_t M = 0; uint32_t N = 0; uint32_t batch = 1; if (b_rank == 2) { - ORT_ENFORCE(narrow(b_shape[0]) == K, - "MatMul contraction dim mismatch: A K=", K, " vs B rows=", b_shape[0]); + if (narrow(b_shape[0]) != K) { + return false; + } M = narrow(a_shape.Size() / static_cast(K)); N = narrow(b_shape[1]); } else { - if (a_rank != b_rank) { - return Status::OK(); + if (a_rank != b_rank || narrow(b_shape[b_rank - 2]) != K) { + return false; } - ORT_ENFORCE(narrow(b_shape[b_rank - 2]) == K, - "MatMul contraction dim mismatch: A K=", K, - " vs B rows=", b_shape[b_rank - 2]); M = narrow(a_shape[a_rank - 2]); N = narrow(b_shape[b_rank - 1]); for (size_t i = 0; i + 2 < a_rank; ++i) { if (a_shape[i] != b_shape[i]) { - return Status::OK(); + return false; } } batch = narrow(a_shape.SizeToDimension(a_rank - 2)); } if (M == 0 || N == 0) { - return Status::OK(); + return false; } const std::optional tiling = tiling_selector_(context, M, N, K, batch); if (!tiling) { - return Status::OK(); + return false; } + const auto& config = config_; + const bool needs_padded_b = N % 2 != 0; + if (needs_padded_b && (!b_is_constant || N == std::numeric_limits::max())) { + return false; + } + const uint32_t n_b = needs_padded_b ? N + 1 : N; + return config.K != 0 && K % config.K == 0 && + M >= tiling->tile_m && n_b >= tiling->tile_n; + } + + Status Compute(ComputeContext& context, + const std::vector& inputs, + Tensor* output, + const Activation& activation, + bool is_channels_last, + bool b_is_constant) override { + ORT_RETURN_IF_NOT(CanApply(context, inputs, is_channels_last, b_is_constant), + "MatMul algorithm subgroup_matrix does not support these inputs or this device."); + + const auto* a = inputs[0]; + const auto* b = inputs[1]; + const auto& a_shape = a->Shape(); + const auto& b_shape = b->Shape(); + const bool has_bias = inputs.size() > 2; + const size_t a_rank = a_shape.NumDimensions(); + const size_t b_rank = b_shape.NumDimensions(); + const uint32_t K = narrow(a_shape[a_rank - 1]); + + uint32_t M = 0; + uint32_t N = 0; + uint32_t batch = 1; + if (b_rank == 2) { + ORT_ENFORCE(narrow(b_shape[0]) == K, + "MatMul contraction dim mismatch: A K=", K, " vs B rows=", b_shape[0]); + M = narrow(a_shape.Size() / static_cast(K)); + N = narrow(b_shape[1]); + } else { + ORT_ENFORCE(narrow(b_shape[b_rank - 2]) == K, + "MatMul contraction dim mismatch: A K=", K, + " vs B rows=", b_shape[b_rank - 2]); + M = narrow(a_shape[a_rank - 2]); + N = narrow(b_shape[b_rank - 1]); + for (size_t i = 0; i + 2 < a_rank; ++i) { + ORT_ENFORCE(a_shape[i] == b_shape[i]); + } + batch = narrow(a_shape.SizeToDimension(a_rank - 2)); + } + + const std::optional tiling = tiling_selector_(context, M, N, K, batch); + ORT_ENFORCE(tiling.has_value()); + const auto& config = config_; const bool needs_padded_b = N % 2 != 0; // Require whole subgroup-matrix K blocks. An odd-width B must be constant // because its padded copy is cached by this implementation. - if (config.K == 0 || K % config.K != 0 || - (needs_padded_b && !b_is_constant)) { - return Status::OK(); - } + ORT_ENFORCE(config.K != 0 && K % config.K == 0 && + (!needs_padded_b || b_is_constant)); // N_b is just N rounded up to even - compute it before doing any padding work so // the tile-fit check below can bail out without a wasted pad dispatch. @@ -161,9 +203,7 @@ class SubgroupMatrixMatMulImpl final : public MatMulOptImpl { // The kernel keeps its operand loads in bounds by shifting a trailing partial // tile back, which is only possible when the tile fits within M and N. - if (M < tiling->tile_m || N_b < tiling->tile_n) { - return Status::OK(); - } + ORT_ENFORCE(M >= tiling->tile_m && N_b >= tiling->tile_n); // The optimized path will run: now materialize the even-strided B for odd N. const Tensor* b_used = b; @@ -208,7 +248,6 @@ class SubgroupMatrixMatMulImpl final : public MatMulOptImpl { } ORT_RETURN_IF_ERROR(context.RunProgram(program)); - handled = true; return Status::OK(); } @@ -320,7 +359,7 @@ Status SubgroupMatrixMatMulProgram::GenerateShaderCode(ShaderHelper& shader) con "Unsupported subgroup matrix config dimensions."); } -std::unique_ptr CreateSubgroupMatrixMatMulImpl(const ComputeContextBase& context) { +std::unique_ptr CreateSubgroupMatrixMatMulImpl(const ComputeContextBase& context) { // Only run on devices that report the 8x16x16 F16 subgroup-matrix config this // kernel is implemented for and can provide its required subgroup size. constexpr auto kF16 = wgpu::SubgroupMatrixComponentType::F16; @@ -336,7 +375,7 @@ std::unique_ptr CreateSubgroupMatrixMatMulImpl(const ComputeConte if (!tiling_selector) { return nullptr; } - return std::make_unique(*config, std::move(tiling_selector)); + return std::make_unique(*config, std::move(tiling_selector)); } } // namespace webgpu diff --git a/onnxruntime/core/providers/webgpu/nn/conv.cc b/onnxruntime/core/providers/webgpu/nn/conv.cc index 4b7937e02f240..c522d5bdec6d6 100644 --- a/onnxruntime/core/providers/webgpu/nn/conv.cc +++ b/onnxruntime/core/providers/webgpu/nn/conv.cc @@ -9,7 +9,6 @@ #include "core/providers/webgpu/tensor/transpose.h" #include "core/providers/webgpu/nn/grouped_conv.h" #include "core/providers/webgpu/webgpu_utils.h" -#include "core/providers/webgpu/math/matmul.h" namespace onnxruntime { namespace webgpu { @@ -281,8 +280,8 @@ Status Conv::ComputeInternal(ComputeContext& context if (has_bias) { matmul_inputs.push_back(bias); } - return ComputeMatMul(&context, activation_, matmul_inputs, output, is_channels_last, - matmul_compute_cache_, matmul_b_is_constant); + return matmul_compute_dispatcher_.Compute(context, activation_, matmul_inputs, output, is_channels_last, + matmul_b_is_constant); } // Transpose weights when necessary Tensor transposed_kernel; diff --git a/onnxruntime/core/providers/webgpu/nn/conv.h b/onnxruntime/core/providers/webgpu/nn/conv.h index a64a206f1ef34..6684b7ca5259c 100644 --- a/onnxruntime/core/providers/webgpu/nn/conv.h +++ b/onnxruntime/core/providers/webgpu/nn/conv.h @@ -8,7 +8,7 @@ #include "core/providers/cpu/nn/conv_attributes.h" #include "core/providers/webgpu/program.h" #include "core/providers/webgpu/shader_helper.h" -#include "core/providers/webgpu/math/matmul.h" +#include "core/providers/webgpu/math/matmul_compute_dispatcher.h" #include "core/providers/webgpu/nn/fuse_utils.h" namespace onnxruntime { @@ -49,7 +49,7 @@ class Conv : public WebGpuKernel { // Layout of the tensor ComputeInternal ends up consuming -- `prepacked_kernel_` when it // is set, otherwise input 1. Stays `OIHW` while `prepacked_kernel_` is null. KernelLayout kernel_layout_{KernelLayout::OIHW}; - mutable MatMulOptImplCache matmul_compute_cache_; + mutable MatMulComputeDispatcher matmul_compute_dispatcher_; }; Status TransposeKernel(ComputeContext& context, const Tensor* kernel, const TensorShape& kernel_shape, Tensor* transposed_kernel, const InlinedVector& perm); diff --git a/onnxruntime/core/providers/webgpu/vendor/intel/math/gemm_subgroup.cc b/onnxruntime/core/providers/webgpu/vendor/intel/math/gemm_subgroup.cc index 5b17f06e08070..b4e6c195881e5 100644 --- a/onnxruntime/core/providers/webgpu/vendor/intel/math/gemm_subgroup.cc +++ b/onnxruntime/core/providers/webgpu/vendor/intel/math/gemm_subgroup.cc @@ -161,6 +161,18 @@ int64_t ElementsPerThreadY(ComputeContext& context, uint32_t M) { return M <= 8 ? 1 : (M <= 16 ? 2 : (M <= 32 ? 4 : (is_xe_lpg_or_xe_3lpg ? 4 : 8))); } +bool CanUseAVec4CooperativeLoad(std::string_view architecture, + uint32_t dim_inner, + int64_t elements_per_thread_y) { + // A 32-wide subgroup has four 8-lane cooperative-load groups. Every generated + // subgroup-size branch must distribute rows evenly across its lane groups. + constexpr int64_t max_lane_groups = 4; + return architecture == gpu_arch::kXe3Lpg && + dim_inner % 4 == 0 && + elements_per_thread_y > 0 && + elements_per_thread_y % max_lane_groups == 0; +} + Status MakeMatMulSubgroupSource(ShaderHelper& shader, const InlinedVector& elements_per_thread, const ShaderIndicesHelper* batch_dims, diff --git a/onnxruntime/core/providers/webgpu/vendor/intel/math/gemm_subgroup.h b/onnxruntime/core/providers/webgpu/vendor/intel/math/gemm_subgroup.h index 72babd38cd06d..dbe6b645af03b 100644 --- a/onnxruntime/core/providers/webgpu/vendor/intel/math/gemm_subgroup.h +++ b/onnxruntime/core/providers/webgpu/vendor/intel/math/gemm_subgroup.h @@ -4,16 +4,12 @@ #pragma once #include "core/providers/webgpu/shader_helper.h" +#include "core/providers/webgpu/vendor/intel/math/gemm_subgroup_utils.h" namespace onnxruntime { namespace webgpu { namespace intel { -namespace gpu_arch { -inline constexpr std::string_view kXeLpg = "xe-lpg"; -inline constexpr std::string_view kXe3Lpg = "xe-3lpg"; -} // namespace gpu_arch - const uint32_t kSubgroupLogicalWorkGroupSizeX = 32; const uint32_t kSubgroupLogicalWorkGroupSizeY = 8; const uint32_t kSubgroupLogicalWorkGroupSizeZ = 1; diff --git a/onnxruntime/core/providers/webgpu/vendor/intel/math/gemm_subgroup_utils.h b/onnxruntime/core/providers/webgpu/vendor/intel/math/gemm_subgroup_utils.h new file mode 100644 index 0000000000000..c1d6139f16b0f --- /dev/null +++ b/onnxruntime/core/providers/webgpu/vendor/intel/math/gemm_subgroup_utils.h @@ -0,0 +1,29 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include +#include +#include + +namespace onnxruntime { +namespace webgpu { +namespace intel { + +namespace gpu_arch { +inline constexpr std::string_view kXeLpg = "xe-lpg"; +inline constexpr std::string_view kXe3Lpg = "xe-3lpg"; +} // namespace gpu_arch + +bool CanUseAVec4CooperativeLoad(std::string_view architecture, + uint32_t dim_inner, + int64_t elements_per_thread_y); + +std::optional SelectMatMulSubgroupSize(uint32_t adapter_min_subgroup_size, + uint32_t adapter_max_subgroup_size, + bool has_subgroup_size_control); + +} // namespace intel +} // namespace webgpu +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/webgpu/vendor/intel/math/matmul.cc b/onnxruntime/core/providers/webgpu/vendor/intel/math/matmul.cc index 9bb66fc551242..e3b19088cd193 100644 --- a/onnxruntime/core/providers/webgpu/vendor/intel/math/matmul.cc +++ b/onnxruntime/core/providers/webgpu/vendor/intel/math/matmul.cc @@ -12,6 +12,12 @@ namespace onnxruntime { namespace webgpu { namespace intel { +namespace { + +constexpr std::array kSupportedMatMulSubgroupSizes{32, 16, 8}; + +} // namespace + Status MatMulSubgroupProgram::GenerateShaderCode(ShaderHelper& shader) const { const auto& a = shader.AddInput("a", ShaderUsage::UseUniform | ShaderUsage::UseIndicesTypeAlias | ShaderUsage::UseValueTypeAlias | ShaderUsage::UseElementTypeAlias); @@ -36,15 +42,46 @@ Status MatMulSubgroupProgram::GenerateShaderCode(ShaderHelper& shader) const { return Status::OK(); } -bool CanApplyMatMulIntel(const ComputeContext& context, int64_t M, int64_t N, int64_t K) { - return CanApplySubgroup(context, M, N, K); +std::optional SelectMatMulSubgroupSize(uint32_t adapter_min_subgroup_size, + uint32_t adapter_max_subgroup_size, + bool has_subgroup_size_control) { + if (adapter_min_subgroup_size == adapter_max_subgroup_size) { + for (const uint32_t size : kSupportedMatMulSubgroupSizes) { + if (adapter_min_subgroup_size == size) { + return size; + } + } + return std::nullopt; + } + + if (has_subgroup_size_control) { + for (const uint32_t size : kSupportedMatMulSubgroupSizes) { + if (adapter_min_subgroup_size <= size && size <= adapter_max_subgroup_size) { + return size; + } + } + } + return std::nullopt; } -Status ApplyMatMulIntel(ComputeContext& context, - const Activation& activation, - const std::vector& inputs, - Tensor* output, - bool is_channels_last) { +std::optional SelectMatMulSubgroupSize(const ComputeContext& context) { + if (!context.HasFeature(wgpu::FeatureName::Subgroups)) { + return std::nullopt; + } + + const auto& adapter_info = context.AdapterInfo(); + return SelectMatMulSubgroupSize( + adapter_info.subgroupMinSize, + adapter_info.subgroupMaxSize, + context.HasFeature(wgpu::FeatureName::SubgroupSizeControl)); +} + +Status ApplyMatMulSubgroup(ComputeContext& context, + const Activation& activation, + const std::vector& inputs, + Tensor* output, + bool is_channels_last, + uint32_t subgroup_size) { const auto* a = inputs[0]; const auto* b = inputs[1]; bool has_bias = inputs.size() > 2; @@ -106,11 +143,12 @@ Status ApplyMatMulIntel(ComputeContext& context, const bool is_vec4 = dim_b_outer % 4 == 0; // vec4 A loads and double-buffering of the B tile are only enabled on Xe-3LPG. const bool is_xe_3lpg = arch == gpu_arch::kXe3Lpg; - // Load A from global memory as vec4 when K is a multiple of 4; otherwise fall back to scalar load. - const bool a_vec4 = is_xe_3lpg && dim_inner % 4 == 0; // Double-buffering of the B tile (held in workgroup memory) is only enabled for float16 B inputs. const bool b_is_fp16 = is_xe_3lpg && b->GetElementType() == ONNX_NAMESPACE::TensorProto_DataType_FLOAT16; InlinedVector elements_per_thread = InlinedVector({4, ElementsPerThreadY(context, dim_a_outer), 1}); + // Fall back to scalar A loads when rows cannot be distributed evenly among + // the cooperative vec4 lane groups (for example, forced low-M execution). + const bool a_vec4 = CanUseAVec4CooperativeLoad(arch, dim_inner, elements_per_thread[1]); const uint32_t dispatch_x = narrow((dim_b_outer + kSubgroupLogicalWorkGroupSizeX * elements_per_thread[0] - 1) / (kSubgroupLogicalWorkGroupSizeX * elements_per_thread[0])); @@ -129,6 +167,9 @@ Status ApplyMatMulIntel(ComputeContext& context, MatMulSubgroupProgram program{activation, has_bias, is_vec4, a_vec4, b_is_fp16, is_channels_last, elements_per_thread}; + if (context.HasFeature(wgpu::FeatureName::SubgroupSizeControl)) { + program.SetSubgroupSize(subgroup_size); + } program .CacheHint(activation.CacheKey(), absl::StrJoin(elements_per_thread, "-"), a_vec4, b_is_fp16, is_channels_last) diff --git a/onnxruntime/core/providers/webgpu/vendor/intel/math/matmul.h b/onnxruntime/core/providers/webgpu/vendor/intel/math/matmul.h index 0887e0539a8b0..8fc1816d6baba 100644 --- a/onnxruntime/core/providers/webgpu/vendor/intel/math/matmul.h +++ b/onnxruntime/core/providers/webgpu/vendor/intel/math/matmul.h @@ -3,6 +3,9 @@ #pragma once +#include +#include + #include "core/providers/webgpu/webgpu_kernel.h" #include "core/providers/webgpu/shader_helper.h" #include "core/providers/webgpu/program.h" @@ -12,6 +15,8 @@ namespace onnxruntime { namespace webgpu { namespace intel { +// TODO: Move this common subgroup implementation out of vendor/intel. Keep only +// Intel-specific automatic-selection thresholds and tuning under the vendor directory. class MatMulSubgroupProgram final : public Program { public: MatMulSubgroupProgram(const Activation& activation, @@ -46,13 +51,14 @@ class MatMulSubgroupProgram final : public Program { const InlinedVector elements_per_thread_; }; -bool CanApplyMatMulIntel(const ComputeContext& context, int64_t M, int64_t N, int64_t K); +std::optional SelectMatMulSubgroupSize(const ComputeContext& context); -Status ApplyMatMulIntel(ComputeContext& context, - const Activation& activation, - const std::vector& inputs, - Tensor* output, - bool is_channels_last); +Status ApplyMatMulSubgroup(ComputeContext& context, + const Activation& activation, + const std::vector& inputs, + Tensor* output, + bool is_channels_last, + uint32_t subgroup_size); } // namespace intel } // namespace webgpu diff --git a/onnxruntime/core/providers/webgpu/vendor/intel/math/matmul_algorithm_scheduler.cc b/onnxruntime/core/providers/webgpu/vendor/intel/math/matmul_algorithm_scheduler.cc new file mode 100644 index 0000000000000..0f859c6a93d01 --- /dev/null +++ b/onnxruntime/core/providers/webgpu/vendor/intel/math/matmul_algorithm_scheduler.cc @@ -0,0 +1,29 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/providers/webgpu/vendor/intel/math/matmul_algorithm_scheduler.h" + +#include + +namespace onnxruntime { +namespace webgpu { +namespace intel { + +IntelMatMulAlgorithmScheduler::IntelMatMulAlgorithmScheduler(SplitKConfig split_k_config) + : MatMulAlgorithmScheduler{std::move(split_k_config)} {} + +std::optional IntelMatMulAlgorithmScheduler::SelectVendorAlgorithm( + const MatMulAlgorithmSelectionParams& params) const { + if (params.can_use_subgroup_matrix) { + return MatMulAlgorithm::SubgroupMatrix; + } + if (params.has_subgroup_capability && + params.m >= 64 && params.n >= 512 && params.k >= 32) { + return MatMulAlgorithm::Subgroup; + } + return std::nullopt; +} + +} // namespace intel +} // namespace webgpu +} // namespace onnxruntime \ No newline at end of file diff --git a/onnxruntime/core/providers/webgpu/vendor/intel/math/matmul_algorithm_scheduler.h b/onnxruntime/core/providers/webgpu/vendor/intel/math/matmul_algorithm_scheduler.h new file mode 100644 index 0000000000000..477e22a34ed91 --- /dev/null +++ b/onnxruntime/core/providers/webgpu/vendor/intel/math/matmul_algorithm_scheduler.h @@ -0,0 +1,24 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include "core/providers/webgpu/math/matmul_algorithm_scheduler.h" + +namespace onnxruntime { +namespace webgpu { +namespace intel { + +class IntelMatMulAlgorithmScheduler final : public MatMulAlgorithmScheduler { + public: + IntelMatMulAlgorithmScheduler() = default; + explicit IntelMatMulAlgorithmScheduler(SplitKConfig split_k_config); + + protected: + std::optional SelectVendorAlgorithm( + const MatMulAlgorithmSelectionParams& params) const override; +}; + +} // namespace intel +} // namespace webgpu +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/webgpu/vendor/intel/math/split_k_config.cc b/onnxruntime/core/providers/webgpu/vendor/intel/math/split_k_config.cc new file mode 100644 index 0000000000000..4c5f8f02e45a5 --- /dev/null +++ b/onnxruntime/core/providers/webgpu/vendor/intel/math/split_k_config.cc @@ -0,0 +1,41 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/providers/webgpu/vendor/intel/math/split_k_config.h" + +namespace onnxruntime { +namespace webgpu { +namespace intel { + +SplitKConfig CreateSplitKConfig(std::string_view architecture) { + // Disable Split-K on old Intel GPUs. + if (architecture == "gen-7" || architecture == "gen-8" || + architecture == "gen-9" || architecture == "gen-11") { + return {}; + } + + constexpr uint32_t max_batch_size = 8; + constexpr uint32_t split_dim_inner = 256; + constexpr uint32_t min_dim_inner_with_split_k = split_dim_inner * 2; + + if (architecture == "xe-2lpg" || architecture == "xe-2hpg" || + architecture == "gen-12hp") { + // These thresholds are verified on Intel discrete GPUs and Lunar Lake iGPUs. + return SplitKConfig{ + max_batch_size, split_dim_inner, min_dim_inner_with_split_k, {{768, 52.0}, {2304, 35.0}, {3072, 21.5}, {4096, 16.0}}}; + } + + if (architecture == "xe-3lpg") { + // These thresholds are verified on Intel Panther Lake iGPUs (12Xe). + return SplitKConfig{ + max_batch_size, split_dim_inner, min_dim_inner_with_split_k, {{768, 40.0}, {1792, 22.0}, {3072, 18.0}, {4096, 10.0}}}; + } + + // Default thresholds for newer Intel GPUs, chosen on a gen-12lp GPU with 32 EUs. + return SplitKConfig{ + max_batch_size, split_dim_inner, min_dim_inner_with_split_k, {{768, 20.0}, {1792, 13.0}, {3072, 8.0}, {4096, 6.0}}}; +} + +} // namespace intel +} // namespace webgpu +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/webgpu/vendor/intel/math/split_k_config.h b/onnxruntime/core/providers/webgpu/vendor/intel/math/split_k_config.h new file mode 100644 index 0000000000000..c2b826126222c --- /dev/null +++ b/onnxruntime/core/providers/webgpu/vendor/intel/math/split_k_config.h @@ -0,0 +1,18 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include + +#include "core/providers/webgpu/webgpu_utils.h" + +namespace onnxruntime { +namespace webgpu { +namespace intel { + +SplitKConfig CreateSplitKConfig(std::string_view architecture); + +} // namespace intel +} // namespace webgpu +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/webgpu/webgpu_context.cc b/onnxruntime/core/providers/webgpu/webgpu_context.cc index a9cf224f3d70a..f00412ab0328d 100644 --- a/onnxruntime/core/providers/webgpu/webgpu_context.cc +++ b/onnxruntime/core/providers/webgpu/webgpu_context.cc @@ -293,7 +293,7 @@ void WebGpuContext::Initialize(const WebGpuContextConfig& config) { program_mgr_ = std::make_unique(*this); // create split-k config - split_k_config_ = std::make_unique(adapter_info_); + split_k_config_ = std::make_unique(CreateSplitKConfig(adapter_info_)); // set query type #if !defined(__wasm__) diff --git a/onnxruntime/core/providers/webgpu/webgpu_execution_provider.cc b/onnxruntime/core/providers/webgpu/webgpu_execution_provider.cc index 57df6dd0417fa..6d01176469f73 100644 --- a/onnxruntime/core/providers/webgpu/webgpu_execution_provider.cc +++ b/onnxruntime/core/providers/webgpu/webgpu_execution_provider.cc @@ -617,6 +617,7 @@ WebGpuExecutionProvider::WebGpuExecutionProvider(int context_id, multi_rotary_cache_concat_offset_{config.multi_rotary_cache_concat_offset}, kv_cache_quantization_bits_{config.kv_cache_quantization_bits}, enable_matmul_fp32_accumulation_{config.enable_matmul_fp32_accumulation}, + forced_matmul_algorithm_{config.forced_matmul_algorithm}, recording_{std::make_unique()}, prepack_allocator_{CreateWebGpuAllocator( /*device_free=*/!context.HasDevice(), diff --git a/onnxruntime/core/providers/webgpu/webgpu_execution_provider.h b/onnxruntime/core/providers/webgpu/webgpu_execution_provider.h index 823d3ee416b4c..3eb4079bbfc4e 100644 --- a/onnxruntime/core/providers/webgpu/webgpu_execution_provider.h +++ b/onnxruntime/core/providers/webgpu/webgpu_execution_provider.h @@ -6,6 +6,7 @@ #include #include +#include #include #include #include @@ -17,6 +18,7 @@ #include "core/graph/constants.h" #include "core/providers/providers.h" #include "core/providers/webgpu/buffer_manager.h" +#include "core/providers/webgpu/math/matmul_algorithm.h" #include "core/providers/webgpu/session_buffer_pool.h" #if defined(ENABLE_PIX_FOR_WEBGPU_EP) @@ -63,6 +65,7 @@ struct WebGpuExecutionProviderConfig { // This is the single line that decides the shipped default for the // "enableMatmulFp32Accumulation" provider option. bool enable_matmul_fp32_accumulation{false}; + std::optional forced_matmul_algorithm; std::vector force_cpu_node_names{}; }; @@ -130,6 +133,7 @@ class WebGpuExecutionProvider : public IExecutionProvider { uint32_t KvCacheQuantizationBits() const { return kv_cache_quantization_bits_; } bool KvCacheQuantizationEnabled() const { return kv_cache_quantization_bits_ != 0; } bool EnableMatmulFp32Accumulation() const { return enable_matmul_fp32_accumulation_; } + std::optional ForcedMatMulAlgorithm() const { return forced_matmul_algorithm_; } #if defined(ORT_USE_EP_API_ADAPTERS) onnxruntime::ep::adapter::Logger& GetEpLogger() const; @@ -152,6 +156,7 @@ class WebGpuExecutionProvider : public IExecutionProvider { uint32_t multi_rotary_cache_concat_offset_ = 0; uint32_t kv_cache_quantization_bits_ = 0; bool enable_matmul_fp32_accumulation_ = false; + std::optional forced_matmul_algorithm_; std::unordered_map graph_id_to_run_count_; // Required regular runs before graph capture for any necessary allocations. const int min_num_runs_before_graph_capture_ = 0; diff --git a/onnxruntime/core/providers/webgpu/webgpu_provider_factory.cc b/onnxruntime/core/providers/webgpu/webgpu_provider_factory.cc index b68392874aeec..1d1e99481018a 100644 --- a/onnxruntime/core/providers/webgpu/webgpu_provider_factory.cc +++ b/onnxruntime/core/providers/webgpu/webgpu_provider_factory.cc @@ -132,6 +132,14 @@ WebGpuExecutionProviderConfig ParseEpConfig(const ConfigOptions& config_options) } } + if (std::string forced_matmul_algorithm_str; + config_options.TryGetConfigEntry(kForceMatMulAlgorithm, forced_matmul_algorithm_str)) { + webgpu_ep_config.forced_matmul_algorithm = ParseMatMulAlgorithm(forced_matmul_algorithm_str); + ORT_ENFORCE(webgpu_ep_config.forced_matmul_algorithm.has_value(), + "Invalid forced MatMul algorithm: ", forced_matmul_algorithm_str, + ". Must be one of: subgroup_matrix, naive, subgroup, packed, packed_split_k."); + } + // parse force CPU node names // The force CPU node names are separated by EOL (\n or \r\n) in the config entry. // each line is a node name that will be forced to run on CPU. diff --git a/onnxruntime/core/providers/webgpu/webgpu_provider_options.h b/onnxruntime/core/providers/webgpu/webgpu_provider_options.h index 76317b53455c5..1d91cd2d6783c 100644 --- a/onnxruntime/core/providers/webgpu/webgpu_provider_options.h +++ b/onnxruntime/core/providers/webgpu/webgpu_provider_options.h @@ -32,6 +32,8 @@ constexpr const char* kKvCacheQuantizationBits = "ep.webgpuexecutionprovider.kvC // Today this covers MatMulNBits and its fused variants; the unquantized MatMul family is planned // as follow-up work under the same option. constexpr const char* kEnableMatmulFp32Accumulation = "ep.webgpuexecutionprovider.enableMatmulFp32Accumulation"; +// Internal test option for selecting a concrete unquantized MatMul implementation. +constexpr const char* kForceMatMulAlgorithm = "ep.webgpuexecutionprovider.forceMatmulAlgorithm"; constexpr const char* kDawnProcTable = "ep.webgpuexecutionprovider.dawnProcTable"; diff --git a/onnxruntime/core/providers/webgpu/webgpu_utils.cc b/onnxruntime/core/providers/webgpu/webgpu_utils.cc index 62d7b9934e164..a22b93a0b3cab 100644 --- a/onnxruntime/core/providers/webgpu/webgpu_utils.cc +++ b/onnxruntime/core/providers/webgpu/webgpu_utils.cc @@ -3,7 +3,10 @@ #include "core/providers/webgpu/webgpu_utils.h" #include +#include + #include "core/providers/webgpu/shader_variable.h" +#include "core/providers/webgpu/vendor/intel/math/split_k_config.h" namespace onnxruntime { namespace webgpu { @@ -25,55 +28,32 @@ TensorShape ReduceShapeByComponents(const TensorShape& shape, int64_t components return TensorShape(shape_vector); } -SplitKConfig::SplitKConfig(const wgpu::AdapterInfo& adapter_info) { - if (adapter_info.vendor == std::string_view{"intel"}) { - // Disable Split-K on old Intel GPUs. - if (adapter_info.architecture == std::string_view{"gen-7"} || - adapter_info.architecture == std::string_view{"gen-8"} || - adapter_info.architecture == std::string_view{"gen-9"} || - adapter_info.architecture == std::string_view{"gen-11"}) { - enable_split_k_ = false; - } else if (adapter_info.architecture == std::string_view{"xe-2lpg"} || - adapter_info.architecture == std::string_view{"xe-2hpg"} || - adapter_info.architecture == std::string_view{"gen-12hp"}) { - // Below thresholds are only verified on Intel discrete GPUs and Lunar Lake iGPUs. - enable_split_k_ = true; - - max_batch_size_ = 8; - split_dim_inner_ = 256; - min_dim_inner_with_split_k_ = split_dim_inner_ * 2; - - configs_per_dim_inner_range_.emplace_back(768, 52.0); - configs_per_dim_inner_range_.emplace_back(2304, 35.0); - configs_per_dim_inner_range_.emplace_back(3072, 21.5); - configs_per_dim_inner_range_.emplace_back(4096, 16.0); - } else if (adapter_info.architecture == std::string_view{"xe-3lpg"}) { - // Below thresholds are only verified on Intel Panther Lake iGPUs (12Xe). - enable_split_k_ = true; - - max_batch_size_ = 8; - split_dim_inner_ = 256; - min_dim_inner_with_split_k_ = split_dim_inner_ * 2; - - configs_per_dim_inner_range_.emplace_back(768, 40.0); - configs_per_dim_inner_range_.emplace_back(1792, 22.0); - configs_per_dim_inner_range_.emplace_back(3072, 18.0); - configs_per_dim_inner_range_.emplace_back(4096, 10.0); - } else { - // Below are the default thresholds on newer Intel GPUs. These values are chosen on - // Intel "gen-12lp" GPU with 32EUs. - enable_split_k_ = true; - - max_batch_size_ = 8; - split_dim_inner_ = 256; - min_dim_inner_with_split_k_ = split_dim_inner_ * 2; - - configs_per_dim_inner_range_.emplace_back(768, 20.0); - configs_per_dim_inner_range_.emplace_back(1792, 13.0); - configs_per_dim_inner_range_.emplace_back(3072, 8.0); - configs_per_dim_inner_range_.emplace_back(4096, 6.0); - } +SplitKConfig::SplitKConfig( + uint32_t max_batch_size, + uint32_t split_dim_inner, + uint32_t min_dim_inner_with_split_k, + std::initializer_list> configs_per_dim_inner_range) + : enable_split_k_{true}, + split_dim_inner_{split_dim_inner}, + min_dim_inner_with_split_k_{min_dim_inner_with_split_k}, + max_batch_size_{max_batch_size} { + configs_per_dim_inner_range_.reserve(configs_per_dim_inner_range.size()); + for (const auto& [max_dim_inner, rate] : configs_per_dim_inner_range) { + configs_per_dim_inner_range_.emplace_back(max_dim_inner, rate); + } +} + +SplitKConfig CreateSplitKConfig(const wgpu::AdapterInfo& adapter_info) { + return CreateSplitKConfig( + std::string_view{adapter_info.vendor}, + std::string_view{adapter_info.architecture}); +} + +SplitKConfig CreateSplitKConfig(std::string_view vendor, std::string_view architecture) { + if (vendor == "intel") { + return intel::CreateSplitKConfig(architecture); } + return {}; } SplitKConfig::ConfigAtRange::ConfigAtRange(uint32_t max_dim_inner, double rate) @@ -88,9 +68,9 @@ bool SplitKConfig::UseSplitK( bool is_vec4, ActivationKind activation_kind, uint64_t batch_size, - uint32_t dim_a_outer, - uint32_t dim_b_outer, - uint32_t dim_inner, + uint64_t dim_a_outer, + uint64_t dim_b_outer, + uint64_t dim_inner, bool is_channels_last) const { if (!enable_split_k_) { return false; diff --git a/onnxruntime/core/providers/webgpu/webgpu_utils.h b/onnxruntime/core/providers/webgpu/webgpu_utils.h index d4bb245e3e9e8..587bdd9f45d83 100644 --- a/onnxruntime/core/providers/webgpu/webgpu_utils.h +++ b/onnxruntime/core/providers/webgpu/webgpu_utils.h @@ -4,6 +4,11 @@ #pragma once #include +#include +#include +#include +#include + #include "core/common/common.h" #include "core/framework/tensor.h" #include "core/framework/tensor_shape.h" @@ -103,11 +108,17 @@ inline Tensor CreateTensorView(const Tensor& tensor, MLDataType new_data_type, c */ class SplitKConfig { public: - explicit SplitKConfig(const wgpu::AdapterInfo& adapter_info); + SplitKConfig() = default; + SplitKConfig( + uint32_t max_batch_size, + uint32_t split_dim_inner, + uint32_t min_dim_inner_with_split_k, + std::initializer_list> configs_per_dim_inner_range); bool UseSplitK( bool is_vec4, ActivationKind activation_kind, uint64_t batch_size, - uint32_t dim_a_outer, uint32_t dim_b_outer, uint32_t dim_inner, bool is_channels_last = true) const; + uint64_t dim_a_outer, uint64_t dim_b_outer, uint64_t dim_inner, + bool is_channels_last = true) const; uint32_t GetSplitDimInner() const; @@ -127,6 +138,9 @@ class SplitKConfig { std::vector configs_per_dim_inner_range_; }; +SplitKConfig CreateSplitKConfig(const wgpu::AdapterInfo& adapter_info); +SplitKConfig CreateSplitKConfig(std::string_view vendor, std::string_view architecture); + /** * Generates WGSL (WebGPU Shading Language) code for performing an atomic add operation * on a non-integer value (e.g., floating-point) in a shader. diff --git a/onnxruntime/test/providers/webgpu/matmul_algorithm_scheduler_test.cc b/onnxruntime/test/providers/webgpu/matmul_algorithm_scheduler_test.cc new file mode 100644 index 0000000000000..7313e22c040b3 --- /dev/null +++ b/onnxruntime/test/providers/webgpu/matmul_algorithm_scheduler_test.cc @@ -0,0 +1,487 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include +#include + +#include "gtest/gtest.h" + +#include "core/providers/webgpu/math/matmul_algorithm.h" +#include "core/providers/webgpu/math/matmul_algorithm_scheduler.h" +#include "core/providers/webgpu/math/matmul_compute_dispatcher.h" +#include "core/providers/webgpu/vendor/intel/math/gemm_subgroup_utils.h" +#include "core/providers/webgpu/vendor/intel/math/matmul_algorithm_scheduler.h" +#include "core/providers/webgpu/vendor/intel/math/split_k_config.h" + +namespace onnxruntime { +namespace webgpu { +namespace test { + +namespace { + +static_assert(!std::is_copy_constructible_v); +static_assert(!std::is_copy_assignable_v); +static_assert(!std::is_move_constructible_v); +static_assert(!std::is_move_assignable_v); +static_assert(!std::is_copy_constructible_v); +static_assert(!std::is_copy_assignable_v); +static_assert(!std::is_move_constructible_v); +static_assert(!std::is_move_assignable_v); +static_assert(!std::is_copy_constructible_v); +static_assert(!std::is_copy_assignable_v); +static_assert(!std::is_move_constructible_v); +static_assert(!std::is_move_assignable_v); + +class AlwaysPackedVendorScheduler final : public MatMulAlgorithmScheduler { + protected: + std::optional SelectVendorAlgorithm( + const MatMulAlgorithmSelectionParams& /*params*/) const override { + return MatMulAlgorithm::Packed; + } +}; + +class TunedPackedVendorScheduler final : public MatMulAlgorithmScheduler { + protected: + std::optional SelectVendorAlgorithm( + const MatMulAlgorithmSelectionParams& /*params*/) const override { + return MatMulAlgorithm::Naive; + } + + std::optional SelectVendorConfiguration( + MatMulAlgorithm algorithm, + const MatMulAlgorithmSelectionParams& params) const override { + const bool is_packed_algorithm = algorithm == MatMulAlgorithm::Packed || + algorithm == MatMulAlgorithm::PackedSplitK; + if (!is_packed_algorithm || + params.adapter_architecture != "test-architecture" || + params.a_data_type != 10 || params.b_data_type != 10) { + return std::nullopt; + } + + MatMulPackedConfiguration configuration{}; + configuration.workgroup_size = {16, 4, 1}; + configuration.elements_per_thread = {4, 2, 1}; + configuration.tile_inner = 16; + configuration.split_dim_inner = 128; + return configuration; + } +}; + +class VendorSplitKThresholdScheduler final : public MatMulAlgorithmScheduler { + public: + VendorSplitKThresholdScheduler() + : MatMulAlgorithmScheduler{SplitKConfig{ + /*max_batch_size=*/8, + /*split_dim_inner=*/128, + /*min_dim_inner_with_split_k=*/256, + {{4096, 16.0}}}} {} + + protected: + std::optional SelectVendorAlgorithm( + const MatMulAlgorithmSelectionParams& params) const override { + if (params.batch_size <= 16 && params.is_vec4 && + !params.deterministic_compute && !params.has_fused_activation && + (!params.has_bias || params.is_channels_last)) { + return MatMulAlgorithm::PackedSplitK; + } + return std::nullopt; + } +}; + +} // namespace + +TEST(MatMulAlgorithmParsingTest, RoundTripsEveryAlgorithmName) { + struct TestCase { + std::string_view name; + MatMulAlgorithm algorithm; + }; + + constexpr TestCase test_cases[] = { + {"subgroup_matrix", MatMulAlgorithm::SubgroupMatrix}, + {"naive", MatMulAlgorithm::Naive}, + {"subgroup", MatMulAlgorithm::Subgroup}, + {"packed", MatMulAlgorithm::Packed}, + {"packed_split_k", MatMulAlgorithm::PackedSplitK}, + }; + + for (const auto& test_case : test_cases) { + SCOPED_TRACE(test_case.name); + EXPECT_EQ(ParseMatMulAlgorithm(test_case.name), test_case.algorithm); + EXPECT_EQ(MatMulAlgorithmName(test_case.algorithm), test_case.name); + } +} + +TEST(MatMulAlgorithmParsingTest, RejectsUnknownAlgorithmName) { + EXPECT_EQ(ParseMatMulAlgorithm("unknown"), std::nullopt); +} + +TEST(SplitKConfigTest, IntelArchitectureProfilesPreserveCurrentBoundaries) { + const SplitKConfig discrete_config = intel::CreateSplitKConfig("xe-2lpg"); + EXPECT_EQ(discrete_config.GetSplitDimInner(), 256u); + EXPECT_TRUE(discrete_config.UseSplitK( + /*is_vec4=*/true, ActivationKind::None, /*batch_size=*/1, + /*dim_a_outer=*/192, /*dim_b_outer=*/160, /*dim_inner=*/1024)); + + const SplitKConfig xe3_config = intel::CreateSplitKConfig("xe-3lpg"); + EXPECT_FALSE(xe3_config.UseSplitK( + /*is_vec4=*/true, ActivationKind::None, /*batch_size=*/1, + /*dim_a_outer=*/192, /*dim_b_outer=*/160, /*dim_inner=*/1024)); + EXPECT_TRUE(xe3_config.UseSplitK( + /*is_vec4=*/true, ActivationKind::None, /*batch_size=*/1, + /*dim_a_outer=*/128, /*dim_b_outer=*/128, /*dim_inner=*/1024)); + + const SplitKConfig default_config = intel::CreateSplitKConfig("gen-12lp"); + EXPECT_FALSE(default_config.UseSplitK( + /*is_vec4=*/true, ActivationKind::None, /*batch_size=*/1, + /*dim_a_outer=*/128, /*dim_b_outer=*/128, /*dim_inner=*/1024)); + + const SplitKConfig legacy_config = intel::CreateSplitKConfig("gen-9"); + EXPECT_EQ(legacy_config.GetSplitDimInner(), 0u); + EXPECT_FALSE(legacy_config.UseSplitK( + /*is_vec4=*/true, ActivationKind::None, /*batch_size=*/1, + /*dim_a_outer=*/1, /*dim_b_outer=*/1, /*dim_inner=*/1024)); +} + +TEST(SplitKConfigTest, FactoryRoutesOnlySupportedVendorProfiles) { + EXPECT_EQ(CreateSplitKConfig("intel", "xe-2lpg").GetSplitDimInner(), 256u); + EXPECT_EQ(CreateSplitKConfig("nvidia", "pascal").GetSplitDimInner(), 0u); +} + +TEST(MatMulAlgorithmSchedulerTest, ForcedAlgorithmTakesPrecedence) { + AlwaysPackedVendorScheduler scheduler; + MatMulAlgorithmSelectionParams params{}; + params.can_use_subgroup_matrix = true; + + EXPECT_EQ(scheduler.Select(params, MatMulAlgorithm::Naive), MatMulAlgorithm::Naive); +} + +TEST(MatMulAlgorithmSchedulerTest, ReevaluatesSelectionForEachRuntimeShape) { + MatMulAlgorithmScheduler scheduler; + MatMulAlgorithmSelectionParams params{}; + params.m = 4; + params.packed_m = 4; + params.n = 7; + params.k = 7; + + EXPECT_EQ(scheduler.CreateExecutionPlan(params).algorithm, MatMulAlgorithm::Naive); + + params.m = 64; + params.packed_m = 64; + params.n = 64; + params.k = 64; + EXPECT_EQ(scheduler.CreateExecutionPlan(params).algorithm, MatMulAlgorithm::Packed); +} + +TEST(MatMulAlgorithmSchedulerTest, VendorCanTuneForcedAlgorithmConfiguration) { + TunedPackedVendorScheduler scheduler; + MatMulAlgorithmSelectionParams params{}; + params.adapter_architecture = "test-architecture"; + params.a_data_type = 10; + params.b_data_type = 10; + + const MatMulExecutionPlan plan = + scheduler.CreateExecutionPlan(params, MatMulAlgorithm::Packed); + + EXPECT_EQ(plan.algorithm, MatMulAlgorithm::Packed); + const auto& configuration = + std::get(plan.configuration); + EXPECT_EQ(configuration.workgroup_size, + (std::array{16, 4, 1})); + EXPECT_EQ(configuration.elements_per_thread, + (std::array{4, 2, 1})); + EXPECT_EQ(configuration.tile_inner, 16u); + EXPECT_EQ(configuration.split_dim_inner, 128u); +} + +TEST(MatMulAlgorithmSchedulerTest, VendorCanSetIndependentSplitKThresholds) { + VendorSplitKThresholdScheduler scheduler; + MatMulAlgorithmSelectionParams params{}; + params.m = 32; + params.n = 64; + params.k = 1024; + params.batch_size = 16; + params.packed_batch_size = 16; + params.is_vec4 = true; + params.is_channels_last = true; + + const MatMulExecutionPlan plan = scheduler.CreateExecutionPlan(params); + EXPECT_EQ(plan.algorithm, MatMulAlgorithm::PackedSplitK); + EXPECT_EQ(std::get(plan.configuration).split_dim_inner, 128u); + + params.batch_size = 17; + params.packed_batch_size = 17; + EXPECT_EQ(scheduler.Select(params), MatMulAlgorithm::Packed); +} + +TEST(MatMulAlgorithmSchedulerTest, InjectedSplitKPolicyCreatesPackedSplitKPlan) { + intel::IntelMatMulAlgorithmScheduler scheduler{intel::CreateSplitKConfig("xe-2lpg")}; + MatMulAlgorithmSelectionParams params{}; + params.m = 192; + params.packed_m = 192; + params.n = 160; + params.k = 1024; + params.is_vec4 = true; + + const MatMulExecutionPlan plan = scheduler.CreateExecutionPlan(params); + + EXPECT_EQ(plan.algorithm, MatMulAlgorithm::PackedSplitK); + const auto& configuration = + std::get(plan.configuration); + EXPECT_EQ(configuration.split_dim_inner, 256u); +} + +TEST(MatMulAlgorithmSchedulerTest, CommonPackedConfigurationPreservesCurrentTuning) { + MatMulAlgorithmScheduler scheduler; + MatMulAlgorithmSelectionParams params{}; + params.m = 8; + params.packed_m = 8; + params.n = 64; + params.k = 64; + + const MatMulExecutionPlan small_m_plan = scheduler.CreateExecutionPlan(params); + const auto& small_m_configuration = + std::get(small_m_plan.configuration); + EXPECT_EQ(small_m_configuration.workgroup_size, + (std::array{8, 8, 1})); + EXPECT_EQ(small_m_configuration.elements_per_thread, + (std::array{4, 1, 1})); + EXPECT_EQ(small_m_configuration.tile_inner, 32u); + + params.m = 9; + params.packed_m = 9; + const MatMulExecutionPlan large_m_plan = scheduler.CreateExecutionPlan(params); + const auto& large_m_configuration = + std::get(large_m_plan.configuration); + EXPECT_EQ(large_m_configuration.elements_per_thread, + (std::array{4, 4, 1})); +} + +TEST(MatMulAlgorithmConfigurationTest, PackedConfigurationKeepsBatchAxesUntiled) { + MatMulPackedConfiguration configuration{}; + EXPECT_TRUE(IsMatMulPackedConfigurationValid(configuration, /*use_split_k=*/false)); + + configuration.workgroup_size[2] = 2; + EXPECT_FALSE(IsMatMulPackedConfigurationValid(configuration, /*use_split_k=*/false)); + + configuration.workgroup_size[2] = 1; + configuration.elements_per_thread[2] = 2; + EXPECT_FALSE(IsMatMulPackedConfigurationValid(configuration, /*use_split_k=*/false)); +} + +TEST(MatMulAlgorithmConfigurationTest, SubgroupConfigurationCarriesSelectedSize) { + MatMulAlgorithmScheduler scheduler; + MatMulAlgorithmSelectionParams params{}; + params.subgroup_size = 16; + + const auto plan = scheduler.CreateExecutionPlan(params, MatMulAlgorithm::Subgroup); + const auto* configuration = std::get_if(&plan.configuration); + ASSERT_NE(configuration, nullptr); + EXPECT_EQ(configuration->subgroup_size, 16u); +} + +TEST(MatMulAlgorithmConfigurationTest, SubgroupSizeSelectionRequiresSupportedSize) { + EXPECT_EQ(intel::SelectMatMulSubgroupSize(8, 8, false), 8u); + EXPECT_EQ(intel::SelectMatMulSubgroupSize(16, 16, false), 16u); + EXPECT_EQ(intel::SelectMatMulSubgroupSize(32, 32, false), 32u); + EXPECT_FALSE(intel::SelectMatMulSubgroupSize(64, 64, false).has_value()); + EXPECT_FALSE(intel::SelectMatMulSubgroupSize(8, 32, false).has_value()); + EXPECT_EQ(intel::SelectMatMulSubgroupSize(8, 32, true), 32u); + EXPECT_EQ(intel::SelectMatMulSubgroupSize(8, 16, true), 16u); + EXPECT_FALSE(intel::SelectMatMulSubgroupSize(4, 4, true).has_value()); +} + +TEST(MatMulAlgorithmConfigurationTest, SplitKConfigurationRequiresTileAlignedSplits) { + MatMulPackedConfiguration configuration{}; + configuration.tile_inner = 32; + configuration.split_dim_inner = 64; + EXPECT_TRUE(IsMatMulPackedConfigurationValid(configuration, /*use_split_k=*/true)); + + configuration.split_dim_inner = 48; + EXPECT_FALSE(IsMatMulPackedConfigurationValid(configuration, /*use_split_k=*/true)); +} + +TEST(MatMulAlgorithmConfigurationTest, PackedDispatchArithmeticIsOverflowSafe) { + const auto maximum_tuning_dispatch = TryGetMatMulPackedDispatchGroupCount( + std::numeric_limits::max(), + std::numeric_limits::max(), + std::numeric_limits::max()); + ASSERT_TRUE(maximum_tuning_dispatch.has_value()); + EXPECT_EQ(*maximum_tuning_dispatch, 1u); + + EXPECT_EQ(TryGetMatMulPackedDispatchGroupCount( + std::numeric_limits::max(), 1, 1), + std::nullopt); +} + +TEST(MatMulAlgorithmSchedulerTest, CommonFallbackPrefersSubgroupMatrix) { + MatMulAlgorithmScheduler scheduler{intel::CreateSplitKConfig("xe-2lpg")}; + MatMulAlgorithmSelectionParams params{}; + params.m = 64; + params.packed_m = 64; + params.n = 64; + params.k = 1024; + params.can_use_subgroup_matrix = true; + params.is_vec4 = true; + + EXPECT_EQ(scheduler.Select(params), MatMulAlgorithm::SubgroupMatrix); +} + +TEST(MatMulAlgorithmSchedulerTest, VendorPolicyPrecedesCommonHeuristics) { + AlwaysPackedVendorScheduler scheduler; + MatMulAlgorithmSelectionParams params{}; + params.n = 4; + params.k = 4; + params.can_use_subgroup_matrix = true; + + EXPECT_EQ(scheduler.Select(params), MatMulAlgorithm::Packed); +} + +TEST(MatMulAlgorithmSchedulerTest, ZeroContractionDimensionPrecedesVendorPolicy) { + AlwaysPackedVendorScheduler scheduler; + MatMulAlgorithmSelectionParams params{}; + params.n = 8; + params.k = 0; + + EXPECT_EQ(scheduler.Select(params), MatMulAlgorithm::Naive); +} + +TEST(MatMulAlgorithmSchedulerTest, NaiveUsesStrictSmallDimensionBoundaries) { + MatMulAlgorithmScheduler scheduler; + + MatMulAlgorithmSelectionParams small_params{}; + small_params.n = 7; + small_params.k = 7; + EXPECT_EQ(scheduler.Select(small_params), MatMulAlgorithm::Naive); + + MatMulAlgorithmSelectionParams n_boundary = small_params; + n_boundary.n = 8; + EXPECT_EQ(scheduler.Select(n_boundary), MatMulAlgorithm::Packed); + + MatMulAlgorithmSelectionParams k_boundary = small_params; + k_boundary.k = 8; + EXPECT_EQ(scheduler.Select(k_boundary), MatMulAlgorithm::Packed); +} + +TEST(MatMulAlgorithmSchedulerTest, ZeroContractionDimensionUsesNaive) { + MatMulAlgorithmScheduler scheduler; + MatMulAlgorithmSelectionParams params{}; + params.n = 8; + params.k = 0; + + EXPECT_EQ(scheduler.Select(params), MatMulAlgorithm::Naive); +} + +TEST(MatMulAlgorithmSchedulerTest, IntelSchedulerAppliesCurrentVendorRule) { + intel::IntelMatMulAlgorithmScheduler scheduler; + MatMulAlgorithmSelectionParams params{}; + params.m = 64; + params.n = 512; + params.k = 32; + params.has_subgroup_capability = true; + + EXPECT_EQ(scheduler.Select(params), MatMulAlgorithm::Subgroup); + + params.n = 511; + EXPECT_EQ(scheduler.Select(params), MatMulAlgorithm::Packed); +} + +TEST(MatMulAlgorithmSchedulerTest, IntelSchedulerPreservesSubgroupMatrixPrecedence) { + intel::IntelMatMulAlgorithmScheduler scheduler; + MatMulAlgorithmSelectionParams params{}; + params.m = 64; + params.n = 512; + params.k = 32; + params.can_use_subgroup_matrix = true; + params.has_subgroup_capability = true; + + EXPECT_EQ(scheduler.Select(params), MatMulAlgorithm::SubgroupMatrix); +} + +TEST(MatMulAlgorithmSchedulerTest, SplitKPrecedesPackedFallback) { + MatMulAlgorithmScheduler scheduler{intel::CreateSplitKConfig("xe-2lpg")}; + MatMulAlgorithmSelectionParams params{}; + params.m = 64; + params.packed_m = 64; + params.n = 64; + params.k = 1024; + params.is_vec4 = true; + EXPECT_EQ(scheduler.Select(params), MatMulAlgorithm::PackedSplitK); + + MatMulAlgorithmScheduler disabled_scheduler; + EXPECT_EQ(disabled_scheduler.Select(params), MatMulAlgorithm::Packed); +} + +TEST(MatMulAlgorithmPrerequisiteTest, SplitKRejectsEachHardConstraint) { + MatMulAlgorithmPrerequisites prerequisites{}; + prerequisites.has_nonzero_k = true; + prerequisites.split_k_configured = true; + prerequisites.is_vec4 = true; + prerequisites.split_k_bias_layout_supported = true; + EXPECT_TRUE(MeetsMatMulAlgorithmPrerequisites(MatMulAlgorithm::PackedSplitK, prerequisites)); + + prerequisites.deterministic_compute = true; + EXPECT_FALSE(MeetsMatMulAlgorithmPrerequisites(MatMulAlgorithm::PackedSplitK, prerequisites)); + prerequisites.deterministic_compute = false; + + prerequisites.is_vec4 = false; + EXPECT_FALSE(MeetsMatMulAlgorithmPrerequisites(MatMulAlgorithm::PackedSplitK, prerequisites)); + prerequisites.is_vec4 = true; + + prerequisites.has_fused_activation = true; + EXPECT_FALSE(MeetsMatMulAlgorithmPrerequisites(MatMulAlgorithm::PackedSplitK, prerequisites)); + prerequisites.has_fused_activation = false; + + prerequisites.split_k_bias_layout_supported = false; + EXPECT_FALSE(MeetsMatMulAlgorithmPrerequisites(MatMulAlgorithm::PackedSplitK, prerequisites)); +} + +TEST(MatMulAlgorithmPrerequisiteTest, PackedAlgorithmsRejectZeroContractionDimension) { + MatMulAlgorithmPrerequisites prerequisites{}; + prerequisites.split_k_configured = true; + prerequisites.is_vec4 = true; + prerequisites.split_k_bias_layout_supported = true; + + EXPECT_FALSE(MeetsMatMulAlgorithmPrerequisites(MatMulAlgorithm::Packed, prerequisites)); + EXPECT_FALSE(MeetsMatMulAlgorithmPrerequisites(MatMulAlgorithm::PackedSplitK, prerequisites)); + + prerequisites.has_nonzero_k = true; + EXPECT_TRUE(MeetsMatMulAlgorithmPrerequisites(MatMulAlgorithm::Packed, prerequisites)); + EXPECT_TRUE(MeetsMatMulAlgorithmPrerequisites(MatMulAlgorithm::PackedSplitK, prerequisites)); +} + +TEST(MatMulAlgorithmPrerequisiteTest, IntelAVec4RequiresCompatibleRowsPerThread) { + EXPECT_FALSE(intel::CanUseAVec4CooperativeLoad(intel::gpu_arch::kXe3Lpg, 32, 1)); + EXPECT_FALSE(intel::CanUseAVec4CooperativeLoad(intel::gpu_arch::kXe3Lpg, 32, 2)); + EXPECT_TRUE(intel::CanUseAVec4CooperativeLoad(intel::gpu_arch::kXe3Lpg, 32, 4)); + EXPECT_FALSE(intel::CanUseAVec4CooperativeLoad(intel::gpu_arch::kXe3Lpg, 31, 4)); + EXPECT_FALSE(intel::CanUseAVec4CooperativeLoad(intel::gpu_arch::kXeLpg, 32, 4)); +} + +TEST(MatMulAlgorithmPrerequisiteTest, IntelCapabilityDoesNotIncludeAutomaticThresholds) { + MatMulAlgorithmPrerequisites prerequisites{}; + prerequisites.has_subgroup_capability = true; + prerequisites.has_nonzero_k = true; + EXPECT_TRUE(MeetsMatMulAlgorithmPrerequisites(MatMulAlgorithm::Subgroup, prerequisites)); + + MatMulAlgorithmSelectionParams below_heuristic_threshold{}; + below_heuristic_threshold.m = 1; + below_heuristic_threshold.n = 1; + below_heuristic_threshold.k = 1; + below_heuristic_threshold.has_subgroup_capability = true; + intel::IntelMatMulAlgorithmScheduler scheduler; + EXPECT_NE(scheduler.Select(below_heuristic_threshold), MatMulAlgorithm::Subgroup); +} + +TEST(MatMulAlgorithmPrerequisiteTest, SubgroupRejectsZeroContractionDimension) { + MatMulAlgorithmPrerequisites prerequisites{}; + prerequisites.has_subgroup_capability = true; + + EXPECT_FALSE(MeetsMatMulAlgorithmPrerequisites(MatMulAlgorithm::Subgroup, prerequisites)); + + prerequisites.has_nonzero_k = true; + EXPECT_TRUE(MeetsMatMulAlgorithmPrerequisites(MatMulAlgorithm::Subgroup, prerequisites)); +} + +} // namespace test +} // namespace webgpu +} // namespace onnxruntime diff --git a/onnxruntime/test/providers/webgpu/matmul_large_test.cc b/onnxruntime/test/providers/webgpu/matmul_large_test.cc index 3375fc5e6b53b..ba8548ecb3431 100644 --- a/onnxruntime/test/providers/webgpu/matmul_large_test.cc +++ b/onnxruntime/test/providers/webgpu/matmul_large_test.cc @@ -3,12 +3,30 @@ #include #include +#include #include +#include +#include +#include +#include #include #include "gtest/gtest.h" -#include "test/providers/provider_test_utils.h" + +#include "core/graph/onnx_protobuf.h" +#include "core/providers/webgpu/math/matmul_algorithm.h" +#if !defined(ORT_USE_EP_API_ADAPTERS) +#include "core/providers/webgpu/math/subgroup_matrix_config.h" +#include "core/providers/webgpu/vendor/intel/math/gemm_subgroup_utils.h" +#include "core/providers/webgpu/webgpu_context.h" +#endif +#include "core/providers/webgpu/webgpu_provider_options.h" #include "test/common/tensor_op_test_utils.h" +#include "test/providers/provider_test_utils.h" +#include "test/test_environment.h" +#include "test/unittest_util/framework_test_utils.h" +#include "test/util/include/asserts.h" +#include "test/util/include/inference_session_wrapper.h" #include "default_providers.h" namespace onnxruntime { @@ -72,15 +90,98 @@ static void ComputeExpectedResult(std::initializer_list a_dims, if (!b_is_vector) output_dims.push_back(N); } +static std::optional GetForcedAlgorithmUnsupportedReason( + const IExecutionProvider& ep, + webgpu::MatMulAlgorithm algorithm) { +#if defined(ORT_USE_EP_API_ADAPTERS) + ORT_UNUSED_PARAMETER(ep); + switch (algorithm) { + case webgpu::MatMulAlgorithm::Subgroup: + case webgpu::MatMulAlgorithm::PackedSplitK: + case webgpu::MatMulAlgorithm::SubgroupMatrix: + return "hardware-specific forced MatMul tests require direct adapter capability inspection."; + case webgpu::MatMulAlgorithm::Naive: + case webgpu::MatMulAlgorithm::Packed: + return std::nullopt; + } + return std::nullopt; +#else + auto& context = webgpu::WebGpuContextFactory::GetContext(ep.GetDeviceId()); + + switch (algorithm) { + case webgpu::MatMulAlgorithm::Subgroup: { + if (!context.DeviceHasFeature(wgpu::FeatureName::Subgroups)) { + return "subgroup requires the WebGPU Subgroups feature."; + } + const auto& adapter_info = context.AdapterInfo(); + if (!webgpu::intel::SelectMatMulSubgroupSize( + adapter_info.subgroupMinSize, + adapter_info.subgroupMaxSize, + context.DeviceHasFeature(wgpu::FeatureName::SubgroupSizeControl)) + .has_value()) { + return "subgroup requires a supported subgroup size (8, 16, or 32)."; + } + break; + } + case webgpu::MatMulAlgorithm::PackedSplitK: + if (context.GetSplitKConfig().GetSplitDimInner() == 0) { + return "packed_split_k is not configured for the selected adapter."; + } + break; + case webgpu::MatMulAlgorithm::SubgroupMatrix: { + if (!context.DeviceHasFeature(wgpu::FeatureName::ChromiumExperimentalSubgroupMatrix)) { + return "subgroup_matrix requires the WebGPU subgroup-matrix feature."; + } + + const auto& adapter_info = context.AdapterInfo(); + const auto& device_configs = context.SubgroupMatrixConfigs(); + constexpr auto kF16 = wgpu::SubgroupMatrixComponentType::F16; + if (!webgpu::detail::SelectSubgroupMatrixConfigFromAdapterConfigs( + {device_configs.configs, device_configs.configCount}, + adapter_info.subgroupMinSize, adapter_info.subgroupMaxSize, + context.DeviceHasFeature(wgpu::FeatureName::SubgroupSizeControl), + {{kF16, kF16, 8, 16, 16, 32, false}})) { + return "subgroup_matrix requires an 8x16x16 F16 configuration with subgroup size 32."; + } + break; + } + case webgpu::MatMulAlgorithm::Naive: + case webgpu::MatMulAlgorithm::Packed: + break; + } + + return std::nullopt; +#endif +} + template void RunTestTyped(std::initializer_list a_dims, std::initializer_list b_dims, - bool b_is_constant = false) { + bool b_is_constant = false, + std::optional forced_algorithm = std::nullopt, + OpTester::ExpectResult expected_result = OpTester::ExpectResult::kExpectSuccess, + const char* expected_error = nullptr) { static_assert(std::is_same_v || std::is_same_v, "unexpected type for T"); - auto webgpu_ep = DefaultWebGpuExecutionProvider(); + std::unique_ptr webgpu_ep; + if (forced_algorithm.has_value()) { + ConfigOptions config_options{}; + const std::string algorithm_name{webgpu::MatMulAlgorithmName(*forced_algorithm)}; + ASSERT_STATUS_OK(config_options.AddConfigEntry( + webgpu::options::kForceMatMulAlgorithm, + algorithm_name.c_str())); + webgpu_ep = WebGpuExecutionProviderWithOptions(config_options); + } else { + webgpu_ep = DefaultWebGpuExecutionProvider(); + } if (!webgpu_ep) { GTEST_SKIP() << "WebGPU execution provider is not available."; } + if (forced_algorithm.has_value()) { + if (const auto reason = GetForcedAlgorithmUnsupportedReason(*webgpu_ep, *forced_algorithm); + reason.has_value()) { + GTEST_SKIP() << *reason; + } + } RandomValueGenerator random{1234}; std::vector a_vals(random.Gaussian(AsSpan(a_dims), 0.0f, 0.25f)); @@ -103,7 +204,9 @@ void RunTestTyped(std::initializer_list a_dims, std::initializer_list @@ -123,6 +226,10 @@ TEST(MatMulNaiveProgramTest, VectorExecution) { RunTestTyped({3}, {3}); } +TEST(MatMulNaiveProgramTest, ZeroContractionDimensionExecution) { + RunTestTyped({1, 0}, {0, 8}); +} + TEST(MatMulProgramTest, VectorFallbackExecution) { RunTestTyped({8}, {8, 3}); RunTestTyped({2, 8}, {8}); @@ -131,6 +238,152 @@ TEST(MatMulProgramTest, VectorFallbackExecution) { RunTestTyped({2, 2, 8}, {8}); } +TEST(WebGpuMatMulAlgorithmTest, RejectsUnknownForcedAlgorithm) { + ConfigOptions valid_config_options{}; + if (!WebGpuExecutionProviderWithOptions(valid_config_options)) { + GTEST_SKIP() << "WebGPU execution provider is unavailable in this build."; + } + + ConfigOptions config_options{}; + ASSERT_STATUS_OK(config_options.AddConfigEntry(webgpu::options::kForceMatMulAlgorithm, "unknown")); + EXPECT_THROW(WebGpuExecutionProviderWithOptions(config_options), OnnxRuntimeException); +} + +static std::string BuildDynamicMatMulModelBytes() { + ONNX_NAMESPACE::ModelProto model; + model.set_ir_version(ONNX_NAMESPACE::IR_VERSION); + auto* opset = model.add_opset_import(); + opset->set_domain(""); + opset->set_version(13); + + auto* graph = model.mutable_graph(); + graph->set_name("dynamic_matmul"); + + auto set_float_shape = [](ONNX_NAMESPACE::ValueInfoProto* value_info, + const char* first_dimension, + const char* second_dimension) { + auto* tensor_type = value_info->mutable_type()->mutable_tensor_type(); + tensor_type->set_elem_type(ONNX_NAMESPACE::TensorProto_DataType_FLOAT); + auto* shape = tensor_type->mutable_shape(); + shape->add_dim()->set_dim_param(first_dimension); + if (second_dimension != nullptr) { + shape->add_dim()->set_dim_param(second_dimension); + } else { + shape->add_dim()->set_dim_value(8); + } + }; + + auto* a = graph->add_input(); + a->set_name("A"); + set_float_shape(a, "M", "K"); + + auto* b = graph->add_input(); + b->set_name("B"); + set_float_shape(b, "K", nullptr); + + auto* y = graph->add_output(); + y->set_name("Y"); + set_float_shape(y, "M", nullptr); + + auto* node = graph->add_node(); + node->set_name("MatMul"); + node->set_op_type("MatMul"); + node->add_input("A"); + node->add_input("B"); + node->add_output("Y"); + + std::string bytes; + model.SerializeToString(&bytes); + return bytes; +} + +TEST(WebGpuMatMulAlgorithmTest, ForcedNaive) { + RunTestTyped({8, 8}, {8, 8}, false, webgpu::MatMulAlgorithm::Naive); +} + +TEST(WebGpuMatMulAlgorithmTest, ForcedPacked) { + RunTestTyped({2, 2}, {2, 2}, false, webgpu::MatMulAlgorithm::Packed); +} + +TEST(WebGpuMatMulAlgorithmTest, ForcedPackedRejectsZeroContractionDimension) { + RunTestTyped({1, 0}, {0, 1}, false, webgpu::MatMulAlgorithm::Packed, + OpTester::ExpectResult::kExpectFailure, + "MatMul algorithm packed"); +} + +TEST(WebGpuMatMulAlgorithmTest, ForcedSubgroup) { + RunTestTyped({8, 32}, {32, 64}, false, webgpu::MatMulAlgorithm::Subgroup); +} + +TEST(WebGpuMatMulAlgorithmTest, ForcedSubgroupRejectsZeroContractionDimension) { + RunTestTyped({1, 0}, {0, 1}, false, webgpu::MatMulAlgorithm::Subgroup, + OpTester::ExpectResult::kExpectFailure, + "MatMul algorithm subgroup"); +} + +TEST(WebGpuMatMulAlgorithmTest, ForcedPackedSplitK) { + RunTestTyped({1, 1024}, {1024, 16}, false, webgpu::MatMulAlgorithm::PackedSplitK); +} + +TEST(WebGpuMatMulAlgorithmTest, ForcedSubgroupMatrix) { + RunTestTyped({32, 16}, {16, 32}, false, webgpu::MatMulAlgorithm::SubgroupMatrix); +} + +TEST(WebGpuMatMulAlgorithmTest, ForcedSubgroupMatrixRejectsFloatInputs) { + RunTestTyped({32, 16}, {16, 32}, false, webgpu::MatMulAlgorithm::SubgroupMatrix, + OpTester::ExpectResult::kExpectFailure, + "MatMul algorithm subgroup_matrix"); +} + +TEST(WebGpuMatMulAlgorithmTest, ReselectsAlgorithmForEachDynamicShape) { + auto webgpu_ep = DefaultWebGpuExecutionProvider(); + if (!webgpu_ep) { + GTEST_SKIP() << "WebGPU execution provider is not available."; + } + + SessionOptions session_options; + session_options.session_logid = "WebGpuDynamicMatMulAlgorithmSelection"; + InferenceSessionWrapper session(session_options, GetEnvironment()); + ASSERT_STATUS_OK(session.RegisterExecutionProvider(std::move(webgpu_ep))); + + const std::string model_bytes = BuildDynamicMatMulModelBytes(); + ASSERT_STATUS_OK(session.Load(model_bytes.data(), static_cast(model_bytes.size()))); + ASSERT_STATUS_OK(session.Initialize()); + + const std::vector output_names{"Y"}; + auto run = [&](const TensorShape& a_shape, std::vector a_data, + const TensorShape& b_shape, std::vector b_data, + const TensorShape& expected_shape, float expected_value) { + OrtValue a_value; + OrtValue b_value; + CreateMLValue(a_shape.GetDims(), a_data.data(), OrtMemoryInfo(), &a_value); + CreateMLValue(b_shape.GetDims(), b_data.data(), OrtMemoryInfo(), &b_value); + + NameMLValMap feeds; + feeds.emplace("A", std::move(a_value)); + feeds.emplace("B", std::move(b_value)); + std::vector fetches; + ASSERT_STATUS_OK(session.Run(feeds, output_names, &fetches)); + ASSERT_EQ(fetches.size(), 1u); + const Tensor& output = fetches[0].Get(); + ASSERT_EQ(output.Shape(), expected_shape); + for (float value : output.DataAsSpan()) { + EXPECT_EQ(value, expected_value); + } + }; + + // The first invocation uses Packed. The second has K=0 and must reselect Naive on the same + // session and kernel; reusing the first execution plan would fail Packed's nonzero-K prerequisite. + run(TensorShape({8, 8}), std::vector(64, 1.0f), + TensorShape({8, 8}), std::vector(64, 1.0f), + TensorShape({8, 8}), 8.0f); + + const float unused_storage = 0.0f; + run(TensorShape({1, 0}), std::vector{unused_storage}, + TensorShape({0, 8}), std::vector{unused_storage}, + TensorShape({1, 8}), 0.0f); +} + // 2D aligned baseline shapes. TEST(MatMul_Large, DISABLED_Aligned) { RunBothTypes({128, 64}, {64, 1024}); @@ -159,7 +412,7 @@ TEST(MatMul_Large, DISABLED_Broadcast4D) { } // Batched B (true bmm): A [..., M, K] x B [..., K, N] with matching batch. On the -// Intel subgroup path each (A, B) slice is dispatched on z. Covers small and +// Subgroup path each (A, B) slice is dispatched on z. Covers small and // larger batch counts with tile-aligned per-slice shapes. TEST(MatMul_Large, DISABLED_BatchedB) { RunBothTypes({2, 128, 64}, {2, 64, 1024}); @@ -212,7 +465,7 @@ TEST(MatMul_Large, DISABLED_ConstantWeightOddN) { // Broadcasted batch dims that are NOT identical but share the same batch // *product* (A=[2,1,...], B=[1,2,...] -> [2,2,...]; A=[1,4,...], B=[4,1,...] -> // [4,4,...]). A product-only batch check would wrongly route these onto the -// Intel subgroup path, which pairs slice i of A with slice i of B and copies A's +// Subgroup path, which pairs slice i of A with slice i of B and copies A's // shape to the output - producing the wrong output shape and mismatched pairing. // N is even and every per-slice shape is tile-aligned, so only the // identical-batch-dims guard (not the odd-N or partial-tile fallbacks) keeps them