Skip to content
Merged
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
42 changes: 42 additions & 0 deletions TECHNICAL_REPORTS/pr-2097-pool-backed-kv-trim.en.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
# PR #2097: Pool-backed `KVCache::trim` rewinds the block table

**Date**: 2026-10-06
**Status**: Reviewed and verified on CUDA (GB10); M5 tile-aligned path verified by the contributor only
**Risk**: Low

## Summary

`KVCache::trim` returned `0` for a pool-backed cache (`paged_backing.is_some()`), so every caller that dropped prefill pad rows with `trim(excess)` left those rows in the pool and left `offset` at the padded length. Decode then attended over the pad rows and took its RoPE position from the padded `offset`. The PR makes `trim` call `PagedBlockPool::rewind_tokens` through the cache's own backing and subtract the removed count from `offset`, and routes the decode lookahead teardown (`apply_lookahead_trim`) through the same call instead of its pool-only `rewind_paged_tokens` branch, which rewound the pool but left `offset` one ahead.

Closes #2096. Contributor: rapsealk.

## Review

- Every production `trim` caller on a pool-backed cache (the four pad trims in `scheduler/prefill.rs` and the lookahead teardown) expects a real trim. No caller pairs `trim` with a separate pool rewind, so there is no double rewind. `trim_to`/`DetachedCacheSet::truncate_to` have no production caller.
- `sync_paged_state_with_dense` returns early for pool-backed sequences, so the `sync_sequence_storage` call after the teardown does not re-trim the block table.
- The `live_len()` clamp still runs before the pool branch, and `offset` moves by the count the pool actually removed, so pool `len` and `offset` cannot diverge through this path.
- The pool branch uses `.expect` where the old lookahead branch logged a warning. This matches the neighbouring `write_paged` and `update_and_fetch_paged` calls on the same pool, and a rewind failure there means the block table is already inconsistent. Recorded as MEDIUM, not changed.
- `CachePool::rewind_paged_tokens` has no non-test caller left; it was kept.

No CRITICAL or HIGH findings; no code changes were made on top of the contributor's commits.

## Verification on GB10 (CUDA, release profile, `--features cuda`)

Branch head merged with `origin/main` at `2bbf192f`.

- New test `trim_drops_padded_prefill_rows_from_a_pool_backed_cache`: passes with the fix. With `cache.rs` and `decode_tick.rs` reverted to `origin/main` it fails at the first assertion (`left: 0, right: 27`).
- Neighbouring mlxcel-core suites, `--test-threads=1` under `gpu-lock`: `cache::paged_batch_decode` 23 passed, `cache::paged` 137 passed, `cache::tests` 83 passed, `cache::detach` 51 passed, `speculative` 175 passed. Main crate: `lookahead` 10 passed (including the teardown position tests), `block_reclaim` 23 passed, `server::batch::scheduler` 167 passed (7 ignored).
- `cargo fmt --all -- --check`, `cargo clippy -p mlxcel-core --lib --tests --features cuda -- -D warnings`, and the same for `-p mlxcel`: clean.
- Server: `mlxcel-server -m llama-3.2-1b-4bit --no-prompt-cache`, five distinct chat prompts (40, 42, 51, 52, 64 prompt tokens) sent concurrently, three rounds, `temperature 0`, `max_tokens 32`. Default backend (auto, resolves to pool-backed paged on this host) compared with `--decode-storage-backend dense`. Debug logs confirm the batched padded prefill ran (`batched prefill: 3 requests, padded to 52`, etc.).

| Build | Default vs dense, identical continuations | First batched decode gather |
|---|---|---|
| `origin/main` (fix reverted) | 5 of 15 | 224 visible KV tokens across 4 requests (pad rows retained) |
| PR #2097 | 15 of 15 | 189 visible KV tokens across 4 requests (real rows plus one each) |

Before the fix the default backend answered `Repeat: one two three` with `One!Two!Three!` and the France prompt with a rambling `The capital of France is... Paris! The City of Light, ...`, matching the contributor's M5 report. The dense reference output was byte-identical between the two builds. This confirms the issue's claim that the batched-prefill longest-row padding corrupts pool-backed output on hardware without a Neural Accelerator.

## Not verified on this host

- The M5 tile-aligned prefill path (`should_align_prefill()` true) needs Apple M5 hardware. The contributor's M5 Pro results (Llama-3.2-1B, SmolLM2-135M, Qwen3-0.6B, before 0-1 of 5, after 5 of 5) are taken as reported.
- The #1760 whole-prompt replay abort on M5 is out of scope for this PR.
42 changes: 42 additions & 0 deletions TECHNICAL_REPORTS/pr-2097-pool-backed-kv-trim.ko.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
# PR #2097: 풀 기반 `KVCache::trim`의 블록 테이블 되감기

**작성일**: 2026-10-06
**상태**: 리뷰 완료, CUDA(GB10)에서 검증; M5 타일 정렬 경로는 기여자 검증만 있음
**위험도**: 낮음

## 요약

`KVCache::trim`은 풀 기반 캐시(`paged_backing.is_some()`)에 대해 `0`을 반환했습니다. 그래서 `trim(excess)`로 prefill 패딩 행을 제거하던 모든 호출부가 패딩 행을 풀에 남겼고 `offset`도 패딩된 길이에 머물렀습니다. 그 결과 decode가 패딩 행까지 attention 대상으로 삼았고, RoPE 위치도 패딩된 `offset`에서 가져왔습니다. 이 PR은 `trim`이 캐시 자신의 backing을 통해 `PagedBlockPool::rewind_tokens`를 호출하고 실제로 제거된 개수만큼 `offset`을 줄이도록 바꿉니다. 또한 decode lookahead 해제(`apply_lookahead_trim`)가 풀 전용 `rewind_paged_tokens` 분기 대신 같은 호출을 쓰도록 합니다. 기존 분기는 풀만 되감고 `offset`은 한 칸 앞에 남겨 두었습니다.

Closes #2096. 기여자: rapsealk.

## 리뷰

- 풀 기반 캐시에 대한 모든 실제 `trim` 호출부(`scheduler/prefill.rs`의 패딩 trim 네 곳과 lookahead 해제)는 실제 trim을 기대합니다. `trim`과 별도의 풀 되감기를 함께 호출하는 곳이 없으므로 이중 되감기는 생기지 않습니다. `trim_to`/`DetachedCacheSet::truncate_to`에는 실제 호출부가 없습니다.
- `sync_paged_state_with_dense`는 풀 기반 시퀀스에서 바로 반환하므로, 해제 뒤의 `sync_sequence_storage` 호출이 블록 테이블을 다시 자르지 않습니다.
- `live_len()` 제한은 풀 분기보다 먼저 적용되고, `offset`은 풀이 실제로 제거한 개수만큼 움직이므로 이 경로로 풀 `len`과 `offset`이 어긋날 수 없습니다.
- 풀 분기는 `.expect`를 쓰지만 기존 lookahead 분기는 경고 로그만 남겼습니다. 같은 풀을 쓰는 인접한 `write_paged`, `update_and_fetch_paged` 호출과 같은 방식이며, 여기서 되감기가 실패한다면 블록 테이블이 이미 불일치 상태입니다. MEDIUM으로 기록만 하고 수정하지 않았습니다.
- `CachePool::rewind_paged_tokens`는 이제 테스트 외 호출부가 없지만 그대로 두었습니다.

CRITICAL/HIGH 문제는 없었고, 기여자 커밋 위에 코드 변경을 추가하지 않았습니다.

## GB10 검증 (CUDA, release 프로필, `--features cuda`)

브랜치 헤드를 `origin/main`의 `2bbf192f`와 병합한 상태에서 검증했습니다.

- 새 테스트 `trim_drops_padded_prefill_rows_from_a_pool_backed_cache`: 수정 적용 시 통과합니다. `cache.rs`와 `decode_tick.rs`를 `origin/main`으로 되돌리면 첫 단언에서 실패합니다(`left: 0, right: 27`).
- 인접 mlxcel-core 테스트, `gpu-lock` 아래 `--test-threads=1`: `cache::paged_batch_decode` 23개, `cache::paged` 137개, `cache::tests` 83개, `cache::detach` 51개, `speculative` 175개 모두 통과. 메인 크레이트: `lookahead` 10개(해제 위치 테스트 포함), `block_reclaim` 23개, `server::batch::scheduler` 167개(7개 ignored) 통과.
- `cargo fmt --all -- --check`, `cargo clippy -p mlxcel-core --lib --tests --features cuda -- -D warnings`, `-p mlxcel` 동일 명령 모두 경고 없음.
- 서버: `mlxcel-server -m llama-3.2-1b-4bit --no-prompt-cache`로 서로 다른 채팅 프롬프트 5개(프롬프트 토큰 40, 42, 51, 52, 64)를 동시에 보내는 라운드를 3회 반복했습니다(`temperature 0`, `max_tokens 32`). 기본 백엔드(auto, 이 호스트에서는 풀 기반 paged로 결정)와 `--decode-storage-backend dense`를 비교했습니다. 디버그 로그로 패딩된 batched prefill이 실제로 실행됐음을 확인했습니다(`batched prefill: 3 requests, padded to 52` 등).

| 빌드 | 기본 vs dense, 동일한 출력 | 첫 batched decode gather |
|---|---|---|
| `origin/main` (수정 되돌림) | 15개 중 5개 | 요청 4개에 걸쳐 KV 토큰 224개 (패딩 행 잔존) |
| PR #2097 | 15개 중 15개 | 요청 4개에 걸쳐 KV 토큰 189개 (실제 행과 요청당 1개) |

수정 전 기본 백엔드는 `Repeat: one two three`에 `One!Two!Three!`로, 프랑스 질문에는 `The capital of France is... Paris! The City of Light, ...`처럼 장황하게 답했으며, 이는 기여자의 M5 보고와 일치합니다. dense 기준 출력은 두 빌드에서 바이트 단위로 같았습니다. 이로써 Neural Accelerator가 없는 하드웨어에서도 batched prefill의 최장 행 패딩이 풀 기반 출력을 오염시킨다는 이슈의 주장을 확인했습니다.

## 이 호스트에서 검증하지 못한 항목

- M5 타일 정렬 prefill 경로(`should_align_prefill()`가 true)는 Apple M5 하드웨어가 필요합니다. 기여자의 M5 Pro 결과(Llama-3.2-1B, SmolLM2-135M, Qwen3-0.6B, 수정 전 5개 중 0-1개, 수정 후 5개 중 5개)는 보고된 대로 받아들입니다.
- M5에서의 #1760 전체 프롬프트 재생 abort는 이 PR의 범위가 아닙니다.
26 changes: 19 additions & 7 deletions src/lib/mlxcel-core/src/cache.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2337,7 +2337,8 @@ impl KVCache {
/// touched. This mirrors speculative decoding's "rewind one block" pattern
/// from `update_turbo4_asym` and avoids paying for a re-quantize on the
/// common short-rewind case.
/// Used by: speculative decoding cache rewinds
/// A pool-backed cache rewinds its block table in the shared pool instead.
/// Used by: speculative decoding cache rewinds, padded-prefill pad trims
pub fn trim(&mut self, n: i32) -> i32 {
// Clamp against the live window length, not the monotonic offset:
// after a `trim_front`-induced live_start advance we must not roll
Expand All @@ -2350,12 +2351,23 @@ impl KVCache {
return 0;
}
// Pool-backed caches keep no dense `keys`/`values` buffers (#121); the
// block table is the authoritative store and is trimmed through the
// pool API (`CachePool::trim_paged_tokens` / `rewind_paged_tokens`),
// never the dense buffer slicing below (which would `unwrap` a `None`
// buffer). Treat a dense-side trim as a no-op for them.
if self.paged_backing.is_some() {
return 0;
// block table is the authoritative store, so rewind it instead of the
// dense buffer slicing below (which would `unwrap` a `None` buffer).
// `offset` moves with it: it is the next RoPE position and has to agree
// with the pool's write position.
if let Some(backing) = self.paged_backing.as_ref() {
let trimmed = backing
.pool
.borrow_mut()
.rewind_tokens(
&mut backing.state.borrow_mut(),
backing.layer_idx,
n as usize,
)
.expect("PagedBlockPool::rewind_tokens failed for pool-backed cache")
as i32;
self.offset -= trimmed;
return trimmed;
}
// Turbo4Delegated: hot-first trim. Tokens to remove from cold = max(0, n - hot_len).
// We adjust cold_offset and offset, then fall through to the per-mode buffer slicing
Expand Down
56 changes: 56 additions & 0 deletions src/lib/mlxcel-core/src/cache/paged_batch_decode_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -230,6 +230,62 @@ fn run_step(
(to_vec_f32(&out), to_vec_f32(&reference), stats)
}

/// A tile-padded prefill writes its pad rows and then trims them. The trim has
/// to reach the pool: while it was a no-op, decode read the pad rows as context
/// and rotated the next token at the padded position.
#[test]
fn trim_drops_padded_prefill_rows_from_a_pool_backed_cache() {
let (kv_heads, head_dim, real, padded) = (2, 64, 37, 64);
let (pool, states) = fresh_pool(1, kv_heads as usize, head_dim as usize);
let mut cache = KVCache::new_paged(pool.clone(), states[0].clone(), 0);
let mut rng = Rng::new(7);
let k = random_array(&mut rng, &[1, kv_heads, padded, head_dim]);
let v = random_array(&mut rng, &[1, kv_heads, padded, head_dim]);
let real_rows = |a: &MlxArray| {
to_vec_f32(&ffi::slice(
a,
&[0, 0, 0, 0],
&[1, kv_heads, real, head_dim],
))
};
let (k_real, v_real) = (real_rows(&k), real_rows(&v));
cache.update(k, v);

assert_eq!(cache.trim(padded - real), padded - real);
assert_eq!(
cache.offset, real,
"the next RoPE position is the real length"
);
assert_eq!(states[0].borrow().layer(0).unwrap().len, real as usize);
assert_eq!(
pool.borrow().allocated_block_count(),
2,
"37 tokens span two 32-token blocks"
);

// The next token lands right after the real rows, and the window decode
// gathers is the real rows plus that token.
let k_next = random_array(&mut rng, &[1, kv_heads, 1, head_dim]);
let v_next = random_array(&mut rng, &[1, kv_heads, 1, head_dim]);
let (k_next_rows, v_next_rows) = (to_vec_f32(&k_next), to_vec_f32(&v_next));
let (k_seen, v_seen) = cache.update_and_fetch(k_next, v_next);
assert_eq!(
ffi::array_shape(&k_seen),
vec![1, kv_heads, real + 1, head_dim]
);
assert_eq!(real_rows(&k_seen), k_real);
assert_eq!(real_rows(&v_seen), v_real);
let last_row = |a: &MlxArray| {
to_vec_f32(&ffi::slice(
a,
&[0, 0, real, 0],
&[1, kv_heads, real + 1, head_dim],
))
};
assert_eq!(last_row(&k_seen), k_next_rows);
assert_eq!(last_row(&v_seen), v_next_rows);
}

#[test]
fn batched_decode_matches_the_gather_path_above_the_floor() {
crate::test_support::kernel_ports::require_paged_attention_port!();
Expand Down
40 changes: 8 additions & 32 deletions src/server/batch/scheduler/decode_tick.rs
Original file line number Diff line number Diff line change
Expand Up @@ -422,10 +422,10 @@ impl BatchScheduler {
/// teardown before the synchronous path, a completion, or a prompt-cache
/// donation runs, so slot reuse and detach always see clean caches.
///
/// Pool-backed paged sequences rewind through the pool block table (one
/// token per layer, releasing any tail block); dense (and dense-natural
/// paged mirror) sequences trim the dense KV tail and re-mirror the shorter
/// length into the paged bookkeeping state.
/// Every cache trims through [`KVCache::trim`], which rewinds the pool block
/// table for a pool-backed sequence and the dense KV tail otherwise, moving
/// `offset` with it in both cases. Dense-natural paged mirrors then re-mirror
/// the shorter length into the paged bookkeeping state.
///
/// `positions` is the number of speculative appends to unwind: `1` for a
/// teardown before the step-n+1 prime forward has run (admission,
Expand All @@ -436,33 +436,9 @@ impl BatchScheduler {
if positions == 0 {
return;
}
let num_layers = self.model.num_layers();
let want = positions as i32;
for &seq_id in ids {
let paged_backed = self
.cache_pool
.get(seq_id)
.map(|s| s.caches.iter().any(|c| c.is_paged_backed()))
.unwrap_or(false);
if paged_backed {
for layer in 0..num_layers {
// A failed rewind silently leaks the speculative KV
// position(s), which would corrupt a later donation of this
// sequence's cache; surface it so the leak is diagnosable.
if let Err(err) = self
.cache_pool
.rewind_paged_tokens(seq_id, layer, positions)
{
tracing::warn!(
seq_id = %seq_id,
layer,
positions,
"lookahead teardown: paged rewind failed, speculative KV \
position may leak: {err}"
);
}
}
} else if let Some(caches) = self.cache_pool.get_caches_mut(seq_id) {
let want = positions as i32;
if let Some(caches) = self.cache_pool.get_caches_mut(seq_id) {
for (layer, cache) in caches.iter_mut().enumerate() {
// KVCache::trim clamps to the live window and returns the
// count actually removed; a short trim means a speculative
Expand All @@ -475,13 +451,13 @@ impl BatchScheduler {
layer,
requested = want,
trimmed,
"lookahead teardown: dense trim removed fewer positions \
"lookahead teardown: trim removed fewer positions \
than requested, KV may be out of sync"
);
}
}
// Re-mirror the shorter dense length into any paged bookkeeping
// (no-op for a pure dense pool).
// (no-op for a pure dense pool and for pool-backed sequences).
self.sync_sequence_storage(seq_id);
}
}
Expand Down
Loading