[llvm] [CodeGenPrepare] Extend SimplifyCFG to lower s/ucmp + 3-arm switch (PR #223268)

via llvm-commits llvm-commits at lists.llvm.org
Sun Sep 13 12:00:19 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-llvm-transforms

Author: Joe Isaacs (joseph-isaacs)

<details>
<summary>Changes</summary>

Fixes https://github.com/llvm/llvm-project/issues/176492

Lowers ucmp + switch into icmp respecting the weights then the original switch order. 

```llvm
%c = call i8 @<!-- -->llvm.ucmp.i8.i8(i8 %x, i8 %y)
switch i8 %c, label %unr [
  i8 -1, label %lt
  i8 0, label %end
  i8 1, label %gt
]
```

into 

```llvm
  %is.lt = icmp ult i8 %x, %y 
  br i1 %is.lt, label %lt, label %cmp.next

cmp.next:
  %is.eq = icmp eq i8 %x, %y
  br i1 %is.eq, label %end, label %gt
```

This lowering produces strictly better code than a lowering of  ucmp to`%c = ucmp` and switch to `icmp %c`.

## Implementation 

Add `optimizeSwitchOfCmpIntrinsic` from `CodeGenPrepare` to work with 3 live successor blocks. 1 or 2 handled previously. 

## Related 

Similar to https://github.com/llvm/llvm-project/pull/176582, but this approach uses `SimplifyCFG` over `CodeGenPrepare`, since its a cost model is not needed here. Instead run this on the final `SimplifyCFG`.

Previous Attempt: https://github.com/llvm/llvm-project/pull/222786

Fixes https://github.com/rust-lang/rust/issues/86511

Co-Authored-By: codex

---
Full diff: https://github.com/llvm/llvm-project/pull/223268.diff


2 Files Affected:

- (modified) llvm/lib/CodeGen/CodeGenPrepare.cpp (+146-4) 
- (added) llvm/test/Transforms/CodeGenPrepare/switch-on-cmp.ll (+256) 


``````````diff
diff --git a/llvm/lib/CodeGen/CodeGenPrepare.cpp b/llvm/lib/CodeGen/CodeGenPrepare.cpp
index 52ed1255ab9c4..431234d999105 100644
--- a/llvm/lib/CodeGen/CodeGenPrepare.cpp
+++ b/llvm/lib/CodeGen/CodeGenPrepare.cpp
@@ -436,9 +436,10 @@ class CodeGenPrepare {
   bool optimizeFunnelShift(IntrinsicInst *Fsh);
   bool optimizeSelectInst(SelectInst *SI);
   bool optimizeShuffleVectorInst(ShuffleVectorInst *SVI);
+  bool optimizeSwitchOfCmpIntrinsic(SwitchInst *SI);
   bool optimizeSwitchType(SwitchInst *SI);
   bool optimizeSwitchPhiConstants(SwitchInst *SI);
-  bool optimizeSwitchInst(SwitchInst *SI);
+  bool optimizeSwitchInst(SwitchInst *SI, ModifyDT &ModifiedDT);
   bool optimizeExtractElementInst(Instruction *Inst);
   bool dupRetToEnableTailCallOpts(BasicBlock *BB, ModifyDT &ModifiedDT);
   bool fixupDbgVariableRecord(DbgVariableRecord &I);
@@ -8062,6 +8063,142 @@ bool CodeGenPrepare::tryToSinkFreeOperands(Instruction *I) {
   return Changed;
 }
 
+/// Lower a switch that sends each result of a comparison intrinsic to a
+/// different block into two branches comparing the intrinsic's operands.
+bool CodeGenPrepare::optimizeSwitchOfCmpIntrinsic(SwitchInst *SI) {
+  auto *Cmp = dyn_cast<CmpIntrinsic>(SI->getCondition());
+  if (!Cmp || !Cmp->hasOneUse())
+    return false;
+
+  // Must have three real successors.
+  // If two switch cases require a default reachable block.
+  unsigned NumCases = SI->getNumCases();
+  if (NumCases + !SI->defaultDestUnreachable() != 3)
+    return false;
+
+  // Each result must go to a different block.
+  // Make the later PHI update simpler.
+  SmallPtrSet<BasicBlock *, 4> Succs(from_range, successors(SI));
+  if (Succs.size() != SI->getNumSuccessors())
+    return false;
+
+  BasicBlock *SwitchBB = SI->getParent();
+  BasicBlock *Default = SI->getDefaultDest();
+
+  // Each arm is one comparison result
+  struct Arm {
+    BasicBlock *Dest;
+    CmpInst::Predicate Pred;
+    BranchProbability Prob;
+    uint32_t Weight;
+  };
+  SmallVector<Arm, 3> Arms;
+  SmallVector<uint32_t, 4> Weights;
+  bool HasWeights =
+      extractBranchWeights(getValidBranchWeightMDNode(*SI), Weights);
+  Weights.resize(4);
+  bool ExpectedWeights = HasWeights && hasBranchWeightOrigin(*SI);
+  auto AddArm = [&](BasicBlock *Dest, int8_t Result, unsigned Index) {
+    CmpInst::Predicate Pred = Result < 0   ? Cmp->getLTPredicate()
+                              : Result > 0 ? Cmp->getGTPredicate()
+                                           : CmpInst::ICMP_EQ;
+    Arms.push_back({Dest, Pred, BPI->getEdgeProbability(SwitchBB, Index),
+                     Weights[Index]});
+  };
+
+  // The results -1, 0 and 1 sum to zero, so with two cases the default takes
+  // the negated sum of the case values.
+  int8_t Sum = 0;
+  for (auto &Case : SI->cases()) {
+    std::optional<int64_t> Val = Case.getCaseValue()->getValue().trySExtValue();
+    if (!Val || *Val < -1 || *Val > 1)
+      return false;
+    Sum += *Val;
+    AddArm(Case.getCaseSuccessor(), *Val, Case.getSuccessorIndex());
+  }
+
+  if (NumCases == 2)
+    AddArm(Default, -Sum, 0);
+
+  // Test the most likely arm first.
+  stable_sort(Arms, [](const Arm &A, const Arm &B) { return A.Prob > B.Prob; });
+
+
+  // Check if any BBs contain a loop backedge.
+  MDNode *LoopMD = SI->getMetadata(LLVMContext::MD_loop);
+  auto IsBackedge = [&](const Arm &A) {
+    Loop *L = LI->getLoopFor(A.Dest);
+    return L && L->getHeader() == A.Dest && L->contains(SwitchBB);
+  };
+
+  MDNode *Unpredictable = SI->getMetadata(LLVMContext::MD_unpredictable);
+
+  IRBuilder<> Builder(SI);
+  auto CreateCheck = [&](const Arm &A, BasicBlock *FalseDest,
+                         uint64_t FalseWeight, bool IsLatch) {
+    Value *Cond = Builder.CreateICmp(A.Pred, Cmp->getLHS(), Cmp->getRHS());
+    auto *Br =
+        Builder.CreateCondBr(Cond, A.Dest, FalseDest, nullptr, Unpredictable);
+    if (HasWeights)
+      setFittedBranchWeights(*Br, {A.Weight, FalseWeight}, ExpectedWeights);
+    if (IsLatch)
+      Br->setMetadata(LLVMContext::MD_loop, LoopMD);
+  };
+  BasicBlock *Next =
+      BasicBlock::Create(SwitchBB->getContext(), "cmp.next",
+                         SwitchBB->getParent(), SwitchBB->getNextNode());
+  // The first icmp branch to mostly likely successor.
+  // Other two go to new Next branch.
+  CreateCheck(Arms[0], Next, uint64_t(Arms[1].Weight) + Arms[2].Weight,
+              IsBackedge(Arms[0]));
+  Builder.SetInsertPoint(Next);
+  CreateCheck(Arms[1], Arms[2].Dest, Arms[2].Weight,
+              IsBackedge(Arms[1]) || IsBackedge(Arms[2]));
+
+  // With three cases the unreachable default is no longer a successor.
+  SmallVector<DominatorTree::UpdateType, 6> Updates = {
+      {DominatorTree::Insert, SwitchBB, Next}};
+  if (NumCases == 3) {
+    Default->removePredecessor(SwitchBB);
+    Updates.push_back({DominatorTree::Delete, SwitchBB, Default});
+  }
+
+  // The first arm keeps its edge from SwitchBB. The other two move to Next.
+  for (const Arm &A : drop_begin(Arms)) {
+    A.Dest->replacePhiUsesWith(SwitchBB, Next);
+    Updates.push_back({DominatorTree::Delete, SwitchBB, A.Dest});
+    Updates.push_back({DominatorTree::Insert, Next, A.Dest});
+  }
+  SI->eraseFromParent();
+  Cmp->eraseFromParent();
+  DTU->applyUpdates(Updates);
+
+  // Split the BPI probability over all branches, freqencies are unchanged.
+  SmallVector<BranchProbability, 3> Probs = {Arms[0].Prob, Arms[1].Prob,
+                                             Arms[2].Prob};
+  BranchProbability::normalizeProbabilities(Probs);
+  BranchProbability RestProb = Probs[0].getCompl();
+  BPI->setEdgeProbability(SwitchBB, {Probs[0], RestProb});
+  SmallVector<BranchProbability, 2> NextProbs = {Probs[1], Probs[2]};
+  BranchProbability::normalizeProbabilities(NextProbs);
+  BPI->setEdgeProbability(Next, NextProbs);
+  BFI->setBlockFreq(Next, BFI->getBlockFreq(SwitchBB) * RestProb);
+
+  // Next belongs to the innermost loop that contains SwitchBB and at least one
+  // of Next's successors. If both successors leave SwitchBB's loop, Next is an
+  // exiting path and sits outside it.
+  Loop *L = LI->getLoopFor(SwitchBB);
+  while (L && !L->contains(Arms[1].Dest) && !L->contains(Arms[2].Dest))
+    L = L->getParentLoop();
+  if (L)
+    L->addBasicBlockToLoop(Next, *LI);
+
+  // Next BB to explore
+  if (IsHugeFunc)
+    FreshBBs.insert(Next);
+  return true;
+}
+
 bool CodeGenPrepare::optimizeSwitchType(SwitchInst *SI) {
   Value *Cond = SI->getCondition();
   Type *OldType = Cond->getType();
@@ -8190,7 +8327,12 @@ bool CodeGenPrepare::optimizeSwitchPhiConstants(SwitchInst *SI) {
   return Changed;
 }
 
-bool CodeGenPrepare::optimizeSwitchInst(SwitchInst *SI) {
+bool CodeGenPrepare::optimizeSwitchInst(SwitchInst *SI, ModifyDT &ModifiedDT) {
+  // Lower a comparison switch before optimizeSwitchType widens its condition.
+  if (optimizeSwitchOfCmpIntrinsic(SI)) {
+    ModifiedDT = ModifyDT::ModifyBBDT;
+    return true;
+  }
   bool Changed = optimizeSwitchType(SI);
   Changed |= optimizeSwitchPhiConstants(SI);
   return Changed;
@@ -9125,7 +9267,7 @@ bool CodeGenPrepare::optimizeInst(Instruction *I, ModifyDT &ModifiedDT) {
   case Instruction::ShuffleVector:
     return optimizeShuffleVectorInst(cast<ShuffleVectorInst>(I));
   case Instruction::Switch:
-    return optimizeSwitchInst(cast<SwitchInst>(I));
+    return optimizeSwitchInst(cast<SwitchInst>(I), ModifiedDT);
   case Instruction::ExtractElement:
     return optimizeExtractElementInst(cast<ExtractElementInst>(I));
   case Instruction::CondBr:
@@ -9167,7 +9309,7 @@ bool CodeGenPrepare::optimizeBlock(BasicBlock &BB, ModifyDT &ModifiedDT) {
     while (CurInstIterator != BB.end()) {
       MadeChange |= optimizeInst(&*CurInstIterator++, ModifiedDT);
       if (ModifiedDT != ModifyDT::NotModifyDT) {
-        // For huge function we tend to quickly go though the inner optmization
+        // For huge function we tend to quickly go though the inner optimization
         // opportunities in the BB. So we go back to the BB head to re-optimize
         // each instruction instead of go back to the function head.
         if (IsHugeFunc)
diff --git a/llvm/test/Transforms/CodeGenPrepare/switch-on-cmp.ll b/llvm/test/Transforms/CodeGenPrepare/switch-on-cmp.ll
new file mode 100644
index 0000000000000..ad2fcf9f6b76a
--- /dev/null
+++ b/llvm/test/Transforms/CodeGenPrepare/switch-on-cmp.ll
@@ -0,0 +1,256 @@
+; NOTE: Assertions have been autogenerated by utils/update_test_checks.py
+; RUN: opt -S -codegenprepare -mtriple=x86_64 < %s | FileCheck %s
+; REQUIRES: x86-registered-target
+
+define void @ucmp_gt_not_same_succ(i32 %a, i32 %b) {
+; CHECK-LABEL: define void @ucmp_gt_not_same_succ(
+; CHECK-SAME: i32 [[A:%.*]], i32 [[B:%.*]]) {
+; CHECK-NEXT:    [[TMP1:%.*]] = icmp ult i32 [[A]], [[B]]
+; CHECK-NEXT:    br i1 [[TMP1]], label %[[BB2:.*]], label %[[CMP_NEXT:.*]]
+; CHECK:       [[CMP_NEXT]]:
+; CHECK-NEXT:    [[TMP2:%.*]] = icmp eq i32 [[A]], [[B]]
+; CHECK-NEXT:    br i1 [[TMP2]], label %[[BB3:.*]], label %[[BB1:.*]]
+; CHECK:       [[BB1]]:
+; CHECK-NEXT:    call void @foo()
+; CHECK-NEXT:    br label %[[BB2]]
+; CHECK:       [[BB3]]:
+; CHECK-NEXT:    call void @foo()
+; CHECK-NEXT:    br label %[[BB2]]
+; CHECK:       [[BB2]]:
+; CHECK-NEXT:    ret void
+;
+  %res = call i8 @llvm.ucmp.i8.i32(i32 %a, i32 %b)
+  switch i8 %res, label %bb1 [
+  i8 -1, label %bb2
+  i8 0, label %bb3
+  ]
+
+bb1:
+  call void @foo()
+  br label %bb2
+
+bb3:
+  call void @foo()
+  br label %bb2
+
+bb2:
+  ret void
+}
+
+define void @ucmp_gt_unreachable_no_two_equal_cases(i32 %a, i32 %b) {
+; CHECK-LABEL: define void @ucmp_gt_unreachable_no_two_equal_cases(
+; CHECK-SAME: i32 [[A:%.*]], i32 [[B:%.*]]) {
+; CHECK-NEXT:    [[TMP1:%.*]] = icmp ult i32 [[A]], [[B]]
+; CHECK-NEXT:    br i1 [[TMP1]], label %[[BB3:.*]], label %[[CMP_NEXT:.*]]
+; CHECK:       [[CMP_NEXT]]:
+; CHECK-NEXT:    [[TMP2:%.*]] = icmp eq i32 [[A]], [[B]]
+; CHECK-NEXT:    br i1 [[TMP2]], label %[[BB2:.*]], label %[[BB1:.*]]
+; CHECK:       [[BB1]]:
+; CHECK-NEXT:    call void @foo()
+; CHECK-NEXT:    br label %[[BB2]]
+; CHECK:       [[BB3]]:
+; CHECK-NEXT:    call void @foo()
+; CHECK-NEXT:    br label %[[BB2]]
+; CHECK:       [[BB2]]:
+; CHECK-NEXT:    ret void
+; CHECK:       [[UNREACHABLE:.*:]]
+; CHECK-NEXT:    unreachable
+;
+  %res = call i8 @llvm.ucmp.i8.i32(i32 %a, i32 %b)
+  switch i8 %res, label %unreachable [
+  i8 -1, label %bb3
+  i8 0, label %bb2
+  i8 1, label %bb1
+  ]
+
+bb1:
+  call void @foo()
+  br label %bb2
+
+bb3:
+  call void @foo()
+  br label %bb2
+
+bb2:
+  ret void
+
+unreachable:
+  unreachable
+}
+
+define i64 @ucmp_three_destinations_phi_select(i8 %x, i8 %y) {
+; CHECK-LABEL: define i64 @ucmp_three_destinations_phi_select(
+; CHECK-SAME: i8 [[X:%.*]], i8 [[Y:%.*]]) {
+; CHECK-NEXT:  [[START:.*:]]
+; CHECK-NEXT:    [[TMP0:%.*]] = icmp ult i8 [[X]], [[Y]]
+; CHECK-NEXT:    br i1 [[TMP0]], label %[[LT:.*]], label %[[CMP_NEXT:.*]]
+; CHECK:       [[CMP_NEXT]]:
+; CHECK-NEXT:    [[TMP1:%.*]] = icmp eq i8 [[X]], [[Y]]
+; CHECK-NEXT:    br i1 [[TMP1]], label %[[END:.*]], label %[[GT:.*]]
+; CHECK:       [[UNR:.*:]]
+; CHECK-NEXT:    unreachable
+; CHECK:       [[LT]]:
+; CHECK-NEXT:    [[Z:%.*]] = zext i8 [[X]] to i64
+; CHECK-NEXT:    br label %[[END]]
+; CHECK:       [[GT]]:
+; CHECK-NEXT:    br label %[[END]]
+; CHECK:       [[END]]:
+; CHECK-NEXT:    [[R:%.*]] = phi i64 [ [[Z]], %[[LT]] ], [ 2, %[[GT]] ], [ 4, %[[CMP_NEXT]] ]
+; CHECK-NEXT:    ret i64 [[R]]
+;
+start:
+  %c = call i8 @llvm.ucmp.i8.i8(i8 %x, i8 %y)
+  switch i8 %c, label %unr [
+  i8 -1, label %lt
+  i8 0, label %end
+  i8 1, label %gt
+  ]
+
+unr:
+  unreachable
+
+lt:
+  %z = zext i8 %x to i64
+  br label %end
+
+gt:
+  br label %end
+
+end:
+  %r = phi i64 [ %z, %lt ], [ 2, %gt ], [ 4, %start ]
+  ret i64 %r
+}
+
+define void @three_destinations_weighted_signed(i32 %a, i32 %b) {
+; CHECK-LABEL: define void @three_destinations_weighted_signed(
+; CHECK-SAME: i32 [[A:%.*]], i32 [[B:%.*]]) {
+; CHECK-NEXT:    [[TMP1:%.*]] = icmp slt i32 [[A]], [[B]]
+; CHECK-NEXT:    br i1 [[TMP1]], label %[[LESS:.*]], label %[[CMP_NEXT:.*]], !prof [[PROF0:![0-9]+]]
+; CHECK:       [[CMP_NEXT]]:
+; CHECK-NEXT:    [[TMP2:%.*]] = icmp sgt i32 [[A]], [[B]]
+; CHECK-NEXT:    br i1 [[TMP2]], label %[[GREATER:.*]], label %[[EQUAL:.*]], !prof [[PROF1:![0-9]+]]
+; CHECK:       [[EQUAL]]:
+; CHECK-NEXT:    call void @equal()
+; CHECK-NEXT:    ret void
+; CHECK:       [[GREATER]]:
+; CHECK-NEXT:    call void @greater()
+; CHECK-NEXT:    ret void
+; CHECK:       [[LESS]]:
+; CHECK-NEXT:    call void @less()
+; CHECK-NEXT:    ret void
+; CHECK:       [[DEAD:.*:]]
+; CHECK-NEXT:    unreachable
+;
+  %cmp = call i8 @llvm.scmp.i8.i32(i32 %a, i32 %b)
+  switch i8 %cmp, label %dead [
+  i8 0, label %equal
+  i8 1, label %greater
+  i8 -1, label %less
+  ], !prof !{!"branch_weights", i32 0, i32 10, i32 30, i32 60}
+equal:
+  call void @equal()
+  ret void
+greater:
+  call void @greater()
+  ret void
+less:
+  call void @less()
+  ret void
+dead:
+  unreachable
+}
+
+declare void @use(i8)
+declare void @foo()
+declare void @equal()
+declare void @greater()
+declare void @less()
+
+; Exercise the second comparison block on a loop backedge. The weights make the
+; exit to %less the first check, so the backedge and its loop metadata move to
+; cmp.next.
+define i32 @ucmp_loop_default_phi(i32 %a, i32 %b) {
+; CHECK-LABEL: define i32 @ucmp_loop_default_phi(
+; CHECK-SAME: i32 [[A:%.*]], i32 [[B:%.*]]) {
+; CHECK-NEXT:  [[ENTRY:.*]]:
+; CHECK-NEXT:    br label %[[LOOP:.*]]
+; CHECK:       [[LOOP]]:
+; CHECK-NEXT:    [[X:%.*]] = phi i32 [ [[A]], %[[ENTRY]] ], [ [[NEXT:%.*]], %[[CMP_NEXT:.*]] ]
+; CHECK-NEXT:    [[NEXT]] = add i32 [[X]], 1
+; CHECK-NEXT:    [[TMP0:%.*]] = icmp ult i32 [[NEXT]], [[B]]
+; CHECK-NEXT:    br i1 [[TMP0]], label %[[LESS:.*]], label %[[CMP_NEXT]], !prof [[PROF2:![0-9]+]]
+; CHECK:       [[CMP_NEXT]]:
+; CHECK-NEXT:    [[TMP1:%.*]] = icmp ugt i32 [[NEXT]], [[B]]
+; CHECK-NEXT:    br i1 [[TMP1]], label %[[LOOP]], label %[[EQUAL:.*]], !prof [[PROF3:![0-9]+]], !llvm.loop [[LOOP4:![0-9]+]]
+; CHECK:       [[LESS]]:
+; CHECK-NEXT:    call void @less()
+; CHECK-NEXT:    ret i32 [[NEXT]]
+; CHECK:       [[EQUAL]]:
+; CHECK-NEXT:    call void @equal()
+; CHECK-NEXT:    ret i32 [[X]]
+;
+entry:
+  br label %loop
+loop:
+  %x = phi i32 [ %a, %entry ], [ %next, %loop ]
+  %next = add i32 %x, 1
+  %cmp = call i8 @llvm.ucmp.i8.i32(i32 %next, i32 %b)
+  switch i8 %cmp, label %loop [
+  i8 -1, label %less
+  i8 0, label %equal
+  ], !prof !{!"branch_weights", i32 10, i32 100, i32 1}, !llvm.loop !0
+less:
+  call void @less()
+  ret i32 %next
+equal:
+  call void @equal()
+  ret i32 %x
+}
+
+; Without profile data, BPI still orders the checks: the arm calling a cold
+; function is tested last.
+define void @ucmp_cold_arm_last(i32 %a, i32 %b) {
+; CHECK-LABEL: define void @ucmp_cold_arm_last(
+; CHECK-SAME: i32 [[A:%.*]], i32 [[B:%.*]]) {
+; CHECK-NEXT:    [[TMP1:%.*]] = icmp eq i32 [[A]], [[B]]
+; CHECK-NEXT:    br i1 [[TMP1]], label %[[EQUAL:.*]], label %[[CMP_NEXT:.*]]
+; CHECK:       [[CMP_NEXT]]:
+; CHECK-NEXT:    [[TMP2:%.*]] = icmp ugt i32 [[A]], [[B]]
+; CHECK-NEXT:    br i1 [[TMP2]], label %[[GREATER:.*]], label %[[LESS:.*]]
+; CHECK:       [[LESS]]:
+; CHECK-NEXT:    call void @cold_fn()
+; CHECK-NEXT:    ret void
+; CHECK:       [[EQUAL]]:
+; CHECK-NEXT:    call void @equal()
+; CHECK-NEXT:    ret void
+; CHECK:       [[GREATER]]:
+; CHECK-NEXT:    call void @greater()
+; CHECK-NEXT:    ret void
+;
+  %cmp = call i8 @llvm.ucmp.i8.i32(i32 %a, i32 %b)
+  switch i8 %cmp, label %greater [
+  i8 -1, label %less
+  i8 0, label %equal
+  ]
+less:
+  call void @cold_fn()
+  ret void
+equal:
+  call void @equal()
+  ret void
+greater:
+  call void @greater()
+  ret void
+}
+
+declare void @cold_fn() cold
+
+!0 = distinct !{!0}
+
+;.
+; CHECK: [[PROF0]] = !{!"branch_weights", i32 60, i32 40}
+; CHECK: [[PROF1]] = !{!"branch_weights", i32 30, i32 10}
+; CHECK: [[PROF2]] = !{!"branch_weights", i32 100, i32 11}
+; CHECK: [[PROF3]] = !{!"branch_weights", i32 10, i32 1}
+; CHECK: [[LOOP4]] = distinct !{[[LOOP4]]}
+;.

``````````

</details>


https://github.com/llvm/llvm-project/pull/223268


More information about the llvm-commits mailing list