[clang] [HLSL][Matrix] Add comparison operator support (PR #216791)
via cfe-commits
cfe-commits at lists.llvm.org
Mon Aug 17 10:45:19 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-clang
Author: Farzon Lotfi (farzonl)
<details>
<summary>Changes</summary>
fixes #<!-- -->216786
This change add matric cmp support similar to that done for vector types.
---
Full diff: https://github.com/llvm/llvm-project/pull/216791.diff
6 Files Affected:
- (modified) clang/include/clang/Sema/Sema.h (+3)
- (modified) clang/lib/AST/Type.cpp (+2)
- (modified) clang/lib/CodeGen/CGExprScalar.cpp (+3)
- (modified) clang/lib/Sema/SemaExpr.cpp (+33)
- (added) clang/test/CodeGenHLSL/BasicFeatures/MatrixComparisonOperators.hlsl (+83)
- (added) clang/test/SemaHLSL/Operators/matrix-comparisons.hlsl (+54)
``````````diff
diff --git a/clang/include/clang/Sema/Sema.h b/clang/include/clang/Sema/Sema.h
index d931e70cb2342..32949ca1d6998 100644
--- a/clang/include/clang/Sema/Sema.h
+++ b/clang/include/clang/Sema/Sema.h
@@ -7886,6 +7886,9 @@ class Sema final : public SemaBase {
QualType CheckSizelessVectorCompareOperands(ExprResult &LHS, ExprResult &RHS,
SourceLocation Loc,
BinaryOperatorKind Opc);
+ QualType CheckMatrixCompareOperands(ExprResult &LHS, ExprResult &RHS,
+ SourceLocation Loc,
+ BinaryOperatorKind Opc);
QualType CheckVectorLogicalOperands(ExprResult &LHS, ExprResult &RHS,
SourceLocation Loc,
BinaryOperatorKind Opc);
diff --git a/clang/lib/AST/Type.cpp b/clang/lib/AST/Type.cpp
index 5069b587c1f8b..8b10b0347329e 100644
--- a/clang/lib/AST/Type.cpp
+++ b/clang/lib/AST/Type.cpp
@@ -2340,6 +2340,8 @@ bool Type::isSignedIntegerOrEnumerationType() const {
bool Type::hasSignedIntegerRepresentation() const {
if (const auto *VT = dyn_cast<VectorType>(CanonicalType))
return VT->getElementType()->isSignedIntegerOrEnumerationType();
+ if (const auto *MT = dyn_cast<MatrixType>(CanonicalType))
+ return MT->getElementType()->isSignedIntegerOrEnumerationType();
if (const auto *BT = dyn_cast<BuiltinType>(CanonicalType)) {
switch (BT->getKind()) {
diff --git a/clang/lib/CodeGen/CGExprScalar.cpp b/clang/lib/CodeGen/CGExprScalar.cpp
index 67014904ffb37..342e3113b1121 100644
--- a/clang/lib/CodeGen/CGExprScalar.cpp
+++ b/clang/lib/CodeGen/CGExprScalar.cpp
@@ -5368,6 +5368,9 @@ Value *ScalarExprEmitter::EmitCompare(const BinaryOperator *E,
if (LHSTy->isVectorType() || LHSTy->isSveVLSBuiltinType())
return Builder.CreateSExt(Result, ConvertType(E->getType()), "sext");
+ if (LHSTy->isMatrixType())
+ return Result;
+
} else {
// Complex Comparison: can only be an equality comparison.
CodeGenFunction::ComplexPairTy LHS, RHS;
diff --git a/clang/lib/Sema/SemaExpr.cpp b/clang/lib/Sema/SemaExpr.cpp
index 0908841dca8bf..021024a6bfd93 100644
--- a/clang/lib/Sema/SemaExpr.cpp
+++ b/clang/lib/Sema/SemaExpr.cpp
@@ -12963,6 +12963,10 @@ QualType Sema::CheckCompareOperands(ExprResult &LHS, ExprResult &RHS,
CheckPtrComparisonWithNullChar(RHS, LHS);
}
+ if (getLangOpts().HLSL && (LHS.get()->getType()->isConstantMatrixType() ||
+ RHS.get()->getType()->isConstantMatrixType()))
+ return CheckMatrixCompareOperands(LHS, RHS, Loc, Opc);
+
// Handle vector comparisons separately.
if (LHS.get()->getType()->isVectorType() ||
RHS.get()->getType()->isVectorType())
@@ -13538,6 +13542,35 @@ QualType Sema::CheckVectorCompareOperands(ExprResult &LHS, ExprResult &RHS,
return GetSignedVectorType(vType);
}
+QualType Sema::CheckMatrixCompareOperands(ExprResult &LHS, ExprResult &RHS,
+ SourceLocation Loc,
+ BinaryOperatorKind Opc) {
+ assert(getLangOpts().HLSL && "matrix comparisons are only supported in HLSL");
+ assert(Opc != BO_Cmp && "three-way comparisons are not supported in HLSL");
+
+ QualType MatrixTy =
+ CheckMatrixElementwiseOperands(LHS, RHS, Loc, /*IsCompAssign=*/false);
+ if (MatrixTy.isNull())
+ return QualType();
+
+ if (!LHS.get()->getType()->isMatrixType()) {
+ LHS = prepareMatrixSplat(MatrixTy, LHS.get());
+ if (LHS.isInvalid())
+ return QualType();
+ LHS = ImpCastExprToType(LHS.get(), MatrixTy, CK_HLSLAggregateSplatCast);
+ }
+ if (!RHS.get()->getType()->isMatrixType()) {
+ RHS = prepareMatrixSplat(MatrixTy, RHS.get());
+ if (RHS.isInvalid())
+ return QualType();
+ RHS = ImpCastExprToType(RHS.get(), MatrixTy, CK_HLSLAggregateSplatCast);
+ }
+
+ const auto *MT = MatrixTy->castAs<ConstantMatrixType>();
+ return Context.getConstantMatrixType(Context.BoolTy, MT->getNumRows(),
+ MT->getNumColumns());
+}
+
QualType Sema::CheckSizelessVectorCompareOperands(ExprResult &LHS,
ExprResult &RHS,
SourceLocation Loc,
diff --git a/clang/test/CodeGenHLSL/BasicFeatures/MatrixComparisonOperators.hlsl b/clang/test/CodeGenHLSL/BasicFeatures/MatrixComparisonOperators.hlsl
new file mode 100644
index 0000000000000..94ea87166d3b0
--- /dev/null
+++ b/clang/test/CodeGenHLSL/BasicFeatures/MatrixComparisonOperators.hlsl
@@ -0,0 +1,83 @@
+// RUN: %clang_cc1 -finclude-default-header -x hlsl -triple dxil-pc-shadermodel6.3-library %s -emit-llvm -disable-llvm-passes -o - | FileCheck %s
+
+// ==================================================================
+// Float Matrix vs Float Matrix
+// ==================================================================
+
+// CHECK-LABEL: define {{.*}}test_float_lt
+// CHECK: fcmp {{.*}}olt <4 x float>
+bool2x2 test_float_lt(float2x2 A, float2x2 B) {
+ return A < B;
+}
+
+// CHECK-LABEL: define {{.*}}test_float_gt
+// CHECK: fcmp {{.*}}ogt <9 x float>
+bool3x3 test_float_gt(float3x3 A, float3x3 B) {
+ return A > B;
+}
+
+// CHECK-LABEL: define {{.*}}test_float_le
+// CHECK: fcmp {{.*}}ole <16 x float>
+bool4x4 test_float_le(float4x4 A, float4x4 B) {
+ return A <= B;
+}
+
+// CHECK-LABEL: define {{.*}}test_float_ge
+// CHECK: fcmp {{.*}}oge <6 x float>
+bool2x3 test_float_ge(float2x3 A, float2x3 B) {
+ return A >= B;
+}
+
+// CHECK-LABEL: define {{.*}}test_float_eq
+// CHECK: fcmp {{.*}}oeq <4 x float>
+bool2x2 test_float_eq(float2x2 A, float2x2 B) {
+ return A == B;
+}
+
+// CHECK-LABEL: define {{.*}}test_float_neq
+// CHECK: fcmp {{.*}}une <4 x float>
+bool2x2 test_float_neq(float2x2 A, float2x2 B) {
+ return A != B;
+}
+
+// ==================================================================
+// Integer Matrix vs Integer Matrix
+// ==================================================================
+
+// CHECK-LABEL: define {{.*}}test_int_lt
+// CHECK: icmp slt <4 x i32>
+bool2x2 test_int_lt(int2x2 A, int2x2 B) {
+ return A < B;
+}
+
+// CHECK-LABEL: define {{.*}}test_int_ge
+// CHECK: icmp sge <4 x i32>
+bool2x2 test_int_ge(int2x2 A, int2x2 B) {
+ return A >= B;
+}
+
+// CHECK-LABEL: define {{.*}}test_int_eq
+// CHECK: icmp eq <4 x i32>
+bool2x2 test_int_eq(int2x2 A, int2x2 B) {
+ return A == B;
+}
+
+// ==================================================================
+// Matrix vs Scalar (Broadcast)
+// ==================================================================
+
+// CHECK-LABEL: define {{.*}}test_scalar_lt
+// CHECK: [[SPLAT:%.*]] = insertelement <4 x float> poison, float {{%.*}}, i64 0
+// CHECK: [[B_MAT:%.*]] = shufflevector <4 x float> [[SPLAT]], <4 x float> poison, <4 x i32> zeroinitializer
+// CHECK: fcmp {{.*}}olt <4 x float> {{%.*}}, [[B_MAT]]
+bool2x2 test_scalar_lt(float2x2 A, float B) {
+ return A < B;
+}
+
+// CHECK-LABEL: define {{.*}}test_scalar_neq
+// CHECK: [[SPLAT:%.*]] = insertelement <4 x float> poison, float {{%.*}}, i64 0
+// CHECK: [[B_MAT:%.*]] = shufflevector <4 x float> [[SPLAT]], <4 x float> poison, <4 x i32> zeroinitializer
+// CHECK: fcmp {{.*}}une <4 x float> {{%.*}}, [[B_MAT]]
+bool2x2 test_scalar_neq(float2x2 A, float B) {
+ return A != B;
+}
diff --git a/clang/test/SemaHLSL/Operators/matrix-comparisons.hlsl b/clang/test/SemaHLSL/Operators/matrix-comparisons.hlsl
new file mode 100644
index 0000000000000..e2e67bce4fd52
--- /dev/null
+++ b/clang/test/SemaHLSL/Operators/matrix-comparisons.hlsl
@@ -0,0 +1,54 @@
+// RUN: %clang_cc1 -finclude-default-header -triple dxil-pc-shadermodel6.3-library -Wno-implicit-int-float-conversion %s -ast-dump -ast-dump-filter=test | FileCheck %s
+// RUN: %clang_cc1 -finclude-default-header -triple dxil-pc-shadermodel6.3-library -Wno-implicit-int-float-conversion %s -DERRORS -verify
+
+// CHECK-LABEL: FunctionDecl {{.*}} test_matrix_matrix 'bool2x2 (float2x2, float2x2)'
+// CHECK: BinaryOperator {{.*}} 'matrix<bool, 2, 2>' '<'
+// CHECK-NEXT: ImplicitCastExpr {{.*}} 'float2x2':'matrix<float, 2, 2>' <LValueToRValue>
+// CHECK-NEXT: DeclRefExpr {{.*}} 'float2x2':'matrix<float, 2, 2>' lvalue ParmVar {{.*}} 'a'
+// CHECK-NEXT: ImplicitCastExpr {{.*}} 'float2x2':'matrix<float, 2, 2>' <LValueToRValue>
+// CHECK-NEXT: DeclRefExpr {{.*}} 'float2x2':'matrix<float, 2, 2>' lvalue ParmVar {{.*}} 'b'
+bool2x2 test_matrix_matrix(float2x2 a, float2x2 b) {
+ return a < b;
+}
+
+// CHECK-LABEL: FunctionDecl {{.*}} test_matrix_scalar 'bool2x2 (float2x2, int)'
+// CHECK: BinaryOperator {{.*}} 'matrix<bool, 2, 2>' '=='
+// CHECK-NEXT: ImplicitCastExpr {{.*}} 'float2x2':'matrix<float, 2, 2>' <LValueToRValue>
+// CHECK-NEXT: DeclRefExpr {{.*}} 'float2x2':'matrix<float, 2, 2>' lvalue ParmVar {{.*}} 'a'
+// CHECK-NEXT: ImplicitCastExpr {{.*}} 'float2x2':'matrix<float, 2, 2>' <HLSLAggregateSplatCast>
+// CHECK-NEXT: ImplicitCastExpr {{.*}} 'float' <IntegralToFloating>
+// CHECK-NEXT: ImplicitCastExpr {{.*}} 'int' <LValueToRValue>
+// CHECK-NEXT: DeclRefExpr {{.*}} 'int' lvalue ParmVar {{.*}} 'b'
+bool2x2 test_matrix_scalar(float2x2 a, int b) {
+ return a == b;
+}
+
+// CHECK-LABEL: FunctionDecl {{.*}} test_scalar_matrix 'bool2x2 (int, float2x2)'
+// CHECK: BinaryOperator {{.*}} 'matrix<bool, 2, 2>' '>='
+// CHECK-NEXT: ImplicitCastExpr {{.*}} 'float2x2':'matrix<float, 2, 2>' <HLSLAggregateSplatCast>
+// CHECK-NEXT: ImplicitCastExpr {{.*}} 'float' <IntegralToFloating>
+// CHECK-NEXT: ImplicitCastExpr {{.*}} 'int' <LValueToRValue>
+// CHECK-NEXT: DeclRefExpr {{.*}} 'int' lvalue ParmVar {{.*}} 'a'
+// CHECK-NEXT: ImplicitCastExpr {{.*}} 'float2x2':'matrix<float, 2, 2>' <LValueToRValue>
+// CHECK-NEXT: DeclRefExpr {{.*}} 'float2x2':'matrix<float, 2, 2>' lvalue ParmVar {{.*}} 'b'
+bool2x2 test_scalar_matrix(int a, float2x2 b) {
+ return a >= b;
+}
+
+#ifdef ERRORS
+
+bool2x2 test_dimension_mismatch(float2x2 a, float3x3 b) {
+ return a != b; // expected-error {{invalid operands to binary expression ('float2x2' (aka 'matrix<float, 2, 2>') and 'float3x3' (aka 'matrix<float, 3, 3>'))}}
+}
+
+bool2x2 test_element_mismatch(float2x2 a, int2x2 b) {
+ return a < b; // expected-error {{invalid operands to binary expression ('float2x2' (aka 'matrix<float, 2, 2>') and 'int2x2' (aka 'matrix<int, 2, 2>'))}}
+}
+
+struct Unsupported {};
+
+bool2x2 test_unsupported_operand(float2x2 a, Unsupported b) {
+ return a > b; // expected-error {{invalid operands to binary expression ('float2x2' (aka 'matrix<float, 2, 2>') and 'Unsupported')}} expected-error {{cannot initialize a value of type 'float' with an rvalue of type 'Unsupported'}}
+}
+
+#endif
``````````
</details>
https://github.com/llvm/llvm-project/pull/216791
More information about the cfe-commits
mailing list