Jlm
Loading...
Searching...
No Matches
ScalarEvolution.cpp
Go to the documentation of this file.
1/*
2 * Copyright 2025 Andreas Lilleby Hjulstad <andreas.lilleby.hjulstad@gmail.com>
3 * See COPYING for terms of redistribution.
4 */
5
11#include <jlm/llvm/ir/Trace.hpp>
14#include <jlm/rvsdg/theta.hpp>
17
18#include <algorithm>
19#include <cmath>
20#include <queue>
21
22namespace jlm::llvm
23{
24
26{
27public:
28 ~Context() = default;
29
30 Context() = default;
31
32 Context(const Context &) = delete;
33
34 Context(Context &&) = delete;
35
36 Context &
37 operator=(const Context &) = delete;
38
39 Context &
40 operator=(Context &&) = delete;
41
42 void
44 {
45 LoopVars_.insert(&var);
46 }
47
48 size_t
50 {
51 return LoopVars_.size();
52 }
53
54 static std::unique_ptr<Context>
56 {
57 return std::make_unique<Context>();
58 }
59
60 std::unique_ptr<SCEVChainRecurrence>
62 {
63 const auto it = ChrecMap_.find(&output);
64 if (it == ChrecMap_.end() || !it->second)
65 return nullptr;
66
67 return SCEV::CloneAs<SCEVChainRecurrence>(*it->second);
68 }
69
70 std::unique_ptr<SCEV>
72 {
73 const auto it = SCEVMap_.find(&output);
74 if (it == SCEVMap_.end() || !it->second)
75 return nullptr;
76
77 return it->second->Clone();
78 }
79
80 void
81 InsertChrec(rvsdg::Output & output, const std::unique_ptr<SCEVChainRecurrence> & chrec)
82 {
83 ChrecMap_.insert_or_assign(&output, SCEV::CloneAs<SCEVChainRecurrence>(*chrec));
84 }
85
86 const std::unordered_map<rvsdg::Output *, std::unique_ptr<SCEVChainRecurrence>> &
87 GetChrecMap() const noexcept
88 {
89 return ChrecMap_;
90 }
91
92 const std::unordered_map<rvsdg::Output *, std::unique_ptr<SCEV>> &
93 GetSCEVMap() const noexcept
94 {
95 return SCEVMap_;
96 }
97
98 int
100 {
101 int count = 0;
102 for (auto & [out, chrec] : ChrecMap_)
103 {
104 // Count induction variables (loop variables with a computed recurrence) with specific order
106 && out->Type()->Kind() != rvsdg::TypeKind::State)
107 {
108 if (chrec->GetOperands().size() == n + 1 && !IsUnknown(*chrec))
109 count++;
110 }
111 }
112 return count;
113 }
114
115 size_t
117 {
118 int count = 0;
119 for (auto & [out, chrec] : ChrecMap_)
120 {
122 && out->Type()->Kind() != rvsdg::TypeKind::State)
123 {
124 // Only count chrecs that are not unknown
125 if (!IsUnknown(*chrec))
126 count++;
127 }
128 }
129 return count;
130 }
131
132 void
133 InsertSCEV(rvsdg::Output & output, const std::unique_ptr<SCEV> & scev)
134 {
135 SCEVMap_.insert_or_assign(&output, scev->Clone());
136 }
137
138 void
140 {
141 NumLoops_++;
142 }
143
144 size_t
146 {
147 return NumLoops_;
148 }
149
150 void
151 SetTripCount(const rvsdg::ThetaNode & thetaNode, const size_t tripCount)
152 {
153 TripCountMap_.insert_or_assign(&thetaNode, tripCount);
154 }
155
156 size_t
157 GetTripCount(const rvsdg::ThetaNode & thetaNode) const
158 {
159 return TripCountMap_.at(&thetaNode);
160 }
161
162 const std::unordered_map<const rvsdg::ThetaNode *, size_t> &
163 GetTripCountMap() const noexcept
164 {
165 return TripCountMap_;
166 }
167
168private:
169 std::unordered_map<rvsdg::Output *, std::unique_ptr<SCEVChainRecurrence>> ChrecMap_;
170 std::unordered_map<rvsdg::Output *, std::unique_ptr<SCEV>> SCEVMap_;
171 std::unordered_map<const rvsdg::ThetaNode *, size_t> TripCountMap_;
172 std::unordered_set<const rvsdg::Output *> LoopVars_;
173
174 size_t NumLoops_ = 0;
175};
176
178{
179
180public:
181 ~Statistics() noexcept override = default;
182
183 explicit Statistics(const util::FilePath & sourceFile)
184 : util::Statistics(Id::ScalarEvolution, sourceFile)
185 {}
186
187 void
188 Start() noexcept
189 {
190 AddTimer(Label::Timer).start();
191 }
192
193 void
194 Stop(const Context & context) noexcept
195 {
196 GetTimer(Label::Timer).stop();
197 AddMeasurement(Label::NumTotalInductionVariables, context.GetNumTotalInductionVariables());
199 Label::NumConstantInductionVariables,
200 context.GetNumInductionVariablesWithOrder(0));
202 Label::NumFirstOrderInductionVariables,
203 context.GetNumInductionVariablesWithOrder(1));
205 Label::NumSecondOrderInductionVariables,
206 context.GetNumInductionVariablesWithOrder(2));
207 AddMeasurement(Label::NumLoopVariablesTotal, context.GetNumTotalLoopVars());
208 AddMeasurement(Label::NumLoops, context.GetNumLoops());
209 AddMeasurement(Label::TripCounts, GetTripCountString(context.GetTripCountMap()));
210 }
211
212 static std::string
213 GetTripCountString(const std::unordered_map<const rvsdg::ThetaNode *, size_t> & tripCountMap)
214 {
215 std::string s = "";
216 bool first = true;
217 for (auto & [thetaNode, tripCount] : tripCountMap)
218 {
219 if (!first)
220 s += ',';
221 first = false;
222
223 s += "ID(" + std::to_string(thetaNode->subregion()->getRegionId())
224 + ")=" + std::to_string(tripCount);
225 }
226 return s;
227 }
228
229 static std::unique_ptr<Statistics>
230 Create(const util::FilePath & sourceFile)
231 {
232 return std::make_unique<Statistics>(sourceFile);
233 }
234};
235
237 : rvsdg::Transformation("ScalarEvolution")
238{}
239
240ScalarEvolution::~ScalarEvolution() noexcept = default;
241
242std::unordered_map<const rvsdg::Output *, std::unique_ptr<SCEVChainRecurrence>>
243ScalarEvolution::GetChrecMap() const
244{
245 std::unordered_map<const rvsdg::Output *, std::unique_ptr<SCEVChainRecurrence>> mapCopy{};
246 for (auto & [output, chrec] : Context_->GetChrecMap())
247 {
248 mapCopy.emplace(output, SCEV::CloneAs<SCEVChainRecurrence>(*chrec));
249 }
250 return mapCopy;
251}
252
253std::unordered_map<const rvsdg::Output *, std::unique_ptr<SCEV>>
255{
256 std::unordered_map<const rvsdg::Output *, std::unique_ptr<SCEV>> mapCopy{};
257 for (auto & [output, scev] : Context_->GetSCEVMap())
258 {
259 mapCopy.emplace(output, scev->Clone());
260 }
261 return mapCopy;
262}
263
264std::unordered_map<const rvsdg::ThetaNode *, size_t>
266{
267 return Context_->GetTripCountMap();
268}
269
270void
272 rvsdg::RvsdgModule & rvsdgModule,
274{
275 auto statistics = Statistics::Create(rvsdgModule.SourceFilePath().value());
276 statistics->Start();
277
279 rvsdg::Region & rootRegion = rvsdgModule.Rvsdg().GetRootRegion();
280 AnalyzeRegion(rootRegion);
282
283 statistics->Stop(*Context_);
284 statisticsCollector.CollectDemandedStatistics(std::move(statistics));
285};
286
287void
289{
290 for (auto & node : region.Nodes())
291 {
292 if (auto structuralNode = dynamic_cast<rvsdg::StructuralNode *>(&node))
293 {
294 for (auto & subregion : structuralNode->Subregions())
295 {
296 AnalyzeRegion(subregion);
297 }
298 if (auto thetaNode = dynamic_cast<rvsdg::ThetaNode *>(structuralNode))
299 {
300 Context_->AddLoopToCount();
301 // Add number of loop vars in theta (for statistics)
302 for (const auto loopVar : thetaNode->GetLoopVars())
303 {
304 if (loopVar.pre->Type()->Kind() != rvsdg::TypeKind::State)
305 {
306 // Only add loop variables that are not states
307 Context_->AddLoopVar(*loopVar.pre);
308 }
309 }
310
311 PerformSCEVAnalysis(*thetaNode);
312
313 auto tripCount = GetPredictedTripCount(*thetaNode);
314
315 if (tripCount.has_value())
316 Context_->SetTripCount(*thetaNode, *tripCount);
317 }
318 }
319 }
320}
321
322bool
324{
325 if (auto constantStep = dynamic_cast<const SCEVConstant *>(&stepSCEV))
326 {
327 return constantStep->GetValue() < 0;
328 }
329 if (auto recurrenceStep = dynamic_cast<const SCEVChainRecurrence *>(&stepSCEV))
330 {
332
333 const auto start = dynamic_cast<const SCEVConstant *>(recurrenceStep->GetStartValue());
334 auto stepPtr = recurrenceStep->GetStep();
335 const auto step = dynamic_cast<const SCEVConstant *>(stepPtr->get());
336
337 if (!start || !step)
338 throw std::logic_error("Step can only contain constant SCEVs!");
339
340 const auto a = start->GetValue();
341 const auto b = step->GetValue();
342
343 return a <= 0 && b <= 0 && !(a == 0 && b == 0);
344 }
345 throw std::logic_error("Wrong type for step!");
346}
347
348bool
350{
351 if (auto constantStep = dynamic_cast<const SCEVConstant *>(&stepSCEV))
352 {
353 return constantStep->GetValue() > 0;
354 }
355 if (auto recurrenceStep = dynamic_cast<const SCEVChainRecurrence *>(&stepSCEV))
356 {
358
359 const auto start = dynamic_cast<const SCEVConstant *>(recurrenceStep->GetStartValue());
360 auto stepPtr = recurrenceStep->GetStep();
361 const auto step = dynamic_cast<const SCEVConstant *>(stepPtr->get());
362
363 if (!start || !step)
364 throw std::logic_error("Step can only contain constant SCEVs!");
365
366 const auto a = start->GetValue();
367 const auto b = step->GetValue();
368
369 return a >= 0 && b >= 0 && !(a == 0 && b == 0);
370 }
371 throw std::logic_error("Wrong type for step!");
372}
373
374bool
376{
377 if (auto constantStep = dynamic_cast<const SCEVConstant *>(&stepSCEV))
378 {
379 return constantStep->GetValue() == 0;
380 }
381 if (auto recurrenceStep = dynamic_cast<const SCEVChainRecurrence *>(&stepSCEV))
382 {
384
385 const auto start = dynamic_cast<const SCEVConstant *>(recurrenceStep->GetStartValue());
386 auto stepPtr = recurrenceStep->GetStep();
387 const auto step = dynamic_cast<const SCEVConstant *>(stepPtr->get());
388
389 if (!start || !step)
390 throw std::logic_error("Step can only contain constant SCEVs!");
391
392 const auto a = start->GetValue();
393 const auto b = step->GetValue();
394
395 return a == 0 && b == 0;
396 }
397 throw std::logic_error("Wrong type for step!");
398}
399
400std::optional<size_t>
402{
403 const auto pred = thetaNode.predicate();
404 const auto & [node, matchOperation] =
406 if (!matchOperation)
407 return std::nullopt;
408
409 JLM_ASSERT(node->ninputs() == 1); // Match node only has 1 input
410
411 const auto origin = node->input(0)->origin();
412 const auto comparisonNode = rvsdg::TryGetOwnerNode<rvsdg::SimpleNode>(*origin);
413 if (!comparisonNode)
414 return std::nullopt;
415
416 const auto * comparisonOperation = &comparisonNode->GetOperation();
417 if (!(rvsdg::is<IntegerSltOperation>(*comparisonOperation)
418 || rvsdg::is<IntegerSleOperation>(*comparisonOperation)
419 || rvsdg::is<IntegerUltOperation>(*comparisonOperation)
420 || rvsdg::is<IntegerUleOperation>(*comparisonOperation)
421 || rvsdg::is<IntegerSgtOperation>(*comparisonOperation)
422 || rvsdg::is<IntegerSgeOperation>(*comparisonOperation)
423 || rvsdg::is<IntegerUgtOperation>(*comparisonOperation)
424 || rvsdg::is<IntegerUgeOperation>(*comparisonOperation)
425 || rvsdg::is<IntegerNeOperation>(*comparisonOperation)
426 || rvsdg::is<IntegerEqOperation>(*comparisonOperation)))
427 return std::nullopt;
428
429 auto & lhs = *comparisonNode->input(0)->origin();
430 auto & rhs = *comparisonNode->input(1)->origin();
431 auto lhsChrec = Context_->TryGetChrecForOutput(lhs);
432 auto rhsChrec = Context_->TryGetChrecForOutput(rhs);
433
434 if (!lhsChrec)
435 lhsChrec = GetOrCreateChainRecurrence(lhs, *GetOrCreateSCEVForOutput(lhs), thetaNode);
436
437 if (!rhsChrec)
438 rhsChrec = GetOrCreateChainRecurrence(rhs, *GetOrCreateSCEVForOutput(rhs), thetaNode);
439
440 int64_t bound = 0;
441 std::unique_ptr<SCEVChainRecurrence> chrec{};
442
443 if (SCEVChainRecurrence::IsConstant(*lhsChrec))
444 {
445 const auto constantSCEV = dynamic_cast<SCEVConstant *>(lhsChrec->GetOperand(0));
446 if (!constantSCEV)
447 return std::nullopt;
448
449 bound = constantSCEV->GetValue();
450 chrec = SCEV::CloneAs<SCEVChainRecurrence>(*rhsChrec);
451 }
452 else if (SCEVChainRecurrence::IsConstant(*rhsChrec))
453 {
454 const auto constantSCEV = dynamic_cast<SCEVConstant *>(rhsChrec->GetOperand(0));
455 if (!constantSCEV)
456 return std::nullopt;
457
458 bound = constantSCEV->GetValue();
459 chrec = SCEV::CloneAs<SCEVChainRecurrence>(*lhsChrec);
460 }
461 else
462 {
463 // None of them are invariant, we can't reliably compute the backedge taken count
464 return std::nullopt;
465 }
466
468 {
469 // We can only compute the trip count reliably for affine and quadratic recurrences. In other
470 // cases, we return nullopt
471 return std::nullopt;
472 }
473
474 for (const auto op : chrec->GetOperands())
475 {
476 if (!dynamic_cast<const SCEVConstant *>(op))
477 {
478 // If any of the operands is not a constant, we cannot compute the trip count, and should
479 // return early
480 return std::nullopt;
481 }
482 }
483
484 const auto start = dynamic_cast<const SCEVConstant *>(chrec->GetStartValue())->GetValue();
485 const auto stepOpt = chrec->GetStep();
486 if (!stepOpt)
487 return std::nullopt;
488
489 const auto & stepSCEV = **stepOpt;
490
491 if (rvsdg::is<IntegerSltOperation>(*comparisonOperation)
492 || rvsdg::is<IntegerUltOperation>(*comparisonOperation))
493 {
494 // Trivial case (backedge is not taken and the only iteration is the first one)
495 if (start >= bound)
496 return 1;
497 if (start < bound && IsStepPositive(stepSCEV))
498 {
499 const auto backedgeTakenCount =
500 ComputeBackedgeTakenCountForChrec(*chrec, bound, comparisonOperation);
501 if (backedgeTakenCount.has_value())
502 {
503 // The trip count for a loop is the backedge taken count plus one
504 return *backedgeTakenCount + 1;
505 }
506 }
507 }
508 if (rvsdg::is<IntegerSleOperation>(*comparisonOperation)
509 || rvsdg::is<IntegerUleOperation>(*comparisonOperation))
510 {
511 if (start > bound)
512 return 1;
513 if (start <= bound && IsStepPositive(stepSCEV))
514 {
515 const auto backedgeTakenCount =
516 ComputeBackedgeTakenCountForChrec(*chrec, bound, comparisonOperation);
517 if (backedgeTakenCount.has_value())
518 {
519 return *backedgeTakenCount + 1;
520 }
521 }
522 }
523 if (rvsdg::is<IntegerSgtOperation>(*comparisonOperation)
524 || rvsdg::is<IntegerUgtOperation>(*comparisonOperation))
525 {
526 if (start <= bound)
527 return 1;
528 if (start > bound && IsStepNegative(stepSCEV))
529 {
530 const auto backedgeTakenCount =
531 ComputeBackedgeTakenCountForChrec(*chrec, bound, comparisonOperation);
532 if (backedgeTakenCount.has_value())
533 {
534 return *backedgeTakenCount + 1;
535 }
536 }
537 }
538 if (rvsdg::is<IntegerSgeOperation>(*comparisonOperation)
539 || rvsdg::is<IntegerUgeOperation>(*comparisonOperation))
540 {
541 if (start < bound)
542 return 1;
543 if (start >= bound && IsStepNegative(stepSCEV))
544 {
545 const auto backedgeTakenCount =
546 ComputeBackedgeTakenCountForChrec(*chrec, bound, comparisonOperation);
547 if (backedgeTakenCount.has_value())
548 {
549 return *backedgeTakenCount + 1;
550 }
551 }
552 }
553
554 if (rvsdg::is<IntegerNeOperation>(*comparisonOperation))
555 {
557 {
558 // With Ne and Eq comparisons, we only compute non-trivial backedge counts for affine
559 // recurrences as there is no general way to compute it for quadratic recurrences.
560 const auto step = dynamic_cast<const SCEVConstant *>(&stepSCEV)->GetValue();
561 if (IsStepPositive(stepSCEV))
562 {
563 const auto backedgeTakenCount =
564 ComputeBackedgeTakenCountForChrec(*chrec, bound, comparisonOperation);
565 // We need to make sure that it does not pass the bound value (results infinite loop)
566 if (start <= bound && (bound - start) % step == 0)
567 return *backedgeTakenCount + 1;
568 }
569 if (IsStepNegative(stepSCEV))
570 {
571 const auto backedgeTakenCount =
572 ComputeBackedgeTakenCountForChrec(*chrec, bound, comparisonOperation);
573 if (start >= bound && (bound - start) % step == 0)
574 return *backedgeTakenCount + 1;
575 }
576 }
577 if (start == bound)
578 return 1;
579 }
580
581 if (rvsdg::is<IntegerEqOperation>(*comparisonOperation))
582 {
583 if (start == bound)
584 {
585 if (!IsStepZero(stepSCEV))
586 return 2; // Backedge taken once
587 }
588 else
589 return 1;
590 }
591
593 {
594 // For quadratic recurrences, if the step is neither positive, negative or zero, we are not able
595 // to accurately compute the trip count.
596 if (!(IsStepPositive(stepSCEV) || IsStepNegative(stepSCEV) || IsStepZero(stepSCEV)))
597 {
598 return std::nullopt;
599 }
600 }
601
602 // If we have not returned a value at this point, we have an infinite loop.
603 return std::nullopt;
604}
605
606std::optional<size_t>
608 const SCEVChainRecurrence & chrec,
609 const int64_t bound,
610 const rvsdg::SimpleOperation * comparisonOperation)
611{
612 const auto start = dynamic_cast<const SCEVConstant *>(chrec.GetStartValue())->GetValue();
613 const auto stepOpt = chrec.GetStep();
614 if (!stepOpt)
615 return std::nullopt;
616
617 const auto & stepSCEV = *stepOpt;
618
619 bool isEqualsComparison = rvsdg::is<IntegerSleOperation>(*comparisonOperation)
620 || rvsdg::is<IntegerUleOperation>(*comparisonOperation)
621 || rvsdg::is<IntegerSgeOperation>(*comparisonOperation)
622 || rvsdg::is<IntegerUgeOperation>(*comparisonOperation);
623
624 // Check the size of the step recurrence: 1 -> Affine, 2 -> Quadratic
625 // We can only compute the backedge taken count for these two cases
627 {
628 const auto stepConstant = dynamic_cast<const SCEVConstant *>(stepSCEV.get());
629 const auto step = stepConstant->GetValue();
630
631 // f(i) = a + b * i
632 // f(i) = k => a + b * i = k => i = (k - a)/b
633 size_t result = std::ceil(static_cast<double>(bound - start) / step);
634
635 if (isEqualsComparison)
636 {
637 // If we have an equals comparison and the value of the difference between the bound and the
638 // start is a whole multiple of the step size, we get another backedge taken
639 if ((bound - start) % step == 0)
640 result += 1;
641 }
642 return result;
643 }
645 {
646 // Create a quadratic equation for the recurrence {a,+,b,+,c}
647 // The start value is a, and the increments are b, b+c, b+2c, ..., so the accumulated values are
648 // a+b, (a+b)+(b+c), (a+b)+(b+c)+(b+2c), ..., that is,
649 // a+b, a+2b+c, a+3b+3c, ...
650 // After i iterations the value is a + ib + i(i-1)/2 c = f(i).
651 const auto stepRecurrence = dynamic_cast<const SCEVChainRecurrence *>(stepSCEV.get());
652 const int64_t stepFirst =
653 dynamic_cast<const SCEVConstant *>(stepRecurrence->GetStartValue())->GetValue();
654
655 const int64_t stepSecond =
656 dynamic_cast<const SCEVConstant *>(stepRecurrence->GetStep()->get())->GetValue();
657
658 // Let f(i) = a + ib + i(i-1)/2 c
659 //
660 // We want to find out when this polynomial is equal to the compare value, i.e. f(i) = k.
661 // This is equivalent with the expression f(i) - k "switching sign" from positive to negative.
662 // Conversely, this is also when the predicate condition will no longer hold.
663 //
664 // The equation f(i) - k = 0 is written as:
665 // a + ib + i(i-1)/2 c - k = 0, or 2(a-k) + 2b i + i(i-1) c = 0.
666 // In a quadratic form it becomes:
667 // c i^2 + (2b - c) i + 2(a - k) = 0.
668 //
669 // We use the quadratic formula to solve this.
670
671 const int64_t a = stepSecond;
672 const int64_t b = 2 * stepFirst - stepSecond;
673 const int64_t c = 2 * (start - bound);
674
675 const auto quadraticResult = SolveQuadraticEquation(a, b, c);
676 if (!quadraticResult.has_value())
677 return std::nullopt;
678
679 size_t result = *quadraticResult;
680
681 if (isEqualsComparison)
682 {
683 // Same as for affine, but instead of checking using modulo, we evaluate the value at the
684 // result and check
685 const int64_t valueAtResult =
686 start + result * stepFirst + result * (result - 1) / 2 * stepSecond;
687 if (valueAtResult == bound)
688 result += 1;
689 }
690 return result;
691 }
692 return std::nullopt;
693}
694
695std::optional<size_t>
696ScalarEvolution::SolveQuadraticEquation(int64_t a, int64_t b, int64_t c)
697{
698 // If a is negative, negate all the coefficients to simplify the math
699 if (a < 0)
700 {
701 a = -a;
702 b = -b;
703 c = -c;
704 }
705
706 const auto d = b * b - 4 * a * c; // Discriminant
707
708 if (d < 0)
709 return std::nullopt;
710
711 // Integer square root of the discriminant
712 int64_t sq = std::floor(std::sqrt(d));
713
714 // Check if square root is exact
715 const bool inexactSq = (sq * sq != d);
716
717 // Adjust if sq^2 > discriminant (shouldn't happen with floor, but just to be safe)
718 if (sq * sq > d)
719 sq -= 1;
720
721 int64_t x = 0;
722 int64_t rem = 0;
723
724 // The vertex (min/max value) of the parabola f(x) = Ax^2 + Bx + C is at -B/2A. Since A > 0, the
725 // vertex is at a non-positive x location iff B >= 0. In that case the first zero crossing is the
726 // greater root. If B < 0, the vertex is at a positive x location, meaning both roots are positive
727 // and the smaller root is the first crossing.
728 if (b < 0)
729 {
730 // The square root is rounded down, so the roots may be inexact. When using the quadratic
731 // formula, the low root could be greater than the exact one. To make sure this does not happen,
732 // we add 1 if the root is inexact when calculating the low root.
733 x = (-b - (sq + (inexactSq ? 1 : 0))) / (2 * a);
734 rem = (-b - sq) % (2 * a);
735 }
736 else
737 {
738 x = (-b + sq) / (2 * a);
739 rem = (-b + sq) % (2 * a);
740 }
741
742 // Result should be non-negative
743 if (x < 0)
744 x = 0;
745
746 // Check for exact solution
747 if (!inexactSq && rem == 0)
748 {
749 return x;
750 }
751
752 // The exact value of the square root should be between sq and sq + 1
753 // Check for sign change between f(x) and f(x+1)
754 const int64_t valueAtX = (a * x + b) * x + c;
755 const int64_t valueAtXPlusOne = (a * (x + 1) + b) * (x + 1) + c;
756
757 const bool signChange =
758 ((valueAtX < 0) != (valueAtXPlusOne < 0)) || ((valueAtX == 0) != (valueAtXPlusOne == 0));
759 // Sign did not change, not a valid solution
760 if (!signChange)
761 return std::nullopt;
762
763 x += 1;
764 return x;
765}
766
767void
769{
770 bool changed{};
771 do
772 {
773 changed = false;
774
775 std::vector<std::pair<rvsdg::Output *, std::unique_ptr<SCEV>>> pending;
776 for (auto & [output, chrec] : Context_->GetChrecMap())
777 {
778 if (auto newSCEV = TryReplaceInitForSCEV(*chrec, *output))
779 {
780 pending.emplace_back(output, std::move(*newSCEV));
781 changed = true;
782 }
783 }
784
785 for (auto & [output, scev] : pending)
786 {
787 // Check if the result is actually a chrec
788 if (auto * chrec = dynamic_cast<SCEVChainRecurrence *>(scev.get()))
789 {
790 Context_->InsertChrec(*output, SCEV::CloneAs<SCEVChainRecurrence>(*chrec));
791 }
792 else
793 {
794 // The transformation produced a non-chrec SCEV (n-ary expression), store it in the SCEV
795 // map instead
796 Context_->InsertSCEV(*output, std::move(scev));
797 }
798 }
799 } while (changed);
800}
801
802std::optional<std::unique_ptr<SCEV>>
804{
805 if (const auto initSCEV = dynamic_cast<const SCEVInit *>(&scev))
806 {
807 // Found an Init node, find the origin of its input value and get or create its chain
808 // recurrence
809 const auto & initPrePointer = initSCEV->GetPrePointer();
810 if (const auto innerTheta = rvsdg::TryGetRegionParentNode<rvsdg::ThetaNode>(initPrePointer))
811 {
812 const auto correspondingInput = innerTheta->MapPreLoopVar(initPrePointer).input;
813 auto & inputOrigin = rvsdg::traceOutput(*correspondingInput->origin(), false);
814 if (const auto originSCEV = Context_->TryGetSCEVForOutput(inputOrigin))
815 {
816 // We have found a SCEV for the origin of the input, find the corresponding theta node so
817 // we can create a recurrence for it
818 const auto thetaParent = rvsdg::TryGetOwnerNode<rvsdg::ThetaNode>(inputOrigin);
819 const auto outerTheta =
820 thetaParent ? thetaParent
821 : util::assertedCast<rvsdg::ThetaNode>(inputOrigin.region()->node());
822
823 const auto chrec = GetOrCreateChainRecurrence(inputOrigin, *originSCEV, *outerTheta);
824
825 // Create a chain recurrence for the SCEV, with the outer theta as the loop
826 return chrec->Clone();
827 }
828 }
829 }
830 if (const auto nArySCEV = dynamic_cast<const SCEVNAryExpr *>(&scev))
831 {
832 // An n-ary scev is any scev with an arbitrary number of operands: chain recurrence, n-ary add
833 // and n-ary mult. We want to recursively check all it's operands for Init nodes
834 auto clone = SCEV::CloneAs<SCEVNAryExpr>(*nArySCEV);
835 const auto operands = nArySCEV->GetOperands();
836 bool changed = false;
837 for (size_t i = 0; i < operands.size(); ++i)
838 {
839 if (auto result = TryReplaceInitForSCEV(*operands[i], output))
840 {
841 if (*result)
842 {
843 // Replace the Init operand with the chrec
844 changed = true;
845 clone->ReplaceOperand(i, std::move(*result));
846 }
847 }
848 }
849 if (!changed)
850 return std::nullopt;
851
852 if (dynamic_cast<const SCEVChainRecurrence *>(&scev))
853 {
854 // Result is a new chain recurrence, return it
855 return clone;
856 }
857 // If it is an n-ary expression (Add or Mul), we try to fold the operands into themselves,
858 // e.g. if, after replacing Init nodes with recurrences, we have ({0,+,1} + {1,+,2}) in an
859 // n-ary add expression, we can fold this into {1,+,3}.
860 return FoldNAryExpression(*clone, output);
861 }
862 // Default is to just return nothing
863 return std::nullopt;
864}
865
866void
868{
869 for (const auto loopVar : thetaNode.GetLoopVars())
870 {
871 // In some cases (e.g. with store operations), we still want to create a SCEV tree for the loop
872 // variable even though it is a state variable. However, we still want to filter out state
873 // variables that are purely for scaffolding as they are uninteresting for the analysis.
874 if (loopVar.pre->Type()->Kind() == rvsdg::TypeKind::State
876 {
877 continue;
878 }
879 const auto post = loopVar.post;
880 // We compute the SCEV for each loop variable in a recursive bottom up fashion,
881 // starting at the post's origin
882 auto scev = GetOrCreateSCEVForOutput(*post->origin());
883 Context_->InsertSCEV(*loopVar.output, scev); // Save the SCEV at the theta outputs as well
884 }
885
886 auto dependencyGraph = CreateDependencyGraph(thetaNode);
887
889 for (const auto & [output, deps] : dependencyGraph)
890 {
891 if (CanCreateChainRecurrence(*output, dependencyGraph))
892 validOutputs.insert(output);
893 }
894
895 // Filter the dependency graph to only contain the outputs of the SCEVs that are valid chain
896 // recurrences and update dependencies accordingly
897 auto filteredDependencyGraph = dependencyGraph;
898 for (auto it = filteredDependencyGraph.begin(); it != filteredDependencyGraph.end();)
899 {
900 if (!validOutputs.Contains(it->first))
901 {
902 for (auto & [node, deps] : filteredDependencyGraph)
903 deps.erase(it->first);
904 it = filteredDependencyGraph.erase(it);
905 }
906 else
907 ++it;
908 }
909
910 const auto order = TopologicalSort(filteredDependencyGraph);
911
912 for (auto output : order)
913 {
914 std::unique_ptr<SCEV> scev{};
915 if (const auto theta = rvsdg::TryGetRegionParentNode<rvsdg::ThetaNode>(*output);
916 &thetaNode == theta)
917 {
918 // For loop variables, we need to retrieve and use the SCEV saved at the post's origin,
919 // equivalent to a "backedge" which describes how the value at the pre pointer is updated
920 auto & newOutput = *thetaNode.MapPreLoopVar(*output).post->origin();
921 scev = Context_->TryGetSCEVForOutput(newOutput);
922 }
923 else
924 scev = Context_->TryGetSCEVForOutput(*output);
925
926 JLM_ASSERT(scev);
927
928 auto chrec = GetOrCreateChainRecurrence(*output, *scev, thetaNode);
929 Context_->InsertChrec(*output, chrec);
930 }
931
932 for (auto & [output, scev] : Context_->GetSCEVMap())
933 {
934 if (std::find(order.begin(), order.end(), output) == order.end())
935 {
936 auto unknownChainRecurrence =
938 Context_->InsertChrec(*output, unknownChainRecurrence);
939 }
940 }
941}
942
943std::unique_ptr<SCEV>
945{
946 if (const auto existing = Context_->TryGetSCEVForOutput(output))
947 return existing->Clone();
948
949 std::unique_ptr<SCEV> result{};
951 {
952 // We know this is a loop variable, create a placeholder SCEV for now, and compute the
953 // expression later
954 result = SCEVPlaceholder::Create(output);
955 }
956
957 const auto & [simpleNode, simpleOperation] =
959
960 if (simpleNode)
961 {
962 if (rvsdg::is<MemoryHoistBarrierOperation>(*simpleOperation))
963 {
964 const auto barredInputOrigin =
966 result = GetOrCreateSCEVForOutput(*barredInputOrigin);
967 }
968 else if (
969 rvsdg::is<SExtOperation>(*simpleOperation) || rvsdg::is<ZExtOperation>(*simpleOperation))
970 {
971 JLM_ASSERT(simpleNode->ninputs() == 1);
972 result = GetOrCreateSCEVForOutput(*simpleNode->input(0)->origin());
973 }
974 else if (const auto gepOp = dynamic_cast<const GetElementPtrOperation *>(&*simpleOperation))
975 {
976 JLM_ASSERT(simpleNode->ninputs() >= 2);
977 const auto baseIndex = simpleNode->input(0)->origin();
978 JLM_ASSERT(is<PointerType>(baseIndex->Type()));
979
980 const auto & pointeeType = gepOp->getPointeeType();
981
982 auto baseScev = GetOrCreateSCEVForOutput(*baseIndex);
983
984 auto wholeTypeIndex = GetOrCreateSCEVForOutput(*simpleNode->input(1)->origin());
985 const auto wholeTypeSize = GetTypeAllocSize(*pointeeType);
986
987 std::unique_ptr<SCEV> offset =
988 SCEVMulExpr::Create(std::move(wholeTypeIndex), SCEVConstant::Create(wholeTypeSize));
989 if (auto innerOffset = ComputeSCEVForGepInnerOffset(*simpleNode, 2, *pointeeType))
990 offset = SCEVAddExpr::Create(std::move(offset), std::move(innerOffset));
991
992 result = SCEVAddExpr::Create(std::move(baseScev), std::move(offset));
993 }
994 else if (const auto constOp = dynamic_cast<const IntegerConstantOperation *>(&*simpleOperation))
995 {
996 const auto value = constOp->Representation().to_int();
997 result = SCEVConstant::Create(value);
998 }
999 else if (rvsdg::is<IntegerBinaryOperation>(*simpleOperation))
1000 {
1001 JLM_ASSERT(simpleNode->ninputs() == 2);
1002 const auto lhs = simpleNode->input(0)->origin();
1003 const auto rhs = simpleNode->input(1)->origin();
1004
1005 auto lhsScev = GetOrCreateSCEVForOutput(*lhs);
1006 auto rhsScev = GetOrCreateSCEVForOutput(*rhs);
1007 if (rvsdg::is<IntegerAddOperation>(*simpleOperation))
1008 {
1009 result = SCEVAddExpr::Create(std::move(lhsScev), std::move(rhsScev));
1010 }
1011 else if (rvsdg::is<IntegerSubOperation>(*simpleOperation))
1012 {
1013 auto rhsNegativeScev = GetNegativeSCEV(*rhsScev);
1014
1015 result = SCEVAddExpr::Create(std::move(lhsScev), std::move(rhsNegativeScev));
1016 }
1017 else if (rvsdg::is<IntegerMulOperation>(*simpleOperation))
1018 {
1019 result = SCEVMulExpr::Create(std::move(lhsScev), std::move(rhsScev));
1020 }
1021 else if (rvsdg::is<IntegerShlOperation>(*simpleOperation))
1022 {
1023 if (const auto * rhsConst = dynamic_cast<SCEVConstant *>(rhsScev.get()))
1024 {
1025 const auto shiftAmount = rhsConst->GetValue();
1026 auto factor = SCEVConstant::Create(1ULL << shiftAmount);
1027 result = SCEVMulExpr::Create(std::move(lhsScev), std::move(factor));
1028 }
1029 }
1030 }
1031 else
1032 {
1033 // Unknown operation, we traverse through to it's inputs
1034 for (auto & input : simpleNode->Inputs())
1035 {
1036 GetOrCreateSCEVForOutput(*input.origin());
1037 }
1038 }
1039 }
1040
1041 if (!result)
1042 // If none of the cases match, return an unknown SCEV expression
1043 result = SCEVUnknown::Create();
1044
1045 // Save the result in the cache
1046 Context_->InsertSCEV(output, result);
1047
1048 return result;
1049}
1050
1051std::unique_ptr<SCEV>
1053 const rvsdg::SimpleNode & gepNode,
1054 const size_t inputIndex,
1055 const rvsdg::Type & type)
1056{
1057 JLM_ASSERT(inputIndex >= 2);
1058
1059 if (inputIndex >= gepNode.ninputs())
1060 {
1061 return nullptr;
1062 }
1063
1064 const auto gepInput = gepNode.input(inputIndex);
1065 if (const auto arrayType = dynamic_cast<const ArrayType *>(&type))
1066 {
1067 const auto & elementType = *arrayType->GetElementType();
1068 auto offset = SCEVMulExpr::Create(
1069 GetOrCreateSCEVForOutput(*gepInput->origin()),
1071
1072 auto subOffset = ComputeSCEVForGepInnerOffset(gepNode, inputIndex + 1, elementType);
1073
1074 if (!subOffset)
1075 return offset;
1076
1077 return SCEVAddExpr::Create(std::move(offset), std::move(subOffset));
1078 }
1079 if (const auto structType = dynamic_cast<const StructType *>(&type))
1080 {
1081 const auto indexingValue = tryGetConstantSignedInteger(*gepInput->origin());
1082
1083 if (!indexingValue.has_value())
1084 return nullptr;
1085
1086 const auto & fieldType = structType->getElementType(*indexingValue);
1087
1088 auto offset = SCEVConstant::Create(structType->GetFieldOffset(*indexingValue));
1089
1090 auto subOffset = ComputeSCEVForGepInnerOffset(gepNode, inputIndex + 1, *fieldType);
1091
1092 if (!subOffset)
1093 return offset;
1094
1095 return SCEVAddExpr::Create(std::move(offset), std::move(subOffset));
1096 }
1097 throw std::logic_error("Unknown GEP type!");
1098}
1099
1100void
1102 const SCEV & scev,
1103 DependencyMap & dependencies,
1104 const DependencyOp op = DependencyOp::None)
1105{
1106 if (const auto placeholderSCEV = dynamic_cast<const SCEVPlaceholder *>(&scev))
1107 {
1108 auto & dependency = placeholderSCEV->GetPrePointer();
1109 // Retrieves dependency info struct from the map
1110 // In the case where the dependency does not already exist, a new struct is created with the
1111 // default count being 0 and the default operation being None
1112 auto & depInfo = dependencies[&dependency];
1113 depInfo.operation = op;
1114 depInfo.count++;
1115 }
1116
1117 if (const auto addSCEV = dynamic_cast<const SCEVAddExpr *>(&scev))
1118 {
1119 FindDependenciesForSCEV(*addSCEV->GetLeftOperand(), dependencies, DependencyOp::Add);
1120 FindDependenciesForSCEV(*addSCEV->GetRightOperand(), dependencies, DependencyOp::Add);
1121 }
1122
1123 if (const auto mulSCEV = dynamic_cast<const SCEVMulExpr *>(&scev))
1124 {
1125 FindDependenciesForSCEV(*mulSCEV->GetLeftOperand(), dependencies, DependencyOp::Mul);
1126 FindDependenciesForSCEV(*mulSCEV->GetRightOperand(), dependencies, DependencyOp::Mul);
1127 }
1128}
1129
1132{
1133 DependencyGraph graph{};
1134
1135 for (const auto & [output, scev] : Context_->GetSCEVMap())
1136 {
1137 DependencyMap dependencies{};
1138 if (const auto theta = rvsdg::TryGetRegionParentNode<rvsdg::ThetaNode>(*output);
1139 theta == &thetaNode)
1140 {
1141 // We know this is a pre pointer, so we map it to loop var and use the SCEV for the
1142 // post's origin (backedge) instead
1143 const auto loopVar = theta->MapPreLoopVar(*output);
1144 auto newScev = Context_->TryGetSCEVForOutput(*loopVar.post->origin());
1145
1146 FindDependenciesForSCEV(*newScev.get(), dependencies);
1147 }
1148 else
1149 FindDependenciesForSCEV(*scev.get(), dependencies);
1150
1151 graph[output] = dependencies;
1152 }
1153 return graph;
1154}
1155
1156// Implementation of Kahn's algorithm for topological sort
1157std::vector<rvsdg::Output *>
1159{
1160 const size_t numVertices = dependencyGraph.size();
1161 std::unordered_map<const rvsdg::Output *, int> indegree(numVertices);
1162 std::queue<rvsdg::Output *> q{};
1163 for (auto & [node, deps] : dependencyGraph)
1164 {
1165 for (auto & dep : deps)
1166 {
1167 if (const auto ptr = dep.first; ptr == node)
1168 continue; // Ignore self-edges
1169 // To begin with, the indegree is just the number of incoming edges
1170 indegree[node] += 1;
1171 }
1172 if (indegree[node] == 0)
1173 {
1174 // Add nodes with no incoming edges to the queue, we know that these have no dependencies
1175 q.push(node);
1176 }
1177 }
1178
1179 std::vector<rvsdg::Output *> result{};
1180 while (!q.empty())
1181 {
1182 rvsdg::Output * currentNode = q.front();
1183 q.pop();
1184 result.push_back(currentNode);
1185
1186 for (const auto & [node, deps] : dependencyGraph)
1187 {
1188 if (node == currentNode)
1189 continue;
1190
1191 for (const auto & dep : deps)
1192 {
1193 const auto ptr = dep.first;
1194 if (ptr == node)
1195 continue; // Skip self-edges
1196 if (ptr == currentNode)
1197 {
1198 // Update the indegree of nodes depending on this one
1199 indegree[node] -= 1;
1200 if (indegree[node] == 0)
1201 q.push(node);
1202 }
1203 }
1204 }
1205 }
1206 JLM_ASSERT(result.size() == numVertices);
1207 return result;
1208}
1209
1210std::unique_ptr<SCEVChainRecurrence>
1212 rvsdg::Output & output,
1213 const SCEV & scev,
1214 rvsdg::ThetaNode & thetaNode)
1215{
1216 if (const auto existing = Context_->TryGetChrecForOutput(output))
1217 {
1218 return SCEV::CloneAs<SCEVChainRecurrence>(*existing);
1219 }
1220
1221 auto stepRecurrence = GetOrCreateStepForSCEV(output, scev, thetaNode);
1222
1223 if (const auto theta = rvsdg::TryGetRegionParentNode<rvsdg::ThetaNode>(output);
1224 theta == &thetaNode)
1225 {
1226 // Find the start value for the recurrence
1227 const auto inputOrigin = thetaNode.MapPreLoopVar(output).input->origin();
1228 if (const auto constantInteger = tryGetConstantSignedInteger(*inputOrigin))
1229 {
1230 // If the input value is a constant, create a SCEV representation and set it as start
1231 // value (first operand in rec)
1232 stepRecurrence->AddOperandToFront(SCEVConstant::Create(*constantInteger));
1233 }
1234 else
1235 {
1236 // If not, create a SCEVInit node representing the start value
1237 stepRecurrence->AddOperandToFront(SCEVInit::Create(output));
1238 }
1239 }
1240 return stepRecurrence;
1241}
1242
1243std::unique_ptr<SCEVChainRecurrence>
1245 rvsdg::Output & output,
1246 const SCEV & scevTree,
1247 rvsdg::ThetaNode & thetaNode)
1248{
1249 if (const auto scevConstant = dynamic_cast<const SCEVConstant *>(&scevTree))
1250 {
1251 // This is a constant, we add it as the only operand
1252 return SCEVChainRecurrence::Create(thetaNode, output, scevConstant->Clone());
1253 }
1254 if (const auto scevPlaceholder = dynamic_cast<const SCEVPlaceholder *>(&scevTree))
1255 {
1256 if (&scevPlaceholder->GetPrePointer() == &output)
1257 {
1258 // Since we are only interested in the step value, and not the initial value, we can ignore
1259 // ourselves by returning an empty chain recurrence (treated as the identity element - 0 for
1260 // addition and 1 for multiplication)
1261 return SCEVChainRecurrence::Create(thetaNode, output);
1262 }
1263 if (auto storedRec = Context_->TryGetChrecForOutput(scevPlaceholder->GetPrePointer()))
1264 {
1265 // We have a dependency of another IV
1266 // Get it's saved value. This is safe to do due to the topological ordering
1267 return storedRec;
1268 }
1269 return SCEVChainRecurrence::Create(thetaNode, output, SCEVUnknown::Create());
1270 }
1271 if (const auto scevAddExpr = dynamic_cast<const SCEVAddExpr *>(&scevTree))
1272 {
1273 const auto lhsStep = GetOrCreateStepForSCEV(output, *scevAddExpr->GetLeftOperand(), thetaNode);
1274 const auto rhsStep = GetOrCreateStepForSCEV(output, *scevAddExpr->GetRightOperand(), thetaNode);
1275
1276 return SCEV::CloneAs<SCEVChainRecurrence>(
1277 *ApplyAddFolding(lhsStep.get(), rhsStep.get(), output));
1278 }
1279 if (const auto scevMulExpr = dynamic_cast<const SCEVMulExpr *>(&scevTree))
1280 {
1281 const auto lhsStep = GetOrCreateStepForSCEV(output, *scevMulExpr->GetLeftOperand(), thetaNode);
1282 const auto rhsStep = GetOrCreateStepForSCEV(output, *scevMulExpr->GetRightOperand(), thetaNode);
1283
1284 return SCEV::CloneAs<SCEVChainRecurrence>(
1285 *ApplyMulFolding(lhsStep.get(), rhsStep.get(), output));
1286 }
1287 return SCEVChainRecurrence::Create(thetaNode, output, SCEVUnknown::Create());
1288}
1289
1290std::unique_ptr<SCEV>
1292{
1293 // In some cases, we end up with an n-ary expression like (1 + Init(a1) + 2).
1294 // This method folds the constant operands, turning it into (3 + Init(a1)).
1295 bool folded{};
1296 do
1297 {
1298 folded = false;
1299 for (size_t i = 0; i < expression.NumOperands(); ++i)
1300 {
1301 std::vector<SCEV *> ops = expression.GetOperands();
1302 if (dynamic_cast<const SCEVInit *>(ops[i]))
1303 continue; // Cannot fold init
1304 for (size_t j = i + 1; j < expression.NumOperands(); ++j)
1305 {
1306 if (dynamic_cast<const SCEVInit *>(ops[j]))
1307 continue;
1308
1309 // Both are foldable (constants or recurrences) fold them according to the rules
1310 std::unique_ptr<SCEV> foldedOperand{};
1311 if (dynamic_cast<SCEVNAryAddExpr *>(&expression))
1312 {
1313 foldedOperand = ApplyAddFolding(ops[i], ops[j], output);
1314 }
1315 else if (dynamic_cast<SCEVNAryMulExpr *>(&expression))
1316 {
1317 foldedOperand = ApplyMulFolding(ops[i], ops[j], output);
1318 }
1319 else
1320 {
1321 throw std::logic_error("Invalid n-ary SCEV expression type in FoldNAryExpression!");
1322 }
1323 expression.RemoveOperand(j);
1324 expression.ReplaceOperand(i, foldedOperand);
1325 folded = true;
1326 break;
1327 }
1328 if (folded)
1329 break;
1330 }
1331 } while (folded);
1332
1333 if (expression.NumOperands() == 1)
1334 {
1335 // If there is only one operand in the n-ary expression, we just return the operand
1336 return expression.GetOperand(0)->Clone();
1337 }
1338
1339 return expression.Clone();
1340}
1341
1342std::unique_ptr<SCEV>
1343ScalarEvolution::ApplyAddFolding(SCEV * lhsOperand, SCEV * rhsOperand, rvsdg::Output & output)
1344{
1345 // We have the following folding rules from the CR algebra:
1346 // G + {e,+,f} => {G + e,+,f} (1)
1347 // {e,+,f} + {g,+,h} => {e + g,+,f + h} (2)
1348 //
1349 // And by generalizing rule 2, we have that:
1350 // {G,+,0} + {e,+,f} = {G + e,+,0 + f} = {G + e,+,f}
1351 //
1352 // Since we represent constants in the SCEVTree as recurrences consisting of only a SCEVConstant
1353 // node, we can therefore pad the constant recurrence with however many zeroes we need for the
1354 // length of the other recurrence. This effectively lets us apply both rules in one go.
1355 //
1356 // For constants and unknowns this is trivial, however it becomes a bit complicated when we
1357 // factor in SCEVInit nodes. These nodes represent the initial value of an IV in the case where
1358 // the exact value is unknown at compile time. E.g. function argument or result from a
1359 // call-instruction. In the cases where we have to fold one or more of these init-nodes, we
1360 // create an n-ary add expression (add expression with an arbitrary number of operands), and add
1361 // this to the chrec. Folding two of these n-ary add expressions will result in another n-ary
1362 // add expression, which consists of all the operands in both the left and the right expression.
1363
1364 // The if-chain below goes through each of the possible combinations of lhs and rhs values
1365 if (const auto *lhsUnknown = dynamic_cast<const SCEVUnknown *>(lhsOperand),
1366 *rhsUnknown = dynamic_cast<const SCEVUnknown *>(rhsOperand);
1367 lhsUnknown || rhsUnknown)
1368 {
1369 // If one of the sides is unknown. Return unknown
1370 return SCEVUnknown::Create();
1371 }
1372
1373 auto lhsChrec = dynamic_cast<SCEVChainRecurrence *>(lhsOperand);
1374 auto rhsChrec = dynamic_cast<SCEVChainRecurrence *>(rhsOperand);
1375 if (lhsChrec && rhsChrec)
1376 {
1377 if (&lhsChrec->GetLoop() != &rhsChrec->GetLoop())
1378 {
1380 lhsChrec->GetLoop(),
1381 output,
1382 SCEVNAryAddExpr::Create(lhsChrec->Clone(), rhsChrec->Clone()));
1383 }
1384
1385 auto newChrec = SCEVChainRecurrence::Create(lhsChrec->GetLoop(), output);
1386 const auto lhsSize = lhsChrec->NumOperands();
1387 const auto rhsSize = rhsChrec->NumOperands();
1388 for (size_t i = 0; i < std::max(lhsSize, rhsSize); ++i)
1389 {
1390 SCEV * lhs{};
1391 SCEV * rhs{};
1392 if (i < lhsSize)
1393 lhs = lhsChrec->GetOperand(i);
1394
1395 if (i < rhsSize)
1396 rhs = rhsChrec->GetOperand(i);
1397 newChrec->AddOperand(ApplyAddFolding(lhs, rhs, output));
1398 }
1399 return newChrec;
1400 }
1401
1402 // Chrec + any other operand
1403 // This handles Init, Constant, and any other SCEV type uniformly
1404 if (lhsChrec || rhsChrec)
1405 {
1406 auto * chrec = lhsChrec ? lhsChrec : rhsChrec;
1407 auto * otherOperand = lhsChrec ? rhsOperand : lhsOperand;
1408
1409 // Skip if otherOperand is zero constant (identity for addition)
1410 if (const auto constant = dynamic_cast<const SCEVConstant *>(otherOperand))
1411 {
1412 if (!SCEVConstant::IsNonZero(constant))
1413 {
1414 return chrec->Clone();
1415 }
1416 }
1417 auto newChrec = SCEVChainRecurrence::Create(chrec->GetLoop(), output);
1418 const auto chrecOperands = chrec->GetOperands();
1419
1420 bool isFirst = true;
1421 for (const auto operand : chrecOperands)
1422 {
1423 if (isFirst)
1424 {
1425 // Recursively fold the start value with the other operand
1426 newChrec->AddOperand(ApplyAddFolding(operand, otherOperand, output));
1427 isFirst = false;
1428 }
1429 else
1430 {
1431 newChrec->AddOperand(operand->Clone());
1432 }
1433 }
1434 return newChrec;
1435 }
1436
1437 const auto lhsNAryMulExpr = dynamic_cast<const SCEVNAryMulExpr *>(lhsOperand);
1438 const auto rhsNAryMulExpr = dynamic_cast<const SCEVNAryMulExpr *>(rhsOperand);
1439 // Handle n-ary multiply expressions - they become terms in an n-ary add expression
1440 if (lhsNAryMulExpr && rhsNAryMulExpr)
1441 {
1442 // Two multiply expressions - create add expression with both
1443 return SCEVNAryAddExpr::Create(lhsNAryMulExpr->Clone(), rhsNAryMulExpr->Clone());
1444 }
1445
1446 const auto lhsNAryAddExpr = dynamic_cast<const SCEVNAryAddExpr *>(lhsOperand);
1447 const auto rhsNAryAddExpr = dynamic_cast<const SCEVNAryAddExpr *>(rhsOperand);
1448 if ((lhsNAryMulExpr && rhsNAryAddExpr) || (rhsNAryMulExpr && lhsNAryAddExpr))
1449 {
1450 // Multiply expression with add expression - Clone the add expression and add the multiply as
1451 // a term
1452 const auto * mulExpr = lhsNAryMulExpr ? lhsNAryMulExpr : rhsNAryMulExpr;
1453 auto * addExpr = lhsNAryAddExpr ? lhsNAryAddExpr : rhsNAryAddExpr;
1454 auto newAddExpr = SCEV::CloneAs<SCEVNAryExpr>(*addExpr);
1455 newAddExpr->AddOperand(mulExpr->Clone());
1456 return newAddExpr->Clone();
1457 }
1458
1459 const auto lhsInit = dynamic_cast<const SCEVInit *>(lhsOperand);
1460 const auto rhsInit = dynamic_cast<const SCEVInit *>(rhsOperand);
1461 if ((lhsNAryMulExpr && rhsInit) || (rhsNAryMulExpr && lhsInit))
1462 {
1463 // Multiply expression with init - create add expression
1464 const auto * mulExpr = lhsNAryMulExpr ? lhsNAryMulExpr : rhsNAryMulExpr;
1465 const auto * init = lhsInit ? lhsInit : rhsInit;
1466 return SCEVNAryAddExpr::Create(mulExpr->Clone(), init->Clone());
1467 }
1468
1469 const auto lhsConstant = dynamic_cast<SCEVConstant *>(lhsOperand);
1470 const auto rhsConstant = dynamic_cast<SCEVConstant *>(rhsOperand);
1471 if ((lhsNAryMulExpr && SCEVConstant::IsNonZero(rhsConstant))
1472 || (rhsNAryMulExpr && SCEVConstant::IsNonZero(lhsConstant)))
1473 {
1474 // Multiply expression with nonzero constant - create add expression
1475 const auto * mulExpr = lhsNAryMulExpr ? lhsNAryMulExpr : rhsNAryMulExpr;
1476 const auto * constant = lhsConstant ? lhsConstant : rhsConstant;
1477 return SCEVNAryAddExpr::Create(mulExpr->Clone(), constant->Clone());
1478 }
1479
1480 if (lhsNAryMulExpr || rhsNAryMulExpr)
1481 {
1482 // Single multiply expression, no folding necessary
1483 const auto * mulExpr = lhsNAryMulExpr ? lhsNAryMulExpr : rhsNAryMulExpr;
1484 return mulExpr->Clone();
1485 }
1486
1487 if (lhsInit && rhsInit)
1488 {
1489 // We have two init nodes. Create a nAryAdd with lhsInit and rhsInit
1490 return SCEVNAryAddExpr::Create(lhsInit->Clone(), rhsInit->Clone());
1491 }
1492
1493 if ((lhsInit && rhsNAryAddExpr) || (rhsInit && lhsNAryAddExpr))
1494 {
1495 // We have an init and an add expr. Clone the add expression and add the init as an operand
1496 const auto * init = lhsInit ? lhsInit : rhsInit;
1497 auto * nAryAddExpr = lhsNAryAddExpr ? lhsNAryAddExpr : rhsNAryAddExpr;
1498 auto newAddExpr = SCEV::CloneAs<SCEVNAryAddExpr>(*nAryAddExpr);
1499 newAddExpr->AddOperand(init->Clone());
1500 return newAddExpr->Clone();
1501 }
1502
1503 if ((lhsInit && SCEVConstant::IsNonZero(rhsConstant))
1504 || (rhsInit && SCEVConstant::IsNonZero(lhsConstant)))
1505 {
1506 // We have an init and a nonzero constant. Create a nAryAdd with init and constant
1507 const auto * init = lhsInit ? lhsInit : rhsInit;
1508 const auto * constant = lhsConstant ? lhsConstant : rhsConstant;
1509 return SCEVNAryAddExpr::Create(init->Clone(), constant->Clone());
1510 }
1511
1512 if (lhsInit || rhsInit)
1513 {
1514 // Only one operand. Add it
1515 const auto * init = lhsInit ? lhsInit : rhsInit;
1516 return init->Clone();
1517 }
1518
1519 if (lhsNAryAddExpr && rhsNAryAddExpr)
1520 {
1521 // We have two add expressions. Clone the lhs and add the rhs operands
1522 auto lhsNewNAryAddExpr = SCEV::CloneAs<SCEVNAryAddExpr>(*lhsNAryAddExpr);
1523 for (auto op : rhsNAryAddExpr->GetOperands())
1524 {
1525 lhsNewNAryAddExpr->AddOperand(op->Clone());
1526 }
1527 return lhsNewNAryAddExpr;
1528 }
1529
1530 if ((lhsNAryAddExpr && SCEVConstant::IsNonZero(rhsConstant))
1531 || (rhsNAryAddExpr && SCEVConstant::IsNonZero(lhsConstant)))
1532 {
1533 // We have an add expr and a nonzero constant. Clone the add expr and add the constant
1534 auto * nAryAddExpr = lhsNAryAddExpr ? lhsNAryAddExpr : rhsNAryAddExpr;
1535 auto * constant = lhsConstant ? lhsConstant : rhsConstant;
1536 auto newNAryAddExpr = SCEV::CloneAs<SCEVNAryAddExpr>(*nAryAddExpr);
1537
1538 // Check if there is already a constant operand in the n-ary expression
1539 // If so, fold the new constant with the old one instead of adding it as an operand
1540 bool folded = false;
1541 for (size_t i = 0; i < newNAryAddExpr->NumOperands(); ++i)
1542 {
1543 if (auto existingConstant = dynamic_cast<SCEVConstant *>(newNAryAddExpr->GetOperands()[i]))
1544 {
1545 // Fold the two constants together directly
1546 auto foldedConstant = ApplyAddFolding(existingConstant, constant, output);
1547 newNAryAddExpr->ReplaceOperand(i, foldedConstant);
1548 folded = true;
1549 break;
1550 }
1551 }
1552
1553 if (!folded)
1554 {
1555 // No existing constant to fold with, just append
1556 newNAryAddExpr->AddOperand(constant->Clone());
1557 }
1558
1559 return newNAryAddExpr;
1560 }
1561
1562 if (lhsNAryAddExpr || rhsNAryAddExpr)
1563 {
1564 const auto * nAryAddExpr = lhsNAryAddExpr ? lhsNAryAddExpr : rhsNAryAddExpr;
1565 return nAryAddExpr->Clone();
1566 }
1567 if (lhsConstant && rhsConstant)
1568 {
1569 // Two constants, get their value, and combine them (fold)
1570 const auto lhsValue = lhsConstant->GetValue();
1571 const auto rhsValue = rhsConstant->GetValue();
1572
1573 return SCEVConstant::Create(lhsValue + rhsValue);
1574 }
1575
1576 if (lhsConstant || rhsConstant)
1577 {
1578 const auto * constant = lhsConstant ? lhsConstant : rhsConstant;
1579 return constant->Clone();
1580 }
1581
1582 return SCEVUnknown::Create();
1583}
1584
1585std::unique_ptr<SCEVChainRecurrence>
1587 SCEVChainRecurrence * lhsChrec,
1588 SCEVChainRecurrence * rhsChrec,
1589 rvsdg::Output & output)
1590{
1591 const auto lhsSize = lhsChrec->NumOperands();
1592 const auto rhsSize = rhsChrec->NumOperands();
1593
1594 if (rhsSize == 0)
1595 return SCEV::CloneAs<SCEVChainRecurrence>(*lhsChrec);
1596 if (lhsSize == 0)
1597 return SCEV::CloneAs<SCEVChainRecurrence>(*rhsChrec);
1598
1599 // Handle G * {e,+,f,...} where G is loop invariant
1600 if (lhsSize == 1)
1601 {
1602 auto newChrec = SCEVChainRecurrence::Create(lhsChrec->GetLoop(), output);
1603 // G * {e,+,f,...} = {G * e,+,G * f,...}
1604 auto lhs = lhsChrec->GetOperand(0);
1605
1606 for (auto rhs : rhsChrec->GetOperands())
1607 {
1608 newChrec->AddOperand(ApplyMulFolding(lhs, rhs, output));
1609 }
1610 return newChrec;
1611 }
1612 if (rhsSize == 1)
1613 {
1614 auto newChrec = SCEVChainRecurrence::Create(lhsChrec->GetLoop(), output);
1615 // {e,+,f,...} * G = {e * G,+,f * G,...}
1616 auto rhs = rhsChrec->GetOperand(0);
1617
1618 for (auto lhs : lhsChrec->GetOperands())
1619 {
1620 newChrec->AddOperand(ApplyMulFolding(lhs, rhs, output));
1621 }
1622 return newChrec;
1623 }
1624
1625 // Below is an implementation of the algorithm CRProd from Bachmann et al., ‘Chains of recurrences
1626 // — a method to expedite the evaluation of closed-form functions’
1627 // (https://doi.org/10.1145/190347.190423)
1628 //
1629 // Let lhs = F, rhs = G be CR’s of length k and l.
1630 //
1631 // The product of F and G, S = F * G can be constructed using the algorithm CRProd given below
1632 //
1633 // Algorithm CRProd: Let
1634 // F = {a0, +, a1, +, ..., +, ak} and G = {b0, +, b1, +, ..., +, bl}
1635 // with k ≥ l. This algorithm returns a simple CR S of length k + l such that F * G = S.
1636 //
1637 // P1 [Base case]
1638 // If l = 1 return {a0*b0, +, a1*b0, +, ..., +, ak*b0}
1639 //
1640 // P2 [Prepare recursive calls] Let
1641 // f = {a1, +, a2, +, ..., +, ak}
1642 // g = {b1, +, b2, +, ..., +, bl}
1643 //
1644 // G' = G + g = {b0 + b1, +, b1 + b2, +, ..., +, bl}
1645 //
1646 // P3 [Recursive calls] Set
1647 // {x1'', +, x2'', +, ..., +, x(k+l)''} ← CRProd(F, g)
1648 // {x1', +, x2', +, ..., +, x(k+l)'} ← CRProd(f, G')
1649 //
1650 // P4 [Fold the results together and return]
1651 // return {a0*b0, +, x1' + x1'', +, ..., +, x(k+l)' + x(k+l)''}
1652
1653 JLM_ASSERT(lhsSize >= 2);
1654 JLM_ASSERT(rhsSize >= 2);
1655
1656 if (rhsSize > lhsSize)
1657 std::swap(lhsChrec, rhsChrec);
1658
1659 std::unique_ptr<SCEVChainRecurrence> lhsStepRecurrence, rhsStepRecurrence;
1660
1661 const auto lhsStep = *lhsChrec->GetStep();
1662 if (!lhsStep)
1663 {
1664 // This should not happen since we check size above
1665 throw std::logic_error("Could not get step for LHS in ComputeProductOfChrecs!");
1666 }
1667
1668 const auto rhsStep = *rhsChrec->GetStep();
1669 if (!rhsStep)
1670 {
1671 throw std::logic_error("Could not get step for RHS in ComputeProductOfChrecs!");
1672 }
1673
1674 if (dynamic_cast<SCEVChainRecurrence *>(lhsStep.get()))
1675 lhsStepRecurrence = SCEV::CloneAs<SCEVChainRecurrence>(*lhsStep);
1676 else
1677 lhsStepRecurrence = SCEVChainRecurrence::Create(lhsChrec->GetLoop(), output, lhsStep->Clone());
1678
1679 if (dynamic_cast<SCEVChainRecurrence *>(rhsStep.get()))
1680 rhsStepRecurrence = SCEV::CloneAs<SCEVChainRecurrence>(*rhsStep);
1681 else
1682 rhsStepRecurrence = SCEVChainRecurrence::Create(rhsChrec->GetLoop(), output, rhsStep->Clone());
1683
1684 const auto rhsMarked = SCEV::CloneAs<SCEVChainRecurrence>(
1685 *ApplyAddFolding(rhsChrec, rhsStepRecurrence.get(), output));
1686
1687 const auto res1 = ComputeProductOfChrecs(lhsChrec, rhsStepRecurrence.get(), output);
1688 const auto res2 = ComputeProductOfChrecs(rhsMarked.get(), lhsStepRecurrence.get(), output);
1689
1690 auto resFolded =
1691 SCEV::CloneAs<SCEVChainRecurrence>(*ApplyAddFolding(res1.get(), res2.get(), output));
1692
1693 const auto first = ApplyMulFolding(lhsChrec->GetOperand(0), rhsChrec->GetOperand(0), output);
1694 resFolded->AddOperandToFront(first);
1695
1696 return resFolded;
1697}
1698
1699std::unique_ptr<SCEV>
1700ScalarEvolution::ApplyMulFolding(SCEV * lhsOperand, SCEV * rhsOperand, rvsdg::Output & output)
1701{
1702 // We have the following folding rules from the CR algebra:
1703 // G * {e,+,f} => {G * e,+,G * f}
1704 // {e,+,f} * {g,+,h} => {e * g,+,e * h + f * g + f * h,+,2*f*h}
1705 //
1706 // Similar to addition, we need to handle SCEVInit nodes and n-ary expressions.
1707 // For multiplication with init nodes, we create n-ary multiply expressions.
1708
1709 if (const auto *lhsUnknown = dynamic_cast<const SCEVUnknown *>(lhsOperand),
1710 *rhsUnknown = dynamic_cast<const SCEVUnknown *>(rhsOperand);
1711 lhsUnknown || rhsUnknown)
1712 {
1713 return SCEVUnknown::Create();
1714 }
1715
1716 auto lhsChrec = dynamic_cast<SCEVChainRecurrence *>(lhsOperand);
1717 auto rhsChrec = dynamic_cast<SCEVChainRecurrence *>(rhsOperand);
1718 if (lhsChrec && rhsChrec)
1719 {
1720 if (&lhsChrec->GetLoop() != &rhsChrec->GetLoop())
1721 {
1723 lhsChrec->GetLoop(),
1724 output,
1725 SCEVNAryMulExpr::Create(lhsChrec->Clone(), rhsChrec->Clone()));
1726 }
1727
1728 return ComputeProductOfChrecs(lhsChrec, rhsChrec, output);
1729 }
1730
1731 // Chrec * any other operand
1732 // This handles Init, Constant, and any other SCEV type uniformly
1733 if (lhsChrec || rhsChrec)
1734 {
1735 auto * chrec = lhsChrec ? lhsChrec : rhsChrec;
1736 auto * otherOperand = lhsChrec ? rhsOperand : lhsOperand;
1737
1738 if (auto constant = dynamic_cast<const SCEVConstant *>(otherOperand))
1739 {
1740 if (constant->GetValue() == 1)
1741 {
1742 // Dont fold if operand is constant one (identity for multiplication)
1743 return chrec->Clone();
1744 }
1745
1746 if (constant->GetValue() == 0)
1747 {
1748 // Fold to zero
1749 return SCEVConstant::Create(0);
1750 }
1751 }
1752 auto newChrec = SCEVChainRecurrence::Create(chrec->GetLoop(), output);
1753 const auto chrecOperands = chrec->GetOperands();
1754
1755 for (auto & operand : chrecOperands)
1756 {
1757 // Recursively fold the start value with the other operand
1758 newChrec->AddOperand(ApplyMulFolding(operand, otherOperand, output));
1759 }
1760 return newChrec;
1761 }
1762
1763 const auto lhsNAryAddExpr = dynamic_cast<const SCEVNAryAddExpr *>(lhsOperand);
1764 const auto rhsNAryAddExpr = dynamic_cast<const SCEVNAryAddExpr *>(rhsOperand);
1765 if (lhsNAryAddExpr || rhsNAryAddExpr)
1766 {
1767 // Handle n-ary add expressions - distribute multiplication
1768 // (a + b + c) × G = a×G + b×G + c×G
1769 const auto nAryAddExpr = lhsNAryAddExpr ? lhsNAryAddExpr : rhsNAryAddExpr;
1770 const auto other = lhsNAryAddExpr ? rhsOperand : lhsOperand;
1771
1772 auto resultAddExpr = SCEVNAryAddExpr::Create();
1773 for (auto operand : nAryAddExpr->GetOperands())
1774 {
1775 auto product = ApplyMulFolding(operand, other, output);
1776 resultAddExpr->AddOperand(std::move(product));
1777 }
1778 return resultAddExpr;
1779 }
1780
1781 const auto lhsInit = dynamic_cast<const SCEVInit *>(lhsOperand);
1782 const auto rhsInit = dynamic_cast<const SCEVInit *>(rhsOperand);
1783 if (lhsInit && rhsInit)
1784 {
1785 // Two init nodes - create n-ary multiply expression
1786 return SCEVNAryMulExpr::Create(lhsInit->Clone(), rhsInit->Clone());
1787 }
1788
1789 const auto lhsNAryMulExpr = dynamic_cast<const SCEVNAryMulExpr *>(lhsOperand);
1790 const auto rhsNAryMulExpr = dynamic_cast<const SCEVNAryMulExpr *>(rhsOperand);
1791 if ((lhsInit && rhsNAryMulExpr) || (rhsInit && lhsNAryMulExpr))
1792 {
1793 // Init node with n-ary multiply expression - Clone mult expr and add init as an operand
1794 const auto * init = lhsInit ? lhsInit : rhsInit;
1795 auto * nAryMulExpr = lhsNAryMulExpr ? lhsNAryMulExpr : rhsNAryMulExpr;
1796 auto newNAryMulExpr = SCEV::CloneAs<SCEVNAryMulExpr>(*nAryMulExpr);
1797 newNAryMulExpr->AddOperand(init->Clone());
1798 return newNAryMulExpr->Clone();
1799 }
1800
1801 auto lhsConstant = dynamic_cast<SCEVConstant *>(lhsOperand);
1802 auto rhsConstant = dynamic_cast<SCEVConstant *>(rhsOperand);
1803 if ((lhsInit && rhsConstant && rhsConstant->GetValue() != 1)
1804 || (rhsInit && lhsConstant && lhsConstant->GetValue() != 1))
1805 {
1806 // Init node with non-one constant - create n-ary multiply expression
1807 const auto * init = lhsInit ? lhsInit : rhsInit;
1808 const auto * constant = lhsConstant ? lhsConstant : rhsConstant;
1809 return SCEVNAryMulExpr::Create(init->Clone(), constant->Clone());
1810 }
1811
1812 if (lhsInit || rhsInit)
1813 {
1814 // Single init node, no folding necessary
1815 const auto * init = lhsInit ? lhsInit : rhsInit;
1816 return init->Clone();
1817 }
1818
1819 if (lhsNAryMulExpr && rhsNAryMulExpr)
1820 {
1821 // Two n-ary mult expressions - combine operands
1822 auto lhsNewNAryMulExpr = SCEV::CloneAs<SCEVNAryMulExpr>(*lhsNAryMulExpr);
1823 for (auto op : rhsNAryMulExpr->GetOperands())
1824 {
1825 lhsNewNAryMulExpr->AddOperand(op->Clone());
1826 }
1827 return lhsNewNAryMulExpr;
1828 }
1829
1830 if ((lhsNAryMulExpr && rhsConstant && rhsConstant->GetValue() != 1)
1831 || (rhsNAryMulExpr && lhsConstant && lhsConstant->GetValue() != 1))
1832 {
1833 // N-ary mult expression with non-one constant - Clone mult expression and add constant
1834 auto * nAryMulExpr = lhsNAryMulExpr ? lhsNAryMulExpr : rhsNAryMulExpr;
1835 auto * constant = lhsConstant ? lhsConstant : rhsConstant;
1836
1837 auto newNAryMulExpr = SCEV::CloneAs<SCEVNAryMulExpr>(*nAryMulExpr);
1838
1839 bool folded = false;
1840 for (size_t i = 0; i < newNAryMulExpr->NumOperands(); ++i)
1841 {
1842 if (auto existingConstant = dynamic_cast<SCEVConstant *>(newNAryMulExpr->GetOperands()[i]))
1843 {
1844 // Fold the two constants together directly
1845 auto foldedConstant = ApplyMulFolding(existingConstant, constant, output);
1846 newNAryMulExpr->ReplaceOperand(i, foldedConstant);
1847 folded = true;
1848 break;
1849 }
1850 }
1851
1852 if (!folded)
1853 {
1854 // No existing constant to fold with, just append
1855 newNAryMulExpr->AddOperand(constant->Clone());
1856 }
1857
1858 return newNAryMulExpr;
1859 }
1860
1861 if (lhsNAryMulExpr || rhsNAryMulExpr)
1862 {
1863 const auto * nAryMulExpr = lhsNAryMulExpr ? lhsNAryMulExpr : rhsNAryMulExpr;
1864 return nAryMulExpr->Clone();
1865 }
1866
1867 if (lhsConstant && rhsConstant)
1868 {
1869 // Two constants - fold by multiplying values together
1870 const auto lhsValue = lhsConstant->GetValue();
1871 const auto rhsValue = rhsConstant->GetValue();
1872 return SCEVConstant::Create(lhsValue * rhsValue);
1873 }
1874
1875 if (lhsConstant || rhsConstant)
1876 {
1877 const auto * constant = lhsConstant ? lhsConstant : rhsConstant;
1878 return constant->Clone();
1879 }
1880
1881 return SCEVUnknown::Create();
1882}
1883
1884std::unique_ptr<SCEV>
1886{
1887 // -(c)
1888 if (const auto c = dynamic_cast<const SCEVConstant *>(&scev))
1889 {
1890 const auto value = c->GetValue();
1891 return SCEVConstant::Create(-value);
1892 }
1893 // -(-x) -> x
1894 if (const auto mul = dynamic_cast<const SCEVMulExpr *>(&scev))
1895 {
1896 if (const auto c = dynamic_cast<const SCEVConstant *>(mul->GetLeftOperand());
1897 c && c->GetValue() == -1)
1898 {
1899 return mul->GetRightOperand()->Clone();
1900 }
1901 if (const auto c = dynamic_cast<const SCEVConstant *>(mul->GetRightOperand());
1902 c && c->GetValue() == -1)
1903 {
1904 return mul->GetLeftOperand()->Clone();
1905 }
1906 } // -(x + y) -> (-x) + (-y)
1907 if (const auto add = dynamic_cast<const SCEVAddExpr *>(&scev))
1908 {
1909 return SCEVAddExpr::Create(
1910 GetNegativeSCEV(*add->GetLeftOperand()),
1911 GetNegativeSCEV(*add->GetRightOperand()));
1912 }
1913 // General case: -(x) -> (-1) * x
1915}
1916
1917bool
1919{
1920 auto deps = dependencyGraph[&output];
1921 if (deps.find(&output) != deps.end())
1922 {
1923 if (deps[&output].count != 1)
1924 {
1925 // First check that variable has only one self-reference
1926 return false;
1927 }
1928 if (deps[&output].operation == DependencyOp::Mul)
1929 {
1930 // A variable cannot have a self-depencency via multiplication (results in a geometric
1931 // induction variable)
1932 return false;
1933 }
1934 }
1935
1936 // Then check for cycles through other variables
1937 std::unordered_set<const rvsdg::Output *> visited{};
1938 std::unordered_set<const rvsdg::Output *> recursionStack{};
1939 return !HasCycleThroughOthers(output, output, dependencyGraph, visited, recursionStack);
1940}
1941
1942bool
1944 rvsdg::Output & currentOutput,
1945 const rvsdg::Output & originalOutput,
1946 DependencyGraph & dependencyGraph,
1947 std::unordered_set<const rvsdg::Output *> & visited,
1948 std::unordered_set<const rvsdg::Output *> & recursionStack)
1949{
1950 visited.insert(&currentOutput);
1951 recursionStack.insert(&currentOutput);
1952
1953 for (const auto & [depPtr, depCount] : dependencyGraph[&currentOutput])
1954 {
1955 // Ignore self-references
1956 if (depPtr == &currentOutput)
1957 continue;
1958
1959 // Found a cycle back to the ORIGINAL node we started from
1960 // This means the original output is explicitly part of the cycle
1961 if (depPtr == &originalOutput)
1962 return true;
1963
1964 // Already explored this branch, no cycle containing the original output
1965 if (visited.find(depPtr) != visited.end())
1966 continue;
1967
1968 // Recursively check dependencies, keeping track of the original node
1969 if (HasCycleThroughOthers(*depPtr, originalOutput, dependencyGraph, visited, recursionStack))
1970 return true;
1971 }
1972
1973 recursionStack.erase(&currentOutput);
1974 return false;
1975}
1976
1977bool
1979{
1980 if (dynamic_cast<const SCEVUnknown *>(&scev))
1981 return true;
1982
1983 if (dynamic_cast<const SCEVInit *>(&scev) || dynamic_cast<const SCEVConstant *>(&scev)
1984 || dynamic_cast<const SCEVPlaceholder *>(&scev))
1985 {
1986 return false;
1987 }
1988
1989 if (auto * binaryExpr = dynamic_cast<const SCEVBinaryExpr *>(&scev))
1990 {
1991 return IsUnknown(*binaryExpr->GetLeftOperand()) || IsUnknown(*binaryExpr->GetLeftOperand());
1992 }
1993
1994 if (auto * nAryExpr = dynamic_cast<const SCEVNAryExpr *>(&scev))
1995 {
1996 for (const auto operand : nAryExpr->GetOperands())
1997 {
1998 if (IsUnknown(*operand))
1999 return true;
2000 }
2001 return false;
2002 }
2003
2004 throw std::logic_error("Invalid SCEV type in IsUnknown!\n");
2005}
2006
2007bool
2009{
2010 if (typeid(a) != typeid(b))
2011 return false;
2012
2013 if (dynamic_cast<const SCEVUnknown *>(&a))
2014 return true;
2015
2016 if (auto * constantA = dynamic_cast<const SCEVConstant *>(&a))
2017 {
2018 auto * constantB = dynamic_cast<const SCEVConstant *>(&b);
2019 return constantA->GetValue() == constantB->GetValue();
2020 }
2021
2022 if (auto * initA = dynamic_cast<const SCEVInit *>(&a))
2023 {
2024 auto * initB = dynamic_cast<const SCEVInit *>(&b);
2025 return &initA->GetPrePointer() == &initB->GetPrePointer();
2026 }
2027
2028 if (auto * binaryExprA = dynamic_cast<const SCEVBinaryExpr *>(&a))
2029 {
2030 auto * binaryExprB = dynamic_cast<const SCEVBinaryExpr *>(&b);
2031 return StructurallyEqual(*binaryExprA->GetLeftOperand(), *binaryExprB->GetLeftOperand())
2032 && StructurallyEqual(*binaryExprA->GetRightOperand(), *binaryExprB->GetRightOperand());
2033 }
2034
2035 if (auto * chrecA = dynamic_cast<const SCEVChainRecurrence *>(&a))
2036 {
2037 auto * chrecB = dynamic_cast<const SCEVChainRecurrence *>(&b);
2038 if (&chrecA->GetLoop() != &chrecB->GetLoop())
2039 return false;
2040 if (&chrecA->GetOutput() != &chrecB->GetOutput())
2041 return false;
2042 if (chrecA->NumOperands() != chrecB->NumOperands())
2043 return false;
2044 for (size_t i = 0; i < chrecA->NumOperands(); ++i)
2045 {
2046 if (!StructurallyEqual(*chrecA->GetOperands()[i], *chrecB->GetOperands()[i]))
2047 return false;
2048 }
2049 return true;
2050 }
2051
2052 if (auto * nAryExprA = dynamic_cast<const SCEVNAryExpr *>(&a))
2053 {
2054 auto * nAryExprB = dynamic_cast<const SCEVNAryExpr *>(&b);
2055 if (nAryExprA->NumOperands() != nAryExprB->NumOperands())
2056 return false;
2057 for (size_t i = 0; i < nAryExprA->NumOperands(); ++i)
2058 {
2059 if (!StructurallyEqual(*nAryExprA->GetOperands()[i], *nAryExprB->GetOperands()[i]))
2060 return false;
2061 }
2062 return true;
2063 }
2064
2065 return false;
2066}
2067}
static rvsdg::Input & getAddressInput(const rvsdg::Node &node) noexcept
static std::unique_ptr< SCEVAddExpr > Create(std::unique_ptr< SCEV > left, std::unique_ptr< SCEV > right)
rvsdg::ThetaNode & GetLoop() const noexcept
static std::unique_ptr< SCEVChainRecurrence > Create(rvsdg::ThetaNode &loop, rvsdg::Output &output)
static bool IsQuadratic(const SCEVChainRecurrence &chrec)
static bool IsConstant(const SCEVChainRecurrence &chrec)
std::optional< std::unique_ptr< SCEV > > GetStep() const
static bool IsAffine(const SCEVChainRecurrence &chrec)
static std::unique_ptr< SCEVConstant > Create(const int64_t value)
static bool IsNonZero(const SCEVConstant *c)
static std::unique_ptr< SCEVInit > Create(rvsdg::Output &prePointer)
rvsdg::Output & GetPrePointer() const noexcept
static std::unique_ptr< SCEVMulExpr > Create(std::unique_ptr< SCEV > left, std::unique_ptr< SCEV > right)
static std::unique_ptr< SCEVNAryAddExpr > Create(Args &&... operands)
void RemoveOperand(const size_t index)
SCEV * GetOperand(const size_t index) const
void ReplaceOperand(const size_t index, const std::unique_ptr< SCEV > &operand)
std::vector< SCEV * > GetOperands() const
static std::unique_ptr< SCEVNAryMulExpr > Create(Args &&... operands)
static std::unique_ptr< SCEVPlaceholder > Create(rvsdg::Output &PrePointer_)
static std::unique_ptr< SCEVUnknown > Create()
virtual std::unique_ptr< SCEV > Clone() const =0
size_t GetTripCount(const rvsdg::ThetaNode &thetaNode) const
std::unordered_map< const rvsdg::ThetaNode *, size_t > TripCountMap_
const std::unordered_map< rvsdg::Output *, std::unique_ptr< SCEV > > & GetSCEVMap() const noexcept
std::unordered_set< const rvsdg::Output * > LoopVars_
void SetTripCount(const rvsdg::ThetaNode &thetaNode, const size_t tripCount)
const std::unordered_map< const rvsdg::ThetaNode *, size_t > & GetTripCountMap() const noexcept
void InsertSCEV(rvsdg::Output &output, const std::unique_ptr< SCEV > &scev)
std::unordered_map< rvsdg::Output *, std::unique_ptr< SCEVChainRecurrence > > ChrecMap_
std::unique_ptr< SCEVChainRecurrence > TryGetChrecForOutput(rvsdg::Output &output) const
const std::unordered_map< rvsdg::Output *, std::unique_ptr< SCEVChainRecurrence > > & GetChrecMap() const noexcept
Context(const Context &)=delete
void AddLoopVar(const rvsdg::Output &var)
std::unordered_map< rvsdg::Output *, std::unique_ptr< SCEV > > SCEVMap_
Context & operator=(Context &&)=delete
Context & operator=(const Context &)=delete
std::unique_ptr< SCEV > TryGetSCEVForOutput(rvsdg::Output &output) const
void InsertChrec(rvsdg::Output &output, const std::unique_ptr< SCEVChainRecurrence > &chrec)
static std::unique_ptr< Context > Create()
int GetNumInductionVariablesWithOrder(const size_t n) const
~Statistics() noexcept override=default
static std::string GetTripCountString(const std::unordered_map< const rvsdg::ThetaNode *, size_t > &tripCountMap)
static std::unique_ptr< Statistics > Create(const util::FilePath &sourceFile)
void Stop(const Context &context) noexcept
std::unordered_map< rvsdg::Output *, DependencyInfo > DependencyMap
std::unordered_map< rvsdg::Output *, DependencyMap > DependencyGraph
static std::unique_ptr< SCEVChainRecurrence > ComputeProductOfChrecs(SCEVChainRecurrence *lhsChrec, SCEVChainRecurrence *rhsChrec, rvsdg::Output &output)
std::optional< std::unique_ptr< SCEV > > TryReplaceInitForSCEV(const SCEV &scev, rvsdg::Output &output)
std::unique_ptr< SCEVChainRecurrence > GetOrCreateStepForSCEV(rvsdg::Output &output, const SCEV &scevTree, rvsdg::ThetaNode &thetaNode)
void PerformSCEVAnalysis(rvsdg::ThetaNode &thetaNode)
void Run(rvsdg::RvsdgModule &rvsdgModule, util::StatisticsCollector &statisticsCollector) override
Perform RVSDG transformation.
static void FindDependenciesForSCEV(const SCEV &scev, DependencyMap &dependencies, DependencyOp op)
static bool IsStepZero(const SCEV &stepSCEV)
static std::unique_ptr< SCEV > ApplyMulFolding(SCEV *lhsOperand, SCEV *rhsOperand, rvsdg::Output &output)
Apply folding rules for multiplication to combine two SCEV operands into one.
static std::unique_ptr< SCEV > FoldNAryExpression(SCEVNAryExpr &expression, rvsdg::Output &output)
Try to combine the constants in an n-ary expression (Add or Mul) into themselves.
static std::optional< size_t > SolveQuadraticEquation(int64_t a, int64_t b, int64_t c)
Tries to find a solution to the quadratic equation a^2 x + b x + c = 0 using integer arithmetic.
DependencyGraph CreateDependencyGraph(const rvsdg::ThetaNode &thetaNode) const
static std::unique_ptr< SCEV > ApplyAddFolding(SCEV *lhsOperand, SCEV *rhsOperand, rvsdg::Output &output)
Apply folding rules for addition to combine two SCEV operands into one.
static bool HasCycleThroughOthers(rvsdg::Output &currentOutput, const rvsdg::Output &originalOutput, DependencyGraph &dependencyGraph, std::unordered_set< const rvsdg::Output * > &visited, std::unordered_set< const rvsdg::Output * > &recursionStack)
~ScalarEvolution() noexcept override
std::unique_ptr< Context > Context_
static bool IsStepPositive(const SCEV &stepSCEV)
static bool CanCreateChainRecurrence(rvsdg::Output &output, DependencyGraph &dependencyGraph)
std::unique_ptr< SCEV > ComputeSCEVForGepInnerOffset(const rvsdg::SimpleNode &gepNode, size_t inputIndex, const rvsdg::Type &type)
std::optional< size_t > GetPredictedTripCount(rvsdg::ThetaNode &thetaNode)
static bool StructurallyEqual(const SCEV &a, const SCEV &b)
static bool IsUnknown(const SCEV &scev)
std::unordered_map< const rvsdg::Output *, std::unique_ptr< SCEV > > GetSCEVMap() const
static bool IsStepNegative(const SCEV &stepSCEV)
static std::vector< rvsdg::Output * > TopologicalSort(DependencyGraph &dependencyGraph)
std::unique_ptr< SCEV > GetOrCreateSCEVForOutput(rvsdg::Output &output)
std::unique_ptr< SCEVChainRecurrence > GetOrCreateChainRecurrence(rvsdg::Output &output, const SCEV &scev, rvsdg::ThetaNode &thetaNode)
std::unordered_map< const rvsdg::ThetaNode *, size_t > GetTripCountMap() const noexcept
static std::unique_ptr< SCEV > GetNegativeSCEV(const SCEV &scev)
void AnalyzeRegion(rvsdg::Region &region)
static std::optional< size_t > ComputeBackedgeTakenCountForChrec(const SCEVChainRecurrence &chrec, int64_t bound, const rvsdg::SimpleOperation *comparisonOperation)
StructType class.
Definition types.hpp:184
Region & GetRootRegion() const noexcept
Definition graph.hpp:99
Output * origin() const noexcept
Definition node.hpp:58
size_t ninputs() const noexcept
Definition node.hpp:609
Represent acyclic RVSDG subgraphs.
Definition region.hpp:213
NodeRange Nodes() noexcept
Definition region.hpp:375
const std::optional< util::FilePath > & SourceFilePath() const noexcept
Graph & Rvsdg() noexcept
NodeInput * input(size_t index) const noexcept
LoopVar MapPreLoopVar(const rvsdg::Output &argument) const
Maps variable at start of loop iteration to full varibale description.
Definition theta.cpp:140
std::vector< LoopVar > GetLoopVars() const
Returns all loop variables.
Definition theta.cpp:193
RegionResult * predicate() const noexcept
Definition theta.hpp:96
bool insert(ItemType item)
Definition HashSet.hpp:210
void CollectDemandedStatistics(std::unique_ptr< Statistics > statistics)
Statistics Interface.
util::Timer & GetTimer(const std::string &name)
util::Timer & AddTimer(std::string name)
void AddMeasurement(std::string name, T value)
void start() noexcept
Definition time.hpp:54
void stop() noexcept
Definition time.hpp:67
#define JLM_ASSERT(x)
Definition common.hpp:16
Global memory state passed between functions.
size_t GetTypeAllocSize(const rvsdg::Type &type)
Definition types.cpp:473
static util::StatisticsCollector statisticsCollector
std::optional< int64_t > tryGetConstantSignedInteger(const rvsdg::Output &output)
Definition Trace.cpp:97
Output & traceOutput(Output &output, bool mayEnterSubregions, const Region *withinRegion)
Definition Trace.cpp:454
static bool ThetaLoopVarIsInvariant(const ThetaNode::LoopVar &loopVar) noexcept
Definition theta.hpp:266
@ State
Designate a state type.
NodeType * TryGetOwnerNode(const rvsdg::Input &input) noexcept
Checks if this is an input to a node of specified type.
Definition node.hpp:872
rvsdg::Input * input
Variable at loop entry (input to theta).
Definition theta.hpp:54
rvsdg::Input * post
Variable after iteration (output result from subregion).
Definition theta.hpp:62