https://github.com/farzonl created 
https://github.com/llvm/llvm-project/pull/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

>From 0ec864ef56b19d8933f8ee5b83d6e4790a5e68a1 Mon Sep 17 00:00:00 2001
From: Farzon Lotfi <[email protected]>
Date: Wed, 23 Sep 2026 01:25:41 -0400
Subject: [PATCH] [HLSL] Extract matrix elementwise cast operands from SSA

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
---
 clang/lib/CodeGen/CGExprScalar.cpp            | 69 ++++++++++++++-----
 .../BasicFeatures/VectorElementwiseCast.hlsl  | 61 ++++++----------
 2 files changed, 71 insertions(+), 59 deletions(-)

diff --git a/clang/lib/CodeGen/CGExprScalar.cpp 
b/clang/lib/CodeGen/CGExprScalar.cpp
index 0c310816bc268e..488df3345e98cc 100644
--- a/clang/lib/CodeGen/CGExprScalar.cpp
+++ b/clang/lib/CodeGen/CGExprScalar.cpp
@@ -2548,31 +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 = CGF.Builder.CreateLoad(
+      CGF.CreateIRTempWithoutCast(DestTy, "flatcast.tmp"));
+  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 = CGF.Builder.CreateLoad(
-        CGF.CreateIRTempWithoutCast(DestTy, "flatcast.tmp"));
-    // 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 "
@@ -3190,6 +3201,26 @@ 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");
+            Element = CGF.EmitFromMemory(Element, SrcMatTy->getElementType());
+            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 5156dccb26eca2..4bb162a4cb5c02 100644
--- a/clang/test/CodeGenHLSL/BasicFeatures/VectorElementwiseCast.hlsl
+++ b/clang/test/CodeGenHLSL/BasicFeatures/VectorElementwiseCast.hlsl
@@ -131,31 +131,24 @@ 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
 // CHECK-NEXT:    [[FLATCAST_TMP:%.*]] = alloca <4 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
+// ROW-CHECK-NEXT:    [[M_COL:%.*]] = call {{.*}} <4 x float> 
@llvm.matrix.transpose.v4f32(<4 x float> [[TMP0]], i32 2, i32 2)
 // CHECK-NEXT:    [[TMP1:%.*]] = load <4 x float>, ptr [[FLATCAST_TMP]], align 
4
-// CHECK-NEXT:    [[TMP2:%.*]] = load <4 x float>, ptr [[MATRIX_GEP]], align 4
-// CHECK-NEXT:    [[MATRIXEXT:%.*]] = extractelement <4 x float> [[TMP2]], i32 0
+// 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> [[TMP1]], 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
@@ -168,27 +161,21 @@ 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
 // CHECK-NEXT:    [[FLATCAST_TMP:%.*]] = alloca <3 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
+// ROW-CHECK-NEXT:    [[M_COL:%.*]] = call <3 x i32> 
@llvm.matrix.transpose.v3i32(<3 x i32> [[TMP0]], i32 1, i32 3)
 // CHECK-NEXT:    [[TMP1:%.*]] = load <3 x i32>, ptr [[FLATCAST_TMP]], align 4
-// CHECK-NEXT:    [[TMP2:%.*]] = load <3 x i32>, ptr [[MATRIX_GEP]], align 4
-// CHECK-NEXT:    [[MATRIXEXT:%.*]] = extractelement <3 x i32> [[TMP2]], i32 0
+// 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> [[TMP1]], 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
@@ -201,26 +188,20 @@ 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
 // CHECK-NEXT:    [[FLATCAST_TMP:%.*]] = alloca <2 x i1>, 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
-// COL-CHECK-NEXT:    store <2 x i32> [[TMP1]], ptr [[HLSL_EWCAST_SRC]], align 
4
-// ROW-CHECK-NEXT:    [[TMP1_COL:%.*]] = call <2 x i32> 
@llvm.matrix.transpose.v2i32(<2 x i32> [[TMP1]], i32 2, i32 1)
-// ROW-CHECK-NEXT:    [[TMP1_ROW:%.*]] = call <2 x i32> 
@llvm.matrix.transpose.v2i32(<2 x i32> [[TMP1_COL]], i32 1, i32 2)
-// ROW-CHECK-NEXT:    store <2 x i32> [[TMP1_ROW]], ptr [[HLSL_EWCAST_SRC]], 
align 4
-// CHECK-NEXT:    [[MATRIX_GEP:%.*]] = getelementptr inbounds <2 x i32>, ptr 
[[HLSL_EWCAST_SRC]], i32 0
+// ROW-CHECK-NEXT:    [[M_COL:%.*]] = call <2 x i32> 
@llvm.matrix.transpose.v2i32(<2 x i32> [[TMP1]], i32 2, i32 1)
 // CHECK-NEXT:    [[TMP2:%.*]] = load <2 x i1>, ptr [[FLATCAST_TMP]], align 4
-// CHECK-NEXT:    [[TMP3:%.*]] = load <2 x i32>, ptr [[MATRIX_GEP]], align 4
-// CHECK-NEXT:    [[MATRIXEXT:%.*]] = extractelement <2 x i32> [[TMP3]], i32 0
+// COL-CHECK-NEXT:    [[MATRIXEXT:%.*]] = extractelement <2 x i32> [[TMP1]], 
i64 0
+// ROW-CHECK-NEXT:    [[MATRIXEXT:%.*]] = extractelement <2 x i32> [[M_COL]], 
i64 0
 // CHECK-NEXT:    [[LOADEDV:%.*]] = icmp ne i32 [[MATRIXEXT]], 0
 // CHECK-NEXT:    [[TMP4:%.*]] = insertelement <2 x i1> [[TMP2]], 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
+// COL-CHECK-NEXT:    [[MATRIXEXT1:%.*]] = extractelement <2 x i32> [[TMP1]], 
i64 1
+// ROW-CHECK-NEXT:    [[MATRIXEXT1:%.*]] = extractelement <2 x i32> [[M_COL]], 
i64 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>

_______________________________________________
cfe-commits mailing list
[email protected]
https://lists.llvm.org/cgi-bin/mailman/listinfo/cfe-commits

Reply via email to