1//===-- AMDGPULowerKernelArguments.cpp ------------------------------------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9/// \file This pass replaces accesses to kernel arguments with loads from
10/// offsets from the kernarg base pointer.
11//
12//===----------------------------------------------------------------------===//
13
14#include "AMDGPU.h"
15#include "GCNSubtarget.h"
16#include "llvm/Analysis/AliasAnalysis.h"
17#include "llvm/Analysis/CaptureTracking.h"
18#include "llvm/Analysis/ScopedNoAliasAA.h"
19#include "llvm/Analysis/ValueTracking.h"
20#include "llvm/CodeGen/TargetPassConfig.h"
21#include "llvm/IR/Argument.h"
22#include "llvm/IR/Attributes.h"
23#include "llvm/IR/Dominators.h"
24#include "llvm/IR/IRBuilder.h"
25#include "llvm/IR/InstIterator.h"
26#include "llvm/IR/Instruction.h"
27#include "llvm/IR/Instructions.h"
28#include "llvm/IR/IntrinsicsAMDGPU.h"
29#include "llvm/IR/LLVMContext.h"
30#include "llvm/IR/MDBuilder.h"
31#include "llvm/Target/TargetMachine.h"
32#include <optional>
33
34#define DEBUG_TYPE "amdgpu-lower-kernel-arguments"
35
36using namespace llvm;
37
38namespace {
39
40class AMDGPULowerKernelArguments : public FunctionPass {
41public:
42 static char ID;
43
44 AMDGPULowerKernelArguments() : FunctionPass(ID) {}
45
46 bool runOnFunction(Function &F) override;
47
48 void getAnalysisUsage(AnalysisUsage &AU) const override {
49 AU.addRequired<TargetPassConfig>();
50 AU.addRequired<DominatorTreeWrapperPass>();
51 AU.setPreservesAll();
52 }
53};
54
55} // end anonymous namespace
56
57// skip allocas
58static BasicBlock::iterator getInsertPt(BasicBlock &BB) {
59 BasicBlock::iterator InsPt = BB.getFirstInsertionPt();
60 for (BasicBlock::iterator E = BB.end(); InsPt != E; ++InsPt) {
61 AllocaInst *AI = dyn_cast<AllocaInst>(Val: &*InsPt);
62
63 // If this is a dynamic alloca, the value may depend on the loaded kernargs,
64 // so loads will need to be inserted before it.
65 if (!AI || !AI->isStaticAlloca())
66 break;
67 }
68
69 return InsPt;
70}
71
72static void addAliasScopeMetadata(Function &F, const DataLayout &DL,
73 DominatorTree &DT) {
74 // Collect noalias arguments.
75 SmallVector<const Argument *, 4u> NoAliasArgs;
76
77 for (Argument &Arg : F.args())
78 if (Arg.hasNoAliasAttr() && !Arg.use_empty())
79 NoAliasArgs.push_back(Elt: &Arg);
80
81 if (NoAliasArgs.empty())
82 return;
83
84 // Add alias scopes for each noalias argument.
85 MDBuilder MDB(F.getContext());
86 DenseMap<const Argument *, MDNode *> NewScopes;
87 MDNode *NewDomain = MDB.createAnonymousAliasScopeDomain(Name: F.getName());
88
89 for (unsigned I = 0u; I < NoAliasArgs.size(); ++I) {
90 const Argument *Arg = NoAliasArgs[I];
91 MDNode *NewScope = MDB.createAnonymousAliasScope(Domain: NewDomain, Name: Arg->getName());
92 NewScopes.insert(KV: {Arg, NewScope});
93 }
94
95 // Iterate over all instructions.
96 for (inst_iterator Inst = inst_begin(F), InstEnd = inst_end(F);
97 Inst != InstEnd; ++Inst) {
98 // If instruction accesses memory, collect its pointer arguments.
99 Instruction *I = &(*Inst);
100 SmallVector<const Value *, 2u> PtrArgs;
101 // May reach a noalias argument via a copy captured before the call.
102 bool IsUnrestrictedCall = false;
103
104 if (std::optional<MemoryLocation> MO = MemoryLocation::getOrNone(Inst: I))
105 PtrArgs.push_back(Elt: MO->Ptr);
106 else if (const CallBase *Call = dyn_cast<CallBase>(Val: I)) {
107 MemoryEffects ME = Call->getMemoryEffects();
108 if (ME.doesNotAccessMemory())
109 continue;
110
111 // Inaccessible memory cannot alias any IR-visible pointer.
112 if (ME.onlyAccessesInaccessibleMem())
113 continue;
114 IsUnrestrictedCall = !ME.onlyAccessesArgPointees();
115
116 for (Value *Arg : Call->args()) {
117 if (!Arg->getType()->isPointerTy())
118 continue;
119
120 PtrArgs.push_back(Elt: Arg);
121 }
122 } else {
123 // Not a memory access and not a call — nothing to annotate.
124 continue;
125 }
126
127 // Collect underlying objects of pointer arguments.
128 SmallVector<Metadata *, 4u> Scopes;
129 SmallPtrSet<const Value *, 4u> ObjSet;
130 SmallVector<Metadata *, 4u> NoAliases;
131
132 for (const Value *Val : PtrArgs) {
133 SmallVector<const Value *, 4u> Objects;
134 getUnderlyingObjects(V: Val, Objects);
135 ObjSet.insert_range(R&: Objects);
136 }
137
138 bool RequiresNoCaptureBefore = false;
139 bool UsesUnknownObject = false;
140 bool UsesAliasingPtr = false;
141
142 for (const Value *Val : ObjSet) {
143 if (isa<ConstantData>(Val))
144 continue;
145
146 if (const Argument *Arg = dyn_cast<Argument>(Val)) {
147 if (!Arg->hasAttribute(Kind: Attribute::NoAlias))
148 UsesAliasingPtr = true;
149 } else
150 UsesAliasingPtr = true;
151
152 if (isEscapeSource(V: Val)) {
153 // Can only alias a noalias argument if captured beforehand.
154 RequiresNoCaptureBefore = true;
155 } else if (!isa<Argument>(Val) && !isIdentifiedObject(V: Val)) {
156 // Unknown provenance: assume nothing.
157 UsesUnknownObject = true;
158 }
159 }
160
161 if (UsesUnknownObject)
162 continue;
163
164 if (IsUnrestrictedCall)
165 RequiresNoCaptureBefore = true;
166
167 // Collect noalias scopes for instruction.
168 for (const Argument *Arg : NoAliasArgs) {
169 if (ObjSet.contains(Ptr: Arg))
170 continue;
171
172 if (!RequiresNoCaptureBefore ||
173 !capturesAnything(CC: PointerMayBeCapturedBefore(
174 V: Arg, ReturnCaptures: false, I, DT: &DT, IncludeI: false, Mask: CaptureComponents::Provenance)))
175 NoAliases.push_back(Elt: NewScopes[Arg]);
176 }
177
178 // Collect scopes for alias.scope metadata. Skip unrestricted calls: they
179 // may touch memory beyond their pointer arguments' pointees.
180 if (!UsesAliasingPtr && !IsUnrestrictedCall)
181 for (const Argument *Arg : NoAliasArgs) {
182 if (ObjSet.count(Ptr: Arg))
183 Scopes.push_back(Elt: NewScopes[Arg]);
184 }
185
186 // Add noalias metadata to instruction.
187 if (!NoAliases.empty()) {
188 MDNode *NewMD =
189 MDNode::concatenate(A: Inst->getMetadata(KindID: LLVMContext::MD_noalias),
190 B: MDNode::get(Context&: F.getContext(), MDs: NoAliases));
191 Inst->setMetadata(KindID: LLVMContext::MD_noalias, Node: NewMD);
192 }
193
194 // Add alias.scope metadata to instruction.
195 if (!Scopes.empty()) {
196 MDNode *NewMD =
197 MDNode::concatenate(A: Inst->getMetadata(KindID: LLVMContext::MD_alias_scope),
198 B: MDNode::get(Context&: F.getContext(), MDs: Scopes));
199 Inst->setMetadata(KindID: LLVMContext::MD_alias_scope, Node: NewMD);
200 }
201 }
202}
203
204static bool lowerKernelArguments(Function &F, const TargetMachine &TM,
205 DominatorTree &DT) {
206 CallingConv::ID CC = F.getCallingConv();
207 if (CC != CallingConv::AMDGPU_KERNEL || F.arg_empty())
208 return false;
209
210 const GCNSubtarget &ST = TM.getSubtarget<GCNSubtarget>(F);
211 LLVMContext &Ctx = F.getContext();
212 const DataLayout &DL = F.getDataLayout();
213 BasicBlock &EntryBlock = *F.begin();
214 IRBuilder<> Builder(&EntryBlock, getInsertPt(BB&: EntryBlock));
215
216 const Align KernArgBaseAlign(16); // FIXME: Increase if necessary
217 const uint64_t BaseOffset = ST.getExplicitKernelArgOffset();
218
219 Align MaxAlign;
220 // FIXME: Alignment is broken with explicit arg offset.;
221 const uint64_t TotalKernArgSize = ST.getKernArgSegmentSize(F, MaxAlign);
222 if (TotalKernArgSize == 0)
223 return false;
224
225 CallInst *KernArgSegment = Builder.CreateIntrinsicWithoutFolding(
226 ID: Intrinsic::amdgcn_kernarg_segment_ptr, Args: {}, FMFSource: nullptr,
227 Name: F.getName() + ".kernarg.segment");
228 KernArgSegment->addRetAttr(Kind: Attribute::NonNull);
229 KernArgSegment->addRetAttr(
230 Attr: Attribute::getWithDereferenceableBytes(Context&: Ctx, Bytes: TotalKernArgSize));
231
232 uint64_t ExplicitArgOffset = 0;
233
234 addAliasScopeMetadata(F, DL: F.getParent()->getDataLayout(), DT);
235
236 for (Argument &Arg : F.args()) {
237 const bool IsByRef = Arg.hasByRefAttr();
238 Type *ArgTy = IsByRef ? Arg.getParamByRefType() : Arg.getType();
239 MaybeAlign ParamAlign = IsByRef ? Arg.getParamAlign() : std::nullopt;
240 Align ABITypeAlign = DL.getValueOrABITypeAlignment(Alignment: ParamAlign, Ty: ArgTy);
241
242 uint64_t Size = DL.getTypeSizeInBits(Ty: ArgTy);
243 uint64_t AllocSize = DL.getTypeAllocSize(Ty: ArgTy);
244
245 uint64_t EltOffset = alignTo(Size: ExplicitArgOffset, A: ABITypeAlign) + BaseOffset;
246 ExplicitArgOffset = alignTo(Size: ExplicitArgOffset, A: ABITypeAlign) + AllocSize;
247
248 // Skip inreg arguments which should be preloaded.
249 if (Arg.use_empty() || Arg.hasInRegAttr())
250 continue;
251
252 // If this is byval, the loads are already explicit in the function. We just
253 // need to rewrite the pointer values.
254 if (IsByRef) {
255 Value *ArgOffsetPtr = Builder.CreateConstInBoundsGEP1_64(
256 Ty: Builder.getInt8Ty(), Ptr: KernArgSegment, Idx0: EltOffset,
257 Name: Arg.getName() + ".byval.kernarg.offset");
258
259 Value *CastOffsetPtr =
260 Builder.CreateAddrSpaceCast(V: ArgOffsetPtr, DestTy: Arg.getType());
261 Arg.replaceAllUsesWith(V: CastOffsetPtr);
262 continue;
263 }
264
265 if (PointerType *PT = dyn_cast<PointerType>(Val: ArgTy)) {
266 // FIXME: Hack. We rely on AssertZext to be able to fold DS addressing
267 // modes on SI to know the high bits are 0 so pointer adds don't wrap. We
268 // can't represent this with range metadata because it's only allowed for
269 // integer types.
270 if ((PT->getAddressSpace() == AMDGPUAS::LOCAL_ADDRESS ||
271 PT->getAddressSpace() == AMDGPUAS::REGION_ADDRESS) &&
272 !ST.hasUsableDSOffset())
273 continue;
274 }
275
276 auto *VT = dyn_cast<FixedVectorType>(Val: ArgTy);
277 bool IsV3 = VT && VT->getNumElements() == 3;
278 bool DoShiftOpt = Size < 32 && !ArgTy->isAggregateType();
279
280 VectorType *V4Ty = nullptr;
281
282 int64_t AlignDownOffset = alignDown(Value: EltOffset, Align: 4);
283 int64_t OffsetDiff = EltOffset - AlignDownOffset;
284 Align AdjustedAlign = commonAlignment(
285 A: KernArgBaseAlign, Offset: DoShiftOpt ? AlignDownOffset : EltOffset);
286
287 Value *ArgPtr;
288 Type *AdjustedArgTy;
289 if (DoShiftOpt) { // FIXME: Handle aggregate types
290 // Since we don't have sub-dword scalar loads, avoid doing an extload by
291 // loading earlier than the argument address, and extracting the relevant
292 // bits.
293 // TODO: Update this for GFX12 which does have scalar sub-dword loads.
294 //
295 // Additionally widen any sub-dword load to i32 even if suitably aligned,
296 // so that CSE between different argument loads works easily.
297 ArgPtr = Builder.CreateConstInBoundsGEP1_64(
298 Ty: Builder.getInt8Ty(), Ptr: KernArgSegment, Idx0: AlignDownOffset,
299 Name: Arg.getName() + ".kernarg.offset.align.down");
300 AdjustedArgTy = Builder.getInt32Ty();
301 } else {
302 ArgPtr = Builder.CreateConstInBoundsGEP1_64(
303 Ty: Builder.getInt8Ty(), Ptr: KernArgSegment, Idx0: EltOffset,
304 Name: Arg.getName() + ".kernarg.offset");
305 AdjustedArgTy = ArgTy;
306 }
307
308 if (IsV3 && Size >= 32) {
309 V4Ty = FixedVectorType::get(ElementType: VT->getElementType(), NumElts: 4);
310 // Use the hack that clang uses to avoid SelectionDAG ruining v3 loads
311 AdjustedArgTy = V4Ty;
312 }
313
314 LoadInst *Load =
315 Builder.CreateAlignedLoad(Ty: AdjustedArgTy, Ptr: ArgPtr, Align: AdjustedAlign);
316 Load->setMetadata(KindID: LLVMContext::MD_invariant_load, Node: MDNode::get(Context&: Ctx, MDs: {}));
317
318 MDBuilder MDB(Ctx);
319
320 if (Arg.hasAttribute(Kind: Attribute::NoUndef) && AdjustedArgTy == ArgTy)
321 Load->setMetadata(KindID: LLVMContext::MD_noundef, Node: MDNode::get(Context&: Ctx, MDs: {}));
322
323 if (Arg.hasAttribute(Kind: Attribute::Range) && AdjustedArgTy == ArgTy) {
324 const ConstantRange &Range =
325 Arg.getAttribute(Kind: Attribute::Range).getValueAsConstantRange();
326 Load->setMetadata(KindID: LLVMContext::MD_range,
327 Node: MDB.createRange(Lo: Range.getLower(), Hi: Range.getUpper()));
328 }
329
330 if (Arg.hasAttribute(Kind: Attribute::NoFPClass) && AdjustedArgTy == ArgTy) {
331 FPClassTest Mask = Arg.getNoFPClass();
332 Load->setMetadata(
333 KindID: LLVMContext::MD_nofpclass,
334 Node: MDNode::get(Context&: Ctx, MDs: ConstantAsMetadata::get(
335 C: ConstantInt::get(Ty: Type::getInt32Ty(C&: Ctx), V: Mask))));
336 }
337
338 if (isa<PointerType>(Val: ArgTy)) {
339 if (Arg.hasNonNullAttr())
340 Load->setMetadata(KindID: LLVMContext::MD_nonnull, Node: MDNode::get(Context&: Ctx, MDs: {}));
341
342 uint64_t DerefBytes = Arg.getDereferenceableBytes();
343 if (DerefBytes != 0) {
344 Load->setMetadata(
345 KindID: LLVMContext::MD_dereferenceable,
346 Node: MDNode::get(Context&: Ctx,
347 MDs: MDB.createConstant(
348 C: ConstantInt::get(Ty: Builder.getInt64Ty(), V: DerefBytes))));
349 }
350
351 uint64_t DerefOrNullBytes = Arg.getDereferenceableOrNullBytes();
352 if (DerefOrNullBytes != 0) {
353 Load->setMetadata(
354 KindID: LLVMContext::MD_dereferenceable_or_null,
355 Node: MDNode::get(Context&: Ctx,
356 MDs: MDB.createConstant(C: ConstantInt::get(Ty: Builder.getInt64Ty(),
357 V: DerefOrNullBytes))));
358 }
359
360 if (MaybeAlign ParamAlign = Arg.getParamAlign()) {
361 Load->setMetadata(
362 KindID: LLVMContext::MD_align,
363 Node: MDNode::get(Context&: Ctx, MDs: MDB.createConstant(C: ConstantInt::get(
364 Ty: Builder.getInt64Ty(), V: ParamAlign->value()))));
365 }
366 }
367
368 if (DoShiftOpt) {
369 Value *ExtractBits = OffsetDiff == 0 ?
370 Load : Builder.CreateLShr(LHS: Load, RHS: OffsetDiff * 8);
371
372 IntegerType *ArgIntTy = Builder.getIntNTy(N: Size);
373 Value *Trunc = Builder.CreateTrunc(V: ExtractBits, DestTy: ArgIntTy);
374 Value *NewVal = Builder.CreateBitCast(V: Trunc, DestTy: ArgTy,
375 Name: Arg.getName() + ".load");
376 Arg.replaceAllUsesWith(V: NewVal);
377 } else if (IsV3) {
378 Value *Shuf = Builder.CreateShuffleVector(V: Load, Mask: ArrayRef<int>{0, 1, 2},
379 Name: Arg.getName() + ".load");
380 Arg.replaceAllUsesWith(V: Shuf);
381 } else {
382 Load->setName(Arg.getName() + ".load");
383 Arg.replaceAllUsesWith(V: Load);
384 }
385 }
386
387 KernArgSegment->addRetAttr(
388 Attr: Attribute::getWithAlignment(Context&: Ctx, Alignment: std::max(a: KernArgBaseAlign, b: MaxAlign)));
389
390 return true;
391}
392
393bool AMDGPULowerKernelArguments::runOnFunction(Function &F) {
394 auto &TPC = getAnalysis<TargetPassConfig>();
395 const TargetMachine &TM = TPC.getTM<TargetMachine>();
396 DominatorTree &DT = getAnalysis<DominatorTreeWrapperPass>().getDomTree();
397 return lowerKernelArguments(F, TM, DT);
398}
399
400INITIALIZE_PASS_BEGIN(AMDGPULowerKernelArguments, DEBUG_TYPE,
401 "AMDGPU Lower Kernel Arguments", false, false)
402INITIALIZE_PASS_END(AMDGPULowerKernelArguments, DEBUG_TYPE, "AMDGPU Lower Kernel Arguments",
403 false, false)
404
405char AMDGPULowerKernelArguments::ID = 0;
406
407FunctionPass *llvm::createAMDGPULowerKernelArgumentsPass() {
408 return new AMDGPULowerKernelArguments();
409}
410
411PreservedAnalyses
412AMDGPULowerKernelArgumentsPass::run(Function &F, FunctionAnalysisManager &AM) {
413 DominatorTree &DT = AM.getResult<DominatorTreeAnalysis>(IR&: F);
414 bool Changed = lowerKernelArguments(F, TM, DT);
415 if (Changed) {
416 // TODO: Preserves a lot more.
417 PreservedAnalyses PA;
418 PA.preserveSet<CFGAnalyses>();
419 return PA;
420 }
421
422 return PreservedAnalyses::all();
423}
424