-
Notifications
You must be signed in to change notification settings - Fork 4.1k
[Native WebGPU] Add Matmul #24046
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Merged
Merged
[Native WebGPU] Add Matmul #24046
Changes from 4 commits
Commits
Show all changes
11 commits
Select commit
Hold shift + click to select a range
e246fb5
Implement MatMul support in WebGPU with shader code generation and ut…
vraspar 27aeb65
Refactor and Lint
vraspar 463adc9
Matmul test fix liting error
vraspar fa3e804
Fix output shape calculation and adjust offset in MatMul shader code
vraspar a051eb6
Rename MatMulNativeProgram to MatMulNaiveProgram and update parameter…
vraspar d466b08
Fix: MatMul shader code with correct offset calculations and fix type…
vraspar 5034fbf
Fix data type references in MatMul shader code and update shader inpu…
vraspar 1b52d26
Refactor MatMul code: update shape creation functions, improve type h…
vraspar 6de4d8a
Merge branch 'main' into vraspar/webgpu-native-matmul
vraspar 23d6042
Apply feedback, add comments
vraspar 1681e35
Add conditional compilation for WebGPU in matmul test cases
vraspar File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,233 @@ | ||
| // Copyright (c) Microsoft Corporation. All rights reserved. | ||
| // Licensed under the MIT License. | ||
|
|
||
| #include "core/common/inlined_containers.h" | ||
| #include "core/providers/webgpu/math/matmul.h" | ||
|
vraspar marked this conversation as resolved.
|
||
| #include "core/providers/cpu/tensor/utils.h" | ||
| #include "core/providers/webgpu/shader_helper.h" | ||
| #include "core/providers/webgpu/webgpu_supported_types.h" | ||
|
|
||
| #include "core/providers/webgpu/data_transfer.h" | ||
| namespace onnxruntime { | ||
| namespace webgpu { | ||
|
|
||
| ONNX_OPERATOR_VERSIONED_KERNEL_EX( | ||
| MatMul, | ||
| kOnnxDomain, | ||
| 1, 12, | ||
| kWebGpuExecutionProvider, | ||
| (*KernelDefBuilder::Create()) | ||
| .TypeConstraint("T", WebGpuSupportedNumberTypes()), | ||
| MatMul); | ||
|
|
||
| ONNX_OPERATOR_KERNEL_EX( | ||
| MatMul, | ||
| kOnnxDomain, | ||
| 13, | ||
| kWebGpuExecutionProvider, | ||
| (*KernelDefBuilder::Create()) | ||
| .TypeConstraint("T", WebGpuSupportedNumberTypes()), | ||
| MatMul); | ||
|
|
||
| std::string CalcResult(int components, int a_components, int output_number) { | ||
| std::ostringstream oss; | ||
| oss << "var a_data: a_value_t;\n"; | ||
| for (int i = 0; i < a_components; ++i) { | ||
| oss << "let b_data" << i << " = b[(b_offset + (k + " << i << ") * uniforms.N + col) / " << components << "];\n"; | ||
| } | ||
| for (int i = 0; i < output_number; ++i) { | ||
| oss << "a_data = a[(a_offset + (row + " << i << ") * uniforms.K + k) / " << a_components << "];\n"; | ||
|
|
||
| for (int j = 0; j < a_components; j++) { | ||
| oss << "values[" << i << "] = fma(b_value_t(a_data" << (a_components == 1 ? "" : "[" + std::to_string(j) + "]") << "), b_data" << j << ", values[" << i << "]);\n"; | ||
| } | ||
| } | ||
| return oss.str(); | ||
| } | ||
|
|
||
| Status MatMulNativeProgram::GenerateShaderCode(ShaderHelper& shader) const { | ||
| LOGS_DEFAULT(VERBOSE) << "MatMulNativeProgram: Start generating shader code"; | ||
| 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); | ||
|
|
||
| std::string process_bias; | ||
| if (has_bias_) { | ||
| shader.AddInput("bias", ShaderUsage::UseUniform); | ||
| process_bias = "value += output_value_t(bias[row +i]);"; | ||
|
vraspar marked this conversation as resolved.
Outdated
|
||
| } | ||
|
|
||
| const auto& output = shader.AddOutput("output", ShaderUsage::UseUniform | | ||
| ShaderUsage::UseIndicesTypeAlias | ShaderUsage::UseValueTypeAlias); | ||
| const auto& batch_dims = shader.AddIndices("batch_dims"); | ||
|
|
||
| int a_components = a.NumComponents(); | ||
| int components = b.NumComponents(); // components of N | ||
|
vraspar marked this conversation as resolved.
|
||
|
|
||
| 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_size_ != 2) { | ||
|
vraspar marked this conversation as resolved.
Outdated
|
||
| 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") << ";\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") << ";\n" | ||
| << "var values: array<output_value_t, " << output_number_ << ">;\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" | ||
| << " var value = values[i];\n" | ||
| << process_bias << "\n" | ||
| << " let cur_indices = output_indices_t(batch, row + i, col);\n" | ||
| << " let offset = " << output.IndicesToOffset("cur_indices") << ";\n" | ||
| << output.SetByOffset("offset / " + std::to_string(components), "value") | ||
| << "}\n"; | ||
|
|
||
| return Status::OK(); | ||
| } | ||
|
|
||
| Status MatMul::ComputeInternal(ComputeContext& context) const { | ||
| LOGS_DEFAULT(VERBOSE) << "Running MatMul WebGPU kernel"; | ||
|
|
||
| // calculate output shape | ||
| MatMulComputeHelper helper; | ||
| const auto* a = context.Input(0); | ||
| const auto* b = context.Input(1); | ||
|
|
||
| ORT_RETURN_IF_ERROR(helper.Compute(a->Shape(), b->Shape())); | ||
| auto* output_tensor = context.Output(0, helper.OutputShape()); | ||
|
|
||
| const uint32_t m = static_cast<uint32_t>(helper.M()); | ||
|
vraspar marked this conversation as resolved.
Outdated
|
||
| const uint32_t n = static_cast<uint32_t>(helper.N()); | ||
| const uint32_t k = static_cast<uint32_t>(helper.K()); | ||
|
|
||
| LOGS_DEFAULT(VERBOSE) << "MatMulProgram: m: " << m; | ||
| LOGS_DEFAULT(VERBOSE) << "MatMulProgram: n: " << n; | ||
| LOGS_DEFAULT(VERBOSE) << "MatMulProgram: k: " << k; | ||
|
|
||
| bool has_bias = context.InputCount() > 2; | ||
| LOGS_DEFAULT(VERBOSE) << "MatMulProgram: has_bias: " << has_bias; | ||
|
|
||
| if (n < 8 && k < 8) { // call MatMulNativeProgram | ||
|
|
||
| LOGS_DEFAULT(VERBOSE) << "Running MatMulNativeProgram"; | ||
| const int components = GetMaxComponents(n); | ||
|
vraspar marked this conversation as resolved.
Outdated
|
||
| const int a_components = GetMaxComponents(k); | ||
|
|
||
| const int output_number = GetMaxComponents(m); | ||
| uint32_t output_size = static_cast<uint32_t>(helper.OutputShape().Size() / components / output_number); | ||
|
|
||
| const size_t output_rank = helper.OutputShape().NumDimensions(); | ||
| TensorShape outer_dims = output_rank > 2 ? helper.OutputShape().Slice(0, output_rank - 2) : TensorShape({}); | ||
| const int64_t batch_size = outer_dims.Size(); | ||
| TensorShape output_shape_shader({batch_size, helper.M(), helper.N()}); | ||
|
|
||
| MatMulNativeProgram program{output_size, output_number, has_bias}; | ||
| program | ||
| .CacheHint(std::to_string(components), std::to_string(a_components), std::to_string(output_number)) | ||
| .AddInputs({{a, ProgramTensorMetadataDependency::TypeAndRank,a->Shape(), a_components}, | ||
| {b, ProgramTensorMetadataDependency::TypeAndRank,b->Shape(), components}}); | ||
|
|
||
|
vraspar marked this conversation as resolved.
|
||
| if (has_bias) { | ||
| const auto* bias = context.Input(2); | ||
| program.AddInput({bias, ProgramTensorMetadataDependency::Rank, 1}); | ||
| } | ||
| program | ||
| .AddOutputs({{output_tensor, ProgramTensorMetadataDependency::None, output_shape_shader, components}}) | ||
| .SetDispatchGroupSize(ceil(static_cast<float>(output_size) / 64)) | ||
| .AddIndices(outer_dims) | ||
| .AddUniformVariables({{output_size}, {m}, {n}, {k}}); | ||
|
|
||
| return context.RunProgram(program); | ||
| } | ||
|
|
||
| int64_t batchA = a->Shape().SizeToDimension(a->Shape().NumDimensions() - 2); | ||
| int64_t batchB = b->Shape().SizeToDimension(b->Shape().NumDimensions() - 2); | ||
|
|
||
| TensorShape a_shape = a->Shape(); | ||
| TensorShape b_shape = b->Shape(); | ||
| TensorShape output_shape = helper.OutputShape(); | ||
|
|
||
| const int64_t m_value = output_shape[output_shape.NumDimensions() - 2]; | ||
| // check if A is batch of vector (bach is not 1, M is 1) and B is a matrix (batch is 1) | ||
| if (batchA != 1 && m_value == 1 && batchB == 1) { | ||
| // optimization for batched vector matrix multiplication | ||
| // dimensions of A: [1,`batchA`,K] | ||
| TensorShapeVector dims_a = {1, batchA, helper.K()}; | ||
| // dimensions of B: [1,K,N] | ||
| TensorShapeVector dims_b = {1, helper.K(), helper.N()}; | ||
|
|
||
| a_shape = TensorShape(dims_a); | ||
| b_shape = TensorShape(dims_b); | ||
| output_shape = {1, batchA, 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 int64_t batch_size = outer_dims.Size(); | ||
|
|
||
| // Get dimensions for matrix multiplication from TensorShape | ||
| const int32_t dim_a_outer = static_cast<int32_t>(a_shape[a_shape.NumDimensions() - 2]); // M dimension | ||
|
vraspar marked this conversation as resolved.
Outdated
|
||
| const int32_t dim_inner = static_cast<int32_t>(a_shape[a_shape.NumDimensions() - 1]); // K dimension | ||
| const int32_t dim_b_outer = static_cast<int32_t>(b_shape[b_shape.NumDimensions() - 1]); // N dimension | ||
|
vraspar marked this conversation as resolved.
Outdated
|
||
|
|
||
| const bool is_vec4 = dim_inner % 4 == 0 && dim_b_outer % 4 == 0; | ||
|
|
||
| InlinedVector<int64_t> elements_per_thread = dim_a_outer <= 8 | ||
| ? InlinedVector<int64_t>({4, 1, 1}) | ||
| : InlinedVector<int64_t>({4, 4, 1}); | ||
|
|
||
| const uint32_t dispatch_x = ceil(static_cast<float>(dim_b_outer) / MATMUL_PACKED_WORKGROUP_SIZE_X / elements_per_thread[0]); | ||
| const uint32_t dispatch_y = ceil(static_cast<float>(dim_a_outer) / MATMUL_PACKED_WORKGROUP_SIZE_Y / elements_per_thread[1]); | ||
| const uint32_t dispatch_z = ceil(static_cast<float>(batch_size) / MATMUL_PACKED_WORKGROUP_SIZE_Z / elements_per_thread[2]); | ||
|
|
||
| const int components = is_vec4 ? 4 : 1; | ||
| const TensorShape a_shape_temp = BuildTempShapeVector(outer_dims_a, dim_a_outer, dim_inner, components); | ||
| const TensorShape b_shape_temp = BuildTempShapeVector(outer_dims_b, dim_inner, dim_b_outer, components); | ||
| const TensorShape output_shape_temp = TensorShape({batch_size, dim_a_outer, dim_b_outer / components}); | ||
|
|
||
| MatMulProgram program{has_bias, is_vec4, elements_per_thread}; | ||
| program | ||
| .CacheHint(absl::StrJoin(elements_per_thread, "-"), std::to_string(is_vec4)) | ||
| .AddInputs({{a, ProgramTensorMetadataDependency::TypeAndRank, a_shape_temp, components}, | ||
| {b, ProgramTensorMetadataDependency::TypeAndRank, b_shape_temp, components}}) | ||
| .AddOutputs({{output_tensor, ProgramTensorMetadataDependency::Rank, output_shape_temp, components}}) | ||
| .AddUniformVariables({{dim_a_outer}, {dim_b_outer}, {dim_inner}}) | ||
| .AddIndices(outer_dims) | ||
| .SetDispatchGroupSize(dispatch_x, dispatch_y, dispatch_z) | ||
| .SetWorkgroupSize(MATMUL_PACKED_WORKGROUP_SIZE_X, MATMUL_PACKED_WORKGROUP_SIZE_Y, MATMUL_PACKED_WORKGROUP_SIZE_Z); | ||
|
|
||
| if (has_bias) { | ||
| const auto* bias = context.Input(2); | ||
| program.AddInput({bias, ProgramTensorMetadataDependency::Rank, 1}); | ||
| } | ||
|
|
||
| return context.RunProgram(program); | ||
| } | ||
|
|
||
| } // namespace webgpu | ||
| } // namespace onnxruntime | ||
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,47 @@ | ||
| // Copyright (c) Microsoft Corporation. All rights reserved. | ||
| // Licensed under the MIT License. | ||
|
|
||
| #pragma once | ||
|
|
||
| #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_utils.h" | ||
| #include "core/providers/webgpu/math/matmul_packed.h" | ||
| #include "core/providers/webgpu/webgpu_utils.h" | ||
|
|
||
| namespace onnxruntime { | ||
| namespace webgpu { | ||
|
|
||
| class MatMul final : public WebGpuKernel { | ||
| public: | ||
| MatMul(const OpKernelInfo& info) : WebGpuKernel{info} {} | ||
|
|
||
| Status ComputeInternal(ComputeContext& context) const override; | ||
| Status PrintGPUTensor(ComputeContext& context, const Tensor& tensor) const; | ||
|
vraspar marked this conversation as resolved.
Outdated
|
||
| constexpr static uint32_t MATMUL_PACKED_WORKGROUP_SIZE_X = 8; | ||
| constexpr static uint32_t MATMUL_PACKED_WORKGROUP_SIZE_Y = 8; | ||
| constexpr static uint32_t MATMUL_PACKED_WORKGROUP_SIZE_Z = 1; | ||
| }; | ||
|
|
||
| class MatMulNativeProgram final : public Program<MatMulNativeProgram> { | ||
|
vraspar marked this conversation as resolved.
Outdated
|
||
| public: | ||
| MatMulNativeProgram(const int64_t output_size, int output_number, bool has_bias) | ||
| : Program{"MatMulNative"}, output_size_(output_size), output_number_(output_number), has_bias_{has_bias} { | ||
| } | ||
|
|
||
| Status GenerateShaderCode(ShaderHelper& sh) const override; | ||
|
|
||
| WEBGPU_PROGRAM_DEFINE_UNIFORM_VARIABLES({"output_size", ProgramUniformVariableDataType::Uint32}, | ||
| {"M", ProgramUniformVariableDataType::Uint32}, | ||
| {"N", ProgramUniformVariableDataType::Uint32}, | ||
| {"K", ProgramUniformVariableDataType::Uint32}); | ||
|
|
||
| private: | ||
| const int64_t output_size_; | ||
| const int output_number_; | ||
| const bool has_bias_; | ||
| }; | ||
|
|
||
| } // namespace webgpu | ||
| } // namespace onnxruntime | ||
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
Uh oh!
There was an error while loading. Please reload this page.