diff --git a/tools/clang/unittests/HLSLExec/LinAlgTests.cpp b/tools/clang/unittests/HLSLExec/LinAlgTests.cpp index 01a7b4ce5e..6bac2510a7 100644 --- a/tools/clang/unittests/HLSLExec/LinAlgTests.cpp +++ b/tools/clang/unittests/HLSLExec/LinAlgTests.cpp @@ -1676,6 +1676,7 @@ class LinAlgCPUOracleTests { END_TEST_CLASS() TEST_METHOD(TypedMatrixBufferRoundTrip); + TEST_METHOD(MatrixProductOracle); TEST_METHOD(UntouchedByteVerification); TEST_METHOD(ViewBoundedElements); TEST_METHOD(ViewBoundedStoreBytes); @@ -2314,8 +2315,13 @@ class DxilConf_SM610_LinAlg { // Matrix Matrix Arithmetic TEST_METHOD(MatMatMul_Wave_16x16x16_F16); + TEST_METHOD(MatMatMul_Wave_8x32x16_F16_NonUniform); + TEST_METHOD(MatMatMul_Wave_8x32x16_F16_ToF32); + TEST_METHOD(MatMatMul_Wave_16x16x16_I32); TEST_METHOD(MatMatMulAccum_Wave_16x16x16_F16); + TEST_METHOD(MatMatMulAccum_Wave_8x32x16_F16_ToF32_NonUniform); TEST_METHOD(MatAccum_Wave_16x16_F16); + TEST_METHOD(MatAccum_Wave_8x32_F16_BUse_NonUniform); // Matrix Vector Arithmetic TEST_METHOD(MatVecMul_Thread_16x16_F16); @@ -4310,6 +4316,702 @@ void DxilConf_SM610_LinAlg::MatAccum_Wave_16x16_F16() { /*LHSFill=*/2.0f, /*RHSFill=*/3.0f, SelectedWaveSize); } +namespace cpu_oracle { + +static std::vector +multiplyIntegerMatrices(MatrixDim M, MatrixDim K, MatrixDim N, + const std::vector &MatrixA, + const std::vector &MatrixB, + const std::vector *Accumulator = nullptr) { + std::vector Result(static_cast(M) * N); + for (MatrixDim Row = 0; Row < M; ++Row) { + for (MatrixDim Column = 0; Column < N; ++Column) { + const size_t ResultIndex = static_cast(Row) * N + Column; + int64_t Sum = Accumulator ? (*Accumulator)[ResultIndex] : 0; + for (MatrixDim Inner = 0; Inner < K; ++Inner) + Sum += MatrixA[static_cast(Row) * K + Inner] * + MatrixB[static_cast(Inner) * N + Column]; + Result[ResultIndex] = Sum; + } + } + return Result; +} + +static std::optional> +encodeLogicalMatrixBuffer(const MatrixParams &Matrix, + const std::vector &Values) { + if (Matrix.M == 0 || Matrix.N == 0) + return std::nullopt; + if (Values.size() != static_cast(Matrix.M) * Matrix.N) + return std::nullopt; + if (Matrix.Layout != MatrixLayout::RowMajor && + Matrix.Layout != MatrixLayout::ColumnMajor) + return std::nullopt; + + std::optional LogicalMatrix; + switch (Matrix.CompType) { + case ComponentType::F16: { + std::vector TypedValues; + TypedValues.reserve(Values.size()); + for (int64_t Value : Values) { + if (Value < -2048 || Value > 2048) + return std::nullopt; + TypedValues.emplace_back(static_cast(Value)); + } + LogicalMatrix = makeTypedMatrix(Matrix.M, Matrix.N, std::move(TypedValues)); + break; + } + case ComponentType::F32: { + static constexpr int64_t MaxExactInteger = int64_t(1) << 24; + std::vector TypedValues; + TypedValues.reserve(Values.size()); + for (int64_t Value : Values) { + if (Value < -MaxExactInteger || Value > MaxExactInteger) + return std::nullopt; + TypedValues.push_back(static_cast(Value)); + } + LogicalMatrix = makeTypedMatrix(Matrix.M, Matrix.N, std::move(TypedValues)); + break; + } + case ComponentType::I32: { + std::vector TypedValues; + TypedValues.reserve(Values.size()); + for (int64_t Value : Values) { + if (Value < std::numeric_limits::min() || + Value > std::numeric_limits::max()) + return std::nullopt; + TypedValues.push_back(static_cast(Value)); + } + LogicalMatrix = makeTypedMatrix(Matrix.M, Matrix.N, std::move(TypedValues)); + break; + } + case ComponentType::U32: { + std::vector TypedValues; + TypedValues.reserve(Values.size()); + for (int64_t Value : Values) { + if (Value < 0 || + static_cast(Value) > std::numeric_limits::max()) + return std::nullopt; + TypedValues.push_back(static_cast(Value)); + } + LogicalMatrix = makeTypedMatrix(Matrix.M, Matrix.N, std::move(TypedValues)); + break; + } + default: + return std::nullopt; + } + if (!LogicalMatrix) + return std::nullopt; + + const MatrixBufferLayout Layout = { + Matrix.Layout, + /*OffsetBytes=*/0, + /*StrideBytes=*/Matrix.strideBytes(), + }; + std::optional BufferSize = + getMatrixBufferSize(*LogicalMatrix, Layout); + if (!BufferSize) + return std::nullopt; + + std::vector Buffer(*BufferSize, 0); + if (!writeMatrixBuffer(*LogicalMatrix, Layout, Buffer)) + return std::nullopt; + return Buffer; +} + +} // namespace cpu_oracle + +void LinAlgCPUOracleTests::MatrixProductOracle() { + using namespace cpu_oracle; + + const std::vector Product = + multiplyIntegerMatrices(/*M=*/2, /*K=*/3, /*N=*/2, + /*MatrixA=*/{1, 2, 3, 4, 5, 6}, + /*MatrixB=*/{7, 8, 9, 10, 11, 12}); + VERIFY_IS_TRUE(Product == std::vector({58, 64, 139, 154})); + + const std::vector InitialAccumulator = {1, -1, 2, -2}; + const std::vector Accumulated = multiplyIntegerMatrices( + /*M=*/2, /*K=*/3, /*N=*/2, + /*MatrixA=*/{1, 2, 3, 4, 5, 6}, + /*MatrixB=*/{7, 8, 9, 10, 11, 12}, &InitialAccumulator); + VERIFY_IS_TRUE(Accumulated == std::vector({59, 63, 141, 152})); +} + +static bool waveArithmeticNeeds16BitTypes(ComponentType CompType) { + return CompType == ComponentType::F16 || CompType == ComponentType::I16 || + CompType == ComponentType::U16; +} + +static MatrixParams makeWaveArithmeticParams(ComponentType CompType, + MatrixDim M, MatrixDim N, + MatrixUse Use, UINT WaveSize) { + MatrixParams Params = {}; + Params.CompType = CompType; + Params.M = M; + Params.N = N; + Params.Use = Use; + Params.Scope = MatrixScope::Wave; + Params.Layout = MatrixLayout::RowMajor; + Params.NumThreads = static_cast(WaveSize); + Params.Enable16Bit = waveArithmeticNeeds16BitTypes(CompType); + return Params; +} + +static std::vector makeWaveArithmeticPattern(MatrixDim M, MatrixDim N, + int64_t RowScale, + int64_t ColumnScale, + int64_t Modulus, + int64_t Center) { + VERIFY_IS_TRUE(M != 0 && N != 0 && Modulus > 0); + if (M == 0 || N == 0 || Modulus <= 0) + return {}; + + std::vector Values(static_cast(M) * N); + for (MatrixDim Row = 0; Row < M; ++Row) { + for (MatrixDim Column = 0; Column < N; ++Column) { + Values[static_cast(Row) * N + Column] = + (static_cast(Row) * RowScale + + static_cast(Column) * ColumnScale) % + Modulus - + Center; + } + } + return Values; +} + +enum class WaveMultiplyOperation { + Multiply, + MultiplyAccumulate, +}; + +struct WaveMultiplyCase { + ComponentType MatrixAType = ComponentType::Invalid; + ComponentType MatrixBType = ComponentType::Invalid; + ComponentType AccumulatorType = ComponentType::Invalid; + MatrixDim M = 0; + MatrixDim K = 0; + MatrixDim N = 0; + WaveMultiplyOperation Operation = WaveMultiplyOperation::Multiply; + std::vector MatrixAValues; + std::vector MatrixBValues; + std::vector AccumulatorValues; + std::wstring PublicRule; + + bool accumulates() const { + return Operation == WaveMultiplyOperation::MultiplyAccumulate; + } +}; + +static bool isWaveMultiplyCaseValid(const WaveMultiplyCase &Case) { + if (Case.M == 0 || Case.K == 0 || Case.N == 0) + return false; + if (Case.MatrixAValues.size() != static_cast(Case.M) * Case.K) + return false; + if (Case.MatrixBValues.size() != static_cast(Case.K) * Case.N) + return false; + if (Case.accumulates() && + Case.AccumulatorValues.size() != static_cast(Case.M) * Case.N) + return false; + if (!Case.accumulates() && !Case.AccumulatorValues.empty()) + return false; + if (!toCapabilityDataType(Case.MatrixAType)) + return false; + if (!toCapabilityDataType(Case.MatrixBType)) + return false; + if (!toCapabilityDataType(Case.AccumulatorType)) + return false; + return !Case.PublicRule.empty(); +} + +static HRESULT selectWaveArithmeticMultiplyWaveSize( + ID3D12Device *Device, const WaveMultiplyCase &Case, LPCWSTR CaseName, + bool &Supported, UINT &SelectedWaveSize) { + Supported = false; + SelectedWaveSize = 0; + if (!Device || !CaseName || !isWaveMultiplyCaseValid(Case) || + !linalg_test::isLegalScope( + linalg_abi::D3D12_LINEAR_ALGEBRA_OPERATION_TYPE_WAVE_MATRIX_MULTIPLY, + MatrixScope::Wave)) + return E_INVALIDARG; + + const linalg_abi::D3D12_LINEAR_ALGEBRA_DATATYPE MatrixAType = + *toCapabilityDataType(Case.MatrixAType); + const linalg_abi::D3D12_LINEAR_ALGEBRA_DATATYPE MatrixBType = + *toCapabilityDataType(Case.MatrixBType); + const linalg_abi::D3D12_LINEAR_ALGEBRA_DATATYPE AccumulatorType = + *toCapabilityDataType(Case.AccumulatorType); + + linalg_test::TierSupport Tier; + HRESULT HR = linalg_test::queryTierSupport(Device, Tier); + if (FAILED(HR) || !Tier.supported()) + return HR; + + UINT MinWaveSize = 0; + UINT MaxWaveSize = 0; + HR = queryLaunchableWaveSizes(Device, MinWaveSize, MaxWaveSize); + if (FAILED(HR)) + return HR; + if (MinWaveSize == 0) + return S_OK; + + const linalg_abi::D3D12_LINEAR_ALGEBRA_MATRIX_SHAPE Shape = { + Case.M, + Case.K, + Case.N, + }; + struct ConstructionRole { + linalg_abi::D3D12_LINEAR_ALGEBRA_DATATYPE Type; + MatrixUse Use; + UINT Rows; + UINT Columns; + }; + const ConstructionRole ConstructionRoles[] = { + {MatrixAType, MatrixUse::A, Case.M, Case.K}, + {MatrixBType, MatrixUse::B, Case.K, Case.N}, + {AccumulatorType, MatrixUse::Accumulator, Case.M, Case.N}, + }; + + for (UINT WaveSize = 4; WaveSize <= 128; WaveSize *= 2) { + if (WaveSize < MinWaveSize || WaveSize > MaxWaveSize) + continue; + + bool AllRolesConstructible = true; + for (const ConstructionRole &Role : ConstructionRoles) { + bool RoleConstructible = false; + HR = supportsMatrixShape(Device, Role.Type, WaveSize, Role.Use, Role.Rows, + Role.Columns, RoleConstructible); + if (FAILED(HR)) + return HR; + if (!RoleConstructible) { + AllRolesConstructible = false; + break; + } + } + if (!AllRolesConstructible) + continue; + + linalg_test::WaveMatrixMultiplySupport Multiply; + HR = linalg_test::queryWaveMatrixMultiply( + Device, {{WaveSize, MatrixAType, MatrixBType, AccumulatorType}, Shape}, + Multiply); + if (FAILED(HR)) + return HR; + if (!Multiply.supported()) + continue; + + hlsl_test::LogCommentFmt( + L"Wave arithmetic capability matched wave=%u, shape=(%u,%u,%u) for %s", + WaveSize, Case.M, Case.K, Case.N, CaseName); + Supported = true; + SelectedWaveSize = WaveSize; + return S_OK; + } + + hlsl_test::LogCommentFmt( + L"No required MatrixConstruction roles and WaveMatrixMultiply " + L"capability intersect for %s", + CaseName); + return S_OK; +} + +static const char WaveMultiplyShader[] = R"( + #define USE_A 0 + #define USE_B 1 + #define USE_ACC 2 + #define SCOPE_WAVE 1 + #define LAYOUT_ROW_MAJOR 0 + + ByteAddressBuffer MatrixAInput : register(t0); + ByteAddressBuffer MatrixBInput : register(t1); +#if DO_ACCUMULATE + ByteAddressBuffer AccumulatorInput : register(t2); + RWByteAddressBuffer Output : register(u3); +#else + RWByteAddressBuffer Output : register(u2); +#endif + + [WaveSize(FORCED_WAVE_SIZE)] + [numthreads(NUMTHREADS, 1, 1)] + void main() { + __builtin_LinAlgMatrix + [[__LinAlgMatrix_Attributes( + MATRIX_A_COMP_TYPE, M_DIM, K_DIM, USE_A, SCOPE_WAVE)]] + MatA; + __builtin_LinAlg_MatrixLoadFromDescriptor( + MatA, MatrixAInput, 0, MATRIX_A_STRIDE, LAYOUT_ROW_MAJOR, 128); + + __builtin_LinAlgMatrix + [[__LinAlgMatrix_Attributes( + MATRIX_B_COMP_TYPE, K_DIM, N_DIM, USE_B, SCOPE_WAVE)]] + MatB; + __builtin_LinAlg_MatrixLoadFromDescriptor( + MatB, MatrixBInput, 0, MATRIX_B_STRIDE, LAYOUT_ROW_MAJOR, 128); + + __builtin_LinAlgMatrix + [[__LinAlgMatrix_Attributes( + ACCUMULATOR_COMP_TYPE, M_DIM, N_DIM, USE_ACC, SCOPE_WAVE)]] + Result; +#if DO_ACCUMULATE + __builtin_LinAlgMatrix + [[__LinAlgMatrix_Attributes( + ACCUMULATOR_COMP_TYPE, M_DIM, N_DIM, USE_ACC, SCOPE_WAVE)]] + Accumulator; + __builtin_LinAlg_MatrixLoadFromDescriptor( + Accumulator, AccumulatorInput, 0, ACCUMULATOR_STRIDE, + LAYOUT_ROW_MAJOR, 128); + __builtin_LinAlg_MatrixMatrixMultiplyAccumulate( + Result, MatA, MatB, Accumulator); +#else + __builtin_LinAlg_MatrixMatrixMultiply(Result, MatA, MatB); +#endif + + __builtin_LinAlg_MatrixStoreToDescriptor( + Result, Output, 0, ACCUMULATOR_STRIDE, LAYOUT_ROW_MAJOR, 128); + } +)"; + +static std::optional +buildWaveMultiplyCompilerArgs(const WaveMultiplyCase &Case, UINT WaveSize) { + if (!isWaveMultiplyCaseValid(Case) || WaveSize == 0) + return std::nullopt; + + const MatrixParams MatrixA = makeWaveArithmeticParams( + Case.MatrixAType, Case.M, Case.K, MatrixUse::A, WaveSize); + const MatrixParams MatrixB = makeWaveArithmeticParams( + Case.MatrixBType, Case.K, Case.N, MatrixUse::B, WaveSize); + const MatrixParams Accumulator = makeWaveArithmeticParams( + Case.AccumulatorType, Case.M, Case.N, MatrixUse::Accumulator, WaveSize); + + std::stringstream SS; + SS << "-HV 202x"; + SS << " -DMATRIX_A_COMP_TYPE=" << static_cast(Case.MatrixAType); + SS << " -DMATRIX_B_COMP_TYPE=" << static_cast(Case.MatrixBType); + SS << " -DACCUMULATOR_COMP_TYPE=" << static_cast(Case.AccumulatorType); + SS << " -DM_DIM=" << Case.M; + SS << " -DK_DIM=" << Case.K; + SS << " -DN_DIM=" << Case.N; + SS << " -DMATRIX_A_STRIDE=" << MatrixA.strideBytes(); + SS << " -DMATRIX_B_STRIDE=" << MatrixB.strideBytes(); + SS << " -DACCUMULATOR_STRIDE=" << Accumulator.strideBytes(); + SS << " -DNUMTHREADS=" << WaveSize; + SS << " -DFORCED_WAVE_SIZE=" << WaveSize; + SS << " -DDO_ACCUMULATE=" << static_cast(Case.accumulates()); + if (waveArithmeticNeeds16BitTypes(Case.MatrixAType) || + waveArithmeticNeeds16BitTypes(Case.MatrixBType) || + waveArithmeticNeeds16BitTypes(Case.AccumulatorType)) + SS << " -enable-16bit-types"; + return SS.str(); +} + +static bool +verifyWaveArithmeticMatrix(const void *Actual, size_t ActualSize, + const MatrixParams &Params, + const std::vector &ExpectedValues, + const std::wstring &PublicRule, bool Verbose) { + const std::optional> ExpectedBuffer = + cpu_oracle::encodeLogicalMatrixBuffer(Params, ExpectedValues); + VERIFY_IS_TRUE(ExpectedBuffer.has_value()); + if (!ExpectedBuffer) + return false; + + const cpu_oracle::MatrixBufferLayout Layout = { + Params.Layout, + /*OffsetBytes=*/0, + /*StrideBytes=*/Params.strideBytes(), + }; + const std::optional Expected = + cpu_oracle::decodeMatrixBuffer(Params.CompType, Params.M, Params.N, + Layout, ExpectedBuffer->data(), + ExpectedBuffer->size()); + VERIFY_IS_TRUE(Expected.has_value()); + if (!Expected) + return false; + + const cpu_oracle::MatrixResultOracle Oracle = + cpu_oracle::exactResult(*Expected, PublicRule); + return cpu_oracle::verifyMatrixBuffer(Actual, ActualSize, Layout, Oracle, + Verbose); +} + +static void runWaveMultiplyCase(ID3D12Device *Device, + dxc::SpecificDllLoader &DxcSupport, + const WaveMultiplyCase &Case, LPCWSTR CaseName, + bool Verbose) { + VERIFY_IS_TRUE(isWaveMultiplyCaseValid(Case)); + if (!isWaveMultiplyCaseValid(Case)) + return; + + bool Supported = false; + UINT SelectedWaveSize = 0; + const HRESULT QueryResult = selectWaveArithmeticMultiplyWaveSize( + Device, Case, CaseName, Supported, SelectedWaveSize); + if (!applyApplicability( + linalg_test::classifyApplicability( + QueryResult, Supported, + linalg_test::CapabilityRequirement::CapabilityGated), + CaseName)) + return; + VERIFY_IS_TRUE(SelectedWaveSize != 0); + if (SelectedWaveSize == 0) + return; + + const MatrixParams MatrixA = makeWaveArithmeticParams( + Case.MatrixAType, Case.M, Case.K, MatrixUse::A, SelectedWaveSize); + const MatrixParams MatrixB = makeWaveArithmeticParams( + Case.MatrixBType, Case.K, Case.N, MatrixUse::B, SelectedWaveSize); + const MatrixParams Accumulator = + makeWaveArithmeticParams(Case.AccumulatorType, Case.M, Case.N, + MatrixUse::Accumulator, SelectedWaveSize); + + const std::optional> MatrixABuffer = + cpu_oracle::encodeLogicalMatrixBuffer(MatrixA, Case.MatrixAValues); + const std::optional> MatrixBBuffer = + cpu_oracle::encodeLogicalMatrixBuffer(MatrixB, Case.MatrixBValues); + const std::optional> AccumulatorBuffer = + Case.accumulates() ? cpu_oracle::encodeLogicalMatrixBuffer( + Accumulator, Case.AccumulatorValues) + : std::optional>(); + const std::vector Expected = cpu_oracle::multiplyIntegerMatrices( + Case.M, Case.K, Case.N, Case.MatrixAValues, Case.MatrixBValues, + Case.accumulates() ? &Case.AccumulatorValues : nullptr); + const std::optional Args = + buildWaveMultiplyCompilerArgs(Case, SelectedWaveSize); + VERIFY_IS_TRUE(MatrixABuffer.has_value()); + VERIFY_IS_TRUE(MatrixBBuffer.has_value()); + VERIFY_IS_TRUE(!Case.accumulates() || AccumulatorBuffer.has_value()); + VERIFY_IS_TRUE(Args.has_value()); + if (!MatrixABuffer || !MatrixBBuffer || + (Case.accumulates() && !AccumulatorBuffer) || !Args) + return; + + const std::optional> ExpectedBuffer = + cpu_oracle::encodeLogicalMatrixBuffer(Accumulator, Expected); + VERIFY_IS_TRUE(ExpectedBuffer.has_value()); + if (!ExpectedBuffer) + return; + + const char *RootSignature = Case.accumulates() + ? "SRV(t0), SRV(t1), SRV(t2), UAV(u3)" + : "SRV(t0), SRV(t1), UAV(u2)"; + compileShader(DxcSupport, WaveMultiplyShader, "cs_6_10", *Args, Verbose); + + auto Op = createComputeOp(WaveMultiplyShader, "cs_6_10", RootSignature, + Args->c_str()); + addSRVBuffer(Op.get(), "MatrixAInput", MatrixABuffer->size(), "byname"); + addSRVBuffer(Op.get(), "MatrixBInput", MatrixBBuffer->size(), "byname"); + if (Case.accumulates()) + addSRVBuffer(Op.get(), "AccumulatorInput", AccumulatorBuffer->size(), + "byname"); + addUAVBuffer(Op.get(), "Output", ExpectedBuffer->size(), true); + addRootView(Op.get(), 0, "MatrixAInput"); + addRootView(Op.get(), 1, "MatrixBInput"); + if (Case.accumulates()) { + addRootView(Op.get(), 2, "AccumulatorInput"); + addRootView(Op.get(), 3, "Output"); + } else { + addRootView(Op.get(), 2, "Output"); + } + + auto Result = runShaderOp( + Device, DxcSupport, std::move(Op), + [MatrixABuffer, MatrixBBuffer, AccumulatorBuffer, + &Case](LPCSTR Name, std::vector &Data, st::ShaderOp *) { + const std::vector *Source = nullptr; + if (_stricmp(Name, "MatrixAInput") == 0) + Source = &*MatrixABuffer; + else if (_stricmp(Name, "MatrixBInput") == 0) + Source = &*MatrixBBuffer; + else if (Case.accumulates() && _stricmp(Name, "AccumulatorInput") == 0) + Source = &*AccumulatorBuffer; + if (!Source) + return; + VERIFY_IS_TRUE(Data.size() == Source->size()); + if (Data.size() != Source->size()) + return; + std::memcpy(Data.data(), Source->data(), Data.size()); + }); + + MappedData OutData; + Result->Test->GetReadBackData("Output", &OutData); + VERIFY_IS_TRUE(verifyWaveArithmeticMatrix(OutData.data(), OutData.size(), + Accumulator, Expected, + Case.PublicRule, Verbose)); +} + +static WaveMultiplyCase +makeRectangularF16WaveMultiplyCase(ComponentType AccumulatorType, + WaveMultiplyOperation Operation) { + WaveMultiplyCase Case = {}; + Case.MatrixAType = ComponentType::F16; + Case.MatrixBType = ComponentType::F16; + Case.AccumulatorType = AccumulatorType; + Case.M = 8; + Case.K = 32; + Case.N = 16; + Case.Operation = Operation; + Case.MatrixAValues = makeWaveArithmeticPattern(Case.M, Case.K, 3, 2, 5, 2); + Case.MatrixBValues = makeWaveArithmeticPattern(Case.K, Case.N, 1, 3, 7, 3); + if (Case.accumulates()) { + Case.AccumulatorValues = + makeWaveArithmeticPattern(Case.M, Case.N, 2, 1, 5, 2); + Case.PublicRule = + L"Exact non-uniform F16 product plus an independent F32 accumulator"; + } else if (AccumulatorType == ComponentType::F32) { + Case.PublicRule = + L"Exact non-uniform F16 matrix product stored in an F32 accumulator"; + } else { + Case.PublicRule = L"Exact non-uniform rectangular F16 matrix product"; + } + return Case; +} + +void DxilConf_SM610_LinAlg::MatMatMul_Wave_8x32x16_F16_NonUniform() { + const WaveMultiplyCase Case = makeRectangularF16WaveMultiplyCase( + ComponentType::F16, WaveMultiplyOperation::Multiply); + runWaveMultiplyCase(D3DDevice, DxcSupport, Case, + L"MatMatMul_Wave_8x32x16_F16_NonUniform", VerboseLogging); +} + +void DxilConf_SM610_LinAlg::MatMatMul_Wave_8x32x16_F16_ToF32() { + const WaveMultiplyCase Case = makeRectangularF16WaveMultiplyCase( + ComponentType::F32, WaveMultiplyOperation::Multiply); + runWaveMultiplyCase(D3DDevice, DxcSupport, Case, + L"MatMatMul_Wave_8x32x16_F16_ToF32", VerboseLogging); +} + +void DxilConf_SM610_LinAlg::MatMatMulAccum_Wave_8x32x16_F16_ToF32_NonUniform() { + const WaveMultiplyCase Case = makeRectangularF16WaveMultiplyCase( + ComponentType::F32, WaveMultiplyOperation::MultiplyAccumulate); + runWaveMultiplyCase(D3DDevice, DxcSupport, Case, + L"MatMatMulAccum_Wave_8x32x16_F16_ToF32_NonUniform", + VerboseLogging); +} + +void DxilConf_SM610_LinAlg::MatMatMul_Wave_16x16x16_I32() { + WaveMultiplyCase Case = {}; + Case.MatrixAType = ComponentType::I32; + Case.MatrixBType = ComponentType::I32; + Case.AccumulatorType = ComponentType::I32; + Case.M = 16; + Case.K = 16; + Case.N = 16; + Case.MatrixAValues = makeWaveArithmeticPattern(Case.M, Case.K, 3, 2, 5, 2); + Case.MatrixBValues = makeWaveArithmeticPattern(Case.K, Case.N, 1, 3, 7, 3); + Case.PublicRule = L"Exact non-uniform I32 matrix product"; + runWaveMultiplyCase(D3DDevice, DxcSupport, Case, + L"MatMatMul_Wave_16x16x16_I32", VerboseLogging); +} + +static const char WaveAccumulateBUseShader[] = R"( + #define USE_ACC 2 + + ByteAddressBuffer AccumulatorInput : register(t0); + ByteAddressBuffer RHSInput : register(t1); + RWByteAddressBuffer Output : register(u2); + + [WaveSize(FORCED_WAVE_SIZE)] + [numthreads(NUMTHREADS, 1, 1)] + void main() { + __builtin_LinAlgMatrix + [[__LinAlgMatrix_Attributes( + COMP_TYPE, M_DIM, N_DIM, USE_ACC, SCOPE)]] + Accumulator; + __builtin_LinAlg_MatrixLoadFromDescriptor( + Accumulator, AccumulatorInput, 0, STRIDE, LAYOUT, 128); + + __builtin_LinAlgMatrix + [[__LinAlgMatrix_Attributes(COMP_TYPE, M_DIM, N_DIM, USE, SCOPE)]] + RHS; + __builtin_LinAlg_MatrixLoadFromDescriptor( + RHS, RHSInput, 0, STRIDE, LAYOUT, 128); + + __builtin_LinAlgMatrix + [[__LinAlgMatrix_Attributes( + COMP_TYPE, M_DIM, N_DIM, USE_ACC, SCOPE)]] + Result; + __builtin_LinAlg_MatrixAccumulate(Result, Accumulator, RHS); + __builtin_LinAlg_MatrixStoreToDescriptor( + Result, Output, 0, STRIDE, LAYOUT, 128); + } +)"; + +void DxilConf_SM610_LinAlg::MatAccum_Wave_8x32_F16_BUse_NonUniform() { + MatrixParams Params = makeWaveArithmeticParams( + ComponentType::F16, /*M=*/8, /*N=*/32, MatrixUse::B, /*WaveSize=*/128); + + UINT SelectedWaveSize = 0; + if (!matrixConstructionApplicable( + D3DDevice, Params, {MatrixUse::Accumulator, MatrixUse::B}, + L"MatAccum_Wave_8x32_F16_BUse_NonUniform", SelectedWaveSize)) + return; + Params.NumThreads = static_cast(SelectedWaveSize); + + const std::vector AccumulatorValues = + makeWaveArithmeticPattern(Params.M, Params.N, 2, 1, 7, 3); + const std::vector RHSValues = + makeWaveArithmeticPattern(Params.M, Params.N, 1, 3, 5, 2); + std::vector ExpectedValues(AccumulatorValues.size()); + for (size_t I = 0; I < ExpectedValues.size(); ++I) + ExpectedValues[I] = AccumulatorValues[I] + RHSValues[I]; + + const MatrixParams AccumulatorParams = + makeWaveArithmeticParams(ComponentType::F16, Params.M, Params.N, + MatrixUse::Accumulator, SelectedWaveSize); + const std::optional> AccumulatorBuffer = + cpu_oracle::encodeLogicalMatrixBuffer(AccumulatorParams, + AccumulatorValues); + const std::optional> RHSBuffer = + cpu_oracle::encodeLogicalMatrixBuffer(Params, RHSValues); + const std::optional> ExpectedBuffer = + cpu_oracle::encodeLogicalMatrixBuffer(AccumulatorParams, ExpectedValues); + VERIFY_IS_TRUE(AccumulatorBuffer.has_value()); + VERIFY_IS_TRUE(RHSBuffer.has_value()); + VERIFY_IS_TRUE(ExpectedBuffer.has_value()); + if (!AccumulatorBuffer || !RHSBuffer || !ExpectedBuffer) + return; + + std::stringstream ExtraDefs; + ExtraDefs << " -DFORCED_WAVE_SIZE=" << SelectedWaveSize; + const std::string Args = buildCompilerArgs(Params, ExtraDefs.str().c_str()); + compileShader(DxcSupport, WaveAccumulateBUseShader, "cs_6_10", Args, + VerboseLogging); + + auto Op = createComputeOp(WaveAccumulateBUseShader, "cs_6_10", + "SRV(t0), SRV(t1), UAV(u2)", Args.c_str()); + addSRVBuffer(Op.get(), "AccumulatorInput", AccumulatorBuffer->size(), + "byname"); + addSRVBuffer(Op.get(), "RHSInput", RHSBuffer->size(), "byname"); + addUAVBuffer(Op.get(), "Output", ExpectedBuffer->size(), true); + addRootView(Op.get(), 0, "AccumulatorInput"); + addRootView(Op.get(), 1, "RHSInput"); + addRootView(Op.get(), 2, "Output"); + + auto Result = + runShaderOp(D3DDevice, DxcSupport, std::move(Op), + [AccumulatorBuffer, RHSBuffer]( + LPCSTR Name, std::vector &Data, st::ShaderOp *) { + const std::vector *Source = nullptr; + if (_stricmp(Name, "AccumulatorInput") == 0) + Source = &*AccumulatorBuffer; + else if (_stricmp(Name, "RHSInput") == 0) + Source = &*RHSBuffer; + if (!Source) + return; + VERIFY_IS_TRUE(Data.size() == Source->size()); + if (Data.size() != Source->size()) + return; + std::memcpy(Data.data(), Source->data(), Data.size()); + }); + + MappedData OutData; + Result->Test->GetReadBackData("Output", &OutData); + VERIFY_IS_TRUE(verifyWaveArithmeticMatrix( + OutData.data(), OutData.size(), AccumulatorParams, ExpectedValues, + L"Exact non-uniform F16 accumulator plus a B-use F16 matrix", + VerboseLogging)); +} + static const char MatVecMulShader[] = R"( #define USE_A 0 #define SCOPE_THREAD 0 @@ -4807,7 +5509,7 @@ void DxilConf_SM610_LinAlg::OuterProduct_Thread_16x16_F16() { #endif // defined(DIRECT3D_LINEAR_ALGEBRA) } -static const char QueryAccumLayoutShader[] = R"( +static const char QueryAccumLayoutValueShader[] = R"( RWByteAddressBuffer Output : register(u0); [numthreads(1, 1, 1)] @@ -4817,38 +5519,170 @@ static const char QueryAccumLayoutShader[] = R"( } )"; +static void runQueryAccumLayoutValue(ID3D12Device *Device, + dxc::SpecificDllLoader &DxcSupport, + bool Verbose) { + const std::string Args = "-HV 202x"; + const size_t BufferSize = sizeof(uint32_t); + + compileShader(DxcSupport, QueryAccumLayoutValueShader, "cs_6_10", Args, + Verbose); + + auto Op = createComputeOp(QueryAccumLayoutValueShader, "cs_6_10", "UAV(u0)", + Args.c_str()); + addUAVBuffer(Op.get(), "Output", BufferSize, true, "byname"); + addRootView(Op.get(), 0, "Output"); + + auto Result = + runShaderOp(Device, DxcSupport, std::move(Op), + [](LPCSTR Name, std::vector &Data, st::ShaderOp *) { + if (_stricmp(Name, "Output") == 0) + cpu_oracle::fillPoison(Data.data(), Data.size()); + }); + + MappedData OutData; + Result->Test->GetReadBackData("Output", &OutData); + VERIFY_IS_TRUE(OutData.size() == BufferSize); + if (OutData.size() != BufferSize) + return; + + uint32_t Layout; + std::memcpy(&Layout, OutData.data(), sizeof(Layout)); + VERIFY_IS_TRUE(Layout == static_cast(MatrixUse::A) || + Layout == static_cast(MatrixUse::B)); + if (Verbose) + hlsl_test::LogCommentFmt(L"AccumulatorLayout = %u", Layout); +} + +static const char QueryAccumLayoutShader[] = R"( + #define USE_A 0 + #define USE_B 1 + #define USE_ACC 2 + + RWByteAddressBuffer Output : register(u0); + + [WaveSize(FORCED_WAVE_SIZE)] + [numthreads(NUMTHREADS, 1, 1)] + void main() { + uint Layout = __builtin_LinAlg_MatrixQueryAccumulatorLayout(); + + __builtin_LinAlgMatrix + [[__LinAlgMatrix_Attributes( + COMP_TYPE, M_DIM, N_DIM, USE_ACC, SCOPE)]] + Accumulator; + __builtin_LinAlg_FillMatrix(Accumulator, 2.0); + + __builtin_LinAlgMatrix + [[__LinAlgMatrix_Attributes(COMP_TYPE, M_DIM, N_DIM, USE_A, SCOPE)]] + MatrixA; + __builtin_LinAlg_FillMatrix(MatrixA, 3.0); + + __builtin_LinAlgMatrix + [[__LinAlgMatrix_Attributes(COMP_TYPE, M_DIM, N_DIM, USE_B, SCOPE)]] + MatrixB; + __builtin_LinAlg_FillMatrix(MatrixB, 7.0); + + __builtin_LinAlgMatrix + [[__LinAlgMatrix_Attributes( + COMP_TYPE, M_DIM, N_DIM, USE_ACC, SCOPE)]] + Result; + if (Layout == USE_A) + __builtin_LinAlg_MatrixAccumulate(Result, Accumulator, MatrixA); + else + __builtin_LinAlg_MatrixAccumulate(Result, Accumulator, MatrixB); + + __builtin_LinAlg_MatrixStoreToDescriptor( + Result, Output, 0, STRIDE, LAYOUT, 128); + + // Only lane zero publishes the queried layout, avoiding a UAV write race. + if (WaveGetLaneIndex() == 0) + Output.Store(LAYOUT_OFFSET, Layout); + } +)"; + static void runQueryAccumLayout(ID3D12Device *Device, dxc::SpecificDllLoader &DxcSupport, + MatrixParams Params, UINT SelectedWaveSize, bool Verbose) { - std::string Args = "-HV 202x"; - size_t BufferSize = elementSize(ComponentType::I32); + Params.NumThreads = static_cast(SelectedWaveSize); + const size_t MatrixBytes = Params.totalBytes(); + const size_t BufferSize = MatrixBytes + sizeof(uint32_t); + + std::stringstream ExtraDefs; + ExtraDefs << " -DFORCED_WAVE_SIZE=" << SelectedWaveSize; + ExtraDefs << " -DLAYOUT_OFFSET=" << MatrixBytes; + const std::string Args = buildCompilerArgs(Params, ExtraDefs.str().c_str()); compileShader(DxcSupport, QueryAccumLayoutShader, "cs_6_10", Args, Verbose); auto Op = createComputeOp(QueryAccumLayoutShader, "cs_6_10", "UAV(u0)", Args.c_str()); - addUAVBuffer(Op.get(), "Output", BufferSize, true); + addUAVBuffer(Op.get(), "Output", BufferSize, true, "byname"); addRootView(Op.get(), 0, "Output"); - auto Result = runShaderOp(Device, DxcSupport, std::move(Op)); + auto Result = + runShaderOp(Device, DxcSupport, std::move(Op), + [](LPCSTR Name, std::vector &Data, st::ShaderOp *) { + if (_stricmp(Name, "Output") == 0) + cpu_oracle::fillPoison(Data.data(), Data.size()); + }); MappedData OutData; Result->Test->GetReadBackData("Output", &OutData); - const uint32_t *Out = static_cast(OutData.data()); + VERIFY_IS_TRUE(OutData.size() == BufferSize); + if (OutData.size() != BufferSize) + return; + + uint32_t Layout; + std::memcpy(&Layout, static_cast(OutData.data()) + MatrixBytes, + sizeof(Layout)); + VERIFY_IS_TRUE(Layout == static_cast(MatrixUse::A) || + Layout == static_cast(MatrixUse::B)); + if (Layout != static_cast(MatrixUse::A) && + Layout != static_cast(MatrixUse::B)) + return; - // Accum Layout must be A or B - VERIFY_IS_TRUE(Out[0] == static_cast(MatrixUse::A) || - Out[0] == static_cast(MatrixUse::B)); + const int64_t ExpectedValue = + Layout == static_cast(MatrixUse::A) ? 5 : 9; + const std::vector ExpectedValues(Params.totalElements(), + ExpectedValue); + VERIFY_IS_TRUE(verifyWaveArithmeticMatrix( + OutData.data(), OutData.size(), Params, ExpectedValues, + L"Accumulator layout selects the matching A-use or B-use accumulate " + L"path", + Verbose)); if (Verbose) - hlsl_test::LogCommentFmt(L"AccumulatorLayout = %u", Out[0]); + hlsl_test::LogCommentFmt(L"AccumulatorLayout = %u", Layout); } void DxilConf_SM610_LinAlg::QueryAccumLayout() { - // Constructs no matrix, so tier support is the only capability it needs. if (!linAlgTierApplicable(D3DDevice, L"QueryAccumLayout")) return; - runQueryAccumLayout(D3DDevice, DxcSupport, VerboseLogging); + MatrixParams Params = + makeWaveArithmeticParams(ComponentType::F16, /*M=*/4, /*N=*/8, + MatrixUse::Accumulator, /*WaveSize=*/128); + + bool Supported = false; + UINT SelectedWaveSize = 0; + const HRESULT QueryResult = selectMatrixConstructionWaveSize( + D3DDevice, Params, {MatrixUse::Accumulator, MatrixUse::A, MatrixUse::B}, + Supported, SelectedWaveSize); + VERIFY_SUCCEEDED(QueryResult); + if (FAILED(QueryResult)) + return; + + if (Supported) { + VERIFY_IS_TRUE(SelectedWaveSize != 0); + if (SelectedWaveSize == 0) + return; + runQueryAccumLayout(D3DDevice, DxcSupport, Params, SelectedWaveSize, + VerboseLogging); + } else { + // The query itself is matrix-free, so preserve its tier-only coverage when + // this observable F16 tile is unavailable. + runQueryAccumLayoutValue(D3DDevice, DxcSupport, VerboseLogging); + } } static const char LoadMemoryShader[] = R"(