[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