Repository navigation
[AIROCMLIR-1214] Add address-based redundant check in rock-narrow-redundant-loads - #467
Conversation
There was a problem hiding this comment.
🟡 Changes recommended
The new “drop mask then reapply via select” rewrite can turn predicated loads into unconditional memory accesses, risking invalid/OOB loads when the original mask fully disables a collapsed group.
Once you've addressed the issues Copilot identified, you can request another Copilot review.
Pull request overview
This PR enhances the rock-narrow-redundant-loads TTIR transform to recognize redundant masked loads whose addresses are invariant even when the loaded result is not (because the mask is folded into Triton’s result constancy). This targets fused-convolution epilogues (e.g., per-channel bias) where avoiding full-tile layout conversions can significantly reduce LDS traffic and improve tuning outcomes.
Changes:
- Add an address-constancy-based narrowing path (
getBroadcastNarrowShape) and plumb aMaskPolicythrough rewriting to support post-broadcast mask restoration viaarith.select. - Prefer the existing result-constancy narrowing path when it applies, and fall back to the new address-based path otherwise.
- Extend MLIR tests to cover the new “narrow + broadcast + select” rewrite pattern and add a per-channel bias reproducer.
File summaries
| File | Description |
|---|---|
| mlir/lib/Dialect/Rock/Transforms/NarrowRedundantLoads.cpp | Adds address-based narrowing and a mask-handling policy to rewrite masked broadcast-like loads. |
| mlir/test/Dialect/Rock/narrow-redundant-loads.mlir | Updates/extends FileCheck coverage for the new narrowing behavior (including per-channel bias). |
Review details
- Files reviewed: 2/2 changed files
- Comments generated: 1
- Review effort level: Lite
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
There was a problem hiding this comment.
Verdict: REQUEST_CHANGES -- submitted as COMMENT (automated reviews are advisory) · Findings: 4 (0 Critical, 1 Major, 3 Minor)
Scope
Adds a second narrowing path to rock-narrow-redundant-loads: when the result constancy is polluted by a varying mask, the pass now falls back to address constancy (getBroadcastNarrowShape), narrows the load with the mask dropped (MaskPolicy::ReapplyAfterBroadcast), broadcasts back, and re-applies the mask with arith.select to restore other. Two files: mlir/lib/Dialect/Rock/Transforms/NarrowRedundantLoads.cpp and mlir/test/Dialect/Rock/narrow-redundant-loads.mlir (two new cases, one flipped from negative to positive).
Findings
NarrowRedundantLoads.cpp:307(Major) -- the new path turns a predicated read into an unconditional one. The guard conditions imply the mask depends only on narrowed dims and the address only on surviving dims, so the original reads all addresses or none; the rewrite always reads all of them. The doc comment's "never touches extra addresses" (and the same claim in the PR description) does not hold when every lane is masked off.NarrowRedundantLoads.cpp:329(Minor) -- the constancy loop is a verbatim copy ofgetNarrowShape:279-284; worth a shared helper.NarrowRedundantLoads.cpp:288(Minor) -- new comment lines exceed the 80-column limit.narrow-redundant-loads.mlir:333(Minor) -- no lit case exercises the mask-without-othersub-branch atNarrowRedundantLoads.cpp:381.
Notes
The design reads well: preferring the exact-slice path first, requiring a splat-like other, and rejecting masks that vary along a surviving dim are all the right guards, and the mask_varies_along_surviving_dim negative test pins the last one. The arith.select is built before replaceOp at an insertion point the mask/other dominate, so ordering is sound.
Back-port: mlir/lib/Dialect/Rock/Transforms/ is nominally shared with ROCm/rocMLIR, but this pass is TTIR-only (tt::LoadOp, tt::ModuleAxisInfoAnalysis) and cannot exist upstream, so I did not raise a Major for the missing back-port note. A one-line "Triton-only pass, not applicable to rocMLIR" in the description would close the checklist item explicitly.
Please run git clang-format --diff origin/develop before merge; a couple of the new lines and the wrapping in the module.walk lambda look like they will move.
CI status
detect is CANCELLED; review and copilot-pull-request-reviewer are the review pipelines themselves and are still in progress. py-checks passed. No real test signal (check-rocmlir) is visible in the check set yet.
|
Note: here's the upstream PR of the Or we can update the PR once this one is approved |
|
|
||
| rewriter.replaceOp(load, broadcast); | ||
| // Restore masked-out lanes. Skip if there was no `other`: they were undefined. | ||
| if (reapplyMask && load.getOther()) |
There was a problem hiding this comment.
I think there's a weird thing going on with other. We always define it to be zero, so in practice it doesn't matter.
The docs say "other (Block, optional) – if mask[idx] is false, return other[idx]. If other is None, the masked-out value is undefined." So, this is correct. But in practice, originally we didn't pass other in rocmlirTriton and triton effectively treats other=nullptr as zero at least on AMD path.
So, I wouldn't worry about this for our codebase, because we always set zero anyway. But for the upstream PR this might change the actual behavior on AMD path. But the docs say it's undefined behavior so, maybe I'm overthinking this.
There was a problem hiding this comment.
I can look into this a bit more next week. I think this doesn't hurt the perf run today so I will investigate this later
There was a problem hiding this comment.
I looked into this and I think the concern is real. I checked the Triton side code: on the AMD path the masked-out lanes aren't actually undefined. LoadStoreOpToLLVM.cpp initialises the fill to a zero vector and only overwrites it when other is present:
Value falseVal = createZeroVector(rewriter, loc, cast<VectorType>(vecTy));
if (otherElems.size() != 0)
falseVal = packElementRangeIntoVector(...);
so we added extra handling: narrowLoad now always re-applies the mask, using other when the load has one. Apply a zero constant when it doesn't
| SmallVector<Value> operands; | ||
| for (Value operand : {load.getPtr(), load.getMask(), load.getOther()}) { | ||
| if (!operand) { | ||
| if (!operand || (reapplyMask && operand != load.getPtr())) { |
There was a problem hiding this comment.
I think we want to still pass other to the tt.load, because the non-redundant load axis could be masked as well, right?
There was a problem hiding this comment.
This is fixed with the masked load change because the non-redundant load axis now keeps its part of the mask on the narrowed tt.load instead of having it dropped.
There was a problem hiding this comment.
isn't this if dropping other?
Replying here:
Yes, and only when reapplyMask is true, and this is intentional. In that mode we don't need other because the select after the broadcast already fills every masked lane with the original other, so putting it on the narrowed load too would just get overwritten.
| return std::nullopt; | ||
| // Re-applying the mask needs a uniform `other`, or none. | ||
| if (load.getOther() && !isSplatLike(load.getOther())) | ||
| return std::nullopt; |
There was a problem hiding this comment.
this part is common between two functions
There was a problem hiding this comment.
Refactored code to avoid duplications
| /// Narrowing shape when the address is constant along some dims even though | ||
| /// the result is not (the mask varies per lane). Drops the mask: it cannot | ||
| /// ride on a load that no longer has those dims. Safe if the mask is constant | ||
| /// on surviving dims, so the narrowed load never touches extra addresses. |
There was a problem hiding this comment.
why do we need the mask to be constant in surviving dims? can we just do a masked load?
There was a problem hiding this comment.
Yes, this is now fixed. The mask is split into its conjuncts: any conjunct that's constant along the narrowed dims slices onto the narrowed load, so it stays a masked load. The constancy requirement only applies to what's left over. A conjunct that varies along a narrowed dim can't fit the smaller mask, so it comes off the load and stops gating the reads. If it's constant on the surviving dims it's all-on or all-off across every address, so dropping it can't reach a new one. If it's constant on neither, dropping it could read out of bound.
pabloantoniom
left a comment
There was a problem hiding this comment.
FYI, this PR makes performance worse in some cases. I was looking at the case reported in AIROCMLIR-1246, and, without retuning, I get 8.7us vs 8.5us in develop. If you have some time it would be good to understand why and if we should maybe skip this in cases where it makes performance worse
I can bring up a follow up PR once it is merged. I think that will be more clear and easier to review |
I will look into this next week. It might be noise since it's only a ~2% drop but I can confirm that next week |
Any update on these perf investigations? Was it just noise? Or are there additional cases that need to be considered for the heuristic? |
I checked the 1246 kernel with the perfConfig reported in the ticket. No regressions found after an averaged run on both old base and the latest develop. The 0.2 us is confirmed to be noise
|
| Value other = load.getOther(); | ||
| if (!other) | ||
| other = arith::ConstantOp::create(rewriter, load.getLoc(), type, | ||
| rewriter.getZeroAttr(type)); |
There was a problem hiding this comment.
rewriter.getZeroAttr(type) returns a null TypedAttr for a shaped type whose element type is not float, integer, or index -- a tt.load of a pointer-of-pointer tensor (tensor<...x!tt.ptr<f32>>) is legal in Triton and would land here, since nothing upstream restricts the element type. Building arith::ConstantOp from a null attribute asserts in a build with assertions and produces an invalid op otherwise. Cheapest fix is to reject the case before any IR is rewritten: in getBroadcastNarrowShape, bail out when the load has no other and the result element type has no zero attribute, so narrowLoad is only reached for types this fallback can actually construct. Checklist: "Major -- Missing null/failure checks on API calls that can fail", raised at Minor because the input shape is unusual in this pipeline.
There was a problem hiding this comment.
checking AI reviews now. Spent some time on Pablo's PR earilier
There was a problem hiding this comment.
This one is not reachable. A pointer-element load asserts earlier in Triton's axis-info analysis at isIntOrFloat(), before narrowLoad ever runs. updated the comment for this
There was a problem hiding this comment.
Verdict: COMMENT · New findings: 2 (0 Critical, 0 Major, 2 Minor)
Scope
Adds a second narrowing path to rock-narrow-redundant-loads that decides redundancy from the constancy of tt.load's address rather than its result, so a per-channel bias read whose mask varies along the collapsed axis can still be narrowed. Two files: mlir/lib/Dialect/Rock/Transforms/NarrowRedundantLoads.cpp and mlir/test/Dialect/Rock/narrow-redundant-loads.mlir.
Findings
Two Minor items, both in NarrowRedundantLoads.cpp:
NarrowRedundantLoads.cpp:333—isConstantInsideNarrowedDimsis a near-verbatim copy ofisConstantOutsideNarrowedDims(line 316); the bodies differ only in thenarrowShape[dim] == 1/!= 1test.NarrowRedundantLoads.cpp:457—getZeroAttron a tensor type whose element type is neither float, integer, nor index yields a null attribute, so theother-less fallback can build an invalidarith.constant.
The four findings from the previous review round are addressed in this revision: the duplicated constancy walk is now factored into getConstantDimsNarrowShape, the over-80-column comment lines are wrapped, the mask-safety argument is documented on getBroadcastNarrowShape, and @mask_varies_along_dim_no_other covers the no-other path.
Notes
- The conjunct-splitting scheme reads correctly: conjuncts that are constant along the collapsed dims stay on the narrowed load, the rest are dropped and re-applied via
arith.selecton the full mask, so every lane the narrowed load skipped is overwritten. The residual behaviour the doc comment calls out — speculating on dereferenceability when a dropped conjunct masks every lane of a row — is the one case where the narrowed load touches an address the original never read. Worth keeping in mind if this pass is ever applied to loads off non-padded allocations. - The PR description still carries the earlier justification ("Dropping the mask from the load is safe because the pass refuses to narrow whenever the mask varies along a dimension it keeps") and states the select is skipped when there was no
other. Both were superseded by the conjunct-splitting and zero-fill changes; refreshing the description before merge would help future readers. getBroadcastNarrowShaperequiresotherto be splat-like, but the select re-applies the full-shapeotherunchanged, so an arbitraryothertensor would work too. Loosening that is a reasonable follow-up, not a blocker.- Lit coverage is solid: positive cases for the bias read and the 2-D bounds-check conjunction, plus two negatives that pin down when narrowing must not fire.
- rocMLIR back-port: the file operates on
tt::LoadOpand includes Triton'sAxisInfoanalysis, so it has no rocMLIR counterpart — option (c) applies and no back-port note is needed.
CI status
No checks in a failed or cancelled state. py-checks and detect passed; copilot-pull-request-reviewer and this pipeline's own review check are still in progress, which is expected.
pabloantoniom
left a comment
There was a problem hiding this comment.
This is an interesting PR. As far as I understand, what this is doing is narrowing the bias read from a 128x128 tile to 128x1, so the ttg.convert_layout is never created.
Related is my PR #483 which attacks the same problem in a different way: it extends support of upstream's OptimizeEpilogue to FMA ops. Note that the problems where you report speedup are f32 (on gfx1101 I assume), which makes sense, because OptimizeEpilogue simply skips FMA kernels. So my PR would improve the epilogue so that, it deletes the ttg.convert_layout by moving the store into the accumulator's layout instead.
Let me check if both PRs are fighting for the same performance, or if they can work well together
umangyadav
left a comment
There was a problem hiding this comment.
Changes look good but check the claude-review/copilot agent reported review comments on this PR.
@erizheng-amd I measured the cases that my PR improves and compared to this branch (same perfConfig on both). Both achieve the same speedup. I also merged both PRs into one branch and measured again. Same speedup. From this experiment I would conclude that both PRs does similar things, just at different levels: yours is at rock level, whereas mine is at Triton level. Not sure what other people think, but we would need to think if it's worth merging both. Performance-wise it's not a win, but also not a performance lose. From my side I'll try to upstream the Triton fix. If it's get rejected, I would vote for merging this PR and closing mine (so we don't need to live with the extra Triton patch) Still something I don't understand is why the mlir/test/rocmlir-driver/oob-buffer-store-fold-split-soffset.mlir does not fail on your PR (it fails on mine). That's probably a sign that both PRs are somewhat different, but the speedup does not seem to accumulate. |
I think they are different, I guess they have a similar effect because the kernels you are looking at have fusions with broadcasted tensors? what if they aren't broadcasted? I understand your PR should still work. |
Not really, I am testing plain convs with no fusions. Actually, I was running some experiments earlier where I was checking the ASM and I saw that my PR drops the number of instructions and VGPRs usage significantly (sometimes 30-40% because it cuts down the LDS round-trip) so now that I think about it, could be a good way to understand how both PRs differ. I'll report the numbers for both PRs later |
Isn't this PR trying to optimize loads and #483 trying to optimize stores ? How are they both similar then ? |
The test you changed |
@pabloantoniom From my understanding, our two PRs have a similar effect at the lower level. They both remove a redundant data movement, and both trade against contiguous I actually found a case where our changes produced a better result. For test |
| {load, std::move(*narrowShape), MaskPolicy::Narrow, {}}); | ||
| return; | ||
| } | ||
| SmallVector<Value> loadMaskConjuncts; |
There was a problem hiding this comment.
It seems like we are holding the conjuncts across rewrites, so could we potentially have a case where a conjunct is a result of a tt.load? In that case, when it gets erased that could cause a crash right?
There was a problem hiding this comment.
This is a real issue and it is now fixed. Now we store conjunct positions instead of values and re-splitting the live mask at rewrite time. When a load gets replaced, MLIR updates the mask that used it, so reading the conjuncts back out of the mask right before we need them always gives us values that are still alive. An index cannot go stale the way a saved value can, and it still points at the right conjunct because a rewrite only swaps one leaf of the mask, never changing how many conjuncts there are or what order they come in.
| Value other = load.getOther(); | ||
| if (!other) | ||
| other = arith::ConstantOp::create(rewriter, load.getLoc(), type, | ||
| rewriter.getZeroAttr(type)); |
Could be fixed by #473 |
| SmallVector<Value> operands; | ||
| for (Value operand : {load.getPtr(), load.getMask(), load.getOther()}) { | ||
| if (!operand) { | ||
| if (!operand || (reapplyMask && operand != load.getPtr())) { |
There was a problem hiding this comment.
isn't this if dropping other?
Replying here:
Yes, and only when reapplyMask is true, and this is intentional. In that mode we don't need other because the select after the broadcast already fills every masked lane with the original other, so putting it on the narrowed load too would just get overwritten.
| } | ||
| if (!isConstantAlongDims(conjunct, shape, *narrowShape, DimSet::Kept, | ||
| axisInfo)) | ||
| return std::nullopt; |
There was a problem hiding this comment.
in what case do we reach here? if they are not collapsed, it must be kept?
There was a problem hiding this comment.
We reach here when a conjunct isn't constant along the dims we collapse. Both checks are about the conjunct, not the dimensions, so failing the first doesn't imply passing the second: a diagonal check like m + n < lim varies in every direction and fails both. Such a conjunct can't stay on the narrowed load, and dropping it isn't safe either, since the load would then read addresses the original never touched, so we give up.
| // rewrites only replace leaves of the `andi` tree, never reshape it. | ||
| SmallVector<Value> conjuncts; | ||
| if (load.getMask()) | ||
| collectMaskConjuncts(load.getMask(), conjuncts); |
There was a problem hiding this comment.
why do we need conjuncts? is it because each of the masks that we AND actually only masks a single dimension?
There was a problem hiding this comment.
I think this deserves a comment in the code
There was a problem hiding this comment.
Comment added. We need conjuncts because a mask is usually a few bounds checks together, one per dimension. After narrowing, they can't all live in the same place: a check that doesn't vary along the dimensions we collapse stays on the smaller load, so it still skips the same elements the original did. The others no longer fit that load, so we apply them to the select after the broadcast.
Motivation
The per-channel bias read in a fused convolution epilogue was not narrowed in
rock-narrow-redundant-loadspass. This pass decided redundancy from the constancy of the loaded result, and Triton's axis analysis folds the mask into that result.The bias read has a repeating address along the column axis but a varying mask, so its result constancy came out as 1 on both
axes and the load looked non-redundant.
The bias therefore stayed a full
128x128tile through layout assignment. That does not match the accumulator's distribution, so Triton insertedttg.convert_layouteither side of the epilogue add, and on AMD those lower to a full-tile LDS round trip. The resulting LDS ceiling kept the tuner out of the small-tile, high-occupancy configurations these convolutions want.This PR adds one more decision path on address constancy for recognizing the redundant loads, after the fix, we see some good performance improvement as shown below:
1x256x202x202_384x256x3x31x128x402x402_256x128x3x3Technical Details
This PR adds a second narrowing path based on address constancy rather than result constancy, plus a
MaskPolicyenum that tellsnarrowLoadhow to reconstruct the original semantics.getBroadcastNarrowShapequeriesAxisInfoonload.getPtr()instead ofload.getResult(). It applies only when the load has a mask,otheris splat-like or absent, and the mask is constant along every dimension that survives narrowing.MaskPolicy::ReapplyAfterBroadcastmakesnarrowLoadbuild an unmasked narrowed load, broadcast it back, and re-apply the original mask witharith.select. A mask that varies along a collapsed dimension has no per-lane bit left to ride on, so it cannot be sliced onto the narrow load. The select is skipped when there was noother, since those lanes were undefined anyway.For the bias case this turns a
128x128masked load into a128x1load, a broadcast, and a select, so layout assignment can keep one layout across add/ReLU/store.Dropping the mask from the load is safe because the pass refuses to narrow whenever the mask varies along a dimension it keeps. Every row therefore has the same mask pattern, so the address the narrowed load reads is one the original load already read.
Test Plan
-- ninja check-rocmlir
-- PR CI
Test Result
-- ninja check-rocmlir passed
-- PR CI passed
-- PR nightly CI passed
Submission Checklist