[clang] 3a2f13f - [HLSL] Extract matrix elementwise cast operands from SSA (#227065)

via cfe-commits cfe-commits at lists.llvm.org
Tue Sep 29 16:56:41 PDT 2026


Author: Farzon Lotfi
Date: 2026-09-29T23:56:31Z
New Revision: 3a2f13ff993961e7bd745a37d6f163fe46657150

URL: https://github.com/llvm/llvm-project/commit/3a2f13ff993961e7bd745a37d6f163fe46657150
DIFF: https://github.com/llvm/llvm-project/commit/3a2f13ff993961e7bd745a37d6f163fe46657150.diff

LOG: [HLSL] Extract matrix elementwise cast operands from SSA (#227065)

Avoid materializing matrix operands in the hlsl.ewcast.src temporary
when performing matrix-to-vector elementwise casts. Extract each element
directly from the canonical matrix SSA value using its column-major
flattened index.

Factor vector result construction into a shared helper so aggregate and
matrix sources use the same conversion and insertion logic.

Update the matrix-to-vector CodeGen checks for both row-major and
column-major layouts.

Assisted by Copilot using GPT 5.6 Sol

Added: 
    

Modified: 
    clang/lib/CodeGen/CGExprScalar.cpp
    clang/test/CodeGenHLSL/BasicFeatures/VectorElementwiseCast.hlsl

Removed: 
    


################################################################################
diff  --git a/clang/lib/CodeGen/CGExprScalar.cpp b/clang/lib/CodeGen/CGExprScalar.cpp
index c91ed144be005..82168e1aa3ef5 100644
--- a/clang/lib/CodeGen/CGExprScalar.cpp
+++ b/clang/lib/CodeGen/CGExprScalar.cpp
@@ -2548,30 +2548,42 @@ bool CodeGenFunction::ShouldNullCheckClassCastValue(const CastExpr *CE) {
   return true;
 }
 
+template <typename GetElementTy>
+static Value *
+EmitHLSLElementwiseCastToVector(CodeGenFunction &CGF, QualType DestTy,
+                                unsigned NumSrcElements,
+                                GetElementTy GetElement, SourceLocation Loc) {
+  const auto *VecTy = DestTy->castAs<VectorType>();
+  assert(NumSrcElements >= VecTy->getNumElements() &&
+         "Flattened type on RHS must have the same number or more elements "
+         "than vector on LHS.");
+  Value *V = llvm::PoisonValue::get(CGF.ConvertType(DestTy));
+  for (unsigned I = 0, E = VecTy->getNumElements(); I < E; ++I) {
+    auto [Element, ElementTy] = GetElement(I);
+    Value *Cast = CGF.EmitScalarConversion(Element, ElementTy,
+                                           VecTy->getElementType(), Loc);
+    V = CGF.Builder.CreateInsertElement(V, Cast, I);
+  }
+  return V;
+}
+
 // RHS is an aggregate type
 static Value *EmitHLSLElementwiseCast(CodeGenFunction &CGF, LValue SrcVal,
                                       QualType DestTy, SourceLocation Loc) {
   SmallVector<LValue, 16> LoadList;
   CGF.FlattenAccessAndTypeLValue(SrcVal, LoadList);
   // Dest is either a vector, constant matrix, or a builtin
-  // if its a vector create a temp alloca to store into and return that
-  if (auto *VecTy = DestTy->getAs<VectorType>()) {
-    assert(LoadList.size() >= VecTy->getNumElements() &&
-           "Flattened type on RHS must have the same number or more elements "
-           "than vector on LHS.");
-    llvm::Value *V = llvm::PoisonValue::get(CGF.ConvertType(DestTy));
-    // write to V.
-    for (unsigned I = 0, E = VecTy->getNumElements(); I < E; I++) {
-      RValue RVal = CGF.EmitLoadOfLValue(LoadList[I], Loc);
-      assert(RVal.isScalar() &&
-             "All flattened source values should be scalars.");
-      llvm::Value *Cast =
-          CGF.EmitScalarConversion(RVal.getScalarVal(), LoadList[I].getType(),
-                                   VecTy->getElementType(), Loc);
-      V = CGF.Builder.CreateInsertElement(V, Cast, I);
-    }
-    return V;
-  }
+  if (DestTy->isVectorType())
+    return EmitHLSLElementwiseCastToVector(
+        CGF, DestTy, LoadList.size(),
+        [&](unsigned I) {
+          RValue RVal = CGF.EmitLoadOfLValue(LoadList[I], Loc);
+          assert(RVal.isScalar() &&
+                 "All flattened source values should be scalars.");
+          return std::pair(RVal.getScalarVal(), LoadList[I].getType());
+        },
+        Loc);
+
   if (auto *MatTy = DestTy->getAs<ConstantMatrixType>()) {
     assert(LoadList.size() >= MatTy->getNumElementsFlattened() &&
            "Flattened type on RHS must have the same number or more elements "
@@ -3188,6 +3200,24 @@ Value *ScalarExprEmitter::VisitCastExpr(CastExpr *CE) {
     RValue RV = CGF.EmitAnyExpr(E);
     SourceLocation Loc = CE->getExprLoc();
 
+    if (const auto *SrcMatTy = E->getType()->getAs<ConstantMatrixType>()) {
+      assert(DestTy->isVectorType() &&
+             "Matrix elementwise cast destination must be a vector");
+      assert(RV.isScalar() && "Matrix rvalue must have scalar representation");
+      Value *SrcVal = RV.getScalarVal();
+      return EmitHLSLElementwiseCastToVector(
+          CGF, DestTy, SrcMatTy->getNumElementsFlattened(),
+          [&](unsigned I) {
+            unsigned Row = I / SrcMatTy->getNumColumns();
+            unsigned Col = I % SrcMatTy->getNumColumns();
+            unsigned Idx = SrcMatTy->getColumnMajorFlattenedIndex(Row, Col);
+            Value *Element =
+                Builder.CreateExtractElement(SrcVal, Idx, "matrixext");
+            return std::pair(Element, SrcMatTy->getElementType());
+          },
+          Loc);
+    }
+
     Address SrcAddr = Address::invalid();
 
     if (RV.isAggregate()) {

diff  --git a/clang/test/CodeGenHLSL/BasicFeatures/VectorElementwiseCast.hlsl b/clang/test/CodeGenHLSL/BasicFeatures/VectorElementwiseCast.hlsl
index 99d0bcd9e7497..09af847e25f62 100644
--- a/clang/test/CodeGenHLSL/BasicFeatures/VectorElementwiseCast.hlsl
+++ b/clang/test/CodeGenHLSL/BasicFeatures/VectorElementwiseCast.hlsl
@@ -125,29 +125,22 @@ export void call6(Derived D) {
 // CHECK-LABEL: call7
 // CHECK:    [[M_ADDR:%.*]] = alloca [2 x <2 x float>], align 4
 // CHECK-NEXT:    [[V:%.*]] = alloca <4 x float>, align 4
-// CHECK-NEXT:    [[HLSL_EWCAST_SRC:%.*]] = alloca [2 x <2 x float>], align 4
 // COL-CHECK-NEXT:    store <4 x float> %M, ptr [[M_ADDR]], align 4
 // ROW-CHECK-NEXT:    [[M_ROW:%.*]] = call {{.*}} <4 x float> @llvm.matrix.transpose.v4f32(<4 x float> %M, i32 2, i32 2)
 // ROW-CHECK-NEXT:    store <4 x float> [[M_ROW]], ptr [[M_ADDR]], align 4
 // CHECK-NEXT:    [[TMP0:%.*]] = load <4 x float>, ptr [[M_ADDR]], align 4
-// COL-CHECK-NEXT:    store <4 x float> [[TMP0]], ptr [[HLSL_EWCAST_SRC]], align 4
-// ROW-CHECK-NEXT:    [[TMP0_COL:%.*]] = call {{.*}} <4 x float> @llvm.matrix.transpose.v4f32(<4 x float> [[TMP0]], i32 2, i32 2)
-// ROW-CHECK-NEXT:    [[TMP0_ROW:%.*]] = call {{.*}} <4 x float> @llvm.matrix.transpose.v4f32(<4 x float> [[TMP0_COL]], i32 2, i32 2)
-// ROW-CHECK-NEXT:    store <4 x float> [[TMP0_ROW]], ptr [[HLSL_EWCAST_SRC]], align 4
-// CHECK-NEXT:    [[MATRIX_GEP:%.*]] = getelementptr inbounds <4 x float>, ptr [[HLSL_EWCAST_SRC]], i32 0
-// CHECK-NEXT:    [[TMP2:%.*]] = load <4 x float>, ptr [[MATRIX_GEP]], align 4
-// CHECK-NEXT:    [[MATRIXEXT:%.*]] = extractelement <4 x float> [[TMP2]], i32 0
+// ROW-CHECK-NEXT:    [[M_COL:%.*]] = call {{.*}} <4 x float> @llvm.matrix.transpose.v4f32(<4 x float> [[TMP0]], i32 2, i32 2)
+// COL-CHECK-NEXT:    [[MATRIXEXT:%.*]] = extractelement <4 x float> [[TMP0]], i64 0
+// ROW-CHECK-NEXT:    [[MATRIXEXT:%.*]] = extractelement <4 x float> [[M_COL]], i64 0
 // CHECK-NEXT:    [[TMP3:%.*]] = insertelement <4 x float> poison, float [[MATRIXEXT]], i64 0
-// CHECK-NEXT:    [[TMP4:%.*]] = load <4 x float>, ptr [[MATRIX_GEP]], align 4
-// COL-CHECK-NEXT:    [[MATRIXEXT1:%.*]] = extractelement <4 x float> [[TMP4]], i32 2
-// ROW-CHECK-NEXT:    [[MATRIXEXT1:%.*]] = extractelement <4 x float> [[TMP4]], i32 1
+// COL-CHECK-NEXT:    [[MATRIXEXT1:%.*]] = extractelement <4 x float> [[TMP0]], i64 2
+// ROW-CHECK-NEXT:    [[MATRIXEXT1:%.*]] = extractelement <4 x float> [[M_COL]], i64 2
 // CHECK-NEXT:    [[TMP5:%.*]] = insertelement <4 x float> [[TMP3]], float [[MATRIXEXT1]], i64 1
-// CHECK-NEXT:    [[TMP6:%.*]] = load <4 x float>, ptr [[MATRIX_GEP]], align 4
-// COL-CHECK-NEXT:    [[MATRIXEXT2:%.*]] = extractelement <4 x float> [[TMP6]], i32 1
-// ROW-CHECK-NEXT:    [[MATRIXEXT2:%.*]] = extractelement <4 x float> [[TMP6]], i32 2
+// COL-CHECK-NEXT:    [[MATRIXEXT2:%.*]] = extractelement <4 x float> [[TMP0]], i64 1
+// ROW-CHECK-NEXT:    [[MATRIXEXT2:%.*]] = extractelement <4 x float> [[M_COL]], i64 1
 // CHECK-NEXT:    [[TMP7:%.*]] = insertelement <4 x float> [[TMP5]], float [[MATRIXEXT2]], i64 2
-// CHECK-NEXT:    [[TMP8:%.*]] = load <4 x float>, ptr [[MATRIX_GEP]], align 4
-// CHECK-NEXT:    [[MATRIXEXT3:%.*]] = extractelement <4 x float> [[TMP8]], i32 3
+// COL-CHECK-NEXT:    [[MATRIXEXT3:%.*]] = extractelement <4 x float> [[TMP0]], i64 3
+// ROW-CHECK-NEXT:    [[MATRIXEXT3:%.*]] = extractelement <4 x float> [[M_COL]], i64 3
 // CHECK-NEXT:    [[TMP9:%.*]] = insertelement <4 x float> [[TMP7]], float [[MATRIXEXT3]], i64 3
 // CHECK-NEXT:    store <4 x float> [[TMP9]], ptr [[V]], align 4
 // CHECK-NEXT:    ret void
@@ -160,25 +153,19 @@ export void call7(float2x2 M) {
 // COL-CHECK:    [[M_ADDR:%.*]] = alloca [1 x <3 x i32>], align 4
 // ROW-CHECK:    [[M_ADDR:%.*]] = alloca [3 x <1 x i32>], align 4
 // CHECK-NEXT:    [[V:%.*]] = alloca <3 x i32>, align 4
-// COL-CHECK-NEXT:    [[HLSL_EWCAST_SRC:%.*]] = alloca [1 x <3 x i32>], align 4
-// ROW-CHECK-NEXT:    [[HLSL_EWCAST_SRC:%.*]] = alloca [3 x <1 x i32>], align 4
 // COL-CHECK-NEXT:    store <3 x i32> %M, ptr [[M_ADDR]], align 4
 // ROW-CHECK-NEXT:    [[M_ROW:%.*]] = call <3 x i32> @llvm.matrix.transpose.v3i32(<3 x i32> %M, i32 3, i32 1)
 // ROW-CHECK-NEXT:    store <3 x i32> [[M_ROW]], ptr [[M_ADDR]], align 4
 // CHECK-NEXT:    [[TMP0:%.*]] = load <3 x i32>, ptr [[M_ADDR]], align 4
-// COL-CHECK-NEXT:    store <3 x i32> [[TMP0]], ptr [[HLSL_EWCAST_SRC]], align 4
-// ROW-CHECK-NEXT:    [[TMP0_COL:%.*]] = call <3 x i32> @llvm.matrix.transpose.v3i32(<3 x i32> [[TMP0]], i32 1, i32 3)
-// ROW-CHECK-NEXT:    [[TMP0_ROW:%.*]] = call <3 x i32> @llvm.matrix.transpose.v3i32(<3 x i32> [[TMP0_COL]], i32 3, i32 1)
-// ROW-CHECK-NEXT:    store <3 x i32> [[TMP0_ROW]], ptr [[HLSL_EWCAST_SRC]], align 4
-// CHECK-NEXT:    [[MATRIX_GEP:%.*]] = getelementptr inbounds <3 x i32>, ptr [[HLSL_EWCAST_SRC]], i32 0
-// CHECK-NEXT:    [[TMP2:%.*]] = load <3 x i32>, ptr [[MATRIX_GEP]], align 4
-// CHECK-NEXT:    [[MATRIXEXT:%.*]] = extractelement <3 x i32> [[TMP2]], i32 0
+// ROW-CHECK-NEXT:    [[M_COL:%.*]] = call <3 x i32> @llvm.matrix.transpose.v3i32(<3 x i32> [[TMP0]], i32 1, i32 3)
+// COL-CHECK-NEXT:    [[MATRIXEXT:%.*]] = extractelement <3 x i32> [[TMP0]], i64 0
+// ROW-CHECK-NEXT:    [[MATRIXEXT:%.*]] = extractelement <3 x i32> [[M_COL]], i64 0
 // CHECK-NEXT:    [[TMP3:%.*]] = insertelement <3 x i32> poison, i32 [[MATRIXEXT]], i64 0
-// CHECK-NEXT:    [[TMP4:%.*]] = load <3 x i32>, ptr [[MATRIX_GEP]], align 4
-// CHECK-NEXT:    [[MATRIXEXT1:%.*]] = extractelement <3 x i32> [[TMP4]], i32 1
+// COL-CHECK-NEXT:    [[MATRIXEXT1:%.*]] = extractelement <3 x i32> [[TMP0]], i64 1
+// ROW-CHECK-NEXT:    [[MATRIXEXT1:%.*]] = extractelement <3 x i32> [[M_COL]], i64 1
 // CHECK-NEXT:    [[TMP5:%.*]] = insertelement <3 x i32> [[TMP3]], i32 [[MATRIXEXT1]], i64 1
-// CHECK-NEXT:    [[TMP6:%.*]] = load <3 x i32>, ptr [[MATRIX_GEP]], align 4
-// CHECK-NEXT:    [[MATRIXEXT2:%.*]] = extractelement <3 x i32> [[TMP6]], i32 2
+// COL-CHECK-NEXT:    [[MATRIXEXT2:%.*]] = extractelement <3 x i32> [[TMP0]], i64 2
+// ROW-CHECK-NEXT:    [[MATRIXEXT2:%.*]] = extractelement <3 x i32> [[M_COL]], i64 2
 // CHECK-NEXT:    [[TMP7:%.*]] = insertelement <3 x i32> [[TMP5]], i32 [[MATRIXEXT2]], i64 2
 // CHECK-NEXT:    store <3 x i32> [[TMP7]], ptr [[V]], align 4
 // CHECK-NEXT:    ret void
@@ -191,30 +178,21 @@ export void call8(int3x1 M) {
 // COL-CHECK:    [[M_ADDR:%.*]] = alloca [2 x <1 x i32>], align 4
 // ROW-CHECK:    [[M_ADDR:%.*]] = alloca [1 x <2 x i32>], align 4
 // CHECK-NEXT:    [[V:%.*]] = alloca <2 x i32>, align 4
-// COL-CHECK-NEXT:    [[HLSL_EWCAST_SRC:%.*]] = alloca [2 x <1 x i32>], align 4
-// ROW-CHECK-NEXT:    [[HLSL_EWCAST_SRC:%.*]] = alloca [1 x <2 x i32>], align 4
 // COL-CHECK-NEXT:    [[TMP0:%.*]] = zext <2 x i1> %M to <2 x i32>
 // ROW-CHECK-NEXT:    [[M_ROW:%.*]] = call <2 x i1> @llvm.matrix.transpose.v2i1(<2 x i1> %M, i32 1, i32 2)
 // ROW-CHECK-NEXT:    [[TMP0:%.*]] = zext <2 x i1> [[M_ROW]] to <2 x i32>
 // CHECK-NEXT:    store <2 x i32> [[TMP0]], ptr [[M_ADDR]], align 4
 // CHECK-NEXT:    [[TMP1:%.*]] = load <2 x i32>, ptr [[M_ADDR]], align 4
 // CHECK-NEXT:    [[M_LOADEDV:%.*]] = icmp ne <2 x i32> [[TMP1]], zeroinitializer
-// COL-CHECK-NEXT:    [[M_EXT:%.*]] = zext <2 x i1> [[M_LOADEDV]] to <2 x i32>
-// ROW-CHECK-NEXT:    [[M_LOADEDV_COL:%.*]] = call <2 x i1> @llvm.matrix.transpose.v2i1(<2 x i1> [[M_LOADEDV]], i32 2, i32 1)
-// ROW-CHECK-NEXT:    [[M_LOADEDV_ROW:%.*]] = call <2 x i1> @llvm.matrix.transpose.v2i1(<2 x i1> [[M_LOADEDV_COL]], i32 1, i32 2)
-// ROW-CHECK-NEXT:    [[M_EXT:%.*]] = zext <2 x i1> [[M_LOADEDV_ROW]] to <2 x i32>
-// CHECK-NEXT:    store <2 x i32> [[M_EXT]], ptr [[HLSL_EWCAST_SRC]], align 4
-// CHECK-NEXT:    [[MATRIX_GEP:%.*]] = getelementptr inbounds <2 x i32>, ptr [[HLSL_EWCAST_SRC]], i32 0
-// CHECK-NEXT:    [[TMP3:%.*]] = load <2 x i32>, ptr [[MATRIX_GEP]], align 4
-// CHECK-NEXT:    [[MATRIXEXT:%.*]] = extractelement <2 x i32> [[TMP3]], i32 0
-// CHECK-NEXT:    [[LOADEDV:%.*]] = icmp ne i32 [[MATRIXEXT]], 0
-// CHECK-NEXT:    [[TMP4:%.*]] = insertelement <2 x i1> poison, i1 [[LOADEDV]], i64 0
-// CHECK-NEXT:    [[TMP5:%.*]] = load <2 x i32>, ptr [[MATRIX_GEP]], align 4
-// CHECK-NEXT:    [[MATRIXEXT1:%.*]] = extractelement <2 x i32> [[TMP5]], i32 1
-// CHECK-NEXT:    [[LOADEDV2:%.*]] = icmp ne i32 [[MATRIXEXT1]], 0
-// CHECK-NEXT:    [[TMP6:%.*]] = insertelement <2 x i1> [[TMP4]], i1 [[LOADEDV2]], i64 1
-// CHECK-NEXT:    [[TMP7:%.*]] = zext <2 x i1> [[TMP6]] to <2 x i32>
-// CHECK-NEXT:    store <2 x i32> [[TMP7]], ptr [[V]], align 4
+// ROW-CHECK-NEXT:    [[M_COL:%.*]] = call <2 x i1> @llvm.matrix.transpose.v2i1(<2 x i1> [[M_LOADEDV]], i32 2, i32 1)
+// COL-CHECK-NEXT:    [[MATRIXEXT:%.*]] = extractelement <2 x i1> [[M_LOADEDV]], i64 0
+// ROW-CHECK-NEXT:    [[MATRIXEXT:%.*]] = extractelement <2 x i1> [[M_COL]], i64 0
+// CHECK-NEXT:    [[TMP3:%.*]] = insertelement <2 x i1> poison, i1 [[MATRIXEXT]], i64 0
+// COL-CHECK-NEXT:    [[MATRIXEXT1:%.*]] = extractelement <2 x i1> [[M_LOADEDV]], i64 1
+// ROW-CHECK-NEXT:    [[MATRIXEXT1:%.*]] = extractelement <2 x i1> [[M_COL]], i64 1
+// CHECK-NEXT:    [[TMP5:%.*]] = insertelement <2 x i1> [[TMP3]], i1 [[MATRIXEXT1]], i64 1
+// CHECK-NEXT:    [[TMP6:%.*]] = zext <2 x i1> [[TMP5]] to <2 x i32>
+// CHECK-NEXT:    store <2 x i32> [[TMP6]], ptr [[V]], align 4
 // CHECK-NEXT:    ret void
 export void call9(bool1x2 M) {
     bool2 V = (bool2)M;


        


More information about the cfe-commits mailing list