1#ifndef PROTEUS_KERNELARGVISITOR_H
2#define PROTEUS_KERNELARGVISITOR_H
9#include <llvm/Analysis/PtrUseVisitor.h>
10#include <llvm/Analysis/ValueTracking.h>
12#include <llvm/ADT/DenseMap.h>
13#include <llvm/ADT/DenseSet.h>
14#include <llvm/ADT/Hashing.h>
15#include <llvm/ADT/SmallPtrSet.h>
16#include <llvm/ADT/SmallVector.h>
18#include <llvm/Analysis/ValueTracking.h>
19#include <llvm/IR/CFG.h>
20#include <llvm/IR/Constants.h>
21#include <llvm/IR/DataLayout.h>
22#include <llvm/IR/DebugInfo.h>
23#include <llvm/IR/Function.h>
24#include <llvm/IR/InstrTypes.h>
25#include <llvm/IR/Instruction.h>
26#include <llvm/IR/Instructions.h>
27#include <llvm/IR/Metadata.h>
28#include <llvm/IR/Module.h>
29#include <llvm/IR/PassManager.h>
30#include <llvm/IR/Type.h>
39 if (
auto *TermInst = dyn_cast<ReturnInst>(BB.getTerminator()))
46 ArrayRef<unsigned> Indices) {
47 LLVMContext &Ctx = AggTy->getContext();
48 SmallVector<Value *, 8> GEPIndices;
49 GEPIndices.push_back(ConstantInt::get(Type::getInt32Ty(Ctx), 0));
50 for (
unsigned Idx : Indices)
51 GEPIndices.push_back(ConstantInt::get(Type::getInt32Ty(Ctx), Idx));
52 return DL.getIndexedOffsetInType(AggTy, GEPIndices);
63 int64_t LHSOffset = 0;
64 int64_t RHSOffset = 0;
65 Value *LHSBase = GetPointerBaseWithConstantOffset(LHS, LHSOffset, DL);
66 Value *RHSBase = GetPointerBaseWithConstantOffset(RHS, RHSOffset, DL);
67 return LHSBase && RHSBase && LHSBase == RHSBase && LHSOffset == RHSOffset;
75 Instruction *Before, Value *Address,
76 SmallPtrSetImpl<BasicBlock *> &Visited,
78 if (!Visited.insert(BB).second) {
83 for (Instruction *I = Before ? Before->getPrevNode() : BB->getTerminator(); I;
84 I = I->getPrevNode()) {
85 auto *SI = dyn_cast<StoreInst>(I);
86 if (SI && SI->getValueOperand()->getType()->isPointerTy() &&
88 Result.
Values.push_back(SI->getValueOperand());
97 for (BasicBlock *Pred : predecessors(BB))
108 SmallPtrSet<BasicBlock *, 8> Visited;
115 SmallPtrSetImpl<Value *> &Visited);
122 SmallPtrSetImpl<Value *> &Visited) {
129 for (Value *V : drop_begin(Stores.
Values)) {
141 SmallPtrSetImpl<Value *> &Visited) {
142 auto *LI = dyn_cast<LoadInst>(V);
143 if (!LI || !LI->getType()->isPointerTy())
145 if (!Visited.insert(V).second)
158 SmallPtrSet<Value *, 8> Visited;
170 if (!LI.getType()->isPointerTy())
172 const Value *Storage = getUnderlyingObject(LI.getPointerOperand());
173 auto *Slot = dyn_cast_or_null<AllocaInst>(Storage);
174 return Slot && Slot->getAllocatedType()->isPointerTy();
203 const DataLayout &DL;
204 SmallVector<WorkItem> WorkList;
205 SmallDenseSet<Value *> Seen;
208 uint32_t KernelArg = 0;
209 Function *KernelFunction =
nullptr;
212 std::optional<RuntimeConstantType> ChangedRC = std::nullopt;
213 bool AnalysisSuccess =
false;
214 bool AnalysisFailed =
false;
218 int64_t Off,
const DataLayout &Dl)
219 : LambdaCB(LambdaCBArg), DL(Dl),
Offset(Off) {
220 WorkList.push_back({Start, LastSeen});
223 std::optional<Value *> getCallBaseIdentityArgOperand(CallBase &CB) {
224 auto *CalledFunction = CB.getCalledFunction();
226 <<
"Checking if function is identity " << *CalledFunction <<
"\n");
227 auto RetInstOpt =
getRetInst(*CalledFunction);
230 auto *RetInst = *RetInstOpt;
233 << *RetInst->getReturnValue() <<
"\n");
234 for (
size_t ArgNum = 0; ArgNum < CalledFunction->arg_size(); ++ArgNum) {
236 <<
"Called Fn arg " << *CalledFunction->getArg(ArgNum) <<
"\n");
237 if (RetInst->getReturnValue() == CalledFunction->getArg(ArgNum))
238 return CB.getArgOperand(ArgNum);
254 auto back() {
return WorkList.back(); }
256 bool seen(Value *Val) {
return Seen.contains(Val); }
258 bool empty() {
return WorkList.empty(); }
264 inline std::optional<LambdaKernelArgAnalysis>
266 int64_t StartOffset) {
269 while (!Visitor.empty() && !Visitor.success() && !Visitor.failed()) {
270 auto [V, AccessedFrom] = Visitor.back();
271 Visitor.MemoryAnalysisPtrUse = AccessedFrom;
276 Visitor.markAsSeen(V);
278 if (
auto *I = dyn_cast<Instruction>(V))
280 else if (
auto *A = dyn_cast<Argument>(V))
281 Visitor.visitArgument(*A);
285 if (!Visitor.success() || Visitor.failed())
287 return Visitor.getKernelArgInfo();
290 inline std::optional<FunctionAnalysis> analyzeFunction(CallBase &CB,
291 int64_t StartOffset) {
293 auto &F = *CB.getCalledFunction();
299 while (!Visitor.empty() && !Visitor.success() && !Visitor.failed()) {
300 auto [V, AccessedFrom] = Visitor.back();
302 <<
"Function analysis visiting " << *V <<
" with offset "
303 << Visitor.getOffset() <<
"\n");
304 Visitor.MemoryAnalysisPtrUse = AccessedFrom;
309 Visitor.markAsSeen(V);
311 if (
auto *I = dyn_cast<Instruction>(V))
313 else if (
auto *A = dyn_cast<Argument>(V)) {
314 if (A->getParent() == CB.getCalledFunction()) {
315 Result.
PtrArgToCB = CB.getArgOperand(A->getArgNo());
316 Result.
Offset = Visitor.Offset;
319 <<
"Function analysis found termination case "
320 << *Result.
PtrArgToCB <<
" with offset " << Visitor.getOffset()
324 Visitor.visitArgument(*A);
335 : LambdaCB(LambdaCB), DL(M.getDataLayout()),
Offset(0) {
336 auto *ClosurePtr = LambdaCB->getArgOperand(0);
337 WorkList.push_back({ClosurePtr, LambdaCB});
341 WorkList.push_back({SI.getValueOperand(), &SI});
345 if (!CB.getCalledFunction() || CB.getCalledFunction()->isDeclaration()) {
347 <<
"[Lambda arg analysis]: Cannot trace indirect or declaration "
350 AnalysisFailed =
true;
351 AnalysisSuccess =
false;
356 auto SubAnalysis = analyzeFunction(CB,
Offset);
358 WorkList.push_back({SubAnalysis.value().PtrArgToCB, &CB});
359 Offset = SubAnalysis.value().Offset;
363 <<
"Function analysis of \n"
364 << *CB.getCalledFunction() <<
" failed \n");
365 AnalysisFailed =
true;
366 AnalysisSuccess =
false;
376 SmallPtrSet<Value *, 8> Visited;
377 if (Value *StoredPointer =
379 WorkList.push_back({StoredPointer, &LI});
384 <<
"[Lambda arg analysis]: Pointer spill load has ambiguous "
387 AnalysisFailed =
true;
388 AnalysisSuccess =
false;
393 AnalysisFailed =
true;
394 AnalysisSuccess =
false;
397 WorkList.push_back({Res->DominatingWrite, &LI});
401 WorkList.push_back({LI.getPointerOperand(), &LI});
405 APInt StepOffset(DL.getIndexTypeSizeInBits(GEP.getType()), 0);
406 if (!GEP.accumulateConstantOffset(DL, StepOffset)) {
407 AnalysisFailed =
true;
408 AnalysisSuccess =
false;
411 Offset += StepOffset.getSExtValue();
412 WorkList.push_back({GEP.getPointerOperand(), &GEP});
417 DL, EVI.getAggregateOperand()->getType(), EVI.getIndices());
419 WorkList.push_back({EVI.getAggregateOperand(), &EVI});
427 auto *AggregateOperand = IVI.getAggregateOperand();
430 while (Cur && AggregateOperand) {
432 DL, Cur->getAggregateOperand()->getType(), Cur->getIndices());
433 TypeSize InsertedSize =
434 DL.getTypeAllocSize(Cur->getInsertedValueOperand()->getType());
436 <<
"Curr offset " << CurOffset <<
"\nOffset " <<
Offset <<
"\n");
437 if (!InsertedSize.isScalable() &&
Offset >= CurOffset &&
438 static_cast<uint64_t
>(
Offset - CurOffset) <
439 InsertedSize.getFixedValue()) {
445 WorkList.push_back({Cur->getInsertedValueOperand(), &IVI});
448 Cur = dyn_cast<InsertValueInst>(AggregateOperand);
450 AggregateOperand = Cur->getAggregateOperand();
452 AnalysisFailed =
true;
453 AnalysisSuccess =
false;
463 WorkList.push_back({Res->DominatingWrite, &Alloca});
476 WorkList.push_back({Res->DominatingWrite, &BC});
483 WorkList.push_back({ASC.getPointerOperand(), &ASC});
489 WorkList.push_back({Res->DominatingWrite, &ASC});
496 auto *IntegerVal = ITP.getOperand(0);
497 auto *Ptr = dyn_cast<PtrToIntInst>(IntegerVal);
499 AnalysisSuccess =
false;
500 AnalysisFailed =
true;
505 WorkList.push_back({Ptr->getPointerOperand(), &ITP});
509 auto *MT = dyn_cast<MemTransferInst>(&I);
512 AnalysisFailed =
true;
516 int64_t DstOff = 0, SrcOff = 0;
518 GetPointerBaseWithConstantOffset(MT->getRawDest(), DstOff, DL);
520 GetPointerBaseWithConstantOffset(MT->getRawSource(), SrcOff, DL);
521 if (!DstBase || !SrcBase) {
522 AnalysisFailed =
true;
528 if (
auto *LenC = dyn_cast<ConstantInt>(MT->getLength())) {
529 uint64_t Len = LenC->getZExtValue();
530 if (
Offset < DstOff || uint64_t(
Offset - DstOff) >= Len) {
531 AnalysisFailed =
true;
536 WorkList.push_back({SrcBase->stripPointerCasts(), &I});
541 AnalysisFailed =
true;
546 WorkList.push_back({TI.getOperand(0), &TI});
550 Function *F = A.getParent();
552 <<
"Visiting argument with parent function = \n"
554 auto ArgNum = A.getArgNo();
558 if (F->getCallingConv() == CallingConv::AMDGPU_KERNEL ||
559 F->getCallingConv() == CallingConv::PTX_Kernel ||
560 (F->hasMetadata(
"proteus.jit") &&
561 !F->hasMetadata(
"proteus.wrapper_call") &&
562 !F->hasMetadata(
"proteus.registered_lambda"))) {
564 <<
"Found termination case from function " << F->getName() <<
"\n");
567 AnalysisSuccess =
true;
573 for (User *U : F->users()) {
574 auto *CB = dyn_cast<CallBase>(U);
578 <<
"Analysis crossed interprocedural boundary at "
579 << *CB->getArgOperand(ArgNum) <<
"\n");
580 WorkList.push_back({CB->getArgOperand(ArgNum), CB});
586 <<
"[Lambda arg analysis]: Unhandled instruction "
587 << I.getOpcodeName() <<
": " << I <<
"\n");
588 AnalysisFailed =
true;
593 if (P.getNumIncomingValues() == 0) {
594 AnalysisFailed =
true;
600 if (!FirstAnalysis) {
601 AnalysisFailed =
true;
602 AnalysisSuccess =
false;
605 auto BaseSlot = FirstAnalysis->KernelArgIndex;
606 auto BaseOffset = FirstAnalysis->Offset;
608 for (
size_t Idx = 1; Idx < P.getNumIncomingValues(); ++Idx) {
609 auto Analysis = cloneAndAnalyze(P.getIncomingValue(Idx),
611 if (!Analysis || Analysis->KernelArgIndex != BaseSlot ||
612 Analysis->Offset != BaseOffset) {
613 AnalysisFailed =
true;
614 AnalysisSuccess =
false;
618 KernelFunction = FirstAnalysis->KernelFunction;
620 KernelArg = BaseSlot;
621 ChangedRC = FirstAnalysis->ChangedRCLayout;
622 AnalysisSuccess =
true;
623 AnalysisFailed =
false;
635 if (!TrueAnalysis || !FalseAnalysis ||
636 TrueAnalysis->KernelFunction != FalseAnalysis->KernelFunction ||
637 TrueAnalysis->KernelArgIndex != FalseAnalysis->KernelArgIndex ||
638 TrueAnalysis->Offset != FalseAnalysis->Offset ||
639 TrueAnalysis->ChangedRCLayout != FalseAnalysis->ChangedRCLayout) {
641 <<
"[Lambda arg analysis]: Select arms do not resolve to the "
642 "same kernel argument and offset: "
644 AnalysisFailed =
true;
645 AnalysisSuccess =
false;
649 KernelFunction = TrueAnalysis->KernelFunction;
650 KernelArg = TrueAnalysis->KernelArgIndex;
651 Offset = TrueAnalysis->Offset;
652 ChangedRC = TrueAnalysis->ChangedRCLayout;
653 AnalysisSuccess =
true;
654 AnalysisFailed =
false;
660 DenseMap<CallBase *, LambdaKernelArgAnalysis> &CallBaseToArgOffset,
661 const SmallVector<CallBase *> &CBToAnalyze) {
663 for (
auto *FunctorCB : CBToAnalyze) {
666 auto [V, LastSeen] = Visitor.
back();
675 if (
auto *I = dyn_cast<Instruction>(V))
677 else if (
auto *A = dyn_cast<Argument>(V))
684 <<
"[WARNING]: Kernel arg analysis failed for functor beginning at "
685 << *FunctorCB <<
"\n");
691 CallBaseToArgOffset[FunctorCB] = Info;
uint32_t int32_t Type
Definition CompilerInterfaceDevice.cpp:98
uint64_t uint32_t uint32_t int64_t Offset
Definition CompilerInterfaceDevice.cpp:75
#define DEBUG(x)
Definition Helpers.h:14
Definition KernelArgVisitor.h:200
void visitIntToPtr(IntToPtrInst &ITP)
Definition KernelArgVisitor.h:495
void visitInsertValueInst(InsertValueInst &IVI)
Definition KernelArgVisitor.h:426
void visitLoadInst(LoadInst &LI)
Definition KernelArgVisitor.h:369
LambdaKernelArgAnalysis getKernelArgInfo()
Definition KernelArgVisitor.h:244
void markAsSeen(Value *Val)
Definition KernelArgVisitor.h:257
void visitMemIntrinsic(MemIntrinsic &I)
Definition KernelArgVisitor.h:508
bool failed()
Definition KernelArgVisitor.h:260
void visitInstruction(Instruction &I)
Definition KernelArgVisitor.h:584
auto getOffset()
Definition KernelArgVisitor.h:261
auto back()
Definition KernelArgVisitor.h:254
void visitStoreInst(StoreInst &SI)
Definition KernelArgVisitor.h:340
void visitSelectInst(SelectInst &S)
Definition KernelArgVisitor.h:626
void visitCallBase(CallBase &CB)
Definition KernelArgVisitor.h:344
void visitBitCastInst(BitCastInst &BC)
Definition KernelArgVisitor.h:471
void visitArgument(Argument &A)
Definition KernelArgVisitor.h:549
void visitExtractValueInst(ExtractValueInst &EVI)
Definition KernelArgVisitor.h:415
void visitIntrinsicInst(IntrinsicInst &)
Definition KernelArgVisitor.h:540
void popBack()
Definition KernelArgVisitor.h:255
LambdaArgVisitor(CallBase *LambdaCB, Module &M)
Definition KernelArgVisitor.h:334
bool success()
Definition KernelArgVisitor.h:259
bool empty()
Definition KernelArgVisitor.h:258
bool seen(Value *Val)
Definition KernelArgVisitor.h:256
void visitAddrSpaceCastInst(AddrSpaceCastInst &ASC)
Definition KernelArgVisitor.h:482
void visitGetElementPtrInst(GetElementPtrInst &GEP)
Definition KernelArgVisitor.h:404
void visitAllocaInst(AllocaInst &Alloca)
Definition KernelArgVisitor.h:457
void visitPHINode(PHINode &P)
Definition KernelArgVisitor.h:592
Value * MemoryAnalysisPtrUse
Definition KernelArgVisitor.h:253
void visitTruncInst(TruncInst &TI)
Definition KernelArgVisitor.h:545
static llvm::raw_ostream & logs(const std::string &Name)
Definition Logger.h:19
Definition CompiledLibrary.h:8
Definition MemoryCache.h:27
std::optional< LambdaPtrUseAnalysis > getDominatingUse(const DataLayout &DL, Value *ValueNeedingAnalysis, Value *SeenUse, int64_t TargetOffset, CallBase *LambdaCB=nullptr)
Definition KernelArgPtrUseVisitor.h:557
Value * getUniqueReachingPointer(const DataLayout &DL, const ReachingPointerStores &Stores, SmallPtrSetImpl< Value * > &Visited)
Definition KernelArgVisitor.h:120
bool analyzeLambdaUses(llvm::Module &M, DenseMap< CallBase *, LambdaKernelArgAnalysis > &CallBaseToArgOffset, const SmallVector< CallBase * > &CBToAnalyze)
Definition KernelArgVisitor.h:658
Value * getPointerLoadOrigin(const DataLayout &DL, Value *V, SmallPtrSetImpl< Value * > &Visited)
Definition KernelArgVisitor.h:140
int64_t getValueIndicesOffset(const DataLayout &DL, Type *AggTy, ArrayRef< unsigned > Indices)
Definition KernelArgVisitor.h:45
RuntimeConstantType convertTypeToRuntimeConstantType(Type *Ty)
Definition RuntimeConstantTypeHelpers.h:21
void collectReachingPointerStores(const DataLayout &DL, BasicBlock *BB, Instruction *Before, Value *Address, SmallPtrSetImpl< BasicBlock * > &Visited, ReachingPointerStores &Result)
Definition KernelArgVisitor.h:74
bool isSamePointerAddress(const DataLayout &DL, Value *LHS, Value *RHS)
Definition KernelArgVisitor.h:62
bool hasAmbiguousReachingPointers(const DataLayout &DL, const ReachingPointerStores &Stores)
Definition KernelArgVisitor.h:154
std::optional< ReturnInst * > getRetInst(Function &F)
Definition KernelArgVisitor.h:37
ReachingPointerStores getReachingPointerStores(const DataLayout &DL, LoadInst &LI)
Definition KernelArgVisitor.h:105
bool isPointerSpillLoad(const LoadInst &LI)
Definition KernelArgVisitor.h:169
Definition KernelArgVisitor.h:186
std::optional< RuntimeConstantType > ChangedRCLayout
Definition KernelArgVisitor.h:192
int64_t Offset
Definition KernelArgVisitor.h:189
Value * PtrArgToCB
Definition KernelArgVisitor.h:187
uint32_t ArgIndex
Definition KernelArgVisitor.h:188
Definition KernelArgVisitor.h:177
uint32_t KernelArgIndex
Definition KernelArgVisitor.h:179
std::optional< RuntimeConstantType > ChangedRCLayout
Definition KernelArgVisitor.h:183
Function * KernelFunction
Definition KernelArgVisitor.h:178
int64_t Offset
Definition KernelArgVisitor.h:180
Definition KernelArgVisitor.h:55
SmallVector< Value *, 4 > Values
Definition KernelArgVisitor.h:56
bool Complete
Definition KernelArgVisitor.h:57
Definition KernelArgVisitor.h:195
Value * CurVal
Definition KernelArgVisitor.h:196
Value * Src
Definition KernelArgVisitor.h:197