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
16#include "llvm/Transforms/Scalar/ScalarizeMaskedMemIntrin.h"
17#include "llvm/ADT/Twine.h"
18#include "llvm/Analysis/DomTreeUpdater.h"
19#include "llvm/Analysis/TargetTransformInfo.h"
20#include "llvm/Analysis/VectorUtils.h"
21#include "llvm/IR/BasicBlock.h"
22#include "llvm/IR/Constant.h"
23#include "llvm/IR/Constants.h"
24#include "llvm/IR/DerivedTypes.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"
29#include "llvm/IR/Instructions.h"
30#include "llvm/IR/IntrinsicInst.h"
31#include "llvm/IR/Metadata.h"
32#include "llvm/IR/ProfDataUtils.h"
33#include "llvm/IR/Type.h"
34#include "llvm/IR/Value.h"
35#include "llvm/InitializePasses.h"
36#include "llvm/Pass.h"
37#include "llvm/Support/Casting.h"
38#include "llvm/Transforms/Scalar.h"
39#include "llvm/Transforms/Utils/BasicBlockUtils.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) {
54 initializeScalarizeMaskedMemIntrinLegacyPassPass(
55 *PassRegistry::getPassRegistry());
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 {
65 AU.addRequired<TargetTransformInfoWrapperPass>();
66 AU.addPreserved<DominatorTreeWrapperPass>();
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,
76 const TargetTransformInfo &TTI,
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)
85INITIALIZE_PASS_DEPENDENCY(TargetTransformInfoWrapperPass)
86INITIALIZE_PASS_DEPENDENCY(DominatorTreeWrapperPass)
87INITIALIZE_PASS_END(ScalarizeMaskedMemIntrinLegacyPass, DEBUG_TYPE,
88 "Scalarize unsupported masked memory intrinsics", false,
89 false)
90
91FunctionPass *llvm::createScalarizeMaskedMemIntrinLegacyPass() {
92 return new ScalarizeMaskedMemIntrinLegacyPass();
93}
94
95static bool isConstantIntVector(Value *Mask) {
96 Constant *C = dyn_cast<Constant>(Val: Mask);
97 if (!C)
98 return false;
99
100 unsigned NumElts = cast<FixedVectorType>(Val: Mask->getType())->getNumElements();
101 for (unsigned i = 0; i != NumElts; ++i) {
102 Constant *CElt = C->getAggregateElement(Elt: i);
103 if (!CElt || !isa<ConstantInt>(Val: 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(KindID: 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>(MD: CacheHint->getOperand(I: 0));
124 if (OperandNo->getZExtValue() != SourcePtrOperand)
125 return;
126
127 Metadata *DestOperandNo = ConstantAsMetadata::get(
128 C: ConstantInt::get(Ty: Type::getInt32Ty(C&: Dest.getContext()), V: DestPtrOperand));
129 Dest.setMetadata(KindID: LLVMContext::MD_mem_cache_hint,
130 Node: MDNode::get(Context&: Dest.getContext(),
131 MDs: {DestOperandNo, CacheHint->getOperand(I: 1)}));
132}
133
134static void copyMetadataForMemoryAccess(
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(SrcInst: Source,
142 WL: {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(Offset: *ByteOffset, AccessTy: 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
161static void copyMetadataForScalarizedLoad(LoadInst &Dest,
162 const Instruction &Source,
163 const DataLayout &DL,
164 unsigned SourcePtrOperand,
165 std::optional<size_t> ByteOffset) {
166 copyMetadataForMemoryAccess(Dest, Source, DL, SourcePtrOperand,
167 DestPtrOperand: Dest.getPointerOperandIndex(), AccessType: Dest.getType(),
168 IsWholeAccess: 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(SrcInst: Source, WL: {LLVMContext::MD_fpmath, LLVMContext::MD_range,
173 LLVMContext::MD_invariant_load});
174}
175
176static void copyMetadataForScalarizedStore(StoreInst &Dest,
177 const Instruction &Source,
178 const DataLayout &DL,
179 unsigned SourcePtrOperand,
180 std::optional<size_t> ByteOffset) {
181 copyMetadataForMemoryAccess(
182 Dest, Source, DL, SourcePtrOperand, DestPtrOperand: Dest.getPointerOperandIndex(),
183 AccessType: Dest.getValueOperand()->getType(),
184 IsWholeAccess: Dest.getValueOperand()->getType() == Source.getOperand(i: 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(i: 0);
224 Value *Mask = CI->getArgOperand(i: 1);
225 Value *Src0 = CI->getArgOperand(i: 2);
226
227 const Align AlignVal = CI->getParamAlign(ArgNo: 0).valueOrOne();
228 VectorType *VecType = cast<FixedVectorType>(Val: 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>(Val: Mask) && cast<Constant>(Val: Mask)->isAllOnesValue()) {
238 LoadInst *NewI = Builder.CreateAlignedLoad(Ty: VecType, Ptr, Align: AlignVal);
239 copyMetadataForScalarizedLoad(Dest&: *NewI, Source: *CI, DL, /*SourcePtrOperand=*/0,
240 ByteOffset: std::nullopt);
241 NewI->takeName(V: CI);
242 CI->replaceAllUsesWith(V: NewI);
243 CI->eraseFromParent();
244 return;
245 }
246
247 // Adjust alignment for the scalar instruction.
248 const Align AdjustedAlignVal =
249 commonAlignment(A: AlignVal, Offset: EltTy->getPrimitiveSizeInBits() / 8);
250 unsigned VectorWidth = cast<FixedVectorType>(Val: 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>(Val: Mask)->getAggregateElement(Elt: Idx)->isNullValue())
258 continue;
259 Value *Gep = Builder.CreateConstInBoundsGEP1_32(Ty: EltTy, Ptr, Idx0: Idx);
260 LoadInst *Load = Builder.CreateAlignedLoad(Ty: EltTy, Ptr: Gep, Align: AdjustedAlignVal);
261 copyMetadataForScalarizedLoad(
262 Dest&: *Load, Source: *CI, DL, /*SourcePtrOperand=*/0,
263 ByteOffset: Idx * DL.getTypeAllocSize(Ty: EltTy).getFixedValue());
264 VResult = Builder.CreateInsertElement(Vec: VResult, NewElt: Load, Idx);
265 }
266 CI->replaceAllUsesWith(V: 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(V: Mask, /*Index=*/0)) {
275 Value *Predicate = Builder.CreateExtractElement(Vec: Mask, Idx: uint64_t(0ull),
276 Name: 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(Cond: Predicate, SplitBefore: InsertPt, /*Unreachable=*/false,
281 BranchWeights: getExplicitlyUnknownBranchWeightsIfProfiled(
282 F&: *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(Ty: VecType, Ptr, Align: AlignVal,
289 Name: CI->getName() + ".cond.load");
290 copyMetadataForScalarizedLoad(Dest&: *Load, Source: *CI, DL, /*SourcePtrOperand=*/0,
291 ByteOffset: std::nullopt);
292
293 BasicBlock *PostLoad = ThenTerm->getSuccessor(Idx: 0);
294 Builder.SetInsertPoint(PostLoad->begin());
295 PHINode *Phi = Builder.CreatePHI(Ty: VecType, /*NumReservedValues=*/2);
296 Phi->addIncoming(V: Load, BB: CondBlock);
297 Phi->addIncoming(V: Src0, BB: IfBlock);
298 Phi->takeName(V: CI);
299
300 CI->replaceAllUsesWith(V: 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(N: VectorWidth);
311 SclrMask = Builder.CreateBitCast(V: Mask, DestTy: SclrMaskTy, Name: "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
324 Value *Predicate;
325 if (SclrMask != nullptr) {
326 Value *Mask = Builder.getInt(AI: APInt::getOneBitSet(
327 numBits: VectorWidth, BitNo: adjustForEndian(DL, VectorWidth, Idx)));
328 Predicate = Builder.CreateICmpNE(LHS: Builder.CreateAnd(LHS: SclrMask, RHS: Mask),
329 RHS: Builder.getIntN(N: VectorWidth, C: 0));
330 } else {
331 Predicate = Builder.CreateExtractElement(Vec: 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(Cond: Predicate, SplitBefore: InsertPt, /*Unreachable=*/false,
344 BranchWeights: getExplicitlyUnknownBranchWeightsIfProfiled(
345 F&: *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(Ty: EltTy, Ptr, Idx0: Idx);
353 LoadInst *Load = Builder.CreateAlignedLoad(Ty: EltTy, Ptr: Gep, Align: AdjustedAlignVal);
354 copyMetadataForScalarizedLoad(
355 Dest&: *Load, Source: *CI, DL, /*SourcePtrOperand=*/0,
356 ByteOffset: Idx * DL.getTypeAllocSize(Ty: EltTy).getFixedValue());
357 Value *NewVResult = Builder.CreateInsertElement(Vec: VResult, NewElt: Load, Idx);
358
359 // Create "else" block, fill it in the next iteration
360 BasicBlock *NewIfBlock = ThenTerm->getSuccessor(Idx: 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->begin());
367 PHINode *Phi = Builder.CreatePHI(Ty: VecType, NumReservedValues: 2, Name: "res.phi.else");
368 Phi->addIncoming(V: NewVResult, BB: CondBlock);
369 Phi->addIncoming(V: VResult, BB: PrevIfBlock);
370 VResult = Phi;
371 }
372
373 CI->replaceAllUsesWith(V: 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(i: 0);
409 Value *Ptr = CI->getArgOperand(i: 1);
410 Value *Mask = CI->getArgOperand(i: 2);
411
412 const Align AlignVal = CI->getParamAlign(ArgNo: 1).valueOrOne();
413 auto *VecType = cast<VectorType>(Val: 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>(Val: Mask) && cast<Constant>(Val: Mask)->isAllOnesValue()) {
422 StoreInst *Store = Builder.CreateAlignedStore(Val: Src, Ptr, Align: AlignVal);
423 Store->takeName(V: CI);
424 copyMetadataForScalarizedStore(Dest&: *Store, Source: *CI, DL, /*SourcePtrOperand=*/1,
425 ByteOffset: std::nullopt);
426 // This is a one-to-one replacement, so the assignment link remains valid.
427 Store->copyMetadata(SrcInst: *CI, WL: LLVMContext::MD_DIAssignID);
428 CI->eraseFromParent();
429 return;
430 }
431
432 // Adjust alignment for the scalar instruction.
433 const Align AdjustedAlignVal =
434 commonAlignment(A: AlignVal, Offset: EltTy->getPrimitiveSizeInBits() / 8);
435 unsigned VectorWidth = cast<FixedVectorType>(Val: VecType)->getNumElements();
436
437 if (isConstantIntVector(Mask)) {
438 for (unsigned Idx = 0; Idx < VectorWidth; ++Idx) {
439 if (cast<Constant>(Val: Mask)->getAggregateElement(Elt: Idx)->isNullValue())
440 continue;
441 Value *OneElt = Builder.CreateExtractElement(Vec: Src, Idx);
442 Value *Gep = Builder.CreateConstInBoundsGEP1_32(Ty: EltTy, Ptr, Idx0: Idx);
443 StoreInst *Store =
444 Builder.CreateAlignedStore(Val: OneElt, Ptr: Gep, Align: AdjustedAlignVal);
445 copyMetadataForScalarizedStore(
446 Dest&: *Store, Source: *CI, DL, /*SourcePtrOperand=*/1,
447 ByteOffset: Idx * DL.getTypeAllocSize(Ty: 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(V: Mask, /*Index=*/0)) {
457 Value *Predicate = Builder.CreateExtractElement(Vec: Mask, Idx: uint64_t(0ull),
458 Name: 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(Cond: Predicate, SplitBefore: InsertPt, /*Unreachable=*/false,
463 BranchWeights: getExplicitlyUnknownBranchWeightsIfProfiled(
464 F&: *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(Val: Src, Ptr, Align: AlignVal);
471 Store->takeName(V: CI);
472 copyMetadataForScalarizedStore(Dest&: *Store, Source: *CI, DL, /*SourcePtrOperand=*/1,
473 ByteOffset: std::nullopt);
474 // This is a one-to-one replacement, so the assignment link remains valid.
475 Store->copyMetadata(SrcInst: *CI, WL: 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(N: VectorWidth);
488 SclrMask = Builder.CreateBitCast(V: Mask, DestTy: SclrMaskTy, Name: "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
501 Value *Predicate;
502 if (SclrMask != nullptr) {
503 Value *Mask = Builder.getInt(AI: APInt::getOneBitSet(
504 numBits: VectorWidth, BitNo: adjustForEndian(DL, VectorWidth, Idx)));
505 Predicate = Builder.CreateICmpNE(LHS: Builder.CreateAnd(LHS: SclrMask, RHS: Mask),
506 RHS: Builder.getIntN(N: VectorWidth, C: 0));
507 } else {
508 Predicate = Builder.CreateExtractElement(Vec: 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(Cond: Predicate, SplitBefore: InsertPt, /*Unreachable=*/false,
521 BranchWeights: getExplicitlyUnknownBranchWeightsIfProfiled(
522 F&: *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(Vec: Src, Idx);
530 Value *Gep = Builder.CreateConstInBoundsGEP1_32(Ty: EltTy, Ptr, Idx0: Idx);
531 StoreInst *Store =
532 Builder.CreateAlignedStore(Val: OneElt, Ptr: Gep, Align: AdjustedAlignVal);
533 copyMetadataForScalarizedStore(
534 Dest&: *Store, Source: *CI, DL, /*SourcePtrOperand=*/1,
535 ByteOffset: Idx * DL.getTypeAllocSize(Ty: EltTy).getFixedValue());
536
537 // Create "else" block, fill it in the next iteration
538 BasicBlock *NewIfBlock = ThenTerm->getSuccessor(Idx: 0);
539 NewIfBlock->setName("else");
540
541 Builder.SetInsertPoint(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
577static void scalarizeMaskedGather(const DataLayout &DL,
578 bool HasBranchDivergence, CallInst *CI,
579 DomTreeUpdater *DTU, bool &ModifiedDT) {
580 Value *Ptrs = CI->getArgOperand(i: 0);
581 Value *Mask = CI->getArgOperand(i: 1);
582 Value *Src0 = CI->getArgOperand(i: 2);
583
584 auto *VecType = cast<FixedVectorType>(Val: 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(ArgNo: 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>(Val: Mask)->getAggregateElement(Elt: Idx)->isNullValue())
600 continue;
601 Value *Ptr = Builder.CreateExtractElement(Vec: Ptrs, Idx, Name: "Ptr" + Twine(Idx));
602 LoadInst *Load =
603 Builder.CreateAlignedLoad(Ty: EltTy, Ptr, Align: AlignVal, Name: "Load" + Twine(Idx));
604 copyMetadataForScalarizedLoad(Dest&: *Load, Source: *CI, DL, /*SourcePtrOperand=*/0,
605 /*ByteOffset=*/0);
606 VResult =
607 Builder.CreateInsertElement(Vec: VResult, NewElt: Load, Idx, Name: "Res" + Twine(Idx));
608 }
609 CI->replaceAllUsesWith(V: 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(N: VectorWidth);
620 SclrMask = Builder.CreateBitCast(V: Mask, DestTy: SclrMaskTy, Name: "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
634 Value *Predicate;
635 if (SclrMask != nullptr) {
636 Value *Mask = Builder.getInt(AI: APInt::getOneBitSet(
637 numBits: VectorWidth, BitNo: adjustForEndian(DL, VectorWidth, Idx)));
638 Predicate = Builder.CreateICmpNE(LHS: Builder.CreateAnd(LHS: SclrMask, RHS: Mask),
639 RHS: Builder.getIntN(N: VectorWidth, C: 0));
640 } else {
641 Predicate = Builder.CreateExtractElement(Vec: Mask, Idx, Name: "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(Cond: Predicate, SplitBefore: InsertPt, /*Unreachable=*/false,
654 BranchWeights: getExplicitlyUnknownBranchWeightsIfProfiled(
655 F&: *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(Vec: Ptrs, Idx, Name: "Ptr" + Twine(Idx));
663 LoadInst *Load =
664 Builder.CreateAlignedLoad(Ty: EltTy, Ptr, Align: AlignVal, Name: "Load" + Twine(Idx));
665 copyMetadataForScalarizedLoad(Dest&: *Load, Source: *CI, DL, /*SourcePtrOperand=*/0,
666 /*ByteOffset=*/0);
667 Value *NewVResult =
668 Builder.CreateInsertElement(Vec: VResult, NewElt: Load, Idx, Name: "Res" + Twine(Idx));
669
670 // Create "else" block, fill it in the next iteration
671 BasicBlock *NewIfBlock = ThenTerm->getSuccessor(Idx: 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->begin());
678 PHINode *Phi = Builder.CreatePHI(Ty: VecType, NumReservedValues: 2, Name: "res.phi.else");
679 Phi->addIncoming(V: NewVResult, BB: CondBlock);
680 Phi->addIncoming(V: VResult, BB: PrevIfBlock);
681 VResult = Phi;
682 }
683
684 CI->replaceAllUsesWith(V: 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// . . .
716static void scalarizeMaskedScatter(const DataLayout &DL,
717 bool HasBranchDivergence, CallInst *CI,
718 DomTreeUpdater *DTU, bool &ModifiedDT) {
719 Value *Src = CI->getArgOperand(i: 0);
720 Value *Ptrs = CI->getArgOperand(i: 1);
721 Value *Mask = CI->getArgOperand(i: 2);
722
723 auto *SrcFVTy = cast<FixedVectorType>(Val: 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(ArgNo: 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>(Val: Mask)->getAggregateElement(Elt: Idx)->isNullValue())
740 continue;
741 Value *OneElt =
742 Builder.CreateExtractElement(Vec: Src, Idx, Name: "Elt" + Twine(Idx));
743 Value *Ptr = Builder.CreateExtractElement(Vec: Ptrs, Idx, Name: "Ptr" + Twine(Idx));
744 StoreInst *Store = Builder.CreateAlignedStore(Val: OneElt, Ptr, Align: AlignVal);
745 copyMetadataForScalarizedStore(Dest&: *Store, Source: *CI, DL,
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(N: VectorWidth);
758 SclrMask = Builder.CreateBitCast(V: Mask, DestTy: SclrMaskTy, Name: "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
771 Value *Predicate;
772 if (SclrMask != nullptr) {
773 Value *Mask = Builder.getInt(AI: APInt::getOneBitSet(
774 numBits: VectorWidth, BitNo: adjustForEndian(DL, VectorWidth, Idx)));
775 Predicate = Builder.CreateICmpNE(LHS: Builder.CreateAnd(LHS: SclrMask, RHS: Mask),
776 RHS: Builder.getIntN(N: VectorWidth, C: 0));
777 } else {
778 Predicate = Builder.CreateExtractElement(Vec: Mask, Idx, Name: "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(Cond: Predicate, SplitBefore: InsertPt, /*Unreachable=*/false,
791 BranchWeights: getExplicitlyUnknownBranchWeightsIfProfiled(
792 F&: *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(Vec: Src, Idx, Name: "Elt" + Twine(Idx));
800 Value *Ptr = Builder.CreateExtractElement(Vec: Ptrs, Idx, Name: "Ptr" + Twine(Idx));
801 StoreInst *Store = Builder.CreateAlignedStore(Val: OneElt, Ptr, Align: AlignVal);
802 copyMetadataForScalarizedStore(Dest&: *Store, Source: *CI, DL,
803 /*SourcePtrOperand=*/1,
804 /*ByteOffset=*/0);
805
806 // Create "else" block, fill it in the next iteration
807 BasicBlock *NewIfBlock = ThenTerm->getSuccessor(Idx: 0);
808 NewIfBlock->setName("else");
809
810 Builder.SetInsertPoint(NewIfBlock->begin());
811 }
812 CI->eraseFromParent();
813
814 ModifiedDT = true;
815}
816
817static void scalarizeMaskedExpandLoad(const DataLayout &DL,
818 bool HasBranchDivergence, CallInst *CI,
819 DomTreeUpdater *DTU, bool &ModifiedDT) {
820 Value *Ptr = CI->getArgOperand(i: 0);
821 Value *Mask = CI->getArgOperand(i: 1);
822 Value *PassThru = CI->getArgOperand(i: 2);
823 Align Alignment = CI->getParamAlign(ArgNo: 0).valueOrOne();
824
825 auto *VecType = cast<FixedVectorType>(Val: 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(A: Alignment, Offset: 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(T: VecType);
848 SmallVector<int, 16> ShuffleMask(VectorWidth, PoisonMaskElem);
849 for (unsigned Idx = 0; Idx < VectorWidth; ++Idx) {
850 Value *InsertElt;
851 if (cast<Constant>(Val: Mask)->getAggregateElement(Elt: Idx)->isNullValue()) {
852 InsertElt = PoisonValue::get(T: EltTy);
853 ShuffleMask[Idx] = Idx + VectorWidth;
854 } else {
855 Value *NewPtr =
856 Builder.CreateConstInBoundsGEP1_32(Ty: EltTy, Ptr, Idx0: MemIndex);
857 LoadInst *Load = Builder.CreateAlignedLoad(
858 Ty: EltTy, Ptr: NewPtr, Align: AdjustedAlignment, Name: "Load" + Twine(Idx));
859 copyMetadataForScalarizedLoad(
860 Dest&: *Load, Source: *CI, DL, /*SourcePtrOperand=*/0,
861 ByteOffset: MemIndex * DL.getTypeAllocSize(Ty: EltTy).getFixedValue());
862 InsertElt = Load;
863 ShuffleMask[Idx] = Idx;
864 ++MemIndex;
865 }
866 VResult = Builder.CreateInsertElement(Vec: VResult, NewElt: InsertElt, Idx,
867 Name: "Res" + Twine(Idx));
868 }
869 VResult = Builder.CreateShuffleVector(V1: VResult, V2: PassThru, Mask: ShuffleMask);
870 CI->replaceAllUsesWith(V: 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(N: VectorWidth);
881 SclrMask = Builder.CreateBitCast(V: Mask, DestTy: SclrMaskTy, Name: "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
895 Value *Predicate;
896 if (SclrMask != nullptr) {
897 Value *Mask = Builder.getInt(AI: APInt::getOneBitSet(
898 numBits: VectorWidth, BitNo: adjustForEndian(DL, VectorWidth, Idx)));
899 Predicate = Builder.CreateICmpNE(LHS: Builder.CreateAnd(LHS: SclrMask, RHS: Mask),
900 RHS: Builder.getIntN(N: VectorWidth, C: 0));
901 } else {
902 Predicate = Builder.CreateExtractElement(Vec: Mask, Idx, Name: "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(Cond: Predicate, SplitBefore: InsertPt, /*Unreachable=*/false,
915 BranchWeights: getExplicitlyUnknownBranchWeightsIfProfiled(
916 F&: *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(Ty: EltTy, Ptr, Align: AdjustedAlignment);
924 copyMetadataForScalarizedLoad(Dest&: *Load, Source: *CI, DL, /*SourcePtrOperand=*/0,
925 ByteOffset: std::nullopt);
926 Value *NewVResult = Builder.CreateInsertElement(Vec: VResult, NewElt: 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(Ty: EltTy, Ptr, Idx0: 1);
932
933 // Create "else" block, fill it in the next iteration
934 BasicBlock *NewIfBlock = ThenTerm->getSuccessor(Idx: 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->begin());
941 PHINode *ResultPhi = Builder.CreatePHI(Ty: VecType, NumReservedValues: 2, Name: "res.phi.else");
942 ResultPhi->addIncoming(V: NewVResult, BB: CondBlock);
943 ResultPhi->addIncoming(V: VResult, BB: 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(Ty: Ptr->getType(), NumReservedValues: 2, Name: "ptr.phi.else");
949 PtrPhi->addIncoming(V: NewPtr, BB: CondBlock);
950 PtrPhi->addIncoming(V: Ptr, BB: PrevIfBlock);
951 Ptr = PtrPhi;
952 }
953 }
954
955 CI->replaceAllUsesWith(V: VResult);
956 CI->eraseFromParent();
957
958 ModifiedDT = true;
959}
960
961static void scalarizeMaskedCompressStore(const DataLayout &DL,
962 bool HasBranchDivergence, CallInst *CI,
963 DomTreeUpdater *DTU,
964 bool &ModifiedDT) {
965 Value *Src = CI->getArgOperand(i: 0);
966 Value *Ptr = CI->getArgOperand(i: 1);
967 Value *Mask = CI->getArgOperand(i: 2);
968 Align Alignment = CI->getParamAlign(ArgNo: 1).valueOrOne();
969
970 auto *VecType = cast<FixedVectorType>(Val: 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(A: Alignment, Offset: 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>(Val: Mask)->getAggregateElement(Elt: Idx)->isNullValue())
989 continue;
990 Value *OneElt =
991 Builder.CreateExtractElement(Vec: Src, Idx, Name: "Elt" + Twine(Idx));
992 Value *NewPtr = Builder.CreateConstInBoundsGEP1_32(Ty: EltTy, Ptr, Idx0: MemIndex);
993 StoreInst *Store =
994 Builder.CreateAlignedStore(Val: OneElt, Ptr: NewPtr, Align: AdjustedAlignment);
995 copyMetadataForScalarizedStore(
996 Dest&: *Store, Source: *CI, DL, /*SourcePtrOperand=*/1,
997 ByteOffset: MemIndex * DL.getTypeAllocSize(Ty: 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(N: VectorWidth);
1010 SclrMask = Builder.CreateBitCast(V: Mask, DestTy: SclrMaskTy, Name: "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
1022 Value *Predicate;
1023 if (SclrMask != nullptr) {
1024 Value *Mask = Builder.getInt(AI: APInt::getOneBitSet(
1025 numBits: VectorWidth, BitNo: adjustForEndian(DL, VectorWidth, Idx)));
1026 Predicate = Builder.CreateICmpNE(LHS: Builder.CreateAnd(LHS: SclrMask, RHS: Mask),
1027 RHS: Builder.getIntN(N: VectorWidth, C: 0));
1028 } else {
1029 Predicate = Builder.CreateExtractElement(Vec: Mask, Idx, Name: "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(Cond: Predicate, SplitBefore: InsertPt, /*Unreachable=*/false,
1042 BranchWeights: getExplicitlyUnknownBranchWeightsIfProfiled(
1043 F&: *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(Vec: Src, Idx);
1051 StoreInst *Store =
1052 Builder.CreateAlignedStore(Val: OneElt, Ptr, Align: AdjustedAlignment);
1053 copyMetadataForScalarizedStore(Dest&: *Store, Source: *CI, DL, /*SourcePtrOperand=*/1,
1054 ByteOffset: 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(Ty: EltTy, Ptr, Idx0: 1);
1060
1061 // Create "else" block, fill it in the next iteration
1062 BasicBlock *NewIfBlock = ThenTerm->getSuccessor(Idx: 0);
1063 NewIfBlock->setName("else");
1064 BasicBlock *PrevIfBlock = IfBlock;
1065 IfBlock = NewIfBlock;
1066
1067 Builder.SetInsertPoint(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(Ty: Ptr->getType(), NumReservedValues: 2, Name: "ptr.phi.else");
1072 PtrPhi->addIncoming(V: NewPtr, BB: CondBlock);
1073 PtrPhi->addIncoming(V: Ptr, BB: PrevIfBlock);
1074 Ptr = PtrPhi;
1075 }
1076 }
1077 CI->eraseFromParent();
1078
1079 ModifiedDT = true;
1080}
1081
1082static void scalarizeMaskedVectorHistogram(const DataLayout &DL, CallInst *CI,
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(i: 0);
1089 Value *Inc = CI->getArgOperand(i: 1);
1090 Value *Mask = CI->getArgOperand(i: 2);
1091
1092 auto *AddrType = cast<FixedVectorType>(Val: 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(LHS: Load, RHS: Inc);
1106 break;
1107 case Intrinsic::experimental_vector_histogram_uadd_sat:
1108 UpdateOp =
1109 Builder.CreateIntrinsic(ID: Intrinsic::uadd_sat, OverloadTypes: {EltTy}, Args: {Load, Inc});
1110 break;
1111 case Intrinsic::experimental_vector_histogram_umin:
1112 UpdateOp = Builder.CreateIntrinsic(ID: Intrinsic::umin, OverloadTypes: {EltTy}, Args: {Load, Inc});
1113 break;
1114 case Intrinsic::experimental_vector_histogram_umax:
1115 UpdateOp = Builder.CreateIntrinsic(ID: Intrinsic::umax, OverloadTypes: {EltTy}, Args: {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>(Val: Mask)->getAggregateElement(Elt: Idx)->isNullValue())
1128 continue;
1129 Value *Ptr = Builder.CreateExtractElement(Vec: Ptrs, Idx, Name: "Ptr" + Twine(Idx));
1130 LoadInst *Load = Builder.CreateLoad(Ty: EltTy, Ptr, Name: "Load" + Twine(Idx));
1131 copyMetadataForScalarizedLoad(Dest&: *Load, Source: *CI, DL, /*SourcePtrOperand=*/0,
1132 /*ByteOffset=*/0);
1133 Value *Update =
1134 CreateHistogramUpdateValue(cast<IntrinsicInst>(Val: CI), Load, Inc);
1135 StoreInst *Store = Builder.CreateStore(Val: Update, Ptr);
1136 copyMetadataForScalarizedStore(Dest&: *Store, Source: *CI, DL,
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(Vec: Mask, Idx, Name: "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(Cond: Predicate, SplitBefore: InsertPt, /*Unreachable=*/false,
1152 BranchWeights: getExplicitlyUnknownBranchWeightsIfProfiled(
1153 F&: *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(Vec: Ptrs, Idx, Name: "Ptr" + Twine(Idx));
1161 LoadInst *Load = Builder.CreateLoad(Ty: EltTy, Ptr, Name: "Load" + Twine(Idx));
1162 copyMetadataForScalarizedLoad(Dest&: *Load, Source: *CI, DL, /*SourcePtrOperand=*/0,
1163 /*ByteOffset=*/0);
1164 Value *UpdateOp =
1165 CreateHistogramUpdateValue(cast<IntrinsicInst>(Val: CI), Load, Inc);
1166 StoreInst *Store = Builder.CreateStore(Val: UpdateOp, Ptr);
1167 copyMetadataForScalarizedStore(Dest&: *Store, Source: *CI, DL,
1168 /*SourcePtrOperand=*/0,
1169 /*ByteOffset=*/0);
1170
1171 // Create "else" block, fill it in the next iteration
1172 BasicBlock *NewIfBlock = ThenTerm->getSuccessor(Idx: 0);
1173 NewIfBlock->setName("else");
1174 Builder.SetInsertPoint(NewIfBlock->begin());
1175 }
1176
1177 CI->eraseFromParent();
1178 ModifiedDT = true;
1179}
1180
1181static bool runImpl(Function &F, const TargetTransformInfo &TTI,
1182 DominatorTree *DT) {
1183 std::optional<DomTreeUpdater> DTU;
1184 if (DT)
1185 DTU.emplace(args&: DT, args: DomTreeUpdater::UpdateStrategy::Lazy);
1186
1187 bool EverMadeChange = false;
1188 bool MadeChange = true;
1189 auto &DL = F.getDataLayout();
1190 bool HasBranchDivergence = TTI.hasBranchDivergence(F: &F);
1191 while (MadeChange) {
1192 MadeChange = false;
1193 for (BasicBlock &BB : llvm::make_early_inc_range(Range&: F)) {
1194 bool ModifiedDTOnIteration = false;
1195 MadeChange |= optimizeBlock(BB, ModifiedDT&: ModifiedDTOnIteration, TTI, DL,
1196 HasBranchDivergence, DTU: 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
1217ScalarizeMaskedMemIntrinPass::run(Function &F, FunctionAnalysisManager &AM) {
1218 auto &TTI = AM.getResult<TargetIRAnalysis>(IR&: F);
1219 auto *DT = AM.getCachedResult<DominatorTreeAnalysis>(IR&: F);
1220 if (!runImpl(F, TTI, DT))
1221 return PreservedAnalyses::all();
1222 PreservedAnalyses PA;
1223 PA.preserve<TargetIRAnalysis>();
1224 PA.preserve<DominatorTreeAnalysis>();
1225 return PA;
1226}
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>(Val: &*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) {
1249 IntrinsicInst *II = dyn_cast<IntrinsicInst>(Val: CI);
1250 if (II) {
1251 // The scalarization code below does not work for scalable vectors.
1252 if (isa<ScalableVectorType>(Val: II->getType()) ||
1253 any_of(Range: II->args(),
1254 P: [](Value *V) { return isa<ScalableVectorType>(Val: 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(AddrType: CI->getArgOperand(i: 0)->getType(),
1264 DataType: CI->getArgOperand(i: 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 DataType: CI->getType(), Alignment: CI->getParamAlign(ArgNo: 0).valueOrOne(),
1272 AddressSpace: cast<PointerType>(Val: CI->getArgOperand(i: 0)->getType())
1273 ->getAddressSpace(),
1274 MaskKind: isConstantIntVector(Mask: CI->getArgOperand(i: 1))
1275 ? TTI::MaskKind::ConstantMask
1276 : TTI::MaskKind::VariableOrConstantMask))
1277 return false;
1278 scalarizeMaskedLoad(DL, HasBranchDivergence, CI, DTU, ModifiedDT);
1279 return true;
1280 case Intrinsic::masked_store:
1281 if (TTI.isLegalMaskedStore(
1282 DataType: CI->getArgOperand(i: 0)->getType(),
1283 Alignment: CI->getParamAlign(ArgNo: 1).valueOrOne(),
1284 AddressSpace: cast<PointerType>(Val: CI->getArgOperand(i: 1)->getType())
1285 ->getAddressSpace(),
1286 MaskKind: isConstantIntVector(Mask: CI->getArgOperand(i: 2))
1287 ? TTI::MaskKind::ConstantMask
1288 : TTI::MaskKind::VariableOrConstantMask))
1289 return false;
1290 scalarizeMaskedStore(DL, HasBranchDivergence, CI, DTU, ModifiedDT);
1291 return true;
1292 case Intrinsic::masked_gather: {
1293 Align Alignment = CI->getParamAlign(ArgNo: 0).valueOrOne();
1294 Type *LoadTy = CI->getType();
1295 if (TTI.isLegalMaskedGather(DataType: LoadTy, Alignment) &&
1296 !TTI.forceScalarizeMaskedGather(Type: cast<VectorType>(Val: 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(ArgNo: 1).valueOrOne();
1303 Type *StoreTy = CI->getArgOperand(i: 0)->getType();
1304 if (TTI.isLegalMaskedScatter(DataType: StoreTy, Alignment) &&
1305 !TTI.forceScalarizeMaskedScatter(Type: cast<VectorType>(Val: 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 DataType: CI->getType(),
1314 Alignment: CI->getAttributes().getParamAttrs(ArgNo: 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 DataType: CI->getArgOperand(i: 0)->getType(),
1321 Alignment: CI->getAttributes().getParamAttrs(ArgNo: 1).getAlignment().valueOrOne()))
1322 return false;
1323 scalarizeMaskedCompressStore(DL, HasBranchDivergence, CI, DTU,
1324 ModifiedDT);
1325 return true;
1326 }
1327 }
1328
1329 return false;
1330}
1331