[llvm] [AMDGPU] Look through bitcast when skipping musttail returns in UnifyDivergentExitNodes (PR #201021)

via llvm-commits llvm-commits at lists.llvm.org
Mon Jun 1 22:58:43 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-backend-amdgpu

Author: Arseniy Obolenskiy (aobolensk)

<details>
<summary>Changes</summary>

The musttail check only inspected the ret's immediate predecessor, missing the optional bitcast the verifier permits between a musttail call and its ret

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


2 Files Affected:

- (modified) llvm/lib/Target/AMDGPU/AMDGPUUnifyDivergentExitNodes.cpp (+2-3) 
- (modified) llvm/test/CodeGen/AMDGPU/do-not-unify-divergent-exit-nodes-with-musttail.ll (+27) 


``````````diff
diff --git a/llvm/lib/Target/AMDGPU/AMDGPUUnifyDivergentExitNodes.cpp b/llvm/lib/Target/AMDGPU/AMDGPUUnifyDivergentExitNodes.cpp
index 0c012681d6d6b..b7f2ee33485ab 100644
--- a/llvm/lib/Target/AMDGPU/AMDGPUUnifyDivergentExitNodes.cpp
+++ b/llvm/lib/Target/AMDGPU/AMDGPUUnifyDivergentExitNodes.cpp
@@ -253,9 +253,8 @@ bool AMDGPUUnifyDivergentExitNodesImpl::run(Function &F, DominatorTree *DT,
 
   for (BasicBlock *BB : PDT.roots()) {
     Instruction *Term = BB->getTerminator();
-    if (auto *RI = dyn_cast<ReturnInst>(Term)) {
-      auto *CI = dyn_cast_or_null<CallInst>(RI->getPrevNode());
-      if (CI && CI->isMustTailCall())
+    if (isa<ReturnInst>(Term)) {
+      if (BB->getTerminatingMustTailCall())
         continue;
       if (HasDivergentExitBlock)
         ReturningBlocks.push_back(BB);
diff --git a/llvm/test/CodeGen/AMDGPU/do-not-unify-divergent-exit-nodes-with-musttail.ll b/llvm/test/CodeGen/AMDGPU/do-not-unify-divergent-exit-nodes-with-musttail.ll
index 076a99ff8588f..2eb7712d6dcdd 100644
--- a/llvm/test/CodeGen/AMDGPU/do-not-unify-divergent-exit-nodes-with-musttail.ll
+++ b/llvm/test/CodeGen/AMDGPU/do-not-unify-divergent-exit-nodes-with-musttail.ll
@@ -4,6 +4,7 @@
 declare void @foo(ptr)
 declare i1 @bar(ptr)
 declare i32 @bar32(ptr)
+declare ptr @bar_ptr(ptr)
 
 define void @musttail_call_without_return_value(ptr %p) {
 ; CHECK-LABEL: define void @musttail_call_without_return_value(
@@ -102,3 +103,29 @@ bb.0:
 bb.1:
   ret i32 %load
 }
+
+define ptr @musttail_call_with_bitcast(ptr %p) {
+; CHECK-LABEL: define ptr @musttail_call_with_bitcast(
+; CHECK-SAME: ptr [[P:%.*]]) #[[ATTR0]] {
+; CHECK-NEXT:  [[ENTRY:.*:]]
+; CHECK-NEXT:    [[LOAD:%.*]] = load i1, ptr [[P]], align 1
+; CHECK-NEXT:    br i1 [[LOAD]], label %[[BB_0:.*]], label %[[BB_1:.*]]
+; CHECK:       [[BB_0]]:
+; CHECK-NEXT:    [[RET:%.*]] = musttail call ptr @bar_ptr(ptr [[P]])
+; CHECK-NEXT:    [[RETC:%.*]] = bitcast ptr [[RET]] to ptr
+; CHECK-NEXT:    ret ptr [[RETC]]
+; CHECK:       [[BB_1]]:
+; CHECK-NEXT:    ret ptr null
+;
+entry:
+  %load = load i1, ptr %p, align 1
+  br i1 %load, label %bb.0, label %bb.1
+
+bb.0:
+  %ret = musttail call ptr @bar_ptr(ptr %p)
+  %retc = bitcast ptr %ret to ptr
+  ret ptr %retc
+
+bb.1:
+  ret ptr null
+}

``````````

</details>


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


More information about the llvm-commits mailing list