[llvm] ca27dc2 - [SPIR-V] Matrix in struct pointer legalization (#193073)
via llvm-commits
llvm-commits at lists.llvm.org
Tue Apr 28 06:37:00 PDT 2026
Author: Steven Perron
Date: 2026-04-28T09:36:55-04:00
New Revision: ca27dc2933ed81835569d3e86c16935d850d38ff
URL: https://github.com/llvm/llvm-project/commit/ca27dc2933ed81835569d3e86c16935d850d38ff
DIFF: https://github.com/llvm/llvm-project/commit/ca27dc2933ed81835569d3e86c16935d850d38ff.diff
LOG: [SPIR-V] Matrix in struct pointer legalization (#193073)
When looking to load an object at the start of a struct, the types do
not always match exactly. When we have an HLSL matrix the type in the
load will not match the type in memory. We need to improve the pointer
legalization pass to look for any "compatible" type at the start of an
aggragate.
A compatible are two types that the pass knows know to convert from one
to another.
This involves a refactoring of the code to make the check more general.
Assisted-by: Gemini
<!-- branch-stack-start -->
<!-- branch-stack-end -->
Added:
llvm/test/CodeGen/SPIRV/pointers/load-store-matrix-in-struct.ll
Modified:
llvm/lib/Target/SPIRV/SPIRVLegalizePointerCast.cpp
Removed:
################################################################################
diff --git a/llvm/lib/Target/SPIRV/SPIRVLegalizePointerCast.cpp b/llvm/lib/Target/SPIRV/SPIRVLegalizePointerCast.cpp
index 74df9a057d71b..4a830795fe798 100644
--- a/llvm/lib/Target/SPIRV/SPIRVLegalizePointerCast.cpp
+++ b/llvm/lib/Target/SPIRV/SPIRVLegalizePointerCast.cpp
@@ -145,26 +145,94 @@ class SPIRVLegalizePointerCastImpl {
return Output;
}
- // Loads the first value in an aggregate pointed by |Source| of containing
- // elements of type |ElementType|. Load flags will be copied from |BadLoad|,
- // which should be the load being legalized. Returns the loaded value.
- Value *loadFirstValueFromAggregate(IRBuilder<> &B, Type *ElementType,
- Value *Source, LoadInst *BadLoad) {
- std::array<Type *, 2> Types = {BadLoad->getPointerOperandType(),
- Source->getType()};
- SmallVector<Value *, 8> Args{/* isInBounds= */ B.getInt1(false), Source};
-
- Type *AggregateType = GR->findDeducedElementType(Source);
- assert(AggregateType && "Could not deduce aggregate type");
- buildGEPIndexChain(B, ElementType, AggregateType, Args);
-
- auto *GEP = B.CreateIntrinsic(Intrinsic::spv_gep, {Types}, {Args});
- GR->buildAssignPtr(B, ElementType, GEP);
-
- LoadInst *LI = B.CreateLoad(ElementType, GEP);
- LI->setAlignment(BadLoad->getAlign());
- buildAssignType(B, ElementType, LI);
- return LI;
+ // Returns true if |FromTy| has a memory layout compatible with loading or
+ // storing |ToTy|.
+ bool isCompatibleMemoryLayout(Type *ToTy, Type *FromTy) {
+ if (ToTy == FromTy)
+ return true;
+ auto *SVT = dyn_cast<FixedVectorType>(FromTy);
+ auto *DVT = dyn_cast<FixedVectorType>(ToTy);
+ if (SVT && DVT)
+ return true;
+ auto *SAT = dyn_cast<ArrayType>(FromTy);
+ if (SAT && DVT) {
+ if (SAT->getElementType() == DVT->getElementType())
+ return true;
+ if (auto *MAT = dyn_cast<FixedVectorType>(SAT->getElementType()))
+ if (MAT->getElementType() == DVT->getElementType())
+ return true;
+ }
+ return false;
+ }
+
+ // Traverses the aggregate type to find the first sub-type that matches
+ // the TargetElemType's memory layout, optionally emitting a GEP intrinsic.
+ std::optional<std::pair<Value *, Type *>>
+ getPointerToFirstCompatibleType(IRBuilder<> &B, Value *BasePtr,
+ Type *PointerType, Type *TargetElemType,
+ bool IsInBounds) {
+ Type *CurrentTy = GR->findDeducedElementType(BasePtr);
+ assert(CurrentTy && "Could not deduce aggregate type");
+ SmallVector<Value *, 8> Args{/* isInBounds= */ B.getInt1(IsInBounds),
+ BasePtr};
+ Args.push_back(B.getInt32(0)); // Pointer offset
+
+ while (!isCompatibleMemoryLayout(TargetElemType, CurrentTy)) {
+ if (auto *ST = dyn_cast<StructType>(CurrentTy)) {
+ if (ST->getNumElements() == 0)
+ return std::nullopt;
+ CurrentTy = ST->getTypeAtIndex(0u);
+ } else if (auto *AT = dyn_cast<ArrayType>(CurrentTy)) {
+ CurrentTy = AT->getElementType();
+ } else if (auto *VT = dyn_cast<FixedVectorType>(CurrentTy)) {
+ CurrentTy = VT->getElementType();
+ } else {
+ return std::nullopt;
+ }
+ Args.push_back(B.getInt32(0));
+ }
+
+ Value *GEP = BasePtr;
+ if (Args.size() > 3) {
+ std::array<Type *, 2> Types = {PointerType, BasePtr->getType()};
+ GEP = B.CreateIntrinsic(Intrinsic::spv_gep, {Types}, {Args});
+ GR->buildAssignPtr(B, CurrentTy, GEP);
+ }
+
+ return std::make_pair(GEP, CurrentTy);
+ }
+
+ // Builds a legalized load from a pointer, drilling down through
+ // memory layouts to find a compatible type. Load flags will be
+ // copied from |BadLoad|, which should be the load being legalized.
+ Value *buildLegalizedLoad(IRBuilder<> &B, Type *ElementType, Value *Source,
+ LoadInst *BadLoad) {
+ auto ResultOpt = getPointerToFirstCompatibleType(
+ B, Source, BadLoad->getPointerOperandType(), ElementType, false);
+ assert(ResultOpt && "Failed to load from aggregate: "
+ "Could not find compatible memory layout.");
+ auto [GEP, CurrentTy] = *ResultOpt;
+
+ auto *SAT = dyn_cast<ArrayType>(CurrentTy);
+ auto *SVT = dyn_cast<FixedVectorType>(CurrentTy);
+ auto *DVT = dyn_cast<FixedVectorType>(ElementType);
+ auto *MAT =
+ SAT ? dyn_cast<FixedVectorType>(SAT->getElementType()) : nullptr;
+
+ if (ElementType == CurrentTy) {
+ LoadInst *LI = B.CreateLoad(ElementType, GEP);
+ LI->setAlignment(BadLoad->getAlign());
+ buildAssignType(B, ElementType, LI);
+ return LI;
+ }
+ if (SVT && DVT)
+ return loadVectorFromVector(B, SVT, DVT, GEP);
+ if (SAT && DVT && SAT->getElementType() == DVT->getElementType())
+ return loadVectorFromArray(B, DVT, GEP);
+ if (MAT && DVT && MAT->getElementType() == DVT->getElementType())
+ return loadVectorFromMatrixArray(B, DVT, GEP, MAT);
+
+ llvm_unreachable("Failed to load from aggregate.");
}
Value *
buildVectorFromLoadedElements(IRBuilder<> &B, FixedVectorType *TargetType,
@@ -302,34 +370,10 @@ class SPIRVLegalizePointerCastImpl {
// operand.
void transformLoad(IRBuilder<> &B, LoadInst *LI, Value *CastedOperand,
Value *OriginalOperand) {
- Type *FromTy = GR->findDeducedElementType(OriginalOperand);
Type *ToTy = GR->findDeducedElementType(CastedOperand);
- Value *Output = nullptr;
-
- auto *SAT = dyn_cast<ArrayType>(FromTy);
- auto *SVT = dyn_cast<FixedVectorType>(FromTy);
- auto *DVT = dyn_cast<FixedVectorType>(ToTy);
- auto *MAT =
- SAT ? dyn_cast<FixedVectorType>(SAT->getElementType()) : nullptr;
-
B.SetInsertPoint(LI);
- // Destination is the element type of some member of FromTy. For example,
- // loading the 1st element of an array:
- // - float a = array[0];
- if (isTypeFirstElementAggregate(ToTy, FromTy))
- Output = loadFirstValueFromAggregate(B, ToTy, OriginalOperand, LI);
- // Destination is a smaller vector than source or
diff erent vector type.
- // - float3 v3 = vector4;
- // - float4 v2 = int4;
- else if (SVT && DVT)
- Output = loadVectorFromVector(B, SVT, DVT, OriginalOperand);
- else if (SAT && DVT && SAT->getElementType() == DVT->getElementType())
- Output = loadVectorFromArray(B, DVT, OriginalOperand);
- else if (MAT && DVT && MAT->getElementType() == DVT->getElementType())
- Output = loadVectorFromMatrixArray(B, DVT, OriginalOperand, MAT);
- else
- llvm_unreachable("Unimplemented implicit down-cast from load.");
+ Value *Output = buildLegalizedLoad(B, ToTy, OriginalOperand, LI);
GR->replaceAllUsesWith(LI, Output, /* DeleteOld= */ true);
DeadInstructions.push_back(LI);
@@ -410,75 +454,49 @@ class SPIRVLegalizePointerCastImpl {
return SI;
}
- void buildGEPIndexChain(IRBuilder<> &B, Type *Search, Type *Aggregate,
- SmallVectorImpl<Value *> &Indices) {
- Indices.push_back(B.getInt32(0));
-
- if (Search == Aggregate)
+ // Builds a legalized store to a pointer, drilling down through
+ // memory layouts to find a compatible type.
+ void buildLegalizedStore(IRBuilder<> &B, Value *Src, Value *Dst,
+ Align Alignment) {
+ auto ResultOpt = getPointerToFirstCompatibleType(B, Dst, Dst->getType(),
+ Src->getType(), true);
+ assert(ResultOpt && "Failed to store to aggregate: "
+ "Could not find compatible memory layout.");
+ auto [GEP, CurrentTy] = *ResultOpt;
+
+ auto *DAT = dyn_cast<ArrayType>(CurrentTy);
+ auto *DVT = dyn_cast<FixedVectorType>(CurrentTy);
+ auto *SVT = dyn_cast<FixedVectorType>(Src->getType());
+ auto *DMAT =
+ DAT ? dyn_cast<FixedVectorType>(DAT->getElementType()) : nullptr;
+
+ if (Src->getType() == CurrentTy) {
+ StoreInst *SI = B.CreateStore(Src, GEP);
+ SI->setAlignment(Alignment);
return;
+ }
+ if (DVT && SVT) {
+ storeVectorFromVector(B, Src, GEP, Alignment);
+ return;
+ }
+ if (DAT && SVT && SVT->getElementType() == DAT->getElementType()) {
+ storeArrayFromVector(B, Src, GEP, DAT, Alignment);
+ return;
+ }
+ if (DMAT && SVT && DMAT->getElementType() == SVT->getElementType()) {
+ storeMatrixArrayFromVector(B, Src, GEP, DAT, Alignment);
+ return;
+ }
- if (auto *ST = dyn_cast<StructType>(Aggregate))
- buildGEPIndexChain(B, Search, ST->getTypeAtIndex(0u), Indices);
- else if (auto *AT = dyn_cast<ArrayType>(Aggregate))
- buildGEPIndexChain(B, Search, AT->getElementType(), Indices);
- else if (auto *VT = dyn_cast<FixedVectorType>(Aggregate))
- buildGEPIndexChain(B, Search, VT->getElementType(), Indices);
- else
- llvm_unreachable("Bad access chain?");
- }
-
- // Stores the given Src value into the first entry of the Dst aggregate.
- Value *storeToFirstValueAggregate(IRBuilder<> &B, Value *Src, Value *Dst,
- Type *DstPointeeType, Align Alignment) {
- std::array<Type *, 2> Types = {Dst->getType(), Dst->getType()};
- SmallVector<Value *, 8> Args{/* isInBounds= */ B.getInt1(true), Dst};
- buildGEPIndexChain(B, Src->getType(), DstPointeeType, Args);
- auto *GEP = B.CreateIntrinsic(Intrinsic::spv_gep, {Types}, {Args});
- GR->buildAssignPtr(B, Src->getType(), GEP);
- StoreInst *SI = B.CreateStore(Src, GEP);
- SI->setAlignment(Alignment);
- return SI;
- }
-
- bool isTypeFirstElementAggregate(Type *Search, Type *Aggregate) {
- if (Search == Aggregate)
- return true;
- if (auto *ST = dyn_cast<StructType>(Aggregate))
- return isTypeFirstElementAggregate(Search, ST->getTypeAtIndex(0u));
- if (auto *VT = dyn_cast<FixedVectorType>(Aggregate))
- return isTypeFirstElementAggregate(Search, VT->getElementType());
- if (auto *AT = dyn_cast<ArrayType>(Aggregate))
- return isTypeFirstElementAggregate(Search, AT->getElementType());
- return false;
+ llvm_unreachable("Failed to store to aggregate.");
}
// Transforms a store instruction (or SPV intrinsic) using a ptrcast as
// operand into a valid logical SPIR-V store with no ptrcast.
void transformStore(IRBuilder<> &B, Instruction *BadStore, Value *Src,
Value *Dst, Align Alignment) {
- Type *ToTy = GR->findDeducedElementType(Dst);
- Type *FromTy = Src->getType();
-
- auto *S_VT = dyn_cast<FixedVectorType>(FromTy);
- auto *D_VT = dyn_cast<FixedVectorType>(ToTy);
- auto *D_AT = dyn_cast<ArrayType>(ToTy);
- auto *D_MAT =
- D_AT ? dyn_cast<FixedVectorType>(D_AT->getElementType()) : nullptr;
-
B.SetInsertPoint(BadStore);
- if (isTypeFirstElementAggregate(FromTy, ToTy))
- storeToFirstValueAggregate(B, Src, Dst, ToTy, Alignment);
- else if (D_VT && S_VT)
- storeVectorFromVector(B, Src, Dst, Alignment);
- else if (D_VT && !S_VT && FromTy == D_VT->getElementType())
- storeToFirstValueAggregate(B, Src, Dst, D_VT, Alignment);
- else if (D_AT && S_VT && S_VT->getElementType() == D_AT->getElementType())
- storeArrayFromVector(B, Src, Dst, D_AT, Alignment);
- else if (D_MAT && S_VT && D_MAT->getElementType() == S_VT->getElementType())
- storeMatrixArrayFromVector(B, Src, Dst, D_AT, Alignment);
- else
- llvm_unreachable("Unsupported ptrcast use in store. Please fix.");
-
+ buildLegalizedStore(B, Src, Dst, Alignment);
DeadInstructions.push_back(BadStore);
}
diff --git a/llvm/test/CodeGen/SPIRV/pointers/load-store-matrix-in-struct.ll b/llvm/test/CodeGen/SPIRV/pointers/load-store-matrix-in-struct.ll
new file mode 100644
index 0000000000000..cfb7ea0297d3f
--- /dev/null
+++ b/llvm/test/CodeGen/SPIRV/pointers/load-store-matrix-in-struct.ll
@@ -0,0 +1,57 @@
+; RUN: llc -verify-machineinstrs -O0 -mtriple=spirv1.6-unknown-vulkan1.3 %s -stop-after=spirv-legalize-pointer-cast -o - | FileCheck %s --check-prefix=IRCHECK
+; RUN: llc -verify-machineinstrs -O0 -mtriple=spirv1.6-unknown-vulkan1.3 %s -o - | FileCheck %s
+; RUN: %if spirv-tools %{ llc -O0 -mtriple=spirv1.6-unknown-vulkan1.3 %s -o - -filetype=obj | spirv-val --target-env vulkan1.3 %}
+
+; CHECK-DAG: [[FLOAT:%[0-9]+]] = OpTypeFloat 32
+; CHECK-DAG: [[VEC2FLOAT:%[0-9]+]] = OpTypeVector [[FLOAT]] 2
+; CHECK-DAG: [[UINT:%[0-9]+]] = OpTypeInt 32 0
+; CHECK-DAG: [[UINT4:%[0-9]+]] = OpConstant [[UINT]] 4
+; CHECK-DAG: [[ARRAY4VEC2:%[0-9]+]] = OpTypeArray [[VEC2FLOAT]] [[UINT4]]
+; CHECK-DAG: [[STRUCT:%[0-9]+]] = OpTypeStruct [[ARRAY4VEC2]]
+
+; Load:
+; CHECK: [[IN_AC_STRUCT:%[0-9]+]] = OpAccessChain %{{[0-9]+}} %{{[0-9]+}} %{{[0-9]+}} %{{[0-9]+}}
+; CHECK: [[IN_AC_ARR:%[0-9]+]] = OpAccessChain %{{[0-9]+}} [[IN_AC_STRUCT]] %{{[0-9]+}}
+; CHECK: [[IN_AC_VEC:%[0-9]+]] = OpAccessChain %{{[0-9]+}} [[IN_AC_ARR]] %{{[0-9]+}}
+; CHECK: OpLoad [[VEC2FLOAT]] [[IN_AC_VEC]]
+
+; Store:
+; CHECK: [[OUT_AC_STRUCT:%[0-9]+]] = OpAccessChain %{{[0-9]+}} %{{[0-9]+}} %{{[0-9]+}} %{{[0-9]+}}
+; CHECK: [[OUT_AC_ARR:%[0-9]+]] = OpInBoundsAccessChain %{{[0-9]+}} [[OUT_AC_STRUCT]] %{{[0-9]+}}
+; CHECK: [[OUT_AC_VEC:%[0-9]+]] = OpAccessChain %{{[0-9]+}} [[OUT_AC_ARR]] %{{[0-9]+}}
+; CHECK: OpStore [[OUT_AC_VEC]]
+
+; IRCHECK-LABEL: define void @main
+
+; IRCHECK: [[IN_PTR:%[0-9]+]] = {{.*}}call {{.*}} ptr addrspace(11) @llvm.spv.resource.getpointer
+; IRCHECK: [[GEP_STRUCT_IN:%[0-9]+]] = call ptr addrspace(11) (i1, ptr addrspace(11), ...) @llvm.spv.gep.p11.p11(i1 false, ptr addrspace(11) [[IN_PTR]], i32 0, i32 0)
+; IRCHECK: [[GEP_ROW_IN:%[0-9]+]] = call ptr addrspace(11) (i1, ptr addrspace(11), ...) @llvm.spv.gep.p11.p11(i1 false, ptr addrspace(11) [[GEP_STRUCT_IN]], i32 0, i32 0)
+; IRCHECK: load <2 x float>, ptr addrspace(11) [[GEP_ROW_IN]]
+
+; IRCHECK: [[OUT_PTR:%[0-9]+]] = {{.*}}call {{.*}} ptr addrspace(11) @llvm.spv.resource.getpointer
+; IRCHECK: [[GEP_STRUCT_OUT:%[0-9]+]] = call ptr addrspace(11) (i1, ptr addrspace(11), ...) @llvm.spv.gep.p11.p11(i1 true, ptr addrspace(11) [[OUT_PTR]], i32 0, i32 0)
+; IRCHECK: [[GEP_ROW_OUT:%[0-9]+]] = call ptr addrspace(11) (i1, ptr addrspace(11), ...) @llvm.spv.gep.p11.p11(i1 false, ptr addrspace(11) [[GEP_STRUCT_OUT]], i32 0, i32 0)
+; IRCHECK: store <2 x float> {{%[0-9]+}}, ptr addrspace(11) [[GEP_ROW_OUT]]
+
+%struct.S = type { [4 x <2 x float>] }
+
+ at .str = private unnamed_addr constant [3 x i8] c"In\00", align 1
+ at .str.2 = private unnamed_addr constant [4 x i8] c"Out\00", align 1
+
+define void @main() #0 {
+entry:
+ %0 = tail call target("spirv.VulkanBuffer", [0 x %struct.S], 12, 0) @llvm.spv.resource.handlefrombinding.tspirv.VulkanBuffer(i32 0, i32 0, i32 1, i32 0, ptr nonnull @.str)
+ %1 = tail call noundef align 4 dereferenceable(32) ptr addrspace(11) @llvm.spv.resource.getpointer.p11.tspirv.VulkanBuffer(target("spirv.VulkanBuffer", [0 x %struct.S], 12, 0) %0, i32 0)
+ %2 = load <8 x float>, ptr addrspace(11) %1, align 4
+
+ %3 = tail call target("spirv.VulkanBuffer", [0 x %struct.S], 12, 0) @llvm.spv.resource.handlefrombinding.tspirv.VulkanBuffer(i32 0, i32 1, i32 1, i32 0, ptr nonnull @.str.2)
+ %4 = tail call noundef align 4 dereferenceable(32) ptr addrspace(11) @llvm.spv.resource.getpointer.p11.tspirv.VulkanBuffer(target("spirv.VulkanBuffer", [0 x %struct.S], 12, 0) %3, i32 0)
+ store <8 x float> %2, ptr addrspace(11) %4, align 4
+
+ ret void
+}
+
+declare target("spirv.VulkanBuffer", [0 x %struct.S], 12, 0) @llvm.spv.resource.handlefrombinding.tspirv.VulkanBuffer(i32, i32, i32, i32, ptr)
+declare ptr addrspace(11) @llvm.spv.resource.getpointer.p11.tspirv.VulkanBuffer(target("spirv.VulkanBuffer", [0 x %struct.S], 12, 0), i32)
+
+attributes #0 = { "hlsl.numthreads"="1,1,1" "hlsl.shader"="compute" }
More information about the llvm-commits
mailing list