Proteus
Programmable JIT compilation and optimization for C/C++ using LLVM
Loading...
Searching...
No Matches
KernelArgVisitor.h
Go to the documentation of this file.
1#ifndef PROTEUS_KERNELARGVISITOR_H
2#define PROTEUS_KERNELARGVISITOR_H
3
4#include "Helpers.h"
9#include <llvm/Analysis/PtrUseVisitor.h>
10#include <llvm/Analysis/ValueTracking.h>
11
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>
17
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>
31#include <memory>
32#include <optional>
33
34namespace proteus {
35using namespace llvm;
36
37std::optional<ReturnInst *> getRetInst(Function &F) {
38 for (auto &BB : F) {
39 if (auto *TermInst = dyn_cast<ReturnInst>(BB.getTerminator()))
40 return TermInst;
41 }
42 return std::nullopt;
43}
44
45inline int64_t getValueIndicesOffset(const DataLayout &DL, Type *AggTy,
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);
53}
54
56 SmallVector<Value *, 4> Values;
57 bool Complete = true;
58};
59
60// Return whether LHS and RHS name exactly the same byte address after peeling
61// constant-offset pointer arithmetic and casts.
62inline bool isSamePointerAddress(const DataLayout &DL, Value *LHS, Value *RHS) {
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;
68}
69
70// Collect the closest pointer-valued store to Address on every CFG path that
71// reaches Before. A path with no store, or a revisited block (such as a loop),
72// marks the result incomplete so callers conservatively decline rather than
73// infer a store that is not guaranteed to reach the load.
74inline void collectReachingPointerStores(const DataLayout &DL, BasicBlock *BB,
75 Instruction *Before, Value *Address,
76 SmallPtrSetImpl<BasicBlock *> &Visited,
77 ReachingPointerStores &Result) {
78 if (!Visited.insert(BB).second) {
79 Result.Complete = false;
80 return;
81 }
82
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() &&
87 isSamePointerAddress(DL, SI->getPointerOperand(), Address)) {
88 Result.Values.push_back(SI->getValueOperand());
89 return;
90 }
91 }
92
93 if (pred_empty(BB)) {
94 Result.Complete = false;
95 return;
96 }
97 for (BasicBlock *Pred : predecessors(BB))
98 collectReachingPointerStores(DL, Pred, nullptr, Address, Visited, Result);
99}
100
101// Resolve pointer spills by walking backwards from the load through the CFG.
102// This finds the nearest store on every incoming path instead of depending on
103// the arbitrary order in which Value::users() happens to enumerate writes.
104// The result is complete only when every incoming path contributes a store.
106 LoadInst &LI) {
108 SmallPtrSet<BasicBlock *, 8> Visited;
109 collectReachingPointerStores(DL, LI.getParent(), &LI, LI.getPointerOperand(),
110 Visited, Result);
111 return Result;
112}
113
114inline Value *getPointerLoadOrigin(const DataLayout &DL, Value *V,
115 SmallPtrSetImpl<Value *> &Visited);
116
117// Return one source pointer when every reaching store has the same origin.
118// Nested pointer-spill loads are recursively resolved; distinct origins or
119// cycles are ambiguous and return nullptr.
120inline Value *getUniqueReachingPointer(const DataLayout &DL,
121 const ReachingPointerStores &Stores,
122 SmallPtrSetImpl<Value *> &Visited) {
123 if (!Stores.Complete || Stores.Values.empty())
124 return nullptr;
125
126 Value *First = getPointerLoadOrigin(DL, Stores.Values.front(), Visited);
127 if (!First)
128 return nullptr;
129 for (Value *V : drop_begin(Stores.Values)) {
130 Value *Origin = getPointerLoadOrigin(DL, V, Visited);
131 if (Origin != First)
132 return nullptr;
133 }
134 return First;
135}
136
137// Resolve a pointer value through nested pointer-spill loads. Non-load pointer
138// values are already origins. Visited prevents cyclic spill graphs from being
139// mistaken for a unique source.
140inline Value *getPointerLoadOrigin(const DataLayout &DL, Value *V,
141 SmallPtrSetImpl<Value *> &Visited) {
142 auto *LI = dyn_cast<LoadInst>(V);
143 if (!LI || !LI->getType()->isPointerTy())
144 return V;
145 if (!Visited.insert(V).second)
146 return nullptr;
147
149 return getUniqueReachingPointer(DL, Stores, Visited);
150}
151
152// Report ambiguity only for complete reaching-store sets. Incomplete sets can
153// still be handled by ordinary backward memory-use analysis.
154inline bool hasAmbiguousReachingPointers(const DataLayout &DL,
155 const ReachingPointerStores &Stores) {
156 if (!Stores.Complete || Stores.Values.empty())
157 return false;
158 SmallPtrSet<Value *, 8> Visited;
159 return !getUniqueReachingPointer(DL, Stores, Visited);
160}
161
162// A compiler spill is a temporary local slot the compiler uses to save an SSA
163// pointer value (for example, `alloca ptr`, followed by `store ptr` and a
164// later `load ptr`). A pointer-valued load is not necessarily such a spill:
165// it can instead read an ordinary pointer field from a closure/context
166// aggregate. The reaching-store recovery below is only valid for a local
167// `alloca ptr` slot; aggregate fields must be traced backwards through their
168// address.
169inline bool isPointerSpillLoad(const LoadInst &LI) {
170 if (!LI.getType()->isPointerTy())
171 return false;
172 const Value *Storage = getUnderlyingObject(LI.getPointerOperand());
173 auto *Slot = dyn_cast_or_null<AllocaInst>(Storage);
174 return Slot && Slot->getAllocatedType()->isPointerTy();
175}
176
178 Function *KernelFunction = nullptr;
179 uint32_t KernelArgIndex = 0;
180 int64_t Offset = 0;
181 // Sometimes instructions like ptrtoint --> inttoptr change the layout of
182 // the kernel args.
183 std::optional<RuntimeConstantType> ChangedRCLayout = std::nullopt;
184};
185
187 Value *PtrArgToCB = nullptr;
188 uint32_t ArgIndex = 0;
189 int64_t Offset = 0;
190 // Sometimes instructions like ptrtoint --> inttoptr change the layout of
191 // the kernel args.
192 std::optional<RuntimeConstantType> ChangedRCLayout = std::nullopt;
193};
194
195struct WorkItem {
196 Value *CurVal;
197 Value *Src;
198};
199
200class LambdaArgVisitor : public InstVisitor<LambdaArgVisitor> {
201private:
202 CallBase *LambdaCB;
203 const DataLayout &DL;
204 SmallVector<WorkItem> WorkList;
205 SmallDenseSet<Value *> Seen;
206
207 int64_t Offset;
208 uint32_t KernelArg = 0;
209 Function *KernelFunction = nullptr;
210 // instructions like inttoptr can change the runtimeconstant type
211 // we to read from the Blob
212 std::optional<RuntimeConstantType> ChangedRC = std::nullopt;
213 bool AnalysisSuccess = false;
214 bool AnalysisFailed = false;
215
216 // Constructor used for cloning and merging branches of phi node analysis
217 LambdaArgVisitor(Value *Start, Value *LastSeen, CallBase *LambdaCBArg,
218 int64_t Off, const DataLayout &Dl)
219 : LambdaCB(LambdaCBArg), DL(Dl), Offset(Off) {
220 WorkList.push_back({Start, LastSeen});
221 }
222
223 std::optional<Value *> getCallBaseIdentityArgOperand(CallBase &CB) {
224 auto *CalledFunction = CB.getCalledFunction();
225 DEBUG(Logger::logs("proteus-pass")
226 << "Checking if function is identity " << *CalledFunction << "\n");
227 auto RetInstOpt = getRetInst(*CalledFunction);
228 if (!RetInstOpt)
229 return std::nullopt;
230 auto *RetInst = *RetInstOpt;
231
232 DEBUG(Logger::logs("proteus-pass") << "CB called function return inst "
233 << *RetInst->getReturnValue() << "\n");
234 for (size_t ArgNum = 0; ArgNum < CalledFunction->arg_size(); ++ArgNum) {
235 DEBUG(Logger::logs("proteus-pass")
236 << "Called Fn arg " << *CalledFunction->getArg(ArgNum) << "\n");
237 if (RetInst->getReturnValue() == CalledFunction->getArg(ArgNum))
238 return CB.getArgOperand(ArgNum);
239 }
240 return std::nullopt;
241 }
242
243public:
245 return LambdaKernelArgAnalysis{KernelFunction, KernelArg, Offset,
246 ChangedRC};
247 }
248 // Whenever the analysis encounters an instruction returning a Ptr
249 // memory analysis is required to identify a dominating write,
250 // which is where the LambdaArgVisitor continues its analysis.
251 // This pointer identifies which Ptr use the main analysis used
252 // to discover the Ptr needing analysis, so as to prevent cycles.
253 Value *MemoryAnalysisPtrUse = nullptr;
254 auto back() { return WorkList.back(); }
255 void popBack() { WorkList.pop_back(); }
256 bool seen(Value *Val) { return Seen.contains(Val); }
257 void markAsSeen(Value *Val) { Seen.insert(Val); }
258 bool empty() { return WorkList.empty(); }
259 bool success() { return AnalysisSuccess; }
260 bool failed() { return AnalysisFailed; }
261 auto getOffset() { return Offset; }
262
263private:
264 inline std::optional<LambdaKernelArgAnalysis>
265 cloneAndAnalyze(Value *Start, Value *MemoryAnalysisPtrUse,
266 int64_t StartOffset) {
267 LambdaArgVisitor Visitor(Start, MemoryAnalysisPtrUse, LambdaCB, StartOffset,
268 DL);
269 while (!Visitor.empty() && !Visitor.success() && !Visitor.failed()) {
270 auto [V, AccessedFrom] = Visitor.back();
271 Visitor.MemoryAnalysisPtrUse = AccessedFrom;
272 Visitor.popBack();
273 // Prevent loops/infinite recursion
274 if (Visitor.seen(V))
275 continue;
276 Visitor.markAsSeen(V);
277 // Analyze the instruction
278 if (auto *I = dyn_cast<Instruction>(V))
279 Visitor.visit(*I);
280 else if (auto *A = dyn_cast<Argument>(V))
281 Visitor.visitArgument(*A);
282 else
283 continue;
284 }
285 if (!Visitor.success() || Visitor.failed())
286 return std::nullopt;
287 return Visitor.getKernelArgInfo();
288 }
289
290 inline std::optional<FunctionAnalysis> analyzeFunction(CallBase &CB,
291 int64_t StartOffset) {
292 FunctionAnalysis Result;
293 auto &F = *CB.getCalledFunction();
294 auto RetInstOpt = getRetInst(F);
295 if (!RetInstOpt)
296 return std::nullopt;
297 LambdaArgVisitor Visitor(RetInstOpt.value()->getReturnValue(),
298 MemoryAnalysisPtrUse, LambdaCB, StartOffset, DL);
299 while (!Visitor.empty() && !Visitor.success() && !Visitor.failed()) {
300 auto [V, AccessedFrom] = Visitor.back();
301 DEBUG(Logger::logs("proteus-pass")
302 << "Function analysis visiting " << *V << " with offset "
303 << Visitor.getOffset() << "\n");
304 Visitor.MemoryAnalysisPtrUse = AccessedFrom;
305 Visitor.popBack();
306 // Prevent loops/infinite recursion
307 if (Visitor.seen(V))
308 continue;
309 Visitor.markAsSeen(V);
310 // Analyze the instruction
311 if (auto *I = dyn_cast<Instruction>(V))
312 Visitor.visit(*I);
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;
317 Result.ArgIndex = A->getArgNo();
318 DEBUG(Logger::logs("proteus-pass")
319 << "Function analysis found termination case "
320 << *Result.PtrArgToCB << " with offset " << Visitor.getOffset()
321 << "\n");
322 return Result;
323 }
324 Visitor.visitArgument(*A);
325
326 } else
327 continue;
328 }
329
330 return std::nullopt;
331 }
332
333public:
334 LambdaArgVisitor(CallBase *LambdaCB, Module &M)
335 : LambdaCB(LambdaCB), DL(M.getDataLayout()), Offset(0) {
336 auto *ClosurePtr = LambdaCB->getArgOperand(0);
337 WorkList.push_back({ClosurePtr, LambdaCB});
338 }
339
340 void visitStoreInst(StoreInst &SI) {
341 WorkList.push_back({SI.getValueOperand(), &SI});
342 }
343
344 void visitCallBase(CallBase &CB) {
345 if (!CB.getCalledFunction() || CB.getCalledFunction()->isDeclaration()) {
346 DEBUG(Logger::logs("proteus-pass")
347 << "[Lambda arg analysis]: Cannot trace indirect or declaration "
348 "call "
349 << CB << "\n");
350 AnalysisFailed = true;
351 AnalysisSuccess = false;
352 return;
353 }
354 // Clone the visitor to determine if this function (a) returns a ptr
355 // and (b) which arg determines the value of that ptr, and at which offset
356 auto SubAnalysis = analyzeFunction(CB, Offset);
357 if (SubAnalysis) {
358 WorkList.push_back({SubAnalysis.value().PtrArgToCB, &CB});
359 Offset = SubAnalysis.value().Offset;
360 return;
361 }
362 DEBUG(Logger::logs("proteus-pass")
363 << "Function analysis of \n"
364 << *CB.getCalledFunction() << " failed \n");
365 AnalysisFailed = true;
366 AnalysisSuccess = false;
367 }
368
369 void visitLoadInst(LoadInst &LI) {
370 DEBUG(Logger::logs("proteus-pass") << "Load inst analysis \n")
371 // Loading a pointer from a spill slot does not change the offset within
372 // the pointee. Resolve the pointer-sized store at offset zero in the slot,
373 // then continue with the original pointee-relative Offset.
374 if (isPointerSpillLoad(LI)) {
376 SmallPtrSet<Value *, 8> Visited;
377 if (Value *StoredPointer =
378 getUniqueReachingPointer(DL, Stores, Visited)) {
379 WorkList.push_back({StoredPointer, &LI});
380 return;
381 }
382 if (hasAmbiguousReachingPointers(DL, Stores)) {
383 DEBUG(Logger::logs("proteus-pass")
384 << "[Lambda arg analysis]: Pointer spill load has ambiguous "
385 "reaching stores: "
386 << LI << "\n");
387 AnalysisFailed = true;
388 AnalysisSuccess = false;
389 return;
390 }
391 auto Res = getDominatingUse(DL, LI.getPointerOperand(), &LI, 0, LambdaCB);
392 if (!Res) {
393 AnalysisFailed = true;
394 AnalysisSuccess = false;
395 return;
396 }
397 WorkList.push_back({Res->DominatingWrite, &LI});
398 return;
399 }
400
401 WorkList.push_back({LI.getPointerOperand(), &LI});
402 }
403
404 void visitGetElementPtrInst(GetElementPtrInst &GEP) {
405 APInt StepOffset(DL.getIndexTypeSizeInBits(GEP.getType()), 0);
406 if (!GEP.accumulateConstantOffset(DL, StepOffset)) {
407 AnalysisFailed = true;
408 AnalysisSuccess = false;
409 return;
410 }
411 Offset += StepOffset.getSExtValue();
412 WorkList.push_back({GEP.getPointerOperand(), &GEP});
413 }
414
415 void visitExtractValueInst(ExtractValueInst &EVI) {
416 int64_t EVIOffset = getValueIndicesOffset(
417 DL, EVI.getAggregateOperand()->getType(), EVI.getIndices());
418 Offset += EVIOffset;
419 WorkList.push_back({EVI.getAggregateOperand(), &EVI});
420 }
421
422 // The analysis always encounters IVI chains in a backwards direction, meaning
423 // we always see the final IVI in a chain of writes. We assert this shape in
424 // our analysis, stopping at the location within the aggregate where we know
425 // our closure ptr lives
426 void visitInsertValueInst(InsertValueInst &IVI) {
427 auto *AggregateOperand = IVI.getAggregateOperand();
428 auto *Cur = &IVI;
429
430 while (Cur && AggregateOperand) {
431 int64_t CurOffset = getValueIndicesOffset(
432 DL, Cur->getAggregateOperand()->getType(), Cur->getIndices());
433 TypeSize InsertedSize =
434 DL.getTypeAllocSize(Cur->getInsertedValueOperand()->getType());
435 DEBUG(Logger::logs("proteus-pass")
436 << "Curr offset " << CurOffset << "\nOffset " << Offset << "\n");
437 if (!InsertedSize.isScalable() && Offset >= CurOffset &&
438 static_cast<uint64_t>(Offset - CurOffset) <
439 InsertedSize.getFixedValue()) {
440 // We are now following the value inserted into this aggregate field,
441 // so make the tracked offset relative to that value. Leaving the
442 // aggregate-field offset in place causes it to be counted again when
443 // the inserted value came from a GEP into another aggregate.
444 Offset -= CurOffset;
445 WorkList.push_back({Cur->getInsertedValueOperand(), &IVI});
446 return;
447 }
448 Cur = dyn_cast<InsertValueInst>(AggregateOperand);
449 if (Cur)
450 AggregateOperand = Cur->getAggregateOperand();
451 }
452 AnalysisFailed = true;
453 AnalysisSuccess = false;
454 }
455
456 // todo: these three methods need to be changed to find a dominating store
457 void visitAllocaInst(AllocaInst &Alloca) {
458 auto Res =
459 getDominatingUse(DL, &Alloca, MemoryAnalysisPtrUse, Offset, LambdaCB);
460 if (!Res)
461 return;
462
463 WorkList.push_back({Res->DominatingWrite, &Alloca});
464 // Res->Offset converts the current allocation-relative byte offset into
465 // the coordinate system of DominatingWrite. For a field store it removes
466 // the field displacement; for a memory transfer it translates destination
467 // displacement into the corresponding source displacement.
468 Offset -= Res->Offset;
469 }
470
471 void visitBitCastInst(BitCastInst &BC) {
472 auto Res =
473 getDominatingUse(DL, &BC, MemoryAnalysisPtrUse, Offset, LambdaCB);
474 if (!Res)
475 return;
476 WorkList.push_back({Res->DominatingWrite, &BC});
477 // Res->Offset converts the current cast-relative byte offset into the
478 // coordinate system of DominatingWrite.
479 Offset -= Res->Offset;
480 }
481
482 void visitAddrSpaceCastInst(AddrSpaceCastInst &ASC) {
483 WorkList.push_back({ASC.getPointerOperand(), &ASC});
484 auto Res =
485 getDominatingUse(DL, &ASC, MemoryAnalysisPtrUse, Offset, LambdaCB);
486 if (!Res)
487 return;
488
489 WorkList.push_back({Res->DominatingWrite, &ASC});
490 // Res->Offset converts the current cast-relative byte offset into the
491 // coordinate system of DominatingWrite.
492 Offset -= Res->Offset;
493 }
494
495 void visitIntToPtr(IntToPtrInst &ITP) {
496 auto *IntegerVal = ITP.getOperand(0);
497 auto *Ptr = dyn_cast<PtrToIntInst>(IntegerVal);
498 if (!Ptr) {
499 AnalysisSuccess = false;
500 AnalysisFailed = true;
501 return;
502 }
503 ChangedRC = convertTypeToRuntimeConstantType(Ptr->getType());
504
505 WorkList.push_back({Ptr->getPointerOperand(), &ITP});
506 }
507
508 void visitMemIntrinsic(MemIntrinsic &I) {
509 auto *MT = dyn_cast<MemTransferInst>(&I); // memcpy/memmove
510 if (!MT) {
511 // memset doesn't preserve any src->dst relationship we can use
512 AnalysisFailed = true;
513 return;
514 }
515
516 int64_t DstOff = 0, SrcOff = 0;
517 Value *DstBase =
518 GetPointerBaseWithConstantOffset(MT->getRawDest(), DstOff, DL);
519 Value *SrcBase =
520 GetPointerBaseWithConstantOffset(MT->getRawSource(), SrcOff, DL);
521 if (!DstBase || !SrcBase) {
522 AnalysisFailed = true;
523 return;
524 }
525
526 // Optional safety: only valid if the tracked byte lies within the copied
527 // region.
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;
532 return;
533 }
534 }
535
536 WorkList.push_back({SrcBase->stripPointerCasts(), &I});
537 Offset = Offset - DstOff + SrcOff;
538 }
539
540 void visitIntrinsicInst(IntrinsicInst &) {
541 AnalysisFailed = true;
542 return;
543 }
544
545 void visitTruncInst(TruncInst &TI) {
546 WorkList.push_back({TI.getOperand(0), &TI});
547 }
548
549 void visitArgument(Argument &A) {
550 Function *F = A.getParent();
551 DEBUG(Logger::logs("proteus-pass")
552 << "Visiting argument with parent function = \n"
553 << *F << "\n");
554 auto ArgNum = A.getArgNo();
555 // termination case: we have reached the parent calling kernel
556 // todo: we could just pass in the kernel pointer here and check equality
557
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"))) {
563 DEBUG(Logger::logs("proteus-pass")
564 << "Found termination case from function " << F->getName() << "\n");
565 DEBUG(Logger::logs("proteus-pass").flush());
566
567 AnalysisSuccess = true;
568 KernelArg = ArgNum;
569 KernelFunction = F;
570 return;
571 }
572
573 for (User *U : F->users()) {
574 auto *CB = dyn_cast<CallBase>(U);
575 if (!CB)
576 continue;
577 DEBUG(Logger::logs("proteus-pass")
578 << "Analysis crossed interprocedural boundary at "
579 << *CB->getArgOperand(ArgNum) << "\n");
580 WorkList.push_back({CB->getArgOperand(ArgNum), CB});
581 }
582 }
583
584 void visitInstruction(Instruction &I) {
585 DEBUG(Logger::logs("proteus-pass")
586 << "[Lambda arg analysis]: Unhandled instruction "
587 << I.getOpcodeName() << ": " << I << "\n");
588 AnalysisFailed = true;
589 return;
590 }
591
592 void visitPHINode(PHINode &P) {
593 if (P.getNumIncomingValues() == 0) {
594 AnalysisFailed = true;
595 return;
596 }
597
598 auto FirstAnalysis =
599 cloneAndAnalyze(P.getIncomingValue(0), MemoryAnalysisPtrUse, Offset);
600 if (!FirstAnalysis) {
601 AnalysisFailed = true;
602 AnalysisSuccess = false;
603 return;
604 }
605 auto BaseSlot = FirstAnalysis->KernelArgIndex;
606 auto BaseOffset = FirstAnalysis->Offset;
607
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;
615 return;
616 }
617 }
618 KernelFunction = FirstAnalysis->KernelFunction;
619 Offset = BaseOffset;
620 KernelArg = BaseSlot;
621 ChangedRC = FirstAnalysis->ChangedRCLayout;
622 AnalysisSuccess = true;
623 AnalysisFailed = false;
624 }
625
626 void visitSelectInst(SelectInst &S) {
627 // A select can merge call results or loads as well as simple GEPs. Base
628 // pointer equality rejects semantically identical paths in those shapes,
629 // so analyze both arms exactly as we do for a PHI and retain the result
630 // only when they resolve to one kernel argument and byte offset.
631 auto TrueAnalysis =
632 cloneAndAnalyze(S.getTrueValue(), MemoryAnalysisPtrUse, Offset);
633 auto FalseAnalysis =
634 cloneAndAnalyze(S.getFalseValue(), MemoryAnalysisPtrUse, Offset);
635 if (!TrueAnalysis || !FalseAnalysis ||
636 TrueAnalysis->KernelFunction != FalseAnalysis->KernelFunction ||
637 TrueAnalysis->KernelArgIndex != FalseAnalysis->KernelArgIndex ||
638 TrueAnalysis->Offset != FalseAnalysis->Offset ||
639 TrueAnalysis->ChangedRCLayout != FalseAnalysis->ChangedRCLayout) {
640 DEBUG(Logger::logs("proteus-pass")
641 << "[Lambda arg analysis]: Select arms do not resolve to the "
642 "same kernel argument and offset: "
643 << S << "\n");
644 AnalysisFailed = true;
645 AnalysisSuccess = false;
646 return;
647 }
648
649 KernelFunction = TrueAnalysis->KernelFunction;
650 KernelArg = TrueAnalysis->KernelArgIndex;
651 Offset = TrueAnalysis->Offset;
652 ChangedRC = TrueAnalysis->ChangedRCLayout;
653 AnalysisSuccess = true;
654 AnalysisFailed = false;
655 }
656};
657
659 llvm::Module &M,
660 DenseMap<CallBase *, LambdaKernelArgAnalysis> &CallBaseToArgOffset,
661 const SmallVector<CallBase *> &CBToAnalyze) {
662 DEBUG(Logger::logs("proteus-pass") << "Beginning analysis " << "\n");
663 for (auto *FunctorCB : CBToAnalyze) {
664 LambdaArgVisitor Visitor(FunctorCB, M);
665 while (!Visitor.empty() && !Visitor.success() && !Visitor.failed()) {
666 auto [V, LastSeen] = Visitor.back();
667 Visitor.MemoryAnalysisPtrUse = LastSeen;
668 Visitor.popBack();
669 // Prevent loops/infinite recursion
670 if (Visitor.seen(V))
671 continue;
672 Visitor.markAsSeen(V);
673 DEBUG(Logger::logs("proteus-pass") << "Visiting value " << *V << "\n");
674 // Analyze the instruction
675 if (auto *I = dyn_cast<Instruction>(V))
676 Visitor.visit(*I);
677 else if (auto *A = dyn_cast<Argument>(V))
678 Visitor.visitArgument(*A);
679 else
680 continue;
681 }
682 if (!Visitor.success() || Visitor.failed()) {
683 DEBUG(Logger::logs("proteus-pass")
684 << "[WARNING]: Kernel arg analysis failed for functor beginning at "
685 << *FunctorCB << "\n");
686 return false;
687 }
689 if (!Info.KernelFunction)
690 return false;
691 CallBaseToArgOffset[FunctorCB] = Info;
692 }
693 return true;
694}
695} // namespace proteus
696
697#endif
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