1//===-- SPIRVLegalizePointerCast.cpp ----------------------*- C++ -*-===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// The LLVM IR has multiple legal patterns we cannot lower to Logical SPIR-V.
10// This pass modifies such loads to have an IR we can directly lower to valid
11// logical SPIR-V.
12// OpenCL can avoid this because they rely on ptrcast, which is not supported
13// by logical SPIR-V.
14//
15// This pass relies on the assign_ptr_type intrinsic to deduce the type of the
16// pointed values, must replace all occurences of `ptrcast`. This is why
17// unhandled cases are reported as unreachable: we MUST cover all cases.
18//
19// 1. Loading the first element of an array
20//
21// %array = [10 x i32]
22// %value = load i32, ptr %array
23//
24// LLVM can skip the GEP instruction, and only request loading the first 4
25// bytes. In logical SPIR-V, we need an OpAccessChain to access the first
26// element. This pass will add a getelementptr instruction before the load.
27//
28//
29// 2. Implicit downcast from load
30//
31// %1 = getelementptr <4 x i32>, ptr %vec4, i64 0
32// %2 = load <3 x i32>, ptr %1
33//
34// The pointer in the GEP instruction is only used for offset computations,
35// but it doesn't NEED to match the pointed type. OpAccessChain however
36// requires this. Also, LLVM loads define the bitwidth of the load, not the
37// pointer. In this example, we can guess %vec4 is a vec4 thanks to the GEP
38// instruction basetype, but we only want to load the first 3 elements, hence
39// do a partial load. In logical SPIR-V, this is not legal. What we must do
40// is load the full vector (basetype), extract 3 elements, and recombine them
41// to form a 3-element vector.
42//
43//===----------------------------------------------------------------------===//
44
45#include "SPIRV.h"
46#include "SPIRVSubtarget.h"
47#include "SPIRVTargetMachine.h"
48#include "SPIRVUtils.h"
49#include "llvm/IR/IRBuilder.h"
50#include "llvm/IR/IntrinsicInst.h"
51#include "llvm/IR/Intrinsics.h"
52#include "llvm/IR/IntrinsicsSPIRV.h"
53#include "llvm/Transforms/Utils/Cloning.h"
54#include "llvm/Transforms/Utils/LowerMemIntrinsics.h"
55
56using namespace llvm;
57
58namespace {
59class SPIRVLegalizePointerCastImpl {
60
61 // Builds the `spv_assign_type` assigning |Ty| to |Value| at the current
62 // builder position.
63 void buildAssignType(IRBuilder<> &B, Type *Ty, Value *Arg) {
64 Value *OfType = PoisonValue::get(T: Ty);
65 CallInst *AssignCI = buildIntrWithMD(IntrID: Intrinsic::spv_assign_type,
66 Types: {Arg->getType()}, Arg: OfType, Arg2: Arg, Imms: {}, B);
67 GR->addAssignPtrTypeInstr(Val: Arg, AssignPtrTyCI: AssignCI);
68 }
69
70 static FixedVectorType *makeVectorFromTotalBits(Type *ElemTy,
71 TypeSize TotalBits) {
72 unsigned ElemBits = ElemTy->getScalarSizeInBits();
73 assert(ElemBits && TotalBits % ElemBits == 0 &&
74 "TotalBits must be divisible by element bit size");
75 return FixedVectorType::get(ElementType: ElemTy, NumElts: TotalBits / ElemBits);
76 }
77
78 Value *resizeVectorBitsWithShuffle(IRBuilder<> &B, Value *V,
79 FixedVectorType *DstTy) {
80 auto *SrcTy = cast<FixedVectorType>(Val: V->getType());
81 assert(SrcTy->getElementType() == DstTy->getElementType() &&
82 "shuffle resize expects identical element types");
83
84 const unsigned NumNeeded = DstTy->getNumElements();
85 const unsigned NumSource = SrcTy->getNumElements();
86
87 SmallVector<int> Mask(NumNeeded);
88 for (unsigned I = 0; I < NumNeeded; ++I)
89 Mask[I] = (I < NumSource) ? static_cast<int>(I) : -1;
90
91 Value *Resized = B.CreateShuffleVector(V1: V, V2: V, Mask);
92 buildAssignType(B, Ty: DstTy, Arg: Resized);
93 return Resized;
94 }
95
96 // Loads parts of the vector of type |SourceType| from the pointer |Source|
97 // and create a new vector of type |TargetType|. |TargetType| must be a vector
98 // type.
99 // Returns the loaded value.
100 Value *loadVectorFromVector(IRBuilder<> &B, FixedVectorType *SourceType,
101 FixedVectorType *TargetType, Value *Source,
102 Align OriginalAlign) {
103 LoadInst *NewLoad = B.CreateLoad(Ty: SourceType, Ptr: Source);
104 NewLoad->setAlignment(OriginalAlign);
105 buildAssignType(B, Ty: SourceType, Arg: NewLoad);
106 Value *AssignValue = NewLoad;
107 if (TargetType->getElementType() != SourceType->getElementType()) {
108 const DataLayout &DL = B.GetInsertBlock()->getModule()->getDataLayout();
109 TypeSize TargetTypeSize = DL.getTypeSizeInBits(Ty: TargetType);
110 TypeSize SourceTypeSize = DL.getTypeSizeInBits(Ty: SourceType);
111
112 Value *BitcastSrcVal = NewLoad;
113 FixedVectorType *BitcastSrcTy =
114 cast<FixedVectorType>(Val: BitcastSrcVal->getType());
115 FixedVectorType *BitcastDstTy = TargetType;
116
117 if (TargetTypeSize != SourceTypeSize) {
118 unsigned TargetElemBits =
119 TargetType->getElementType()->getScalarSizeInBits();
120 if (SourceTypeSize % TargetElemBits == 0) {
121 // No Resize needed. Same total bits as source, but use target element
122 // type.
123 BitcastDstTy = makeVectorFromTotalBits(ElemTy: TargetType->getElementType(),
124 TotalBits: SourceTypeSize);
125 } else {
126 // Resize source to target total bitwidth using source element type.
127 BitcastSrcTy = makeVectorFromTotalBits(ElemTy: SourceType->getElementType(),
128 TotalBits: TargetTypeSize);
129 BitcastSrcVal = resizeVectorBitsWithShuffle(B, V: NewLoad, DstTy: BitcastSrcTy);
130 }
131 }
132 AssignValue =
133 B.CreateIntrinsic(ID: Intrinsic::spv_bitcast,
134 OverloadTypes: {BitcastDstTy, BitcastSrcTy}, Args: {BitcastSrcVal});
135 buildAssignType(B, Ty: BitcastDstTy, Arg: AssignValue);
136 if (BitcastDstTy == TargetType)
137 return AssignValue;
138 }
139
140 auto *AssignVecTy = cast<FixedVectorType>(Val: AssignValue->getType());
141 const unsigned NumTarget = TargetType->getNumElements();
142 const unsigned NumSource = AssignVecTy->getNumElements();
143
144 // Optimizations may widen a narrow load to cover padding (e.g., loading a
145 // <1 x float> column as <4 x float>). Since extra lanes read trailing
146 // padding, insert only the valid lanes into a poison vector to avoid poison
147 // scalars.
148 if (NumTarget > NumSource) {
149 Value *Result = PoisonValue::get(T: TargetType);
150 buildAssignType(B, Ty: TargetType, Arg: Result);
151 for (unsigned I = 0; I < NumSource; ++I) {
152 Value *Scalar = extractScalarFromVector(B, Vector: AssignValue, Index: I);
153 Result = makeInsertElement(B, Vector: Result, Element: Scalar, Index: I);
154 }
155 return Result;
156 }
157
158 assert(NumTarget < NumSource);
159 SmallVector<int> Mask(/* Size= */ NumTarget);
160 for (unsigned I = 0; I < NumTarget; ++I)
161 Mask[I] = I;
162 Value *Output = B.CreateShuffleVector(V1: AssignValue, V2: AssignValue, Mask);
163 buildAssignType(B, Ty: TargetType, Arg: Output);
164 return Output;
165 }
166
167 // Returns true if |FromTy| has a memory layout compatible with loading or
168 // storing |ToTy|.
169 bool isCompatibleMemoryLayout(Type *ToTy, Type *FromTy) {
170 if (ToTy == FromTy)
171 return true;
172 auto *SVT = dyn_cast<FixedVectorType>(Val: FromTy);
173 auto *DVT = dyn_cast<FixedVectorType>(Val: ToTy);
174 if (SVT && DVT)
175 return true;
176 auto *SAT = dyn_cast<ArrayType>(Val: FromTy);
177 if (SAT && DVT) {
178 if (SAT->getElementType() == DVT->getElementType())
179 return true;
180 if (auto *MAT = dyn_cast<FixedVectorType>(Val: SAT->getElementType()))
181 if (MAT->getElementType() == DVT->getElementType())
182 return true;
183 }
184 return false;
185 }
186
187 // Traverses the aggregate type to find the first sub-type that matches
188 // the TargetElemType's memory layout, optionally emitting a GEP intrinsic.
189 std::optional<std::pair<Value *, Type *>>
190 getPointerToFirstCompatibleType(IRBuilder<> &B, Value *BasePtr,
191 Type *PointerType, Type *TargetElemType,
192 bool IsInBounds) {
193 Type *CurrentTy = GR->findDeducedElementType(Val: BasePtr);
194 assert(CurrentTy && "Could not deduce aggregate type");
195 SmallVector<Value *, 8> Args{/* isInBounds= */ B.getInt1(V: IsInBounds),
196 BasePtr};
197 Args.push_back(Elt: B.getInt32(C: 0)); // Pointer offset
198
199 while (!isCompatibleMemoryLayout(ToTy: TargetElemType, FromTy: CurrentTy)) {
200 if (auto *ST = dyn_cast<StructType>(Val: CurrentTy)) {
201 if (ST->getNumElements() == 0)
202 return std::nullopt;
203 CurrentTy = ST->getTypeAtIndex(N: 0u);
204 } else if (auto *AT = dyn_cast<ArrayType>(Val: CurrentTy)) {
205 CurrentTy = AT->getElementType();
206 } else if (auto *VT = dyn_cast<FixedVectorType>(Val: CurrentTy)) {
207 CurrentTy = VT->getElementType();
208 } else {
209 return std::nullopt;
210 }
211 Args.push_back(Elt: B.getInt32(C: 0));
212 }
213
214 Value *GEP = BasePtr;
215 if (Args.size() > 3) {
216 std::array<Type *, 2> Types = {PointerType, BasePtr->getType()};
217 GEP = B.CreateIntrinsic(ID: Intrinsic::spv_gep, OverloadTypes: {Types}, Args: {Args});
218 GR->buildAssignPtr(B, ElemTy: CurrentTy, Arg: GEP);
219 }
220
221 return std::make_pair(x&: GEP, y&: CurrentTy);
222 }
223
224 static IntrinsicInst *getResourceGetPointer(Value *Ptr) {
225 if (auto *II = dyn_cast<IntrinsicInst>(Val: Ptr))
226 if (II->getIntrinsicID() == Intrinsic::spv_resource_getpointer)
227 return II;
228 return nullptr;
229 }
230
231 Value *gepByteOffset(IRBuilder<> &B, Value *BasePtr, unsigned ByteOffset) {
232 if (ByteOffset == 0)
233 return BasePtr;
234
235 if (IntrinsicInst *ResourcePtr = getResourceGetPointer(Ptr: BasePtr)) {
236 Value *Handle = ResourcePtr->getOperand(i_nocapture: 0);
237 Value *BaseOffset = ResourcePtr->getOperand(i_nocapture: 1);
238 Value *NewOffset;
239 if (auto *CI = dyn_cast<ConstantInt>(Val: BaseOffset))
240 NewOffset =
241 ConstantInt::get(Ty: CI->getType(), V: CI->getZExtValue() + ByteOffset);
242 else
243 NewOffset = B.CreateAdd(
244 LHS: BaseOffset, RHS: ConstantInt::get(Ty: BaseOffset->getType(), V: ByteOffset));
245 SmallVector<OperandBundleDef> OpBundles;
246 ResourcePtr->getOperandBundlesAsDefs(Defs&: OpBundles);
247 CallInst *ResourcePtrAtOffset = B.CreateCall(
248 FTy: ResourcePtr->getFunctionType(), Callee: ResourcePtr->getCalledOperand(),
249 Args: {Handle, NewOffset}, OpBundles);
250 ResourcePtrAtOffset->setAttributes(ResourcePtr->getAttributes());
251 ResourcePtrAtOffset->setCallingConv(ResourcePtr->getCallingConv());
252 Type *I8Ty = Type::getInt8Ty(C&: B.getContext());
253 GR->buildAssignPtr(B, ElemTy: I8Ty, Arg: ResourcePtrAtOffset);
254 return ResourcePtrAtOffset;
255 }
256 llvm_unreachable(
257 "byte layout pointer must come from spv.resource.getpointer");
258 }
259
260 Value *scalarToStoreInt(IRBuilder<> &B, Value *Scalar) {
261 Type *Ty = Scalar->getType();
262 const DataLayout &DL = B.GetInsertBlock()->getModule()->getDataLayout();
263 Type *IntTy =
264 IntegerType::get(C&: B.getContext(), NumBits: DL.getTypeStoreSizeInBits(Ty));
265 if (Ty == IntTy)
266 return Scalar;
267 if (Ty->isIntOrIntVectorTy())
268 return B.CreateIntCast(V: Scalar, DestTy: IntTy, /*isSigned=*/false);
269 return B.CreateBitCast(V: Scalar, DestTy: IntTy);
270 }
271
272 Value *storeIntToScalar(IRBuilder<> &B, Value *IntVal, Type *ScalarTy) {
273 if (IntVal->getType() == ScalarTy)
274 return IntVal;
275 if (ScalarTy->isIntOrIntVectorTy())
276 return B.CreateIntCast(V: IntVal, DestTy: ScalarTy, /*isSigned=*/false);
277 return B.CreateBitCast(V: IntVal, DestTy: ScalarTy);
278 }
279
280 void storeScalarToByteLayout(IRBuilder<> &B, Value *Src, Value *Dst,
281 Align Alignment) {
282 LLVMContext &Ctx = B.getContext();
283 Type *I8Ty = Type::getInt8Ty(C&: Ctx);
284 const DataLayout &DL = B.GetInsertBlock()->getModule()->getDataLayout();
285 Value *IntVal = scalarToStoreInt(B, Scalar: Src);
286 unsigned NumBytes = DL.getTypeStoreSize(Ty: Src->getType());
287
288 auto StoreByte = [&](unsigned I, Value *Shifted) {
289 Value *Byte = B.CreateTrunc(V: Shifted, DestTy: I8Ty);
290 buildAssignType(B, Ty: I8Ty, Arg: Byte);
291 Value *Ptr = gepByteOffset(B, BasePtr: Dst, ByteOffset: I);
292 StoreInst *SI = B.CreateStore(Val: Byte, Ptr);
293 SI->setAlignment(commonAlignment(A: Alignment, Offset: I));
294 };
295
296 if (NumBytes > 0)
297 StoreByte(0, IntVal);
298
299 for (unsigned I = 1; I < NumBytes; ++I) {
300 Value *Shifted =
301 B.CreateLShr(LHS: IntVal, RHS: ConstantInt::get(Ty: IntVal->getType(), V: 8 * I));
302 StoreByte(I, Shifted);
303 }
304 }
305
306 Value *loadScalarFromByteLayout(IRBuilder<> &B, Type *AccessTy, Value *Src,
307 Align Alignment) {
308 LLVMContext &Ctx = B.getContext();
309 Type *I8Ty = Type::getInt8Ty(C&: Ctx);
310 const DataLayout &DL = B.GetInsertBlock()->getModule()->getDataLayout();
311 unsigned NumBytes = DL.getTypeStoreSize(Ty: AccessTy);
312 Type *IntTy = IntegerType::get(C&: Ctx, NumBits: DL.getTypeStoreSizeInBits(Ty: AccessTy));
313 Value *IntVal = ConstantInt::get(Ty: IntTy, V: 0);
314
315 for (unsigned I = 0; I < NumBytes; ++I) {
316 Value *Ptr = gepByteOffset(B, BasePtr: Src, ByteOffset: I);
317 LoadInst *LI = B.CreateLoad(Ty: I8Ty, Ptr);
318 LI->setAlignment(commonAlignment(A: Alignment, Offset: I));
319 buildAssignType(B, Ty: I8Ty, Arg: LI);
320 Value *Extended = B.CreateZExt(V: LI, DestTy: IntTy);
321 buildAssignType(B, Ty: IntTy, Arg: Extended);
322
323 if (I == 0) {
324 IntVal = Extended;
325 } else {
326 Value *Shifted = B.CreateShl(LHS: Extended, RHS: ConstantInt::get(Ty: IntTy, V: 8 * I));
327 buildAssignType(B, Ty: IntTy, Arg: Shifted);
328 IntVal = B.CreateOr(LHS: IntVal, RHS: Shifted);
329 }
330 buildAssignType(B, Ty: IntTy, Arg: IntVal);
331 }
332
333 Value *Result = storeIntToScalar(B, IntVal, ScalarTy: AccessTy);
334 if (Result != IntVal)
335 buildAssignType(B, Ty: AccessTy, Arg: Result);
336 return Result;
337 }
338
339 // Classifies a ptrcast reinterpretation: casted pointee matches the access
340 // type but differs from the original storage layout (e.g. i8 byte buffer as
341 // i32). ByteWise means multi-byte access must use per-byte i8 load/store.
342 bool shouldReinterpretByteWise(IRBuilder<> &B, Type *AccessTy,
343 Value *OriginalPtr) {
344 Type *OriginalElemTy = GR->findDeducedElementType(Val: OriginalPtr);
345 const DataLayout &DL = B.GetInsertBlock()->getModule()->getDataLayout();
346 if (OriginalElemTy && OriginalElemTy == Type::getInt8Ty(C&: B.getContext()) &&
347 AccessTy->isSingleValueType() && DL.getTypeStoreSize(Ty: AccessTy) > 1)
348 return true;
349
350 return false;
351 }
352
353 bool tryReinterpretLoad(IRBuilder<> &B, Type *AccessTy, Value *OriginalPtr,
354 Value *CastedPtr, LoadInst *IllegalLoad) {
355 Type *CastedElemTy = GR->findDeducedElementType(Val: CastedPtr);
356 if (!CastedElemTy || CastedElemTy != AccessTy)
357 return false;
358
359 Align Alignment = IllegalLoad->getAlign();
360 if (shouldReinterpretByteWise(B, AccessTy, OriginalPtr)) {
361 const DataLayout &DL = B.GetInsertBlock()->getModule()->getDataLayout();
362 Value *Loaded;
363 if (auto *VT = dyn_cast<FixedVectorType>(Val: AccessTy)) {
364 unsigned ElemSize = DL.getTypeStoreSize(Ty: VT->getElementType());
365 SmallVector<Value *, 4> LoadedElements;
366 for (unsigned I = 0; I < VT->getNumElements(); ++I) {
367 Value *ElemPtr = gepByteOffset(B, BasePtr: OriginalPtr, ByteOffset: I * ElemSize);
368 LoadedElements.push_back(Elt: loadScalarFromByteLayout(
369 B, AccessTy: VT->getElementType(), Src: ElemPtr,
370 Alignment: commonAlignment(A: Alignment, Offset: I * ElemSize)));
371 }
372 Loaded = buildVectorFromLoadedElements(B, TargetType: VT, LoadedElements);
373 } else {
374 Loaded = loadScalarFromByteLayout(B, AccessTy, Src: OriginalPtr, Alignment);
375 buildAssignType(B, Ty: AccessTy, Arg: Loaded);
376 }
377 GR->replaceAllUsesWith(Old: IllegalLoad, New: Loaded, /* DeleteOld= */ true);
378 DeadInstructions.push_back(x: IllegalLoad);
379 return true;
380 }
381
382 GR->buildAssignPtr(B, ElemTy: AccessTy, Arg: OriginalPtr);
383 LoadInst *LI = B.CreateLoad(Ty: AccessTy, Ptr: OriginalPtr);
384 LI->setAlignment(Alignment);
385 buildAssignType(B, Ty: AccessTy, Arg: LI);
386 GR->replaceAllUsesWith(Old: IllegalLoad, New: LI, /* DeleteOld= */ true);
387 DeadInstructions.push_back(x: IllegalLoad);
388 return true;
389 }
390
391 bool tryReinterpretStore(IRBuilder<> &B, Type *AccessTy, Value *OriginalPtr,
392 Value *CastedPtr, Value *StoreSrc, Align Alignment) {
393 Type *CastedElemTy = GR->findDeducedElementType(Val: CastedPtr);
394 if (!CastedElemTy || CastedElemTy != AccessTy)
395 return false;
396
397 if (shouldReinterpretByteWise(B, AccessTy, OriginalPtr)) {
398 const DataLayout &DL = B.GetInsertBlock()->getModule()->getDataLayout();
399 if (auto *VT = dyn_cast<FixedVectorType>(Val: StoreSrc->getType())) {
400 unsigned ElemSize = DL.getTypeStoreSize(Ty: VT->getElementType());
401 for (unsigned I = 0; I < VT->getNumElements(); ++I) {
402 Value *Elem = extractScalarFromVector(B, Vector: StoreSrc, Index: I);
403 Value *ElemPtr = gepByteOffset(B, BasePtr: OriginalPtr, ByteOffset: I * ElemSize);
404 storeScalarToByteLayout(B, Src: Elem, Dst: ElemPtr,
405 Alignment: commonAlignment(A: Alignment, Offset: I * ElemSize));
406 }
407 } else {
408 storeScalarToByteLayout(B, Src: StoreSrc, Dst: OriginalPtr, Alignment);
409 }
410 return true;
411 }
412
413 GR->buildAssignPtr(B, ElemTy: AccessTy, Arg: OriginalPtr);
414 StoreInst *SI = B.CreateStore(Val: StoreSrc, Ptr: OriginalPtr);
415 SI->setAlignment(Alignment);
416 return true;
417 }
418
419 // Builds a legalized load from a pointer, drilling down through
420 // memory layouts to find a compatible type. Load flags will be
421 // copied from |IllegalLoad|, which should be the load being legalized.
422 Value *buildLegalizedLoad(IRBuilder<> &B, Type *ElementType, Value *Source,
423 LoadInst *IllegalLoad, Value *CastedPtr) {
424 auto ResultOpt = getPointerToFirstCompatibleType(
425 B, BasePtr: Source, PointerType: IllegalLoad->getPointerOperandType(), TargetElemType: ElementType, IsInBounds: false);
426 if (!ResultOpt) {
427 if (tryReinterpretLoad(B, AccessTy: ElementType, OriginalPtr: Source, CastedPtr, IllegalLoad))
428 return nullptr;
429 llvm_unreachable("Failed to load from aggregate: "
430 "Could not find compatible memory layout.");
431 }
432 auto [GEP, CurrentTy] = *ResultOpt;
433
434 auto *SAT = dyn_cast<ArrayType>(Val: CurrentTy);
435 auto *SVT = dyn_cast<FixedVectorType>(Val: CurrentTy);
436 auto *DVT = dyn_cast<FixedVectorType>(Val: ElementType);
437 auto *MAT =
438 SAT ? dyn_cast<FixedVectorType>(Val: SAT->getElementType()) : nullptr;
439
440 if (ElementType == CurrentTy) {
441 LoadInst *LI = B.CreateLoad(Ty: ElementType, Ptr: GEP);
442 LI->setAlignment(IllegalLoad->getAlign());
443 buildAssignType(B, Ty: ElementType, Arg: LI);
444 return LI;
445 }
446 if (SVT && DVT)
447 return loadVectorFromVector(B, SourceType: SVT, TargetType: DVT, Source: GEP, OriginalAlign: IllegalLoad->getAlign());
448 if (SAT && DVT && SAT->getElementType() == DVT->getElementType())
449 return loadVectorFromArray(B, TargetType: DVT, Source: GEP, OriginalAlign: IllegalLoad->getAlign());
450 if (MAT && DVT && MAT->getElementType() == DVT->getElementType())
451 return loadVectorFromMatrixArray(B, TargetType: DVT, Source: GEP, ArrElemVecTy: MAT,
452 OriginalAlign: IllegalLoad->getAlign());
453
454 llvm_unreachable("Failed to load from aggregate.");
455 }
456
457 Value *
458 buildVectorFromLoadedElements(IRBuilder<> &B, FixedVectorType *TargetType,
459 SmallVector<Value *, 4> &LoadedElements) {
460 // <1 x T> shares the SPIR-V type with T, so emitting OpCompositeInsert on
461 // a scalar would be invalid. Bridge with spv_bitcast instead unless
462 // SPV_EXT_long_vector is available.
463 bool CanUseAnyVectorRank = TM.getSubtargetImpl()->canUseExtension(
464 E: SPIRV::Extension::SPV_EXT_long_vector);
465 if (TargetType->getNumElements() == 1 && !CanUseAnyVectorRank) {
466 Value *Scalar = LoadedElements[0];
467 Value *NewVector = B.CreateIntrinsic(
468 ID: Intrinsic::spv_bitcast, OverloadTypes: {TargetType, Scalar->getType()}, Args: {Scalar});
469 buildAssignType(B, Ty: TargetType, Arg: NewVector);
470 return NewVector;
471 }
472
473 // Build the vector from the loaded elements.
474 Value *NewVector = PoisonValue::get(T: TargetType);
475 buildAssignType(B, Ty: TargetType, Arg: NewVector);
476
477 for (unsigned I = 0, E = TargetType->getNumElements(); I < E; ++I) {
478 Value *Index = B.getInt32(C: I);
479 SmallVector<Type *, 4> Types = {TargetType, TargetType,
480 TargetType->getElementType(),
481 Index->getType()};
482 SmallVector<Value *> Args = {NewVector, LoadedElements[I], Index};
483 NewVector = B.CreateIntrinsic(ID: Intrinsic::spv_insertelt, OverloadTypes: {Types}, Args: {Args});
484 buildAssignType(B, Ty: TargetType, Arg: NewVector);
485 }
486 return NewVector;
487 }
488
489 // Loads elements from a matrix with an array of vector memory layout and
490 // constructs a vector.
491 Value *loadVectorFromMatrixArray(IRBuilder<> &B, FixedVectorType *TargetType,
492 Value *Source, FixedVectorType *ArrElemVecTy,
493 Align OriginalAlign) {
494 Type *TargetElemTy = TargetType->getElementType();
495 unsigned ScalarsPerArrayElement = ArrElemVecTy->getNumElements();
496 const DataLayout &DL = B.GetInsertBlock()->getModule()->getDataLayout();
497 uint64_t ArrElemVecSize = DL.getTypeAllocSize(Ty: ArrElemVecTy);
498 // Load each element of the array.
499 SmallVector<Value *, 4> LoadedElements;
500 std::array<Type *, 2> Types = {Source->getType(), Source->getType()};
501 for (unsigned I = 0, E = TargetType->getNumElements(); I < E; ++I) {
502 unsigned ArrayIndex = I / ScalarsPerArrayElement;
503 unsigned ElementIndexInArrayElem = I % ScalarsPerArrayElement;
504 // Create a GEP to access the i-th element of the array.
505 std::array<Value *, 4> Args = {
506 B.getInt1(/*Inbounds=*/V: false), Source, B.getInt32(C: 0),
507 ConstantInt::get(Ty: B.getInt32Ty(), V: ArrayIndex)};
508 auto *ElementPtr = B.CreateIntrinsic(ID: Intrinsic::spv_gep, OverloadTypes: {Types}, Args: {Args});
509 GR->buildAssignPtr(B, ElemTy: ArrElemVecTy, Arg: ElementPtr);
510 LoadInst *LoadVec = B.CreateLoad(Ty: ArrElemVecTy, Ptr: ElementPtr);
511 LoadVec->setAlignment(
512 commonAlignment(A: OriginalAlign, Offset: ArrayIndex * ArrElemVecSize));
513 buildAssignType(B, Ty: ArrElemVecTy, Arg: LoadVec);
514 LoadedElements.push_back(Elt: makeExtractElement(B, ElementType: TargetElemTy, Vector: LoadVec,
515 Index: ElementIndexInArrayElem));
516 }
517 return buildVectorFromLoadedElements(B, TargetType, LoadedElements);
518 }
519
520 // Loads elements from an array and constructs a vector.
521 Value *loadVectorFromArray(IRBuilder<> &B, FixedVectorType *TargetType,
522 Value *Source, Align OriginalAlign) {
523 // Load each element of the array.
524 SmallVector<Value *, 4> LoadedElements;
525 std::array<Type *, 2> Types = {Source->getType(), Source->getType()};
526 const DataLayout &DL = B.GetInsertBlock()->getModule()->getDataLayout();
527 uint64_t ElemSize = DL.getTypeAllocSize(Ty: TargetType->getElementType());
528 for (unsigned I = 0, E = TargetType->getNumElements(); I < E; ++I) {
529 // Create a GEP to access the i-th element of the array.
530 std::array<Value *, 4> Args = {B.getInt1(/*Inbounds=*/V: false), Source,
531 B.getInt32(C: 0),
532 ConstantInt::get(Ty: B.getInt32Ty(), V: I)};
533 auto *ElementPtr = B.CreateIntrinsic(ID: Intrinsic::spv_gep, OverloadTypes: {Types}, Args: {Args});
534 GR->buildAssignPtr(B, ElemTy: TargetType->getElementType(), Arg: ElementPtr);
535
536 // Load the value from the element pointer.
537 LoadInst *Load = B.CreateLoad(Ty: TargetType->getElementType(), Ptr: ElementPtr);
538 Load->setAlignment(commonAlignment(A: OriginalAlign, Offset: I * ElemSize));
539 buildAssignType(B, Ty: TargetType->getElementType(), Arg: Load);
540 LoadedElements.push_back(Elt: Load);
541 }
542 return buildVectorFromLoadedElements(B, TargetType, LoadedElements);
543 }
544
545 // Stores elements from a vector into a matrix (an array of vectors).
546 void storeMatrixArrayFromVector(IRBuilder<> &B, Value *SrcVector,
547 Value *DstArrayPtr, ArrayType *ArrTy,
548 Align Alignment) {
549 auto *SrcVecTy = cast<FixedVectorType>(Val: SrcVector->getType());
550 auto *ArrElemVecTy = cast<FixedVectorType>(Val: ArrTy->getElementType());
551 Type *ElemTy = ArrElemVecTy->getElementType();
552 unsigned ScalarsPerArrayElement = ArrElemVecTy->getNumElements();
553 unsigned SrcNumElements = SrcVecTy->getNumElements();
554 assert(
555 SrcNumElements % ScalarsPerArrayElement == 0 &&
556 "Source vector size must be a multiple of array element vector size");
557
558 std::array<Type *, 2> Types = {DstArrayPtr->getType(),
559 DstArrayPtr->getType()};
560 const DataLayout &DL = B.GetInsertBlock()->getModule()->getDataLayout();
561 uint64_t ArrElemVecSize = DL.getTypeAllocSize(Ty: ArrElemVecTy);
562
563 for (unsigned I = 0; I < SrcNumElements; I += ScalarsPerArrayElement) {
564 unsigned ArrayIndex = I / ScalarsPerArrayElement;
565 // Create a GEP to access the array element.
566 std::array<Value *, 4> Args = {
567 B.getInt1(/*Inbounds=*/V: false), DstArrayPtr, B.getInt32(C: 0),
568 ConstantInt::get(Ty: B.getInt32Ty(), V: ArrayIndex)};
569 auto *ElementPtr = B.CreateIntrinsic(ID: Intrinsic::spv_gep, OverloadTypes: {Types}, Args: {Args});
570 GR->buildAssignPtr(B, ElemTy: ArrElemVecTy, Arg: ElementPtr);
571
572 // Extract scalar elements from the source vector for this array slot.
573 SmallVector<Value *, 4> Elements;
574 for (unsigned J = 0; J < ScalarsPerArrayElement; ++J)
575 Elements.push_back(Elt: makeExtractElement(B, ElementType: ElemTy, Vector: SrcVector, Index: I + J));
576
577 // Build a vector from the extracted elements and store it.
578 Value *Vec = buildVectorFromLoadedElements(B, TargetType: ArrElemVecTy, LoadedElements&: Elements);
579 StoreInst *SI = B.CreateStore(Val: Vec, Ptr: ElementPtr);
580 SI->setAlignment(commonAlignment(A: Alignment, Offset: ArrayIndex * ArrElemVecSize));
581 }
582 }
583
584 // Stores elements from a vector into an array.
585 void storeArrayFromVector(IRBuilder<> &B, Value *SrcVector,
586 Value *DstArrayPtr, ArrayType *ArrTy,
587 Align Alignment) {
588 auto *VecTy = cast<FixedVectorType>(Val: SrcVector->getType());
589 Type *ElemTy = ArrTy->getElementType();
590
591 // Ensure the element types of the array and vector are the same.
592 assert(VecTy->getElementType() == ElemTy &&
593 "Element types of array and vector must be the same.");
594 std::array<Type *, 2> Types = {DstArrayPtr->getType(),
595 DstArrayPtr->getType()};
596 const DataLayout &DL = B.GetInsertBlock()->getModule()->getDataLayout();
597 uint64_t ElemSize = DL.getTypeAllocSize(Ty: ElemTy);
598
599 for (unsigned I = 0, E = VecTy->getNumElements(); I < E; ++I) {
600 // Create a GEP to access the i-th element of the array.
601 std::array<Value *, 4> Args = {B.getInt1(/*Inbounds=*/V: false), DstArrayPtr,
602 B.getInt32(C: 0),
603 ConstantInt::get(Ty: B.getInt32Ty(), V: I)};
604 auto *ElementPtr = B.CreateIntrinsic(ID: Intrinsic::spv_gep, OverloadTypes: {Types}, Args: {Args});
605 GR->buildAssignPtr(B, ElemTy, Arg: ElementPtr);
606
607 // Extract the element from the vector and store it.
608 bool CanUseAnyVectorRank = TM.getSubtargetImpl()->canUseExtension(
609 E: SPIRV::Extension::SPV_EXT_long_vector);
610 Value *Element = (E == 1 && !CanUseAnyVectorRank)
611 ? SrcVector
612 : makeExtractElement(B, ElementType: ElemTy, Vector: SrcVector, Index: I);
613 StoreInst *SI = B.CreateStore(Val: Element, Ptr: ElementPtr);
614 SI->setAlignment(commonAlignment(A: Alignment, Offset: I * ElemSize));
615 }
616 }
617
618 // Replaces the load instruction to get rid of the ptrcast used as source
619 // operand.
620 void transformLoad(IRBuilder<> &B, LoadInst *LI, Value *CastedOperand,
621 Value *OriginalOperand) {
622 Type *ToTy = GR->findDeducedElementType(Val: CastedOperand);
623 B.SetInsertPoint(LI);
624
625 Value *Output =
626 buildLegalizedLoad(B, ElementType: ToTy, Source: OriginalOperand, IllegalLoad: LI, CastedPtr: CastedOperand);
627 if (!Output)
628 return;
629
630 GR->replaceAllUsesWith(Old: LI, New: Output, /* DeleteOld= */ true);
631 DeadInstructions.push_back(x: LI);
632 }
633
634 // Creates an spv_insertelt instruction (equivalent to llvm's insertelement).
635 Value *makeInsertElement(IRBuilder<> &B, Value *Vector, Value *Element,
636 unsigned Index) {
637 Type *Int32Ty = Type::getInt32Ty(C&: B.getContext());
638 SmallVector<Type *, 4> Types = {Vector->getType(), Vector->getType(),
639 Element->getType(), Int32Ty};
640 SmallVector<Value *> Args = {Vector, Element, B.getInt32(C: Index)};
641 Value *NewI = B.CreateIntrinsic(ID: Intrinsic::spv_insertelt, OverloadTypes: {Types}, Args: {Args});
642 buildAssignType(B, Ty: Vector->getType(), Arg: NewI);
643 return NewI;
644 }
645
646 // Creates an spv_extractelt instruction (equivalent to llvm's
647 // extractelement).
648 Value *makeExtractElement(IRBuilder<> &B, Type *ElementType, Value *Vector,
649 unsigned Index) {
650 Type *Int32Ty = Type::getInt32Ty(C&: B.getContext());
651 SmallVector<Type *, 3> Types = {ElementType, Vector->getType(), Int32Ty};
652 SmallVector<Value *> Args = {Vector, B.getInt32(C: Index)};
653 Value *NewI = B.CreateIntrinsic(ID: Intrinsic::spv_extractelt, OverloadTypes: {Types}, Args: {Args});
654 buildAssignType(B, Ty: ElementType, Arg: NewI);
655 return NewI;
656 }
657
658 // Extracts scalar element |Index| from |Vector|. A <1 x T> vector shares its
659 // SPIR-V type with the scalar T, so a plain extractelement would be invalid;
660 // bridge it with spv_bitcast instead.
661 Value *extractScalarFromVector(IRBuilder<> &B, Value *Vector,
662 unsigned Index) {
663 auto *VecTy = cast<FixedVectorType>(Val: Vector->getType());
664 Type *ElemTy = VecTy->getElementType();
665 if (VecTy->getNumElements() == 1) {
666 Value *Scalar =
667 B.CreateIntrinsic(ID: Intrinsic::spv_bitcast, OverloadTypes: {ElemTy, VecTy}, Args: {Vector});
668 buildAssignType(B, Ty: ElemTy, Arg: Scalar);
669 return Scalar;
670 }
671 return makeExtractElement(B, ElementType: ElemTy, Vector, Index);
672 }
673
674 // Stores the given Src vector operand into the Dst vector, adjusting the size
675 // if required.
676 Value *storeVectorFromVector(IRBuilder<> &B, Value *Src, Value *Dst,
677 Align Alignment) {
678 FixedVectorType *SrcType = cast<FixedVectorType>(Val: Src->getType());
679 FixedVectorType *DstType =
680 cast<FixedVectorType>(Val: GR->findDeducedElementType(Val: Dst));
681 auto dstNumElements = DstType->getNumElements();
682 auto srcNumElements = SrcType->getNumElements();
683
684 // if the element type differs, it is a bitcast.
685 if (DstType->getElementType() != SrcType->getElementType()) {
686 // Support bitcast between vectors of different sizes only if
687 // the total bitwidth is the same.
688 [[maybe_unused]] auto dstBitWidth =
689 DstType->getElementType()->getScalarSizeInBits() * dstNumElements;
690 [[maybe_unused]] auto srcBitWidth =
691 SrcType->getElementType()->getScalarSizeInBits() * srcNumElements;
692 assert(dstBitWidth == srcBitWidth &&
693 "Unsupported bitcast between vectors of different sizes.");
694
695 Src =
696 B.CreateIntrinsic(ID: Intrinsic::spv_bitcast, OverloadTypes: {DstType, SrcType}, Args: {Src});
697 buildAssignType(B, Ty: DstType, Arg: Src);
698 SrcType = DstType;
699
700 StoreInst *SI = B.CreateStore(Val: Src, Ptr: Dst);
701 SI->setAlignment(Alignment);
702 return SI;
703 }
704
705 assert(DstType->getNumElements() >= SrcType->getNumElements());
706 LoadInst *LI = B.CreateLoad(Ty: DstType, Ptr: Dst);
707 LI->setAlignment(Alignment);
708 Value *OldValues = LI;
709 buildAssignType(B, Ty: OldValues->getType(), Arg: OldValues);
710 Value *NewValues = Src;
711
712 for (unsigned I = 0; I < SrcType->getNumElements(); ++I) {
713 Value *Element =
714 makeExtractElement(B, ElementType: SrcType->getElementType(), Vector: NewValues, Index: I);
715 OldValues = makeInsertElement(B, Vector: OldValues, Element, Index: I);
716 }
717
718 StoreInst *SI = B.CreateStore(Val: OldValues, Ptr: Dst);
719 SI->setAlignment(Alignment);
720 return SI;
721 }
722
723 // Builds a legalized store to a pointer, drilling down through
724 // memory layouts to find a compatible type.
725 void buildLegalizedStore(IRBuilder<> &B, Value *Src, Value *Dst,
726 Align Alignment, Value *CastedPtr,
727 Instruction *IllegalStore) {
728 auto ResultOpt = getPointerToFirstCompatibleType(B, BasePtr: Dst, PointerType: Dst->getType(),
729 TargetElemType: Src->getType(), IsInBounds: true);
730 if (!ResultOpt) {
731 if (tryReinterpretStore(B, AccessTy: Src->getType(), OriginalPtr: Dst, CastedPtr, StoreSrc: Src,
732 Alignment))
733 return;
734 llvm_unreachable("Failed to store to aggregate: "
735 "Could not find compatible memory layout.");
736 }
737 auto [GEP, CurrentTy] = *ResultOpt;
738
739 auto *DAT = dyn_cast<ArrayType>(Val: CurrentTy);
740 auto *DVT = dyn_cast<FixedVectorType>(Val: CurrentTy);
741 auto *SVT = dyn_cast<FixedVectorType>(Val: Src->getType());
742 auto *DMAT =
743 DAT ? dyn_cast<FixedVectorType>(Val: DAT->getElementType()) : nullptr;
744
745 if (Src->getType() == CurrentTy) {
746 StoreInst *SI = B.CreateStore(Val: Src, Ptr: GEP);
747 SI->setAlignment(Alignment);
748 return;
749 }
750 if (DVT && SVT) {
751 storeVectorFromVector(B, Src, Dst: GEP, Alignment);
752 return;
753 }
754 if (DAT && SVT && SVT->getElementType() == DAT->getElementType()) {
755 storeArrayFromVector(B, SrcVector: Src, DstArrayPtr: GEP, ArrTy: DAT, Alignment);
756 return;
757 }
758 if (DMAT && SVT && DMAT->getElementType() == SVT->getElementType()) {
759 storeMatrixArrayFromVector(B, SrcVector: Src, DstArrayPtr: GEP, ArrTy: DAT, Alignment);
760 return;
761 }
762
763 llvm_unreachable("Failed to store to aggregate.");
764 }
765
766 // Transforms a store instruction (or SPV intrinsic) using a ptrcast as
767 // operand into a valid logical SPIR-V store with no ptrcast.
768 void transformStore(IRBuilder<> &B, Instruction *IllegalStore, Value *Src,
769 Value *Dst, Value *CastedOperand, Align Alignment) {
770 B.SetInsertPoint(IllegalStore);
771 buildLegalizedStore(B, Src, Dst, Alignment, CastedPtr: CastedOperand, IllegalStore);
772 DeadInstructions.push_back(x: IllegalStore);
773 }
774
775 void legalizePointerCast(IntrinsicInst *II) {
776 Value *CastedOperand = II;
777 Value *OriginalOperand = II->getOperand(i_nocapture: 0);
778
779 IRBuilder<> B(II->getContext());
780 std::vector<Value *> Users;
781 for (Use &U : II->uses())
782 Users.push_back(x: U.getUser());
783
784 for (Value *User : Users) {
785 if (LoadInst *LI = dyn_cast<LoadInst>(Val: User)) {
786 transformLoad(B, LI, CastedOperand, OriginalOperand);
787 continue;
788 }
789
790 if (StoreInst *SI = dyn_cast<StoreInst>(Val: User)) {
791 transformStore(B, IllegalStore: SI, Src: SI->getValueOperand(), Dst: OriginalOperand,
792 CastedOperand, Alignment: SI->getAlign());
793 continue;
794 }
795
796 if (IntrinsicInst *Intrin = dyn_cast<IntrinsicInst>(Val: User)) {
797 if (Intrin->getIntrinsicID() == Intrinsic::spv_assign_ptr_type) {
798 DeadInstructions.push_back(x: Intrin);
799 continue;
800 }
801
802 if (Intrin->getIntrinsicID() == Intrinsic::spv_gep) {
803 GR->replaceAllUsesWith(Old: CastedOperand, New: OriginalOperand,
804 /* DeleteOld= */ false);
805 continue;
806 }
807
808 if (Intrin->getIntrinsicID() == Intrinsic::spv_store) {
809 Align Alignment;
810 if (ConstantInt *C = dyn_cast<ConstantInt>(Val: Intrin->getOperand(i_nocapture: 3)))
811 Alignment = Align(C->getZExtValue());
812 transformStore(B, IllegalStore: Intrin, Src: Intrin->getArgOperand(i: 0), Dst: OriginalOperand,
813 CastedOperand, Alignment);
814 continue;
815 }
816 }
817
818 llvm_unreachable("Unsupported ptrcast user. Please fix.");
819 }
820
821 DeadInstructions.push_back(x: II);
822 }
823
824public:
825 SPIRVLegalizePointerCastImpl(const SPIRVTargetMachine &TM) : TM(TM) {}
826
827 bool run(Function &F) {
828 const SPIRVSubtarget &ST = TM.getSubtarget<SPIRVSubtarget>(F);
829 GR = ST.getSPIRVGlobalRegistry();
830 DeadInstructions.clear();
831
832 std::vector<IntrinsicInst *> WorkList;
833 for (auto &BB : F) {
834 for (auto &I : BB) {
835 auto *II = dyn_cast<IntrinsicInst>(Val: &I);
836 if (II && II->getIntrinsicID() == Intrinsic::spv_ptrcast)
837 WorkList.push_back(x: II);
838 }
839 }
840
841 for (IntrinsicInst *II : WorkList)
842 legalizePointerCast(II);
843
844 for (Instruction *I : DeadInstructions)
845 I->eraseFromParent();
846
847 return DeadInstructions.size() != 0;
848 }
849
850private:
851 const SPIRVTargetMachine &TM;
852 SPIRVGlobalRegistry *GR = nullptr;
853 std::vector<Instruction *> DeadInstructions;
854};
855
856class SPIRVLegalizePointerCastLegacy : public FunctionPass {
857public:
858 static char ID;
859 SPIRVLegalizePointerCastLegacy(const SPIRVTargetMachine &TM)
860 : FunctionPass(ID), TM(TM) {}
861
862 bool runOnFunction(Function &F) override {
863 return SPIRVLegalizePointerCastImpl(TM).run(F);
864 }
865
866private:
867 const SPIRVTargetMachine &TM;
868};
869} // namespace
870
871PreservedAnalyses
872SPIRVLegalizePointerCastPass::run(Function &F, FunctionAnalysisManager &AM) {
873 return SPIRVLegalizePointerCastImpl(TM).run(F) ? PreservedAnalyses::none()
874 : PreservedAnalyses::all();
875}
876
877char SPIRVLegalizePointerCastLegacy::ID = 0;
878INITIALIZE_PASS(SPIRVLegalizePointerCastLegacy, "spirv-legalize-pointer-cast",
879 "SPIRV legalize pointer cast pass", false, false)
880
881FunctionPass *llvm::createSPIRVLegalizePointerCastPass(SPIRVTargetMachine *TM) {
882 return new SPIRVLegalizePointerCastLegacy(*TM);
883}
884