Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
53 changes: 48 additions & 5 deletions external/llvm-project/llvm/lib/CodeGen/InlineSpiller.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -320,6 +320,21 @@ static void getVDefInterval(const MachineInstr &MI, LiveIntervals &LIS) {
LIS.getInterval(MO.getReg());
}

/// The lanes of LI's register that hold a value at Idx. Live range splitting
/// copies only the lanes that are live, so a sibling may have lanes with no
/// subrange at all.
static LaneBitmask liveLanesAt(const LiveInterval &LI, SlotIndex Idx,
const MachineRegisterInfo &MRI) {
if (!LI.hasSubRanges())
return LI.liveAt(Idx) ? MRI.getMaxLaneMaskForVReg(LI.reg())
: LaneBitmask::getNone();
LaneBitmask LiveLanes = LaneBitmask::getNone();
for (const LiveInterval::SubRange &SR : LI.subranges())
if (SR.liveAt(Idx))
LiveLanes |= SR.LaneMask;
return LiveLanes;
}

/// isSnippet - Identify if a live interval is a snippet that should be spilled.
/// It is assumed that SnipLI is a virtual register with the same original as
/// Edit->getReg().
Expand Down Expand Up @@ -454,6 +469,14 @@ bool InlineSpiller::hoistSpillInsideBB(LiveInterval &SpillLI,
if (DefMBB != CopyMI.getParent() || !SrcQ.isKill())
return false;

// The hoisted store writes all of SrcReg into the slot every sibling shares.
// A lane SrcReg leaves undefined while the original value is live in it
// would overwrite that value, which another sibling may still reload.
LaneBitmask OrigLanes =
liveLanesAt(LIS.getInterval(Original), SrcVNI->def, MRI);
if ((OrigLanes & ~liveLanesAt(SrcLI, SrcVNI->def, MRI)).any())
return false;

MachineBasicBlock *MBB = DefMBB;
MachineBasicBlock::iterator MII;
if (SrcVNI->isPHIDef())
Expand Down Expand Up @@ -1368,20 +1391,40 @@ void InlineSpiller::spillAroundUses(Register Reg) {
// FIXME: Infer regclass from instruction alone.
Register NewVReg = Edit->createFrom(Reg);

if (RI.Reads)
// The spill below stores every lane of NewVReg. When MI defines only some
// of them without reading the rest, lanes the original value holds across
// MI would be overwritten in the slot, where a sibling may still reload
// them. Fill them from the slot first so the store writes that value back.
LaneBitmask DefLanes = LaneBitmask::getNone();
bool hasLiveDef = false;
for (const auto &[OpMI, OpIdx] : Ops) {
const MachineOperand &MO = OpMI->getOperand(OpIdx);
if (!MO.isDef())
continue;
DefLanes |= MO.getSubReg() ? TRI.getSubRegIndexLaneMask(MO.getSubReg())
: MRI.getMaxLaneMaskForVReg(Reg);
hasLiveDef |= !MO.isDead();
}
bool PreserveSlotLanes = false;
if (RI.Writes && !RI.Reads && hasLiveDef &&
StackInt->liveAt(Idx.getBaseIndex())) {
LaneBitmask OrigLanes =
liveLanesAt(LIS.getInterval(Original), Idx.getBaseIndex(), MRI);
PreserveSlotLanes = (OrigLanes & ~DefLanes).any();
}

if (RI.Reads || PreserveSlotLanes)
insertReload(NewVReg, Idx, &MI);

// Rewrite instruction operands.
bool hasLiveDef = false;
for (const auto &OpPair : Ops) {
MachineOperand &MO = OpPair.first->getOperand(OpPair.second);
MO.setReg(NewVReg);
if (MO.isUse()) {
if (!OpPair.first->isRegTiedToDefOperand(OpPair.second))
MO.setIsKill();
} else {
if (!MO.isDead())
hasLiveDef = true;
} else if (PreserveSlotLanes) {
MO.setIsUndef(false);
}
}
LLVM_DEBUG(dbgs() << "\trewrite: " << Idx << '\t' << MI << '\n');
Expand Down
77 changes: 77 additions & 0 deletions llvm-patches/llvm-patch-content.txt
Original file line number Diff line number Diff line change
Expand Up @@ -858,3 +858,80 @@ Applies on top of patch201186.patch and does not apply to a tree without it.

Drop this patch on the next LLVM/Triton bump that includes upstream commit
077ec06128e27e3dce2b97563dab703f89024ddf.

----------------------------------------------------------------------
patch-inline-spiller-partial-sibling-spill.patch
----------------------------------------------------------------------
- Upstream status : Downstream fix. Reported upstream as
https://github.com/llvm/llvm-project/issues/225054
- Upstream title : [InlineSpiller] Do not store undefined lanes into a
shared spill slot
- File Fixed : llvm/lib/CodeGen/InlineSpiller.cpp
- Contents : Adds liveLanesAt(), which returns the lanes of a register
that hold a value at a slot index, counting a lane with no
subrange as dead. hoistSpillInsideBB now declines to hoist
a spill onto a sibling that leaves undefined a lane the
original value is live in at the sibling's def. When
spillAroundUses spills a def that writes only some lanes
without reading the rest, while the stack slot is live and
the original value is live in other lanes across the def,
it reloads the slot first and drops the undef flag, so the
full-width store writes those lanes back unchanged. Lanes
the original value never holds are free to carry garbage,
so upstream's hoisting of partially defined values
(splitkit.mir, splitkit-copy-bundle.mir) is unchanged.

Symptom without the patch
-------------------------
A gfx950 attention kernel from the weekly parameterSweeps run (f16, g=5,
seq_len_q=1, seq_len_k=361, 4 Q heads over 2 KV heads, attention bias,
sliding_window_look_back=262, last_valid_kv_index=141,134,279,29,84; perf
config mPerBlockG0=64, nPerBlockG0=128, kPerBlock=64, numWaves=1,
numStages=3, wavesPerEU=2) dies with a GPU memory access fault on the page
16 KiB past the bias buffer. The kernel needs all 512 VGPRs, so SGPRs are
spilled into VGPR lanes. The K buffer descriptor is spilled before the KV
loop and reloaded at the top of every iteration, but inside the loop the
same four lanes are overwritten with the bias buffer's base address. From
the second iteration on, the K loads address the bias buffer with K's
offsets.

Later rocMLIR changes altered that kernel's IR and it no longer reaches the
hazard. A bf16 kernel from the same sweep (g=6, seq_len_q=1, seq_len_k=151,
16 Q heads over 4 KV heads, attention scale and bias, transV,
sliding_window_look_back=150, last_valid_kv_index=120,143,78,69,87,150; perf
config mPerBlockG0=16, nPerBlockG0=128, nPerBlockG1=256, kPerBlock=512,
numWaves=16, matrixInstrNonkdim=16, numStages=1) faults the same way: a
descriptor rebuilt inside the KV loop from another's sub2_sub3 is stored
full-width over the spill slot of a descriptor reloaded at the top of every
iteration.

Root cause
----------
The bias descriptor borrows num_records and flags (sub2_sub3) from the K
descriptor. Splitting K's live range inside the loop therefore produced a
sibling defined only by `undef %a.sub2_sub3 = COPY %b.sub2_sub3`, since the
base (sub0_sub1) is no longer needed in a register there; its value lives
on in the stack slot that all siblings share, for the next iteration's
reload. When the following piece was spilled, spillAroundUses treated its
defining subregister copy as a sibling copy and hoistSpillInsideBB stored
the whole partial sibling into that slot, writing whatever the physical
registers held in sub0_sub1 over K's base. Without the hoist, insertSpill
would store the undefined lanes the same way. Upstream #177703 fixed the
same hazard for the cross-block hoist in isSpillCandBB; upstream main still
has the in-block path unchanged. The #177703 check only inspects subranges
that exist, so it would accept a partial copy like the one above, which has
no subrange for the lanes it does not define. It did not fire here and is
left as upstream wrote it, since it is also what hoists the partially
defined value in splitkit.mir out of its loop.

CI coverage
-----------
mlir/test/e2e/PrAttentionBF16Gfx950.toml pins the bf16 shape, perf config
and chip counts above and runs them on gfx950 with -pv --pv-f64 on the fixed
input pattern, for both transO values. The kernel LLVM IR it hands the backend
is identical to the faulting sweep's, so without this patch the run dies
with the memory access fault and FileCheck fails on the missing verification
line.

#225054 is still open upstream. Drop or reconcile this patch once a fix for
it lands in the pinned LLVM.
87 changes: 87 additions & 0 deletions llvm-patches/patch-inline-spiller-partial-sibling-spill.patch
Original file line number Diff line number Diff line change
@@ -0,0 +1,87 @@
diff --git a/llvm/lib/CodeGen/InlineSpiller.cpp b/llvm/lib/CodeGen/InlineSpiller.cpp
index 5ed3b3628..b5a2b6fbd 100644
--- a/llvm/lib/CodeGen/InlineSpiller.cpp
+++ b/llvm/lib/CodeGen/InlineSpiller.cpp
@@ -320,6 +320,21 @@ static void getVDefInterval(const MachineInstr &MI, LiveIntervals &LIS) {
LIS.getInterval(MO.getReg());
}

+/// The lanes of LI's register that hold a value at Idx. Live range splitting
+/// copies only the lanes that are live, so a sibling may have lanes with no
+/// subrange at all.
+static LaneBitmask liveLanesAt(const LiveInterval &LI, SlotIndex Idx,
+ const MachineRegisterInfo &MRI) {
+ if (!LI.hasSubRanges())
+ return LI.liveAt(Idx) ? MRI.getMaxLaneMaskForVReg(LI.reg())
+ : LaneBitmask::getNone();
+ LaneBitmask LiveLanes = LaneBitmask::getNone();
+ for (const LiveInterval::SubRange &SR : LI.subranges())
+ if (SR.liveAt(Idx))
+ LiveLanes |= SR.LaneMask;
+ return LiveLanes;
+}
+
/// isSnippet - Identify if a live interval is a snippet that should be spilled.
/// It is assumed that SnipLI is a virtual register with the same original as
/// Edit->getReg().
@@ -454,6 +469,14 @@ bool InlineSpiller::hoistSpillInsideBB(LiveInterval &SpillLI,
if (DefMBB != CopyMI.getParent() || !SrcQ.isKill())
return false;

+ // The hoisted store writes all of SrcReg into the slot every sibling shares.
+ // A lane SrcReg leaves undefined while the original value is live in it
+ // would overwrite that value, which another sibling may still reload.
+ LaneBitmask OrigLanes =
+ liveLanesAt(LIS.getInterval(Original), SrcVNI->def, MRI);
+ if ((OrigLanes & ~liveLanesAt(SrcLI, SrcVNI->def, MRI)).any())
+ return false;
+
MachineBasicBlock *MBB = DefMBB;
MachineBasicBlock::iterator MII;
if (SrcVNI->isPHIDef())
@@ -1368,20 +1391,40 @@ void InlineSpiller::spillAroundUses(Register Reg) {
// FIXME: Infer regclass from instruction alone.
Register NewVReg = Edit->createFrom(Reg);

- if (RI.Reads)
+ // The spill below stores every lane of NewVReg. When MI defines only some
+ // of them without reading the rest, lanes the original value holds across
+ // MI would be overwritten in the slot, where a sibling may still reload
+ // them. Fill them from the slot first so the store writes that value back.
+ LaneBitmask DefLanes = LaneBitmask::getNone();
+ bool hasLiveDef = false;
+ for (const auto &[OpMI, OpIdx] : Ops) {
+ const MachineOperand &MO = OpMI->getOperand(OpIdx);
+ if (!MO.isDef())
+ continue;
+ DefLanes |= MO.getSubReg() ? TRI.getSubRegIndexLaneMask(MO.getSubReg())
+ : MRI.getMaxLaneMaskForVReg(Reg);
+ hasLiveDef |= !MO.isDead();
+ }
+ bool PreserveSlotLanes = false;
+ if (RI.Writes && !RI.Reads && hasLiveDef &&
+ StackInt->liveAt(Idx.getBaseIndex())) {
+ LaneBitmask OrigLanes =
+ liveLanesAt(LIS.getInterval(Original), Idx.getBaseIndex(), MRI);
+ PreserveSlotLanes = (OrigLanes & ~DefLanes).any();
+ }
+
+ if (RI.Reads || PreserveSlotLanes)
insertReload(NewVReg, Idx, &MI);

// Rewrite instruction operands.
- bool hasLiveDef = false;
for (const auto &OpPair : Ops) {
MachineOperand &MO = OpPair.first->getOperand(OpPair.second);
MO.setReg(NewVReg);
if (MO.isUse()) {
if (!OpPair.first->isRegTiedToDefOperand(OpPair.second))
MO.setIsKill();
- } else {
- if (!MO.isDead())
- hasLiveDef = true;
+ } else if (PreserveSlotLanes) {
+ MO.setIsUndef(false);
}
}
LLVM_DEBUG(dbgs() << "\trewrite: " << Idx << '\t' << MI << '\n');
1 change: 1 addition & 0 deletions mlir/test/e2e/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,7 @@ if (ROCMLIR_DRIVER_PR_E2E_TEST_ENABLED)
PrAttentionF32
PrAttentionF16
PrAttentionBF16
PrAttentionBF16Gfx950
PrAttentionI8
PrGemmSplitK
PrGemmElementwiseGemmF32
Expand Down
4 changes: 4 additions & 0 deletions mlir/test/e2e/PrAttentionBF16Gfx950.cfg
Original file line number Diff line number Diff line change
@@ -0,0 +1,4 @@
# The perf config requires 128 KiB of LDS, and the register allocation pinned
# down by these tests is the one produced for gfx950
if "gfx950" not in config.arch:
config.unsupported = True
25 changes: 25 additions & 0 deletions mlir/test/e2e/PrAttentionBF16Gfx950.toml
Original file line number Diff line number Diff line change
@@ -0,0 +1,25 @@
directory = "PrAttentionBF16Gfx950"
prefix = "rocmlir-gen"
# The fault does not depend on the input values. With random data this shape
# misses the bf16 tolerance against the CPU reference for any perf config, so
# the inputs are pinned to the fixed pattern instead of %random_data.
suffix = "--operation attention -t bf16 --arch %arch -pv --pv-f64 -rand fixed %rocmlir_gen_flags | rocmlir-driver --host-pipeline=highlevel | rocmlir-driver -c | mlir-runner -O2 --shared-libs=%linalg_test_lib_dir/libmlir_rocm_runtime%shlibext,%conv_validation_wrapper_library_dir/libconv-validation-wrappers%shlibext,%linalg_test_lib_dir/libmlir_runner_utils%shlibext,%linalg_test_lib_dir/libmlir_float16_utils%shlibext --entry-point-result=void | FileCheck %s --check-prefix="

[[axis]]
name = "transO"
values = ["true", "false"]
prefix = "--transO="

[[suite]]
name = "pr_attention_bf16_gfx950"

# SGPRs spill into VGPR lanes, and inside the KV loop a buffer descriptor is
# rebuilt from num_records/flags (sub2_sub3) copied out of another one. That
# made InlineSpiller store the rebuilt descriptor full-width into the spill
# slot of a descriptor reloaded at the top of every iteration, replacing its
# base address with garbage, so the second iteration's loads fault
# (llvm-patches/patch-inline-spiller-partial-sibling-spill.patch). Changing
# the shape, perf config or chip counts changes register allocation and can
# hide the bug.
[[suite.test]]
config = "--num_cu 256 --num_chiplets 8 -g 6 -seq_len_q 1 -seq_len_k 151 -num_heads_q 16 -num_heads_kv 4 -head_dim_qk 196 -head_dim_v 183 --with-attn-scale --with-attn-bias --transV=true -sliding_window_look_back=150 -last_valid_kv_index=120,143,78,69,87,150 -perf_config attn:mPerBlockG0=16,nPerBlockG0=128,nPerBlockG1=256,kPerBlock=512,kpack=1,numCTAs=1,numWaves=16,matrixInstrNonkdim=16,splitKFactor=1,numStages=1,wavesPerEU=0,gridGroupSize=0"