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
