Skip to content

perf: choose the prefill chunk by whether other sequences are decoding #2228

Description

@inureyes

Part of #2166.

Problem / Background

#2170 (PR #2205) made 2048 the prefill chunk on both front ends, decided by single-stream TTFT: at an 8192-token prompt, 512 was 8 to 29 percent slower (ADR 0007, "Prefill chunk" row). What 2048 costs concurrent decode streams was left for the end-of-run measurement (docs/CONTINUOUS_BATCHING.md:197). PR #2226 measured it (docs/benchmark_results/unified-engine-final-gb10-2026-10-08.md, "Prefill chunk under live decode"; raw admit-llama-c*-r*.txt and admit.sh in docs/benchmark_results/data/unified-engine-final-gb10-2026-10-08/): Llama-3.2-1B 4-bit, mlxcel serve --parallel 8 --ignore-eos, four streams decoding, then one 8192-token request admitted (scripts/bench_mixed_step_admission.py), three rounds.

Chunk Stream ITL p95, quiet Stream ITL p95 during admission Stream ITL mean during admission Admitted request TTFT
2048 (default) 18.6..19.5 ms 127.9..128.9 ms 18.7..19.3 ms 1903..2060 ms
512 17.2..19.4 ms 55.8..56.4 ms 13.4..13.9 ms 3768..4067 ms

Neither value wins both measures. Maintainer decision: the chunk depends on the scheduler state. 2048 when no other sequence is decoding, 512 when at least one is. An explicit --prefill-chunk-size (or its --batch-size / LLAMA_ARG_BATCH alias) or MLXCEL_PREFILL_CHUNK overrides both values.

Current Behavior

  • The policy is one number: mlxcel_core::prefill_plan::prefill_chunk_len() (src/lib/mlxcel-core/src/prefill_plan.rs:75): MLXCEL_PREFILL_CHUNK, else DEFAULT_PREFILL_CHUNK = 2048 (prefill_plan.rs:60). Users: DirectEngine (src/lib/mlxcel-core/src/engine/direct.rs:222), the Gemma 4 MTP prefill, the engine benchmark, and the server default.
  • The server flag is prefill_chunk_size: usize with default_value_t = prefill_chunk_len() (src/main.rs:1662, src/bin/mlx_server.rs:823-826), so an explicit --prefill-chunk-size 2048 cannot be told from the default. resolve_prefill_chunk_size (src/server/cli_input.rs:2276) infers "explicit" from "differs from the default". The value flows through ServerStartupConfig::prefill_chunk_size (src/server/startup.rs:278, default at :698), ServerConfig::prefill_chunk_size (src/server/config.rs:835, :1171) and the scheduler config (src/server/batch/scheduler/config.rs:63) into BatchScheduler::prefill_chunk_size.
  • BatchScheduler::prefill_plan_for (src/server/batch/scheduler/planned_prefill.rs:61) rebuilds the PrefillPlan from the sequence on every tick with that one chunk (PrefillPlan::with_prefix, prefill_plan.rs:253); SequenceInfo::prefill_offset is the cursor, and continue_chunked_prefill aborts when plan.piece_starting_at(seq.prefill_offset) is None (planned_prefill.rs:357-364). Pieces start at the adopted offset (or the history boundary) plus multiples of the chunk, so a cursor on the 512 grid is generally not a piece start of a 2048 plan: switching the chunk mid-prefill with today's constructor would abort the request.
  • run_prefill_pieces (planned_prefill.rs:198) counts a chunk and emits a prompt_progress frame only when plan.chunk().is_some().

Proposed Solution

  1. Policy type in src/lib/mlxcel-core/src/prefill_plan.rs:
    • pub const CONTENDED_PREFILL_CHUNK: usize = 512; next to DEFAULT_PREFILL_CHUNK (which stays 2048 and is the alone value).
    • #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct PrefillChunkPolicy { pub alone: usize, pub contended: usize } with pub fn resolve(explicit: Option<usize>) -> Self (Some(n) gives {n, n}, None gives {DEFAULT_PREFILL_CHUNK, CONTENDED_PREFILL_CHUNK}) and pub fn chunk(&self, others_decoding: bool) -> usize.
    • pub fn prefill_chunk_override() -> Option<usize>: MLXCEL_PREFILL_CHUNK parsed by the existing parse_prefill_chunk rules, read once per process; an unparseable value logs the existing warning and counts as unset. prefill_chunk_len() keeps its signature and returns prefill_chunk_override().unwrap_or(DEFAULT_PREFILL_CHUNK), so DirectEngine, the MTP prefill and mlxcel generate keep using the alone value.
  2. Resumable plan: PrefillPlan::resume_at(prompt_len, adopted, boundary, cursor, chunk, caps) -> PrefillPlan. When cursor <= adopted.min(prompt_len) it returns exactly with_prefix(prompt_len, adopted, boundary, chunk, caps). Otherwise (a continuation; the boundary segment, if any, already ran, so cursor >= boundary, debug-asserted) its pieces cut [cursor, prompt_len) into pieces of at most chunk starting at cursor, padded by the same rule as with_prefix; adopted(), boundary() and forwarded_len() report what with_prefix would; chunk() is Some(chunk) whenever chunk > 0, the model supports chunked prefill and the input is tokens, so the last piece of an already-chunked prefill still counts and reports progress. Property: when the chunk is unchanged, resume_at(.., cursor, chunk, ..).pieces() equals the pieces of with_prefix(.., chunk, ..) from cursor on.
  3. Scheduler: BatchScheduler holds a PrefillChunkPolicy instead of the single prefill_chunk_size used for plans. fn prefill_chunk_for_tick(&self) -> usize returns policy.chunk(!self.active_batch.is_empty()) (the parked chunked prefill lives in chunked_prefill_seq, not in active_batch). prefill_plan_for builds PrefillPlan::resume_at(len, seq.prefill_start_offset, self.history_boundary_split(seq), seq.prefill_offset, self.prefill_chunk_for_tick(), caps). Every caller (run_planned_prefill, start_chunked_prefill, continue_chunked_prefill, chunked_prefill_reserved_blocks at src/server/batch/scheduler/block_reclaim.rs:340) reads the state at call time; the continuation still reserves the blocks of the piece it actually runs before its forward (fix: close residual paged-budget gaps left by #2077 #2088). prefill_plan_for_with_chunk(&seq, 0) (the full-prefill partition, prefill.rs:704) is unchanged. The chunked_prefill_start span's chunk_size field reports the chosen chunk.
  4. Flags and config: --prefill-chunk-size becomes Option<usize> with no default_value_t in both src/main.rs and src/bin/mlx_server.rs, help text [default: 2048 when no other sequence is decoding, 512 when one is; MLXCEL_PREFILL_CHUNK overrides]. resolve_prefill_chunk_size takes Option<usize> and returns explicit: Option<usize> (flag, else --batch-size, else prefill_chunk_override()); the conflict warning fires when the flag and --batch-size are both given and differ. ServerStartupConfig / ServerConfig keep prefill_chunk_size as the alone value and gain prefill_chunk_size_contended; both equal the override when one is set. The probes that pin a chunk (src/server/engine_probe/server_engine.rs:189, src/bin/engine_parity.rs:343) set both.
  5. Unchanged on purpose: the batched-prefill budget default_batched_prefill_token_budget (src/server/batch/prefill_cohort.rs:286) is computed from the alone value and is already capped at BATCHED_PREFILL_ROW_TOKENS = 512 per row (prefill_cohort.rs:267), so its default stays 4096. mlxcel run and the chat REPL run a one-slot in-process server, so they always take the alone value, as today. mlxcel generate always uses the alone value.

Rejected: fixing the chunk per sequence at admission. A prefill admitted next to a stream that then finishes would keep 512 for the rest of the prompt and pay the slower TTFT with nothing to protect.

Scope

In scope: src/lib/mlxcel-core/src/prefill_plan.rs (+ prefill_plan_tests.rs), src/server/batch/scheduler/planned_prefill.rs, mod.rs, config.rs, src/server/cli_input.rs, src/server/startup.rs, src/server/config.rs, src/main.rs, src/commands/serve.rs, src/bin/mlx_server.rs, the two probes, docs/CONTINUOUS_BATCHING.md (flag table row at :34, the plan section at :178, and the "not measured yet" note at :197), and ADR 0007's "Prefill chunk" row (an amendment note citing this issue and the PR #2226 numbers).

Out of scope: the speculative burst prefill (BurstContext::prefill_chunk_size, src/server/batch/speculative_burst.rs:1103, and speculative_finalize.rs:228,343) keeps the alone value; the --prefill-grant-interval policy (#1011); per-request chunk overrides.

Implementation Notes

  • Reuse: PrefillPlan::with_prefix for the piece cutting and padding (factor the shared loop, do not copy it); parse_prefill_chunk for the env value; history_boundary_split.
  • Prompt cache invariant: PrefillPlan::reproduces (prefill_plan.rs:403) and docs/CONTINUOUS_BATCHING.md "One prefill plan, and when a prompt-cache hit reproduces a miss" stay true for prefills that run in the same scheduler state: a request prefilled alone is cut at 2048 from its start, exactly as today, so scheduler_prompt_cache_plan_tests::cache_hit_reproduces_miss_exactly_from_a_plan_split_point keeps passing unchanged. Document that a prefill whose chunk changed mid-prompt is a different partition (the near-tie class the doc already describes).
  • Edge cases: chunk 0 (override) disables chunking for both states; a prompt shorter than the chosen chunk is one piece; an embedding-input (VLM) prefill is never split, as today; a history-boundary segment stays one unpadded forward whatever the chunk; the contended chunk applies only to the pieces run while active_batch is non-empty.
  • Error handling: unchanged; continue_chunked_prefill must no longer reach its "no remaining tokens" abort because of a chunk change (covered by a test below).

Acceptance Criteria

  • On the PR docs: record the unified engine final measurement on GB10 #2226 scenario (admit.sh with the new build and no --prefill-chunk-size), live-stream ITL p95 during admission is within 10 percent of the 512 result (at most 62.0 ms, from 56.4 ms), three rounds.
  • Single-stream 8192-token TTFT on the server path (mlxcel-bench-engine --path server, no chunk flag) is inside the null range of an explicit --prefill-chunk 2048 arm, both models, under scripts/engine_bench_rounds.py (5 rounds, null arm, --hostgate).
  • prefill_plan_tests.rs: resume_at equals the with_prefix suffix at every split point (chunks 512 and 2048, with and without a history boundary, padding on and off); an 8192-token prompt cut at 512 for two pieces and resumed at cursor 1024 with 2048 yields 1024..3072, 3072..5120, 5120..7168, 7168..8192; chunk() is Some on a continuation whose remainder fits one chunk; PrefillChunkPolicy::resolve for None, Some(1024) and Some(0).
  • src/server/cli_input_tests.rs: --prefill-chunk-size 2048 --batch-size 1024 resolves to an explicit 2048 with the conflict flag set (fails today, because 2048 equals the default); no flag gives {2048, 512}; MLXCEL_PREFILL_CHUNK set gives the override for both (under the test env lock).
  • A scheduler test on a stub model that records prefill forward lengths: with one sequence decoding, a 1537-token prompt prefills in pieces of 512, 512, 512, 1; with none decoding, a 3000-token prompt prefills as 2048 then 952; a 5000-token prompt admitted next to one decoding sequence that finishes before the second piece runs prefills as 512, 2048, 2048, 392 and the request completes (no abort). The first case fails on current main.
  • Integrated: mlxcel serve and mlxcel-server use the policy by default, and the startup log line that reports prefill_chunk_size (src/server/startup.rs:3078) reports both values.
  • docs/CONTINUOUS_BATCHING.md and the ADR 0007 note describe the two values, the override, and the measurement.

Verification

cargo fmt --all -- --check
cargo clippy --workspace --all-targets --features cuda -- -D warnings
gpu-lock run --tag test -- cargo test --profile test-fast --features cuda -p mlxcel-core --lib prefill_plan -- --test-threads=1
gpu-lock run --tag test -- cargo test --profile test-fast --features cuda --lib server::cli_input -- --test-threads=1
gpu-lock run --tag verify -- make verify-test-cuda

# Real checkpoint, functional: a long prompt admitted next to a live stream completes
cargo build --release --features cuda --bin mlxcel --bin mlxcel-bench-engine
# Measurement, once, at the end, on a quiet host (F=<results data dir>):
gpu-lock run --tag admit -- env F=$F bash docs/benchmark_results/data/unified-engine-final-gb10-2026-10-08/admit.sh   # edited: new build only, no --prefill-chunk-size
for m in qwen3-1.7b-4bit llama-3.2-1b-instruct-4bit; do
  gpu-lock run --tag engine-bench -- python3 scripts/engine_bench_rounds.py --bin target/release/mlxcel-bench-engine \
    --model models/mlx/$m --prompt-tokens 8192 --rounds 5 --hostgate \
    --arm c2048="--path server --prefill-chunk 2048" --arm policy="--path server" --out ttft-$m.jsonl
done

A pass is p95 during admission at or under 62.0 ms in all three rounds, and a policy vs c2048 TTFT delta inside the null arm's range for both models.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    area:coremlxcel-core: MLX FFI, primitives, KV cache, layersarea:inferenceGeneration, sampling, decoding (incl. speculative, DRY)priority:mediumMedium prioritystatus:readyReady to be worked ontype:performancePerformance improvements

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions