1//===- MLInlineAdvisor.cpp - machine learned InlineAdvisor ----------------===//
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// This file implements the interface between the inliner and a learned model.
10// It delegates model evaluation to either the AOT compiled model (the
11// 'release' mode) or a runtime-loaded model (the 'development' case).
12//
13//===----------------------------------------------------------------------===//
14#include "llvm/Analysis/MLInlineAdvisor.h"
15#include "llvm/ADT/SCCIterator.h"
16#include "llvm/Analysis/AssumptionCache.h"
17#include "llvm/Analysis/BlockFrequencyInfo.h"
18#include "llvm/Analysis/CallGraph.h"
19#include "llvm/Analysis/FunctionPropertiesAnalysis.h"
20#include "llvm/Analysis/InlineCost.h"
21#include "llvm/Analysis/InlineModelFeatureMaps.h"
22#include "llvm/Analysis/LazyCallGraph.h"
23#include "llvm/Analysis/LoopInfo.h"
24#include "llvm/Analysis/MLModelRunner.h"
25#include "llvm/Analysis/OptimizationRemarkEmitter.h"
26#include "llvm/Analysis/ProfileSummaryInfo.h"
27#include "llvm/Analysis/ReleaseModeModelRunner.h"
28#include "llvm/Analysis/TargetTransformInfo.h"
29#include "llvm/Analysis/TensorSpec.h"
30#include "llvm/Analysis/Utils/MLGOUtils.h"
31#include "llvm/IR/Dominators.h"
32#include "llvm/IR/InstIterator.h"
33#include "llvm/IR/Module.h"
34#include "llvm/IR/PassManager.h"
35#include "llvm/Support/CommandLine.h"
36
37using namespace llvm;
38
39static cl::opt<std::string> InteractiveChannelBaseName(
40 "inliner-interactive-channel-base", cl::Hidden,
41 cl::desc(
42 "Base file path for the interactive mode. The incoming filename should "
43 "have the name <inliner-interactive-channel-base>.in, while the "
44 "outgoing name should be <inliner-interactive-channel-base>.out"));
45static const std::string InclDefaultMsg =
46 (Twine("In interactive mode, also send the default policy decision: ") +
47 DefaultDecisionName + ".")
48 .str();
49static cl::opt<bool>
50 InteractiveIncludeDefault("inliner-interactive-include-default", cl::Hidden,
51 cl::desc(InclDefaultMsg));
52
53enum class SkipMLPolicyCriteria { Never, IfCallerIsNotCold };
54
55static cl::opt<SkipMLPolicyCriteria> SkipPolicy(
56 "ml-inliner-skip-policy", cl::Hidden, cl::init(Val: SkipMLPolicyCriteria::Never),
57 cl::values(clEnumValN(SkipMLPolicyCriteria::Never, "never", "never"),
58 clEnumValN(SkipMLPolicyCriteria::IfCallerIsNotCold,
59 "if-caller-not-cold", "if the caller is not cold")));
60
61static cl::opt<std::string> ModelSelector("ml-inliner-model-selector",
62 cl::Hidden, cl::init(Val: ""));
63
64static cl::opt<bool> StopImmediatelyForTest("ml-inliner-stop-immediately",
65 cl::Hidden);
66
67#if defined(LLVM_HAVE_TF_AOT_INLINERSIZEMODEL)
68// codegen-ed file
69#include "InlinerSizeModel.h" // NOLINT
70using CompiledModelType = llvm::InlinerSizeModel;
71#else
72using CompiledModelType = NoopSavedModelImpl;
73#endif
74
75#include "llvm/Analysis/EmitCModelRunner.h"
76#include "llvm/Analysis/InlinerModels.h"
77
78enum class EmitCModelChoice {
79 Default,
80#define MLGO_MODEL(CLASS_NAME, CLI_FLAG) CLASS_NAME,
81#include "llvm/Analysis/InlinerModels.def"
82};
83
84static llvm::cl::opt<EmitCModelChoice> SelectedMLGOModel(
85 "mlgo-model", llvm::cl::desc("Select the MLGO model to execute:"),
86 llvm::cl::init(Val: EmitCModelChoice::Default),
87 llvm::cl::values(clEnumValN(EmitCModelChoice::Default, "default",
88 "Use standard heuristic")
89#define MLGO_MODEL(CLASS_NAME, CLI_FLAG) \
90 , clEnumValN(EmitCModelChoice::CLASS_NAME, CLI_FLAG, \
91 "Use the " CLI_FLAG " MLGO model")
92#include "llvm/Analysis/InlinerModels.def"
93 ));
94
95static std::unique_ptr<MLModelRunner>
96createEmitCModelRunner(LLVMContext &Ctx,
97 const std::vector<TensorSpec> &InputFeatures) {
98 switch (SelectedMLGOModel) {
99 case EmitCModelChoice::Default:
100 return nullptr;
101#define MLGO_MODEL(CLASS_NAME, CLI_FLAG) \
102 case EmitCModelChoice::CLASS_NAME: \
103 return std::make_unique<EmitCModelRunner<CLASS_NAME>>(Ctx, InputFeatures);
104#include "llvm/Analysis/InlinerModels.def"
105 }
106 llvm_unreachable("Unknown MLGO model type!");
107}
108
109std::unique_ptr<InlineAdvisor>
110llvm::getReleaseModeAdvisor(Module &M, ModuleAnalysisManager &MAM,
111 std::function<bool(CallBase &)> GetDefaultAdvice) {
112 if (!isReleaseModelValid<CompiledModelType>(InteractiveChannelBaseName,
113 SelectedModel: SelectedMLGOModel))
114 return nullptr;
115 auto RunnerFactory = [&](const std::vector<TensorSpec> &InputFeatures)
116 -> std::unique_ptr<MLModelRunner> {
117 return createReleaseModeModelRunner<CompiledModelType>(
118 Ctx&: M.getContext(), InputFeatures, DecisionName, InteractiveChannelBaseName,
119 InteractiveDecisionSpec: InlineDecisionSpec, CreateEmitCModelRunner&: createEmitCModelRunner,
120 Options: EmbeddedModelRunnerOptions().setModelSelector(ModelSelector));
121 };
122 return std::make_unique<MLInlineAdvisor>(args&: M, args&: MAM, args&: RunnerFactory,
123 args&: GetDefaultAdvice);
124}
125
126#define DEBUG_TYPE "inline-ml"
127
128static cl::opt<float> SizeIncreaseThreshold(
129 "ml-advisor-size-increase-threshold", cl::Hidden,
130 cl::desc("Maximum factor by which expected native size may increase before "
131 "blocking any further inlining."),
132 cl::init(Val: 2.0));
133
134static cl::opt<bool> KeepFPICache(
135 "ml-advisor-keep-fpi-cache", cl::Hidden,
136 cl::desc(
137 "For test - keep the ML Inline advisor's FunctionPropertiesInfo cache"),
138 cl::init(Val: false));
139
140const std::vector<TensorSpec> &MLInlineAdvisor::getInitialFeatureMap() {
141 // clang-format off
142static std::vector<TensorSpec> FeatureMap{
143#define POPULATE_NAMES(DTYPE, SHAPE, NAME, __) TensorSpec::createSpec<DTYPE>(#NAME, SHAPE),
144// InlineCost features - these must come first
145 INLINE_COST_FEATURE_ITERATOR(POPULATE_NAMES)
146
147// Non-cost features
148 INLINE_FEATURE_ITERATOR(POPULATE_NAMES)
149#undef POPULATE_NAMES
150};
151 // clang-format on
152 return FeatureMap;
153}
154
155const char *const llvm::DecisionName = "inlining_decision";
156const TensorSpec llvm::InlineDecisionSpec =
157 TensorSpec::createSpec<int64_t>(Name: DecisionName, Shape: {1});
158const char *const llvm::DefaultDecisionName = "inlining_default";
159const TensorSpec llvm::DefaultDecisionSpec =
160 TensorSpec::createSpec<int64_t>(Name: DefaultDecisionName, Shape: {1});
161const char *const llvm::RewardName = "delta_size";
162
163CallBase *getInlinableCS(Instruction &I) {
164 if (auto *CS = dyn_cast<CallBase>(Val: &I))
165 if (Function *Callee = CS->getCalledFunction()) {
166 if (!Callee->isDeclaration()) {
167 return CS;
168 }
169 }
170 return nullptr;
171}
172
173MLInlineAdvisor::MLInlineAdvisor(
174 Module &M, ModuleAnalysisManager &MAM,
175 std::function<
176 std::unique_ptr<MLModelRunner>(const std::vector<TensorSpec> &)>
177 GetModelRunner,
178 std::function<bool(CallBase &)> GetDefaultAdvice)
179 : InlineAdvisor(
180 M, MAM.getResult<FunctionAnalysisManagerModuleProxy>(IR&: M).getManager()),
181 GetDefaultAdvice(GetDefaultAdvice), FeatureMap(getInitialFeatureMap()),
182 CG(MAM.getResult<LazyCallGraphAnalysis>(IR&: M)),
183 UseIR2Vec(MAM.getCachedResult<IR2VecVocabAnalysis>(IR&: M) != nullptr),
184 InitialIRSize(getModuleIRSize()), CurrentIRSize(InitialIRSize),
185 PSI(MAM.getResult<ProfileSummaryAnalysis>(IR&: M)) {
186 // Extract the 'call site height' feature - the position of a call site
187 // relative to the farthest statically reachable SCC node. We don't mutate
188 // this value while inlining happens. Empirically, this feature proved
189 // critical in behavioral cloning - i.e. training a model to mimic the manual
190 // heuristic's decisions - and, thus, equally important for training for
191 // improvement.
192 CallGraph CGraph(M);
193 for (auto I = scc_begin(G: &CGraph); !I.isAtEnd(); ++I) {
194 const std::vector<CallGraphNode *> &CGNodes = *I;
195 unsigned Level = 0;
196 for (auto *CGNode : CGNodes) {
197 Function *F = CGNode->getFunction();
198 if (!F || F->isDeclaration())
199 continue;
200 for (auto &I : instructions(F)) {
201 if (auto *CS = getInlinableCS(I)) {
202 auto *Called = CS->getCalledFunction();
203 auto Pos = FunctionLevels.find(x: &CG.get(F&: *Called));
204 // In bottom up traversal, an inlinable callee is either in the
205 // same SCC, or to a function in a visited SCC. So not finding its
206 // level means we haven't visited it yet, meaning it's in this SCC.
207 if (Pos == FunctionLevels.end())
208 continue;
209 Level = std::max(a: Level, b: Pos->second + 1);
210 }
211 }
212 }
213 for (auto *CGNode : CGNodes) {
214 Function *F = CGNode->getFunction();
215 if (F && !F->isDeclaration())
216 FunctionLevels[&CG.get(F&: *F)] = Level;
217 }
218 }
219 for (auto KVP : FunctionLevels) {
220 AllNodes.insert(V: KVP.first);
221 EdgeCount += getLocalCalls(F&: KVP.first->getFunction());
222 }
223 NodeCount = AllNodes.size();
224
225 if (auto *IR2VecVocabResult = MAM.getCachedResult<IR2VecVocabAnalysis>(IR&: M)) {
226 if (!IR2VecVocabResult->isValid()) {
227 M.getContext().emitError(ErrorStr: "IR2VecVocabAnalysis is not valid");
228 return;
229 }
230 // Add the IR2Vec features to the feature map
231 auto IR2VecDim = IR2VecVocabResult->getDimension();
232 FeatureMap.push_back(
233 x: TensorSpec::createSpec<float>(Name: "callee_embedding", Shape: {IR2VecDim}));
234 FeatureMap.push_back(
235 x: TensorSpec::createSpec<float>(Name: "caller_embedding", Shape: {IR2VecDim}));
236 }
237 if (InteractiveIncludeDefault)
238 FeatureMap.push_back(x: DefaultDecisionSpec);
239
240 ModelRunner = GetModelRunner(getFeatureMap());
241 if (!ModelRunner) {
242 M.getContext().emitError(ErrorStr: "Could not create model runner");
243 return;
244 }
245 ModelRunner->switchContext(Name: "");
246 ForceStop = StopImmediatelyForTest;
247}
248
249unsigned MLInlineAdvisor::getInitialFunctionLevel(const Function &F) const {
250 return CG.lookup(F) ? FunctionLevels.at(k: CG.lookup(F)) : 0;
251}
252
253void MLInlineAdvisor::onPassEntry(LazyCallGraph::SCC *CurSCC) {
254 if (!CurSCC || ForceStop)
255 return;
256 FPICache.clear();
257 // Function passes executed between InlinerPass runs may have changed the
258 // module-wide features.
259 // The cgscc pass manager rules are such that:
260 // - if a pass leads to merging SCCs, then the pipeline is restarted on the
261 // merged SCC
262 // - if a pass leads to splitting the SCC, then we continue with one of the
263 // splits
264 // This means that the NodesInLastSCC is a superset (not strict) of the nodes
265 // that subsequent passes would have processed
266 // - in addition, if new Nodes were created by a pass (e.g. CoroSplit),
267 // they'd be adjacent to Nodes in the last SCC. So we just need to check the
268 // boundary of Nodes in NodesInLastSCC for Nodes we haven't seen. We don't
269 // care about the nature of the Edge (call or ref). `FunctionLevels`-wise, we
270 // record them at the same level as the original node (this is a choice, may
271 // need revisiting).
272 // - nodes are only deleted at the end of a call graph walk where they are
273 // batch deleted, so we shouldn't see any dead nodes here.
274 while (!NodesInLastSCC.empty()) {
275 const auto *N = *NodesInLastSCC.begin();
276 assert(!N->isDead());
277 NodesInLastSCC.erase(Ptr: N);
278 EdgeCount += getLocalCalls(F&: N->getFunction());
279 const auto NLevel = FunctionLevels.at(k: N);
280 for (const auto &E : *(*N)) {
281 const auto *AdjNode = &E.getNode();
282 assert(!AdjNode->isDead() && !AdjNode->getFunction().isDeclaration());
283 auto I = AllNodes.insert(V: AdjNode);
284 // We've discovered a new function.
285 if (I.second) {
286 ++NodeCount;
287 NodesInLastSCC.insert(Ptr: AdjNode);
288 FunctionLevels[AdjNode] = NLevel;
289 }
290 }
291 }
292
293 EdgeCount -= EdgesOfLastSeenNodes;
294 EdgesOfLastSeenNodes = 0;
295
296 // (Re)use NodesInLastSCC to remember the nodes in the SCC right now,
297 // in case the SCC is split before onPassExit and some nodes are split out
298 assert(NodesInLastSCC.empty());
299 for (const auto &N : *CurSCC)
300 NodesInLastSCC.insert(Ptr: &N);
301}
302
303void MLInlineAdvisor::onPassExit(LazyCallGraph::SCC *CurSCC) {
304 // No need to keep this around - function passes will invalidate it.
305 if (!KeepFPICache)
306 FPICache.clear();
307 if (!CurSCC || ForceStop)
308 return;
309 // Keep track of the nodes and edges we last saw. Then, in onPassEntry,
310 // we update the node count and edge count from the subset of these nodes that
311 // survived.
312 EdgesOfLastSeenNodes = 0;
313
314 // Check on nodes that were in SCC onPassEntry
315 for (const LazyCallGraph::Node *N : NodesInLastSCC) {
316 assert(!N->isDead());
317 EdgesOfLastSeenNodes += getLocalCalls(F&: N->getFunction());
318 }
319
320 // Check on nodes that may have got added to SCC
321 for (const auto &N : *CurSCC) {
322 assert(!N.isDead());
323 auto I = NodesInLastSCC.insert(Ptr: &N);
324 if (I.second)
325 EdgesOfLastSeenNodes += getLocalCalls(F&: N.getFunction());
326 }
327 assert(NodeCount >= NodesInLastSCC.size());
328 assert(EdgeCount >= EdgesOfLastSeenNodes);
329}
330
331int64_t MLInlineAdvisor::getLocalCalls(Function &F) {
332 return getCachedFPI(F).DirectCallsToDefinedFunctions;
333}
334
335// Update the internal state of the advisor, and force invalidate feature
336// analysis. Currently, we maintain minimal (and very simple) global state - the
337// number of functions and the number of static calls. We also keep track of the
338// total IR size in this module, to stop misbehaving policies at a certain bloat
339// factor (SizeIncreaseThreshold)
340void MLInlineAdvisor::onSuccessfulInlining(const MLInlineAdvice &Advice,
341 bool CalleeWasDeleted) {
342 assert(!ForceStop);
343 Function *Caller = Advice.getCaller();
344 Function *Callee = Advice.getCallee();
345 // The caller features aren't valid anymore.
346 {
347 PreservedAnalyses PA = PreservedAnalyses::all();
348 PA.abandon<FunctionPropertiesAnalysis>();
349 PA.abandon<LoopAnalysis>();
350 FAM.invalidate(IR&: *Caller, PA);
351 }
352 Advice.updateCachedCallerFPI(FAM);
353 if (Caller == Callee) {
354 assert(!CalleeWasDeleted);
355 // We double-counted CallerAndCalleeEdges - since the caller and callee
356 // would be the same
357 assert(Advice.CallerAndCalleeEdges % 2 == 0);
358 CurrentIRSize += getIRSize(F&: *Caller) - Advice.CallerIRSize;
359 EdgeCount += getCachedFPI(*Caller).DirectCallsToDefinedFunctions -
360 Advice.CallerAndCalleeEdges / 2;
361 // The NodeCount would stay the same.
362 } else {
363 int64_t IRSizeAfter =
364 getIRSize(F&: *Caller) + (CalleeWasDeleted ? 0 : Advice.CalleeIRSize);
365 CurrentIRSize += IRSizeAfter - (Advice.CallerIRSize + Advice.CalleeIRSize);
366
367 // We can delta-update module-wide features. We know the inlining only
368 // changed the caller, and maybe the callee (by deleting the latter). Nodes
369 // are simple to update. For edges, we 'forget' the edges that the caller
370 // and callee used to have before inlining, and add back what they currently
371 // have together.
372 int64_t NewCallerAndCalleeEdges =
373 getCachedFPI(*Caller).DirectCallsToDefinedFunctions;
374
375 // A dead function's node is not actually removed from the call graph until
376 // the end of the call graph walk, but the node no longer belongs to any
377 // valid SCC.
378 if (CalleeWasDeleted) {
379 --NodeCount;
380 NodesInLastSCC.erase(Ptr: CG.lookup(F: *Callee));
381 DeadFunctions.insert(V: Callee);
382 } else {
383 NewCallerAndCalleeEdges +=
384 getCachedFPI(*Callee).DirectCallsToDefinedFunctions;
385 }
386 EdgeCount += (NewCallerAndCalleeEdges - Advice.CallerAndCalleeEdges);
387 }
388 if (CurrentIRSize > SizeIncreaseThreshold * InitialIRSize)
389 ForceStop = true;
390
391 assert(CurrentIRSize >= 0 && EdgeCount >= 0 && NodeCount >= 0);
392}
393
394int64_t MLInlineAdvisor::getModuleIRSize() const {
395 int64_t Ret = 0;
396 for (auto &F : M)
397 if (!F.isDeclaration())
398 Ret += getIRSize(F);
399 return Ret;
400}
401
402FunctionPropertiesInfo &MLInlineAdvisor::getCachedFPI(Function &F) const {
403 auto InsertPair = FPICache.try_emplace(k: &F);
404 if (!InsertPair.second)
405 return InsertPair.first->second;
406 InsertPair.first->second = FAM.getResult<FunctionPropertiesAnalysis>(IR&: F);
407 return InsertPair.first->second;
408}
409
410std::unique_ptr<InlineAdvice> MLInlineAdvisor::getAdviceImpl(CallBase &CB) {
411 if (auto Skip = getSkipAdviceIfUnreachableCallsite(CB))
412 return Skip;
413
414 auto &Caller = *CB.getCaller();
415 auto &Callee = *CB.getCalledFunction();
416
417 auto GetAssumptionCache = [&](Function &F) -> AssumptionCache & {
418 return FAM.getResult<AssumptionAnalysis>(IR&: F);
419 };
420 auto &TIR = FAM.getResult<TargetIRAnalysis>(IR&: Callee);
421 auto &ORE = FAM.getResult<OptimizationRemarkEmitterAnalysis>(IR&: Caller);
422
423 if (SkipPolicy == SkipMLPolicyCriteria::IfCallerIsNotCold) {
424 if (!PSI.isFunctionEntryCold(F: &Caller)) {
425 // Return a MLInlineAdvice, despite delegating to the default advice,
426 // because we need to keep track of the internal state. This is different
427 // from the other instances where we return a "default" InlineAdvice,
428 // which happen at points we won't come back to the MLAdvisor for
429 // decisions requiring that state.
430 return ForceStop ? std::make_unique<InlineAdvice>(args: this, args&: CB, args&: ORE,
431 args: GetDefaultAdvice(CB))
432 : std::make_unique<MLInlineAdvice>(args: this, args&: CB, args&: ORE,
433 args: GetDefaultAdvice(CB));
434 }
435 }
436 auto MandatoryKind = InlineAdvisor::getMandatoryKind(CB, FAM, ORE);
437 // If this is a "never inline" case, there won't be any changes to internal
438 // state we need to track, so we can just return the base InlineAdvice, which
439 // will do nothing interesting.
440 // Same thing if this is a recursive case.
441 if (MandatoryKind == InlineAdvisor::MandatoryInliningKind::Never ||
442 &Caller == &Callee)
443 return getMandatoryAdvice(CB, Advice: false);
444
445 bool Mandatory =
446 MandatoryKind == InlineAdvisor::MandatoryInliningKind::Always;
447
448 // If we need to stop, we won't want to track anymore any state changes, so
449 // we just return the base InlineAdvice, which acts as a noop.
450 if (ForceStop) {
451 ORE.emit(RemarkBuilder: [&] {
452 return OptimizationRemarkMissed(DEBUG_TYPE, "ForceStop", &CB)
453 << "Won't attempt inlining because module size grew too much.";
454 });
455 return std::make_unique<InlineAdvice>(args: this, args&: CB, args&: ORE, args&: Mandatory);
456 }
457
458 int CostEstimate = 0;
459 if (!Mandatory) {
460 auto IsCallSiteInlinable =
461 llvm::getInliningCostEstimate(Call&: CB, CalleeTTI&: TIR, GetAssumptionCache);
462 if (!IsCallSiteInlinable) {
463 // We can't inline this for correctness reasons, so return the base
464 // InlineAdvice, as we don't care about tracking any state changes (which
465 // won't happen).
466 return std::make_unique<InlineAdvice>(args: this, args&: CB, args&: ORE, args: false);
467 }
468 CostEstimate = *IsCallSiteInlinable;
469 }
470
471 const auto CostFeatures =
472 llvm::getInliningCostFeatures(Call&: CB, CalleeTTI&: TIR, GetAssumptionCache);
473 if (!CostFeatures) {
474 return std::make_unique<InlineAdvice>(args: this, args&: CB, args&: ORE, args: false);
475 }
476
477 if (Mandatory)
478 return getMandatoryAdvice(CB, Advice: true);
479
480 auto NumCtantParams = 0;
481 for (auto I = CB.arg_begin(), E = CB.arg_end(); I != E; ++I) {
482 NumCtantParams += (isa<Constant>(Val: *I));
483 }
484
485 auto &CallerBefore = getCachedFPI(F&: Caller);
486 auto &CalleeBefore = getCachedFPI(F&: Callee);
487
488 *ModelRunner->getTensor<int64_t>(FeatureID: FeatureIndex::callee_basic_block_count) =
489 CalleeBefore.BasicBlockCount;
490 *ModelRunner->getTensor<int64_t>(FeatureID: FeatureIndex::callsite_height) =
491 getInitialFunctionLevel(F: Caller);
492 *ModelRunner->getTensor<int64_t>(FeatureID: FeatureIndex::node_count) = NodeCount;
493 *ModelRunner->getTensor<int64_t>(FeatureID: FeatureIndex::nr_ctant_params) =
494 NumCtantParams;
495 *ModelRunner->getTensor<int64_t>(FeatureID: FeatureIndex::edge_count) = EdgeCount;
496 *ModelRunner->getTensor<int64_t>(FeatureID: FeatureIndex::caller_users) =
497 CallerBefore.Uses;
498 *ModelRunner->getTensor<int64_t>(
499 FeatureID: FeatureIndex::caller_conditionally_executed_blocks) =
500 CallerBefore.BlocksReachedFromConditionalInstruction;
501 *ModelRunner->getTensor<int64_t>(FeatureID: FeatureIndex::caller_basic_block_count) =
502 CallerBefore.BasicBlockCount;
503 *ModelRunner->getTensor<int64_t>(
504 FeatureID: FeatureIndex::callee_conditionally_executed_blocks) =
505 CalleeBefore.BlocksReachedFromConditionalInstruction;
506 *ModelRunner->getTensor<int64_t>(FeatureID: FeatureIndex::callee_users) =
507 CalleeBefore.Uses;
508 *ModelRunner->getTensor<int64_t>(FeatureID: FeatureIndex::cost_estimate) = CostEstimate;
509 *ModelRunner->getTensor<int64_t>(FeatureID: FeatureIndex::is_callee_avail_external) =
510 Callee.hasAvailableExternallyLinkage();
511 *ModelRunner->getTensor<int64_t>(FeatureID: FeatureIndex::is_caller_avail_external) =
512 Caller.hasAvailableExternallyLinkage();
513
514 if (UseIR2Vec) {
515 // Python side expects float embeddings. The IR2Vec embeddings are doubles
516 // as of now due to the restriction of fromJSON method used by the
517 // readVocabulary method in ir2vec::Embeddings.
518 auto setEmbedding = [&](const ir2vec::Embedding &Embedding,
519 FeatureIndex Index) {
520 llvm::transform(Range: Embedding, d_first: ModelRunner->getTensor<float>(FeatureID: Index),
521 F: [](double Val) { return static_cast<float>(Val); });
522 };
523
524 setEmbedding(CalleeBefore.getFunctionEmbedding(),
525 FeatureIndex::callee_embedding);
526 setEmbedding(CallerBefore.getFunctionEmbedding(),
527 FeatureIndex::caller_embedding);
528 }
529
530 // Add the cost features
531 for (size_t I = 0;
532 I < static_cast<size_t>(InlineCostFeatureIndex::NumberOfFeatures); ++I) {
533 *ModelRunner->getTensor<int64_t>(FeatureID: inlineCostFeatureToMlFeature(
534 Feature: static_cast<InlineCostFeatureIndex>(I))) = CostFeatures->at(n: I);
535 }
536 // This one would have been set up to be right at the end.
537 if (!InteractiveChannelBaseName.empty() && InteractiveIncludeDefault)
538 *ModelRunner->getTensor<int64_t>(FeatureID: getFeatureMap().size() - 1) =
539 GetDefaultAdvice(CB);
540 return getAdviceFromModel(CB, ORE);
541}
542
543std::unique_ptr<MLInlineAdvice>
544MLInlineAdvisor::getAdviceFromModel(CallBase &CB,
545 OptimizationRemarkEmitter &ORE) {
546 return std::make_unique<MLInlineAdvice>(
547 args: this, args&: CB, args&: ORE, args: static_cast<bool>(ModelRunner->evaluate<int64_t>()));
548}
549
550std::unique_ptr<InlineAdvice>
551MLInlineAdvisor::getSkipAdviceIfUnreachableCallsite(CallBase &CB) {
552 if (!FAM.getResult<DominatorTreeAnalysis>(IR&: *CB.getCaller())
553 .isReachableFromEntry(A: CB.getParent()))
554 return std::make_unique<InlineAdvice>(args: this, args&: CB, args&: getCallerORE(CB), args: false);
555 return nullptr;
556}
557
558std::unique_ptr<InlineAdvice> MLInlineAdvisor::getMandatoryAdvice(CallBase &CB,
559 bool Advice) {
560 // Make sure we track inlinings in all cases - mandatory or not.
561 if (auto Skip = getSkipAdviceIfUnreachableCallsite(CB))
562 return Skip;
563 if (Advice && !ForceStop)
564 return getMandatoryAdviceImpl(CB);
565
566 // If this is a "never inline" case, there won't be any changes to internal
567 // state we need to track, so we can just return the base InlineAdvice, which
568 // will do nothing interesting.
569 // Same if we are forced to stop - we don't track anymore.
570 return std::make_unique<InlineAdvice>(args: this, args&: CB, args&: getCallerORE(CB), args&: Advice);
571}
572
573std::unique_ptr<MLInlineAdvice>
574MLInlineAdvisor::getMandatoryAdviceImpl(CallBase &CB) {
575 return std::make_unique<MLInlineAdvice>(args: this, args&: CB, args&: getCallerORE(CB), args: true);
576}
577
578void MLInlineAdvisor::print(raw_ostream &OS) const {
579 OS << "[MLInlineAdvisor] Nodes: " << NodeCount << " Edges: " << EdgeCount
580 << " EdgesOfLastSeenNodes: " << EdgesOfLastSeenNodes << "\n";
581 OS << "[MLInlineAdvisor] FPI:\n";
582 for (auto I : FPICache) {
583 OS << I.first->getName() << ":\n";
584 I.second.print(OS);
585 OS << "\n";
586 }
587 OS << "\n";
588 OS << "[MLInlineAdvisor] FuncLevels:\n";
589 for (auto I : FunctionLevels)
590 OS << (DeadFunctions.contains(V: &I.first->getFunction())
591 ? "<deleted>"
592 : I.first->getFunction().getName())
593 << " : " << I.second << "\n";
594
595 OS << "\n";
596}
597
598MLInlineAdvice::MLInlineAdvice(MLInlineAdvisor *Advisor, CallBase &CB,
599 OptimizationRemarkEmitter &ORE,
600 bool Recommendation)
601 : InlineAdvice(Advisor, CB, ORE, Recommendation),
602 CallerIRSize(Advisor->isForcedToStop() ? 0 : Advisor->getIRSize(F&: *Caller)),
603 CalleeIRSize(Advisor->isForcedToStop() ? 0 : Advisor->getIRSize(F&: *Callee)),
604 CallerAndCalleeEdges(Advisor->isForcedToStop()
605 ? 0
606 : (Advisor->getLocalCalls(F&: *Caller) +
607 Advisor->getLocalCalls(F&: *Callee))),
608 PreInlineCallerFPI(Advisor->getCachedFPI(F&: *Caller)) {
609 if (Recommendation)
610 FPU.emplace(args&: Advisor->getCachedFPI(F&: *getCaller()), args&: CB);
611}
612
613void MLInlineAdvice::reportContextForRemark(
614 DiagnosticInfoOptimizationBase &OR) {
615 using namespace ore;
616 OR << NV("Callee", Callee->getName());
617 for (size_t I = 0; I < getAdvisor()->getFeatureMap().size(); ++I)
618 OR << NV(getAdvisor()->getFeatureMap()[I].name(),
619 *getAdvisor()->getModelRunner().getTensor<int64_t>(FeatureID: I));
620 OR << NV("ShouldInline", isInliningRecommended());
621}
622
623void MLInlineAdvice::updateCachedCallerFPI(FunctionAnalysisManager &FAM) const {
624 FPU->finish(FAM);
625}
626
627void MLInlineAdvice::recordInliningImpl() {
628 ORE.emit(RemarkBuilder: [&]() {
629 OptimizationRemark R(DEBUG_TYPE, "InliningSuccess", DLoc, Block);
630 reportContextForRemark(OR&: R);
631 return R;
632 });
633 getAdvisor()->onSuccessfulInlining(Advice: *this, /*CalleeWasDeleted*/ false);
634}
635
636void MLInlineAdvice::recordInliningWithCalleeDeletedImpl() {
637 ORE.emit(RemarkBuilder: [&]() {
638 OptimizationRemark R(DEBUG_TYPE, "InliningSuccessWithCalleeDeleted", DLoc,
639 Block);
640 reportContextForRemark(OR&: R);
641 return R;
642 });
643 getAdvisor()->onSuccessfulInlining(Advice: *this, /*CalleeWasDeleted*/ true);
644}
645
646void MLInlineAdvice::recordUnsuccessfulInliningImpl(
647 const InlineResult &Result) {
648 getAdvisor()->getCachedFPI(F&: *Caller) = PreInlineCallerFPI;
649 ORE.emit(RemarkBuilder: [&]() {
650 OptimizationRemarkMissed R(DEBUG_TYPE, "InliningAttemptedAndUnsuccessful",
651 DLoc, Block);
652 reportContextForRemark(OR&: R);
653 return R;
654 });
655}
656void MLInlineAdvice::recordUnattemptedInliningImpl() {
657 assert(!FPU);
658 ORE.emit(RemarkBuilder: [&]() {
659 OptimizationRemarkMissed R(DEBUG_TYPE, "IniningNotAttempted", DLoc, Block);
660 reportContextForRemark(OR&: R);
661 return R;
662 });
663}
664