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]>
575 lines
20 KiB
Rust
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);
|
|
}
|
|
}
|