Skip to content

Eliminate concat_past_present - #5184

Merged
pfultz2 merged 19 commits into
developfrom
fuse-concat-past-present
Oct 1, 2026
Merged

pfultz2 merged 19 commits into
developfrom
fuse-concat-past-present

Conversation

@pfultz2

@pfultz2 pfultz2 commented Aug 24, 2026 •

Copy link
Copy Markdown
Collaborator

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_present pass (src/targets/gpu/fuse_concat_past_present.cpp), run after fuse_ops. A matcher finds a precompiled concat_past_present whose present input comes, through single-use reorder views, from a single-use precompiled pointwise or fused_concat writing into its own allocate. 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 becomes identity(cache, producer), which keeps the write ordered before readers of the cache and is dropped by eliminate_identity at the end of the pipeline.

  • Prefill (seq > 1) appends at position zero, so the view is a static slice.
  • Decode (seq == 1, batch 1) appends at a device-computed position. The pass inserts hip::load_scalar, which copies seqlens_k into a pinned host buffer allocated once in finalize and spins on hipStreamQuery to avoid scheduler wake latency, then gpu::slice_at(cache, idx), a size-1 aliasing view along the sequence axis. One load_scalar is shared by every append reading the same seqlens_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).
  • Producer and view chain are moved down to the concat so the view's inputs are defined first; no position map is needed.
  • Gated by the new eliminate_concat_past_present backend option (default true).

hipgraphify. hip::load_scalar is uncapturable, and is_capturable now rejects "runtime views" (aliasing ops with a host-value input, i.e. slice_at, scan_slice, dyn_slice after 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). Copies vec<T, N> along the head axis, with N chosen once by the JIT (widest of 4/2/1 dividing head_size, capped by memory coloring's 4-element alignment) and passed as a template argument. Large copies take four vectors per thread via global_stride since one load per wave cannot hide memory latency.

Shared code. precompile_name moved from fuse_ops.cpp into migraphx/gpu/precompile_op.hpp and reads the inner op's name off the value instead of reconstructing the op. test/gpu/make_precompile_op.hpp gained 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_k were removed.

Changelog Category

Add a CHANGELOG.md entry for any option other than Not Applicable

    • Added: New functionality.
    • Changed: Changes to existing functionality.
    • Removed: Functionality or support that has been removed. (Compared to a previous release)
    • Optimized: Component performance that has been optimized or improved.
    • Resolved Issues: Known issues from a previous version that have been resolved.
    • Not Applicable: This PR is not to be included in the changelog.

Follow the LLVM AI Tool Use Policy for contributions using AI.

@bdevorem
bdevorem self-requested a review August 28, 2026 19:20
@pfultz2 pfultz2 added the llm label Sep 20, 2026
@pfultz2
pfultz2 marked this pull request as ready for review September 29, 2026 21:31
@pfultz2
pfultz2 requested a review from causten as a code owner September 29, 2026 21:31
Copilot AI balanced review requested due to automatic review settings September 29, 2026 21:31

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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 High severity · 1 Medium severity · 2 Low severity

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);

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Also, from the comment here this was fixed in ModelBench. I am going to change this to a throw for out of bounds.

Comment on lines +108 to +109
if(ps != cs or not ps.standard() or not ks.standard())
return false;
Comment on lines +106 to +107
constexpr index_int head_size = params.head_size;
constexpr index_int vec_head_size = head_size / N;
Comment thread test/gpu/hip_graph.cpp Outdated
Comment on lines +379 to +380
static migraphx::instruction_ref
add_gathers(migraphx::module& m, migraphx::instruction_ref x, std::size_t axis, std::size_t n)

@bdevorem bdevorem left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

needs a changelog, and I had a couple qs. otherwise lgtm

}

#if MIGRAPHX_USE_HIPBLASLT

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

was this just a bug before, the rocblas functions being gated by this? or was there a real reason?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Oops, I thought I removed that change.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed.

// 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

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

is this allowing fp8 vectorization? we have that disabled elsewhere, right?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yea thats a good point. Let me update that.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed.

@pfultz2
pfultz2 requested a review from bdevorem September 30, 2026 18:18
@pfultz2
pfultz2 merged commit 010ee28 into develop Oct 1, 2026
33 checks passed
@pfultz2
pfultz2 deleted the fuse-concat-past-present branch October 1, 2026 13:44
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants