//! 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, /// 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` (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, /// 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); } }