78 collectWrapperCallsites(Function &FunctorOperatorFunction,
79 SmallVectorImpl<CallBase *> &CBToAnalyze) {
80 for (
auto *U : FunctorOperatorFunction.users()) {
81 if (
auto *CB = dyn_cast<CallBase>(U))
82 CBToAnalyze.push_back(CB);
85 llvm::sort(CBToAnalyze, [](
const CallBase *L,
const CallBase *R) {
86 const Function *LF = L->getFunction();
87 const Function *RF = R->getFunction();
89 return LF->getName() < RF->getName();
90 return getInstructionOrder(L) < getInstructionOrder(R);
94 static const RuntimeConstant *findArgByOffset(
const JitVariantMap &RCMap,
96 for (
auto &[_, Arg] : RCMap) {
105 auto It = RCMap.find(
Pos);
106 if (It == RCMap.end())
111 static auto traceOut(
int Slot, Constant *C) {
113 raw_svector_ostream OS(S);
114 OS <<
"[LambdaSpec] Replacing slot " << Slot <<
" with " << *C <<
"\n";
119 static void handleLoad(Module &M, LoadInst *LI,
const JitVariantMap &RCVec) {
120 auto *Arg = findArgByPos(RCVec, 0);
124 Constant *C =
getConstant(M.getContext(), LI->getType(), *Arg);
125 LI->replaceAllUsesWith(C);
131 static void handleGEP(Module &M, GetElementPtrInst *GEP,
132 const JitVariantMap &RCVec) {
133 auto *GEPSlot = GEP->getOperand(GEP->getNumOperands() - 1);
134 ConstantInt *CI = dyn_cast<ConstantInt>(GEPSlot);
135 int Slot = CI->getZExtValue();
136 Type *SrcTy = GEP->getSourceElementType();
138 auto *Arg = SrcTy->isStructTy() ? findArgByPos(RCVec, Slot)
139 : findArgByOffset(RCVec, Slot);
143 for (
auto *GEPUser : GEP->users()) {
144 auto *LI = dyn_cast<LoadInst>(GEPUser);
147 Type *LoadType = LI->getType();
148 Constant *C =
getConstant(M.getContext(), LoadType, *Arg);
149 LI->replaceAllUsesWith(C);
156 static Function *findLambdaOperatorForFunctor(Module &M, uint64_t FunctorID) {
157 SmallVector<std::pair<Function *, uint64_t>> LambdaOperators;
160 for (
auto [Lambda, ID] : LambdaOperators)
167 findFunctorFunctorOperatorFunctionOperatorFromID(Module &M,
168 uint64_t FunctorID) {
169 SmallVector<std::pair<Function *, uint64_t>> FunctorOperators;
171 for (
auto [Lambda, ID] : FunctorOperators)
177 static CallBase *findDirectCallTo(Function &Caller, Function *Callee) {
178 for (BasicBlock &BB : Caller) {
179 for (Instruction &I : BB) {
180 auto *CB = dyn_cast<CallBase>(&I);
183 Function *Called = CB->getCalledFunction();
186 if (Called == Callee)
193 static uint64_t getInstructionOrder(
const Instruction *I) {
195 for (
const BasicBlock &BB : *I->getFunction()) {
196 for (
const Instruction &Cur : BB) {
205 static Function *cloneForVariant(Function &F, uint64_t FunctorID,
206 size_t VariantIndex) {
207 ValueToValueMapTy VMap;
209 auto *NewFunc = Function::Create(
210 F.getFunctionType(), F.getLinkage(), F.getAddressSpace(),
211 F.getName() +
".proteus.variant." + Twine(FunctorID) +
"." +
215 NewFunc->copyAttributesFrom(&F);
216 NewFunc->setCallingConv(F.getCallingConv());
218 auto *NewArgIt = NewFunc->arg_begin();
219 for (
auto &OldArg : F.args()) {
220 NewArgIt->setName(OldArg.getName());
221 VMap[&OldArg] = &(*NewArgIt++);
224 SmallVector<ReturnInst *, 8> Returns;
225 CloneFunctionInto(NewFunc, &F, VMap,
226 CloneFunctionChangeType::LocalChangesOnly, Returns);
230 NewFunc->removeFnAttr(Attribute::NoInline);
235 static StructType *inferStructTypeFromGEPValue(Value *Ptr) {
236 auto *GEP = dyn_cast<GetElementPtrInst>(Ptr);
240 Type *CurTy = GEP->getSourceElementType();
241 bool IsFirstIndex =
true;
242 for (Value *Idx : GEP->indices()) {
246 IsFirstIndex =
false;
250 auto *CI = dyn_cast<ConstantInt>(Idx);
254 if (
auto *ST = dyn_cast<StructType>(CurTy)) {
255 uint64_t ElemIdx = CI->getZExtValue();
256 if (ElemIdx >= ST->getNumElements())
258 CurTy = ST->getElementType(ElemIdx);
262 if (
auto *AT = dyn_cast<ArrayType>(CurTy)) {
263 CurTy = AT->getElementType();
267 if (
auto *VT = dyn_cast<VectorType>(CurTy)) {
268 CurTy = VT->getElementType();
275 return dyn_cast<StructType>(CurTy);
278 static void specializeCallOperator(Module &M, Function &CallOp,
279 const JitVariantMap &RCMap) {
280 if (CallOp.arg_empty())
283 auto *LambdaClass = CallOp.getArg(0);
285 <<
"[LambdaSpec] Function: " << CallOp.getName()
286 <<
" RCVec size " << RCMap.size() <<
"\n");
288 for (User *U : LambdaClass->users()) {
289 if (
auto *LI = dyn_cast<LoadInst>(U))
290 handleLoad(M, LI, RCMap);
291 else if (
auto *GEP = dyn_cast<GetElementPtrInst>(U))
292 handleGEP(M, GEP, RCMap);
302 if (LambdaStorageBasePointerOffset < 0)
305 auto *Base =
static_cast<const unsigned char *
>(KernelArgs[
KernelArgIndex]);
309 const unsigned char *StorageBase =
nullptr;
311 StorageBase = Base +
static_cast<size_t>(LambdaStorageBasePointerOffset);
313 const unsigned char *IndirectStoragePtr =
314 Base +
static_cast<size_t>(LambdaStorageBasePointerOffset);
315 uintptr_t StorageAddr = 0;
319 std::memcpy(&Bits, IndirectStoragePtr,
sizeof(Bits));
320 StorageAddr =
static_cast<uintptr_t
>(Bits);
324 void *PtrVal =
nullptr;
325 std::memcpy(&PtrVal, IndirectStoragePtr,
sizeof(PtrVal));
326 StorageAddr =
reinterpret_cast<uintptr_t
>(PtrVal);
332 StorageBase =
reinterpret_cast<const unsigned char *
>(StorageAddr);
337 const unsigned char *Ptr = StorageBase +
static_cast<size_t>(RCOffset);
343 std::memcpy(&RC.Value.BoolVal, Ptr,
sizeof(RC.Value.BoolVal));
346 std::memcpy(&RC.Value.Int8Val, Ptr,
sizeof(RC.Value.Int8Val));
349 std::memcpy(&RC.Value.Int32Val, Ptr,
sizeof(RC.Value.Int32Val));
352 std::memcpy(&RC.Value.Int64Val, Ptr,
sizeof(RC.Value.Int64Val));
355 std::memcpy(&RC.Value.FloatVal, Ptr,
sizeof(RC.Value.FloatVal));
358 std::memcpy(&RC.Value.DoubleVal, Ptr,
sizeof(RC.Value.DoubleVal));
361 std::memcpy(&RC.Value.LongDoubleVal, Ptr,
sizeof(RC.Value.LongDoubleVal));
364 std::memcpy(&RC.Value.PtrVal, Ptr,
sizeof(RC.Value.PtrVal));
374 readRuntimeConstantFromStorage(
const void *StorageBasePtr,
382 const unsigned char *Base =
383 static_cast<const unsigned char *
>(StorageBasePtr);
384 const unsigned char *Ptr = Base +
static_cast<size_t>(RCOffset);
390 std::memcpy(&RC.Value.BoolVal, Ptr,
sizeof(RC.Value.BoolVal));
393 std::memcpy(&RC.Value.Int8Val, Ptr,
sizeof(RC.Value.Int8Val));
396 std::memcpy(&RC.Value.Int32Val, Ptr,
sizeof(RC.Value.Int32Val));
399 std::memcpy(&RC.Value.Int64Val, Ptr,
sizeof(RC.Value.Int64Val));
402 std::memcpy(&RC.Value.FloatVal, Ptr,
sizeof(RC.Value.FloatVal));
405 std::memcpy(&RC.Value.DoubleVal, Ptr,
sizeof(RC.Value.DoubleVal));
408 std::memcpy(&RC.Value.LongDoubleVal, Ptr,
sizeof(RC.Value.LongDoubleVal));
411 std::memcpy(&RC.Value.PtrVal, Ptr,
sizeof(RC.Value.PtrVal));
415 "Unsupported RuntimeConstantType in host-storage reader");
425 const JitVariantMap &VariantSchema) {
427 RuntimeConstants.reserve(VariantSchema.size());
429 for (
auto [Slot, RC] : VariantSchema) {
431 RuntimeConstants.push_back(readRuntimeConstantFromKernelArgs(
433 Location.
StorageType, RC.Type, RC.Pos, RC.Offset));
436 llvm::sort(RuntimeConstants,
440 return RuntimeConstants;
444 void *
const *KernelArgs,
445 ArrayRef<LambdaKernelArgLocation> KernelArgLocations,
446 ArrayRef<LambdaRegistry::JitVariantMap> VariantSchemas) {
448 if (VariantSchemas.empty())
449 return RuntimeVariants;
451 RuntimeVariants.reserve(KernelArgLocations.size());
452 const JitVariantMap &VariantSchema = VariantSchemas.front();
454 for (
const auto &Location : KernelArgLocations) {
455 JitVariantMap RuntimeVariant;
456 for (
auto [Slot, RC] : VariantSchema) {
457 RuntimeVariant[Slot] = readRuntimeConstantFromKernelArgs(
458 KernelArgs, Location.KernelArgIndex, Location.Offset,
459 Location.StorageType, RC.Type, RC.Pos, RC.Offset);
461 RuntimeVariants.push_back(std::move(RuntimeVariant));
464 return RuntimeVariants;
468 const void *StorageBasePtr,
469 ArrayRef<LambdaRegistry::JitVariantMap> VariantSchemas) {
470 JitVariantMap RuntimeVariant;
471 if (VariantSchemas.empty())
472 return RuntimeVariant;
474 const JitVariantMap &VariantSchema = VariantSchemas.front();
475 for (
auto [Slot, RC] : VariantSchema) {
476 RuntimeVariant[Slot] = readRuntimeConstantFromStorage(
477 StorageBasePtr, RC.Type, RC.Pos, RC.Offset);
480 return RuntimeVariant;
484 const JitVariantMap &RuntimeVariant) {
485 specializeCallOperator(M, Lambda, RuntimeVariant);
489 Module &M, uint64_t FunctorID,
491 Function *FunctorOperatorFunction =
492 findFunctorFunctorOperatorFunctionOperatorFromID(M, FunctorID);
493 Function *LambdaOperatorMethod = findLambdaOperatorForFunctor(M, FunctorID);
494 if (!FunctorOperatorFunction) {
496 Logger::trace(
"[LambdaSpec] Skipping lambda specialization: "
497 "no wrapper found for ID " +
498 std::to_string(FunctorID));
502 SmallVector<CallBase *> CBToAnalyze;
503 collectWrapperCallsites(*FunctorOperatorFunction, CBToAnalyze);
505 size_t VariantIndex = 0;
506 for (CallBase *CB : CBToAnalyze) {
512 if (It == CallsiteRuntimeConstants.end())
515 JitVariantMap RuntimeVariant;
516 for (
const auto &RC : It->second) {
517 RuntimeVariant[RC.Pos] = RC;
521 Function *WrapperClone =
522 cloneForVariant(*FunctorOperatorFunction, FunctorID, VariantIndex);
523 specializeCallOperator(M, *WrapperClone, RuntimeVariant);
525 if (!LambdaOperatorMethod) {
529 CB->setCalledFunction(WrapperClone);
533 Function *LambdaClone =
534 cloneForVariant(*LambdaOperatorMethod, FunctorID, VariantIndex);
535 specializeCallOperator(M, *LambdaClone, RuntimeVariant);
537 if (CallBase *WrapperToLambdaCall =
538 findDirectCallTo(*WrapperClone, LambdaOperatorMethod)) {
539 WrapperToLambdaCall->setCalledFunction(LambdaClone);
544 CB->setCalledFunction(WrapperClone);
553 if (FunctorOperatorFunction->users().empty())
554 FunctorOperatorFunction->eraseFromParent();
556 if (LambdaOperatorMethod && LambdaOperatorMethod->users().empty())
557 LambdaOperatorMethod->eraseFromParent();
560#ifdef PROTEUS_TRANSFORM_CONSERVATIVE
561 static void transformConservative(Module &M, uint64_t FunctorID,
562 ArrayRef<JitVariantMap> Variants) {
563 if (Variants.empty())
565 Function *FunctorOperatorFunction =
566 findFunctorFunctorOperatorFunctionOperatorFromID(M, FunctorID);
567 Function *LambdaOperatorMethod = findLambdaOperatorForFunctor(M, FunctorID);
568 if (!FunctorOperatorFunction) {
570 Logger::trace(
"[LambdaSpec] Internal lambda specialization error:"
571 "no wrapper found for ID " +
572 std::to_string(FunctorID));
576 if (!LambdaOperatorMethod) {
580 if (Variants.size() == 1)
581 specializeCallOperator(M, *FunctorOperatorFunction, Variants.front());
585 if (Variants.size() == 1) {
586 specializeCallOperator(M, *LambdaOperatorMethod, Variants.front());
591 findDirectCallTo(*FunctorOperatorFunction, LambdaOperatorMethod);
595 CallingConv::ID CallConv = OrigCall->getCallingConv();
596 AttributeList CallAttrs = OrigCall->getAttributes();
598 SmallVector<Value *, 8> CallArgs;
599 CallArgs.reserve(OrigCall->arg_size());
600 for (Use &U : OrigCall->args())
601 CallArgs.push_back(U.get());
603 Value *LambdaObjPtr = OrigCall->getArgOperand(0);
604 StructType *LambdaTy = inferStructTypeFromGEPValue(LambdaObjPtr);
607 BasicBlock *EntryBB = OrigCall->getParent();
609 EntryBB->splitBasicBlock(OrigCall,
"proteus.after_call");
612 EntryBB->getTerminator()->eraseFromParent();
615 PHINode *ResultPhi =
nullptr;
616 Type *RetTy = OrigCall->getType();
617 if (!RetTy->isVoidTy()) {
619 PHINode::Create(RetTy, Variants.size() + 1,
620 "proteus.lambda_dispatch.result", &*TailBB->begin());
621 OrigCall->replaceAllUsesWith(ResultPhi);
624 Instruction *AfterCallIP = OrigCall->getNextNode();
625 OrigCall->eraseFromParent();
627 LLVMContext &Ctx = M.getContext();
629 BasicBlock *FallbackBB = BasicBlock::Create(
630 Ctx,
"proteus.lambda_dispatch.fallback", FunctorOperatorFunction);
633 BasicBlock *CurCheckBB = EntryBB;
634 for (
size_t I = 0, E = Variants.size(); I < E; ++I) {
635 BasicBlock *VariantBB =
636 BasicBlock::Create(Ctx,
"proteus.lambda_dispatch.variant." + Twine(I),
637 FunctorOperatorFunction);
638 BasicBlock *NextCheckBB =
641 : BasicBlock::Create(
642 Ctx,
"proteus.lambda_dispatch.check." + Twine(I + 1),
643 FunctorOperatorFunction);
646 IRBuilder<> B(CurCheckBB);
647 Value *Match = buildVariantMatch(B, M.getDataLayout(), LambdaObjPtr,
648 LambdaTy, Variants[I]);
652 BranchInst::Create(FallbackBB, CurCheckBB);
654 BranchInst::Create(VariantBB, NextCheckBB, Match, CurCheckBB);
660 IRBuilder<> B(VariantBB);
661 Function *Clone = cloneForVariant(*LambdaOperatorMethod, FunctorID, I);
662 specializeCallOperator(M, *Clone, Variants[I]);
664 CallInst *C = B.CreateCall(Clone, CallArgs);
665 C->setCallingConv(CallConv);
666 C->setAttributes(CallAttrs);
669 ResultPhi->addIncoming(C, VariantBB);
674 CurCheckBB = NextCheckBB;
679 IRBuilder<> B(FallbackBB);
680 CallInst *C = B.CreateCall(LambdaOperatorMethod, CallArgs);
681 C->setCallingConv(CallConv);
682 C->setAttributes(CallAttrs);
685 ResultPhi->addIncoming(C, FallbackBB);
695 static Value *buildSingleCompare(IRBuilder<> &B,
const DataLayout &DL,
696 Value *LambdaObjPtr, StructType *LambdaTy,
698 Value *FieldPtr =
nullptr;
699 if (LambdaTy && RC.
Pos >= 0 &&
700 static_cast<uint64_t
>(RC.
Pos) < LambdaTy->getNumElements()) {
701 FieldPtr = B.CreateStructGEP(LambdaTy, LambdaObjPtr, RC.
Pos);
705 FieldPtr = B.CreateGEP(B.getInt8Ty(), LambdaObjPtr,
706 B.getInt64(
static_cast<uint64_t
>(RC.
Offset)));
711 Value *Loaded = B.CreateAlignedLoad(B.getInt8Ty(), FieldPtr, Align(1));
713 return B.CreateICmpEQ(Loaded, C);
716 Value *Loaded = B.CreateAlignedLoad(B.getInt8Ty(), FieldPtr, Align(1));
717 Value *C = B.getInt8(
static_cast<uint8_t
>(RC.
Value.
Int8Val));
718 return B.CreateICmpEQ(Loaded, C);
721 Value *Loaded = B.CreateAlignedLoad(B.getInt32Ty(), FieldPtr, Align(1));
722 Value *C = B.getInt32(
static_cast<uint32_t
>(RC.
Value.
Int32Val));
723 return B.CreateICmpEQ(Loaded, C);
726 Value *Loaded = B.CreateAlignedLoad(B.getInt64Ty(), FieldPtr, Align(1));
727 Value *C = B.getInt64(
static_cast<uint64_t
>(RC.
Value.
Int64Val));
728 return B.CreateICmpEQ(Loaded, C);
731 Value *LoadedF = B.CreateAlignedLoad(B.getFloatTy(), FieldPtr, Align(1));
732 Value *LoadedBits = B.CreateBitCast(LoadedF, B.getInt32Ty());
735 Value *C = B.getInt32(Bits);
736 return B.CreateICmpEQ(LoadedBits, C);
739 Value *LoadedF = B.CreateAlignedLoad(B.getDoubleTy(), FieldPtr, Align(1));
740 Value *LoadedBits = B.CreateBitCast(LoadedF, B.getInt64Ty());
743 Value *C = B.getInt64(Bits);
744 return B.CreateICmpEQ(LoadedBits, C);
747 Type *IntPtrTy = DL.getIntPtrType(B.getContext());
748 Value *Loaded = B.CreateAlignedLoad(IntPtrTy, FieldPtr, Align(1));
750 static_cast<uint64_t
>(
reinterpret_cast<uintptr_t
>(RC.
Value.
PtrVal));
751 Value *C = ConstantInt::get(IntPtrTy, Bits);
752 return B.CreateICmpEQ(Loaded, C);
759 static Value *buildVariantMatch(IRBuilder<> &B,
const DataLayout &DL,
760 Value *LambdaObjPtr, StructType *LambdaTy,
761 const JitVariantMap &Variant) {
763 return B.getInt1(
true);
765 SmallVector<int32_t, 16> Keys;
766 Keys.reserve(Variant.size());
767 for (
const auto &KV : Variant)
768 Keys.push_back(KV.first);
770 llvm::sort(Keys.begin(), Keys.end());
772 Value *Match = B.getInt1(
true);
773 for (int32_t K : Keys) {
774 auto It = Variant.find(K);
775 if (It == Variant.end())
778 buildSingleCompare(B, DL, LambdaObjPtr, LambdaTy, It->second);
781 Match = B.CreateAnd(Match, Cmp);