diff --git a/docs/DXIL.rst b/docs/DXIL.rst index 20284df0f4..ba0de6e739 100644 --- a/docs/DXIL.rst +++ b/docs/DXIL.rst @@ -3216,12 +3216,16 @@ INSTR.LINALGMATRIXDIMKVECKMISMATCH %0 vector size '%1' must b INSTR.LINALGMATRIXDIMVECTORMISMATCH %0 vector size '%1' must match input matrix M dimension '%2' INSTR.LINALGMATRIXLAYOUTREQSTRIDE %0 with layout '%1' requires stride 0. INSTR.LINALGMATRIXLOADTHREADREQUIRESBAB Loading matrix with Thread scope requires ByteAddressBuffer. +INSTR.LINALGMATRIXMATRIXKDIMMUSTMATCH K dim of A matrix '%0' must match K dim of B matrix '%1'. %2 != %3. +INSTR.LINALGMATRIXMATRIXRESDIMMUSTMATCH %0 matrix dimension '%1' must match A.MxB.N '%2'. INSTR.LINALGMATRIXNOTEXACTMATCH %0 matrix '%1' must exactly match %2 matrix '%3'. INSTR.LINALGMATRIXOUTPUTBIASVECMISMATCH Output vector element type '%0' must match bias vector element type '%1' INSTR.LINALGMATRIXREQUIRESLAYOUT2 %0 requires layout %1 or %2. INSTR.LINALGMATRIXREQUIRESRWBAB %0 requires RWByteAddressBuffer. INSTR.LINALGMATRIXSCOPEMISMATCH %0 matrix scope '%1' does not match expected scope %2. INSTR.LINALGMATRIXSCOPEMISMATCH2 %0 matrix scope '%1' does not match expected scope %2 or %3. +INSTR.LINALGMATRIXSCOPEMUSTMATCH3 Matrix scope must be the same for all matrices. %0 '%1', %2 '%3', %4 '%5'. +INSTR.LINALGMATRIXSCOPEMUSTMATCH4 Matrix scope must be the same for all matrices. %0 '%1', %2 '%3', %4 '%5', %6 '%7'. INSTR.LINALGMATRIXSCOPEREQLAYOUT2 %0 matrix with scope '%1' requires layout %2 or %3 for %4. INSTR.LINALGMATRIXUNSIGNEDFLOATTYPENOTALLOWED Float-like type '%0' must be signed INSTR.LINALGMATRIXUSEMISMATCH %0 matrix use '%1' does not match expected use %2. diff --git a/lib/DxilValidation/DxilValidation.cpp b/lib/DxilValidation/DxilValidation.cpp index b12149ca82..2a3a6e4551 100644 --- a/lib/DxilValidation/DxilValidation.cpp +++ b/lib/DxilValidation/DxilValidation.cpp @@ -1451,12 +1451,198 @@ static void ValidateLinAlgMatrixMultiply(CallInst *CI, ValidationContext &ValCtx) { ValidateLinAlgOpReturnMatrix(CI, ValCtx); ValidateLinAlgOpParameters(CI, ValCtx); + DxilInst_LinAlgMatrixMultiply Op(CI); + std::optional RetMat = + GetCheckedLATT(CI->getType(), ValCtx); + if (!RetMat) + return; + std::optional AMat = + GetCheckedLATT(Op.get_matrixA()->getType(), ValCtx); + if (!AMat) + return; + std::optional BMat = + GetCheckedLATT(Op.get_matrixB()->getType(), ValCtx); + if (!BMat) + return; + + // A is an A matrix + if (AMat->Use != DXIL::MatrixUse::A) + ValCtx.EmitInstrFormatError(CI, + ValidationRule::InstrLinAlgMatrixUseMismatch, + {"A", MatrixUseToString(AMat->Use), "A"}); + + // B is a B matrix + if (BMat->Use != DXIL::MatrixUse::B) + ValCtx.EmitInstrFormatError(CI, + ValidationRule::InstrLinAlgMatrixUseMismatch, + {"B", MatrixUseToString(BMat->Use), "B"}); + + // Ret is an Accumulator matrix + if (RetMat->Use != DXIL::MatrixUse::Accumulator) + ValCtx.EmitInstrFormatError( + CI, ValidationRule::InstrLinAlgMatrixUseMismatch, + {"Return", MatrixUseToString(RetMat->Use), "Accumulator"}); + + // A scope must be Wave or ThreadGroup + if (AMat->Scope != DXIL::MatrixScope::Wave && + AMat->Scope != DXIL::MatrixScope::ThreadGroup) + ValCtx.EmitInstrFormatError( + CI, ValidationRule::InstrLinAlgMatrixScopeMismatch2, + {"A", MatrixScopeToString(AMat->Scope), "Wave", "ThreadGroup"}); + + // B scope must be Wave or ThreadGroup + if (BMat->Scope != DXIL::MatrixScope::Wave && + BMat->Scope != DXIL::MatrixScope::ThreadGroup) + ValCtx.EmitInstrFormatError( + CI, ValidationRule::InstrLinAlgMatrixScopeMismatch2, + {"B", MatrixScopeToString(BMat->Scope), "Wave", "ThreadGroup"}); + + // Ret scope must be Wave or ThreadGroup + if (RetMat->Scope != DXIL::MatrixScope::Wave && + RetMat->Scope != DXIL::MatrixScope::ThreadGroup) + ValCtx.EmitInstrFormatError( + CI, ValidationRule::InstrLinAlgMatrixScopeMismatch2, + {"Return", MatrixScopeToString(RetMat->Scope), "Wave", "ThreadGroup"}); + + // A, B, Ret scope must all be the same + if (AMat->Scope != BMat->Scope || BMat->Scope != RetMat->Scope) + ValCtx.EmitInstrFormatError( + CI, ValidationRule::InstrLinAlgMatrixScopeMustMatch3, + {"A", MatrixScopeToString(AMat->Scope), "B", + MatrixScopeToString(BMat->Scope), "Return", + MatrixScopeToString(RetMat->Scope)}); + + unsigned M = AMat->M; + unsigned AK = AMat->N; + unsigned BK = BMat->M; + unsigned N = BMat->N; + + // K dim must match between A and B + if (AK != BK) + ValCtx.EmitInstrFormatError( + CI, ValidationRule::InstrLinAlgMatrixMatrixKDimMustMatch, + {std::to_string(M) + "x" + std::to_string(AK), + std::to_string(BK) + "x" + std::to_string(N), std::to_string(AK), + std::to_string(BK)}); + + // Return dim must match A.M x B.N + if (RetMat->M != M || RetMat->N != N) + ValCtx.EmitInstrFormatError( + CI, ValidationRule::InstrLinAlgMatrixMatrixResDimMustMatch, + {"Return", std::to_string(RetMat->M) + "x" + std::to_string(RetMat->N), + std::to_string(M) + "x" + std::to_string(N)}); } static void ValidateLinAlgMatrixMultiplyAccumulate(CallInst *CI, ValidationContext &ValCtx) { ValidateLinAlgOpReturnMatrix(CI, ValCtx); ValidateLinAlgOpParameters(CI, ValCtx); + DxilInst_LinAlgMatrixMultiplyAccumulate Op(CI); + std::optional RetMat = + GetCheckedLATT(CI->getType(), ValCtx); + if (!RetMat) + return; + std::optional AMat = + GetCheckedLATT(Op.get_matrixA()->getType(), ValCtx); + if (!AMat) + return; + std::optional BMat = + GetCheckedLATT(Op.get_matrixB()->getType(), ValCtx); + if (!BMat) + return; + std::optional CMat = + GetCheckedLATT(Op.get_matrixC()->getType(), ValCtx); + if (!CMat) + return; + + // A is an A matrix + if (AMat->Use != DXIL::MatrixUse::A) + ValCtx.EmitInstrFormatError(CI, + ValidationRule::InstrLinAlgMatrixUseMismatch, + {"A", MatrixUseToString(AMat->Use), "A"}); + + // B is a B matrix + if (BMat->Use != DXIL::MatrixUse::B) + ValCtx.EmitInstrFormatError(CI, + ValidationRule::InstrLinAlgMatrixUseMismatch, + {"B", MatrixUseToString(BMat->Use), "B"}); + + // C is an Accumulator matrix + if (CMat->Use != DXIL::MatrixUse::Accumulator) + ValCtx.EmitInstrFormatError( + CI, ValidationRule::InstrLinAlgMatrixUseMismatch, + {"C", MatrixUseToString(CMat->Use), "Accumulator"}); + + // Ret is an Accumulator matrix + if (RetMat->Use != DXIL::MatrixUse::Accumulator) + ValCtx.EmitInstrFormatError( + CI, ValidationRule::InstrLinAlgMatrixUseMismatch, + {"Return", MatrixUseToString(RetMat->Use), "Accumulator"}); + + // A scope must be Wave or ThreadGroup + if (AMat->Scope != DXIL::MatrixScope::Wave && + AMat->Scope != DXIL::MatrixScope::ThreadGroup) + ValCtx.EmitInstrFormatError( + CI, ValidationRule::InstrLinAlgMatrixScopeMismatch2, + {"A", MatrixScopeToString(AMat->Scope), "Wave", "ThreadGroup"}); + + // B scope must be Wave or ThreadGroup + if (BMat->Scope != DXIL::MatrixScope::Wave && + BMat->Scope != DXIL::MatrixScope::ThreadGroup) + ValCtx.EmitInstrFormatError( + CI, ValidationRule::InstrLinAlgMatrixScopeMismatch2, + {"B", MatrixScopeToString(BMat->Scope), "Wave", "ThreadGroup"}); + + // C scope must be Wave or ThreadGroup + if (CMat->Scope != DXIL::MatrixScope::Wave && + CMat->Scope != DXIL::MatrixScope::ThreadGroup) + ValCtx.EmitInstrFormatError( + CI, ValidationRule::InstrLinAlgMatrixScopeMismatch2, + {"C", MatrixScopeToString(CMat->Scope), "Wave", "ThreadGroup"}); + + // Ret scope must be Wave or ThreadGroup + if (RetMat->Scope != DXIL::MatrixScope::Wave && + RetMat->Scope != DXIL::MatrixScope::ThreadGroup) + ValCtx.EmitInstrFormatError( + CI, ValidationRule::InstrLinAlgMatrixScopeMismatch2, + {"Return", MatrixScopeToString(RetMat->Scope), "Wave", "ThreadGroup"}); + + // A, B, C, Ret scope must all be the same + if (AMat->Scope != BMat->Scope || BMat->Scope != CMat->Scope || + CMat->Scope != RetMat->Scope) + ValCtx.EmitInstrFormatError( + CI, ValidationRule::InstrLinAlgMatrixScopeMustMatch4, + {"A", MatrixScopeToString(AMat->Scope), "B", + MatrixScopeToString(BMat->Scope), "C", + MatrixScopeToString(CMat->Scope), "Return", + MatrixScopeToString(RetMat->Scope)}); + + unsigned M = AMat->M; + unsigned AK = AMat->N; + unsigned BK = BMat->M; + unsigned N = BMat->N; + + // K dim must match between A and B + if (AK != BK) + ValCtx.EmitInstrFormatError( + CI, ValidationRule::InstrLinAlgMatrixMatrixKDimMustMatch, + {std::to_string(M) + "x" + std::to_string(AK), + std::to_string(BK) + "x" + std::to_string(N), std::to_string(AK), + std::to_string(BK)}); + + // C dim must match A.M x B.N + if (CMat->M != M || CMat->N != N) + ValCtx.EmitInstrFormatError( + CI, ValidationRule::InstrLinAlgMatrixMatrixResDimMustMatch, + {"C", std::to_string(CMat->M) + "x" + std::to_string(CMat->N), + std::to_string(M) + "x" + std::to_string(N)}); + + // Return dim must match A.M x B.N + if (RetMat->M != M || RetMat->N != N) + ValCtx.EmitInstrFormatError( + CI, ValidationRule::InstrLinAlgMatrixMatrixResDimMustMatch, + {"Return", std::to_string(RetMat->M) + "x" + std::to_string(RetMat->N), + std::to_string(M) + "x" + std::to_string(N)}); } static void ValidateLinAlgMatrixOuterProduct(CallInst *CI, diff --git a/tools/clang/test/CodeGenDXIL/hlsl/linalg/builtins/matrixmatrixmultiply/nominal.hlsl b/tools/clang/test/CodeGenDXIL/hlsl/linalg/builtins/matrixmatrixmultiply/nominal.hlsl index 817a4acefe..6c425d8938 100644 --- a/tools/clang/test/CodeGenDXIL/hlsl/linalg/builtins/matrixmatrixmultiply/nominal.hlsl +++ b/tools/clang/test/CodeGenDXIL/hlsl/linalg/builtins/matrixmatrixmultiply/nominal.hlsl @@ -6,14 +6,22 @@ void main() { // CHECK-LABEL: define void @main() - // CHECK: call %dx.types.LinAlgMatrixC4M5N4U1S2 @dx.op.linAlgMatrixMultiply.mC4M5N4U1S2.mC4M5N4U1S2.mC4M5N4U1S2(i32 -2147483625, - // CHECK-SAME: %dx.types.LinAlgMatrixC4M5N4U1S2 %{{.*}}, %dx.types.LinAlgMatrixC4M5N4U1S2 %{{.*}}) ; LinAlgMatrixMultiply(matrixA,matrixB) + // Matrix + __builtin_LinAlgMatrix [[__LinAlgMatrix_Attributes(4, 8, 4, 0, 2)]] matA; + // Matrix + __builtin_LinAlgMatrix [[__LinAlgMatrix_Attributes(4, 4, 8, 1, 2)]] matB; + // Matrix + __builtin_LinAlgMatrix [[__LinAlgMatrix_Attributes(4, 8, 8, 2, 2)]] matC; - // CHECK2: call void @"dx.hl.op..void (i32, %dx.types.LinAlgMatrixC4M5N4U1S2*, %dx.types.LinAlgMatrixC4M5N4U1S2, - // CHECK2-SAME: %dx.types.LinAlgMatrixC4M5N4U1S2)"(i32 412, %dx.types.LinAlgMatrixC4M5N4U1S2* %mat2, - // CHECK2-SAME: %dx.types.LinAlgMatrixC4M5N4U1S2 %{{[0-9]+}}, %dx.types.LinAlgMatrixC4M5N4U1S2 %{{[0-9]+}}) - __builtin_LinAlgMatrix [[__LinAlgMatrix_Attributes(4, 5, 4, 1, 2)]] mat1; - __builtin_LinAlg_FillMatrix(mat1, 1); - __builtin_LinAlgMatrix [[__LinAlgMatrix_Attributes(4, 5, 4, 1, 2)]] mat2; - __builtin_LinAlg_MatrixMatrixMultiply(mat2, mat1, mat1); + __builtin_LinAlg_FillMatrix(matA, 1); + __builtin_LinAlg_FillMatrix(matB, 2); + + // CHECK: call %dx.types.LinAlgMatrixC4M8N8U2S2 @dx.op.linAlgMatrixMultiply.mC4M8N8U2S2.mC4M8N4U0S2.mC4M4N8U1S2(i32 -2147483625, + // CHECK-SAME: %dx.types.LinAlgMatrixC4M8N4U0S2 %{{.*}}, %dx.types.LinAlgMatrixC4M4N8U1S2 %{{.*}}) ; LinAlgMatrixMultiply(matrixA,matrixB) + + // CHECK2: call void @"dx.hl.op..void (i32, %dx.types.LinAlgMatrixC4M8N8U2S2*, %dx.types.LinAlgMatrixC4M8N4U0S2, + // CHECK2-SAME: %dx.types.LinAlgMatrixC4M4N8U1S2)"(i32 412, %dx.types.LinAlgMatrixC4M8N8U2S2* %matC, + // CHECK2-SAME: %dx.types.LinAlgMatrixC4M8N4U0S2 %{{[0-9]+}}, %dx.types.LinAlgMatrixC4M4N8U1S2 %{{[0-9]+}}) + + __builtin_LinAlg_MatrixMatrixMultiply(matC, matA, matB); } diff --git a/tools/clang/test/CodeGenDXIL/hlsl/linalg/builtins/matrixmatrixmultiplyaccumulate/nominal.hlsl b/tools/clang/test/CodeGenDXIL/hlsl/linalg/builtins/matrixmatrixmultiplyaccumulate/nominal.hlsl index 458f08c6a4..1cb0602401 100644 --- a/tools/clang/test/CodeGenDXIL/hlsl/linalg/builtins/matrixmatrixmultiplyaccumulate/nominal.hlsl +++ b/tools/clang/test/CodeGenDXIL/hlsl/linalg/builtins/matrixmatrixmultiplyaccumulate/nominal.hlsl @@ -12,20 +12,28 @@ void main() { // CHECK: ; LinAlgFillMatrix(value) // CHECK: ; LinAlgFillMatrix(value) - // CHECK: call %dx.types.LinAlgMatrixC4M5N3U1S2 - // CHECK-SAME: @dx.op.linAlgMatrixMultiplyAccumulate.mC4M5N3U1S2.mC4M5N4U1S2.mC4M4N3U1S2.mC4M5N3U1S2 - // CHECK-SAME: (i32 -2147483637, %dx.types.LinAlgMatrixC4M5N4U1S2 %{{[0-9]+}}, %dx.types.LinAlgMatrixC4M4N3U1S2 %{{[0-9]+}}, - // CHECK-SAME: %dx.types.LinAlgMatrixC4M5N3U1S2 %{{[0-9]+}}) ; LinAlgMatrixMultiplyAccumulate(matrixA,matrixB,matrixC) + // Matrix + __builtin_LinAlgMatrix [[__LinAlgMatrix_Attributes(4, 8, 4, 0, 2)]] matA; + // Matrix + __builtin_LinAlgMatrix [[__LinAlgMatrix_Attributes(4, 4, 8, 1, 2)]] matB; + // Matrix + __builtin_LinAlgMatrix [[__LinAlgMatrix_Attributes(4, 8, 8, 2, 2)]] matC; + // Matrix + __builtin_LinAlgMatrix [[__LinAlgMatrix_Attributes(4, 8, 8, 2, 2)]] matR; - // CHECK2: call void @"dx.hl.op..void (i32, %dx.types.LinAlgMatrixC4M5N3U1S2*, %dx.types.LinAlgMatrixC4M5N4U1S2, - // CHECK2-SAME: %dx.types.LinAlgMatrixC4M4N3U1S2, %dx.types.LinAlgMatrixC4M5N3U1S2)"(i32 413, - // CHECK2-SAME: %dx.types.LinAlgMatrixC4M5N3U1S2* %{{.*}}, %dx.types.LinAlgMatrixC4M5N4U1S2 %{{[0-9]+}}, - // CHECK2-SAME: %dx.types.LinAlgMatrixC4M4N3U1S2 %{{[0-9]+}}, %dx.types.LinAlgMatrixC4M5N3U1S2 %{{[0-9]+}}) - __builtin_LinAlgMatrix [[__LinAlgMatrix_Attributes(4, 5, 4, 1, 2)]] mat1; - __builtin_LinAlg_FillMatrix(mat1, 1); - __builtin_LinAlgMatrix [[__LinAlgMatrix_Attributes(4, 4, 3, 1, 2)]] mat2; - __builtin_LinAlg_FillMatrix(mat2, 2); - __builtin_LinAlgMatrix [[__LinAlgMatrix_Attributes(4, 5, 3, 1, 2)]] mat3; - __builtin_LinAlg_FillMatrix(mat3, 3); - __builtin_LinAlg_MatrixMatrixMultiplyAccumulate(mat3, mat1, mat2, mat3); + __builtin_LinAlg_FillMatrix(matA, 1); + __builtin_LinAlg_FillMatrix(matB, 2); + __builtin_LinAlg_FillMatrix(matC, 3); + + // CHECK: call %dx.types.LinAlgMatrixC4M8N8U2S2 + // CHECK-SAME: @dx.op.linAlgMatrixMultiplyAccumulate.mC4M8N8U2S2.mC4M8N4U0S2.mC4M4N8U1S2.mC4M8N8U2S2 + // CHECK-SAME: (i32 -2147483637, %dx.types.LinAlgMatrixC4M8N4U0S2 %{{[0-9]+}}, %dx.types.LinAlgMatrixC4M4N8U1S2 %{{[0-9]+}}, + // CHECK-SAME: %dx.types.LinAlgMatrixC4M8N8U2S2 %{{[0-9]+}}) ; LinAlgMatrixMultiplyAccumulate(matrixA,matrixB,matrixC) + + // CHECK2: call void @"dx.hl.op..void (i32, %dx.types.LinAlgMatrixC4M8N8U2S2*, %dx.types.LinAlgMatrixC4M8N4U0S2, + // CHECK2-SAME: %dx.types.LinAlgMatrixC4M4N8U1S2, %dx.types.LinAlgMatrixC4M8N8U2S2)"(i32 413, + // CHECK2-SAME: %dx.types.LinAlgMatrixC4M8N8U2S2* %{{.*}}, %dx.types.LinAlgMatrixC4M8N4U0S2 %{{[0-9]+}}, + // CHECK2-SAME: %dx.types.LinAlgMatrixC4M4N8U1S2 %{{[0-9]+}}, %dx.types.LinAlgMatrixC4M8N8U2S2 %{{[0-9]+}}) + + __builtin_LinAlg_MatrixMatrixMultiplyAccumulate(matR, matA, matB, matC); } diff --git a/tools/clang/test/LitDXILValidation/LinAlgMatrix/linalgmatrix-cs.ll b/tools/clang/test/LitDXILValidation/LinAlgMatrix/linalgmatrix-cs.ll index 8d1c2ccb75..012e4ed609 100644 --- a/tools/clang/test/LitDXILValidation/LinAlgMatrix/linalgmatrix-cs.ll +++ b/tools/clang/test/LitDXILValidation/LinAlgMatrix/linalgmatrix-cs.ll @@ -8,9 +8,9 @@ target triple = "dxil-ms-dx" %dx.types.Handle = type { i8* } %dx.types.ResBind = type { i32, i32, i32, i8 } -%dx.types.LinAlgMatrixC4M5N4U2S2 = type { i8* } -%dx.types.LinAlgMatrixC4M5N4U0S2 = type { i8* } -%dx.types.LinAlgMatrixC4M4N5U1S2 = type { i8* } +%dx.types.LinAlgMatrixC4M4N4U2S2 = type { i8* } +%dx.types.LinAlgMatrixC4M4N4U0S2 = type { i8* } +%dx.types.LinAlgMatrixC4M4N4U1S2 = type { i8* } %dx.types.LinAlgMatrixC4M4N5U2S2 = type { i8* } %dx.types.LinAlgMatrixC4M4N4U0S0 = type { i8* } %dx.types.ResourceProperties = type { i32, i32 } @@ -32,27 +32,29 @@ define void @mainCS() { ; Matrix %mC4M4N4U0S0 = call %dx.types.LinAlgMatrixC4M4N4U0S0 @dx.op.linAlgMatrixLoadFromDescriptor.mC4M4N4U0S0(i32 -2147483634, %dx.types.Handle %bab, i32 0, i32 0, i32 0, i32 128) - ; Matrix - %mC4M5N4U0S2 = call %dx.types.LinAlgMatrixC4M5N4U0S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC4M5N4U0S2(i32 -2147483634, %dx.types.Handle %bab, i32 0, i32 0, i32 0, i32 128) - ; Matrix - %mC4M4N5U1S2 = call %dx.types.LinAlgMatrixC4M4N5U1S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC4M4N5U1S2(i32 -2147483634, %dx.types.Handle %bab, i32 0, i32 0, i32 0, i32 128) + ; Matrix + %mC4M4N4U0S2 = call %dx.types.LinAlgMatrixC4M4N4U0S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC4M4N4U0S2(i32 -2147483634, %dx.types.Handle %bab, i32 0, i32 0, i32 0, i32 128) + ; Matrix + %mC4M4N4U1S2 = call %dx.types.LinAlgMatrixC4M4N4U1S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC4M4N4U1S2(i32 -2147483634, %dx.types.Handle %bab, i32 0, i32 0, i32 0, i32 128) ; Matrix %mC4M4N5U2S2 = call %dx.types.LinAlgMatrixC4M4N5U2S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC4M4N5U2S2(i32 -2147483634, %dx.types.Handle %bab, i32 0, i32 0, i32 0, i32 128) + ; Matrix + %mC4M4N4U2S2 = call %dx.types.LinAlgMatrixC4M4N4U2S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC4M4N4U2S2(i32 -2147483634, %dx.types.Handle %bab, i32 0, i32 0, i32 0, i32 128) ; dx.op.linAlgMatrixAccumulate - %v1 = call %dx.types.LinAlgMatrixC4M4N5U2S2 @dx.op.linAlgMatrixAccumulate.mC4M4N5U2S2.mC4M4N5U2S2.mC4M4N5U1S2(i32 -2147483624, %dx.types.LinAlgMatrixC4M4N5U2S2 %mC4M4N5U2S2, %dx.types.LinAlgMatrixC4M4N5U1S2 %mC4M4N5U1S2) ; LinAlgMatrixAccumulate(matrixLHS,matrixRHS) + %v1 = call %dx.types.LinAlgMatrixC4M4N4U2S2 @dx.op.linAlgMatrixAccumulate.mC4M4N4U2S2.mC4M4N4U2S2.mC4M4N4U1S2(i32 -2147483624, %dx.types.LinAlgMatrixC4M4N4U2S2 %mC4M4N4U2S2, %dx.types.LinAlgMatrixC4M4N4U1S2 %mC4M4N4U1S2) ; LinAlgMatrixAccumulate(matrixLHS,matrixRHS) ; dx.op.linAlgMatrixAccumulateToDescriptor call void @dx.op.linAlgMatrixAccumulateToDescriptor.mC4M4N5U2S2(i32 -2147483621, %dx.types.LinAlgMatrixC4M4N5U2S2 %mC4M4N5U2S2, %dx.types.Handle %rwbab, i32 1, i32 2, i32 0, i32 128) ; LinAlgMatrixAccumulateToDescriptor(matrix,handle,offset,stride,layout,align) ; dx.op.linAlgMatrixLength - %v2 = call i32 @dx.op.linAlgMatrixLength.mC4M5N4U0S2(i32 -2147483632, %dx.types.LinAlgMatrixC4M5N4U0S2 %mC4M5N4U0S2) ; LinAlgMatrixLength(matrix) + %v2 = call i32 @dx.op.linAlgMatrixLength.mC4M4N4U0S2(i32 -2147483632, %dx.types.LinAlgMatrixC4M4N4U0S2 %mC4M4N4U0S2) ; LinAlgMatrixLength(matrix) ; dx.op.linAlgMatrixLoadFromDescriptor - %v3 = call %dx.types.LinAlgMatrixC4M5N4U0S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC4M5N4U0S2(i32 -2147483634, %dx.types.Handle %bab, i32 0, i32 0, i32 0, i32 128) ; LinAlgMatrixLoadFromDescriptor(handle,offset,stride,layout,align) + %v3 = call %dx.types.LinAlgMatrixC4M4N4U0S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC4M4N4U0S2(i32 -2147483634, %dx.types.Handle %bab, i32 0, i32 0, i32 0, i32 128) ; LinAlgMatrixLoadFromDescriptor(handle,offset,stride,layout,align) ; dx.op.linAlgMatrixOuterProduct - %v4 = call %dx.types.LinAlgMatrixC4M5N4U0S2 @dx.op.linAlgMatrixOuterProduct.mC4M5N4U0S2.v4i32.v4i32(i32 -2147483619, <4 x i32> , <4 x i32> ) ; LinAlgMatrixOuterProduct(vectorA,vectorB) + %v4 = call %dx.types.LinAlgMatrixC4M4N4U0S2 @dx.op.linAlgMatrixOuterProduct.mC4M4N4U0S2.v4i32.v4i32(i32 -2147483619, <4 x i32> , <4 x i32> ) ; LinAlgMatrixOuterProduct(vectorA,vectorB) ; dx.op.linAlgMatrixQueryAccumulatorLayout %v5 = call i32 @dx.op.linAlgMatrixQueryAccumulatorLayout(i32 -2147483626) ; LinAlgMatrixQueryAccumulatorLayout() @@ -74,61 +76,64 @@ define void @mainCS() { ; ; dx.op.linAlgCopyConvertMatrix - %v8 = call %dx.types.LinAlgMatrixC4M4N5U1S2 @dx.op.linAlgCopyConvertMatrix.mC4M4N5U1S2.mC4M5N4U0S2(i32 -2147483635, %dx.types.LinAlgMatrixC4M5N4U0S2 %v4, i1 true) ; LinAlgCopyConvertMatrix(srcMatrix,transpose) + %v8 = call %dx.types.LinAlgMatrixC4M4N4U1S2 @dx.op.linAlgCopyConvertMatrix.mC4M4N4U1S2.mC4M4N4U0S2(i32 -2147483635, %dx.types.LinAlgMatrixC4M4N4U0S2 %v4, i1 true) ; LinAlgCopyConvertMatrix(srcMatrix,transpose) ; dx.op.linAlgFillMatrix - %v9 = call %dx.types.LinAlgMatrixC4M5N4U0S2 @dx.op.linAlgFillMatrix.mC4M5N4U0S2.i32(i32 -2147483636, i32 15) ; LinAlgFillMatrix(value) + %v9 = call %dx.types.LinAlgMatrixC4M4N4U0S2 @dx.op.linAlgFillMatrix.mC4M4N4U0S2.i32(i32 -2147483636, i32 15) ; LinAlgFillMatrix(value) ; dx.op.linAlgMatrixGetCoordinate - %v10 = call <2 x i32> @dx.op.linAlgMatrixGetCoordinate.mC4M5N4U0S2(i32 -2147483631, %dx.types.LinAlgMatrixC4M5N4U0S2 %v9, i32 0) ; LinAlgMatrixGetCoordinate(matrix,threadLocalIndex) + %v10 = call <2 x i32> @dx.op.linAlgMatrixGetCoordinate.mC4M4N4U0S2(i32 -2147483631, %dx.types.LinAlgMatrixC4M4N4U0S2 %v9, i32 0) ; LinAlgMatrixGetCoordinate(matrix,threadLocalIndex) ; dx.op.linAlgMatrixGetElement - %v11 = call float @dx.op.linAlgMatrixGetElement.f32.mC4M5N4U0S2(i32 -2147483630, %dx.types.LinAlgMatrixC4M5N4U0S2 %v9, i32 0) ; LinAlgMatrixGetElement(matrix,threadLocalIndex) + %v11 = call float @dx.op.linAlgMatrixGetElement.f32.mC4M4N4U0S2(i32 -2147483630, %dx.types.LinAlgMatrixC4M4N4U0S2 %v9, i32 0) ; LinAlgMatrixGetElement(matrix,threadLocalIndex) ; dx.op.linAlgMatrixMultiply - %v12 = call %dx.types.LinAlgMatrixC4M5N4U2S2 @dx.op.linAlgMatrixMultiply.mC4M5N4U2S2.mC4M5N4U0S2.mC4M4N5U1S2(i32 -2147483625, %dx.types.LinAlgMatrixC4M5N4U0S2 %v9, %dx.types.LinAlgMatrixC4M4N5U1S2 %v8) ; LinAlgMatrixMultiply(matrixA,matrixB) + %v12 = call %dx.types.LinAlgMatrixC4M4N4U2S2 @dx.op.linAlgMatrixMultiply.mC4M4N4U2S2.mC4M4N4U0S2.mC4M4N4U1S2(i32 -2147483625, %dx.types.LinAlgMatrixC4M4N4U0S2 %v9, %dx.types.LinAlgMatrixC4M4N4U1S2 %v8) ; LinAlgMatrixMultiply(matrixA,matrixB) ; dx.op.linAlgMatrixMultiplyAccumulate - %v13 = call %dx.types.LinAlgMatrixC4M5N4U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC4M5N4U2S2.mC4M5N4U0S2.mC4M4N5U1S2.mC4M5N4U2S2(i32 -2147483637, %dx.types.LinAlgMatrixC4M5N4U0S2 %v9, %dx.types.LinAlgMatrixC4M4N5U1S2 %v8, %dx.types.LinAlgMatrixC4M5N4U2S2 %v12) ; LinAlgMatrixMultiplyAccumulate(matrixA,matrixB,matrixC) + %v13 = call %dx.types.LinAlgMatrixC4M4N4U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC4M4N4U2S2.mC4M4N4U0S2.mC4M4N4U1S2.mC4M4N4U2S2(i32 -2147483637, %dx.types.LinAlgMatrixC4M4N4U0S2 %v9, %dx.types.LinAlgMatrixC4M4N4U1S2 %v8, %dx.types.LinAlgMatrixC4M4N4U2S2 %v12) ; LinAlgMatrixMultiplyAccumulate(matrixA,matrixB,matrixC) ; dx.op.linAlgMatrixSetElement - %v14 = call %dx.types.LinAlgMatrixC4M5N4U0S2 @dx.op.linAlgMatrixSetElement.mC4M5N4U0S2.mC4M5N4U0S2.i32(i32 -2147483629, %dx.types.LinAlgMatrixC4M5N4U0S2 %v9, i32 1, i32 1) ; LinAlgMatrixSetElement(matrix,threadLocalIndex,value) + %v14 = call %dx.types.LinAlgMatrixC4M4N4U0S2 @dx.op.linAlgMatrixSetElement.mC4M4N4U0S2.mC4M4N4U0S2.i32(i32 -2147483629, %dx.types.LinAlgMatrixC4M4N4U0S2 %v9, i32 1, i32 1) ; LinAlgMatrixSetElement(matrix,threadLocalIndex,value) ; dx.op.linAlgMatrixStoreToDescriptor - call void @dx.op.linAlgMatrixStoreToDescriptor.mC4M5N4U0S2(i32 -2147483628, %dx.types.LinAlgMatrixC4M5N4U0S2 %v14, %dx.types.Handle %rwbab, i32 1, i32 2, i32 0, i32 128) ; LinAlgMatrixStoreToDescriptor(matrix,handle,offset,stride,layout,align) + call void @dx.op.linAlgMatrixStoreToDescriptor.mC4M4N4U0S2(i32 -2147483628, %dx.types.LinAlgMatrixC4M4N4U0S2 %v14, %dx.types.Handle %rwbab, i32 1, i32 2, i32 0, i32 128) ; LinAlgMatrixStoreToDescriptor(matrix,handle,offset,stride,layout,align) ; dx.op.linAlgMatrixAccumulateToMemory - call void @dx.op.linAlgMatrixAccumulateToMemory.mC4M5N4U0S2.f32(i32 -2147483620, %dx.types.LinAlgMatrixC4M5N4U0S2 %v14, float addrspace(3)* getelementptr inbounds ([64 x float], [64 x float] addrspace(3)* @"\01?SharedArr@@3PAMA", i32 0, i32 0), i32 0, i32 0, i32 0, i32 0) ; LinAlgMatrixAccumulateToMemory(matrix,memory,targetType,offset,stride,layout) + call void @dx.op.linAlgMatrixAccumulateToMemory.mC4M4N4U0S2.f32(i32 -2147483620, %dx.types.LinAlgMatrixC4M4N4U0S2 %v14, float addrspace(3)* getelementptr inbounds ([64 x float], [64 x float] addrspace(3)* @"\01?SharedArr@@3PAMA", i32 0, i32 0), i32 0, i32 0, i32 0, i32 0) ; LinAlgMatrixAccumulateToMemory(matrix,memory,targetType,offset,stride,layout) ; dx.op.linAlgMatrixLoadFromMemory - %v15 = call %dx.types.LinAlgMatrixC4M5N4U0S2 @dx.op.linAlgMatrixLoadFromMemory.mC4M5N4U0S2.f32(i32 -2147483633, float addrspace(3)* getelementptr inbounds ([64 x float], [64 x float] addrspace(3)* @"\01?SharedArr@@3PAMA", i32 0, i32 0), i32 0, i32 0, i32 0) ; LinAlgMatrixLoadFromMemory(memory,offset,stride,layout) + %v15 = call %dx.types.LinAlgMatrixC4M4N4U0S2 @dx.op.linAlgMatrixLoadFromMemory.mC4M4N4U0S2.f32(i32 -2147483633, float addrspace(3)* getelementptr inbounds ([64 x float], [64 x float] addrspace(3)* @"\01?SharedArr@@3PAMA", i32 0, i32 0), i32 0, i32 0, i32 0) ; LinAlgMatrixLoadFromMemory(memory,offset,stride,layout) ; dx.op.linAlgMatrixStoreToMemory - call void @dx.op.linAlgMatrixStoreToMemory.mC4M5N4U0S2.f32(i32 -2147483627, %dx.types.LinAlgMatrixC4M5N4U0S2 %v15, float addrspace(3)* getelementptr inbounds ([64 x float], [64 x float] addrspace(3)* @"\01?SharedArr@@3PAMA", i32 0, i32 0), i32 0, i32 0, i32 0) ; LinAlgMatrixStoreToMemory(matrix,memory,offset,stride,layout) + call void @dx.op.linAlgMatrixStoreToMemory.mC4M4N4U0S2.f32(i32 -2147483627, %dx.types.LinAlgMatrixC4M4N4U0S2 %v15, float addrspace(3)* getelementptr inbounds ([64 x float], [64 x float] addrspace(3)* @"\01?SharedArr@@3PAMA", i32 0, i32 0), i32 0, i32 0, i32 0) ; LinAlgMatrixStoreToMemory(matrix,memory,offset,stride,layout) ret void } ; Function Attrs: nounwind -declare %dx.types.LinAlgMatrixC4M5N4U2S2 @dx.op.linAlgMatrixMultiply.mC4M5N4U2S2.mC4M5N4U0S2.mC4M4N5U1S2(i32, %dx.types.LinAlgMatrixC4M5N4U0S2, %dx.types.LinAlgMatrixC4M4N5U1S2) #0 +declare %dx.types.LinAlgMatrixC4M4N4U2S2 @dx.op.linAlgMatrixMultiply.mC4M4N4U2S2.mC4M4N4U0S2.mC4M4N4U1S2(i32, %dx.types.LinAlgMatrixC4M4N4U0S2, %dx.types.LinAlgMatrixC4M4N4U1S2) #0 ; Function Attrs: nounwind -declare %dx.types.LinAlgMatrixC4M4N5U2S2 @dx.op.linAlgMatrixAccumulate.mC4M4N5U2S2.mC4M4N5U2S2.mC4M4N5U1S2(i32, %dx.types.LinAlgMatrixC4M4N5U2S2, %dx.types.LinAlgMatrixC4M4N5U1S2) #0 +declare %dx.types.LinAlgMatrixC4M4N4U2S2 @dx.op.linAlgMatrixAccumulate.mC4M4N4U2S2.mC4M4N4U2S2.mC4M4N4U1S2(i32, %dx.types.LinAlgMatrixC4M4N4U2S2, %dx.types.LinAlgMatrixC4M4N4U1S2) #0 ; Function Attrs: nounwind -declare void @dx.op.linAlgMatrixStoreToDescriptor.mC4M5N4U0S2(i32, %dx.types.LinAlgMatrixC4M5N4U0S2, %dx.types.Handle, i32, i32, i32, i32) #0 +declare void @dx.op.linAlgMatrixStoreToDescriptor.mC4M4N4U0S2(i32, %dx.types.LinAlgMatrixC4M4N4U0S2, %dx.types.Handle, i32, i32, i32, i32) #0 ; Function Attrs: nounwind declare void @dx.op.linAlgMatrixAccumulateToDescriptor.mC4M4N5U2S2(i32, %dx.types.LinAlgMatrixC4M4N5U2S2, %dx.types.Handle, i32, i32, i32, i32) #0 ; Function Attrs: nounwind -declare i32 @dx.op.linAlgMatrixLength.mC4M5N4U0S2(i32, %dx.types.LinAlgMatrixC4M5N4U0S2) #0 +declare i32 @dx.op.linAlgMatrixLength.mC4M4N4U0S2(i32, %dx.types.LinAlgMatrixC4M4N4U0S2) #0 ; Function Attrs: nounwind -declare %dx.types.LinAlgMatrixC4M5N4U0S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC4M5N4U0S2(i32, %dx.types.Handle, i32, i32, i32, i32) #0 +declare %dx.types.LinAlgMatrixC4M4N4U0S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC4M4N4U0S2(i32, %dx.types.Handle, i32, i32, i32, i32) #0 ; Function Attrs: nounwind -declare %dx.types.LinAlgMatrixC4M4N5U1S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC4M4N5U1S2(i32, %dx.types.Handle, i32, i32, i32, i32) #0 +declare %dx.types.LinAlgMatrixC4M4N4U1S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC4M4N4U1S2(i32, %dx.types.Handle, i32, i32, i32, i32) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC4M4N4U2S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC4M4N4U2S2(i32, %dx.types.Handle, i32, i32, i32, i32) #0 ; Function Attrs: nounwind declare %dx.types.LinAlgMatrixC4M4N5U2S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC4M4N5U2S2(i32, %dx.types.Handle, i32, i32, i32, i32) #0 @@ -137,7 +142,7 @@ declare %dx.types.LinAlgMatrixC4M4N5U2S2 @dx.op.linAlgMatrixLoadFromDescriptor.m declare %dx.types.LinAlgMatrixC4M4N4U0S0 @dx.op.linAlgMatrixLoadFromDescriptor.mC4M4N4U0S0(i32, %dx.types.Handle, i32, i32, i32, i32) #0 ; Function Attrs: nounwind -declare %dx.types.LinAlgMatrixC4M5N4U0S2 @dx.op.linAlgMatrixOuterProduct.mC4M5N4U0S2.v4i32.v4i32(i32, <4 x i32>, <4 x i32>) #0 +declare %dx.types.LinAlgMatrixC4M4N4U0S2 @dx.op.linAlgMatrixOuterProduct.mC4M4N4U0S2.v4i32.v4i32(i32, <4 x i32>, <4 x i32>) #0 ; Function Attrs: nounwind declare i32 @dx.op.linAlgMatrixQueryAccumulatorLayout(i32) #0 @@ -155,31 +160,31 @@ declare <4 x float> @dx.op.linAlgConvert.v4f32.v4i32(i32, <4 x i32>, i32, i32) # declare void @dx.op.linAlgVectorAccumulateToDescriptor.v4f32(i32, %dx.types.Handle, i32, i32, <4 x float>) #0 ; Function Attrs: nounwind -declare %dx.types.LinAlgMatrixC4M4N5U1S2 @dx.op.linAlgCopyConvertMatrix.mC4M4N5U1S2.mC4M5N4U0S2(i32, %dx.types.LinAlgMatrixC4M5N4U0S2, i1) #0 +declare %dx.types.LinAlgMatrixC4M4N4U1S2 @dx.op.linAlgCopyConvertMatrix.mC4M4N4U1S2.mC4M4N4U0S2(i32, %dx.types.LinAlgMatrixC4M4N4U0S2, i1) #0 ; Function Attrs: nounwind -declare %dx.types.LinAlgMatrixC4M5N4U0S2 @dx.op.linAlgFillMatrix.mC4M5N4U0S2.i32(i32, i32) #0 +declare %dx.types.LinAlgMatrixC4M4N4U0S2 @dx.op.linAlgFillMatrix.mC4M4N4U0S2.i32(i32, i32) #0 ; Function Attrs: nounwind -declare <2 x i32> @dx.op.linAlgMatrixGetCoordinate.mC4M5N4U0S2(i32, %dx.types.LinAlgMatrixC4M5N4U0S2, i32) #0 +declare <2 x i32> @dx.op.linAlgMatrixGetCoordinate.mC4M4N4U0S2(i32, %dx.types.LinAlgMatrixC4M4N4U0S2, i32) #0 ; Function Attrs: nounwind -declare float @dx.op.linAlgMatrixGetElement.f32.mC4M5N4U0S2(i32, %dx.types.LinAlgMatrixC4M5N4U0S2, i32) #0 +declare float @dx.op.linAlgMatrixGetElement.f32.mC4M4N4U0S2(i32, %dx.types.LinAlgMatrixC4M4N4U0S2, i32) #0 ; Function Attrs: nounwind -declare %dx.types.LinAlgMatrixC4M5N4U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC4M5N4U2S2.mC4M5N4U0S2.mC4M4N5U1S2.mC4M5N4U2S2(i32, %dx.types.LinAlgMatrixC4M5N4U0S2, %dx.types.LinAlgMatrixC4M4N5U1S2, %dx.types.LinAlgMatrixC4M5N4U2S2) #0 +declare %dx.types.LinAlgMatrixC4M4N4U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC4M4N4U2S2.mC4M4N4U0S2.mC4M4N4U1S2.mC4M4N4U2S2(i32, %dx.types.LinAlgMatrixC4M4N4U0S2, %dx.types.LinAlgMatrixC4M4N4U1S2, %dx.types.LinAlgMatrixC4M4N4U2S2) #0 ; Function Attrs: nounwind -declare %dx.types.LinAlgMatrixC4M5N4U0S2 @dx.op.linAlgMatrixSetElement.mC4M5N4U0S2.mC4M5N4U0S2.i32(i32, %dx.types.LinAlgMatrixC4M5N4U0S2, i32, i32) #0 +declare %dx.types.LinAlgMatrixC4M4N4U0S2 @dx.op.linAlgMatrixSetElement.mC4M4N4U0S2.mC4M4N4U0S2.i32(i32, %dx.types.LinAlgMatrixC4M4N4U0S2, i32, i32) #0 ; Function Attrs: nounwind -declare void @dx.op.linAlgMatrixStoreToMemory.mC4M5N4U0S2.f32(i32, %dx.types.LinAlgMatrixC4M5N4U0S2, float addrspace(3)*, i32, i32, i32) #0 +declare void @dx.op.linAlgMatrixStoreToMemory.mC4M4N4U0S2.f32(i32, %dx.types.LinAlgMatrixC4M4N4U0S2, float addrspace(3)*, i32, i32, i32) #0 ; Function Attrs: nounwind -declare void @dx.op.linAlgMatrixAccumulateToMemory.mC4M5N4U0S2.f32(i32, %dx.types.LinAlgMatrixC4M5N4U0S2, float addrspace(3)*, i32, i32, i32, i32) #0 +declare void @dx.op.linAlgMatrixAccumulateToMemory.mC4M4N4U0S2.f32(i32, %dx.types.LinAlgMatrixC4M4N4U0S2, float addrspace(3)*, i32, i32, i32, i32) #0 ; Function Attrs: nounwind -declare %dx.types.LinAlgMatrixC4M5N4U0S2 @dx.op.linAlgMatrixLoadFromMemory.mC4M5N4U0S2.f32(i32, float addrspace(3)*, i32, i32, i32) #0 +declare %dx.types.LinAlgMatrixC4M4N4U0S2 @dx.op.linAlgMatrixLoadFromMemory.mC4M4N4U0S2.f32(i32, float addrspace(3)*, i32, i32, i32) #0 ; Function Attrs: nounwind readnone declare %dx.types.Handle @dx.op.annotateHandle(i32, %dx.types.Handle, %dx.types.ResourceProperties) #1 @@ -198,9 +203,9 @@ attributes #1 = { nounwind readnone } !dx.resources = !{!6} !dx.entryPoints = !{!9} -!0 = !{%dx.types.LinAlgMatrixC4M5N4U0S2 undef, i32 4, i32 5, i32 4, i32 0, i32 2} -!1 = !{%dx.types.LinAlgMatrixC4M4N5U1S2 undef, i32 4, i32 4, i32 5, i32 1, i32 2} -!2 = !{%dx.types.LinAlgMatrixC4M5N4U2S2 undef, i32 4, i32 5, i32 4, i32 2, i32 2} +!0 = !{%dx.types.LinAlgMatrixC4M4N4U0S2 undef, i32 4, i32 4, i32 4, i32 0, i32 2} +!1 = !{%dx.types.LinAlgMatrixC4M4N4U1S2 undef, i32 4, i32 4, i32 4, i32 1, i32 2} +!2 = !{%dx.types.LinAlgMatrixC4M4N4U2S2 undef, i32 4, i32 4, i32 4, i32 2, i32 2} !3 = !{!"dxc(private) 1.9.0.15241 (main, 1f63535ae)"} !4 = !{i32 1, i32 10} !5 = !{!"cs", i32 6, i32 10} diff --git a/tools/clang/test/LitDXILValidation/LinAlgMatrix/linalgmatrix-matrixmultiply.ll b/tools/clang/test/LitDXILValidation/LinAlgMatrix/linalgmatrix-matrixmultiply.ll new file mode 100644 index 0000000000..4300f84262 --- /dev/null +++ b/tools/clang/test/LitDXILValidation/LinAlgMatrix/linalgmatrix-matrixmultiply.ll @@ -0,0 +1,150 @@ +; REQUIRES: dxil-1-10 +; RUN: not %dxv %s 2>&1 | FileCheck %s + +target datalayout = "e-m:e-p:32:32-i1:32-i8:8-i16:16-i32:32-i64:64-f16:16-f32:32-f64:64-n8:16:32:64" +target triple = "dxil-ms-dx" + +%dx.types.Handle = type { i8* } +%dx.types.ResBind = type { i32, i32, i32, i8 } +%dx.types.ResourceProperties = type { i32, i32 } +%dx.types.LinAlgMatrixC8M8N8U0S2 = type { i8* } +%dx.types.LinAlgMatrixC8M8N8U1S2 = type { i8* } +%dx.types.LinAlgMatrixC8M8N8U0S0 = type { i8* } +%dx.types.LinAlgMatrixC8M8N8U1S0 = type { i8* } +%dx.types.LinAlgMatrixC8M6N8U1S2 = type { i8* } +%dx.types.LinAlgMatrixC8M8N8U2S2 = type { i8* } +%dx.types.LinAlgMatrixC8M8N8U2S0 = type { i8* } +%dx.types.LinAlgMatrixC8M8N6U2S2 = type { i8* } +%struct.ByteAddressBuffer = type { i32 } + +define void @main() { + %1 = call %dx.types.Handle @dx.op.createHandleFromBinding(i32 217, %dx.types.ResBind zeroinitializer, i32 0, i1 false) ; CreateHandleFromBinding(bind,index,nonUniformIndex) + %2 = call %dx.types.Handle @dx.op.annotateHandle(i32 216, %dx.types.Handle %1, %dx.types.ResourceProperties { i32 11, i32 0 }) ; AnnotateHandle(res,props) resource: ByteAddressBuffer + %3 = call %dx.types.LinAlgMatrixC8M8N8U0S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N8U0S2(i32 -2147483634, %dx.types.Handle %2, i32 0, i32 0, i32 0, i32 128) ; LinAlgMatrixLoadFromDescriptor(handle,offset,stride,layout,align) + %4 = call %dx.types.Handle @dx.op.annotateHandle(i32 216, %dx.types.Handle %1, %dx.types.ResourceProperties { i32 11, i32 0 }) ; AnnotateHandle(res,props) resource: ByteAddressBuffer + %5 = call %dx.types.LinAlgMatrixC8M8N8U1S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N8U1S2(i32 -2147483634, %dx.types.Handle %4, i32 0, i32 0, i32 0, i32 128) ; LinAlgMatrixLoadFromDescriptor(handle,offset,stride,layout,align) + %6 = call %dx.types.Handle @dx.op.annotateHandle(i32 216, %dx.types.Handle %1, %dx.types.ResourceProperties { i32 11, i32 0 }) ; AnnotateHandle(res,props) resource: ByteAddressBuffer + %7 = call %dx.types.LinAlgMatrixC8M8N8U0S0 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N8U0S0(i32 -2147483634, %dx.types.Handle %6, i32 0, i32 0, i32 0, i32 128) ; LinAlgMatrixLoadFromDescriptor(handle,offset,stride,layout,align) + %8 = call %dx.types.Handle @dx.op.annotateHandle(i32 216, %dx.types.Handle %1, %dx.types.ResourceProperties { i32 11, i32 0 }) ; AnnotateHandle(res,props) resource: ByteAddressBuffer + %9 = call %dx.types.LinAlgMatrixC8M8N8U1S0 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N8U1S0(i32 -2147483634, %dx.types.Handle %8, i32 0, i32 0, i32 0, i32 128) ; LinAlgMatrixLoadFromDescriptor(handle,offset,stride,layout,align) + %10 = call %dx.types.Handle @dx.op.annotateHandle(i32 216, %dx.types.Handle %1, %dx.types.ResourceProperties { i32 11, i32 0 }) ; AnnotateHandle(res,props) resource: ByteAddressBuffer + %11 = call %dx.types.LinAlgMatrixC8M6N8U1S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M6N8U1S2(i32 -2147483634, %dx.types.Handle %10, i32 0, i32 0, i32 0, i32 128) ; LinAlgMatrixLoadFromDescriptor(handle,offset,stride,layout,align) + + ; CHECK: Function: main: error: A matrix use 'B' does not match expected use A. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiply.mC8M8N8U2S2.mC8M8N8U1S2.mC8M8N8U1S2 + %12 = call %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiply.mC8M8N8U2S2.mC8M8N8U1S2.mC8M8N8U1S2(i32 -2147483625, %dx.types.LinAlgMatrixC8M8N8U1S2 %5, %dx.types.LinAlgMatrixC8M8N8U1S2 %5) ; LinAlgMatrixMultiply(matrixA,matrixB) + + ; CHECK-NEXT: Function: main: error: B matrix use 'Accumulator' does not match expected use B. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiply.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U2S2 + %13 = call %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiply.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U2S2(i32 -2147483625, %dx.types.LinAlgMatrixC8M8N8U0S2 %3, %dx.types.LinAlgMatrixC8M8N8U2S2 %12) ; LinAlgMatrixMultiply(matrixA,matrixB) + + ; CHECK-NEXT: Function: main: error: Return matrix use 'A' does not match expected use Accumulator. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiply.mC8M8N8U0S2.mC8M8N8U0S2.mC8M8N8U1S2 + %14 = call %dx.types.LinAlgMatrixC8M8N8U0S2 @dx.op.linAlgMatrixMultiply.mC8M8N8U0S2.mC8M8N8U0S2.mC8M8N8U1S2(i32 -2147483625, %dx.types.LinAlgMatrixC8M8N8U0S2 %3, %dx.types.LinAlgMatrixC8M8N8U1S2 %5) ; LinAlgMatrixMultiply(matrixA,matrixB) + + ; CHECK-NEXT: Function: main: error: A matrix scope 'Thread' does not match expected scope Wave or ThreadGroup. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiply.mC8M8N8U2S2.mC8M8N8U0S0.mC8M8N8U1S2 + ; CHECK-NEXT: Function: main: error: Matrix scope must be the same for all matrices. A 'Thread', B 'ThreadGroup', Return 'ThreadGroup'. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiply.mC8M8N8U2S2.mC8M8N8U0S0.mC8M8N8U1S2 + %15 = call %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiply.mC8M8N8U2S2.mC8M8N8U0S0.mC8M8N8U1S2(i32 -2147483625, %dx.types.LinAlgMatrixC8M8N8U0S0 %7, %dx.types.LinAlgMatrixC8M8N8U1S2 %5) ; LinAlgMatrixMultiply(matrixA,matrixB) + + ; CHECK-NEXT: Function: main: error: B matrix scope 'Thread' does not match expected scope Wave or ThreadGroup. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiply.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U1S0 + ; CHECK-NEXT: Function: main: error: Matrix scope must be the same for all matrices. A 'ThreadGroup', B 'Thread', Return 'ThreadGroup'. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiply.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U1S0 + %16 = call %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiply.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U1S0(i32 -2147483625, %dx.types.LinAlgMatrixC8M8N8U0S2 %14, %dx.types.LinAlgMatrixC8M8N8U1S0 %9) ; LinAlgMatrixMultiply(matrixA,matrixB) + + ; CHECK-NEXT: Function: main: error: A matrix scope 'Thread' does not match expected scope Wave or ThreadGroup. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiply.mC8M8N8U2S0.mC8M8N8U0S0.mC8M8N8U1S0 + ; CHECK-NEXT: Function: main: error: B matrix scope 'Thread' does not match expected scope Wave or ThreadGroup. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiply.mC8M8N8U2S0.mC8M8N8U0S0.mC8M8N8U1S0 + ; CHECK-NEXT: Function: main: error: Return matrix scope 'Thread' does not match expected scope Wave or ThreadGroup. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiply.mC8M8N8U2S0.mC8M8N8U0S0.mC8M8N8U1S0 + %17 = call %dx.types.LinAlgMatrixC8M8N8U2S0 @dx.op.linAlgMatrixMultiply.mC8M8N8U2S0.mC8M8N8U0S0.mC8M8N8U1S0(i32 -2147483625, %dx.types.LinAlgMatrixC8M8N8U0S0 %7, %dx.types.LinAlgMatrixC8M8N8U1S0 %9) ; LinAlgMatrixMultiply(matrixA,matrixB) + + ; CHECK-NEXT: Function: main: error: K dim of A matrix '8x8' must match K dim of B matrix '6x8'. 8 != 6. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiply.mC8M8N8U2S2.mC8M8N8U0S2.mC8M6N8U1S2 + %18 = call %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiply.mC8M8N8U2S2.mC8M8N8U0S2.mC8M6N8U1S2(i32 -2147483625, %dx.types.LinAlgMatrixC8M8N8U0S2 %14, %dx.types.LinAlgMatrixC8M6N8U1S2 %11) ; LinAlgMatrixMultiply(matrixA,matrixB) + + ; CHECK-NEXT: Function: main: error: Return matrix dimension '8x6' must match A.MxB.N '8x8'. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiply.mC8M8N6U2S2.mC8M8N8U0S2.mC8M8N8U1S2 + %19 = call %dx.types.LinAlgMatrixC8M8N6U2S2 @dx.op.linAlgMatrixMultiply.mC8M8N6U2S2.mC8M8N8U0S2.mC8M8N8U1S2(i32 -2147483625, %dx.types.LinAlgMatrixC8M8N8U0S2 %14, %dx.types.LinAlgMatrixC8M8N8U1S2 %5) ; LinAlgMatrixMultiply(matrixA,matrixB) + + ; CHECK-NEXT: Validation failed. + ret void +} + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U0S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N8U0S2(i32, %dx.types.Handle, i32, i32, i32, i32) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U1S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N8U1S2(i32, %dx.types.Handle, i32, i32, i32, i32) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U0S0 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N8U0S0(i32, %dx.types.Handle, i32, i32, i32, i32) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U1S0 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N8U1S0(i32, %dx.types.Handle, i32, i32, i32, i32) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M6N8U1S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M6N8U1S2(i32, %dx.types.Handle, i32, i32, i32, i32) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiply.mC8M8N8U2S2.mC8M8N8U1S2.mC8M8N8U1S2(i32, %dx.types.LinAlgMatrixC8M8N8U1S2, %dx.types.LinAlgMatrixC8M8N8U1S2) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiply.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U2S2(i32, %dx.types.LinAlgMatrixC8M8N8U0S2, %dx.types.LinAlgMatrixC8M8N8U2S2) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U0S2 @dx.op.linAlgMatrixMultiply.mC8M8N8U0S2.mC8M8N8U0S2.mC8M8N8U1S2(i32, %dx.types.LinAlgMatrixC8M8N8U0S2, %dx.types.LinAlgMatrixC8M8N8U1S2) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiply.mC8M8N8U2S2.mC8M8N8U0S0.mC8M8N8U1S2(i32, %dx.types.LinAlgMatrixC8M8N8U0S0, %dx.types.LinAlgMatrixC8M8N8U1S2) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiply.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U1S0(i32, %dx.types.LinAlgMatrixC8M8N8U0S2, %dx.types.LinAlgMatrixC8M8N8U1S0) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U2S0 @dx.op.linAlgMatrixMultiply.mC8M8N8U2S0.mC8M8N8U0S0.mC8M8N8U1S0(i32, %dx.types.LinAlgMatrixC8M8N8U0S0, %dx.types.LinAlgMatrixC8M8N8U1S0) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiply.mC8M8N8U2S2.mC8M8N8U0S2.mC8M6N8U1S2(i32, %dx.types.LinAlgMatrixC8M8N8U0S2, %dx.types.LinAlgMatrixC8M6N8U1S2) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N6U2S2 @dx.op.linAlgMatrixMultiply.mC8M8N6U2S2.mC8M8N8U0S2.mC8M8N8U1S2(i32, %dx.types.LinAlgMatrixC8M8N8U0S2, %dx.types.LinAlgMatrixC8M8N8U1S2) #0 + +; Function Attrs: nounwind readnone +declare %dx.types.Handle @dx.op.annotateHandle(i32, %dx.types.Handle, %dx.types.ResourceProperties) #1 + +; Function Attrs: nounwind readnone +declare %dx.types.Handle @dx.op.createHandleFromBinding(i32, %dx.types.ResBind, i32, i1) #1 + +attributes #0 = { nounwind } +attributes #1 = { nounwind readnone } + +!dx.targetTypes = !{!0, !1, !2, !3, !4, !5, !6, !7} +!llvm.ident = !{!8} +!dx.version = !{!9} +!dx.valver = !{!9} +!dx.shaderModel = !{!10} +!dx.resources = !{!11} +!dx.entryPoints = !{!14} + +!0 = !{%dx.types.LinAlgMatrixC8M8N8U0S2 undef, i32 8, i32 8, i32 8, i32 0, i32 2} +!1 = !{%dx.types.LinAlgMatrixC8M8N8U1S2 undef, i32 8, i32 8, i32 8, i32 1, i32 2} +!2 = !{%dx.types.LinAlgMatrixC8M8N8U2S2 undef, i32 8, i32 8, i32 8, i32 2, i32 2} +!3 = !{%dx.types.LinAlgMatrixC8M6N8U1S2 undef, i32 8, i32 6, i32 8, i32 1, i32 2} +!4 = !{%dx.types.LinAlgMatrixC8M8N6U2S2 undef, i32 8, i32 8, i32 6, i32 2, i32 2} +!5 = !{%dx.types.LinAlgMatrixC8M8N8U0S0 undef, i32 8, i32 8, i32 8, i32 0, i32 0} +!6 = !{%dx.types.LinAlgMatrixC8M8N8U1S0 undef, i32 8, i32 8, i32 8, i32 1, i32 0} +!7 = !{%dx.types.LinAlgMatrixC8M8N8U2S0 undef, i32 8, i32 8, i32 8, i32 2, i32 0} +!8 = !{!"dxc(private) 1.9.0.5450 (linalg-vali-refactor, 8c87acfc6-dirty)"} +!9 = !{i32 1, i32 10} +!10 = !{!"cs", i32 6, i32 10} +!11 = !{!12, null, null, null} +!12 = !{!13} +!13 = !{i32 0, %struct.ByteAddressBuffer* undef, !"", i32 0, i32 0, i32 1, i32 11, i32 0, null} +!14 = !{void ()* @main, !"main", null, !11, !15} +!15 = !{i32 0, i64 8388624, i32 4, !16} +!16 = !{i32 1, i32 1, i32 1} + diff --git a/tools/clang/test/LitDXILValidation/LinAlgMatrix/linalgmatrix-matrixmultiplyaccumulate.ll b/tools/clang/test/LitDXILValidation/LinAlgMatrix/linalgmatrix-matrixmultiplyaccumulate.ll new file mode 100644 index 0000000000..bb34e8273a --- /dev/null +++ b/tools/clang/test/LitDXILValidation/LinAlgMatrix/linalgmatrix-matrixmultiplyaccumulate.ll @@ -0,0 +1,196 @@ +; REQUIRES: dxil-1-10 +; RUN: not %dxv %s 2>&1 | FileCheck %s + +target datalayout = "e-m:e-p:32:32-i1:32-i8:8-i16:16-i32:32-i64:64-f16:16-f32:32-f64:64-n8:16:32:64" +target triple = "dxil-ms-dx" + +%dx.types.Handle = type { i8* } +%dx.types.ResBind = type { i32, i32, i32, i8 } +%dx.types.ResourceProperties = type { i32, i32 } +%dx.types.LinAlgMatrixC8M8N8U0S2 = type { i8* } +%dx.types.LinAlgMatrixC8M8N8U1S2 = type { i8* } +%dx.types.LinAlgMatrixC8M8N8U2S2 = type { i8* } +%dx.types.LinAlgMatrixC8M8N8U0S0 = type { i8* } +%dx.types.LinAlgMatrixC8M8N8U1S0 = type { i8* } +%dx.types.LinAlgMatrixC8M8N8U2S0 = type { i8* } +%dx.types.LinAlgMatrixC8M8N6U0S2 = type { i8* } +%dx.types.LinAlgMatrixC8M8N6U2S2 = type { i8* } +%struct.ByteAddressBuffer = type { i32 } + +define void @main() { + %1 = call %dx.types.Handle @dx.op.createHandleFromBinding(i32 217, %dx.types.ResBind zeroinitializer, i32 0, i1 false) ; CreateHandleFromBinding(bind,index,nonUniformIndex) + %2 = call %dx.types.Handle @dx.op.annotateHandle(i32 216, %dx.types.Handle %1, %dx.types.ResourceProperties { i32 11, i32 0 }) ; AnnotateHandle(res,props) resource: ByteAddressBuffer + %3 = call %dx.types.LinAlgMatrixC8M8N8U0S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N8U0S2(i32 -2147483634, %dx.types.Handle %2, i32 0, i32 0, i32 0, i32 128) ; LinAlgMatrixLoadFromDescriptor(handle,offset,stride,layout,align) + %4 = call %dx.types.Handle @dx.op.annotateHandle(i32 216, %dx.types.Handle %1, %dx.types.ResourceProperties { i32 11, i32 0 }) ; AnnotateHandle(res,props) resource: ByteAddressBuffer + %5 = call %dx.types.LinAlgMatrixC8M8N8U1S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N8U1S2(i32 -2147483634, %dx.types.Handle %4, i32 0, i32 0, i32 0, i32 128) ; LinAlgMatrixLoadFromDescriptor(handle,offset,stride,layout,align) + %6 = call %dx.types.Handle @dx.op.annotateHandle(i32 216, %dx.types.Handle %1, %dx.types.ResourceProperties { i32 11, i32 0 }) ; AnnotateHandle(res,props) resource: ByteAddressBuffer + %7 = call %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N8U2S2(i32 -2147483634, %dx.types.Handle %6, i32 0, i32 0, i32 0, i32 128) ; LinAlgMatrixLoadFromDescriptor(handle,offset,stride,layout,align) + %8 = call %dx.types.Handle @dx.op.annotateHandle(i32 216, %dx.types.Handle %1, %dx.types.ResourceProperties { i32 11, i32 0 }) ; AnnotateHandle(res,props) resource: ByteAddressBuffer + %9 = call %dx.types.LinAlgMatrixC8M8N8U0S0 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N8U0S0(i32 -2147483634, %dx.types.Handle %8, i32 0, i32 0, i32 0, i32 128) ; LinAlgMatrixLoadFromDescriptor(handle,offset,stride,layout,align) + %10 = call %dx.types.Handle @dx.op.annotateHandle(i32 216, %dx.types.Handle %1, %dx.types.ResourceProperties { i32 11, i32 0 }) ; AnnotateHandle(res,props) resource: ByteAddressBuffer + %11 = call %dx.types.LinAlgMatrixC8M8N8U1S0 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N8U1S0(i32 -2147483634, %dx.types.Handle %10, i32 0, i32 0, i32 0, i32 128) ; LinAlgMatrixLoadFromDescriptor(handle,offset,stride,layout,align) + %12 = call %dx.types.Handle @dx.op.annotateHandle(i32 216, %dx.types.Handle %1, %dx.types.ResourceProperties { i32 11, i32 0 }) ; AnnotateHandle(res,props) resource: ByteAddressBuffer + %13 = call %dx.types.LinAlgMatrixC8M8N8U2S0 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N8U2S0(i32 -2147483634, %dx.types.Handle %12, i32 0, i32 0, i32 0, i32 128) ; LinAlgMatrixLoadFromDescriptor(handle,offset,stride,layout,align) + %14 = call %dx.types.Handle @dx.op.annotateHandle(i32 216, %dx.types.Handle %1, %dx.types.ResourceProperties { i32 11, i32 0 }) ; AnnotateHandle(res,props) resource: ByteAddressBuffer + %15 = call %dx.types.LinAlgMatrixC8M8N6U0S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N6U0S2(i32 -2147483634, %dx.types.Handle %14, i32 0, i32 0, i32 0, i32 128) ; LinAlgMatrixLoadFromDescriptor(handle,offset,stride,layout,align) + %16 = call %dx.types.Handle @dx.op.annotateHandle(i32 216, %dx.types.Handle %1, %dx.types.ResourceProperties { i32 11, i32 0 }) ; AnnotateHandle(res,props) resource: ByteAddressBuffer + %17 = call %dx.types.LinAlgMatrixC8M8N6U2S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N6U2S2(i32 -2147483634, %dx.types.Handle %16, i32 0, i32 0, i32 0, i32 128) ; LinAlgMatrixLoadFromDescriptor(handle,offset,stride,layout,align) + + ; Okay + %18 = call %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U1S2.mC8M8N8U2S2(i32 -2147483637, %dx.types.LinAlgMatrixC8M8N8U0S2 %3, %dx.types.LinAlgMatrixC8M8N8U1S2 %5, %dx.types.LinAlgMatrixC8M8N8U2S2 %7) ; LinAlgMatrixMultiplyAccumulate(matrixA,matrixB,matrixC) + + ; CHECK: Function: main: error: A matrix use 'B' does not match expected use A. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U1S2.mC8M8N8U1S2.mC8M8N8U2S2 + %19 = call %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U1S2.mC8M8N8U1S2.mC8M8N8U2S2(i32 -2147483637, %dx.types.LinAlgMatrixC8M8N8U1S2 %5, %dx.types.LinAlgMatrixC8M8N8U1S2 %5, %dx.types.LinAlgMatrixC8M8N8U2S2 %7) ; LinAlgMatrixMultiplyAccumulate(matrixA,matrixB,matrixC) + + ; CHECK-NEXT: Function: main: error: B matrix use 'A' does not match expected use B. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U0S2.mC8M8N8U2S2 + %20 = call %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U0S2.mC8M8N8U2S2(i32 -2147483637, %dx.types.LinAlgMatrixC8M8N8U0S2 %3, %dx.types.LinAlgMatrixC8M8N8U0S2 %3, %dx.types.LinAlgMatrixC8M8N8U2S2 %7) ; LinAlgMatrixMultiplyAccumulate(matrixA,matrixB,matrixC) + + ; CHECK-NEXT: Function: main: error: C matrix use 'A' does not match expected use Accumulator. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U1S2.mC8M8N8U0S2 + %21 = call %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U1S2.mC8M8N8U0S2(i32 -2147483637, %dx.types.LinAlgMatrixC8M8N8U0S2 %3, %dx.types.LinAlgMatrixC8M8N8U1S2 %5, %dx.types.LinAlgMatrixC8M8N8U0S2 %3) ; LinAlgMatrixMultiplyAccumulate(matrixA,matrixB,matrixC) + + ; CHECK-NEXT: Function: main: error: Return matrix use 'B' does not match expected use Accumulator. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U1S2.mC8M8N8U0S2.mC8M8N8U1S2.mC8M8N8U2S2 + %22 = call %dx.types.LinAlgMatrixC8M8N8U1S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U1S2.mC8M8N8U0S2.mC8M8N8U1S2.mC8M8N8U2S2(i32 -2147483637, %dx.types.LinAlgMatrixC8M8N8U0S2 %3, %dx.types.LinAlgMatrixC8M8N8U1S2 %5, %dx.types.LinAlgMatrixC8M8N8U2S2 %7) ; LinAlgMatrixMultiplyAccumulate(matrixA,matrixB,matrixC) + + ; CHECK-NEXT: Function: main: error: A matrix scope 'Thread' does not match expected scope Wave or ThreadGroup. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S0.mC8M8N8U1S2.mC8M8N8U2S2 + ; CHECK-NEXT: Function: main: error: Matrix scope must be the same for all matrices. A 'Thread', B 'ThreadGroup', C 'ThreadGroup', Return 'ThreadGroup'. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S0.mC8M8N8U1S2.mC8M8N8U2S2 + %23 = call %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S0.mC8M8N8U1S2.mC8M8N8U2S2(i32 -2147483637, %dx.types.LinAlgMatrixC8M8N8U0S0 %9, %dx.types.LinAlgMatrixC8M8N8U1S2 %22, %dx.types.LinAlgMatrixC8M8N8U2S2 %7) ; LinAlgMatrixMultiplyAccumulate(matrixA,matrixB,matrixC) + + ; CHECK-NEXT: Function: main: error: B matrix scope 'Thread' does not match expected scope Wave or ThreadGroup. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U1S0.mC8M8N8U2S2 + ; CHECK-NEXT: Function: main: error: Matrix scope must be the same for all matrices. A 'ThreadGroup', B 'Thread', C 'ThreadGroup', Return 'ThreadGroup'. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U1S0.mC8M8N8U2S2 + %24 = call %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U1S0.mC8M8N8U2S2(i32 -2147483637, %dx.types.LinAlgMatrixC8M8N8U0S2 %3, %dx.types.LinAlgMatrixC8M8N8U1S0 %11, %dx.types.LinAlgMatrixC8M8N8U2S2 %7) ; LinAlgMatrixMultiplyAccumulate(matrixA,matrixB,matrixC) + + ; CHECK-NEXT: Function: main: error: C matrix scope 'Thread' does not match expected scope Wave or ThreadGroup. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U1S2.mC8M8N8U2S0 + ; CHECK-NEXT: Function: main: error: Matrix scope must be the same for all matrices. A 'ThreadGroup', B 'ThreadGroup', C 'Thread', Return 'ThreadGroup'. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U1S2.mC8M8N8U2S0 + %25 = call %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U1S2.mC8M8N8U2S0(i32 -2147483637, %dx.types.LinAlgMatrixC8M8N8U0S2 %3, %dx.types.LinAlgMatrixC8M8N8U1S2 %22, %dx.types.LinAlgMatrixC8M8N8U2S0 %13) ; LinAlgMatrixMultiplyAccumulate(matrixA,matrixB,matrixC) + + ; CHECK-NEXT: Function: main: error: A matrix scope 'Thread' does not match expected scope Wave or ThreadGroup. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S0.mC8M8N8U0S0.mC8M8N8U1S0.mC8M8N8U2S0 + ; CHECK-NEXT: Function: main: error: B matrix scope 'Thread' does not match expected scope Wave or ThreadGroup. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S0.mC8M8N8U0S0.mC8M8N8U1S0.mC8M8N8U2S0 + ; CHECK-NEXT: Function: main: error: C matrix scope 'Thread' does not match expected scope Wave or ThreadGroup. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S0.mC8M8N8U0S0.mC8M8N8U1S0.mC8M8N8U2S0 + ; CHECK-NEXT: Function: main: error: Return matrix scope 'Thread' does not match expected scope Wave or ThreadGroup. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S0.mC8M8N8U0S0.mC8M8N8U1S0.mC8M8N8U2S0 + %26 = call %dx.types.LinAlgMatrixC8M8N8U2S0 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S0.mC8M8N8U0S0.mC8M8N8U1S0.mC8M8N8U2S0(i32 -2147483637, %dx.types.LinAlgMatrixC8M8N8U0S0 %9, %dx.types.LinAlgMatrixC8M8N8U1S0 %11, %dx.types.LinAlgMatrixC8M8N8U2S0 %13) ; LinAlgMatrixMultiplyAccumulate(matrixA,matrixB,matrixC) + + ; CHECK-NEXT: Function: main: error: K dim of A matrix '8x6' must match K dim of B matrix '8x8'. 6 != 8. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N6U0S2.mC8M8N8U1S2.mC8M8N8U2S2 + %27 = call %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N6U0S2.mC8M8N8U1S2.mC8M8N8U2S2(i32 -2147483637, %dx.types.LinAlgMatrixC8M8N6U0S2 %15, %dx.types.LinAlgMatrixC8M8N8U1S2 %22, %dx.types.LinAlgMatrixC8M8N8U2S2 %7) ; LinAlgMatrixMultiplyAccumulate(matrixA,matrixB,matrixC) + + ; CHECK-NEXT: Function: main: error: C matrix dimension '8x6' must match A.MxB.N '8x8'. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U1S2.mC8M8N6U2S2 + %28 = call %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U1S2.mC8M8N6U2S2(i32 -2147483637, %dx.types.LinAlgMatrixC8M8N8U0S2 %3, %dx.types.LinAlgMatrixC8M8N8U1S2 %22, %dx.types.LinAlgMatrixC8M8N6U2S2 %17) ; LinAlgMatrixMultiplyAccumulate(matrixA,matrixB,matrixC) + + ; CHECK-NEXT: Function: main: error: Return matrix dimension '8x6' must match A.MxB.N '8x8'. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N6U2S2.mC8M8N8U0S2.mC8M8N8U1S2.mC8M8N8U2S2 + %29 = call %dx.types.LinAlgMatrixC8M8N6U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N6U2S2.mC8M8N8U0S2.mC8M8N8U1S2.mC8M8N8U2S2(i32 -2147483637, %dx.types.LinAlgMatrixC8M8N8U0S2 %3, %dx.types.LinAlgMatrixC8M8N8U1S2 %22, %dx.types.LinAlgMatrixC8M8N8U2S2 %7) ; LinAlgMatrixMultiplyAccumulate(matrixA,matrixB,matrixC) + + ; CHECK-NEXT: Validation failed. + ret void +} + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U0S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N8U0S2(i32, %dx.types.Handle, i32, i32, i32, i32) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U1S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N8U1S2(i32, %dx.types.Handle, i32, i32, i32, i32) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N8U2S2(i32, %dx.types.Handle, i32, i32, i32, i32) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U0S0 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N8U0S0(i32, %dx.types.Handle, i32, i32, i32, i32) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U1S0 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N8U1S0(i32, %dx.types.Handle, i32, i32, i32, i32) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U2S0 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N8U2S0(i32, %dx.types.Handle, i32, i32, i32, i32) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N6U0S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N6U0S2(i32, %dx.types.Handle, i32, i32, i32, i32) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N6U2S2 @dx.op.linAlgMatrixLoadFromDescriptor.mC8M8N6U2S2(i32, %dx.types.Handle, i32, i32, i32, i32) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U1S2.mC8M8N8U2S2(i32, %dx.types.LinAlgMatrixC8M8N8U0S2, %dx.types.LinAlgMatrixC8M8N8U1S2, %dx.types.LinAlgMatrixC8M8N8U2S2) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U1S2.mC8M8N8U1S2.mC8M8N8U2S2(i32, %dx.types.LinAlgMatrixC8M8N8U1S2, %dx.types.LinAlgMatrixC8M8N8U1S2, %dx.types.LinAlgMatrixC8M8N8U2S2) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U0S2.mC8M8N8U2S2(i32, %dx.types.LinAlgMatrixC8M8N8U0S2, %dx.types.LinAlgMatrixC8M8N8U0S2, %dx.types.LinAlgMatrixC8M8N8U2S2) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U1S2.mC8M8N8U0S2(i32, %dx.types.LinAlgMatrixC8M8N8U0S2, %dx.types.LinAlgMatrixC8M8N8U1S2, %dx.types.LinAlgMatrixC8M8N8U0S2) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U1S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U1S2.mC8M8N8U0S2.mC8M8N8U1S2.mC8M8N8U2S2(i32, %dx.types.LinAlgMatrixC8M8N8U0S2, %dx.types.LinAlgMatrixC8M8N8U1S2, %dx.types.LinAlgMatrixC8M8N8U2S2) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S0.mC8M8N8U1S2.mC8M8N8U2S2(i32, %dx.types.LinAlgMatrixC8M8N8U0S0, %dx.types.LinAlgMatrixC8M8N8U1S2, %dx.types.LinAlgMatrixC8M8N8U2S2) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U1S0.mC8M8N8U2S2(i32, %dx.types.LinAlgMatrixC8M8N8U0S2, %dx.types.LinAlgMatrixC8M8N8U1S0, %dx.types.LinAlgMatrixC8M8N8U2S2) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U1S2.mC8M8N8U2S0(i32, %dx.types.LinAlgMatrixC8M8N8U0S2, %dx.types.LinAlgMatrixC8M8N8U1S2, %dx.types.LinAlgMatrixC8M8N8U2S0) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U2S0 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S0.mC8M8N8U0S0.mC8M8N8U1S0.mC8M8N8U2S0(i32, %dx.types.LinAlgMatrixC8M8N8U0S0, %dx.types.LinAlgMatrixC8M8N8U1S0, %dx.types.LinAlgMatrixC8M8N8U2S0) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N6U0S2.mC8M8N8U1S2.mC8M8N8U2S2(i32, %dx.types.LinAlgMatrixC8M8N6U0S2, %dx.types.LinAlgMatrixC8M8N8U1S2, %dx.types.LinAlgMatrixC8M8N8U2S2) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N8U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N8U2S2.mC8M8N8U0S2.mC8M8N8U1S2.mC8M8N6U2S2(i32, %dx.types.LinAlgMatrixC8M8N8U0S2, %dx.types.LinAlgMatrixC8M8N8U1S2, %dx.types.LinAlgMatrixC8M8N6U2S2) #0 + +; Function Attrs: nounwind +declare %dx.types.LinAlgMatrixC8M8N6U2S2 @dx.op.linAlgMatrixMultiplyAccumulate.mC8M8N6U2S2.mC8M8N8U0S2.mC8M8N8U1S2.mC8M8N8U2S2(i32, %dx.types.LinAlgMatrixC8M8N8U0S2, %dx.types.LinAlgMatrixC8M8N8U1S2, %dx.types.LinAlgMatrixC8M8N8U2S2) #0 + +; Function Attrs: nounwind readnone +declare %dx.types.Handle @dx.op.annotateHandle(i32, %dx.types.Handle, %dx.types.ResourceProperties) #1 + +; Function Attrs: nounwind readnone +declare %dx.types.Handle @dx.op.createHandleFromBinding(i32, %dx.types.ResBind, i32, i1) #1 + +attributes #0 = { nounwind } +attributes #1 = { nounwind readnone } + +!dx.targetTypes = !{!0, !1, !2, !3, !4, !5, !6, !7} +!llvm.ident = !{!8} +!dx.version = !{!9} +!dx.valver = !{!9} +!dx.shaderModel = !{!10} +!dx.resources = !{!11} +!dx.entryPoints = !{!14} + +!0 = !{%dx.types.LinAlgMatrixC8M8N8U0S2 undef, i32 8, i32 8, i32 8, i32 0, i32 2} +!1 = !{%dx.types.LinAlgMatrixC8M8N8U1S2 undef, i32 8, i32 8, i32 8, i32 1, i32 2} +!2 = !{%dx.types.LinAlgMatrixC8M8N8U2S2 undef, i32 8, i32 8, i32 8, i32 2, i32 2} +!3 = !{%dx.types.LinAlgMatrixC8M8N8U0S0 undef, i32 8, i32 8, i32 8, i32 0, i32 0} +!4 = !{%dx.types.LinAlgMatrixC8M8N8U1S0 undef, i32 8, i32 8, i32 8, i32 1, i32 0} +!5 = !{%dx.types.LinAlgMatrixC8M8N8U2S0 undef, i32 8, i32 8, i32 8, i32 2, i32 0} +!6 = !{%dx.types.LinAlgMatrixC8M8N6U0S2 undef, i32 8, i32 8, i32 6, i32 0, i32 2} +!7 = !{%dx.types.LinAlgMatrixC8M8N6U2S2 undef, i32 8, i32 8, i32 6, i32 2, i32 2} +!8 = !{!"dxc(private) 1.9.0.5450 (linalg-vali-refactor, 8c87acfc6-dirty)"} +!9 = !{i32 1, i32 10} +!10 = !{!"cs", i32 6, i32 10} +!11 = !{!12, null, null, null} +!12 = !{!13} +!13 = !{i32 0, %struct.ByteAddressBuffer* undef, !"", i32 0, i32 0, i32 1, i32 11, i32 0, null} +!14 = !{void ()* @main, !"main", null, !11, !15} +!15 = !{i32 0, i64 8388624, i32 4, !16} +!16 = !{i32 1, i32 1, i32 1} + diff --git a/utils/hct/hctdb.py b/utils/hct/hctdb.py index 90bc1e5de4..2597402dea 100644 --- a/utils/hct/hctdb.py +++ b/utils/hct/hctdb.py @@ -8713,6 +8713,22 @@ def build_valrules(self): "Instr.LinAlgMatrix2PartsMustMatch", "%0 matrix %1 '%2' must match %3 matrix %4 '%5'.", ) + self.add_valrule( + "Instr.LinAlgMatrixScopeMustMatch3", + "Matrix scope must be the same for all matrices. %0 '%1', %2 '%3', %4 '%5'.", + ) + self.add_valrule( + "Instr.LinAlgMatrixScopeMustMatch4", + "Matrix scope must be the same for all matrices. %0 '%1', %2 '%3', %4 '%5', %6 '%7'.", + ) + self.add_valrule( + "Instr.LinAlgMatrixMatrixKDimMustMatch", + "K dim of A matrix '%0' must match K dim of B matrix '%1'. %2 != %3.", + ) + self.add_valrule( + "Instr.LinAlgMatrixMatrixResDimMustMatch", + "%0 matrix dimension '%1' must match A.MxB.N '%2'.", + ) # Some legacy rules: # - space is only supported for shader targets 5.1 and higher