blob: 48ba5af94e7fd6f0a754301bc091198dd847e00d [file] [edit]
// Copyright 2026 The Dawn & Tint Authors
//
// Redistribution and use in source and binary forms, with or without
// modification, are permitted provided that the following conditions are met:
//
// 1. Redistributions of source code must retain the above copyright notice, this
// list of conditions and the following disclaimer.
//
// 2. Redistributions in binary form must reproduce the above copyright notice,
// this list of conditions and the following disclaimer in the documentation
// and/or other materials provided with the distribution.
//
// 3. Neither the name of the copyright holder nor the names of its
// contributors may be used to endorse or promote products derived from
// this software without specific prior written permission.
//
// THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
// AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
// IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
// DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
// FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
// DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
// SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
// CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
// OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
// OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
#include <ostream>
#include <span>
#include <utility>
#include <vector>
#include "dawn/webgpu_cpp_print.h"
#include "gtest/gtest.h"
#include "src/dawn/common/GPUInfo.h"
#include "src/dawn/native/d3d12/PhysicalDeviceD3D12.h"
#include "src/utils/assert.h"
namespace dawn::native {
std::ostream& operator<<(std::ostream& o, const SubgroupMatrixConfig& config) {
o << config.M << "x" << config.N << "x" << config.K << " " << config.componentType << " -> "
<< config.resultComponentType << " [" << config.minSubgroupSize << ", "
<< config.maxSubgroupSize << "]";
return o;
}
} // namespace dawn::native
namespace dawn::native::d3d12 {
namespace {
using enum wgpu::SubgroupMatrixComponentType;
class SubgroupMatrixConfigTests : public testing::Test {
protected:
struct Shape {
uint32_t M;
uint32_t N;
uint32_t K;
};
struct SubgroupSizeRange {
uint32_t min;
uint32_t max;
};
D3D12_LINEAR_ALGEBRA_DATATYPE ToD3D12(wgpu::SubgroupMatrixComponentType type) {
switch (type) {
case I32:
return D3D12_LINEAR_ALGEBRA_DATATYPE_SINT32;
case U32:
return D3D12_LINEAR_ALGEBRA_DATATYPE_UINT32;
case F16:
return D3D12_LINEAR_ALGEBRA_DATATYPE_FLOAT16;
case F32:
return D3D12_LINEAR_ALGEBRA_DATATYPE_FLOAT32;
case I8:
return D3D12_LINEAR_ALGEBRA_DATATYPE_SINT8;
case U8:
return D3D12_LINEAR_ALGEBRA_DATATYPE_UINT8;
default:
DAWN_UNREACHABLE();
}
}
D3D12DeviceInfo::LinAlgWMMSupport MakeSupport(uint32_t waveSize,
wgpu::SubgroupMatrixComponentType typeAB,
wgpu::SubgroupMatrixComponentType typeAcc,
std::initializer_list<Shape> shapes) {
D3D12_LINEAR_ALGEBRA_WAVE_MATRIX_MULTIPLY_INPUTS inputs{};
inputs.WaveSize = waveSize;
inputs.MatrixAComponentType = ToD3D12(typeAB);
inputs.MatrixBComponentType = ToD3D12(typeAB);
inputs.AccumulatorComponentType = ToD3D12(typeAcc);
std::vector<D3D12_LINEAR_ALGEBRA_MATRIX_MULTIPLY_SHAPE> d3dShapes;
for (const auto& s : shapes) {
D3D12_LINEAR_ALGEBRA_MATRIX_MULTIPLY_SHAPE d3dShape{};
d3dShape.M = s.M;
d3dShape.N = s.N;
d3dShape.K = s.K;
d3dShapes.push_back(d3dShape);
}
return {inputs, D3D12_LINEAR_ALGEBRA_MULTIPLICATION_SUPPORT_FLAG_SUPPORTED,
std::move(d3dShapes)};
}
SubgroupMatrixConfig MakeConfig(wgpu::SubgroupMatrixComponentType componentType,
wgpu::SubgroupMatrixComponentType resultComponentType,
Shape shape,
SubgroupSizeRange subgroupSizeRange) {
SubgroupMatrixConfig config;
config.componentType = componentType;
config.resultComponentType = resultComponentType;
config.M = shape.M;
config.N = shape.N;
config.K = shape.K;
config.minSubgroupSize = subgroupSizeRange.min;
config.maxSubgroupSize = subgroupSizeRange.max;
return config;
}
std::vector<SubgroupMatrixConfig> EnumerateSubgroupMatrixConfigs(
std::span<const D3D12DeviceInfo::LinAlgWMMSupport> supports,
uint32_t vendorId = 0,
uint32_t deviceId = 0,
bool supportsShaderF16 = true) {
return PhysicalDevice::EnumerateSubgroupMatrixConfigs(supports, vendorId, deviceId,
supportsShaderF16);
}
};
// Identical base shapes across consecutive power-of-two wave sizes merge into a single config.
TEST_F(SubgroupMatrixConfigTests, SameShapeConsecutiveWaveSizes) {
std::vector supports = {
MakeSupport(16, F16, F32, {{16, 16, 16}}),
MakeSupport(32, F16, F32, {{16, 16, 16}}),
MakeSupport(64, F16, F32, {{16, 16, 16}}),
};
std::vector expected = {
MakeConfig(F16, F32, {16, 16, 16}, {16, 64}),
};
auto configs = EnumerateSubgroupMatrixConfigs(supports);
EXPECT_EQ(configs, expected);
}
// Identical base shapes with a gap in wave sizes only merge consecutive powers of two.
TEST_F(SubgroupMatrixConfigTests, SameShapeNonConsecutiveWaveSizes) {
std::vector supports = {
MakeSupport(4, F16, F32, {{16, 16, 16}}),
MakeSupport(16, F16, F32, {{16, 16, 16}}),
MakeSupport(32, F16, F32, {{16, 16, 16}}),
};
std::vector expected = {
MakeConfig(F16, F32, {16, 16, 16}, {4, 4}),
MakeConfig(F16, F32, {16, 16, 16}, {16, 32}),
};
auto configs = EnumerateSubgroupMatrixConfigs(supports);
EXPECT_EQ(configs, expected);
}
// When a larger wave size reports a shape that is a multiple of a smaller wave size's shape
// (e.g. WARP reporting 4x4x4 at WaveSize 4 and 8x8x4 at WaveSize 8), the larger shape's
// minSubgroupSize extends down to the smaller wave size, while the smaller shape stays at its own
// wave size (e.g WARP's 8x8x4 range should be [4,8]).
TEST_F(SubgroupMatrixConfigTests, LargerWaveSizeIsMultipleOfSmallerWaveSize) {
std::vector supports = {
MakeSupport(4, F32, F32, {{4, 4, 4}}),
MakeSupport(8, F32, F32, {{8, 8, 4}}),
MakeSupport(16, F32, F32, {{16, 16, 4}}),
};
std::vector expected = {
MakeConfig(F32, F32, {4, 4, 4}, {4, 4}),
MakeConfig(F32, F32, {8, 8, 4}, {4, 8}),
MakeConfig(F32, F32, {16, 16, 4}, {4, 16}),
};
auto configs = EnumerateSubgroupMatrixConfigs(supports);
EXPECT_EQ(configs, expected);
}
// When a smaller wave size reports a shape that is a multiple of a larger wave size's shape, the
// larger shape's maxSubgroupSize extends up to the larger wave size. This doesn't seem to happen in
// practice, but support it in case it does.
TEST_F(SubgroupMatrixConfigTests, SmallerWaveSizeIsMultipleOfLargerWaveSize) {
std::vector supports = {
MakeSupport(4, F32, F32, {{16, 16, 4}}),
MakeSupport(8, F32, F32, {{8, 8, 4}}),
MakeSupport(16, F32, F32, {{4, 4, 4}}),
};
std::vector expected = {
MakeConfig(F32, F32, {16, 16, 4}, {4, 16}),
MakeConfig(F32, F32, {8, 8, 4}, {8, 16}),
MakeConfig(F32, F32, {4, 4, 4}, {16, 16}),
};
auto configs = EnumerateSubgroupMatrixConfigs(supports);
EXPECT_EQ(configs, expected);
}
// Shape multiples across non-consecutive wave sizes do not merge across a missing power of two.
TEST_F(SubgroupMatrixConfigTests, NonConsecutiveMultipleWaveSizes) {
std::vector supports = {
MakeSupport(4, F32, F32, {{4, 4, 4}}),
MakeSupport(16, F32, F32, {{8, 8, 4}}),
};
std::vector expected = {
MakeConfig(F32, F32, {4, 4, 4}, {4, 4}),
MakeConfig(F32, F32, {8, 8, 4}, {16, 16}),
};
auto configs = EnumerateSubgroupMatrixConfigs(supports);
EXPECT_EQ(configs, expected);
}
// Alternating smaller and larger multiple shapes across consecutive wave sizes (combining
// LargerWaveSizeIsMultipleOfSmallerWaveSize and SmallerWaveSizeIsMultipleOfLargerWaveSize).
TEST_F(SubgroupMatrixConfigTests, AlternatingMultipleShapesAcrossWaveSizes) {
{
// 4 -> 8x8x8, 8 -> 16x16x16, 16 -> 8x8x8:
// 16x16x16 is supported at 4, 8, and 16, while 8x8x8 is only supported at 4 and 16.
std::vector supports = {
MakeSupport(4, F32, F32, {{8, 8, 8}}),
MakeSupport(8, F32, F32, {{16, 16, 16}}),
MakeSupport(16, F32, F32, {{8, 8, 8}}),
};
std::vector expected = {
MakeConfig(F32, F32, {8, 8, 8}, {4, 4}),
MakeConfig(F32, F32, {16, 16, 16}, {4, 16}),
MakeConfig(F32, F32, {8, 8, 8}, {16, 16}),
};
auto configs = EnumerateSubgroupMatrixConfigs(supports);
EXPECT_EQ(configs, expected);
}
{
// 4 -> 16x16x16, 8 -> 8x8x8, 16 -> 16x16x16:
// 16x16x16 is supported at 4, 8, and 16 (single merged entry), while 8x8x8 is only
// supported at 8.
std::vector supports = {
MakeSupport(4, F32, F32, {{16, 16, 16}}),
MakeSupport(8, F32, F32, {{8, 8, 8}}),
MakeSupport(16, F32, F32, {{16, 16, 16}}),
};
std::vector expected = {
MakeConfig(F32, F32, {16, 16, 16}, {4, 16}),
MakeConfig(F32, F32, {8, 8, 8}, {8, 8}),
};
auto configs = EnumerateSubgroupMatrixConfigs(supports);
EXPECT_EQ(configs, expected);
}
{
// 4 -> 4x4x4, 8 -> 16x16x4, 16 -> 8x8x4:
// 16x16x4 extends down to 4 (multiple of 4x4x4) and up to 16 (multiple of 8x8x4).
std::vector supports = {
MakeSupport(4, F32, F32, {{4, 4, 4}}),
MakeSupport(8, F32, F32, {{16, 16, 4}}),
MakeSupport(16, F32, F32, {{8, 8, 4}}),
};
std::vector expected = {
MakeConfig(F32, F32, {4, 4, 4}, {4, 4}),
MakeConfig(F32, F32, {16, 16, 4}, {4, 16}),
MakeConfig(F32, F32, {8, 8, 4}, {16, 16}),
};
auto configs = EnumerateSubgroupMatrixConfigs(supports);
EXPECT_EQ(configs, expected);
}
{
// 4 -> 4x4x4, 8 -> 8x8x4, 16 -> 4x4x4, 32 -> 16x16x4:
// 8x8x4 extends down to 4 and up to 16, and 16x16x4 extends all the way down to 4.
std::vector supports = {
MakeSupport(4, F32, F32, {{4, 4, 4}}),
MakeSupport(8, F32, F32, {{8, 8, 4}}),
MakeSupport(16, F32, F32, {{4, 4, 4}}),
MakeSupport(32, F32, F32, {{16, 16, 4}}),
};
std::vector expected = {
MakeConfig(F32, F32, {4, 4, 4}, {4, 4}),
MakeConfig(F32, F32, {8, 8, 4}, {4, 16}),
MakeConfig(F32, F32, {4, 4, 4}, {16, 16}),
MakeConfig(F32, F32, {16, 16, 4}, {4, 32}),
};
auto configs = EnumerateSubgroupMatrixConfigs(supports);
EXPECT_EQ(configs, expected);
}
}
// Shapes where neither is a multiple of the other in all dimensions remain independent.
TEST_F(SubgroupMatrixConfigTests, NonMultipleShapesAreIndependent) {
std::vector supports = {
MakeSupport(4, F32, F32, {{4, 8, 4}}),
MakeSupport(8, F32, F32, {{8, 4, 4}}),
};
std::vector expected = {
MakeConfig(F32, F32, {4, 8, 4}, {4, 4}),
MakeConfig(F32, F32, {8, 4, 4}, {8, 8}),
};
auto configs = EnumerateSubgroupMatrixConfigs(supports);
EXPECT_EQ(configs, expected);
}
// Configs with different component types or result component types do not merge with each other.
TEST_F(SubgroupMatrixConfigTests, DifferentComponentTypesDoNotMerge) {
std::vector supports = {
MakeSupport(4, F32, F32, {{4, 4, 4}}),
// Different componentType and resultComponentType (same shape):
MakeSupport(8, I32, I32, {{4, 4, 4}}),
// Different componentType, same resultComponentType (multiple shape):
MakeSupport(8, F16, F32, {{8, 8, 4}}),
// Same componentType, different resultComponentType (same shape and multiple shape):
MakeSupport(16, F16, F16, {{8, 8, 4}}),
MakeSupport(16, I32, U32, {{8, 8, 4}}),
};
std::vector expected = {
MakeConfig(F32, F32, {4, 4, 4}, {4, 4}), MakeConfig(I32, I32, {4, 4, 4}, {8, 8}),
MakeConfig(F16, F32, {8, 8, 4}, {8, 8}), MakeConfig(F16, F16, {8, 8, 4}, {16, 16}),
MakeConfig(I32, U32, {8, 8, 4}, {16, 16}),
};
auto configs = EnumerateSubgroupMatrixConfigs(supports);
EXPECT_EQ(configs, expected);
}
// Filtering rules for unsupported type combinations, ShaderF16, and WARP 8-bit types.
TEST_F(SubgroupMatrixConfigTests, TypeAndFeatureFiltering) {
std::vector supports = {
// Valid integer and float combinations:
MakeSupport(32, I32, I32, {{16, 16, 16}}),
MakeSupport(32, I8, I32, {{16, 16, 16}}),
MakeSupport(32, U8, U32, {{16, 16, 16}}),
MakeSupport(32, F16, F32, {{16, 16, 16}}),
MakeSupport(32, F32, F32, {{16, 16, 16}}),
// Invalid: mixed int/float:
MakeSupport(32, I32, F32, {{16, 16, 16}}),
MakeSupport(32, F16, I32, {{16, 16, 16}}),
// Invalid: input byte size > accumulator byte size:
MakeSupport(32, F32, F16, {{16, 16, 16}}),
MakeSupport(32, I32, I8, {{16, 16, 16}}),
};
// Default (non-WARP, ShaderF16 supported):
{
std::vector expected = {
MakeConfig(I32, I32, {16, 16, 16}, {32, 32}),
MakeConfig(I8, I32, {16, 16, 16}, {32, 32}),
MakeConfig(U8, U32, {16, 16, 16}, {32, 32}),
MakeConfig(F16, F32, {16, 16, 16}, {32, 32}),
MakeConfig(F32, F32, {16, 16, 16}, {32, 32}),
};
auto configs = EnumerateSubgroupMatrixConfigs(supports, /*vendorId=*/0,
/*deviceId=*/0,
/*supportsShaderF16=*/true);
EXPECT_EQ(configs, expected);
}
// Without ShaderF16: float configs are filtered out.
{
std::vector expected = {
MakeConfig(I32, I32, {16, 16, 16}, {32, 32}),
MakeConfig(I8, I32, {16, 16, 16}, {32, 32}),
MakeConfig(U8, U32, {16, 16, 16}, {32, 32}),
};
auto configs = EnumerateSubgroupMatrixConfigs(supports, /*vendorId=*/0,
/*deviceId=*/0,
/*supportsShaderF16=*/false);
EXPECT_EQ(configs, expected);
}
// On WARP: 8-bit configs are filtered out.
{
constexpr uint32_t kDeviceID_WARP = 0x8c;
ASSERT_TRUE(gpu_info::IsMicrosoftWARP(gpu_info::kVendorID_Microsoft, kDeviceID_WARP));
std::vector expected = {
MakeConfig(I32, I32, {16, 16, 16}, {32, 32}),
MakeConfig(F16, F32, {16, 16, 16}, {32, 32}),
MakeConfig(F32, F32, {16, 16, 16}, {32, 32}),
};
auto configs =
EnumerateSubgroupMatrixConfigs(supports, gpu_info::kVendorID_Microsoft, kDeviceID_WARP,
/*supportsShaderF16=*/true);
EXPECT_EQ(configs, expected);
}
}
} // namespace
} // namespace dawn::native::d3d12