| 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 | |
| 32 | using namespace llvm; |
| 33 | namespace { |
| 34 | |
| 35 | class DXILFlattenArraysLegacy : public ModulePass { |
| 36 | |
| 37 | public: |
| 38 | bool runOnModule(Module &M) override; |
| 39 | DXILFlattenArraysLegacy() : ModulePass(ID) {} |
| 40 | |
| 41 | static char ID; // Pass identification. |
| 42 | }; |
| 43 | |
| 44 | struct GEPInfo { |
| 45 | ArrayType *RootFlattenedArrayType; |
| 46 | Value *RootPointerOperand; |
| 47 | SmallMapVector<Value *, APInt, 4> VariableOffsets; |
| 48 | APInt ConstantOffset; |
| 49 | }; |
| 50 | |
| 51 | class DXILFlattenArraysVisitor |
| 52 | : public InstVisitor<DXILFlattenArraysVisitor, bool> { |
| 53 | public: |
| 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 (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 | |
| 81 | private: |
| 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 | |
| 95 | bool DXILFlattenArraysVisitor::finish() { |
| 96 | GEPChainInfoMap.clear(); |
| 97 | RecursivelyDeleteTriviallyDeadInstructionsPermissive(DeadInsts&: PotentiallyDeadInstrs); |
| 98 | return true; |
| 99 | } |
| 100 | |
| 101 | bool 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 | |
| 107 | std::pair<unsigned, Type *> |
| 108 | DXILFlattenArraysVisitor::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 | |
| 118 | ConstantInt *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 | |
| 135 | Value *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 | |
| 153 | bool 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 | |
| 176 | bool 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 | |
| 198 | bool 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 | |
| 215 | bool 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 | IRBuilder<> Builder(&GEP); |
| 236 | SmallVector<Value *> Indices(GEP.indices()); |
| 237 | Value *NewGEP = |
| 238 | Builder.CreateGEP(Ty: GEP.getSourceElementType(), Ptr: OldGEPI, IdxList: Indices, |
| 239 | Name: GEP.getName(), NW: GEP.getNoWrapFlags()); |
| 240 | assert(isa<GetElementPtrInst>(NewGEP) && |
| 241 | "Expected newly-created GEP to be an instruction" ); |
| 242 | GetElementPtrInst *NewGEPI = cast<GetElementPtrInst>(Val: NewGEP); |
| 243 | |
| 244 | GEP.replaceAllUsesWith(V: NewGEPI); |
| 245 | GEP.eraseFromParent(); |
| 246 | visitGetElementPtrInst(GEP&: *OldGEPI); |
| 247 | visitGetElementPtrInst(GEP&: *NewGEPI); |
| 248 | return true; |
| 249 | } |
| 250 | |
| 251 | // Construct GEPInfo for this GEP |
| 252 | GEPInfo Info; |
| 253 | |
| 254 | // Obtain the variable and constant byte offsets computed by this GEP |
| 255 | const DataLayout &DL = GEP.getDataLayout(); |
| 256 | unsigned BitWidth = DL.getIndexTypeSizeInBits(Ty: GEP.getType()); |
| 257 | Info.ConstantOffset = {BitWidth, 0}; |
| 258 | [[maybe_unused]] bool Success = GEP.collectOffset( |
| 259 | DL, BitWidth, VariableOffsets&: Info.VariableOffsets, ConstantOffset&: Info.ConstantOffset); |
| 260 | assert(Success && "Failed to collect offsets for GEP" ); |
| 261 | |
| 262 | // If there is a parent GEP, inherit the root array type and pointer, and |
| 263 | // merge the byte offsets. Otherwise, this GEP is itself the root of a GEP |
| 264 | // chain and we need to deterine the root array type |
| 265 | if (auto *PtrOpGEP = dyn_cast<GEPOperator>(Val: PtrOperand)) { |
| 266 | |
| 267 | // If the parent GEP was not processed, then we do not want to process its |
| 268 | // descendants. This can happen if the GEP chain is for an unsupported type |
| 269 | // such as a struct -- we do not flatten structs nor GEP chains for structs |
| 270 | if (!GEPChainInfoMap.contains(Val: PtrOpGEP)) |
| 271 | return false; |
| 272 | |
| 273 | GEPInfo &PGEPInfo = GEPChainInfoMap[PtrOpGEP]; |
| 274 | Info.RootFlattenedArrayType = PGEPInfo.RootFlattenedArrayType; |
| 275 | Info.RootPointerOperand = PGEPInfo.RootPointerOperand; |
| 276 | for (auto &VariableOffset : PGEPInfo.VariableOffsets) |
| 277 | Info.VariableOffsets.insert(KV: VariableOffset); |
| 278 | Info.ConstantOffset += PGEPInfo.ConstantOffset; |
| 279 | } else { |
| 280 | Info.RootPointerOperand = PtrOperand; |
| 281 | |
| 282 | // We should try to determine the type of the root from the pointer rather |
| 283 | // than the GEP's source element type because this could be a scalar GEP |
| 284 | // into an array-typed pointer from an Alloca or Global Variable. |
| 285 | Type *RootTy = GEP.getSourceElementType(); |
| 286 | if (auto *GlobalVar = dyn_cast<GlobalVariable>(Val: PtrOperand)) { |
| 287 | if (GlobalMap.contains(Val: GlobalVar)) |
| 288 | GlobalVar = GlobalMap[GlobalVar]; |
| 289 | Info.RootPointerOperand = GlobalVar; |
| 290 | RootTy = GlobalVar->getValueType(); |
| 291 | } else if (auto *Alloca = dyn_cast<AllocaInst>(Val: PtrOperand)) |
| 292 | RootTy = Alloca->getAllocatedType(); |
| 293 | assert(!isMultiDimensionalArray(RootTy) && |
| 294 | "Expected root array type to be flattened" ); |
| 295 | |
| 296 | // If the root type is not an array, we don't need to do any flattening |
| 297 | if (!isa<ArrayType>(Val: RootTy)) |
| 298 | return false; |
| 299 | |
| 300 | Info.RootFlattenedArrayType = cast<ArrayType>(Val: RootTy); |
| 301 | } |
| 302 | |
| 303 | // GEPs without users or GEPs with non-GEP users should be replaced such that |
| 304 | // the chain of GEPs they are a part of are collapsed to a single GEP into a |
| 305 | // flattened array. |
| 306 | bool ReplaceThisGEP = GEP.users().empty(); |
| 307 | for (Value *User : GEP.users()) |
| 308 | if (!isa<GetElementPtrInst>(Val: User)) |
| 309 | ReplaceThisGEP = true; |
| 310 | |
| 311 | if (ReplaceThisGEP) { |
| 312 | unsigned BytesPerElem = |
| 313 | DL.getTypeAllocSize(Ty: Info.RootFlattenedArrayType->getArrayElementType()); |
| 314 | assert(isPowerOf2_32(BytesPerElem) && |
| 315 | "Bytes per element should be a power of 2" ); |
| 316 | |
| 317 | // Compute the 32-bit index for this flattened GEP from the constant and |
| 318 | // variable byte offsets in the GEPInfo |
| 319 | IRBuilder<> Builder(&GEP); |
| 320 | Value *ZeroIndex = Builder.getInt32(C: 0); |
| 321 | uint64_t ConstantOffset = |
| 322 | Info.ConstantOffset.udiv(RHS: BytesPerElem).getZExtValue(); |
| 323 | assert(ConstantOffset < UINT32_MAX && |
| 324 | "Constant byte offset for flat GEP index must fit within 32 bits" ); |
| 325 | Value *FlattenedIndex = Builder.getInt32(C: ConstantOffset); |
| 326 | for (auto [VarIndex, Multiplier] : Info.VariableOffsets) { |
| 327 | assert(Multiplier.getActiveBits() <= 32 && |
| 328 | "The multiplier for a flat GEP index must fit within 32 bits" ); |
| 329 | assert(VarIndex->getType()->isIntegerTy(32) && |
| 330 | "Expected i32-typed GEP indices" ); |
| 331 | Value *VI; |
| 332 | if (Multiplier.getZExtValue() % BytesPerElem != 0) { |
| 333 | // This can happen, e.g., with i8 GEPs. To handle this we just divide |
| 334 | // by BytesPerElem using an instruction after multiplying VarIndex by |
| 335 | // Multiplier. |
| 336 | VI = Builder.CreateMul(LHS: VarIndex, |
| 337 | RHS: Builder.getInt32(C: Multiplier.getZExtValue())); |
| 338 | VI = Builder.CreateLShr(LHS: VI, RHS: Builder.getInt32(C: Log2_32(Value: BytesPerElem))); |
| 339 | } else |
| 340 | VI = Builder.CreateMul( |
| 341 | LHS: VarIndex, |
| 342 | RHS: Builder.getInt32(C: Multiplier.getZExtValue() / BytesPerElem)); |
| 343 | FlattenedIndex = Builder.CreateAdd(LHS: FlattenedIndex, RHS: VI); |
| 344 | } |
| 345 | |
| 346 | // Construct a new GEP for the flattened array to replace the current GEP |
| 347 | Value *NewGEP = Builder.CreateGEP( |
| 348 | Ty: Info.RootFlattenedArrayType, Ptr: Info.RootPointerOperand, |
| 349 | IdxList: {ZeroIndex, FlattenedIndex}, Name: GEP.getName(), NW: GEP.getNoWrapFlags()); |
| 350 | |
| 351 | // If the pointer operand is a global variable and all indices are 0, |
| 352 | // IRBuilder::CreateGEP will return the global variable instead of creating |
| 353 | // a GEP instruction or GEP ConstantExpr. In this case we have to create and |
| 354 | // insert our own GEP instruction. |
| 355 | if (!isa<GEPOperator>(Val: NewGEP)) |
| 356 | NewGEP = GetElementPtrInst::Create( |
| 357 | PointeeType: Info.RootFlattenedArrayType, Ptr: Info.RootPointerOperand, |
| 358 | IdxList: {ZeroIndex, FlattenedIndex}, NW: GEP.getNoWrapFlags(), NameStr: GEP.getName(), |
| 359 | InsertBefore: Builder.GetInsertPoint()); |
| 360 | |
| 361 | // Replace the current GEP with the new GEP. Store GEPInfo into the map |
| 362 | // for later use in case this GEP was not the end of the chain |
| 363 | GEPChainInfoMap.insert(KV: {cast<GEPOperator>(Val: NewGEP), std::move(Info)}); |
| 364 | GEP.replaceAllUsesWith(V: NewGEP); |
| 365 | GEP.eraseFromParent(); |
| 366 | return true; |
| 367 | } |
| 368 | |
| 369 | // This GEP is potentially dead at the end of the pass since it may not have |
| 370 | // any users anymore after GEP chains have been collapsed. We retain store |
| 371 | // GEPInfo for GEPs down the chain to use to compute their indices. |
| 372 | GEPChainInfoMap.insert(KV: {cast<GEPOperator>(Val: &GEP), std::move(Info)}); |
| 373 | PotentiallyDeadInstrs.emplace_back(Args: &GEP); |
| 374 | return false; |
| 375 | } |
| 376 | |
| 377 | bool DXILFlattenArraysVisitor::visit(Function &F) { |
| 378 | bool MadeChange = false; |
| 379 | ReversePostOrderTraversal<Function *> RPOT(&F); |
| 380 | for (BasicBlock *BB : make_early_inc_range(Range&: RPOT)) { |
| 381 | for (Instruction &I : make_early_inc_range(Range&: *BB)) |
| 382 | MadeChange |= InstVisitor::visit(I); |
| 383 | } |
| 384 | finish(); |
| 385 | return MadeChange; |
| 386 | } |
| 387 | |
| 388 | static void collectElements(Constant *Init, |
| 389 | SmallVectorImpl<Constant *> &Elements) { |
| 390 | // Base case: If Init is not an array, add it directly to the vector. |
| 391 | auto *ArrayTy = dyn_cast<ArrayType>(Val: Init->getType()); |
| 392 | if (!ArrayTy) { |
| 393 | Elements.push_back(Elt: Init); |
| 394 | return; |
| 395 | } |
| 396 | unsigned ArrSize = ArrayTy->getNumElements(); |
| 397 | if (isa<ConstantAggregateZero>(Val: Init)) { |
| 398 | for (unsigned I = 0; I < ArrSize; ++I) |
| 399 | Elements.push_back(Elt: Constant::getNullValue(Ty: ArrayTy->getElementType())); |
| 400 | return; |
| 401 | } |
| 402 | |
| 403 | // Recursive case: Process each element in the array. |
| 404 | if (auto *ArrayConstant = dyn_cast<ConstantArray>(Val: Init)) { |
| 405 | for (unsigned I = 0; I < ArrayConstant->getNumOperands(); ++I) { |
| 406 | collectElements(Init: ArrayConstant->getOperand(i_nocapture: I), Elements); |
| 407 | } |
| 408 | } else if (auto *DataArrayConstant = dyn_cast<ConstantDataArray>(Val: Init)) { |
| 409 | for (unsigned I = 0; I < DataArrayConstant->getNumElements(); ++I) { |
| 410 | collectElements(Init: DataArrayConstant->getElementAsConstant(i: I), Elements); |
| 411 | } |
| 412 | } else { |
| 413 | llvm_unreachable( |
| 414 | "Expected a ConstantArray or ConstantDataArray for array initializer!" ); |
| 415 | } |
| 416 | } |
| 417 | |
| 418 | static Constant *transformInitializer(Constant *Init, Type *OrigType, |
| 419 | ArrayType *FlattenedType, |
| 420 | LLVMContext &Ctx) { |
| 421 | // Handle ConstantAggregateZero (zero-initialized constants) |
| 422 | if (isa<ConstantAggregateZero>(Val: Init)) |
| 423 | return ConstantAggregateZero::get(Ty: FlattenedType); |
| 424 | |
| 425 | // Handle UndefValue (undefined constants) |
| 426 | if (isa<UndefValue>(Val: Init)) |
| 427 | return UndefValue::get(T: FlattenedType); |
| 428 | |
| 429 | if (!isa<ArrayType>(Val: OrigType)) |
| 430 | return Init; |
| 431 | |
| 432 | SmallVector<Constant *> FlattenedElements; |
| 433 | collectElements(Init, Elements&: FlattenedElements); |
| 434 | assert(FlattenedType->getNumElements() == FlattenedElements.size() && |
| 435 | "The number of collected elements should match the FlattenedType" ); |
| 436 | return ConstantArray::get(T: FlattenedType, V: FlattenedElements); |
| 437 | } |
| 438 | |
| 439 | static void flattenGlobalArrays( |
| 440 | Module &M, SmallDenseMap<GlobalVariable *, GlobalVariable *> &GlobalMap) { |
| 441 | LLVMContext &Ctx = M.getContext(); |
| 442 | for (GlobalVariable &G : M.globals()) { |
| 443 | Type *OrigType = G.getValueType(); |
| 444 | if (!DXILFlattenArraysVisitor::isMultiDimensionalArray(T: OrigType)) |
| 445 | continue; |
| 446 | |
| 447 | ArrayType *ArrType = cast<ArrayType>(Val: OrigType); |
| 448 | auto [TotalElements, BaseType] = |
| 449 | DXILFlattenArraysVisitor::getElementCountAndType(ArrayTy: ArrType); |
| 450 | ArrayType *FattenedArrayType = ArrayType::get(ElementType: BaseType, NumElements: TotalElements); |
| 451 | |
| 452 | // Create a new global variable with the updated type |
| 453 | // Note: Initializer is set via transformInitializer |
| 454 | GlobalVariable *NewGlobal = |
| 455 | new GlobalVariable(M, FattenedArrayType, G.isConstant(), G.getLinkage(), |
| 456 | /*Initializer=*/nullptr, G.getName() + ".1dim" , &G, |
| 457 | G.getThreadLocalMode(), G.getAddressSpace(), |
| 458 | G.isExternallyInitialized()); |
| 459 | |
| 460 | // Copy relevant attributes |
| 461 | NewGlobal->setUnnamedAddr(G.getUnnamedAddr()); |
| 462 | if (G.getAlign()) { |
| 463 | NewGlobal->setAlignment(G.getAlign()); |
| 464 | } |
| 465 | |
| 466 | if (G.hasInitializer()) { |
| 467 | Constant *Init = G.getInitializer(); |
| 468 | Constant *NewInit = |
| 469 | transformInitializer(Init, OrigType, FlattenedType: FattenedArrayType, Ctx); |
| 470 | NewGlobal->setInitializer(NewInit); |
| 471 | } |
| 472 | GlobalMap[&G] = NewGlobal; |
| 473 | } |
| 474 | } |
| 475 | |
| 476 | static bool flattenArrays(Module &M) { |
| 477 | bool MadeChange = false; |
| 478 | SmallDenseMap<GlobalVariable *, GlobalVariable *> GlobalMap; |
| 479 | flattenGlobalArrays(M, GlobalMap); |
| 480 | DXILFlattenArraysVisitor Impl(GlobalMap); |
| 481 | for (auto &F : make_early_inc_range(Range: M.functions())) { |
| 482 | if (F.isDeclaration()) |
| 483 | continue; |
| 484 | MadeChange |= Impl.visit(F); |
| 485 | } |
| 486 | for (auto &[Old, New] : GlobalMap) { |
| 487 | Old->replaceAllUsesWith(V: New); |
| 488 | Old->eraseFromParent(); |
| 489 | MadeChange = true; |
| 490 | } |
| 491 | return MadeChange; |
| 492 | } |
| 493 | |
| 494 | PreservedAnalyses DXILFlattenArrays::run(Module &M, ModuleAnalysisManager &) { |
| 495 | bool MadeChanges = flattenArrays(M); |
| 496 | if (!MadeChanges) |
| 497 | return PreservedAnalyses::all(); |
| 498 | PreservedAnalyses PA; |
| 499 | return PA; |
| 500 | } |
| 501 | |
| 502 | bool DXILFlattenArraysLegacy::runOnModule(Module &M) { |
| 503 | return flattenArrays(M); |
| 504 | } |
| 505 | |
| 506 | char DXILFlattenArraysLegacy::ID = 0; |
| 507 | |
| 508 | INITIALIZE_PASS_BEGIN(DXILFlattenArraysLegacy, DEBUG_TYPE, |
| 509 | "DXIL Array Flattener" , false, false) |
| 510 | INITIALIZE_PASS_END(DXILFlattenArraysLegacy, DEBUG_TYPE, "DXIL Array Flattener" , |
| 511 | false, false) |
| 512 | |
| 513 | ModulePass *llvm::createDXILFlattenArraysLegacyPass() { |
| 514 | return new DXILFlattenArraysLegacy(); |
| 515 | } |
| 516 | |