Files
rustytorch/crates/production/rtx-inference/src/chunked_prefill.rs
T
osobhandClaude Sonnet 5 4aaa36a57a style: cargo fmt --workspace (whitespace/wrapping only, no semantic change)
Whole-workspace rustfmt pass picked up while iterating on Mamba GPU
backward work. Verified formatting-only via diff sampling; no logic
changed.

Co-Authored-By: Claude Sonnet 5 <[email protected]>
2026-08-10 07:09:36 -07:00

575 lines
20 KiB
Rust

//! Chunked prefill scheduler — interleaves prefill chunks with decode steps.
//!
//! For long prompts, processing the entire prompt in one shot causes quadratic
//! attention memory (`O(seq^2)`) and starves decode requests waiting in queue.
//! Chunked prefill (cf. arXiv:2309.06180) solves this by:
//!
//! - Splitting each prompt into fixed-size *chunks* (default: 512 tokens).
//! - Scheduling one chunk per step, optionally co-scheduling active decode
//! requests in the same GPU step.
//! - Bounding first-token latency for long prompts while sustaining throughput.
//!
//! # Example
//!
//! ```rust
//! use rtx_inference::chunked_prefill::{
//! ChunkedPrefillConfig, ChunkedPrefillScheduler, PrefillChunkState,
//! };
//!
//! let config = ChunkedPrefillConfig {
//! chunk_size: 512,
//! max_decode_tokens: 128,
//! interleave_decode: true,
//! };
//!
//! let mut scheduler = ChunkedPrefillScheduler::new(config);
//!
//! // Enqueue a 1024-token prompt.
//! scheduler.enqueue(42, 1024);
//!
//! // First step: 512 prefill tokens + up to 128 decode tokens.
//! let step = scheduler.next_step(8);
//! assert_eq!(step.prefill_tokens, 512);
//! assert_eq!(step.decode_tokens, 8); // min(128, 8)
//! assert_eq!(step.prefill_request_id, Some(42));
//!
//! // Second step: remaining 512 tokens.
//! let step = scheduler.next_step(8);
//! assert_eq!(step.prefill_tokens, 512);
//!
//! // Prefill is now complete; drain and confirm.
//! assert_eq!(scheduler.drain_completed(), 1);
//! assert!(scheduler.is_idle());
//! ```
/// Configuration for chunked prefill scheduling.
///
/// All fields have sensible defaults via [`Default`]:
/// - `chunk_size = 512`
/// - `max_decode_tokens = 128`
/// - `interleave_decode = true`
#[derive(Debug, Clone)]
pub struct ChunkedPrefillConfig {
/// Maximum number of tokens to process per prefill chunk.
///
/// Smaller values reduce peak attention memory at the cost of more
/// scheduling overhead. 512 is a good default for 80 GB A100s.
pub chunk_size: usize,
/// Maximum decode tokens to process alongside a prefill chunk.
///
/// When `interleave_decode` is `true`, each step runs at most this many
/// decode tokens in parallel with the prefill chunk. The actual count is
/// `min(max_decode_tokens, active_decode_count)`.
pub max_decode_tokens: usize,
/// Whether to interleave decode requests with prefill chunks.
///
/// Set to `false` to dedicate each GPU step entirely to prefill (useful
/// when decode latency is not a concern, e.g. offline batch processing).
pub interleave_decode: bool,
}
impl Default for ChunkedPrefillConfig {
fn default() -> Self {
Self {
chunk_size: 512,
max_decode_tokens: 128,
interleave_decode: true,
}
}
}
/// Progress state for a single request that is being chunked through prefill.
///
/// Created by [`ChunkedPrefillScheduler::enqueue`] and updated by each call
/// to [`ChunkedPrefillScheduler::next_step`].
#[derive(Debug, Clone)]
pub struct PrefillChunkState {
/// Unique identifier supplied by the caller (mirrors request IDs in the
/// surrounding scheduler infrastructure).
pub request_id: u64,
/// Total prompt tokens for this request.
pub total_tokens: usize,
/// How many tokens have already been prefilled.
pub processed_tokens: usize,
/// `true` once `processed_tokens == total_tokens`.
pub is_complete: bool,
}
impl PrefillChunkState {
/// Create a new state for `request_id` with `total_tokens` to process.
///
/// # Panics
///
/// Panics in debug builds if `total_tokens == 0`, which would be a
/// programming error (zero-length prompts should be rejected upstream).
#[must_use]
pub fn new(request_id: u64, total_tokens: usize) -> Self {
debug_assert!(total_tokens > 0, "total_tokens must be > 0");
Self {
request_id,
total_tokens,
processed_tokens: 0,
is_complete: false,
}
}
/// Advance by `n` tokens.
///
/// Clamps to `total_tokens` so callers do not need to guard against
/// overshooting. Sets `is_complete` when all tokens are processed.
pub fn advance(&mut self, n: usize) {
self.processed_tokens = self
.processed_tokens
.saturating_add(n)
.min(self.total_tokens);
self.is_complete = self.processed_tokens >= self.total_tokens;
}
/// Number of tokens still to be prefilled.
#[must_use]
pub fn remaining(&self) -> usize {
self.total_tokens.saturating_sub(self.processed_tokens)
}
/// Fraction of the prompt that has been processed, in `0.0..=1.0`.
///
/// Returns `1.0` for zero-length prompts to avoid a division-by-zero.
#[must_use]
pub fn progress(&self) -> f32 {
if self.total_tokens == 0 {
return 1.0;
}
self.processed_tokens as f32 / self.total_tokens as f32
}
/// Size of the next chunk to schedule, given `config`.
///
/// Always `<= config.chunk_size` and `<= self.remaining()`.
/// Returns `0` when `is_complete`.
#[must_use]
pub fn next_chunk_size(&self, config: &ChunkedPrefillConfig) -> usize {
self.remaining().min(config.chunk_size)
}
}
/// One scheduling decision returned by [`ChunkedPrefillScheduler::next_step`].
///
/// The executing engine should:
/// 1. Run `prefill_tokens` tokens of prompt `prefill_request_id`.
/// 2. Run `decode_tokens` tokens for the currently active decode requests.
#[derive(Debug, Clone)]
pub struct ChunkedStep {
/// Request whose prefill chunk is scheduled this step.
///
/// `None` when there is no active prefill (pure decode step).
pub prefill_request_id: Option<u64>,
/// Number of prefill tokens to process this step.
///
/// `0` when there is no active prefill request.
pub prefill_tokens: usize,
/// Number of decode tokens to process this step.
///
/// `0` when `interleave_decode` is `false` or `active_decode_count == 0`.
pub decode_tokens: usize,
/// Monotonically increasing step counter (zero-based).
pub step_index: usize,
}
/// Scheduler that splits long prompts into fixed-size chunks and optionally
/// interleaves decode requests within each step.
///
/// # Design notes
///
/// - Uses a `Vec<PrefillChunkState>` (not a `VecDeque`) because the active
/// set is typically very small (< 10 requests). FIFO ordering is preserved
/// by always picking `active[0]` as the head.
/// - Completed states are retained until [`drain_completed`] is called, which
/// mirrors the explicit lifecycle in production schedulers.
/// - All methods are `&mut self`; the struct is **not** `Send` by default
/// because its callers embed it inside a larger struct that owns the lock.
pub struct ChunkedPrefillScheduler {
config: ChunkedPrefillConfig,
/// Requests currently being chunked (includes completed ones until
/// [`drain_completed`] is called).
active: Vec<PrefillChunkState>,
/// Monotonically increasing step counter.
step_counter: usize,
/// Cumulative number of prefill chunks emitted.
chunks_scheduled: usize,
/// Cumulative number of tokens prefilled across all requests.
tokens_prefilled: usize,
}
impl ChunkedPrefillScheduler {
/// Create a new scheduler with the given configuration.
#[must_use]
pub fn new(config: ChunkedPrefillConfig) -> Self {
Self {
config,
active: Vec::new(),
step_counter: 0,
chunks_scheduled: 0,
tokens_prefilled: 0,
}
}
/// Enqueue `request_id` for chunked prefill of `total_tokens` prompt tokens.
///
/// The request will be scheduled FIFO after any currently active prefill
/// requests.
pub fn enqueue(&mut self, request_id: u64, total_tokens: usize) {
self.active
.push(PrefillChunkState::new(request_id, total_tokens));
}
/// Return the next [`ChunkedStep`] to execute.
///
/// **Prefill logic**: picks the first incomplete request in `active` and
/// schedules its next chunk. The state is advanced immediately so the
/// caller does not need to report completion back.
///
/// **Decode interleave**: when `config.interleave_decode` is `true` and
/// `active_decode_count > 0`, `decode_tokens` is set to
/// `min(config.max_decode_tokens, active_decode_count)`.
///
/// Returns a **pure decode step** (prefill fields zeroed, `prefill_request_id
/// = None`) when there is no active prefill work.
pub fn next_step(&mut self, active_decode_count: usize) -> ChunkedStep {
let step_index = self.step_counter;
self.step_counter += 1;
// Decode token count is independent of whether prefill is active.
let decode_tokens = if self.config.interleave_decode {
active_decode_count.min(self.config.max_decode_tokens)
} else {
0
};
// Find the first incomplete prefill request.
let head = self.active.iter_mut().find(|s| !s.is_complete);
match head {
Some(state) => {
let chunk = state.next_chunk_size(&self.config);
let request_id = state.request_id;
state.advance(chunk);
self.chunks_scheduled += 1;
self.tokens_prefilled += chunk;
ChunkedStep {
prefill_request_id: Some(request_id),
prefill_tokens: chunk,
decode_tokens,
step_index,
}
}
None => {
// No active prefill — pure decode step.
ChunkedStep {
prefill_request_id: None,
prefill_tokens: 0,
decode_tokens,
step_index,
}
}
}
}
/// Number of requests currently tracked (including completed ones not yet
/// drained).
#[must_use]
pub fn active_count(&self) -> usize {
self.active.len()
}
/// `true` when there are no incomplete prefill requests.
///
/// Note: returns `true` on an empty scheduler (nothing to do).
#[must_use]
pub fn is_idle(&self) -> bool {
self.active.iter().all(|s| s.is_complete)
}
/// Remove all completed requests and return how many were removed.
pub fn drain_completed(&mut self) -> usize {
let before = self.active.len();
self.active.retain(|s| !s.is_complete);
before - self.active.len()
}
/// Cumulative number of prefill chunks emitted by [`next_step`].
#[must_use]
pub fn chunks_scheduled(&self) -> usize {
self.chunks_scheduled
}
/// Cumulative number of tokens prefilled across all requests.
#[must_use]
pub fn tokens_prefilled(&self) -> usize {
self.tokens_prefilled
}
}
// ── Unit tests ────────────────────────────────────────────────────────────────
#[cfg(test)]
mod tests {
use super::*;
fn default_config() -> ChunkedPrefillConfig {
ChunkedPrefillConfig::default()
}
// ------------------------------------------------------------------
// PrefillChunkState tests
// ------------------------------------------------------------------
/// A 100-token request with chunk_size=512 should complete in exactly one
/// step.
#[test]
fn test_single_short_request_one_chunk() {
let config = default_config(); // chunk_size = 512
let mut scheduler = ChunkedPrefillScheduler::new(config);
scheduler.enqueue(1, 100);
let step = scheduler.next_step(0);
assert_eq!(step.prefill_request_id, Some(1));
assert_eq!(step.prefill_tokens, 100);
assert_eq!(step.step_index, 0);
// After one step the request is complete.
assert!(scheduler.is_idle());
assert_eq!(scheduler.tokens_prefilled(), 100);
assert_eq!(scheduler.chunks_scheduled(), 1);
}
/// A 1500-token request with chunk_size=512 must require exactly 3 steps:
/// 512 + 512 + 476 = 1500.
#[test]
fn test_long_request_multiple_chunks() {
let config = default_config(); // chunk_size = 512
let mut scheduler = ChunkedPrefillScheduler::new(config);
scheduler.enqueue(7, 1500);
let s0 = scheduler.next_step(0);
assert_eq!(s0.prefill_tokens, 512);
assert!(!scheduler.is_idle());
let s1 = scheduler.next_step(0);
assert_eq!(s1.prefill_tokens, 512);
assert!(!scheduler.is_idle());
let s2 = scheduler.next_step(0);
assert_eq!(s2.prefill_tokens, 476); // 1500 - 512 - 512
assert!(scheduler.is_idle());
assert_eq!(scheduler.tokens_prefilled(), 1500);
assert_eq!(scheduler.chunks_scheduled(), 3);
}
/// After advancing 256 of 1024 tokens, progress() should be exactly 0.25.
#[test]
fn test_progress_fraction() {
let mut state = PrefillChunkState::new(99, 1024);
state.advance(256);
let got = state.progress();
assert!((got - 0.25_f32).abs() < 1e-6, "expected 0.25, got {got}");
}
/// remaining() must equal total_tokens - processed_tokens after advancing.
#[test]
fn test_remaining_tokens() {
let mut state = PrefillChunkState::new(3, 800);
state.advance(300);
assert_eq!(state.remaining(), 500);
state.advance(500);
assert_eq!(state.remaining(), 0);
assert!(state.is_complete);
}
/// When only a partial chunk remains, next_chunk_size must return the
/// smaller remainder rather than the full chunk_size.
#[test]
fn test_next_chunk_size_last_chunk() {
let config = ChunkedPrefillConfig {
chunk_size: 512,
..Default::default()
};
let mut state = PrefillChunkState::new(5, 700);
state.advance(512); // first chunk consumed
let next = state.next_chunk_size(&config);
assert_eq!(next, 188); // 700 - 512 = 188 < 512
}
// ------------------------------------------------------------------
// Decode interleave tests
// ------------------------------------------------------------------
/// With interleave_decode=true, decode_tokens should be
/// min(max_decode_tokens, active_decode_count).
#[test]
fn test_interleave_decode_tokens() {
let config = ChunkedPrefillConfig {
chunk_size: 512,
max_decode_tokens: 128,
interleave_decode: true,
};
let mut scheduler = ChunkedPrefillScheduler::new(config);
scheduler.enqueue(10, 1000);
let step = scheduler.next_step(5); // 5 active decode requests
assert_eq!(step.decode_tokens, 5); // min(128, 5)
let step2 = scheduler.next_step(200); // more than max_decode_tokens
assert_eq!(step2.decode_tokens, 128); // clamped to max
}
/// With interleave_decode=false, decode_tokens must always be 0.
#[test]
fn test_no_decode_when_disabled() {
let config = ChunkedPrefillConfig {
chunk_size: 512,
max_decode_tokens: 128,
interleave_decode: false,
};
let mut scheduler = ChunkedPrefillScheduler::new(config);
scheduler.enqueue(20, 600);
for _ in 0..2 {
let step = scheduler.next_step(99);
assert_eq!(step.decode_tokens, 0, "decode must be 0 when disabled");
}
}
// ------------------------------------------------------------------
// Lifecycle / bookkeeping tests
// ------------------------------------------------------------------
/// After a request is fully prefilled, drain_completed() should remove it
/// and return 1.
#[test]
fn test_drain_completed_removes_done_requests() {
let config = default_config();
let mut scheduler = ChunkedPrefillScheduler::new(config);
scheduler.enqueue(42, 100); // completes in one step (< 512)
scheduler.next_step(0); // completes the request
assert!(scheduler.is_idle());
let removed = scheduler.drain_completed();
assert_eq!(removed, 1);
assert_eq!(scheduler.active_count(), 0);
}
/// active_count() must reflect the number of tracked requests (including
/// completed ones not yet drained).
#[test]
fn test_active_count() {
let config = default_config();
let mut scheduler = ChunkedPrefillScheduler::new(config);
assert_eq!(scheduler.active_count(), 0);
scheduler.enqueue(1, 50);
scheduler.enqueue(2, 50);
assert_eq!(scheduler.active_count(), 2);
// Complete both requests.
scheduler.next_step(0); // req 1 done
scheduler.next_step(0); // req 2 done
// Still 2 until explicitly drained.
assert_eq!(scheduler.active_count(), 2);
scheduler.drain_completed();
assert_eq!(scheduler.active_count(), 0);
}
/// is_idle() must return true once all requests have been drained.
#[test]
fn test_idle_after_all_complete() {
let config = default_config();
let mut scheduler = ChunkedPrefillScheduler::new(config);
scheduler.enqueue(55, 256);
assert!(!scheduler.is_idle());
scheduler.next_step(0); // chunk_size=512 > 256, so completes in one step
assert!(scheduler.is_idle());
scheduler.drain_completed();
// Empty scheduler is also idle.
assert!(scheduler.is_idle());
}
/// step_index must increment by 1 for every call to next_step.
#[test]
fn test_step_index_increments() {
let config = default_config();
let mut scheduler = ChunkedPrefillScheduler::new(config);
scheduler.enqueue(100, 2000);
for expected_idx in 0..5_usize {
let step = scheduler.next_step(0);
assert_eq!(
step.step_index, expected_idx,
"step_index should be {expected_idx}"
);
}
}
/// tokens_prefilled() must equal the sum of all chunk sizes emitted.
#[test]
fn test_cumulative_tokens_prefilled() {
let config = ChunkedPrefillConfig {
chunk_size: 200,
..Default::default()
};
let mut scheduler = ChunkedPrefillScheduler::new(config);
// Request A: 400 tokens → 2 chunks of 200 each.
scheduler.enqueue(1, 400);
// Request B: 300 tokens → 2 chunks of 200+100.
scheduler.enqueue(2, 300);
// Drive all chunks to completion.
while !scheduler.is_idle() {
scheduler.next_step(0);
}
assert_eq!(scheduler.tokens_prefilled(), 700); // 400 + 300
assert_eq!(scheduler.chunks_scheduled(), 4); // 2 + 2
}
// ------------------------------------------------------------------
// Extra edge-case tests
// ------------------------------------------------------------------
/// advance() must clamp and not overflow even if n >> total_tokens.
#[test]
fn test_advance_clamps_to_total() {
let mut state = PrefillChunkState::new(9, 100);
state.advance(999); // deliberately overshooting
assert_eq!(state.processed_tokens, 100);
assert!(state.is_complete);
assert_eq!(state.remaining(), 0);
}
/// A pure decode step (no enqueued prefill) must have zero prefill fields.
#[test]
fn test_pure_decode_step_when_no_prefill() {
let config = default_config();
let mut scheduler = ChunkedPrefillScheduler::new(config);
let step = scheduler.next_step(10);
assert_eq!(step.prefill_request_id, None);
assert_eq!(step.prefill_tokens, 0);
assert_eq!(step.decode_tokens, 10);
}
}