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 "SPIRVLegalizePointerCast.h"
46#include "SPIRV.h"
47#include "SPIRVSubtarget.h"
48#include "SPIRVTargetMachine.h"
49#include "SPIRVUtils.h"
50#include "llvm/IR/IRBuilder.h"
51#include "llvm/IR/IntrinsicInst.h"
52#include "llvm/IR/Intrinsics.h"
53#include "llvm/IR/IntrinsicsSPIRV.h"
54#include "llvm/Transforms/Utils/Cloning.h"
55#include "llvm/Transforms/Utils/LowerMemIntrinsics.h"
56
57using namespace llvm;
58
59namespace {
60class SPIRVLegalizePointerCastImpl {
61
62 // Builds the `spv_assign_type` assigning |Ty| to |Value| at the current
63 // builder position.
64 void buildAssignType(IRBuilder<> &B, Type *Ty, Value *Arg) {
65 Value *OfType = PoisonValue::get(T: Ty);
66 CallInst *AssignCI = buildIntrWithMD(IntrID: Intrinsic::spv_assign_type,
67 Types: {Arg->getType()}, Arg: OfType, Arg2: Arg, Imms: {}, B);
68 GR->addAssignPtrTypeInstr(Val: Arg, AssignPtrTyCI: AssignCI);
69 }
70
71 static FixedVectorType *makeVectorFromTotalBits(Type *ElemTy,
72 TypeSize TotalBits) {
73 unsigned ElemBits = ElemTy->getScalarSizeInBits();
74 assert(ElemBits && TotalBits % ElemBits == 0 &&
75 "TotalBits must be divisible by element bit size");
76 return FixedVectorType::get(ElementType: ElemTy, NumElts: TotalBits / ElemBits);
77 }
78
79 Value *resizeVectorBitsWithShuffle(IRBuilder<> &B, Value *V,
80 FixedVectorType *DstTy) {
81 auto *SrcTy = cast<FixedVectorType>(Val: V->getType());
82 assert(SrcTy->getElementType() == DstTy->getElementType() &&
83 "shuffle resize expects identical element types");
84
85 const unsigned NumNeeded = DstTy->getNumElements();
86 const unsigned NumSource = SrcTy->getNumElements();
87
88 SmallVector<int> Mask(NumNeeded);
89 for (unsigned I = 0; I < NumNeeded; ++I)
90 Mask[I] = (I < NumSource) ? static_cast<int>(I) : -1;
91
92 Value *Resized = B.CreateShuffleVector(V1: V, V2: V, Mask);
93 buildAssignType(B, Ty: DstTy, Arg: Resized);
94 return Resized;
95 }
96
97 // Loads parts of the vector of type |SourceType| from the pointer |Source|
98 // and create a new vector of type |TargetType|. |TargetType| must be a vector
99 // type.
100 // Returns the loaded value.
101 Value *loadVectorFromVector(IRBuilder<> &B, FixedVectorType *SourceType,
102 FixedVectorType *TargetType, Value *Source,
103 Align OriginalAlign) {
104 LoadInst *NewLoad = B.CreateLoad(Ty: SourceType, Ptr: Source);
105 NewLoad->setAlignment(OriginalAlign);
106 buildAssignType(B, Ty: SourceType, Arg: NewLoad);
107 Value *AssignValue = NewLoad;
108 if (TargetType->getElementType() != SourceType->getElementType()) {
109 const DataLayout &DL = B.GetInsertBlock()->getModule()->getDataLayout();
110 TypeSize TargetTypeSize = DL.getTypeSizeInBits(Ty: TargetType);
111 TypeSize SourceTypeSize = DL.getTypeSizeInBits(Ty: SourceType);
112
113 Value *BitcastSrcVal = NewLoad;
114 FixedVectorType *BitcastSrcTy =
115 cast<FixedVectorType>(Val: BitcastSrcVal->getType());
116 FixedVectorType *BitcastDstTy = TargetType;
117
118 if (TargetTypeSize != SourceTypeSize) {
119 unsigned TargetElemBits =
120 TargetType->getElementType()->getScalarSizeInBits();
121 if (SourceTypeSize % TargetElemBits == 0) {
122 // No Resize needed. Same total bits as source, but use target element
123 // type.
124 BitcastDstTy = makeVectorFromTotalBits(ElemTy: TargetType->getElementType(),
125 TotalBits: SourceTypeSize);
126 } else {
127 // Resize source to target total bitwidth using source element type.
128 BitcastSrcTy = makeVectorFromTotalBits(ElemTy: SourceType->getElementType(),
129 TotalBits: TargetTypeSize);
130 BitcastSrcVal = resizeVectorBitsWithShuffle(B, V: NewLoad, DstTy: BitcastSrcTy);
131 }
132 }
133 AssignValue =
134 B.CreateIntrinsic(ID: Intrinsic::spv_bitcast,
135 OverloadTypes: {BitcastDstTy, BitcastSrcTy}, Args: {BitcastSrcVal});
136 buildAssignType(B, Ty: BitcastDstTy, Arg: AssignValue);
137 if (BitcastDstTy == TargetType)
138 return AssignValue;
139 }
140
141 assert(TargetType->getNumElements() < SourceType->getNumElements());
142 SmallVector<int> Mask(/* Size= */ TargetType->getNumElements());
143 for (unsigned I = 0; I < TargetType->getNumElements(); ++I)
144 Mask[I] = I;
145 Value *Output = B.CreateShuffleVector(V1: AssignValue, V2: AssignValue, Mask);
146 buildAssignType(B, Ty: TargetType, Arg: Output);
147 return Output;
148 }
149
150 // Returns true if |FromTy| has a memory layout compatible with loading or
151 // storing |ToTy|.
152 bool isCompatibleMemoryLayout(Type *ToTy, Type *FromTy) {
153 if (ToTy == FromTy)
154 return true;
155 auto *SVT = dyn_cast<FixedVectorType>(Val: FromTy);
156 auto *DVT = dyn_cast<FixedVectorType>(Val: ToTy);
157 if (SVT && DVT)
158 return true;
159 auto *SAT = dyn_cast<ArrayType>(Val: FromTy);
160 if (SAT && DVT) {
161 if (SAT->getElementType() == DVT->getElementType())
162 return true;
163 if (auto *MAT = dyn_cast<FixedVectorType>(Val: SAT->getElementType()))
164 if (MAT->getElementType() == DVT->getElementType())
165 return true;
166 }
167 return false;
168 }
169
170 // Traverses the aggregate type to find the first sub-type that matches
171 // the TargetElemType's memory layout, optionally emitting a GEP intrinsic.
172 std::optional<std::pair<Value *, Type *>>
173 getPointerToFirstCompatibleType(IRBuilder<> &B, Value *BasePtr,
174 Type *PointerType, Type *TargetElemType,
175 bool IsInBounds) {
176 Type *CurrentTy = GR->findDeducedElementType(Val: BasePtr);
177 assert(CurrentTy && "Could not deduce aggregate type");
178 SmallVector<Value *, 8> Args{/* isInBounds= */ B.getInt1(V: IsInBounds),
179 BasePtr};
180 Args.push_back(Elt: B.getInt32(C: 0)); // Pointer offset
181
182 while (!isCompatibleMemoryLayout(ToTy: TargetElemType, FromTy: CurrentTy)) {
183 if (auto *ST = dyn_cast<StructType>(Val: CurrentTy)) {
184 if (ST->getNumElements() == 0)
185 return std::nullopt;
186 CurrentTy = ST->getTypeAtIndex(N: 0u);
187 } else if (auto *AT = dyn_cast<ArrayType>(Val: CurrentTy)) {
188 CurrentTy = AT->getElementType();
189 } else if (auto *VT = dyn_cast<FixedVectorType>(Val: CurrentTy)) {
190 CurrentTy = VT->getElementType();
191 } else {
192 return std::nullopt;
193 }
194 Args.push_back(Elt: B.getInt32(C: 0));
195 }
196
197 Value *GEP = BasePtr;
198 if (Args.size() > 3) {
199 std::array<Type *, 2> Types = {PointerType, BasePtr->getType()};
200 GEP = B.CreateIntrinsic(ID: Intrinsic::spv_gep, OverloadTypes: {Types}, Args: {Args});
201 GR->buildAssignPtr(B, ElemTy: CurrentTy, Arg: GEP);
202 }
203
204 return std::make_pair(x&: GEP, y&: CurrentTy);
205 }
206
207 // Builds a legalized load from a pointer, drilling down through
208 // memory layouts to find a compatible type. Load flags will be
209 // copied from |BadLoad|, which should be the load being legalized.
210 Value *buildLegalizedLoad(IRBuilder<> &B, Type *ElementType, Value *Source,
211 LoadInst *BadLoad) {
212 auto ResultOpt = getPointerToFirstCompatibleType(
213 B, BasePtr: Source, PointerType: BadLoad->getPointerOperandType(), TargetElemType: ElementType, IsInBounds: false);
214 assert(ResultOpt && "Failed to load from aggregate: "
215 "Could not find compatible memory layout.");
216 auto [GEP, CurrentTy] = *ResultOpt;
217
218 auto *SAT = dyn_cast<ArrayType>(Val: CurrentTy);
219 auto *SVT = dyn_cast<FixedVectorType>(Val: CurrentTy);
220 auto *DVT = dyn_cast<FixedVectorType>(Val: ElementType);
221 auto *MAT =
222 SAT ? dyn_cast<FixedVectorType>(Val: SAT->getElementType()) : nullptr;
223
224 if (ElementType == CurrentTy) {
225 LoadInst *LI = B.CreateLoad(Ty: ElementType, Ptr: GEP);
226 LI->setAlignment(BadLoad->getAlign());
227 buildAssignType(B, Ty: ElementType, Arg: LI);
228 return LI;
229 }
230 if (SVT && DVT)
231 return loadVectorFromVector(B, SourceType: SVT, TargetType: DVT, Source: GEP, OriginalAlign: BadLoad->getAlign());
232 if (SAT && DVT && SAT->getElementType() == DVT->getElementType())
233 return loadVectorFromArray(B, TargetType: DVT, Source: GEP, OriginalAlign: BadLoad->getAlign());
234 if (MAT && DVT && MAT->getElementType() == DVT->getElementType())
235 return loadVectorFromMatrixArray(B, TargetType: DVT, Source: GEP, ArrElemVecTy: MAT, OriginalAlign: BadLoad->getAlign());
236
237 llvm_unreachable("Failed to load from aggregate.");
238 }
239 Value *
240 buildVectorFromLoadedElements(IRBuilder<> &B, FixedVectorType *TargetType,
241 SmallVector<Value *, 4> &LoadedElements) {
242 // <1 x T> shares the SPIR-V type with T, so emitting OpCompositeInsert on
243 // a scalar would be invalid. Bridge with spv_bitcast instead.
244 if (TargetType->getNumElements() == 1) {
245 Value *Scalar = LoadedElements[0];
246 Value *NewVector = B.CreateIntrinsic(
247 ID: Intrinsic::spv_bitcast, OverloadTypes: {TargetType, Scalar->getType()}, Args: {Scalar});
248 buildAssignType(B, Ty: TargetType, Arg: NewVector);
249 return NewVector;
250 }
251
252 // Build the vector from the loaded elements.
253 Value *NewVector = PoisonValue::get(T: TargetType);
254 buildAssignType(B, Ty: TargetType, Arg: NewVector);
255
256 for (unsigned I = 0, E = TargetType->getNumElements(); I < E; ++I) {
257 Value *Index = B.getInt32(C: I);
258 SmallVector<Type *, 4> Types = {TargetType, TargetType,
259 TargetType->getElementType(),
260 Index->getType()};
261 SmallVector<Value *> Args = {NewVector, LoadedElements[I], Index};
262 NewVector = B.CreateIntrinsic(ID: Intrinsic::spv_insertelt, OverloadTypes: {Types}, Args: {Args});
263 buildAssignType(B, Ty: TargetType, Arg: NewVector);
264 }
265 return NewVector;
266 }
267
268 // Loads elements from a matrix with an array of vector memory layout and
269 // constructs a vector.
270 Value *loadVectorFromMatrixArray(IRBuilder<> &B, FixedVectorType *TargetType,
271 Value *Source, FixedVectorType *ArrElemVecTy,
272 Align OriginalAlign) {
273 Type *TargetElemTy = TargetType->getElementType();
274 unsigned ScalarsPerArrayElement = ArrElemVecTy->getNumElements();
275 const DataLayout &DL = B.GetInsertBlock()->getModule()->getDataLayout();
276 uint64_t ArrElemVecSize = DL.getTypeAllocSize(Ty: ArrElemVecTy);
277 // Load each element of the array.
278 SmallVector<Value *, 4> LoadedElements;
279 std::array<Type *, 2> Types = {Source->getType(), Source->getType()};
280 for (unsigned I = 0, E = TargetType->getNumElements(); I < E; ++I) {
281 unsigned ArrayIndex = I / ScalarsPerArrayElement;
282 unsigned ElementIndexInArrayElem = I % ScalarsPerArrayElement;
283 // Create a GEP to access the i-th element of the array.
284 std::array<Value *, 4> Args = {
285 B.getInt1(/*Inbounds=*/V: false), Source, B.getInt32(C: 0),
286 ConstantInt::get(Ty: B.getInt32Ty(), V: ArrayIndex)};
287 auto *ElementPtr = B.CreateIntrinsic(ID: Intrinsic::spv_gep, OverloadTypes: {Types}, Args: {Args});
288 GR->buildAssignPtr(B, ElemTy: ArrElemVecTy, Arg: ElementPtr);
289 LoadInst *LoadVec = B.CreateLoad(Ty: ArrElemVecTy, Ptr: ElementPtr);
290 LoadVec->setAlignment(
291 commonAlignment(A: OriginalAlign, Offset: ArrayIndex * ArrElemVecSize));
292 buildAssignType(B, Ty: ArrElemVecTy, Arg: LoadVec);
293 LoadedElements.push_back(Elt: makeExtractElement(B, ElementType: TargetElemTy, Vector: LoadVec,
294 Index: ElementIndexInArrayElem));
295 }
296 return buildVectorFromLoadedElements(B, TargetType, LoadedElements);
297 }
298 // Loads elements from an array and constructs a vector.
299 Value *loadVectorFromArray(IRBuilder<> &B, FixedVectorType *TargetType,
300 Value *Source, Align OriginalAlign) {
301 // Load each element of the array.
302 SmallVector<Value *, 4> LoadedElements;
303 std::array<Type *, 2> Types = {Source->getType(), Source->getType()};
304 const DataLayout &DL = B.GetInsertBlock()->getModule()->getDataLayout();
305 uint64_t ElemSize = DL.getTypeAllocSize(Ty: TargetType->getElementType());
306 for (unsigned I = 0, E = TargetType->getNumElements(); I < E; ++I) {
307 // Create a GEP to access the i-th element of the array.
308 std::array<Value *, 4> Args = {B.getInt1(/*Inbounds=*/V: false), Source,
309 B.getInt32(C: 0),
310 ConstantInt::get(Ty: B.getInt32Ty(), V: I)};
311 auto *ElementPtr = B.CreateIntrinsic(ID: Intrinsic::spv_gep, OverloadTypes: {Types}, Args: {Args});
312 GR->buildAssignPtr(B, ElemTy: TargetType->getElementType(), Arg: ElementPtr);
313
314 // Load the value from the element pointer.
315 LoadInst *Load = B.CreateLoad(Ty: TargetType->getElementType(), Ptr: ElementPtr);
316 Load->setAlignment(commonAlignment(A: OriginalAlign, Offset: I * ElemSize));
317 buildAssignType(B, Ty: TargetType->getElementType(), Arg: Load);
318 LoadedElements.push_back(Elt: Load);
319 }
320 return buildVectorFromLoadedElements(B, TargetType, LoadedElements);
321 }
322
323 // Stores elements from a vector into a matrix (an array of vectors).
324 void storeMatrixArrayFromVector(IRBuilder<> &B, Value *SrcVector,
325 Value *DstArrayPtr, ArrayType *ArrTy,
326 Align Alignment) {
327 auto *SrcVecTy = cast<FixedVectorType>(Val: SrcVector->getType());
328 auto *ArrElemVecTy = cast<FixedVectorType>(Val: ArrTy->getElementType());
329 Type *ElemTy = ArrElemVecTy->getElementType();
330 unsigned ScalarsPerArrayElement = ArrElemVecTy->getNumElements();
331 unsigned SrcNumElements = SrcVecTy->getNumElements();
332 assert(
333 SrcNumElements % ScalarsPerArrayElement == 0 &&
334 "Source vector size must be a multiple of array element vector size");
335
336 std::array<Type *, 2> Types = {DstArrayPtr->getType(),
337 DstArrayPtr->getType()};
338 const DataLayout &DL = B.GetInsertBlock()->getModule()->getDataLayout();
339 uint64_t ArrElemVecSize = DL.getTypeAllocSize(Ty: ArrElemVecTy);
340
341 for (unsigned I = 0; I < SrcNumElements; I += ScalarsPerArrayElement) {
342 unsigned ArrayIndex = I / ScalarsPerArrayElement;
343 // Create a GEP to access the array element.
344 std::array<Value *, 4> Args = {
345 B.getInt1(/*Inbounds=*/V: false), DstArrayPtr, B.getInt32(C: 0),
346 ConstantInt::get(Ty: B.getInt32Ty(), V: ArrayIndex)};
347 auto *ElementPtr = B.CreateIntrinsic(ID: Intrinsic::spv_gep, OverloadTypes: {Types}, Args: {Args});
348 GR->buildAssignPtr(B, ElemTy: ArrElemVecTy, Arg: ElementPtr);
349
350 // Extract scalar elements from the source vector for this array slot.
351 SmallVector<Value *, 4> Elements;
352 for (unsigned J = 0; J < ScalarsPerArrayElement; ++J)
353 Elements.push_back(Elt: makeExtractElement(B, ElementType: ElemTy, Vector: SrcVector, Index: I + J));
354
355 // Build a vector from the extracted elements and store it.
356 Value *Vec = buildVectorFromLoadedElements(B, TargetType: ArrElemVecTy, LoadedElements&: Elements);
357 StoreInst *SI = B.CreateStore(Val: Vec, Ptr: ElementPtr);
358 SI->setAlignment(commonAlignment(A: Alignment, Offset: ArrayIndex * ArrElemVecSize));
359 }
360 }
361
362 // Stores elements from a vector into an array.
363 void storeArrayFromVector(IRBuilder<> &B, Value *SrcVector,
364 Value *DstArrayPtr, ArrayType *ArrTy,
365 Align Alignment) {
366 auto *VecTy = cast<FixedVectorType>(Val: SrcVector->getType());
367 Type *ElemTy = ArrTy->getElementType();
368
369 // Ensure the element types of the array and vector are the same.
370 assert(VecTy->getElementType() == ElemTy &&
371 "Element types of array and vector must be the same.");
372 std::array<Type *, 2> Types = {DstArrayPtr->getType(),
373 DstArrayPtr->getType()};
374 const DataLayout &DL = B.GetInsertBlock()->getModule()->getDataLayout();
375 uint64_t ElemSize = DL.getTypeAllocSize(Ty: ElemTy);
376
377 for (unsigned I = 0, E = VecTy->getNumElements(); I < E; ++I) {
378 // Create a GEP to access the i-th element of the array.
379 std::array<Value *, 4> Args = {B.getInt1(/*Inbounds=*/V: false), DstArrayPtr,
380 B.getInt32(C: 0),
381 ConstantInt::get(Ty: B.getInt32Ty(), V: I)};
382 auto *ElementPtr = B.CreateIntrinsic(ID: Intrinsic::spv_gep, OverloadTypes: {Types}, Args: {Args});
383 GR->buildAssignPtr(B, ElemTy, Arg: ElementPtr);
384
385 // Extract the element from the vector and store it.
386 Value *Element =
387 E == 1 ? SrcVector : makeExtractElement(B, ElementType: ElemTy, Vector: SrcVector, Index: I);
388 StoreInst *SI = B.CreateStore(Val: Element, Ptr: ElementPtr);
389 SI->setAlignment(commonAlignment(A: Alignment, Offset: I * ElemSize));
390 }
391 }
392
393 // Replaces the load instruction to get rid of the ptrcast used as source
394 // operand.
395 void transformLoad(IRBuilder<> &B, LoadInst *LI, Value *CastedOperand,
396 Value *OriginalOperand) {
397 Type *ToTy = GR->findDeducedElementType(Val: CastedOperand);
398 B.SetInsertPoint(LI);
399
400 Value *Output = buildLegalizedLoad(B, ElementType: ToTy, Source: OriginalOperand, BadLoad: LI);
401
402 GR->replaceAllUsesWith(Old: LI, New: Output, /* DeleteOld= */ true);
403 DeadInstructions.push_back(x: LI);
404 }
405
406 // Creates an spv_insertelt instruction (equivalent to llvm's insertelement).
407 Value *makeInsertElement(IRBuilder<> &B, Value *Vector, Value *Element,
408 unsigned Index) {
409 Type *Int32Ty = Type::getInt32Ty(C&: B.getContext());
410 SmallVector<Type *, 4> Types = {Vector->getType(), Vector->getType(),
411 Element->getType(), Int32Ty};
412 SmallVector<Value *> Args = {Vector, Element, B.getInt32(C: Index)};
413 Value *NewI = B.CreateIntrinsic(ID: Intrinsic::spv_insertelt, OverloadTypes: {Types}, Args: {Args});
414 buildAssignType(B, Ty: Vector->getType(), Arg: NewI);
415 return NewI;
416 }
417
418 // Creates an spv_extractelt instruction (equivalent to llvm's
419 // extractelement).
420 Value *makeExtractElement(IRBuilder<> &B, Type *ElementType, Value *Vector,
421 unsigned Index) {
422 Type *Int32Ty = Type::getInt32Ty(C&: B.getContext());
423 SmallVector<Type *, 3> Types = {ElementType, Vector->getType(), Int32Ty};
424 SmallVector<Value *> Args = {Vector, B.getInt32(C: Index)};
425 Value *NewI = B.CreateIntrinsic(ID: Intrinsic::spv_extractelt, OverloadTypes: {Types}, Args: {Args});
426 buildAssignType(B, Ty: ElementType, Arg: NewI);
427 return NewI;
428 }
429
430 // Stores the given Src vector operand into the Dst vector, adjusting the size
431 // if required.
432 Value *storeVectorFromVector(IRBuilder<> &B, Value *Src, Value *Dst,
433 Align Alignment) {
434 FixedVectorType *SrcType = cast<FixedVectorType>(Val: Src->getType());
435 FixedVectorType *DstType =
436 cast<FixedVectorType>(Val: GR->findDeducedElementType(Val: Dst));
437 auto dstNumElements = DstType->getNumElements();
438 auto srcNumElements = SrcType->getNumElements();
439
440 // if the element type differs, it is a bitcast.
441 if (DstType->getElementType() != SrcType->getElementType()) {
442 // Support bitcast between vectors of different sizes only if
443 // the total bitwidth is the same.
444 [[maybe_unused]] auto dstBitWidth =
445 DstType->getElementType()->getScalarSizeInBits() * dstNumElements;
446 [[maybe_unused]] auto srcBitWidth =
447 SrcType->getElementType()->getScalarSizeInBits() * srcNumElements;
448 assert(dstBitWidth == srcBitWidth &&
449 "Unsupported bitcast between vectors of different sizes.");
450
451 Src =
452 B.CreateIntrinsic(ID: Intrinsic::spv_bitcast, OverloadTypes: {DstType, SrcType}, Args: {Src});
453 buildAssignType(B, Ty: DstType, Arg: Src);
454 SrcType = DstType;
455
456 StoreInst *SI = B.CreateStore(Val: Src, Ptr: Dst);
457 SI->setAlignment(Alignment);
458 return SI;
459 }
460
461 assert(DstType->getNumElements() >= SrcType->getNumElements());
462 LoadInst *LI = B.CreateLoad(Ty: DstType, Ptr: Dst);
463 LI->setAlignment(Alignment);
464 Value *OldValues = LI;
465 buildAssignType(B, Ty: OldValues->getType(), Arg: OldValues);
466 Value *NewValues = Src;
467
468 for (unsigned I = 0; I < SrcType->getNumElements(); ++I) {
469 Value *Element =
470 makeExtractElement(B, ElementType: SrcType->getElementType(), Vector: NewValues, Index: I);
471 OldValues = makeInsertElement(B, Vector: OldValues, Element, Index: I);
472 }
473
474 StoreInst *SI = B.CreateStore(Val: OldValues, Ptr: Dst);
475 SI->setAlignment(Alignment);
476 return SI;
477 }
478
479 // Builds a legalized store to a pointer, drilling down through
480 // memory layouts to find a compatible type.
481 void buildLegalizedStore(IRBuilder<> &B, Value *Src, Value *Dst,
482 Align Alignment) {
483 auto ResultOpt = getPointerToFirstCompatibleType(B, BasePtr: Dst, PointerType: Dst->getType(),
484 TargetElemType: Src->getType(), IsInBounds: true);
485 assert(ResultOpt && "Failed to store to aggregate: "
486 "Could not find compatible memory layout.");
487 auto [GEP, CurrentTy] = *ResultOpt;
488
489 auto *DAT = dyn_cast<ArrayType>(Val: CurrentTy);
490 auto *DVT = dyn_cast<FixedVectorType>(Val: CurrentTy);
491 auto *SVT = dyn_cast<FixedVectorType>(Val: Src->getType());
492 auto *DMAT =
493 DAT ? dyn_cast<FixedVectorType>(Val: DAT->getElementType()) : nullptr;
494
495 if (Src->getType() == CurrentTy) {
496 StoreInst *SI = B.CreateStore(Val: Src, Ptr: GEP);
497 SI->setAlignment(Alignment);
498 return;
499 }
500 if (DVT && SVT) {
501 storeVectorFromVector(B, Src, Dst: GEP, Alignment);
502 return;
503 }
504 if (DAT && SVT && SVT->getElementType() == DAT->getElementType()) {
505 storeArrayFromVector(B, SrcVector: Src, DstArrayPtr: GEP, ArrTy: DAT, Alignment);
506 return;
507 }
508 if (DMAT && SVT && DMAT->getElementType() == SVT->getElementType()) {
509 storeMatrixArrayFromVector(B, SrcVector: Src, DstArrayPtr: GEP, ArrTy: DAT, Alignment);
510 return;
511 }
512
513 llvm_unreachable("Failed to store to aggregate.");
514 }
515
516 // Transforms a store instruction (or SPV intrinsic) using a ptrcast as
517 // operand into a valid logical SPIR-V store with no ptrcast.
518 void transformStore(IRBuilder<> &B, Instruction *BadStore, Value *Src,
519 Value *Dst, Align Alignment) {
520 B.SetInsertPoint(BadStore);
521 buildLegalizedStore(B, Src, Dst, Alignment);
522 DeadInstructions.push_back(x: BadStore);
523 }
524
525 void legalizePointerCast(IntrinsicInst *II) {
526 Value *CastedOperand = II;
527 Value *OriginalOperand = II->getOperand(i_nocapture: 0);
528
529 IRBuilder<> B(II->getContext());
530 std::vector<Value *> Users;
531 for (Use &U : II->uses())
532 Users.push_back(x: U.getUser());
533
534 for (Value *User : Users) {
535 if (LoadInst *LI = dyn_cast<LoadInst>(Val: User)) {
536 transformLoad(B, LI, CastedOperand, OriginalOperand);
537 continue;
538 }
539
540 if (StoreInst *SI = dyn_cast<StoreInst>(Val: User)) {
541 transformStore(B, BadStore: SI, Src: SI->getValueOperand(), Dst: OriginalOperand,
542 Alignment: SI->getAlign());
543 continue;
544 }
545
546 if (IntrinsicInst *Intrin = dyn_cast<IntrinsicInst>(Val: User)) {
547 if (Intrin->getIntrinsicID() == Intrinsic::spv_assign_ptr_type) {
548 DeadInstructions.push_back(x: Intrin);
549 continue;
550 }
551
552 if (Intrin->getIntrinsicID() == Intrinsic::spv_gep) {
553 GR->replaceAllUsesWith(Old: CastedOperand, New: OriginalOperand,
554 /* DeleteOld= */ false);
555 continue;
556 }
557
558 if (Intrin->getIntrinsicID() == Intrinsic::spv_store) {
559 Align Alignment;
560 if (ConstantInt *C = dyn_cast<ConstantInt>(Val: Intrin->getOperand(i_nocapture: 3)))
561 Alignment = Align(C->getZExtValue());
562 transformStore(B, BadStore: Intrin, Src: Intrin->getArgOperand(i: 0), Dst: OriginalOperand,
563 Alignment);
564 continue;
565 }
566 }
567
568 llvm_unreachable("Unsupported ptrcast user. Please fix.");
569 }
570
571 DeadInstructions.push_back(x: II);
572 }
573
574public:
575 SPIRVLegalizePointerCastImpl(const SPIRVTargetMachine &TM) : TM(TM) {}
576
577 bool run(Function &F) {
578 const SPIRVSubtarget &ST = TM.getSubtarget<SPIRVSubtarget>(F);
579 GR = ST.getSPIRVGlobalRegistry();
580 DeadInstructions.clear();
581
582 std::vector<IntrinsicInst *> WorkList;
583 for (auto &BB : F) {
584 for (auto &I : BB) {
585 auto *II = dyn_cast<IntrinsicInst>(Val: &I);
586 if (II && II->getIntrinsicID() == Intrinsic::spv_ptrcast)
587 WorkList.push_back(x: II);
588 }
589 }
590
591 for (IntrinsicInst *II : WorkList)
592 legalizePointerCast(II);
593
594 for (Instruction *I : DeadInstructions)
595 I->eraseFromParent();
596
597 return DeadInstructions.size() != 0;
598 }
599
600private:
601 const SPIRVTargetMachine &TM;
602 SPIRVGlobalRegistry *GR = nullptr;
603 std::vector<Instruction *> DeadInstructions;
604};
605
606class SPIRVLegalizePointerCastLegacy : public FunctionPass {
607public:
608 static char ID;
609 SPIRVLegalizePointerCastLegacy(const SPIRVTargetMachine &TM)
610 : FunctionPass(ID), TM(TM) {}
611
612 bool runOnFunction(Function &F) override {
613 return SPIRVLegalizePointerCastImpl(TM).run(F);
614 }
615
616private:
617 const SPIRVTargetMachine &TM;
618};
619} // namespace
620
621PreservedAnalyses SPIRVLegalizePointerCast::run(Function &F,
622 FunctionAnalysisManager &AM) {
623 return SPIRVLegalizePointerCastImpl(TM).run(F) ? PreservedAnalyses::none()
624 : PreservedAnalyses::all();
625}
626
627char SPIRVLegalizePointerCastLegacy::ID = 0;
628INITIALIZE_PASS(SPIRVLegalizePointerCastLegacy, "spirv-legalize-pointer-cast",
629 "SPIRV legalize pointer cast pass", false, false)
630
631FunctionPass *llvm::createSPIRVLegalizePointerCastPass(SPIRVTargetMachine *TM) {
632 return new SPIRVLegalizePointerCastLegacy(*TM);
633}
634