Proteus
Programmable JIT compilation and optimization for C/C++ using LLVM
Loading...
Searching...
No Matches
TransformLambdaSpecialization.h
Go to the documentation of this file.
1//===-- TransformLambdaSpecialization.h -- Specialize arguments --===//
2//
3// Part of the Proteus Project, under the Apache License v2.0 with LLVM
4// Exceptions. See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9//===----------------------------------------------------------------------===//
10
11#ifndef PROTEUS_TRANSFORM_LAMBDA_SPECIALIZATION_H
12#define PROTEUS_TRANSFORM_LAMBDA_SPECIALIZATION_H
13
15#include "proteus/Error.h"
17#include "proteus/impl/Debug.h"
20#include "proteus/impl/Utils.h"
21
22#include <llvm/ADT/STLExtras.h>
23#include <llvm/Demangle/Demangle.h>
24#include <llvm/IR/Attributes.h>
25#include <llvm/IR/IRBuilder.h>
26#include <llvm/IR/Instructions.h>
27#include <llvm/IR/LLVMContext.h>
28#include <llvm/IR/Module.h>
29#include <llvm/IR/Type.h>
30#include <llvm/Support/Casting.h>
31#include <llvm/Support/Debug.h>
32#include <llvm/Transforms/Utils/Cloning.h>
33
34#include <cstdint>
35#include <cstring>
36#include <memory>
37
38namespace proteus {
39
40using namespace llvm;
41
42inline Constant *getConstant(LLVMContext &Ctx, Type *ArgType,
43 const RuntimeConstant &RC) {
44 switch (RC.Type) {
46 return ConstantInt::get(ArgType, RC.Value.BoolVal);
48 return ConstantInt::get(ArgType, RC.Value.Int8Val);
50 return ConstantInt::get(ArgType, RC.Value.Int32Val);
52 return ConstantInt::get(ArgType, RC.Value.Int64Val);
54 return ConstantFP::get(ArgType, RC.Value.FloatVal);
56 return ConstantFP::get(ArgType, RC.Value.DoubleVal);
58 return ConstantFP::get(ArgType, RC.Value.LongDoubleVal);
60 auto *IntC = ConstantInt::get(Type::getInt64Ty(Ctx), RC.Value.Int64Val);
61 return ConstantExpr::getIntToPtr(IntC, ArgType);
62 }
63 default:
64 std::string TypeString;
65 raw_string_ostream TypeOstream(TypeString);
66 ArgType->print(TypeOstream);
67 reportFatalError("JIT Incompatible type in runtime constant: " +
68 TypeOstream.str());
69 }
70}
71
73private:
74 using JitVariantMap = LambdaRegistry::JitVariantMap;
75 using JitVariantVec = LambdaRegistry::JitVariantVec;
76
77 static void
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);
83 }
84
85 llvm::sort(CBToAnalyze, [](const CallBase *L, const CallBase *R) {
86 const Function *LF = L->getFunction();
87 const Function *RF = R->getFunction();
88 if (LF != RF)
89 return LF->getName() < RF->getName();
90 return getInstructionOrder(L) < getInstructionOrder(R);
91 });
92 }
93
94 static const RuntimeConstant *findArgByOffset(const JitVariantMap &RCMap,
95 int32_t Offset) {
96 for (auto &[_, Arg] : RCMap) {
97 if (Arg.Offset == Offset)
98 return &Arg;
99 }
100 return nullptr;
101 };
102
103 static const RuntimeConstant *findArgByPos(const JitVariantMap &RCMap,
104 int32_t Pos) {
105 auto It = RCMap.find(Pos);
106 if (It == RCMap.end())
107 return nullptr;
108 return &It->second;
109 };
110
111 static auto traceOut(int Slot, Constant *C) {
112 SmallString<128> S;
113 raw_svector_ostream OS(S);
114 OS << "[LambdaSpec] Replacing slot " << Slot << " with " << *C << "\n";
115
116 return S;
117 };
118
119 static void handleLoad(Module &M, LoadInst *LI, const JitVariantMap &RCVec) {
120 auto *Arg = findArgByPos(RCVec, 0);
121 if (!Arg)
122 return;
123
124 Constant *C = getConstant(M.getContext(), LI->getType(), *Arg);
125 LI->replaceAllUsesWith(C);
126 PROTEUS_DBG(Logger::logs("proteus") << traceOut(Arg->Pos, C));
127 if (Config::get().traceSpecializations())
128 Logger::trace(traceOut(Arg->Pos, C));
129 }
130
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();
137
138 auto *Arg = SrcTy->isStructTy() ? findArgByPos(RCVec, Slot)
139 : findArgByOffset(RCVec, Slot);
140 if (!Arg)
141 return;
142
143 for (auto *GEPUser : GEP->users()) {
144 auto *LI = dyn_cast<LoadInst>(GEPUser);
145 if (!LI)
146 reportFatalError("Expected load instruction");
147 Type *LoadType = LI->getType();
148 Constant *C = getConstant(M.getContext(), LoadType, *Arg);
149 LI->replaceAllUsesWith(C);
150 PROTEUS_DBG(Logger::logs("proteus") << traceOut(Arg->Pos, C));
151 if (Config::get().traceSpecializations())
152 Logger::trace(traceOut(Arg->Pos, C));
153 }
154 }
155
156 static Function *findLambdaOperatorForFunctor(Module &M, uint64_t FunctorID) {
157 SmallVector<std::pair<Function *, uint64_t>> LambdaOperators;
158 findFunctionsWithU64Metadata(M, "proteus.registered_lambda",
159 LambdaOperators);
160 for (auto [Lambda, ID] : LambdaOperators)
161 if (ID == FunctorID)
162 return Lambda;
163 return nullptr;
164 }
165
166 static Function *
167 findFunctorFunctorOperatorFunctionOperatorFromID(Module &M,
168 uint64_t FunctorID) {
169 SmallVector<std::pair<Function *, uint64_t>> FunctorOperators;
170 findFunctionsWithU64Metadata(M, "proteus.wrapper_call", FunctorOperators);
171 for (auto [Lambda, ID] : FunctorOperators)
172 if (ID == FunctorID)
173 return Lambda;
174 return nullptr;
175 }
176
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);
181 if (!CB)
182 continue;
183 Function *Called = CB->getCalledFunction();
184 if (!Called)
185 continue;
186 if (Called == Callee)
187 return CB;
188 }
189 }
190 return nullptr;
191 }
192
193 static uint64_t getInstructionOrder(const Instruction *I) {
194 uint64_t Order = 0;
195 for (const BasicBlock &BB : *I->getFunction()) {
196 for (const Instruction &Cur : BB) {
197 if (&Cur == I)
198 return Order;
199 ++Order;
200 }
201 }
202 reportFatalError("Instruction not found in parent function");
203 }
204
205 static Function *cloneForVariant(Function &F, uint64_t FunctorID,
206 size_t VariantIndex) {
207 ValueToValueMapTy VMap;
208
209 auto *NewFunc = Function::Create(
210 F.getFunctionType(), F.getLinkage(), F.getAddressSpace(),
211 F.getName() + ".proteus.variant." + Twine(FunctorID) + "." +
212 Twine(VariantIndex),
213 F.getParent());
214
215 NewFunc->copyAttributesFrom(&F);
216 NewFunc->setCallingConv(F.getCallingConv());
217
218 auto *NewArgIt = NewFunc->arg_begin();
219 for (auto &OldArg : F.args()) {
220 NewArgIt->setName(OldArg.getName());
221 VMap[&OldArg] = &(*NewArgIt++);
222 }
223
224 SmallVector<ReturnInst *, 8> Returns;
225 CloneFunctionInto(NewFunc, &F, VMap,
226 CloneFunctionChangeType::LocalChangesOnly, Returns);
227
228 // The original lambda operator is marked noinline to preserve a distinct
229 // call target. Allow the specialized clones to be inlined by later O3.
230 NewFunc->removeFnAttr(Attribute::NoInline);
231
232 return NewFunc;
233 }
234
235 static StructType *inferStructTypeFromGEPValue(Value *Ptr) {
236 auto *GEP = dyn_cast<GetElementPtrInst>(Ptr);
237 if (!GEP)
238 return nullptr;
239
240 Type *CurTy = GEP->getSourceElementType();
241 bool IsFirstIndex = true;
242 for (Value *Idx : GEP->indices()) {
243 // The first GEP index selects an element of the "source element" type
244 // as an array; it does not descend into a struct field.
245 if (IsFirstIndex) {
246 IsFirstIndex = false;
247 continue;
248 }
249
250 auto *CI = dyn_cast<ConstantInt>(Idx);
251 if (!CI)
252 return nullptr;
253
254 if (auto *ST = dyn_cast<StructType>(CurTy)) {
255 uint64_t ElemIdx = CI->getZExtValue();
256 if (ElemIdx >= ST->getNumElements())
257 return nullptr;
258 CurTy = ST->getElementType(ElemIdx);
259 continue;
260 }
261
262 if (auto *AT = dyn_cast<ArrayType>(CurTy)) {
263 CurTy = AT->getElementType();
264 continue;
265 }
266
267 if (auto *VT = dyn_cast<VectorType>(CurTy)) {
268 CurTy = VT->getElementType();
269 continue;
270 }
271
272 return nullptr;
273 }
274
275 return dyn_cast<StructType>(CurTy);
276 }
277
278 static void specializeCallOperator(Module &M, Function &CallOp,
279 const JitVariantMap &RCMap) {
280 if (CallOp.arg_empty())
281 return;
282
283 auto *LambdaClass = CallOp.getArg(0);
284 PROTEUS_DBG(Logger::logs("proteus")
285 << "[LambdaSpec] Function: " << CallOp.getName()
286 << " RCVec size " << RCMap.size() << "\n");
287
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);
293 }
294 }
295
296 static inline RuntimeConstant readRuntimeConstantFromKernelArgs(
297 void *const *KernelArgs, uint32_t KernelArgIndex,
298 int64_t LambdaStorageBasePointerOffset, RuntimeConstantType StorageType,
299 RuntimeConstantType Type, int32_t Pos, int32_t RCOffset) {
300 if (!KernelArgs)
301 reportFatalError("KernelArgs is null");
302 if (LambdaStorageBasePointerOffset < 0)
303 reportFatalError("Negative KernelByteOffset");
304
305 auto *Base = static_cast<const unsigned char *>(KernelArgs[KernelArgIndex]);
306 if (!Base)
307 reportFatalError("KernelArgs[KernelArgIndex] is null");
308
309 const unsigned char *StorageBase = nullptr;
311 StorageBase = Base + static_cast<size_t>(LambdaStorageBasePointerOffset);
312 } else {
313 const unsigned char *IndirectStoragePtr =
314 Base + static_cast<size_t>(LambdaStorageBasePointerOffset);
315 uintptr_t StorageAddr = 0;
316 switch (StorageType) {
318 uint64_t Bits = 0;
319 std::memcpy(&Bits, IndirectStoragePtr, sizeof(Bits));
320 StorageAddr = static_cast<uintptr_t>(Bits);
321 break;
322 }
324 void *PtrVal = nullptr;
325 std::memcpy(&PtrVal, IndirectStoragePtr, sizeof(PtrVal));
326 StorageAddr = reinterpret_cast<uintptr_t>(PtrVal);
327 break;
328 }
329 default:
330 reportFatalError("Unsupported kernel lambda storage type");
331 }
332 StorageBase = reinterpret_cast<const unsigned char *>(StorageAddr);
333 if (!StorageBase)
334 reportFatalError("Null indirect kernel lambda storage pointer");
335 }
336
337 const unsigned char *Ptr = StorageBase + static_cast<size_t>(RCOffset);
338
339 RuntimeConstant RC{Type, Pos, RCOffset};
340
341 switch (Type) {
343 std::memcpy(&RC.Value.BoolVal, Ptr, sizeof(RC.Value.BoolVal));
344 break;
346 std::memcpy(&RC.Value.Int8Val, Ptr, sizeof(RC.Value.Int8Val));
347 break;
349 std::memcpy(&RC.Value.Int32Val, Ptr, sizeof(RC.Value.Int32Val));
350 break;
352 std::memcpy(&RC.Value.Int64Val, Ptr, sizeof(RC.Value.Int64Val));
353 break;
355 std::memcpy(&RC.Value.FloatVal, Ptr, sizeof(RC.Value.FloatVal));
356 break;
358 std::memcpy(&RC.Value.DoubleVal, Ptr, sizeof(RC.Value.DoubleVal));
359 break;
361 std::memcpy(&RC.Value.LongDoubleVal, Ptr, sizeof(RC.Value.LongDoubleVal));
362 break;
364 std::memcpy(&RC.Value.PtrVal, Ptr, sizeof(RC.Value.PtrVal));
365 break;
366 default:
367 reportFatalError("Unsupported RuntimeConstantType in kernel-arg reader");
368 }
369
370 return RC;
371 }
372
373 static inline RuntimeConstant
374 readRuntimeConstantFromStorage(const void *StorageBasePtr,
376 int32_t RCOffset) {
377 if (!StorageBasePtr)
378 reportFatalError("StorageBasePtr is null");
379 if (RCOffset < 0)
380 reportFatalError("Negative runtime constant offset");
381
382 const unsigned char *Base =
383 static_cast<const unsigned char *>(StorageBasePtr);
384 const unsigned char *Ptr = Base + static_cast<size_t>(RCOffset);
385
386 RuntimeConstant RC{Type, Pos, RCOffset};
387
388 switch (Type) {
390 std::memcpy(&RC.Value.BoolVal, Ptr, sizeof(RC.Value.BoolVal));
391 break;
393 std::memcpy(&RC.Value.Int8Val, Ptr, sizeof(RC.Value.Int8Val));
394 break;
396 std::memcpy(&RC.Value.Int32Val, Ptr, sizeof(RC.Value.Int32Val));
397 break;
399 std::memcpy(&RC.Value.Int64Val, Ptr, sizeof(RC.Value.Int64Val));
400 break;
402 std::memcpy(&RC.Value.FloatVal, Ptr, sizeof(RC.Value.FloatVal));
403 break;
405 std::memcpy(&RC.Value.DoubleVal, Ptr, sizeof(RC.Value.DoubleVal));
406 break;
408 std::memcpy(&RC.Value.LongDoubleVal, Ptr, sizeof(RC.Value.LongDoubleVal));
409 break;
411 std::memcpy(&RC.Value.PtrVal, Ptr, sizeof(RC.Value.PtrVal));
412 break;
413 default:
415 "Unsupported RuntimeConstantType in host-storage reader");
416 }
417
418 return RC;
419 }
420
421public:
423 readRuntimeConstantsForCallsite(void *const *KernelArgs,
424 const LambdaKernelArgLocation &Location,
425 const JitVariantMap &VariantSchema) {
426 LambdaCallsiteRuntimeConstants RuntimeConstants;
427 RuntimeConstants.reserve(VariantSchema.size());
428
429 for (auto [Slot, RC] : VariantSchema) {
430 (void)Slot;
431 RuntimeConstants.push_back(readRuntimeConstantFromKernelArgs(
432 KernelArgs, Location.KernelArgIndex, Location.Offset,
433 Location.StorageType, RC.Type, RC.Pos, RC.Offset));
434 }
435
436 llvm::sort(RuntimeConstants,
437 [](const RuntimeConstant &L, const RuntimeConstant &R) {
438 return L.Pos < R.Pos;
439 });
440 return RuntimeConstants;
441 }
442
444 void *const *KernelArgs,
445 ArrayRef<LambdaKernelArgLocation> KernelArgLocations,
446 ArrayRef<LambdaRegistry::JitVariantMap> VariantSchemas) {
447 LambdaRegistry::JitVariantVec RuntimeVariants;
448 if (VariantSchemas.empty())
449 return RuntimeVariants;
450
451 RuntimeVariants.reserve(KernelArgLocations.size());
452 const JitVariantMap &VariantSchema = VariantSchemas.front();
453
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);
460 }
461 RuntimeVariants.push_back(std::move(RuntimeVariant));
462 }
463
464 return RuntimeVariants;
465 }
466
467 static JitVariantMap readRuntimeVariantFromStorage(
468 const void *StorageBasePtr,
469 ArrayRef<LambdaRegistry::JitVariantMap> VariantSchemas) {
470 JitVariantMap RuntimeVariant;
471 if (VariantSchemas.empty())
472 return RuntimeVariant;
473
474 const JitVariantMap &VariantSchema = VariantSchemas.front();
475 for (auto [Slot, RC] : VariantSchema) {
476 RuntimeVariant[Slot] = readRuntimeConstantFromStorage(
477 StorageBasePtr, RC.Type, RC.Pos, RC.Offset);
478 }
479
480 return RuntimeVariant;
481 }
482
483 static void transformHostFunction(Module &M, Function &Lambda,
484 const JitVariantMap &RuntimeVariant) {
485 specializeCallOperator(M, Lambda, RuntimeVariant);
486 }
487
489 Module &M, uint64_t FunctorID,
490 const LambdaCallsiteRuntimeConstantsMap &CallsiteRuntimeConstants) {
491 Function *FunctorOperatorFunction =
492 findFunctorFunctorOperatorFunctionOperatorFromID(M, FunctorID);
493 Function *LambdaOperatorMethod = findLambdaOperatorForFunctor(M, FunctorID);
494 if (!FunctorOperatorFunction) {
495 if (Config::get().traceSpecializations())
496 Logger::trace("[LambdaSpec] Skipping lambda specialization: "
497 "no wrapper found for ID " +
498 std::to_string(FunctorID));
499 return;
500 }
501
502 SmallVector<CallBase *> CBToAnalyze;
503 collectWrapperCallsites(*FunctorOperatorFunction, CBToAnalyze);
504
505 size_t VariantIndex = 0;
506 for (CallBase *CB : CBToAnalyze) {
507 auto CallsiteIndex = getLambdaCallsiteIndex(*CB, FunctorID);
508 if (!CallsiteIndex)
509 reportFatalError("Missing lambda callsite metadata for callsite");
510
511 auto It = CallsiteRuntimeConstants.find(*CallsiteIndex);
512 if (It == CallsiteRuntimeConstants.end())
513 reportFatalError("Missing lambda runtime constants for callsite");
514
515 JitVariantMap RuntimeVariant;
516 for (const auto &RC : It->second) {
517 RuntimeVariant[RC.Pos] = RC;
518 }
519
520 ++VariantIndex;
521 Function *WrapperClone =
522 cloneForVariant(*FunctorOperatorFunction, FunctorID, VariantIndex);
523 specializeCallOperator(M, *WrapperClone, RuntimeVariant);
524
525 if (!LambdaOperatorMethod) {
526 // On host (and sometimes device) the lambda call operator may be fully
527 // inlined into the functor FunctorOperatorFunction call operator,
528 // leaving no separate `lambda::operator()` function to tag.
529 CB->setCalledFunction(WrapperClone);
530 continue;
531 }
532
533 Function *LambdaClone =
534 cloneForVariant(*LambdaOperatorMethod, FunctorID, VariantIndex);
535 specializeCallOperator(M, *LambdaClone, RuntimeVariant);
536
537 if (CallBase *WrapperToLambdaCall =
538 findDirectCallTo(*WrapperClone, LambdaOperatorMethod)) {
539 WrapperToLambdaCall->setCalledFunction(LambdaClone);
540 } else {
541 reportFatalError("Expected wrapper clone to call registered lambda");
542 }
543
544 CB->setCalledFunction(WrapperClone);
545 }
546 // Remove the old wrapper/lambda methods from the module if they are unused.
547 // If we actually specialized anything, we cloned and removed the original
548 // methods from the IR. Pruning is important so that downstream
549 // optimizations like TransformSharedArray don't operate on dead code,
550 // containing calls to shared_array without a constant argument. Delete the
551 // wrapper first, because usually the only remaining user of the lambda is
552 // the wrapper
553 if (FunctorOperatorFunction->users().empty())
554 FunctorOperatorFunction->eraseFromParent();
555
556 if (LambdaOperatorMethod && LambdaOperatorMethod->users().empty())
557 LambdaOperatorMethod->eraseFromParent();
558 }
559
560#ifdef PROTEUS_TRANSFORM_CONSERVATIVE
561 static void transformConservative(Module &M, uint64_t FunctorID,
562 ArrayRef<JitVariantMap> Variants) {
563 if (Variants.empty())
564 return;
565 Function *FunctorOperatorFunction =
566 findFunctorFunctorOperatorFunctionOperatorFromID(M, FunctorID);
567 Function *LambdaOperatorMethod = findLambdaOperatorForFunctor(M, FunctorID);
568 if (!FunctorOperatorFunction) {
569 if (Config::get().traceSpecializations())
570 Logger::trace("[LambdaSpec] Internal lambda specialization error:"
571 "no wrapper found for ID " +
572 std::to_string(FunctorID));
573 return;
574 }
575
576 if (!LambdaOperatorMethod) {
577 // On host (and sometimes device) the lambda call operator may be fully
578 // inlined into the functor FunctorOperatorFunction call operator, leaving
579 // no separate `lambda::operator()` function to tag.
580 if (Variants.size() == 1)
581 specializeCallOperator(M, *FunctorOperatorFunction, Variants.front());
582 return;
583 }
584
585 if (Variants.size() == 1) {
586 specializeCallOperator(M, *LambdaOperatorMethod, Variants.front());
587 return;
588 }
589
590 CallBase *OrigCall =
591 findDirectCallTo(*FunctorOperatorFunction, LambdaOperatorMethod);
592 if (!OrigCall)
593 return;
594
595 CallingConv::ID CallConv = OrigCall->getCallingConv();
596 AttributeList CallAttrs = OrigCall->getAttributes();
597
598 SmallVector<Value *, 8> CallArgs;
599 CallArgs.reserve(OrigCall->arg_size());
600 for (Use &U : OrigCall->args())
601 CallArgs.push_back(U.get());
602
603 Value *LambdaObjPtr = OrigCall->getArgOperand(0);
604 StructType *LambdaTy = inferStructTypeFromGEPValue(LambdaObjPtr);
605
606 // Split out the original call (and everything after it) into a tail block.
607 BasicBlock *EntryBB = OrigCall->getParent();
608 BasicBlock *TailBB =
609 EntryBB->splitBasicBlock(OrigCall, "proteus.after_call");
610
611 // Remove the unconditional branch produced by splitBasicBlock.
612 EntryBB->getTerminator()->eraseFromParent();
613
614 // Create a PHI in the tail block for any call result.
615 PHINode *ResultPhi = nullptr;
616 Type *RetTy = OrigCall->getType();
617 if (!RetTy->isVoidTy()) {
618 ResultPhi =
619 PHINode::Create(RetTy, Variants.size() + 1,
620 "proteus.lambda_dispatch.result", &*TailBB->begin());
621 OrigCall->replaceAllUsesWith(ResultPhi);
622 }
623
624 Instruction *AfterCallIP = OrigCall->getNextNode();
625 OrigCall->eraseFromParent();
626
627 LLVMContext &Ctx = M.getContext();
628
629 BasicBlock *FallbackBB = BasicBlock::Create(
630 Ctx, "proteus.lambda_dispatch.fallback", FunctorOperatorFunction);
631
632 // Build the chain of check blocks (entry -> check0 -> ... -> fallback).
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 =
639 (I + 1 == E)
640 ? FallbackBB
641 : BasicBlock::Create(
642 Ctx, "proteus.lambda_dispatch.check." + Twine(I + 1),
643 FunctorOperatorFunction);
644
645 {
646 IRBuilder<> B(CurCheckBB);
647 Value *Match = buildVariantMatch(B, M.getDataLayout(), LambdaObjPtr,
648 LambdaTy, Variants[I]);
649 if (!Match) {
650 // Unsupported runtime-constant type for dispatch; fall back to the
651 // unspecialized call.
652 BranchInst::Create(FallbackBB, CurCheckBB);
653 } else {
654 BranchInst::Create(VariantBB, NextCheckBB, Match, CurCheckBB);
655 }
656 }
657
658 // Emit the specialized call in the variant block.
659 {
660 IRBuilder<> B(VariantBB);
661 Function *Clone = cloneForVariant(*LambdaOperatorMethod, FunctorID, I);
662 specializeCallOperator(M, *Clone, Variants[I]);
663
664 CallInst *C = B.CreateCall(Clone, CallArgs);
665 C->setCallingConv(CallConv);
666 C->setAttributes(CallAttrs);
667
668 if (ResultPhi)
669 ResultPhi->addIncoming(C, VariantBB);
670
671 B.CreateBr(TailBB);
672 }
673
674 CurCheckBB = NextCheckBB;
675 }
676
677 // Emit the fallback call to the original, unspecialized operator.
678 {
679 IRBuilder<> B(FallbackBB);
680 CallInst *C = B.CreateCall(LambdaOperatorMethod, CallArgs);
681 C->setCallingConv(CallConv);
682 C->setAttributes(CallAttrs);
683
684 if (ResultPhi)
685 ResultPhi->addIncoming(C, FallbackBB);
686
687 B.CreateBr(TailBB);
688 }
689
690 // If the call was immediately followed by a terminator, splitBasicBlock
691 // still produced a valid tail block. No further fixup required.
692 (void)AfterCallIP;
693 }
694
695 static Value *buildSingleCompare(IRBuilder<> &B, const DataLayout &DL,
696 Value *LambdaObjPtr, StructType *LambdaTy,
697 const RuntimeConstant &RC) {
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);
702 } else {
703 if (RC.Offset < 0)
704 return nullptr;
705 FieldPtr = B.CreateGEP(B.getInt8Ty(), LambdaObjPtr,
706 B.getInt64(static_cast<uint64_t>(RC.Offset)));
707 }
708
709 switch (RC.Type) {
711 Value *Loaded = B.CreateAlignedLoad(B.getInt8Ty(), FieldPtr, Align(1));
712 Value *C = B.getInt8(RC.Value.BoolVal ? 1 : 0);
713 return B.CreateICmpEQ(Loaded, C);
714 }
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);
719 }
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);
724 }
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);
729 }
731 Value *LoadedF = B.CreateAlignedLoad(B.getFloatTy(), FieldPtr, Align(1));
732 Value *LoadedBits = B.CreateBitCast(LoadedF, B.getInt32Ty());
733 uint32_t Bits = 0;
734 std::memcpy(&Bits, &RC.Value.FloatVal, sizeof(Bits));
735 Value *C = B.getInt32(Bits);
736 return B.CreateICmpEQ(LoadedBits, C);
737 }
739 Value *LoadedF = B.CreateAlignedLoad(B.getDoubleTy(), FieldPtr, Align(1));
740 Value *LoadedBits = B.CreateBitCast(LoadedF, B.getInt64Ty());
741 uint64_t Bits = 0;
742 std::memcpy(&Bits, &RC.Value.DoubleVal, sizeof(Bits));
743 Value *C = B.getInt64(Bits);
744 return B.CreateICmpEQ(LoadedBits, C);
745 }
747 Type *IntPtrTy = DL.getIntPtrType(B.getContext());
748 Value *Loaded = B.CreateAlignedLoad(IntPtrTy, FieldPtr, Align(1));
749 uint64_t Bits =
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);
753 }
754 default:
755 return nullptr;
756 }
757 }
758
759 static Value *buildVariantMatch(IRBuilder<> &B, const DataLayout &DL,
760 Value *LambdaObjPtr, StructType *LambdaTy,
761 const JitVariantMap &Variant) {
762 if (Variant.empty())
763 return B.getInt1(true);
764
765 SmallVector<int32_t, 16> Keys;
766 Keys.reserve(Variant.size());
767 for (const auto &KV : Variant)
768 Keys.push_back(KV.first);
769
770 llvm::sort(Keys.begin(), Keys.end());
771
772 Value *Match = B.getInt1(true);
773 for (int32_t K : Keys) {
774 auto It = Variant.find(K);
775 if (It == Variant.end())
776 reportFatalError("Internal error: DenseMap key vanished during match");
777 Value *Cmp =
778 buildSingleCompare(B, DL, LambdaObjPtr, LambdaTy, It->second);
779 if (!Cmp)
780 return nullptr;
781 Match = B.CreateAnd(Match, Cmp);
782 }
783
784 return Match;
785 }
786#endif
787};
788
789} // namespace proteus
790
791#endif
uint32_t int32_t Type
Definition CompilerInterfaceDevice.cpp:98
uint64_t uint32_t CallsiteIndex
Definition CompilerInterfaceDevice.cpp:73
uint64_t uint32_t uint32_t int64_t Offset
Definition CompilerInterfaceDevice.cpp:75
uint32_t int32_t int32_t Pos
Definition CompilerInterfaceDevice.cpp:98
uint64_t uint32_t uint32_t KernelArgIndex
Definition CompilerInterfaceDevice.cpp:74
uint64_t uint32_t uint32_t int64_t int32_t StorageType
Definition CompilerInterfaceDevice.cpp:76
#define PROTEUS_DBG(x)
Definition Debug.h:9
static Config & get()
Definition Config.h:371
SmallVector< JitVariantMap, 4 > JitVariantVec
Definition LambdaRegistry.h:38
DenseMap< int32_t, RuntimeConstant > JitVariantMap
Definition LambdaRegistry.h:37
static void trace(llvm::StringRef Msg)
Definition Logger.h:30
static llvm::raw_ostream & logs(const std::string &Name)
Definition Logger.h:19
Definition TransformLambdaSpecialization.h:72
static JitVariantMap readRuntimeVariantFromStorage(const void *StorageBasePtr, ArrayRef< LambdaRegistry::JitVariantMap > VariantSchemas)
Definition TransformLambdaSpecialization.h:467
static LambdaCallsiteRuntimeConstants readRuntimeConstantsForCallsite(void *const *KernelArgs, const LambdaKernelArgLocation &Location, const JitVariantMap &VariantSchema)
Definition TransformLambdaSpecialization.h:423
static void transformDeviceKernel(Module &M, uint64_t FunctorID, const LambdaCallsiteRuntimeConstantsMap &CallsiteRuntimeConstants)
Definition TransformLambdaSpecialization.h:488
static void transformHostFunction(Module &M, Function &Lambda, const JitVariantMap &RuntimeVariant)
Definition TransformLambdaSpecialization.h:483
static LambdaRegistry::JitVariantVec readRuntimeVariantsFromKernelArgs(void *const *KernelArgs, ArrayRef< LambdaKernelArgLocation > KernelArgLocations, ArrayRef< LambdaRegistry::JitVariantMap > VariantSchemas)
Definition TransformLambdaSpecialization.h:443
Definition CompiledLibrary.h:8
Definition MemoryCache.h:27
void findFunctionsWithU64Metadata(llvm::Module &M, llvm::StringRef Key, llvm::SmallVectorImpl< std::pair< llvm::Function *, std::uint64_t > > &Out)
Definition CoreLLVM.h:426
RuntimeConstantType
Definition CompilerInterfaceTypes.h:20
@ NONE
Definition CompilerInterfaceTypes.h:22
@ INT32
Definition CompilerInterfaceTypes.h:25
@ INT64
Definition CompilerInterfaceTypes.h:26
@ FLOAT
Definition CompilerInterfaceTypes.h:27
@ LONG_DOUBLE
Definition CompilerInterfaceTypes.h:29
@ INT8
Definition CompilerInterfaceTypes.h:24
@ BOOL
Definition CompilerInterfaceTypes.h:23
@ PTR
Definition CompilerInterfaceTypes.h:30
@ DOUBLE
Definition CompilerInterfaceTypes.h:28
void reportFatalError(const llvm::Twine &Reason, const char *FILE, unsigned Line)
Definition Error.cpp:14
Constant * getConstant(LLVMContext &Ctx, Type *ArgType, const RuntimeConstant &RC)
Definition TransformLambdaSpecialization.h:42
std::optional< uint32_t > getLambdaCallsiteIndex(const CallBase &CB, uint64_t ExpectedLambdaID)
Definition LambdaCallsite.h:68
DenseMap< uint64_t, LambdaCallsiteRuntimeConstants > LambdaCallsiteRuntimeConstantsMap
Definition LambdaCallsite.h:32
SmallVector< RuntimeConstant, 8 > LambdaCallsiteRuntimeConstants
Definition LambdaCallsite.h:30
Definition LambdaCallsite.h:19
int64_t Offset
Definition LambdaCallsite.h:21
RuntimeConstantType StorageType
Definition LambdaCallsite.h:22
uint32_t KernelArgIndex
Definition LambdaCallsite.h:20
Definition CompilerInterfaceTypes.h:72
RuntimeConstantValue Value
Definition CompilerInterfaceTypes.h:73
RuntimeConstantType Type
Definition CompilerInterfaceTypes.h:74
int32_t Offset
Definition CompilerInterfaceTypes.h:76
int32_t Pos
Definition CompilerInterfaceTypes.h:75
double DoubleVal
Definition CompilerInterfaceTypes.h:65
int64_t Int64Val
Definition CompilerInterfaceTypes.h:63
void * PtrVal
Definition CompilerInterfaceTypes.h:67
bool BoolVal
Definition CompilerInterfaceTypes.h:60
int8_t Int8Val
Definition CompilerInterfaceTypes.h:61
int32_t Int32Val
Definition CompilerInterfaceTypes.h:62
float FloatVal
Definition CompilerInterfaceTypes.h:64
long double LongDoubleVal
Definition CompilerInterfaceTypes.h:66