1 //===- MVEGatherScatterLowering.cpp - Gather/Scatter lowering -------------===// 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 /// This pass custom lowers llvm.gather and llvm.scatter instructions to 10 /// arm.mve.gather and arm.mve.scatter intrinsics, optimising the code to 11 /// produce a better final result as we go. 12 // 13 //===----------------------------------------------------------------------===// 14 15 #include "ARM.h" 16 #include "ARMBaseInstrInfo.h" 17 #include "ARMSubtarget.h" 18 #include "llvm/Analysis/TargetTransformInfo.h" 19 #include "llvm/CodeGen/TargetLowering.h" 20 #include "llvm/CodeGen/TargetPassConfig.h" 21 #include "llvm/CodeGen/TargetSubtargetInfo.h" 22 #include "llvm/InitializePasses.h" 23 #include "llvm/IR/BasicBlock.h" 24 #include "llvm/IR/Constant.h" 25 #include "llvm/IR/Constants.h" 26 #include "llvm/IR/DerivedTypes.h" 27 #include "llvm/IR/Function.h" 28 #include "llvm/IR/InstrTypes.h" 29 #include "llvm/IR/Instruction.h" 30 #include "llvm/IR/Instructions.h" 31 #include "llvm/IR/IntrinsicInst.h" 32 #include "llvm/IR/Intrinsics.h" 33 #include "llvm/IR/IntrinsicsARM.h" 34 #include "llvm/IR/IRBuilder.h" 35 #include "llvm/IR/PatternMatch.h" 36 #include "llvm/IR/Type.h" 37 #include "llvm/IR/Value.h" 38 #include "llvm/Pass.h" 39 #include "llvm/Support/Casting.h" 40 #include <algorithm> 41 #include <cassert> 42 43 using namespace llvm; 44 45 #define DEBUG_TYPE "mve-gather-scatter-lowering" 46 47 cl::opt<bool> EnableMaskedGatherScatters( 48 "enable-arm-maskedgatscat", cl::Hidden, cl::init(false), 49 cl::desc("Enable the generation of masked gathers and scatters")); 50 51 namespace { 52 53 class MVEGatherScatterLowering : public FunctionPass { 54 public: 55 static char ID; // Pass identification, replacement for typeid 56 57 explicit MVEGatherScatterLowering() : FunctionPass(ID) { 58 initializeMVEGatherScatterLoweringPass(*PassRegistry::getPassRegistry()); 59 } 60 61 bool runOnFunction(Function &F) override; 62 63 StringRef getPassName() const override { 64 return "MVE gather/scatter lowering"; 65 } 66 67 void getAnalysisUsage(AnalysisUsage &AU) const override { 68 AU.setPreservesCFG(); 69 AU.addRequired<TargetPassConfig>(); 70 FunctionPass::getAnalysisUsage(AU); 71 } 72 73 private: 74 // Check this is a valid gather with correct alignment 75 bool isLegalTypeAndAlignment(unsigned NumElements, unsigned ElemSize, 76 unsigned Alignment); 77 // Check whether Ptr is hidden behind a bitcast and look through it 78 void lookThroughBitcast(Value *&Ptr); 79 // Check for a getelementptr and deduce base and offsets from it, on success 80 // returning the base directly and the offsets indirectly using the Offsets 81 // argument 82 Value *checkGEP(Value *&Offsets, Type *Ty, Value *Ptr, IRBuilder<> &Builder); 83 // Compute the scale of this gather/scatter instruction 84 int computeScale(unsigned GEPElemSize, unsigned MemoryElemSize); 85 86 bool lowerGather(IntrinsicInst *I); 87 // Create a gather from a base + vector of offsets 88 Value *tryCreateMaskedGatherOffset(IntrinsicInst *I, Value *Ptr, 89 Instruction *&Root, IRBuilder<> &Builder); 90 // Create a gather from a vector of pointers 91 Value *tryCreateMaskedGatherBase(IntrinsicInst *I, Value *Ptr, 92 IRBuilder<> &Builder); 93 94 bool lowerScatter(IntrinsicInst *I); 95 // Create a scatter to a base + vector of offsets 96 Value *tryCreateMaskedScatterOffset(IntrinsicInst *I, Value *Ptr, 97 IRBuilder<> &Builder); 98 // Create a scatter to a vector of pointers 99 Value *tryCreateMaskedScatterBase(IntrinsicInst *I, Value *Ptr, 100 IRBuilder<> &Builder); 101 }; 102 103 } // end anonymous namespace 104 105 char MVEGatherScatterLowering::ID = 0; 106 107 INITIALIZE_PASS(MVEGatherScatterLowering, DEBUG_TYPE, 108 "MVE gather/scattering lowering pass", false, false) 109 110 Pass *llvm::createMVEGatherScatterLoweringPass() { 111 return new MVEGatherScatterLowering(); 112 } 113 114 bool MVEGatherScatterLowering::isLegalTypeAndAlignment(unsigned NumElements, 115 unsigned ElemSize, 116 unsigned Alignment) { 117 if (((NumElements == 4 && 118 (ElemSize == 32 || ElemSize == 16 || ElemSize == 8)) || 119 (NumElements == 8 && (ElemSize == 16 || ElemSize == 8)) || 120 (NumElements == 16 && ElemSize == 8)) && 121 ElemSize / 8 <= Alignment) 122 return true; 123 LLVM_DEBUG(dbgs() << "masked gathers/scatters: instruction does not have " 124 << "valid alignment or vector type \n"); 125 return false; 126 } 127 128 Value *MVEGatherScatterLowering::checkGEP(Value *&Offsets, Type *Ty, Value *Ptr, 129 IRBuilder<> &Builder) { 130 GetElementPtrInst *GEP = dyn_cast<GetElementPtrInst>(Ptr); 131 if (!GEP) { 132 LLVM_DEBUG( 133 dbgs() << "masked gathers/scatters: no getelementpointer found\n"); 134 return nullptr; 135 } 136 LLVM_DEBUG(dbgs() << "masked gathers/scatters: getelementpointer found." 137 << " Looking at intrinsic for base + vector of offsets\n"); 138 Value *GEPPtr = GEP->getPointerOperand(); 139 if (GEPPtr->getType()->isVectorTy()) { 140 return nullptr; 141 } 142 if (GEP->getNumOperands() != 2) { 143 LLVM_DEBUG(dbgs() << "masked gathers/scatters: getelementptr with too many" 144 << " operands. Expanding.\n"); 145 return nullptr; 146 } 147 Offsets = GEP->getOperand(1); 148 // Paranoid check whether the number of parallel lanes is the same 149 assert(Ty->getVectorNumElements() == 150 Offsets->getType()->getVectorNumElements()); 151 // Only <N x i32> offsets can be integrated into an arm gather, any smaller 152 // type would have to be sign extended by the gep - and arm gathers can only 153 // zero extend. Additionally, the offsets do have to originate from a zext of 154 // a vector with element types smaller or equal the type of the gather we're 155 // looking at 156 if (Offsets->getType()->getScalarSizeInBits() != 32) 157 return nullptr; 158 if (ZExtInst *ZextOffs = dyn_cast<ZExtInst>(Offsets)) 159 Offsets = ZextOffs->getOperand(0); 160 else if (!(Offsets->getType()->getVectorNumElements() == 4 && 161 Offsets->getType()->getScalarSizeInBits() == 32)) 162 return nullptr; 163 164 if (Ty != Offsets->getType()) { 165 if ((Ty->getScalarSizeInBits() < 166 Offsets->getType()->getScalarSizeInBits())) { 167 LLVM_DEBUG(dbgs() << "masked gathers/scatters: no correct offset type." 168 << " Can't create intrinsic.\n"); 169 return nullptr; 170 } else { 171 Offsets = Builder.CreateZExt( 172 Offsets, VectorType::getInteger(cast<VectorType>(Ty))); 173 } 174 } 175 // If none of the checks failed, return the gep's base pointer 176 LLVM_DEBUG(dbgs() << "masked gathers/scatters: found correct offsets\n"); 177 return GEPPtr; 178 } 179 180 void MVEGatherScatterLowering::lookThroughBitcast(Value *&Ptr) { 181 // Look through bitcast instruction if #elements is the same 182 if (auto *BitCast = dyn_cast<BitCastInst>(Ptr)) { 183 Type *BCTy = BitCast->getType(); 184 Type *BCSrcTy = BitCast->getOperand(0)->getType(); 185 if (BCTy->getVectorNumElements() == BCSrcTy->getVectorNumElements()) { 186 LLVM_DEBUG( 187 dbgs() << "masked gathers/scatters: looking through bitcast\n"); 188 Ptr = BitCast->getOperand(0); 189 } 190 } 191 } 192 193 int MVEGatherScatterLowering::computeScale(unsigned GEPElemSize, 194 unsigned MemoryElemSize) { 195 // This can be a 32bit load/store scaled by 4, a 16bit load/store scaled by 2, 196 // or a 8bit, 16bit or 32bit load/store scaled by 1 197 if (GEPElemSize == 32 && MemoryElemSize == 32) 198 return 2; 199 else if (GEPElemSize == 16 && MemoryElemSize == 16) 200 return 1; 201 else if (GEPElemSize == 8) 202 return 0; 203 LLVM_DEBUG(dbgs() << "masked gathers/scatters: incorrect scale. Can't " 204 << "create intrinsic\n"); 205 return -1; 206 } 207 208 bool MVEGatherScatterLowering::lowerGather(IntrinsicInst *I) { 209 using namespace PatternMatch; 210 LLVM_DEBUG(dbgs() << "masked gathers: checking transform preconditions\n"); 211 212 // @llvm.masked.gather.*(Ptrs, alignment, Mask, Src0) 213 // Attempt to turn the masked gather in I into a MVE intrinsic 214 // Potentially optimising the addressing modes as we do so. 215 Type *Ty = I->getType(); 216 Value *Ptr = I->getArgOperand(0); 217 unsigned Alignment = cast<ConstantInt>(I->getArgOperand(1))->getZExtValue(); 218 Value *Mask = I->getArgOperand(2); 219 Value *PassThru = I->getArgOperand(3); 220 221 if (!isLegalTypeAndAlignment(Ty->getVectorNumElements(), 222 Ty->getScalarSizeInBits(), Alignment)) 223 return false; 224 lookThroughBitcast(Ptr); 225 assert(Ptr->getType()->isVectorTy() && "Unexpected pointer type"); 226 227 IRBuilder<> Builder(I->getContext()); 228 Builder.SetInsertPoint(I); 229 Builder.SetCurrentDebugLocation(I->getDebugLoc()); 230 231 Instruction *Root = I; 232 Value *Load = tryCreateMaskedGatherOffset(I, Ptr, Root, Builder); 233 if (!Load) 234 Load = tryCreateMaskedGatherBase(I, Ptr, Builder); 235 if (!Load) 236 return false; 237 238 if (!isa<UndefValue>(PassThru) && !match(PassThru, m_Zero())) { 239 LLVM_DEBUG(dbgs() << "masked gathers: found non-trivial passthru - " 240 << "creating select\n"); 241 Load = Builder.CreateSelect(Mask, Load, PassThru); 242 } 243 244 Root->replaceAllUsesWith(Load); 245 Root->eraseFromParent(); 246 if (Root != I) 247 // If this was an extending gather, we need to get rid of the sext/zext 248 // sext/zext as well as of the gather itself 249 I->eraseFromParent(); 250 LLVM_DEBUG(dbgs() << "masked gathers: successfully built masked gather\n"); 251 return true; 252 } 253 254 Value *MVEGatherScatterLowering::tryCreateMaskedGatherBase( 255 IntrinsicInst *I, Value *Ptr, IRBuilder<> &Builder) { 256 using namespace PatternMatch; 257 Type *Ty = I->getType(); 258 LLVM_DEBUG(dbgs() << "masked gathers: loading from vector of pointers\n"); 259 if (Ty->getVectorNumElements() != 4 || Ty->getScalarSizeInBits() != 32) 260 // Can't build an intrinsic for this 261 return nullptr; 262 Value *Mask = I->getArgOperand(2); 263 if (match(Mask, m_One())) 264 return Builder.CreateIntrinsic(Intrinsic::arm_mve_vldr_gather_base, 265 {Ty, Ptr->getType()}, 266 {Ptr, Builder.getInt32(0)}); 267 else 268 return Builder.CreateIntrinsic( 269 Intrinsic::arm_mve_vldr_gather_base_predicated, 270 {Ty, Ptr->getType(), Mask->getType()}, 271 {Ptr, Builder.getInt32(0), Mask}); 272 } 273 274 Value *MVEGatherScatterLowering::tryCreateMaskedGatherOffset( 275 IntrinsicInst *I, Value *Ptr, Instruction *&Root, IRBuilder<> &Builder) { 276 using namespace PatternMatch; 277 278 Type *OriginalTy = I->getType(); 279 Type *ResultTy = OriginalTy; 280 281 unsigned Unsigned = 1; 282 // The size of the gather was already checked in isLegalTypeAndAlignment; 283 // if it was not a full vector width an appropriate extend should follow. 284 auto *Extend = Root; 285 if (OriginalTy->getPrimitiveSizeInBits() < 128) { 286 // Only transform gathers with exactly one use 287 if (!I->hasOneUse()) 288 return nullptr; 289 290 // The correct root to replace is the not the CallInst itself, but the 291 // instruction which extends it 292 Extend = cast<Instruction>(*I->users().begin()); 293 if (isa<SExtInst>(Extend)) { 294 Unsigned = 0; 295 } else if (!isa<ZExtInst>(Extend)) { 296 LLVM_DEBUG(dbgs() << "masked gathers: extend needed but not provided. " 297 << "Expanding\n"); 298 return nullptr; 299 } 300 LLVM_DEBUG(dbgs() << "masked gathers: found an extending gather\n"); 301 ResultTy = Extend->getType(); 302 // The final size of the gather must be a full vector width 303 if (ResultTy->getPrimitiveSizeInBits() != 128) { 304 LLVM_DEBUG(dbgs() << "masked gathers: extending from the wrong type. " 305 << "Expanding\n"); 306 return nullptr; 307 } 308 } 309 310 Value *Offsets; 311 Value *BasePtr = checkGEP(Offsets, ResultTy, Ptr, Builder); 312 if (!BasePtr) 313 return nullptr; 314 315 int Scale = computeScale( 316 BasePtr->getType()->getPointerElementType()->getPrimitiveSizeInBits(), 317 OriginalTy->getScalarSizeInBits()); 318 if (Scale == -1) 319 return nullptr; 320 Root = Extend; 321 322 Value *Mask = I->getArgOperand(2); 323 if (!match(Mask, m_One())) 324 return Builder.CreateIntrinsic( 325 Intrinsic::arm_mve_vldr_gather_offset_predicated, 326 {ResultTy, BasePtr->getType(), Offsets->getType(), Mask->getType()}, 327 {BasePtr, Offsets, Builder.getInt32(OriginalTy->getScalarSizeInBits()), 328 Builder.getInt32(Scale), Builder.getInt32(Unsigned), Mask}); 329 else 330 return Builder.CreateIntrinsic( 331 Intrinsic::arm_mve_vldr_gather_offset, 332 {ResultTy, BasePtr->getType(), Offsets->getType()}, 333 {BasePtr, Offsets, Builder.getInt32(OriginalTy->getScalarSizeInBits()), 334 Builder.getInt32(Scale), Builder.getInt32(Unsigned)}); 335 } 336 337 bool MVEGatherScatterLowering::lowerScatter(IntrinsicInst *I) { 338 using namespace PatternMatch; 339 LLVM_DEBUG(dbgs() << "masked scatters: checking transform preconditions\n"); 340 341 // @llvm.masked.scatter.*(data, ptrs, alignment, mask) 342 // Attempt to turn the masked scatter in I into a MVE intrinsic 343 // Potentially optimising the addressing modes as we do so. 344 Value *Input = I->getArgOperand(0); 345 Value *Ptr = I->getArgOperand(1); 346 unsigned Alignment = cast<ConstantInt>(I->getArgOperand(2))->getZExtValue(); 347 Type *Ty = Input->getType(); 348 349 if (!isLegalTypeAndAlignment(Ty->getVectorNumElements(), 350 Ty->getScalarSizeInBits(), Alignment)) 351 return false; 352 lookThroughBitcast(Ptr); 353 assert(Ptr->getType()->isVectorTy() && "Unexpected pointer type"); 354 355 IRBuilder<> Builder(I->getContext()); 356 Builder.SetInsertPoint(I); 357 Builder.SetCurrentDebugLocation(I->getDebugLoc()); 358 359 Value *Store = tryCreateMaskedScatterOffset(I, Ptr, Builder); 360 if (!Store) 361 Store = tryCreateMaskedScatterBase(I, Ptr, Builder); 362 if (!Store) 363 return false; 364 365 LLVM_DEBUG(dbgs() << "masked scatters: successfully built masked scatter\n"); 366 I->replaceAllUsesWith(Store); 367 I->eraseFromParent(); 368 return true; 369 } 370 371 Value *MVEGatherScatterLowering::tryCreateMaskedScatterBase( 372 IntrinsicInst *I, Value *Ptr, IRBuilder<> &Builder) { 373 using namespace PatternMatch; 374 Value *Input = I->getArgOperand(0); 375 Value *Mask = I->getArgOperand(3); 376 Type *Ty = Input->getType(); 377 // Only QR variants allow truncating 378 if (!(Ty->getVectorNumElements() == 4 && Ty->getScalarSizeInBits() == 32)) { 379 // Can't build an intrinsic for this 380 return nullptr; 381 } 382 // int_arm_mve_vstr_scatter_base(_predicated) addr, offset, data(, mask) 383 LLVM_DEBUG(dbgs() << "masked scatters: storing to a vector of pointers\n"); 384 if (match(Mask, m_One())) 385 return Builder.CreateIntrinsic(Intrinsic::arm_mve_vstr_scatter_base, 386 {Ptr->getType(), Input->getType()}, 387 {Ptr, Builder.getInt32(0), Input}); 388 else 389 return Builder.CreateIntrinsic( 390 Intrinsic::arm_mve_vstr_scatter_base_predicated, 391 {Ptr->getType(), Input->getType(), Mask->getType()}, 392 {Ptr, Builder.getInt32(0), Input, Mask}); 393 } 394 395 Value *MVEGatherScatterLowering::tryCreateMaskedScatterOffset( 396 IntrinsicInst *I, Value *Ptr, IRBuilder<> &Builder) { 397 using namespace PatternMatch; 398 Value *Input = I->getArgOperand(0); 399 Value *Mask = I->getArgOperand(3); 400 Type *InputTy = Input->getType(); 401 Type *MemoryTy = InputTy; 402 LLVM_DEBUG(dbgs() << "masked scatters: getelementpointer found. Storing" 403 << " to base + vector of offsets\n"); 404 // If the input has been truncated, try to integrate that trunc into the 405 // scatter instruction (we don't care about alignment here) 406 if (TruncInst *Trunc = dyn_cast<TruncInst>(Input)) { 407 Value *PreTrunc = Trunc->getOperand(0); 408 Type *PreTruncTy = PreTrunc->getType(); 409 if (PreTruncTy->getPrimitiveSizeInBits() == 128) { 410 Input = PreTrunc; 411 InputTy = PreTruncTy; 412 } 413 } 414 if (InputTy->getPrimitiveSizeInBits() != 128) { 415 LLVM_DEBUG( 416 dbgs() << "masked scatters: cannot create scatters for non-standard" 417 << " input types. Expanding.\n"); 418 return nullptr; 419 } 420 421 Value *Offsets; 422 Value *BasePtr = checkGEP(Offsets, InputTy, Ptr, Builder); 423 if (!BasePtr) 424 return nullptr; 425 int Scale = computeScale( 426 BasePtr->getType()->getPointerElementType()->getPrimitiveSizeInBits(), 427 MemoryTy->getScalarSizeInBits()); 428 if (Scale == -1) 429 return nullptr; 430 431 if (!match(Mask, m_One())) 432 return Builder.CreateIntrinsic( 433 Intrinsic::arm_mve_vstr_scatter_offset_predicated, 434 {BasePtr->getType(), Offsets->getType(), Input->getType(), 435 Mask->getType()}, 436 {BasePtr, Offsets, Input, 437 Builder.getInt32(MemoryTy->getScalarSizeInBits()), 438 Builder.getInt32(Scale), Mask}); 439 else 440 return Builder.CreateIntrinsic( 441 Intrinsic::arm_mve_vstr_scatter_offset, 442 {BasePtr->getType(), Offsets->getType(), Input->getType()}, 443 {BasePtr, Offsets, Input, 444 Builder.getInt32(MemoryTy->getScalarSizeInBits()), 445 Builder.getInt32(Scale)}); 446 } 447 448 bool MVEGatherScatterLowering::runOnFunction(Function &F) { 449 if (!EnableMaskedGatherScatters) 450 return false; 451 auto &TPC = getAnalysis<TargetPassConfig>(); 452 auto &TM = TPC.getTM<TargetMachine>(); 453 auto *ST = &TM.getSubtarget<ARMSubtarget>(F); 454 if (!ST->hasMVEIntegerOps()) 455 return false; 456 SmallVector<IntrinsicInst *, 4> Gathers; 457 SmallVector<IntrinsicInst *, 4> Scatters; 458 for (BasicBlock &BB : F) { 459 for (Instruction &I : BB) { 460 IntrinsicInst *II = dyn_cast<IntrinsicInst>(&I); 461 if (II && II->getIntrinsicID() == Intrinsic::masked_gather) 462 Gathers.push_back(II); 463 else if (II && II->getIntrinsicID() == Intrinsic::masked_scatter) 464 Scatters.push_back(II); 465 } 466 } 467 468 bool Changed = false; 469 for (IntrinsicInst *I : Gathers) 470 Changed |= lowerGather(I); 471 for (IntrinsicInst *I : Scatters) 472 Changed |= lowerScatter(I); 473 474 return Changed; 475 } 476