From 5df21b27092fb8fe8f54122a90a1340b245f18dc Mon Sep 17 00:00:00 2001 From: Ashley Coleman Date: Mon, 17 Aug 2026 10:59:25 -0600 Subject: [PATCH] [SM6.10] LinAlg Validation: MatrixLength --- lib/DxilValidation/DxilValidation.cpp | 11 +++++++++++ .../LinAlgMatrix/linalgmatrix-non-thread-ops.ll | 7 +++++++ 2 files changed, 18 insertions(+) diff --git a/lib/DxilValidation/DxilValidation.cpp b/lib/DxilValidation/DxilValidation.cpp index b12149ca82..7238784f58 100644 --- a/lib/DxilValidation/DxilValidation.cpp +++ b/lib/DxilValidation/DxilValidation.cpp @@ -1124,6 +1124,17 @@ static void ValidateLinAlgOpReturnMatrix(CallInst *CI, static void ValidateLinAlgMatrixLength(CallInst *CI, ValidationContext &ValCtx) { ValidateLinAlgOpParameters(CI, ValCtx); + DxilInst_LinAlgMatrixLength Op(CI); + std::optional Mat = + GetCheckedLATT(Op.get_matrix()->getType(), ValCtx); + if (!Mat) + return; + + if (Mat->Scope != DXIL::MatrixScope::Wave && + Mat->Scope != DXIL::MatrixScope::ThreadGroup) + ValCtx.EmitInstrFormatError( + CI, ValidationRule::InstrLinAlgMatrixScopeMismatch2, + {"Input", MatrixScopeToString(Mat->Scope), "Wave", "ThreadGroup"}); } static void ValidateLinAlgMatrixGetCoordinate(CallInst *CI, diff --git a/tools/clang/test/LitDXILValidation/LinAlgMatrix/linalgmatrix-non-thread-ops.ll b/tools/clang/test/LitDXILValidation/LinAlgMatrix/linalgmatrix-non-thread-ops.ll index 2296d33fd8..2a8aeefb31 100644 --- a/tools/clang/test/LitDXILValidation/LinAlgMatrix/linalgmatrix-non-thread-ops.ll +++ b/tools/clang/test/LitDXILValidation/LinAlgMatrix/linalgmatrix-non-thread-ops.ll @@ -25,6 +25,10 @@ define void @main() { ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixGetCoordinate.mC8M4N4U2S0 %4 = call <2 x i32> @dx.op.linAlgMatrixGetCoordinate.mC8M4N4U2S0(i32 -2147483631, %dx.types.LinAlgMatrixC8M4N4U2S0 %3, i32 0) ; LinAlgMatrixGetCoordinate(matrix,threadLocalIndex) + ; CHECK-NEXT: Function: main: error: Input matrix scope 'Thread' does not match expected scope Wave or ThreadGroup. + ; CHECK-NEXT: note: at {{.*}} @dx.op.linAlgMatrixLength.mC8M4N4U2S0 + %5 = call i32 @dx.op.linAlgMatrixLength.mC8M4N4U2S0(i32 -2147483632, %dx.types.LinAlgMatrixC8M4N4U2S0 %3) ; LinAlgMatrixLength(matrix) + ; CHECK-NEXT: Validation failed. ret void } @@ -41,6 +45,9 @@ declare %dx.types.LinAlgMatrixC8M4N4U2S0 @dx.op.linAlgMatrixSetElement.mC8M4N4U2 ; Function Attrs: nounwind declare <2 x i32> @dx.op.linAlgMatrixGetCoordinate.mC8M4N4U2S0(i32, %dx.types.LinAlgMatrixC8M4N4U2S0, i32) #0 +; Function Attrs: nounwind +declare i32 @dx.op.linAlgMatrixLength.mC8M4N4U2S0(i32, %dx.types.LinAlgMatrixC8M4N4U2S0) #0 + attributes #0 = { nounwind } !dx.targetTypes = !{!0}