[llvm] [NVPTX] Simplify BRX emission avoiding pseduo-instruction chain (PR #209952)
Alex MacLean via llvm-commits
llvm-commits at lists.llvm.org
Thu Jul 16 09:42:30 PDT 2026
https://github.com/AlexMaclean updated https://github.com/llvm/llvm-project/pull/209952
>From 45ec95f26e88d1cc37635a2a818ee78b0ae41b6f Mon Sep 17 00:00:00 2001
From: Alex Maclean <amaclean at nvidia.com>
Date: Wed, 15 Jul 2026 20:34:14 -0700
Subject: [PATCH] [NVPTX] Simplify BRX emission avoiding pseduo-instruction
chain
---
llvm/lib/Target/NVPTX/NVPTXAsmPrinter.cpp | 15 ++++++++
llvm/lib/Target/NVPTX/NVPTXISelDAGToDAG.cpp | 38 -------------------
llvm/lib/Target/NVPTX/NVPTXISelDAGToDAG.h | 1 -
llvm/lib/Target/NVPTX/NVPTXInstrInfo.td | 25 ++++++------
llvm/test/CodeGen/NVPTX/jump-table.ll | 16 +-------
.../test/CodeGen/NVPTX/switch-loop-header.mir | 18 +--------
llvm/test/CodeGen/NVPTX/switch.ll | 10 +----
7 files changed, 32 insertions(+), 91 deletions(-)
diff --git a/llvm/lib/Target/NVPTX/NVPTXAsmPrinter.cpp b/llvm/lib/Target/NVPTX/NVPTXAsmPrinter.cpp
index a22e2cd05a931..7c61c415737e8 100644
--- a/llvm/lib/Target/NVPTX/NVPTXAsmPrinter.cpp
+++ b/llvm/lib/Target/NVPTX/NVPTXAsmPrinter.cpp
@@ -46,6 +46,7 @@
#include "llvm/CodeGen/MachineFrameInfo.h"
#include "llvm/CodeGen/MachineFunction.h"
#include "llvm/CodeGen/MachineInstr.h"
+#include "llvm/CodeGen/MachineJumpTableInfo.h"
#include "llvm/CodeGen/MachineLoopInfo.h"
#include "llvm/CodeGen/MachineModuleInfo.h"
#include "llvm/CodeGen/MachineOperand.h"
@@ -210,6 +211,11 @@ MCOperand NVPTXAsmPrinter::lowerOperand(const MachineOperand &MO) {
MCSymbolRefExpr::create(MO.getMBB()->getSymbol(), OutContext));
case MachineOperand::MO_ExternalSymbol:
return GetSymbolRef(GetExternalSymbolSymbol(MO.getSymbolName()));
+ case MachineOperand::MO_JumpTableIndex:
+ // The jump table index names the .branchtargets list emitted for a brx.idx
+ // (see emitFunctionBodyStart); reference it by that label.
+ return GetSymbolRef(
+ OutContext.getOrCreateSymbol("$L_brx_" + Twine(MO.getIndex())));
case MachineOperand::MO_GlobalAddress:
return GetSymbolRef(getSymbol(MO.getGlobal()));
case MachineOperand::MO_FPImmediate: {
@@ -515,6 +521,15 @@ void NVPTXAsmPrinter::emitFunctionBodyStart() {
for (const auto &[Id, CB] : MFI->getCallPrototypes())
emitCallPrototype(*CB, Id, O);
+ if (const MachineJumpTableInfo *MJTI = MF->getJumpTableInfo())
+ for (const auto &[Idx, JT] : enumerate(MJTI->getJumpTables())) {
+ O << "$L_brx_" << Idx << ": .branchtargets ";
+ interleaveComma(JT.MBBs, O, [&](const MachineBasicBlock *MBB) {
+ MBB->getSymbol()->print(O, MAI);
+ });
+ O << ";\n";
+ }
+
OutStreamer->emitRawText(O.str());
}
diff --git a/llvm/lib/Target/NVPTX/NVPTXISelDAGToDAG.cpp b/llvm/lib/Target/NVPTX/NVPTXISelDAGToDAG.cpp
index ede1deb5400b0..7ffbcff8b5141 100644
--- a/llvm/lib/Target/NVPTX/NVPTXISelDAGToDAG.cpp
+++ b/llvm/lib/Target/NVPTX/NVPTXISelDAGToDAG.cpp
@@ -16,7 +16,6 @@
#include "llvm/ADT/APInt.h"
#include "llvm/Analysis/ValueTracking.h"
#include "llvm/CodeGen/ISDOpcodes.h"
-#include "llvm/CodeGen/MachineJumpTableInfo.h"
#include "llvm/CodeGen/SelectionDAG.h"
#include "llvm/CodeGen/SelectionDAGNodes.h"
#include "llvm/IR/GlobalValue.h"
@@ -191,8 +190,6 @@ void NVPTXDAGToDAGISel::Select(SDNode *N) {
if (tryBF16ArithToFMA(N))
return;
break;
- case ISD::BR_JT:
- return selectBR_JT(N);
default:
break;
}
@@ -2284,38 +2281,3 @@ void NVPTXDAGToDAGISel::selectAtomicSwap128(SDNode *N) {
ReplaceNode(N, ATOM);
}
-
-void NVPTXDAGToDAGISel::selectBR_JT(SDNode *N) {
- assert(Subtarget->hasBrx() &&
- "BR_JT should be expanded during legalization on unsupported targets");
-
- SDLoc DL(N);
- const SDValue InChain = N->getOperand(0);
- const auto *JT = cast<JumpTableSDNode>(N->getOperand(1));
- const SDValue Index = N->getOperand(2);
-
- unsigned JId = JT->getIndex();
- MachineJumpTableInfo *MJTI = CurDAG->getMachineFunction().getJumpTableInfo();
- ArrayRef<MachineBasicBlock *> MBBs = MJTI->getJumpTables()[JId].MBBs;
-
- SDValue IdV = getI32Imm(JId, DL);
-
- // Generate BrxStart node
- MachineSDNode *Chain = CurDAG->getMachineNode(
- NVPTX::BRX_START, DL, {MVT::Other, MVT::Glue}, {IdV, InChain});
-
- // Generate BrxItem nodes
- assert(!MBBs.empty());
- for (MachineBasicBlock *MBB : MBBs.drop_back())
- Chain = CurDAG->getMachineNode(
- NVPTX::BRX_ITEM, DL, {MVT::Other, MVT::Glue},
- {CurDAG->getBasicBlock(MBB), SDValue(Chain, 0), SDValue(Chain, 1)});
-
- // Generate BrxEnd nodes
- MachineSDNode *BrxEnd =
- CurDAG->getMachineNode(NVPTX::BRX_END, DL, MVT::Other,
- {CurDAG->getBasicBlock(MBBs.back()), Index, IdV,
- SDValue(Chain, 0), SDValue(Chain, 1)});
-
- ReplaceNode(N, BrxEnd);
-}
diff --git a/llvm/lib/Target/NVPTX/NVPTXISelDAGToDAG.h b/llvm/lib/Target/NVPTX/NVPTXISelDAGToDAG.h
index fcb5700dcb6d4..055f1ff47306d 100644
--- a/llvm/lib/Target/NVPTX/NVPTXISelDAGToDAG.h
+++ b/llvm/lib/Target/NVPTX/NVPTXISelDAGToDAG.h
@@ -93,7 +93,6 @@ class LLVM_LIBRARY_VISIBILITY NVPTXDAGToDAGISel : public SelectionDAGISel {
void SelectTcgen05Ld(SDNode *N, bool hasOffset = false);
void SelectTcgen05St(SDNode *N, bool hasOffset = false);
void selectAtomicSwap128(SDNode *N);
- void selectBR_JT(SDNode *N);
inline SDValue getI32Imm(unsigned Imm, const SDLoc &DL) {
return CurDAG->getTargetConstant(Imm, DL, MVT::i32);
diff --git a/llvm/lib/Target/NVPTX/NVPTXInstrInfo.td b/llvm/lib/Target/NVPTX/NVPTXInstrInfo.td
index dcf89f9a73956..bab4939932079 100644
--- a/llvm/lib/Target/NVPTX/NVPTXInstrInfo.td
+++ b/llvm/lib/Target/NVPTX/NVPTXInstrInfo.td
@@ -1777,6 +1777,10 @@ def to_tframeindex : SDNodeXForm<frameindex, [{
return CurDAG->getTargetFrameIndex(N->getIndex(), N->getValueType(0));
}]>;
+def to_tjumptable : SDNodeXForm<jumptable, [{
+ return CurDAG->getTargetJumpTable(N->getIndex(), N->getValueType(0));
+}]>;
+
def : Pat<(i32 globaladdr:$dst), (MOV_B32_sym (to_tglobaladdr $dst))>;
def : Pat<(i64 globaladdr:$dst), (MOV_B64_sym (to_tglobaladdr $dst))>;
@@ -2493,21 +2497,16 @@ foreach t = [I32RT, I64RT] in {
//
// BRX
//
+def brx_jt : SDNode<"ISD::BR_JT",
+ SDTypeProfile<0, 2, [SDTCisPtrTy<0>, SDTCisVT<1, i32>]>,
+ [SDNPHasChain]>;
-let isTerminator = 1, isBranch = 1, isIndirectBranch = 1, isNotDuplicable = 1 in {
-
- def BRX_START :
- NVPTXInst<(outs), (ins i32imm:$id), "$$L_brx_$id: .branchtargets">;
+let isTerminator = 1, isBarrier = 1, isBranch = 1, isIndirectBranch = 1,
+ isNotDuplicable = 1 in
+ def BRX_IDX : BasicNVPTXInst<(outs), (ins B32:$val, i32imm:$tgt),
+ "brx.idx">;
- def BRX_ITEM :
- NVPTXInst<(outs), (ins brtarget:$target), "\t$target,">;
-
- def BRX_END :
- NVPTXInst<(outs), (ins brtarget:$target, B32:$val, i32imm:$id),
- "\t$target;\n\tbrx.idx \t$val, $$L_brx_$id;"> {
- let isBarrier = 1;
- }
-}
+def : Pat<(brx_jt jumptable:$jt, i32:$idx), (BRX_IDX $idx, (to_tjumptable $jt))>;
foreach a_type = ["s", "u"] in {
diff --git a/llvm/test/CodeGen/NVPTX/jump-table.ll b/llvm/test/CodeGen/NVPTX/jump-table.ll
index 8f253f5861aea..b18ca2bc9af64 100644
--- a/llvm/test/CodeGen/NVPTX/jump-table.ll
+++ b/llvm/test/CodeGen/NVPTX/jump-table.ll
@@ -13,17 +13,12 @@ define void @foo(i32 %i) {
; PTX60: {
; PTX60-NEXT: .reg .pred %p<2>;
; PTX60-NEXT: .reg .b32 %r<2>;
-; PTX60-EMPTY:
+; PTX60-NEXT: $L_brx_0: .branchtargets $L__BB0_2, $L__BB0_3, $L__BB0_4, $L__BB0_5;
; PTX60-NEXT: // %bb.0: // %entry
; PTX60-NEXT: ld.param.b32 %r1, [foo_param_0];
; PTX60-NEXT: setp.gt.u32 %p1, %r1, 3;
; PTX60-NEXT: @%p1 bra $L__BB0_6;
; PTX60-NEXT: // %bb.1: // %entry
-; PTX60-NEXT: $L_brx_0: .branchtargets
-; PTX60-NEXT: $L__BB0_2,
-; PTX60-NEXT: $L__BB0_3,
-; PTX60-NEXT: $L__BB0_4,
-; PTX60-NEXT: $L__BB0_5;
; PTX60-NEXT: brx.idx %r1, $L_brx_0;
; PTX60-NEXT: $L__BB0_2: // %case0
; PTX60-NEXT: st.global.b32 [out], 0;
@@ -107,7 +102,7 @@ define i32 @test2(i32 %tmp158) {
; PTX60: {
; PTX60-NEXT: .reg .pred %p<6>;
; PTX60-NEXT: .reg .b32 %r<3>;
-; PTX60-EMPTY:
+; PTX60-NEXT: $L_brx_0: .branchtargets $L__BB1_6, $L__BB1_7, $L__BB1_8, $L__BB1_9, $L__BB1_10, $L__BB1_11;
; PTX60-NEXT: // %bb.0: // %entry
; PTX60-NEXT: ld.param.b32 %r1, [test2_param_0];
; PTX60-NEXT: setp.gt.s32 %p1, %r1, 119;
@@ -124,13 +119,6 @@ define i32 @test2(i32 %tmp158) {
; PTX60-NEXT: setp.gt.u32 %p2, %r2, 5;
; PTX60-NEXT: @%p2 bra $L__BB1_5;
; PTX60-NEXT: // %bb.12: // %entry
-; PTX60-NEXT: $L_brx_0: .branchtargets
-; PTX60-NEXT: $L__BB1_6,
-; PTX60-NEXT: $L__BB1_7,
-; PTX60-NEXT: $L__BB1_8,
-; PTX60-NEXT: $L__BB1_9,
-; PTX60-NEXT: $L__BB1_10,
-; PTX60-NEXT: $L__BB1_11;
; PTX60-NEXT: brx.idx %r2, $L_brx_0;
; PTX60-NEXT: $L__BB1_7: // %bb339
; PTX60-NEXT: st.param.b32 [func_retval0], 12;
diff --git a/llvm/test/CodeGen/NVPTX/switch-loop-header.mir b/llvm/test/CodeGen/NVPTX/switch-loop-header.mir
index 1a1f39aab8dc4..eecb3b9618871 100644
--- a/llvm/test/CodeGen/NVPTX/switch-loop-header.mir
+++ b/llvm/test/CodeGen/NVPTX/switch-loop-header.mir
@@ -133,14 +133,7 @@ body: |
; CHECK-NEXT: bb.5:
; CHECK-NEXT: successors: %bb.3(0x3e000000), %bb.1(0x04000000), %bb.6(0x00000000), %bb.2(0x3e000000)
; CHECK-NEXT: {{ $}}
- ; CHECK-NEXT: BRX_START 0
- ; CHECK-NEXT: BRX_ITEM %bb.3
- ; CHECK-NEXT: BRX_ITEM %bb.1
- ; CHECK-NEXT: BRX_ITEM %bb.6
- ; CHECK-NEXT: BRX_ITEM %bb.6
- ; CHECK-NEXT: BRX_ITEM %bb.2
- ; CHECK-NEXT: BRX_ITEM %bb.6
- ; CHECK-NEXT: BRX_END %bb.1, undef [[DEF]], 0
+ ; CHECK-NEXT: BRX_IDX undef [[DEF]], %jump-table.0
; CHECK-NEXT: {{ $}}
; CHECK-NEXT: bb.6:
bb.0:
@@ -169,14 +162,7 @@ body: |
bb.5:
successors: %bb.3(0x3e000000), %bb.1(0x04000000), %bb.6(0x00000000), %bb.2(0x3e000000)
- BRX_START 0
- BRX_ITEM %bb.3
- BRX_ITEM %bb.1
- BRX_ITEM %bb.6
- BRX_ITEM %bb.6
- BRX_ITEM %bb.2
- BRX_ITEM %bb.6
- BRX_END %bb.1, undef %10, 0
+ BRX_IDX undef %10, %jump-table.0
bb.6:
...
diff --git a/llvm/test/CodeGen/NVPTX/switch.ll b/llvm/test/CodeGen/NVPTX/switch.ll
index 328f69c59a800..d9bd8b28100c6 100644
--- a/llvm/test/CodeGen/NVPTX/switch.ll
+++ b/llvm/test/CodeGen/NVPTX/switch.ll
@@ -8,7 +8,7 @@ define void @pr170051(i32 %cond) {
; CHECK: {
; CHECK-NEXT: .reg .pred %p<2>;
; CHECK-NEXT: .reg .b32 %r<4>;
-; CHECK-EMPTY:
+; CHECK-NEXT: $L_brx_0: .branchtargets $L__BB0_2, $L__BB0_3, $L__BB0_5, $L__BB0_5, $L__BB0_1, $L__BB0_5, $L__BB0_3;
; CHECK-NEXT: // %bb.0: // %entry
; CHECK-NEXT: mov.b32 %r2, 0;
; CHECK-NEXT: ld.param.b32 %r1, [pr170051_param_0];
@@ -26,14 +26,6 @@ define void @pr170051(i32 %cond) {
; CHECK-NEXT: // %bb.4: // %BS_LABEL_1
; CHECK-NEXT: // in Loop: Header=BB0_3 Depth=1
; CHECK-NEXT: mov.b32 %r3, %r1;
-; CHECK-NEXT: $L_brx_0: .branchtargets
-; CHECK-NEXT: $L__BB0_2,
-; CHECK-NEXT: $L__BB0_3,
-; CHECK-NEXT: $L__BB0_5,
-; CHECK-NEXT: $L__BB0_5,
-; CHECK-NEXT: $L__BB0_1,
-; CHECK-NEXT: $L__BB0_5,
-; CHECK-NEXT: $L__BB0_3;
; CHECK-NEXT: brx.idx %r1, $L_brx_0;
; CHECK-NEXT: $L__BB0_5: // %unreachable
; CHECK-NEXT: // begin inline asm
More information about the llvm-commits
mailing list