1//===- DXILFlattenArrays.cpp - Flattens DXIL Arrays-----------------------===//
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/// \file This file contains a pass to flatten arrays for the DirectX Backend.
10///
11//===----------------------------------------------------------------------===//
12
13#include "DXILFlattenArrays.h"
14#include "DirectX.h"
15#include "llvm/ADT/PostOrderIterator.h"
16#include "llvm/ADT/STLExtras.h"
17#include "llvm/IR/BasicBlock.h"
18#include "llvm/IR/DerivedTypes.h"
19#include "llvm/IR/IRBuilder.h"
20#include "llvm/IR/InstVisitor.h"
21#include "llvm/IR/ReplaceConstant.h"
22#include "llvm/Support/Casting.h"
23#include "llvm/Support/MathExtras.h"
24#include "llvm/Transforms/Utils/Local.h"
25#include <cassert>
26#include <cstddef>
27#include <cstdint>
28#include <utility>
29
30#define DEBUG_TYPE "dxil-flatten-arrays"
31
32using namespace llvm;
33namespace {
34
35class DXILFlattenArraysLegacy : public ModulePass {
36
37public:
38 bool runOnModule(Module &M) override;
39 DXILFlattenArraysLegacy() : ModulePass(ID) {}
40
41 static char ID; // Pass identification.
42};
43
44struct GEPInfo {
45 ArrayType *RootFlattenedArrayType;
46 Value *RootPointerOperand;
47 SmallMapVector<Value *, APInt, 4> VariableOffsets;
48 APInt ConstantOffset;
49};
50
51class DXILFlattenArraysVisitor
52 : public InstVisitor<DXILFlattenArraysVisitor, bool> {
53public:
54 DXILFlattenArraysVisitor(
55 SmallDenseMap<GlobalVariable *, GlobalVariable *> &GlobalMap)
56 : GlobalMap(GlobalMap) {}
57 bool visit(Function &F);
58 // InstVisitor methods. They return true if the instruction was scalarized,
59 // false if nothing changed.
60 bool visitGetElementPtrInst(GetElementPtrInst &GEPI);
61 bool visitAllocaInst(AllocaInst &AI);
62 bool visitInstruction(Instruction &I) { return false; }
63 bool visitSelectInst(SelectInst &SI) { return false; }
64 bool visitICmpInst(ICmpInst &ICI) { return false; }
65 bool visitFCmpInst(FCmpInst &FCI) { return false; }
66 bool visitUnaryOperator(UnaryOperator &UO) { return false; }
67 bool visitBinaryOperator(BinaryOperator &BO) { return false; }
68 bool visitCastInst(CastInst &CI) { return false; }
69 bool visitBitCastInst(BitCastInst &BCI) { return false; }
70 bool visitInsertElementInst(InsertElementInst &IEI) { return false; }
71 bool visitExtractElementInst(ExtractElementInst &EEI) { return false; }
72 bool visitShuffleVectorInst(ShuffleVectorInst &SVI) { return false; }
73 bool visitPHINode(PHINode &PHI) { return false; }
74 bool visitLoadInst(LoadInst &LI);
75 bool visitStoreInst(StoreInst &SI);
76 bool visitCallInst(CallInst &ICI) { return false; }
77 bool visitFreezeInst(FreezeInst &FI) { return false; }
78 static bool isMultiDimensionalArray(Type *T);
79 static std::pair<unsigned, Type *> getElementCountAndType(Type *ArrayTy);
80
81private:
82 SmallVector<WeakTrackingVH> PotentiallyDeadInstrs;
83 SmallDenseMap<GEPOperator *, GEPInfo> GEPChainInfoMap;
84 SmallDenseMap<GlobalVariable *, GlobalVariable *> &GlobalMap;
85 bool finish();
86 ConstantInt *genConstFlattenIndices(ArrayRef<Value *> Indices,
87 ArrayRef<uint64_t> Dims,
88 IRBuilder<> &Builder);
89 Value *genInstructionFlattenIndices(ArrayRef<Value *> Indices,
90 ArrayRef<uint64_t> Dims,
91 IRBuilder<> &Builder);
92};
93} // namespace
94
95bool DXILFlattenArraysVisitor::finish() {
96 GEPChainInfoMap.clear();
97 RecursivelyDeleteTriviallyDeadInstructionsPermissive(DeadInsts&: PotentiallyDeadInstrs);
98 return true;
99}
100
101bool DXILFlattenArraysVisitor::isMultiDimensionalArray(Type *T) {
102 if (ArrayType *ArrType = dyn_cast<ArrayType>(Val: T))
103 return isa<ArrayType>(Val: ArrType->getElementType());
104 return false;
105}
106
107std::pair<unsigned, Type *>
108DXILFlattenArraysVisitor::getElementCountAndType(Type *ArrayTy) {
109 unsigned TotalElements = 1;
110 Type *CurrArrayTy = ArrayTy;
111 while (auto *InnerArrayTy = dyn_cast<ArrayType>(Val: CurrArrayTy)) {
112 TotalElements *= InnerArrayTy->getNumElements();
113 CurrArrayTy = InnerArrayTy->getElementType();
114 }
115 return std::make_pair(x&: TotalElements, y&: CurrArrayTy);
116}
117
118ConstantInt *DXILFlattenArraysVisitor::genConstFlattenIndices(
119 ArrayRef<Value *> Indices, ArrayRef<uint64_t> Dims, IRBuilder<> &Builder) {
120 assert(Indices.size() == Dims.size() &&
121 "Indicies and dimmensions should be the same");
122 unsigned FlatIndex = 0;
123 unsigned Multiplier = 1;
124
125 for (int I = Indices.size() - 1; I >= 0; --I) {
126 unsigned DimSize = Dims[I];
127 ConstantInt *CIndex = dyn_cast<ConstantInt>(Val: Indices[I]);
128 assert(CIndex && "This function expects all indicies to be ConstantInt");
129 FlatIndex += CIndex->getZExtValue() * Multiplier;
130 Multiplier *= DimSize;
131 }
132 return Builder.getInt32(C: FlatIndex);
133}
134
135Value *DXILFlattenArraysVisitor::genInstructionFlattenIndices(
136 ArrayRef<Value *> Indices, ArrayRef<uint64_t> Dims, IRBuilder<> &Builder) {
137 if (Indices.size() == 1)
138 return Indices[0];
139
140 Value *FlatIndex = Builder.getInt32(C: 0);
141 unsigned Multiplier = 1;
142
143 for (int I = Indices.size() - 1; I >= 0; --I) {
144 unsigned DimSize = Dims[I];
145 Value *VMultiplier = Builder.getInt32(C: Multiplier);
146 Value *ScaledIndex = Builder.CreateMul(LHS: Indices[I], RHS: VMultiplier);
147 FlatIndex = Builder.CreateAdd(LHS: FlatIndex, RHS: ScaledIndex);
148 Multiplier *= DimSize;
149 }
150 return FlatIndex;
151}
152
153bool DXILFlattenArraysVisitor::visitLoadInst(LoadInst &LI) {
154 unsigned NumOperands = LI.getNumOperands();
155 for (unsigned I = 0; I < NumOperands; ++I) {
156 Value *CurrOpperand = LI.getOperand(i_nocapture: I);
157 ConstantExpr *CE = dyn_cast<ConstantExpr>(Val: CurrOpperand);
158 if (CE && CE->getOpcode() == Instruction::GetElementPtr) {
159 GetElementPtrInst *OldGEP =
160 cast<GetElementPtrInst>(Val: CE->getAsInstruction());
161 OldGEP->insertBefore(InsertPos: LI.getIterator());
162
163 IRBuilder<> Builder(&LI);
164 LoadInst *NewLoad =
165 Builder.CreateLoad(Ty: LI.getType(), Ptr: OldGEP, Name: LI.getName());
166 NewLoad->setAlignment(LI.getAlign());
167 LI.replaceAllUsesWith(V: NewLoad);
168 LI.eraseFromParent();
169 visitGetElementPtrInst(GEPI&: *OldGEP);
170 return true;
171 }
172 }
173 return false;
174}
175
176bool DXILFlattenArraysVisitor::visitStoreInst(StoreInst &SI) {
177 unsigned NumOperands = SI.getNumOperands();
178 for (unsigned I = 0; I < NumOperands; ++I) {
179 Value *CurrOpperand = SI.getOperand(i_nocapture: I);
180 ConstantExpr *CE = dyn_cast<ConstantExpr>(Val: CurrOpperand);
181 if (CE && CE->getOpcode() == Instruction::GetElementPtr) {
182 GetElementPtrInst *OldGEP =
183 cast<GetElementPtrInst>(Val: CE->getAsInstruction());
184 OldGEP->insertBefore(InsertPos: SI.getIterator());
185
186 IRBuilder<> Builder(&SI);
187 StoreInst *NewStore = Builder.CreateStore(Val: SI.getValueOperand(), Ptr: OldGEP);
188 NewStore->setAlignment(SI.getAlign());
189 SI.replaceAllUsesWith(V: NewStore);
190 SI.eraseFromParent();
191 visitGetElementPtrInst(GEPI&: *OldGEP);
192 return true;
193 }
194 }
195 return false;
196}
197
198bool DXILFlattenArraysVisitor::visitAllocaInst(AllocaInst &AI) {
199 if (!isMultiDimensionalArray(T: AI.getAllocatedType()))
200 return false;
201
202 ArrayType *ArrType = cast<ArrayType>(Val: AI.getAllocatedType());
203 IRBuilder<> Builder(&AI);
204 auto [TotalElements, BaseType] = getElementCountAndType(ArrayTy: ArrType);
205
206 ArrayType *FattenedArrayType = ArrayType::get(ElementType: BaseType, NumElements: TotalElements);
207 AllocaInst *FlatAlloca =
208 Builder.CreateAlloca(Ty: FattenedArrayType, ArraySize: nullptr, Name: AI.getName() + ".1dim");
209 FlatAlloca->setAlignment(AI.getAlign());
210 AI.replaceAllUsesWith(V: FlatAlloca);
211 AI.eraseFromParent();
212 return true;
213}
214
215bool DXILFlattenArraysVisitor::visitGetElementPtrInst(GetElementPtrInst &GEP) {
216 // Do not visit GEPs more than once
217 if (GEPChainInfoMap.contains(Val: cast<GEPOperator>(Val: &GEP)))
218 return false;
219
220 Value *PtrOperand = GEP.getPointerOperand();
221 // It shouldn't(?) be possible for the pointer operand of a GEP to be a PHI
222 // node unless HLSL has pointers. If this assumption is incorrect or HLSL gets
223 // pointer types, then the handling of this case can be implemented later.
224 assert(!isa<PHINode>(PtrOperand) &&
225 "Pointer operand of GEP should not be a PHI Node");
226
227 // Replace a GEP ConstantExpr pointer operand with a GEP instruction so that
228 // it can be visited
229 if (auto *PtrOpGEPCE = dyn_cast<ConstantExpr>(Val: PtrOperand);
230 PtrOpGEPCE && PtrOpGEPCE->getOpcode() == Instruction::GetElementPtr) {
231 GetElementPtrInst *OldGEPI =
232 cast<GetElementPtrInst>(Val: PtrOpGEPCE->getAsInstruction());
233 OldGEPI->insertBefore(InsertPos: GEP.getIterator());
234
235 SmallVector<Value *> Indices(GEP.indices());
236 GetElementPtrInst *NewGEPI = GetElementPtrInst::Create(
237 PointeeType: GEP.getSourceElementType(), Ptr: OldGEPI, IdxList: Indices, NW: GEP.getNoWrapFlags(),
238 NameStr: GEP.getName(), InsertBefore: GEP.getIterator());
239
240 GEP.replaceAllUsesWith(V: NewGEPI);
241 GEP.eraseFromParent();
242 visitGetElementPtrInst(GEP&: *OldGEPI);
243 visitGetElementPtrInst(GEP&: *NewGEPI);
244 return true;
245 }
246
247 // Construct GEPInfo for this GEP
248 GEPInfo Info;
249
250 // Obtain the variable and constant byte offsets computed by this GEP
251 const DataLayout &DL = GEP.getDataLayout();
252 unsigned BitWidth = DL.getIndexTypeSizeInBits(Ty: GEP.getType());
253 Info.ConstantOffset = {BitWidth, 0};
254 [[maybe_unused]] bool Success = GEP.collectOffset(
255 DL, BitWidth, VariableOffsets&: Info.VariableOffsets, ConstantOffset&: Info.ConstantOffset);
256 assert(Success && "Failed to collect offsets for GEP");
257
258 // If there is a parent GEP, inherit the root array type and pointer, and
259 // merge the byte offsets. Otherwise, this GEP is itself the root of a GEP
260 // chain and we need to deterine the root array type
261 if (auto *PtrOpGEP = dyn_cast<GEPOperator>(Val: PtrOperand)) {
262
263 // If the parent GEP was not processed, then we do not want to process its
264 // descendants. This can happen if the GEP chain is for an unsupported type
265 // such as a struct -- we do not flatten structs nor GEP chains for structs
266 if (!GEPChainInfoMap.contains(Val: PtrOpGEP))
267 return false;
268
269 GEPInfo &PGEPInfo = GEPChainInfoMap[PtrOpGEP];
270 Info.RootFlattenedArrayType = PGEPInfo.RootFlattenedArrayType;
271 Info.RootPointerOperand = PGEPInfo.RootPointerOperand;
272 for (auto &VariableOffset : PGEPInfo.VariableOffsets)
273 Info.VariableOffsets.insert(KV: VariableOffset);
274 Info.ConstantOffset += PGEPInfo.ConstantOffset;
275 } else {
276 Info.RootPointerOperand = PtrOperand;
277
278 // We should try to determine the type of the root from the pointer rather
279 // than the GEP's source element type because this could be a scalar GEP
280 // into an array-typed pointer from an Alloca or Global Variable.
281 Type *RootTy = GEP.getSourceElementType();
282 if (auto *GlobalVar = dyn_cast<GlobalVariable>(Val: PtrOperand)) {
283 if (GlobalMap.contains(Val: GlobalVar))
284 GlobalVar = GlobalMap[GlobalVar];
285 Info.RootPointerOperand = GlobalVar;
286 RootTy = GlobalVar->getValueType();
287 } else if (auto *Alloca = dyn_cast<AllocaInst>(Val: PtrOperand))
288 RootTy = Alloca->getAllocatedType();
289 assert(!isMultiDimensionalArray(RootTy) &&
290 "Expected root array type to be flattened");
291
292 // If the root type is not an array, we don't need to do any flattening
293 if (!isa<ArrayType>(Val: RootTy))
294 return false;
295
296 Info.RootFlattenedArrayType = cast<ArrayType>(Val: RootTy);
297 }
298
299 // GEPs without users or GEPs with non-GEP users should be replaced such that
300 // the chain of GEPs they are a part of are collapsed to a single GEP into a
301 // flattened array.
302 bool ReplaceThisGEP = GEP.users().empty();
303 for (Value *User : GEP.users())
304 if (!isa<GetElementPtrInst>(Val: User))
305 ReplaceThisGEP = true;
306
307 if (ReplaceThisGEP) {
308 unsigned BytesPerElem =
309 DL.getTypeAllocSize(Ty: Info.RootFlattenedArrayType->getArrayElementType());
310 assert(isPowerOf2_32(BytesPerElem) &&
311 "Bytes per element should be a power of 2");
312
313 // Compute the 32-bit index for this flattened GEP from the constant and
314 // variable byte offsets in the GEPInfo
315 IRBuilder<> Builder(&GEP);
316 Value *ZeroIndex = Builder.getInt32(C: 0);
317 uint64_t ConstantOffset =
318 Info.ConstantOffset.udiv(RHS: BytesPerElem).getZExtValue();
319 assert(ConstantOffset < UINT32_MAX &&
320 "Constant byte offset for flat GEP index must fit within 32 bits");
321 Value *FlattenedIndex = Builder.getInt32(C: ConstantOffset);
322 for (auto [VarIndex, Multiplier] : Info.VariableOffsets) {
323 assert(Multiplier.getActiveBits() <= 32 &&
324 "The multiplier for a flat GEP index must fit within 32 bits");
325 assert(VarIndex->getType()->isIntegerTy(32) &&
326 "Expected i32-typed GEP indices");
327 Value *VI;
328 if (Multiplier.getZExtValue() % BytesPerElem != 0) {
329 // This can happen, e.g., with i8 GEPs. To handle this we just divide
330 // by BytesPerElem using an instruction after multiplying VarIndex by
331 // Multiplier.
332 VI = Builder.CreateMul(LHS: VarIndex,
333 RHS: Builder.getInt32(C: Multiplier.getZExtValue()));
334 VI = Builder.CreateLShr(LHS: VI, RHS: Builder.getInt32(C: Log2_32(Value: BytesPerElem)));
335 } else
336 VI = Builder.CreateMul(
337 LHS: VarIndex,
338 RHS: Builder.getInt32(C: Multiplier.getZExtValue() / BytesPerElem));
339 FlattenedIndex = Builder.CreateAdd(LHS: FlattenedIndex, RHS: VI);
340 }
341
342 // Construct a new GEP for the flattened array to replace the current GEP
343 GetElementPtrInst *NewGEP = GetElementPtrInst::Create(
344 PointeeType: Info.RootFlattenedArrayType, Ptr: Info.RootPointerOperand,
345 IdxList: {ZeroIndex, FlattenedIndex}, NW: GEP.getNoWrapFlags(), NameStr: GEP.getName(),
346 InsertBefore: Builder.GetInsertPoint());
347
348 // Replace the current GEP with the new GEP. Store GEPInfo into the map
349 // for later use in case this GEP was not the end of the chain
350 GEPChainInfoMap.insert(KV: {cast<GEPOperator>(Val: NewGEP), std::move(Info)});
351 GEP.replaceAllUsesWith(V: NewGEP);
352 GEP.eraseFromParent();
353 return true;
354 }
355
356 // This GEP is potentially dead at the end of the pass since it may not have
357 // any users anymore after GEP chains have been collapsed. We retain store
358 // GEPInfo for GEPs down the chain to use to compute their indices.
359 GEPChainInfoMap.insert(KV: {cast<GEPOperator>(Val: &GEP), std::move(Info)});
360 PotentiallyDeadInstrs.emplace_back(Args: &GEP);
361 return false;
362}
363
364bool DXILFlattenArraysVisitor::visit(Function &F) {
365 bool MadeChange = false;
366 ReversePostOrderTraversal<Function *> RPOT(&F);
367 for (BasicBlock *BB : make_early_inc_range(Range&: RPOT)) {
368 for (Instruction &I : make_early_inc_range(Range&: *BB))
369 MadeChange |= InstVisitor::visit(I);
370 }
371 finish();
372 return MadeChange;
373}
374
375static void collectElements(Constant *Init,
376 SmallVectorImpl<Constant *> &Elements) {
377 // Base case: If Init is not an array, add it directly to the vector.
378 auto *ArrayTy = dyn_cast<ArrayType>(Val: Init->getType());
379 if (!ArrayTy) {
380 Elements.push_back(Elt: Init);
381 return;
382 }
383 unsigned ArrSize = ArrayTy->getNumElements();
384 if (isa<ConstantAggregateZero>(Val: Init)) {
385 for (unsigned I = 0; I < ArrSize; ++I)
386 Elements.push_back(Elt: Constant::getNullValue(Ty: ArrayTy->getElementType()));
387 return;
388 }
389
390 // Recursive case: Process each element in the array.
391 if (auto *ArrayConstant = dyn_cast<ConstantArray>(Val: Init)) {
392 for (unsigned I = 0; I < ArrayConstant->getNumOperands(); ++I) {
393 collectElements(Init: ArrayConstant->getOperand(i_nocapture: I), Elements);
394 }
395 } else if (auto *DataArrayConstant = dyn_cast<ConstantDataArray>(Val: Init)) {
396 for (unsigned I = 0; I < DataArrayConstant->getNumElements(); ++I) {
397 collectElements(Init: DataArrayConstant->getElementAsConstant(i: I), Elements);
398 }
399 } else {
400 llvm_unreachable(
401 "Expected a ConstantArray or ConstantDataArray for array initializer!");
402 }
403}
404
405static Constant *transformInitializer(Constant *Init, Type *OrigType,
406 ArrayType *FlattenedType,
407 LLVMContext &Ctx) {
408 // Handle ConstantAggregateZero (zero-initialized constants)
409 if (isa<ConstantAggregateZero>(Val: Init))
410 return ConstantAggregateZero::get(Ty: FlattenedType);
411
412 // Handle UndefValue (undefined constants)
413 if (isa<UndefValue>(Val: Init))
414 return UndefValue::get(T: FlattenedType);
415
416 if (!isa<ArrayType>(Val: OrigType))
417 return Init;
418
419 SmallVector<Constant *> FlattenedElements;
420 collectElements(Init, Elements&: FlattenedElements);
421 assert(FlattenedType->getNumElements() == FlattenedElements.size() &&
422 "The number of collected elements should match the FlattenedType");
423 return ConstantArray::get(T: FlattenedType, V: FlattenedElements);
424}
425
426static void flattenGlobalArrays(
427 Module &M, SmallDenseMap<GlobalVariable *, GlobalVariable *> &GlobalMap) {
428 LLVMContext &Ctx = M.getContext();
429 for (GlobalVariable &G : M.globals()) {
430 Type *OrigType = G.getValueType();
431 if (!DXILFlattenArraysVisitor::isMultiDimensionalArray(T: OrigType))
432 continue;
433
434 ArrayType *ArrType = cast<ArrayType>(Val: OrigType);
435 auto [TotalElements, BaseType] =
436 DXILFlattenArraysVisitor::getElementCountAndType(ArrayTy: ArrType);
437 ArrayType *FattenedArrayType = ArrayType::get(ElementType: BaseType, NumElements: TotalElements);
438
439 // Create a new global variable with the updated type
440 // Note: Initializer is set via transformInitializer
441 GlobalVariable *NewGlobal =
442 new GlobalVariable(M, FattenedArrayType, G.isConstant(), G.getLinkage(),
443 /*Initializer=*/nullptr, G.getName() + ".1dim", &G,
444 G.getThreadLocalMode(), G.getAddressSpace(),
445 G.isExternallyInitialized());
446
447 // Copy relevant attributes
448 NewGlobal->setUnnamedAddr(G.getUnnamedAddr());
449 if (G.getAlign()) {
450 NewGlobal->setAlignment(G.getAlign());
451 }
452
453 if (G.hasInitializer()) {
454 Constant *Init = G.getInitializer();
455 Constant *NewInit =
456 transformInitializer(Init, OrigType, FlattenedType: FattenedArrayType, Ctx);
457 NewGlobal->setInitializer(NewInit);
458 }
459 GlobalMap[&G] = NewGlobal;
460 }
461}
462
463static bool flattenArrays(Module &M) {
464 bool MadeChange = false;
465 SmallDenseMap<GlobalVariable *, GlobalVariable *> GlobalMap;
466 flattenGlobalArrays(M, GlobalMap);
467 DXILFlattenArraysVisitor Impl(GlobalMap);
468 for (auto &F : make_early_inc_range(Range: M.functions())) {
469 if (F.isDeclaration())
470 continue;
471 MadeChange |= Impl.visit(F);
472 }
473 for (auto &[Old, New] : GlobalMap) {
474 Old->replaceAllUsesWith(V: New);
475 Old->eraseFromParent();
476 MadeChange = true;
477 }
478 return MadeChange;
479}
480
481PreservedAnalyses DXILFlattenArrays::run(Module &M, ModuleAnalysisManager &) {
482 bool MadeChanges = flattenArrays(M);
483 if (!MadeChanges)
484 return PreservedAnalyses::all();
485 PreservedAnalyses PA;
486 return PA;
487}
488
489bool DXILFlattenArraysLegacy::runOnModule(Module &M) {
490 return flattenArrays(M);
491}
492
493char DXILFlattenArraysLegacy::ID = 0;
494
495INITIALIZE_PASS_BEGIN(DXILFlattenArraysLegacy, DEBUG_TYPE,
496 "DXIL Array Flattener", false, false)
497INITIALIZE_PASS_END(DXILFlattenArraysLegacy, DEBUG_TYPE, "DXIL Array Flattener",
498 false, false)
499
500ModulePass *llvm::createDXILFlattenArraysLegacyPass() {
501 return new DXILFlattenArraysLegacy();
502}
503