Eliminate concat_past_present - #5184
Conversation
There was a problem hiding this comment.
Copilot review overview
🟡 Changes recommended
The fused decode path corrupts cache data for invalid sequence lengths and fails on symbolic dynamic shapes.
Review effort: Balanced
Findings: 1
Open (4)
What changed in this PR
Adds GPU fusion to write concat_past_present producers directly into the KV cache.
Changes:
- Adds decode/prefill fusion and scalar-driven cache slicing.
- Vectorizes the existing copy kernel and updates HIP graph capture boundaries.
- Adds pass, verification, and replay tests.
Review performed as a single pass without agent fan-out; GPU builds and tests were not executed.
| File | Description |
|---|---|
test/verify/test_concat_past_present_mul.cpp |
Adds decode verification. |
test/verify/test_concat_past_present_mul_prefill.cpp |
Adds prefill verification. |
test/gpu/make_precompile_op.hpp |
Adds operation-based test helper. |
test/gpu/hipgraphify.cpp |
Tests runtime-view partitioning. |
test/gpu/hip_graph.cpp |
Tests replay with changing positions. |
test/gpu/fuse_concat_past_present.cpp |
Tests fusion structure. |
src/targets/gpu/target.cpp |
Registers the fusion pass. |
src/targets/gpu/kernels/include/migraphx/kernels/concat_past_present.hpp |
Vectorizes cache copies. |
src/targets/gpu/jit/concat_past_present.cpp |
Selects vector width and launch size. |
src/targets/gpu/include/migraphx/gpu/precompile_op.hpp |
Shares precompile matching logic. |
src/targets/gpu/include/migraphx/gpu/hip.hpp |
Adds scalar loading and spin sync. |
src/targets/gpu/include/migraphx/gpu/fuse_concat_past_present.hpp |
Declares the pass. |
src/targets/gpu/hipgraphify.cpp |
Excludes runtime views from capture. |
src/targets/gpu/hip.cpp |
Registers scalar loading and implements spin sync. |
src/targets/gpu/fuse_ops.cpp |
Reuses the shared matcher. |
src/targets/gpu/fuse_concat_past_present.cpp |
Implements direct-cache fusion. |
src/targets/gpu/device_name.cpp |
Adjusts feature guards. |
src/targets/gpu/CMakeLists.txt |
Builds the new pass. |
💡 Configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
| const auto& s = args[0].get_shape(); | ||
| assert(s.lens()[axis] > 0); | ||
| std::vector<std::size_t> start(s.ndim(), 0); | ||
| start[axis] = std::clamp<std::int64_t>(args[1].at<std::int64_t>(), 0, s.lens()[axis] - 1); |
There was a problem hiding this comment.
So this was added in #4764 as a purely defensive guard during LLM benchmarking. There isnt a real world case for this as it doesnt make sense. Furthermore, ORT's GQA has no such guard. Its CPU kernel computes total_seqlen = SafeInt(seqlens_k) + 1 and writes at past_seqlen unconditionally, so seqlens_k == cache_length is an out-of-bounds write there and a negative value throws from SafeInt.
There was a problem hiding this comment.
Also, from the comment here this was fixed in ModelBench. I am going to change this to a throw for out of bounds.
| if(ps != cs or not ps.standard() or not ks.standard()) | ||
| return false; |
| constexpr index_int head_size = params.head_size; | ||
| constexpr index_int vec_head_size = head_size / N; |
| static migraphx::instruction_ref | ||
| add_gathers(migraphx::module& m, migraphx::instruction_ref x, std::size_t axis, std::size_t n) |
bdevorem
left a comment
There was a problem hiding this comment.
needs a changelog, and I had a couple qs. otherwise lgtm
| } | ||
|
|
||
| #if MIGRAPHX_USE_HIPBLASLT | ||
|
|
There was a problem hiding this comment.
was this just a bug before, the rocblas functions being gated by this? or was there a real reason?
There was a problem hiding this comment.
Oops, I thought I removed that change.
| // Every chunk offset is a multiple of head_size, so copy in the widest | ||
| // vector dividing it, capped at 4 since memory coloring only aligns | ||
| // buffers to 4 elements | ||
| const std::size_t vec_size = params.head_size % 4 == 0 ? 4 |
There was a problem hiding this comment.
is this allowing fp8 vectorization? we have that disabled elsewhere, right?
There was a problem hiding this comment.
Yea thats a good point. Let me update that.



Motivation
During LLM decode, every layer appends the new token's K and V into the kv-cache with
concat_past_present, a copy kernel that moves a tiny{1, H_kv, 1, D}chunk from the producer's buffer into the cache. Two launches per layer for a few hundred bytes each is pure launch overhead on the decode critical path. This change removes the copy kernel by having the producer kernel write directly into the cache slot, and it makes the remaining copy kernel faster for the prefill and unfused cases.Technical Details
gpu::fuse_concat_past_presentpass (src/targets/gpu/fuse_concat_past_present.cpp), run afterfuse_ops. A matcher finds a precompiledconcat_past_presentwhose present input comes, through single-use reorder views, from a single-use precompiledpointwiseorfused_concatwriting into its ownallocate. When the producer's output has the same standard layout as the cache slot, its output buffer is replaced by a view of the cache and the concat becomesidentity(cache, producer), which keeps the write ordered before readers of the cache and is dropped byeliminate_identityat the end of the pipeline.seq > 1) appends at position zero, so the view is a staticslice.seq == 1, batch 1) appends at a device-computed position. The pass insertship::load_scalar, which copiesseqlens_kinto a pinned host buffer allocated once infinalizeand spins onhipStreamQueryto avoid scheduler wake latency, thengpu::slice_at(cache, idx), a size-1 aliasing view along the sequence axis. Oneload_scalaris shared by every append reading the sameseqlens_k. An index outside[0, cache_length)throws: the copy kernel skipped such writes, but a view cannot, and the position is invalid input either way (ORT has no guard there at all).eliminate_concat_past_presentbackend option (default true).hipgraphify.
hip::load_scalaris uncapturable, andis_capturablenow rejects "runtime views" (aliasing ops with a host-value input, i.e.slice_at,scan_slice,dyn_sliceafter lowering) plus any kernel whose input alias path crosses one. A captured kernel would otherwise replay the capture-time slot address. Kernels reading the whole cache afterwards are still captured.Copy kernel (
kernels/concat_past_present.hpp,jit/concat_past_present.cpp). Copiesvec<T, N>along the head axis, withNchosen once by the JIT (widest of 4/2/1 dividinghead_size, capped by memory coloring's 4-element alignment) and passed as a template argument. Large copies take four vectors per thread viaglobal_stridesince one load per wave cannot hide memory latency.Shared code.
precompile_namemoved fromfuse_ops.cppintomigraphx/gpu/precompile_op.hppand reads the inner op's name off the value instead of reconstructing the op.test/gpu/make_precompile_op.hppgained an(op, output_shape)overload.Tests. Pass unit tests (decode, prefill, multi-use skip, option on/off, out-of-range throw); verify tests for decode and prefill through a pointwise producer, plus one with hip_graph enabled; a hipgraphify partition test for the runtime-view boundary; and a hip_graph replay test appending at three positions that fails if the view or writer is captured. The verify tests that asserted skip semantics for negative or too-large
seqlens_kwere removed.Changelog Category
Add a
CHANGELOG.mdentry for any option other thanNot ApplicableFollow the LLVM AI Tool Use Policy for contributions using AI.