LLVM 24.0.0git
ScalarizeMaskedMemIntrin.cpp
Go to the documentation of this file.
1//===- ScalarizeMaskedMemIntrin.cpp - Scalarize unsupported masked mem ----===//
2// intrinsics
3//
4// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
5// See https://llvm.org/LICENSE.txt for license information.
6// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
7//
8//===----------------------------------------------------------------------===//
9//
10// This pass replaces masked memory intrinsics - when unsupported by the target
11// - with a chain of basic blocks, that deal with the elements one-by-one if the
12// appropriate mask bit is set.
13//
14//===----------------------------------------------------------------------===//
15
17#include "llvm/ADT/Twine.h"
21#include "llvm/IR/BasicBlock.h"
22#include "llvm/IR/Constant.h"
23#include "llvm/IR/Constants.h"
25#include "llvm/IR/Dominators.h"
26#include "llvm/IR/Function.h"
27#include "llvm/IR/IRBuilder.h"
28#include "llvm/IR/Instruction.h"
31#include "llvm/IR/Metadata.h"
33#include "llvm/IR/Type.h"
34#include "llvm/IR/Value.h"
36#include "llvm/Pass.h"
40#include <cassert>
41#include <optional>
42
43using namespace llvm;
44
45#define DEBUG_TYPE "scalarize-masked-mem-intrin"
46
47namespace {
48
49class ScalarizeMaskedMemIntrinLegacyPass : public FunctionPass {
50public:
51 static char ID; // Pass identification, replacement for typeid
52
53 explicit ScalarizeMaskedMemIntrinLegacyPass() : FunctionPass(ID) {
56 }
57
58 bool runOnFunction(Function &F) override;
59
60 StringRef getPassName() const override {
61 return "Scalarize Masked Memory Intrinsics";
62 }
63
64 void getAnalysisUsage(AnalysisUsage &AU) const override {
67 }
68};
69
70} // end anonymous namespace
71
72static bool optimizeBlock(BasicBlock &BB, bool &ModifiedDT,
73 const TargetTransformInfo &TTI, const DataLayout &DL,
74 bool HasBranchDivergence, DomTreeUpdater *DTU);
75static bool optimizeCallInst(CallInst *CI, bool &ModifiedDT,
77 const DataLayout &DL, bool HasBranchDivergence,
78 DomTreeUpdater *DTU);
79
80char ScalarizeMaskedMemIntrinLegacyPass::ID = 0;
81
82INITIALIZE_PASS_BEGIN(ScalarizeMaskedMemIntrinLegacyPass, DEBUG_TYPE,
83 "Scalarize unsupported masked memory intrinsics", false,
84 false)
87INITIALIZE_PASS_END(ScalarizeMaskedMemIntrinLegacyPass, DEBUG_TYPE,
88 "Scalarize unsupported masked memory intrinsics", false,
89 false)
90
92 return new ScalarizeMaskedMemIntrinLegacyPass();
93}
94
95static bool isConstantIntVector(Value *Mask) {
97 if (!C)
98 return false;
99
100 unsigned NumElts = cast<FixedVectorType>(Mask->getType())->getNumElements();
101 for (unsigned i = 0; i != NumElts; ++i) {
102 Constant *CElt = C->getAggregateElement(i);
103 if (!CElt || !isa<ConstantInt>(CElt))
104 return false;
105 }
106
107 return true;
108}
109
110static unsigned adjustForEndian(const DataLayout &DL, unsigned VectorWidth,
111 unsigned Idx) {
112 return DL.isBigEndian() ? VectorWidth - 1 - Idx : Idx;
113}
114
115static void copyMemCacheHint(Instruction &Dest, const Instruction &Source,
116 unsigned SourcePtrOperand,
117 unsigned DestPtrOperand) {
118 MDNode *CacheHint = Source.getMetadata(LLVMContext::MD_mem_cache_hint);
119 // These intrinsics have a single memory operand.
120 if (!CacheHint || CacheHint->getNumOperands() != 2)
121 return;
122
123 auto *OperandNo = mdconst::extract<ConstantInt>(CacheHint->getOperand(0));
124 if (OperandNo->getZExtValue() != SourcePtrOperand)
125 return;
126
127 Metadata *DestOperandNo = ConstantAsMetadata::get(
128 ConstantInt::get(Type::getInt32Ty(Dest.getContext()), DestPtrOperand));
129 Dest.setMetadata(LLVMContext::MD_mem_cache_hint,
130 MDNode::get(Dest.getContext(),
131 {DestOperandNo, CacheHint->getOperand(1)}));
132}
133
135 Instruction &Dest, const Instruction &Source, const DataLayout &DL,
136 unsigned SourcePtrOperand, unsigned DestPtrOperand, Type *AccessType,
137 bool IsWholeAccess, std::optional<size_t> ByteOffset) {
138 // Only propagate metadata that is valid on each constituent memory access.
139 // In particular, do not copy metadata whose meaning is tied to the call,
140 // such as !prof or !callsite.
141 Dest.copyMetadata(Source,
142 {LLVMContext::MD_nontemporal,
143 LLVMContext::MD_mem_parallel_loop_access,
144 LLVMContext::MD_access_group, LLVMContext::MD_annotation,
145 LLVMContext::MD_nosanitize, LLVMContext::MD_mmra});
146
147 AAMDNodes AANodes = Source.getAAMetadata();
148 if (IsWholeAccess)
149 Dest.setAAMetadata(AANodes);
150 else if (ByteOffset)
151 Dest.setAAMetadata(AANodes.adjustForAccess(*ByteOffset, AccessType, DL));
152 else {
153 // The packed address is runtime-dependent. The other AA metadata remains
154 // applicable, but !tbaa.struct cannot be adjusted to a known byte range.
155 AANodes.TBAAStruct = nullptr;
156 Dest.setAAMetadata(AANodes);
157 }
158 copyMemCacheHint(Dest, Source, SourcePtrOperand, DestPtrOperand);
159}
160
162 const Instruction &Source,
163 const DataLayout &DL,
164 unsigned SourcePtrOperand,
165 std::optional<size_t> ByteOffset) {
166 copyMetadataForMemoryAccess(Dest, Source, DL, SourcePtrOperand,
167 Dest.getPointerOperandIndex(), Dest.getType(),
168 Dest.getType() == Source.getType(), ByteOffset);
169
170 // !range applies element-wise to vectors, so the same range describes each
171 // scalar result. The other metadata here also describes the loaded result.
172 Dest.copyMetadata(Source, {LLVMContext::MD_fpmath, LLVMContext::MD_range,
173 LLVMContext::MD_invariant_load});
174}
175
177 const Instruction &Source,
178 const DataLayout &DL,
179 unsigned SourcePtrOperand,
180 std::optional<size_t> ByteOffset) {
182 Dest, Source, DL, SourcePtrOperand, Dest.getPointerOperandIndex(),
183 Dest.getValueOperand()->getType(),
184 Dest.getValueOperand()->getType() == Source.getOperand(0)->getType(),
185 ByteOffset);
186}
187
188// Translate a masked load intrinsic like
189// <16 x i32 > @llvm.masked.load( <16 x i32>* %addr,
190// <16 x i1> %mask, <16 x i32> %passthru)
191// to a chain of basic blocks, with loading element one-by-one if
192// the appropriate mask bit is set
193//
194// %1 = bitcast i8* %addr to i32*
195// %2 = extractelement <16 x i1> %mask, i32 0
196// br i1 %2, label %cond.load, label %else
197//
198// cond.load: ; preds = %0
199// %3 = getelementptr i32* %1, i32 0
200// %4 = load i32* %3
201// %5 = insertelement <16 x i32> %passthru, i32 %4, i32 0
202// br label %else
203//
204// else: ; preds = %0, %cond.load
205// %res.phi.else = phi <16 x i32> [ %5, %cond.load ], [ poison, %0 ]
206// %6 = extractelement <16 x i1> %mask, i32 1
207// br i1 %6, label %cond.load1, label %else2
208//
209// cond.load1: ; preds = %else
210// %7 = getelementptr i32* %1, i32 1
211// %8 = load i32* %7
212// %9 = insertelement <16 x i32> %res.phi.else, i32 %8, i32 1
213// br label %else2
214//
215// else2: ; preds = %else, %cond.load1
216// %res.phi.else3 = phi <16 x i32> [ %9, %cond.load1 ], [ %res.phi.else, %else
217// ] %10 = extractelement <16 x i1> %mask, i32 2 br i1 %10, label %cond.load4,
218// label %else5
219//
220static void scalarizeMaskedLoad(const DataLayout &DL, bool HasBranchDivergence,
221 CallInst *CI, DomTreeUpdater *DTU,
222 bool &ModifiedDT) {
223 Value *Ptr = CI->getArgOperand(0);
224 Value *Mask = CI->getArgOperand(1);
225 Value *Src0 = CI->getArgOperand(2);
226
227 const Align AlignVal = CI->getParamAlign(0).valueOrOne();
228 VectorType *VecType = cast<FixedVectorType>(CI->getType());
229
230 Type *EltTy = VecType->getElementType();
231
232 Instruction *InsertPt = CI;
233 IRBuilder<> Builder(InsertPt);
234 BasicBlock *IfBlock = CI->getParent();
235
236 // Short-cut if the mask is all-true.
237 if (isa<Constant>(Mask) && cast<Constant>(Mask)->isAllOnesValue()) {
238 LoadInst *NewI = Builder.CreateAlignedLoad(VecType, Ptr, AlignVal);
239 copyMetadataForScalarizedLoad(*NewI, *CI, DL, /*SourcePtrOperand=*/0,
240 std::nullopt);
241 NewI->takeName(CI);
242 CI->replaceAllUsesWith(NewI);
243 CI->eraseFromParent();
244 return;
245 }
246
247 // Adjust alignment for the scalar instruction.
248 const Align AdjustedAlignVal =
249 commonAlignment(AlignVal, EltTy->getPrimitiveSizeInBits() / 8);
250 unsigned VectorWidth = cast<FixedVectorType>(VecType)->getNumElements();
251
252 // The result vector
253 Value *VResult = Src0;
254
255 if (isConstantIntVector(Mask)) {
256 for (unsigned Idx = 0; Idx < VectorWidth; ++Idx) {
257 if (cast<Constant>(Mask)->getAggregateElement(Idx)->isNullValue())
258 continue;
259 Value *Gep = Builder.CreateConstInBoundsGEP1_32(EltTy, Ptr, Idx);
260 LoadInst *Load = Builder.CreateAlignedLoad(EltTy, Gep, AdjustedAlignVal);
262 *Load, *CI, DL, /*SourcePtrOperand=*/0,
263 Idx * DL.getTypeAllocSize(EltTy).getFixedValue());
264 VResult = Builder.CreateInsertElement(VResult, Load, Idx);
265 }
266 CI->replaceAllUsesWith(VResult);
267 CI->eraseFromParent();
268 return;
269 }
270
271 // Optimize the case where the "masked load" is a predicated load - that is,
272 // where the mask is the splat of a non-constant scalar boolean. In that case,
273 // use that splated value as the guard on a conditional vector load.
274 if (isSplatValue(Mask, /*Index=*/0)) {
275 Value *Predicate = Builder.CreateExtractElement(Mask, uint64_t(0ull),
276 Mask->getName() + ".first");
277 // We mark the branch weights as explicitly unknown given they would only
278 // be derivable from the mask which we do not have VP information for.
279 Instruction *ThenTerm =
280 SplitBlockAndInsertIfThen(Predicate, InsertPt, /*Unreachable=*/false,
282 *CI->getFunction(), DEBUG_TYPE),
283 DTU);
284
285 BasicBlock *CondBlock = ThenTerm->getParent();
286 CondBlock->setName("cond.load");
287 Builder.SetInsertPoint(CondBlock->getTerminator());
288 LoadInst *Load = Builder.CreateAlignedLoad(VecType, Ptr, AlignVal,
289 CI->getName() + ".cond.load");
290 copyMetadataForScalarizedLoad(*Load, *CI, DL, /*SourcePtrOperand=*/0,
291 std::nullopt);
292
293 BasicBlock *PostLoad = ThenTerm->getSuccessor(0);
294 Builder.SetInsertPoint(PostLoad, PostLoad->begin());
295 PHINode *Phi = Builder.CreatePHI(VecType, /*NumReservedValues=*/2);
296 Phi->addIncoming(Load, CondBlock);
297 Phi->addIncoming(Src0, IfBlock);
298 Phi->takeName(CI);
299
300 CI->replaceAllUsesWith(Phi);
301 CI->eraseFromParent();
302 ModifiedDT = true;
303 return;
304 }
305 // If the mask is not v1i1, use scalar bit test operations. This generates
306 // better results on X86 at least. However, don't do this on GPUs and other
307 // machines with divergence, as there each i1 needs a vector register.
308 Value *SclrMask = nullptr;
309 if (VectorWidth != 1 && !HasBranchDivergence) {
310 Type *SclrMaskTy = Builder.getIntNTy(VectorWidth);
311 SclrMask = Builder.CreateBitCast(Mask, SclrMaskTy, "scalar_mask");
312 }
313
314 for (unsigned Idx = 0; Idx < VectorWidth; ++Idx) {
315 // Fill the "else" block, created in the previous iteration
316 //
317 // %res.phi.else3 = phi <16 x i32> [ %11, %cond.load1 ], [ %res.phi.else,
318 // %else ] %mask_1 = and i16 %scalar_mask, i32 1 << Idx %cond = icmp ne i16
319 // %mask_1, 0 br i1 %mask_1, label %cond.load, label %else
320 //
321 // On GPUs, use
322 // %cond = extrectelement %mask, Idx
323 // instead
325 if (SclrMask != nullptr) {
326 Value *Mask = Builder.getInt(APInt::getOneBitSet(
327 VectorWidth, adjustForEndian(DL, VectorWidth, Idx)));
328 Predicate = Builder.CreateICmpNE(Builder.CreateAnd(SclrMask, Mask),
329 Builder.getIntN(VectorWidth, 0));
330 } else {
331 Predicate = Builder.CreateExtractElement(Mask, Idx);
332 }
333
334 // Create "cond" block
335 //
336 // %EltAddr = getelementptr i32* %1, i32 0
337 // %Elt = load i32* %EltAddr
338 // VResult = insertelement <16 x i32> VResult, i32 %Elt, i32 Idx
339 //
340 // We mark the branch weights as explicitly unknown given they would only
341 // be derivable from the mask which we do not have VP information for.
342 Instruction *ThenTerm =
343 SplitBlockAndInsertIfThen(Predicate, InsertPt, /*Unreachable=*/false,
345 *CI->getFunction(), DEBUG_TYPE),
346 DTU);
347
348 BasicBlock *CondBlock = ThenTerm->getParent();
349 CondBlock->setName("cond.load");
350
351 Builder.SetInsertPoint(CondBlock->getTerminator());
352 Value *Gep = Builder.CreateConstInBoundsGEP1_32(EltTy, Ptr, Idx);
353 LoadInst *Load = Builder.CreateAlignedLoad(EltTy, Gep, AdjustedAlignVal);
355 *Load, *CI, DL, /*SourcePtrOperand=*/0,
356 Idx * DL.getTypeAllocSize(EltTy).getFixedValue());
357 Value *NewVResult = Builder.CreateInsertElement(VResult, Load, Idx);
358
359 // Create "else" block, fill it in the next iteration
360 BasicBlock *NewIfBlock = ThenTerm->getSuccessor(0);
361 NewIfBlock->setName("else");
362 BasicBlock *PrevIfBlock = IfBlock;
363 IfBlock = NewIfBlock;
364
365 // Create the phi to join the new and previous value.
366 Builder.SetInsertPoint(NewIfBlock, NewIfBlock->begin());
367 PHINode *Phi = Builder.CreatePHI(VecType, 2, "res.phi.else");
368 Phi->addIncoming(NewVResult, CondBlock);
369 Phi->addIncoming(VResult, PrevIfBlock);
370 VResult = Phi;
371 }
372
373 CI->replaceAllUsesWith(VResult);
374 CI->eraseFromParent();
375
376 ModifiedDT = true;
377}
378
379// Translate a masked store intrinsic, like
380// void @llvm.masked.store(<16 x i32> %src, <16 x i32>* %addr,
381// <16 x i1> %mask)
382// to a chain of basic blocks, that stores element one-by-one if
383// the appropriate mask bit is set
384//
385// %1 = bitcast i8* %addr to i32*
386// %2 = extractelement <16 x i1> %mask, i32 0
387// br i1 %2, label %cond.store, label %else
388//
389// cond.store: ; preds = %0
390// %3 = extractelement <16 x i32> %val, i32 0
391// %4 = getelementptr i32* %1, i32 0
392// store i32 %3, i32* %4
393// br label %else
394//
395// else: ; preds = %0, %cond.store
396// %5 = extractelement <16 x i1> %mask, i32 1
397// br i1 %5, label %cond.store1, label %else2
398//
399// cond.store1: ; preds = %else
400// %6 = extractelement <16 x i32> %val, i32 1
401// %7 = getelementptr i32* %1, i32 1
402// store i32 %6, i32* %7
403// br label %else2
404// . . .
405static void scalarizeMaskedStore(const DataLayout &DL, bool HasBranchDivergence,
406 CallInst *CI, DomTreeUpdater *DTU,
407 bool &ModifiedDT) {
408 Value *Src = CI->getArgOperand(0);
409 Value *Ptr = CI->getArgOperand(1);
410 Value *Mask = CI->getArgOperand(2);
411
412 const Align AlignVal = CI->getParamAlign(1).valueOrOne();
413 auto *VecType = cast<VectorType>(Src->getType());
414
415 Type *EltTy = VecType->getElementType();
416
417 Instruction *InsertPt = CI;
418 IRBuilder<> Builder(InsertPt);
419
420 // Short-cut if the mask is all-true.
421 if (isa<Constant>(Mask) && cast<Constant>(Mask)->isAllOnesValue()) {
422 StoreInst *Store = Builder.CreateAlignedStore(Src, Ptr, AlignVal);
423 Store->takeName(CI);
424 copyMetadataForScalarizedStore(*Store, *CI, DL, /*SourcePtrOperand=*/1,
425 std::nullopt);
426 // This is a one-to-one replacement, so the assignment link remains valid.
427 Store->copyMetadata(*CI, LLVMContext::MD_DIAssignID);
428 CI->eraseFromParent();
429 return;
430 }
431
432 // Adjust alignment for the scalar instruction.
433 const Align AdjustedAlignVal =
434 commonAlignment(AlignVal, EltTy->getPrimitiveSizeInBits() / 8);
435 unsigned VectorWidth = cast<FixedVectorType>(VecType)->getNumElements();
436
437 if (isConstantIntVector(Mask)) {
438 for (unsigned Idx = 0; Idx < VectorWidth; ++Idx) {
439 if (cast<Constant>(Mask)->getAggregateElement(Idx)->isNullValue())
440 continue;
441 Value *OneElt = Builder.CreateExtractElement(Src, Idx);
442 Value *Gep = Builder.CreateConstInBoundsGEP1_32(EltTy, Ptr, Idx);
444 Builder.CreateAlignedStore(OneElt, Gep, AdjustedAlignVal);
446 *Store, *CI, DL, /*SourcePtrOperand=*/1,
447 Idx * DL.getTypeAllocSize(EltTy).getFixedValue());
448 }
449 CI->eraseFromParent();
450 return;
451 }
452
453 // Optimize the case where the "masked store" is a predicated store - that is,
454 // when the mask is the splat of a non-constant scalar boolean. In that case,
455 // optimize to a conditional store.
456 if (isSplatValue(Mask, /*Index=*/0)) {
457 Value *Predicate = Builder.CreateExtractElement(Mask, uint64_t(0ull),
458 Mask->getName() + ".first");
459 // We mark the branch weights as explicitly unknown given they would only
460 // be derivable from the mask which we do not have VP information for.
461 Instruction *ThenTerm =
462 SplitBlockAndInsertIfThen(Predicate, InsertPt, /*Unreachable=*/false,
464 *CI->getFunction(), DEBUG_TYPE),
465 DTU);
466 BasicBlock *CondBlock = ThenTerm->getParent();
467 CondBlock->setName("cond.store");
468 Builder.SetInsertPoint(CondBlock->getTerminator());
469
470 StoreInst *Store = Builder.CreateAlignedStore(Src, Ptr, AlignVal);
471 Store->takeName(CI);
472 copyMetadataForScalarizedStore(*Store, *CI, DL, /*SourcePtrOperand=*/1,
473 std::nullopt);
474 // This is a one-to-one replacement, so the assignment link remains valid.
475 Store->copyMetadata(*CI, LLVMContext::MD_DIAssignID);
476
477 CI->eraseFromParent();
478 ModifiedDT = true;
479 return;
480 }
481
482 // If the mask is not v1i1, use scalar bit test operations. This generates
483 // better results on X86 at least. However, don't do this on GPUs or other
484 // machines with branch divergence, as there each i1 takes up a register.
485 Value *SclrMask = nullptr;
486 if (VectorWidth != 1 && !HasBranchDivergence) {
487 Type *SclrMaskTy = Builder.getIntNTy(VectorWidth);
488 SclrMask = Builder.CreateBitCast(Mask, SclrMaskTy, "scalar_mask");
489 }
490
491 for (unsigned Idx = 0; Idx < VectorWidth; ++Idx) {
492 // Fill the "else" block, created in the previous iteration
493 //
494 // %mask_1 = and i16 %scalar_mask, i32 1 << Idx
495 // %cond = icmp ne i16 %mask_1, 0
496 // br i1 %mask_1, label %cond.store, label %else
497 //
498 // On GPUs, use
499 // %cond = extrectelement %mask, Idx
500 // instead
502 if (SclrMask != nullptr) {
503 Value *Mask = Builder.getInt(APInt::getOneBitSet(
504 VectorWidth, adjustForEndian(DL, VectorWidth, Idx)));
505 Predicate = Builder.CreateICmpNE(Builder.CreateAnd(SclrMask, Mask),
506 Builder.getIntN(VectorWidth, 0));
507 } else {
508 Predicate = Builder.CreateExtractElement(Mask, Idx);
509 }
510
511 // Create "cond" block
512 //
513 // %OneElt = extractelement <16 x i32> %Src, i32 Idx
514 // %EltAddr = getelementptr i32* %1, i32 0
515 // %store i32 %OneElt, i32* %EltAddr
516 //
517 // We mark the branch weights as explicitly unknown given they would only
518 // be derivable from the mask which we do not have VP information for.
519 Instruction *ThenTerm =
520 SplitBlockAndInsertIfThen(Predicate, InsertPt, /*Unreachable=*/false,
522 *CI->getFunction(), DEBUG_TYPE),
523 DTU);
524
525 BasicBlock *CondBlock = ThenTerm->getParent();
526 CondBlock->setName("cond.store");
527
528 Builder.SetInsertPoint(CondBlock->getTerminator());
529 Value *OneElt = Builder.CreateExtractElement(Src, Idx);
530 Value *Gep = Builder.CreateConstInBoundsGEP1_32(EltTy, Ptr, Idx);
532 Builder.CreateAlignedStore(OneElt, Gep, AdjustedAlignVal);
534 *Store, *CI, DL, /*SourcePtrOperand=*/1,
535 Idx * DL.getTypeAllocSize(EltTy).getFixedValue());
536
537 // Create "else" block, fill it in the next iteration
538 BasicBlock *NewIfBlock = ThenTerm->getSuccessor(0);
539 NewIfBlock->setName("else");
540
541 Builder.SetInsertPoint(NewIfBlock, NewIfBlock->begin());
542 }
543 CI->eraseFromParent();
544
545 ModifiedDT = true;
546}
547
548// Translate a masked gather intrinsic like
549// <16 x i32 > @llvm.masked.gather.v16i32( <16 x i32*> %Ptrs, i32 4,
550// <16 x i1> %Mask, <16 x i32> %Src)
551// to a chain of basic blocks, with loading element one-by-one if
552// the appropriate mask bit is set
553//
554// %Ptrs = getelementptr i32, i32* %base, <16 x i64> %ind
555// %Mask0 = extractelement <16 x i1> %Mask, i32 0
556// br i1 %Mask0, label %cond.load, label %else
557//
558// cond.load:
559// %Ptr0 = extractelement <16 x i32*> %Ptrs, i32 0
560// %Load0 = load i32, i32* %Ptr0, align 4
561// %Res0 = insertelement <16 x i32> poison, i32 %Load0, i32 0
562// br label %else
563//
564// else:
565// %res.phi.else = phi <16 x i32>[%Res0, %cond.load], [poison, %0]
566// %Mask1 = extractelement <16 x i1> %Mask, i32 1
567// br i1 %Mask1, label %cond.load1, label %else2
568//
569// cond.load1:
570// %Ptr1 = extractelement <16 x i32*> %Ptrs, i32 1
571// %Load1 = load i32, i32* %Ptr1, align 4
572// %Res1 = insertelement <16 x i32> %res.phi.else, i32 %Load1, i32 1
573// br label %else2
574// . . .
575// %Result = select <16 x i1> %Mask, <16 x i32> %res.phi.select, <16 x i32> %Src
576// ret <16 x i32> %Result
578 bool HasBranchDivergence, CallInst *CI,
579 DomTreeUpdater *DTU, bool &ModifiedDT) {
580 Value *Ptrs = CI->getArgOperand(0);
581 Value *Mask = CI->getArgOperand(1);
582 Value *Src0 = CI->getArgOperand(2);
583
584 auto *VecType = cast<FixedVectorType>(CI->getType());
585 Type *EltTy = VecType->getElementType();
586
587 Instruction *InsertPt = CI;
588 IRBuilder<> Builder(InsertPt);
589 BasicBlock *IfBlock = CI->getParent();
590 Align AlignVal = CI->getParamAlign(0).valueOrOne();
591
592 // The result vector
593 Value *VResult = Src0;
594 unsigned VectorWidth = VecType->getNumElements();
595
596 // Shorten the way if the mask is a vector of constants.
597 if (isConstantIntVector(Mask)) {
598 for (unsigned Idx = 0; Idx < VectorWidth; ++Idx) {
599 if (cast<Constant>(Mask)->getAggregateElement(Idx)->isNullValue())
600 continue;
601 Value *Ptr = Builder.CreateExtractElement(Ptrs, Idx, "Ptr" + Twine(Idx));
602 LoadInst *Load =
603 Builder.CreateAlignedLoad(EltTy, Ptr, AlignVal, "Load" + Twine(Idx));
604 copyMetadataForScalarizedLoad(*Load, *CI, DL, /*SourcePtrOperand=*/0,
605 /*ByteOffset=*/0);
606 VResult =
607 Builder.CreateInsertElement(VResult, Load, Idx, "Res" + Twine(Idx));
608 }
609 CI->replaceAllUsesWith(VResult);
610 CI->eraseFromParent();
611 return;
612 }
613
614 // If the mask is not v1i1, use scalar bit test operations. This generates
615 // better results on X86 at least. However, don't do this on GPUs or other
616 // machines with branch divergence, as there, each i1 takes up a register.
617 Value *SclrMask = nullptr;
618 if (VectorWidth != 1 && !HasBranchDivergence) {
619 Type *SclrMaskTy = Builder.getIntNTy(VectorWidth);
620 SclrMask = Builder.CreateBitCast(Mask, SclrMaskTy, "scalar_mask");
621 }
622
623 for (unsigned Idx = 0; Idx < VectorWidth; ++Idx) {
624 // Fill the "else" block, created in the previous iteration
625 //
626 // %Mask1 = and i16 %scalar_mask, i32 1 << Idx
627 // %cond = icmp ne i16 %mask_1, 0
628 // br i1 %Mask1, label %cond.load, label %else
629 //
630 // On GPUs, use
631 // %cond = extrectelement %mask, Idx
632 // instead
633
635 if (SclrMask != nullptr) {
636 Value *Mask = Builder.getInt(APInt::getOneBitSet(
637 VectorWidth, adjustForEndian(DL, VectorWidth, Idx)));
638 Predicate = Builder.CreateICmpNE(Builder.CreateAnd(SclrMask, Mask),
639 Builder.getIntN(VectorWidth, 0));
640 } else {
641 Predicate = Builder.CreateExtractElement(Mask, Idx, "Mask" + Twine(Idx));
642 }
643
644 // Create "cond" block
645 //
646 // %EltAddr = getelementptr i32* %1, i32 0
647 // %Elt = load i32* %EltAddr
648 // VResult = insertelement <16 x i32> VResult, i32 %Elt, i32 Idx
649 //
650 // We mark the branch weights as explicitly unknown given they would only
651 // be derivable from the mask which we do not have VP information for.
652 Instruction *ThenTerm =
653 SplitBlockAndInsertIfThen(Predicate, InsertPt, /*Unreachable=*/false,
655 *CI->getFunction(), DEBUG_TYPE),
656 DTU);
657
658 BasicBlock *CondBlock = ThenTerm->getParent();
659 CondBlock->setName("cond.load");
660
661 Builder.SetInsertPoint(CondBlock->getTerminator());
662 Value *Ptr = Builder.CreateExtractElement(Ptrs, Idx, "Ptr" + Twine(Idx));
663 LoadInst *Load =
664 Builder.CreateAlignedLoad(EltTy, Ptr, AlignVal, "Load" + Twine(Idx));
665 copyMetadataForScalarizedLoad(*Load, *CI, DL, /*SourcePtrOperand=*/0,
666 /*ByteOffset=*/0);
667 Value *NewVResult =
668 Builder.CreateInsertElement(VResult, Load, Idx, "Res" + Twine(Idx));
669
670 // Create "else" block, fill it in the next iteration
671 BasicBlock *NewIfBlock = ThenTerm->getSuccessor(0);
672 NewIfBlock->setName("else");
673 BasicBlock *PrevIfBlock = IfBlock;
674 IfBlock = NewIfBlock;
675
676 // Create the phi to join the new and previous value.
677 Builder.SetInsertPoint(NewIfBlock, NewIfBlock->begin());
678 PHINode *Phi = Builder.CreatePHI(VecType, 2, "res.phi.else");
679 Phi->addIncoming(NewVResult, CondBlock);
680 Phi->addIncoming(VResult, PrevIfBlock);
681 VResult = Phi;
682 }
683
684 CI->replaceAllUsesWith(VResult);
685 CI->eraseFromParent();
686
687 ModifiedDT = true;
688}
689
690// Translate a masked scatter intrinsic, like
691// void @llvm.masked.scatter.v16i32(<16 x i32> %Src, <16 x i32*>* %Ptrs, i32 4,
692// <16 x i1> %Mask)
693// to a chain of basic blocks, that stores element one-by-one if
694// the appropriate mask bit is set.
695//
696// %Ptrs = getelementptr i32, i32* %ptr, <16 x i64> %ind
697// %Mask0 = extractelement <16 x i1> %Mask, i32 0
698// br i1 %Mask0, label %cond.store, label %else
699//
700// cond.store:
701// %Elt0 = extractelement <16 x i32> %Src, i32 0
702// %Ptr0 = extractelement <16 x i32*> %Ptrs, i32 0
703// store i32 %Elt0, i32* %Ptr0, align 4
704// br label %else
705//
706// else:
707// %Mask1 = extractelement <16 x i1> %Mask, i32 1
708// br i1 %Mask1, label %cond.store1, label %else2
709//
710// cond.store1:
711// %Elt1 = extractelement <16 x i32> %Src, i32 1
712// %Ptr1 = extractelement <16 x i32*> %Ptrs, i32 1
713// store i32 %Elt1, i32* %Ptr1, align 4
714// br label %else2
715// . . .
717 bool HasBranchDivergence, CallInst *CI,
718 DomTreeUpdater *DTU, bool &ModifiedDT) {
719 Value *Src = CI->getArgOperand(0);
720 Value *Ptrs = CI->getArgOperand(1);
721 Value *Mask = CI->getArgOperand(2);
722
723 auto *SrcFVTy = cast<FixedVectorType>(Src->getType());
724
725 assert(
726 isa<VectorType>(Ptrs->getType()) &&
727 isa<PointerType>(cast<VectorType>(Ptrs->getType())->getElementType()) &&
728 "Vector of pointers is expected in masked scatter intrinsic");
729
730 Instruction *InsertPt = CI;
731 IRBuilder<> Builder(InsertPt);
732
733 Align AlignVal = CI->getParamAlign(1).valueOrOne();
734 unsigned VectorWidth = SrcFVTy->getNumElements();
735
736 // Shorten the way if the mask is a vector of constants.
737 if (isConstantIntVector(Mask)) {
738 for (unsigned Idx = 0; Idx < VectorWidth; ++Idx) {
739 if (cast<Constant>(Mask)->getAggregateElement(Idx)->isNullValue())
740 continue;
741 Value *OneElt =
742 Builder.CreateExtractElement(Src, Idx, "Elt" + Twine(Idx));
743 Value *Ptr = Builder.CreateExtractElement(Ptrs, Idx, "Ptr" + Twine(Idx));
744 StoreInst *Store = Builder.CreateAlignedStore(OneElt, Ptr, AlignVal);
746 /*SourcePtrOperand=*/1,
747 /*ByteOffset=*/0);
748 }
749 CI->eraseFromParent();
750 return;
751 }
752
753 // If the mask is not v1i1, use scalar bit test operations. This generates
754 // better results on X86 at least.
755 Value *SclrMask = nullptr;
756 if (VectorWidth != 1 && !HasBranchDivergence) {
757 Type *SclrMaskTy = Builder.getIntNTy(VectorWidth);
758 SclrMask = Builder.CreateBitCast(Mask, SclrMaskTy, "scalar_mask");
759 }
760
761 for (unsigned Idx = 0; Idx < VectorWidth; ++Idx) {
762 // Fill the "else" block, created in the previous iteration
763 //
764 // %Mask1 = and i16 %scalar_mask, i32 1 << Idx
765 // %cond = icmp ne i16 %mask_1, 0
766 // br i1 %Mask1, label %cond.store, label %else
767 //
768 // On GPUs, use
769 // %cond = extrectelement %mask, Idx
770 // instead
772 if (SclrMask != nullptr) {
773 Value *Mask = Builder.getInt(APInt::getOneBitSet(
774 VectorWidth, adjustForEndian(DL, VectorWidth, Idx)));
775 Predicate = Builder.CreateICmpNE(Builder.CreateAnd(SclrMask, Mask),
776 Builder.getIntN(VectorWidth, 0));
777 } else {
778 Predicate = Builder.CreateExtractElement(Mask, Idx, "Mask" + Twine(Idx));
779 }
780
781 // Create "cond" block
782 //
783 // %Elt1 = extractelement <16 x i32> %Src, i32 1
784 // %Ptr1 = extractelement <16 x i32*> %Ptrs, i32 1
785 // %store i32 %Elt1, i32* %Ptr1
786 //
787 // We mark the branch weights as explicitly unknown given they would only
788 // be derivable from the mask which we do not have VP information for.
789 Instruction *ThenTerm =
790 SplitBlockAndInsertIfThen(Predicate, InsertPt, /*Unreachable=*/false,
792 *CI->getFunction(), DEBUG_TYPE),
793 DTU);
794
795 BasicBlock *CondBlock = ThenTerm->getParent();
796 CondBlock->setName("cond.store");
797
798 Builder.SetInsertPoint(CondBlock->getTerminator());
799 Value *OneElt = Builder.CreateExtractElement(Src, Idx, "Elt" + Twine(Idx));
800 Value *Ptr = Builder.CreateExtractElement(Ptrs, Idx, "Ptr" + Twine(Idx));
801 StoreInst *Store = Builder.CreateAlignedStore(OneElt, Ptr, AlignVal);
803 /*SourcePtrOperand=*/1,
804 /*ByteOffset=*/0);
805
806 // Create "else" block, fill it in the next iteration
807 BasicBlock *NewIfBlock = ThenTerm->getSuccessor(0);
808 NewIfBlock->setName("else");
809
810 Builder.SetInsertPoint(NewIfBlock, NewIfBlock->begin());
811 }
812 CI->eraseFromParent();
813
814 ModifiedDT = true;
815}
816
818 bool HasBranchDivergence, CallInst *CI,
819 DomTreeUpdater *DTU, bool &ModifiedDT) {
820 Value *Ptr = CI->getArgOperand(0);
821 Value *Mask = CI->getArgOperand(1);
822 Value *PassThru = CI->getArgOperand(2);
823 Align Alignment = CI->getParamAlign(0).valueOrOne();
824
825 auto *VecType = cast<FixedVectorType>(CI->getType());
826
827 Type *EltTy = VecType->getElementType();
828
829 Instruction *InsertPt = CI;
830 IRBuilder<> Builder(InsertPt);
831 BasicBlock *IfBlock = CI->getParent();
832
833 unsigned VectorWidth = VecType->getNumElements();
834
835 // The result vector
836 Value *VResult = PassThru;
837
838 // Adjust alignment for the scalar instruction.
839 const Align AdjustedAlignment =
840 commonAlignment(Alignment, EltTy->getPrimitiveSizeInBits() / 8);
841
842 // Shorten the way if the mask is a vector of constants.
843 // Create a build_vector pattern, with loads/poisons as necessary and then
844 // shuffle blend with the pass through value.
845 if (isConstantIntVector(Mask)) {
846 unsigned MemIndex = 0;
847 VResult = PoisonValue::get(VecType);
848 SmallVector<int, 16> ShuffleMask(VectorWidth, PoisonMaskElem);
849 for (unsigned Idx = 0; Idx < VectorWidth; ++Idx) {
850 Value *InsertElt;
851 if (cast<Constant>(Mask)->getAggregateElement(Idx)->isNullValue()) {
852 InsertElt = PoisonValue::get(EltTy);
853 ShuffleMask[Idx] = Idx + VectorWidth;
854 } else {
855 Value *NewPtr =
856 Builder.CreateConstInBoundsGEP1_32(EltTy, Ptr, MemIndex);
857 LoadInst *Load = Builder.CreateAlignedLoad(
858 EltTy, NewPtr, AdjustedAlignment, "Load" + Twine(Idx));
860 *Load, *CI, DL, /*SourcePtrOperand=*/0,
861 MemIndex * DL.getTypeAllocSize(EltTy).getFixedValue());
862 InsertElt = Load;
863 ShuffleMask[Idx] = Idx;
864 ++MemIndex;
865 }
866 VResult = Builder.CreateInsertElement(VResult, InsertElt, Idx,
867 "Res" + Twine(Idx));
868 }
869 VResult = Builder.CreateShuffleVector(VResult, PassThru, ShuffleMask);
870 CI->replaceAllUsesWith(VResult);
871 CI->eraseFromParent();
872 return;
873 }
874
875 // If the mask is not v1i1, use scalar bit test operations. This generates
876 // better results on X86 at least. However, don't do this on GPUs or other
877 // machines with branch divergence, as there, each i1 takes up a register.
878 Value *SclrMask = nullptr;
879 if (VectorWidth != 1 && !HasBranchDivergence) {
880 Type *SclrMaskTy = Builder.getIntNTy(VectorWidth);
881 SclrMask = Builder.CreateBitCast(Mask, SclrMaskTy, "scalar_mask");
882 }
883
884 for (unsigned Idx = 0; Idx < VectorWidth; ++Idx) {
885 // Fill the "else" block, created in the previous iteration
886 //
887 // %res.phi.else3 = phi <16 x i32> [ %11, %cond.load1 ], [ %res.phi.else,
888 // %else ] %mask_1 = extractelement <16 x i1> %mask, i32 Idx br i1 %mask_1,
889 // label %cond.load, label %else
890 //
891 // On GPUs, use
892 // %cond = extrectelement %mask, Idx
893 // instead
894
896 if (SclrMask != nullptr) {
897 Value *Mask = Builder.getInt(APInt::getOneBitSet(
898 VectorWidth, adjustForEndian(DL, VectorWidth, Idx)));
899 Predicate = Builder.CreateICmpNE(Builder.CreateAnd(SclrMask, Mask),
900 Builder.getIntN(VectorWidth, 0));
901 } else {
902 Predicate = Builder.CreateExtractElement(Mask, Idx, "Mask" + Twine(Idx));
903 }
904
905 // Create "cond" block
906 //
907 // %EltAddr = getelementptr i32* %1, i32 0
908 // %Elt = load i32* %EltAddr
909 // VResult = insertelement <16 x i32> VResult, i32 %Elt, i32 Idx
910 //
911 // We mark the branch weights as explicitly unknown given they would only
912 // be derivable from the mask which we do not have VP information for.
913 Instruction *ThenTerm =
914 SplitBlockAndInsertIfThen(Predicate, InsertPt, /*Unreachable=*/false,
916 *CI->getFunction(), DEBUG_TYPE),
917 DTU);
918
919 BasicBlock *CondBlock = ThenTerm->getParent();
920 CondBlock->setName("cond.load");
921
922 Builder.SetInsertPoint(CondBlock->getTerminator());
923 LoadInst *Load = Builder.CreateAlignedLoad(EltTy, Ptr, AdjustedAlignment);
924 copyMetadataForScalarizedLoad(*Load, *CI, DL, /*SourcePtrOperand=*/0,
925 std::nullopt);
926 Value *NewVResult = Builder.CreateInsertElement(VResult, Load, Idx);
927
928 // Move the pointer if there are more blocks to come.
929 Value *NewPtr;
930 if ((Idx + 1) != VectorWidth)
931 NewPtr = Builder.CreateConstInBoundsGEP1_32(EltTy, Ptr, 1);
932
933 // Create "else" block, fill it in the next iteration
934 BasicBlock *NewIfBlock = ThenTerm->getSuccessor(0);
935 NewIfBlock->setName("else");
936 BasicBlock *PrevIfBlock = IfBlock;
937 IfBlock = NewIfBlock;
938
939 // Create the phi to join the new and previous value.
940 Builder.SetInsertPoint(NewIfBlock, NewIfBlock->begin());
941 PHINode *ResultPhi = Builder.CreatePHI(VecType, 2, "res.phi.else");
942 ResultPhi->addIncoming(NewVResult, CondBlock);
943 ResultPhi->addIncoming(VResult, PrevIfBlock);
944 VResult = ResultPhi;
945
946 // Add a PHI for the pointer if this isn't the last iteration.
947 if ((Idx + 1) != VectorWidth) {
948 PHINode *PtrPhi = Builder.CreatePHI(Ptr->getType(), 2, "ptr.phi.else");
949 PtrPhi->addIncoming(NewPtr, CondBlock);
950 PtrPhi->addIncoming(Ptr, PrevIfBlock);
951 Ptr = PtrPhi;
952 }
953 }
954
955 CI->replaceAllUsesWith(VResult);
956 CI->eraseFromParent();
957
958 ModifiedDT = true;
959}
960
962 bool HasBranchDivergence, CallInst *CI,
963 DomTreeUpdater *DTU,
964 bool &ModifiedDT) {
965 Value *Src = CI->getArgOperand(0);
966 Value *Ptr = CI->getArgOperand(1);
967 Value *Mask = CI->getArgOperand(2);
968 Align Alignment = CI->getParamAlign(1).valueOrOne();
969
970 auto *VecType = cast<FixedVectorType>(Src->getType());
971
972 Instruction *InsertPt = CI;
973 IRBuilder<> Builder(InsertPt);
974 BasicBlock *IfBlock = CI->getParent();
975
976 Type *EltTy = VecType->getElementType();
977
978 // Adjust alignment for the scalar instruction.
979 const Align AdjustedAlignment =
980 commonAlignment(Alignment, EltTy->getPrimitiveSizeInBits() / 8);
981
982 unsigned VectorWidth = VecType->getNumElements();
983
984 // Shorten the way if the mask is a vector of constants.
985 if (isConstantIntVector(Mask)) {
986 unsigned MemIndex = 0;
987 for (unsigned Idx = 0; Idx < VectorWidth; ++Idx) {
988 if (cast<Constant>(Mask)->getAggregateElement(Idx)->isNullValue())
989 continue;
990 Value *OneElt =
991 Builder.CreateExtractElement(Src, Idx, "Elt" + Twine(Idx));
992 Value *NewPtr = Builder.CreateConstInBoundsGEP1_32(EltTy, Ptr, MemIndex);
994 Builder.CreateAlignedStore(OneElt, NewPtr, AdjustedAlignment);
996 *Store, *CI, DL, /*SourcePtrOperand=*/1,
997 MemIndex * DL.getTypeAllocSize(EltTy).getFixedValue());
998 ++MemIndex;
999 }
1000 CI->eraseFromParent();
1001 return;
1002 }
1003
1004 // If the mask is not v1i1, use scalar bit test operations. This generates
1005 // better results on X86 at least. However, don't do this on GPUs or other
1006 // machines with branch divergence, as there, each i1 takes up a register.
1007 Value *SclrMask = nullptr;
1008 if (VectorWidth != 1 && !HasBranchDivergence) {
1009 Type *SclrMaskTy = Builder.getIntNTy(VectorWidth);
1010 SclrMask = Builder.CreateBitCast(Mask, SclrMaskTy, "scalar_mask");
1011 }
1012
1013 for (unsigned Idx = 0; Idx < VectorWidth; ++Idx) {
1014 // Fill the "else" block, created in the previous iteration
1015 //
1016 // %mask_1 = extractelement <16 x i1> %mask, i32 Idx
1017 // br i1 %mask_1, label %cond.store, label %else
1018 //
1019 // On GPUs, use
1020 // %cond = extrectelement %mask, Idx
1021 // instead
1023 if (SclrMask != nullptr) {
1024 Value *Mask = Builder.getInt(APInt::getOneBitSet(
1025 VectorWidth, adjustForEndian(DL, VectorWidth, Idx)));
1026 Predicate = Builder.CreateICmpNE(Builder.CreateAnd(SclrMask, Mask),
1027 Builder.getIntN(VectorWidth, 0));
1028 } else {
1029 Predicate = Builder.CreateExtractElement(Mask, Idx, "Mask" + Twine(Idx));
1030 }
1031
1032 // Create "cond" block
1033 //
1034 // %OneElt = extractelement <16 x i32> %Src, i32 Idx
1035 // %EltAddr = getelementptr i32* %1, i32 0
1036 // %store i32 %OneElt, i32* %EltAddr
1037 //
1038 // We mark the branch weights as explicitly unknown given they would only
1039 // be derivable from the mask which we do not have VP information for.
1040 Instruction *ThenTerm =
1041 SplitBlockAndInsertIfThen(Predicate, InsertPt, /*Unreachable=*/false,
1043 *CI->getFunction(), DEBUG_TYPE),
1044 DTU);
1045
1046 BasicBlock *CondBlock = ThenTerm->getParent();
1047 CondBlock->setName("cond.store");
1048
1049 Builder.SetInsertPoint(CondBlock->getTerminator());
1050 Value *OneElt = Builder.CreateExtractElement(Src, Idx);
1051 StoreInst *Store =
1052 Builder.CreateAlignedStore(OneElt, Ptr, AdjustedAlignment);
1053 copyMetadataForScalarizedStore(*Store, *CI, DL, /*SourcePtrOperand=*/1,
1054 std::nullopt);
1055
1056 // Move the pointer if there are more blocks to come.
1057 Value *NewPtr;
1058 if ((Idx + 1) != VectorWidth)
1059 NewPtr = Builder.CreateConstInBoundsGEP1_32(EltTy, Ptr, 1);
1060
1061 // Create "else" block, fill it in the next iteration
1062 BasicBlock *NewIfBlock = ThenTerm->getSuccessor(0);
1063 NewIfBlock->setName("else");
1064 BasicBlock *PrevIfBlock = IfBlock;
1065 IfBlock = NewIfBlock;
1066
1067 Builder.SetInsertPoint(NewIfBlock, NewIfBlock->begin());
1068
1069 // Add a PHI for the pointer if this isn't the last iteration.
1070 if ((Idx + 1) != VectorWidth) {
1071 PHINode *PtrPhi = Builder.CreatePHI(Ptr->getType(), 2, "ptr.phi.else");
1072 PtrPhi->addIncoming(NewPtr, CondBlock);
1073 PtrPhi->addIncoming(Ptr, PrevIfBlock);
1074 Ptr = PtrPhi;
1075 }
1076 }
1077 CI->eraseFromParent();
1078
1079 ModifiedDT = true;
1080}
1081
1083 DomTreeUpdater *DTU,
1084 bool &ModifiedDT) {
1085 // If we extend histogram to return a result someday (like the updated vector)
1086 // then we'll need to support it here.
1087 assert(CI->getType()->isVoidTy() && "Histogram with non-void return.");
1088 Value *Ptrs = CI->getArgOperand(0);
1089 Value *Inc = CI->getArgOperand(1);
1090 Value *Mask = CI->getArgOperand(2);
1091
1092 auto *AddrType = cast<FixedVectorType>(Ptrs->getType());
1093 Type *EltTy = Inc->getType();
1094
1095 Instruction *InsertPt = CI;
1096 IRBuilder<> Builder(InsertPt);
1097
1098 // FIXME: Do we need to add an alignment parameter to the intrinsic?
1099 unsigned VectorWidth = AddrType->getNumElements();
1100 auto CreateHistogramUpdateValue = [&](IntrinsicInst *CI, Value *Load,
1101 Value *Inc) -> Value * {
1102 Value *UpdateOp;
1103 switch (CI->getIntrinsicID()) {
1104 case Intrinsic::experimental_vector_histogram_add:
1105 UpdateOp = Builder.CreateAdd(Load, Inc);
1106 break;
1107 case Intrinsic::experimental_vector_histogram_uadd_sat:
1108 UpdateOp =
1109 Builder.CreateIntrinsic(Intrinsic::uadd_sat, {EltTy}, {Load, Inc});
1110 break;
1111 case Intrinsic::experimental_vector_histogram_umin:
1112 UpdateOp = Builder.CreateIntrinsic(Intrinsic::umin, {EltTy}, {Load, Inc});
1113 break;
1114 case Intrinsic::experimental_vector_histogram_umax:
1115 UpdateOp = Builder.CreateIntrinsic(Intrinsic::umax, {EltTy}, {Load, Inc});
1116 break;
1117
1118 default:
1119 llvm_unreachable("Unexpected histogram intrinsic");
1120 }
1121 return UpdateOp;
1122 };
1123
1124 // Shorten the way if the mask is a vector of constants.
1125 if (isConstantIntVector(Mask)) {
1126 for (unsigned Idx = 0; Idx < VectorWidth; ++Idx) {
1127 if (cast<Constant>(Mask)->getAggregateElement(Idx)->isNullValue())
1128 continue;
1129 Value *Ptr = Builder.CreateExtractElement(Ptrs, Idx, "Ptr" + Twine(Idx));
1130 LoadInst *Load = Builder.CreateLoad(EltTy, Ptr, "Load" + Twine(Idx));
1131 copyMetadataForScalarizedLoad(*Load, *CI, DL, /*SourcePtrOperand=*/0,
1132 /*ByteOffset=*/0);
1133 Value *Update =
1134 CreateHistogramUpdateValue(cast<IntrinsicInst>(CI), Load, Inc);
1135 StoreInst *Store = Builder.CreateStore(Update, Ptr);
1137 /*SourcePtrOperand=*/0,
1138 /*ByteOffset=*/0);
1139 }
1140 CI->eraseFromParent();
1141 return;
1142 }
1143
1144 for (unsigned Idx = 0; Idx < VectorWidth; ++Idx) {
1145 Value *Predicate =
1146 Builder.CreateExtractElement(Mask, Idx, "Mask" + Twine(Idx));
1147
1148 // We mark the branch weights as explicitly unknown given they would only
1149 // be derivable from the mask which we do not have VP information for.
1150 Instruction *ThenTerm =
1151 SplitBlockAndInsertIfThen(Predicate, InsertPt, /*Unreachable=*/false,
1153 *CI->getFunction(), DEBUG_TYPE),
1154 DTU);
1155
1156 BasicBlock *CondBlock = ThenTerm->getParent();
1157 CondBlock->setName("cond.histogram.update");
1158
1159 Builder.SetInsertPoint(CondBlock->getTerminator());
1160 Value *Ptr = Builder.CreateExtractElement(Ptrs, Idx, "Ptr" + Twine(Idx));
1161 LoadInst *Load = Builder.CreateLoad(EltTy, Ptr, "Load" + Twine(Idx));
1162 copyMetadataForScalarizedLoad(*Load, *CI, DL, /*SourcePtrOperand=*/0,
1163 /*ByteOffset=*/0);
1164 Value *UpdateOp =
1165 CreateHistogramUpdateValue(cast<IntrinsicInst>(CI), Load, Inc);
1166 StoreInst *Store = Builder.CreateStore(UpdateOp, Ptr);
1168 /*SourcePtrOperand=*/0,
1169 /*ByteOffset=*/0);
1170
1171 // Create "else" block, fill it in the next iteration
1172 BasicBlock *NewIfBlock = ThenTerm->getSuccessor(0);
1173 NewIfBlock->setName("else");
1174 Builder.SetInsertPoint(NewIfBlock, NewIfBlock->begin());
1175 }
1176
1177 CI->eraseFromParent();
1178 ModifiedDT = true;
1179}
1180
1182 DominatorTree *DT) {
1183 std::optional<DomTreeUpdater> DTU;
1184 if (DT)
1185 DTU.emplace(DT, DomTreeUpdater::UpdateStrategy::Lazy);
1186
1187 bool EverMadeChange = false;
1188 bool MadeChange = true;
1189 auto &DL = F.getDataLayout();
1190 bool HasBranchDivergence = TTI.hasBranchDivergence(&F);
1191 while (MadeChange) {
1192 MadeChange = false;
1194 bool ModifiedDTOnIteration = false;
1195 MadeChange |= optimizeBlock(BB, ModifiedDTOnIteration, TTI, DL,
1196 HasBranchDivergence, DTU ? &*DTU : nullptr);
1197
1198 // Restart BB iteration if the dominator tree of the Function was changed
1199 if (ModifiedDTOnIteration)
1200 break;
1201 }
1202
1203 EverMadeChange |= MadeChange;
1204 }
1205 return EverMadeChange;
1206}
1207
1208bool ScalarizeMaskedMemIntrinLegacyPass::runOnFunction(Function &F) {
1209 auto &TTI = getAnalysis<TargetTransformInfoWrapperPass>().getTTI(F);
1210 DominatorTree *DT = nullptr;
1211 if (auto *DTWP = getAnalysisIfAvailable<DominatorTreeWrapperPass>())
1212 DT = &DTWP->getDomTree();
1213 return runImpl(F, TTI, DT);
1214}
1215
1216PreservedAnalyses
1227
1228static bool optimizeBlock(BasicBlock &BB, bool &ModifiedDT,
1229 const TargetTransformInfo &TTI, const DataLayout &DL,
1230 bool HasBranchDivergence, DomTreeUpdater *DTU) {
1231 bool MadeChange = false;
1232
1233 BasicBlock::iterator CurInstIterator = BB.begin();
1234 while (CurInstIterator != BB.end()) {
1235 if (CallInst *CI = dyn_cast<CallInst>(&*CurInstIterator++))
1236 MadeChange |=
1237 optimizeCallInst(CI, ModifiedDT, TTI, DL, HasBranchDivergence, DTU);
1238 if (ModifiedDT)
1239 return true;
1240 }
1241
1242 return MadeChange;
1243}
1244
1245static bool optimizeCallInst(CallInst *CI, bool &ModifiedDT,
1246 const TargetTransformInfo &TTI,
1247 const DataLayout &DL, bool HasBranchDivergence,
1248 DomTreeUpdater *DTU) {
1250 if (II) {
1251 // The scalarization code below does not work for scalable vectors.
1252 if (isa<ScalableVectorType>(II->getType()) ||
1253 any_of(II->args(),
1254 [](Value *V) { return isa<ScalableVectorType>(V->getType()); }))
1255 return false;
1256 switch (II->getIntrinsicID()) {
1257 default:
1258 break;
1259 case Intrinsic::experimental_vector_histogram_add:
1260 case Intrinsic::experimental_vector_histogram_uadd_sat:
1261 case Intrinsic::experimental_vector_histogram_umin:
1262 case Intrinsic::experimental_vector_histogram_umax:
1263 if (TTI.isLegalMaskedVectorHistogram(CI->getArgOperand(0)->getType(),
1264 CI->getArgOperand(1)->getType()))
1265 return false;
1266 scalarizeMaskedVectorHistogram(DL, CI, DTU, ModifiedDT);
1267 return true;
1268 case Intrinsic::masked_load:
1269 // Scalarize unsupported vector masked load
1270 if (TTI.isLegalMaskedLoad(
1271 CI->getType(), CI->getParamAlign(0).valueOrOne(),
1273 ->getAddressSpace(),
1277 return false;
1278 scalarizeMaskedLoad(DL, HasBranchDivergence, CI, DTU, ModifiedDT);
1279 return true;
1280 case Intrinsic::masked_store:
1281 if (TTI.isLegalMaskedStore(
1282 CI->getArgOperand(0)->getType(),
1283 CI->getParamAlign(1).valueOrOne(),
1285 ->getAddressSpace(),
1289 return false;
1290 scalarizeMaskedStore(DL, HasBranchDivergence, CI, DTU, ModifiedDT);
1291 return true;
1292 case Intrinsic::masked_gather: {
1293 Align Alignment = CI->getParamAlign(0).valueOrOne();
1294 Type *LoadTy = CI->getType();
1295 if (TTI.isLegalMaskedGather(LoadTy, Alignment) &&
1296 !TTI.forceScalarizeMaskedGather(cast<VectorType>(LoadTy), Alignment))
1297 return false;
1298 scalarizeMaskedGather(DL, HasBranchDivergence, CI, DTU, ModifiedDT);
1299 return true;
1300 }
1301 case Intrinsic::masked_scatter: {
1302 Align Alignment = CI->getParamAlign(1).valueOrOne();
1303 Type *StoreTy = CI->getArgOperand(0)->getType();
1304 if (TTI.isLegalMaskedScatter(StoreTy, Alignment) &&
1305 !TTI.forceScalarizeMaskedScatter(cast<VectorType>(StoreTy),
1306 Alignment))
1307 return false;
1308 scalarizeMaskedScatter(DL, HasBranchDivergence, CI, DTU, ModifiedDT);
1309 return true;
1310 }
1311 case Intrinsic::masked_expandload:
1312 if (TTI.isLegalMaskedExpandLoad(
1313 CI->getType(),
1314 CI->getAttributes().getParamAttrs(0).getAlignment().valueOrOne()))
1315 return false;
1316 scalarizeMaskedExpandLoad(DL, HasBranchDivergence, CI, DTU, ModifiedDT);
1317 return true;
1318 case Intrinsic::masked_compressstore:
1319 if (TTI.isLegalMaskedCompressStore(
1320 CI->getArgOperand(0)->getType(),
1321 CI->getAttributes().getParamAttrs(1).getAlignment().valueOrOne()))
1322 return false;
1323 scalarizeMaskedCompressStore(DL, HasBranchDivergence, CI, DTU,
1324 ModifiedDT);
1325 return true;
1326 }
1327 }
1328
1329 return false;
1330}
assert(UImm &&(UImm !=~static_cast< T >(0)) &&"Invalid immediate!")
unsigned uint64_t
MachineBasicBlock MachineBasicBlock::iterator DebugLoc DL
static GCRegistry::Add< ShadowStackGC > C("shadow-stack", "Very portable GC for uncooperative code generators")
static bool runImpl(MachineFunction &MF)
Definition CFIFixup.cpp:304
This file contains the declarations for the subclasses of Constant, which represent the different fla...
static bool runOnFunction(Function &F, bool PostInlining)
#define DEBUG_TYPE
#define F(x, y, z)
Definition MD5.cpp:54
This file contains the declarations for metadata subclasses.
uint64_t IntrinsicInst * II
#define INITIALIZE_PASS_DEPENDENCY(depName)
Definition PassSupport.h:42
#define INITIALIZE_PASS_END(passName, arg, name, cfg, analysis)
Definition PassSupport.h:44
#define INITIALIZE_PASS_BEGIN(passName, arg, name, cfg, analysis)
Definition PassSupport.h:39
This file contains the declarations for profiling metadata utility functions.
static void scalarizeMaskedExpandLoad(const DataLayout &DL, bool HasBranchDivergence, CallInst *CI, DomTreeUpdater *DTU, bool &ModifiedDT)
static void scalarizeMaskedVectorHistogram(const DataLayout &DL, CallInst *CI, DomTreeUpdater *DTU, bool &ModifiedDT)
static void copyMemCacheHint(Instruction &Dest, const Instruction &Source, unsigned SourcePtrOperand, unsigned DestPtrOperand)
static void copyMetadataForScalarizedStore(StoreInst &Dest, const Instruction &Source, const DataLayout &DL, unsigned SourcePtrOperand, std::optional< size_t > ByteOffset)
static bool optimizeBlock(BasicBlock &BB, bool &ModifiedDT, const TargetTransformInfo &TTI, const DataLayout &DL, bool HasBranchDivergence, DomTreeUpdater *DTU)
static void scalarizeMaskedScatter(const DataLayout &DL, bool HasBranchDivergence, CallInst *CI, DomTreeUpdater *DTU, bool &ModifiedDT)
static unsigned adjustForEndian(const DataLayout &DL, unsigned VectorWidth, unsigned Idx)
static bool optimizeCallInst(CallInst *CI, bool &ModifiedDT, const TargetTransformInfo &TTI, const DataLayout &DL, bool HasBranchDivergence, DomTreeUpdater *DTU)
static void copyMetadataForScalarizedLoad(LoadInst &Dest, const Instruction &Source, const DataLayout &DL, unsigned SourcePtrOperand, std::optional< size_t > ByteOffset)
static void scalarizeMaskedStore(const DataLayout &DL, bool HasBranchDivergence, CallInst *CI, DomTreeUpdater *DTU, bool &ModifiedDT)
static void scalarizeMaskedCompressStore(const DataLayout &DL, bool HasBranchDivergence, CallInst *CI, DomTreeUpdater *DTU, bool &ModifiedDT)
static void scalarizeMaskedGather(const DataLayout &DL, bool HasBranchDivergence, CallInst *CI, DomTreeUpdater *DTU, bool &ModifiedDT)
static void copyMetadataForMemoryAccess(Instruction &Dest, const Instruction &Source, const DataLayout &DL, unsigned SourcePtrOperand, unsigned DestPtrOperand, Type *AccessType, bool IsWholeAccess, std::optional< size_t > ByteOffset)
static bool runImpl(Function &F, const TargetTransformInfo &TTI, DominatorTree *DT)
static bool isConstantIntVector(Value *Mask)
static void scalarizeMaskedLoad(const DataLayout &DL, bool HasBranchDivergence, CallInst *CI, DomTreeUpdater *DTU, bool &ModifiedDT)
This pass exposes codegen information to IR-level passes.
static APInt getOneBitSet(unsigned numBits, unsigned BitNo)
Return an APInt with exactly one bit set in the result.
Definition APInt.h:235
PassT::Result * getCachedResult(IRUnitT &IR) const
Get the cached result of an analysis pass for a given IR unit.
PassT::Result & getResult(IRUnitT &IR, ExtraArgTs... ExtraArgs)
Get the result of an analysis pass for a given IR unit.
Represent the analysis usage information of a pass.
AnalysisUsage & addRequired()
AnalysisUsage & addPreserved()
Add the specified Pass class to the set of analyses preserved by this pass.
LLVM Basic Block Representation.
Definition BasicBlock.h:62
iterator end()
Definition BasicBlock.h:459
iterator begin()
Instruction iterator methods.
Definition BasicBlock.h:446
InstListType::iterator iterator
Instruction iterators...
Definition BasicBlock.h:170
const Instruction * getTerminator() const LLVM_READONLY
Returns the terminator instruction; assumes that the block is well-formed.
Definition BasicBlock.h:237
MaybeAlign getParamAlign(unsigned ArgNo) const
Extract the alignment for a call or parameter (0=unknown).
Value * getArgOperand(unsigned i) const
LLVM_ABI Intrinsic::ID getIntrinsicID() const
Returns the intrinsic ID of the intrinsic called or Intrinsic::not_intrinsic if the called function i...
AttributeList getAttributes() const
Return the attributes for this call.
This class represents a function call, abstracting a target machine's calling convention.
static ConstantAsMetadata * get(Constant *C)
Definition Metadata.h:559
This is an important base class in LLVM.
Definition Constant.h:43
A parsed version of the target data layout string in and methods for querying it.
Definition DataLayout.h:64
Analysis pass which computes a DominatorTree.
Definition Dominators.h:241
Legacy analysis pass which computes a DominatorTree.
Definition Dominators.h:277
Concrete subclass of DominatorTreeBase that is used to compute a normal dominator tree.
Definition Dominators.h:122
FunctionPass class - This class is used to implement most global optimizations.
Definition Pass.h:314
This provides a uniform API for creating instructions and inserting them into a basic block: either a...
Definition IRBuilder.h:2901
LLVM_ABI void setAAMetadata(const AAMDNodes &N)
Sets the AA metadata on this instruction from the AAMDNodes structure.
LLVM_ABI InstListType::iterator eraseFromParent()
This method unlinks 'this' from the containing basic block and deletes it.
LLVM_ABI const Function * getFunction() const
Return the function this instruction belongs to.
LLVM_ABI BasicBlock * getSuccessor(unsigned Idx) const LLVM_READONLY
Return the specified successor. This instruction must be a terminator.
LLVM_ABI void setMetadata(unsigned KindID, MDNode *Node)
Set the metadata of the specified kind to the specified node.
LLVM_ABI void copyMetadata(const Instruction &SrcInst, ArrayRef< unsigned > WL=ArrayRef< unsigned >())
Copy metadata from SrcInst to this instruction.
A wrapper class for inspecting calls to intrinsic functions.
An instruction for reading from memory.
static unsigned getPointerOperandIndex()
Metadata node.
Definition Metadata.h:1092
const MDOperand & getOperand(unsigned I) const
Definition Metadata.h:1448
static MDTuple * get(LLVMContext &Context, ArrayRef< Metadata * > MDs)
Definition Metadata.h:1590
unsigned getNumOperands() const
Return number of MDNode operands.
Definition Metadata.h:1454
Root of the metadata hierarchy.
Definition Metadata.h:64
void addIncoming(Value *V, BasicBlock *BB)
Add an incoming value to the end of the PHI list.
static LLVM_ABI PassRegistry * getPassRegistry()
getPassRegistry - Access the global registry object, which is automatically initialized at applicatio...
static LLVM_ABI PoisonValue * get(Type *T)
Static factory methods - Return an 'poison' object of the specified type.
A set of analyses that are preserved following a run of a transformation pass.
Definition Analysis.h:112
static PreservedAnalyses all()
Construct a special preserved set that preserves all passes.
Definition Analysis.h:118
PreservedAnalyses & preserve()
Mark an analysis as preserved.
Definition Analysis.h:132
This is a 'vector' (really, a variable-sized array), optimized for the case when the array is small.
An instruction for storing to memory.
Value * getValueOperand()
static unsigned getPointerOperandIndex()
Represent a constant reference to a string, i.e.
Definition StringRef.h:56
Analysis pass providing the TargetTransformInfo.
Wrapper pass for TargetTransformInfo.
This pass provides access to the codegen interfaces that are needed for IR-level transformations.
Twine - A lightweight data structure for efficiently representing the concatenation of temporary valu...
Definition Twine.h:82
The instances of the Type class are immutable: once they are created, they are never changed.
Definition Type.h:46
static LLVM_ABI IntegerType * getInt32Ty(LLVMContext &C)
Definition Type.cpp:299
LLVM_ABI TypeSize getPrimitiveSizeInBits() const LLVM_READONLY
Return the basic size of this type if it is a primitive type.
Definition Type.cpp:187
static LLVM_ABI IntegerType * getIntNTy(LLVMContext &C, unsigned N)
Definition Type.cpp:303
bool isVoidTy() const
Return true if this is 'void'.
Definition Type.h:141
LLVM Value Representation.
Definition Value.h:75
Type * getType() const
All values are typed, get the type of this value.
Definition Value.h:257
LLVM_ABI void setName(const Twine &Name)
Change the name of the value.
Definition Value.cpp:394
LLVM_ABI void replaceAllUsesWith(Value *V)
Change all uses of this to point to a new Value.
Definition Value.cpp:553
LLVMContext & getContext() const
All values hold a context through their type.
Definition Value.h:260
LLVM_ABI StringRef getName() const
Return a constant reference to the value's name.
Definition Value.cpp:319
LLVM_ABI void takeName(Value *V)
Transfer the name from V to this value.
Definition Value.cpp:400
const ParentTy * getParent() const
Definition ilist_node.h:34
#define llvm_unreachable(msg)
Marks that the current location is not supposed to be reachable.
std::enable_if_t< detail::IsValidPointer< X, Y >::value, X * > extract(Y &&MD)
Extract a Value from Metadata.
Definition Metadata.h:690
This is an optimization pass for GlobalISel generic memory operations.
decltype(auto) dyn_cast(const From &Val)
dyn_cast<X> - Return the argument parameter cast to the specified type.
Definition Casting.h:643
@ Load
The value being inserted comes from a load (InsertElement only).
@ Store
The extracted value is stored (ExtractElement only).
iterator_range< early_inc_iterator_impl< detail::IterOfRange< RangeT > > > make_early_inc_range(RangeT &&Range)
Make a range that does early increment to allow mutation of the underlying range without disrupting i...
Definition STLExtras.h:649
LLVM_ABI FunctionPass * createScalarizeMaskedMemIntrinLegacyPass()
bool any_of(R &&range, UnaryPredicate P)
Provide wrappers to std::any_of which take ranges instead of having to pass begin/end explicitly.
Definition STLExtras.h:1762
LLVM_ABI bool isSplatValue(const Value *V, int Index=-1, unsigned Depth=0)
Return true if each element of the vector value V is poisoned or equal to every other non-poisoned el...
LLVM_ABI void initializeScalarizeMaskedMemIntrinLegacyPassPass(PassRegistry &)
bool isa(const From &Val)
isa<X> - Return true if the parameter to the template is an instance of one of the template type argu...
Definition Casting.h:547
LLVM_ABI MDNode * getExplicitlyUnknownBranchWeightsIfProfiled(Function &F, StringRef PassName)
Returns a metadata node containing unknown branch weights if the function has an entry count,...
constexpr int PoisonMaskElem
TargetTransformInfo TTI
decltype(auto) cast(const From &Val)
cast<X> - Return the argument parameter cast to the specified type.
Definition Casting.h:559
Align commonAlignment(Align A, uint64_t Offset)
Returns the alignment that satisfies both alignments.
Definition Alignment.h:201
LLVM_ABI Instruction * SplitBlockAndInsertIfThen(Value *Cond, BasicBlock::iterator SplitBefore, bool Unreachable, MDNode *BranchWeights=nullptr, DomTreeUpdater *DTU=nullptr, LoopInfo *LI=nullptr, BasicBlock *ThenBlock=nullptr)
Split the containing block at the specified instruction - everything before SplitBefore stays in the ...
AnalysisManager< Function > FunctionAnalysisManager
Convenience typedef for the Function analysis manager.
A collection of metadata nodes that might be associated with a memory access used by the alias-analys...
Definition Metadata.h:785
MDNode * TBAAStruct
The tag for type-based alias analysis (tbaa struct).
Definition Metadata.h:805
LLVM_ABI AAMDNodes adjustForAccess(unsigned AccessSize)
Create a new AAMDNode for accessing AccessSize bytes of this AAMDNode.
This struct is a compact representation of a valid (non-zero power of two) alignment.
Definition Alignment.h:39
Align valueOrOne() const
For convenience, returns a valid alignment or 1 if undefined.
Definition Alignment.h:130
LLVM_ABI PreservedAnalyses run(Function &F, FunctionAnalysisManager &AM)