[OPT] Search whole BB for convergence token. (#112728)

The spec for llvm.experimental.convergence.entry says that is must be in
the entry block for a function, and must preceed any other convergent
operation. It does not have to be the first instruction in the entry
block.

Inlining assumes that the call to llvm.experimental.convergence.entry
will be the first instruction after any phi instructions. This commit
modifies inlining to search the entire block for the call.
This commit is contained in:
Steven Perron 2024-10-30 11:19:23 -04:00 committed by GitHub
parent 4015e18d67
commit f405c683ba
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 45 additions and 17 deletions

View File

@ -181,9 +181,21 @@ namespace {
}
}
};
} // end anonymous namespace
static IntrinsicInst *getConvergenceEntry(BasicBlock &BB) {
auto *I = BB.getFirstNonPHI();
while (I) {
if (auto *IntrinsicCall = dyn_cast<ConvergenceControlInst>(I)) {
if (IntrinsicCall->isEntry()) {
return IntrinsicCall;
}
}
I = I->getNextNode();
}
return nullptr;
}
/// Get or create a target for the branch from ResumeInsts.
BasicBlock *LandingPadInliningInfo::getInnerResumeDest() {
if (InnerResumeDest) return InnerResumeDest;
@ -2496,15 +2508,10 @@ llvm::InlineResult llvm::InlineFunction(CallBase &CB, InlineFunctionInfo &IFI,
// fully implements convergence control tokens, there is no mixing of
// controlled and uncontrolled convergent operations in the whole program.
if (CB.isConvergent()) {
auto *I = CalledFunc->getEntryBlock().getFirstNonPHI();
if (auto *IntrinsicCall = dyn_cast<IntrinsicInst>(I)) {
if (IntrinsicCall->getIntrinsicID() ==
Intrinsic::experimental_convergence_entry) {
if (!ConvergenceControlToken) {
return InlineResult::failure(
"convergent call needs convergencectrl operand");
}
}
if (!ConvergenceControlToken &&
getConvergenceEntry(CalledFunc->getEntryBlock())) {
return InlineResult::failure(
"convergent call needs convergencectrl operand");
}
}
@ -2795,13 +2802,10 @@ llvm::InlineResult llvm::InlineFunction(CallBase &CB, InlineFunctionInfo &IFI,
}
if (ConvergenceControlToken) {
auto *I = FirstNewBlock->getFirstNonPHI();
if (auto *IntrinsicCall = dyn_cast<IntrinsicInst>(I)) {
if (IntrinsicCall->getIntrinsicID() ==
Intrinsic::experimental_convergence_entry) {
IntrinsicCall->replaceAllUsesWith(ConvergenceControlToken);
IntrinsicCall->eraseFromParent();
}
IntrinsicInst *IntrinsicCall = getConvergenceEntry(*FirstNewBlock);
if (IntrinsicCall) {
IntrinsicCall->replaceAllUsesWith(ConvergenceControlToken);
IntrinsicCall->eraseFromParent();
}
}

View File

@ -185,6 +185,30 @@ define void @test_two_calls() convergent {
ret void
}
define i32 @token_not_first(i32 %x) convergent alwaysinline {
; CHECK-LABEL: @token_not_first(
; CHECK-NEXT: {{%.*}} = alloca ptr, align 8
; CHECK-NEXT: [[TOKEN:%.*]] = call token @llvm.experimental.convergence.entry()
; CHECK-NEXT: [[Y:%.*]] = call i32 @g(i32 [[X:%.*]]) [ "convergencectrl"(token [[TOKEN]]) ]
; CHECK-NEXT: ret i32 [[Y]]
;
%p = alloca ptr, align 8
%token = call token @llvm.experimental.convergence.entry()
%y = call i32 @g(i32 %x) [ "convergencectrl"(token %token) ]
ret i32 %y
}
define void @test_token_not_first() convergent {
; CHECK-LABEL: @test_token_not_first(
; CHECK-NEXT: [[TOKEN:%.*]] = call token @llvm.experimental.convergence.entry()
; CHECK-NEXT: {{%.*}} = call i32 @g(i32 23) [ "convergencectrl"(token [[TOKEN]]) ]
; CHECK-NEXT: ret void
;
%token = call token @llvm.experimental.convergence.entry()
%x = call i32 @token_not_first(i32 23) [ "convergencectrl"(token %token) ]
ret void
}
declare void @f(i32) convergent
declare i32 @g(i32) convergent