mod request;
mod schedule;
mod speculative;
pub use request::{InferenceRequest, Priority, SequenceGroup, Token};
pub use schedule::{BatchSchedule, BatcherStats, SchedulingPolicy, TokenOutput};
pub use speculative::{ExponentialMovingAverage, SpeculativeDecoder, SpeculativeOutput};
use std::collections::VecDeque;
use std::fmt;
use std::time::Instant;
use crate::paged_kv::SeqId;
#[derive(Debug)]
pub struct ContinuousBatcher {
max_batch_size: usize,
max_seq_len: usize,
running: Vec<SequenceGroup>,
waiting: VecDeque<SequenceGroup>,
swapped: Vec<SequenceGroup>,
policy: SchedulingPolicy,
stats: BatcherStats,
memory_threshold: f64,
}
impl ContinuousBatcher {
pub fn new(max_batch_size: usize, max_seq_len: usize) -> Self {
Self {
max_batch_size,
max_seq_len,
running: Vec::new(),
waiting: VecDeque::new(),
swapped: Vec::new(),
policy: SchedulingPolicy::default(),
stats: BatcherStats {
start_time: Some(Instant::now()),
..Default::default()
},
memory_threshold: 0.9,
}
}
pub fn with_policy(mut self, policy: SchedulingPolicy) -> Self {
self.policy = policy;
self
}
pub fn with_memory_threshold(mut self, threshold: f64) -> Self {
self.memory_threshold = threshold.clamp(0.0, 1.0);
self
}
pub fn policy(&self) -> &SchedulingPolicy {
&self.policy
}
pub fn max_batch_size(&self) -> usize {
self.max_batch_size
}
pub fn max_seq_len(&self) -> usize {
self.max_seq_len
}
pub fn running_count(&self) -> usize {
self.running.len()
}
pub fn waiting_count(&self) -> usize {
self.waiting.len()
}
pub fn swapped_count(&self) -> usize {
self.swapped.len()
}
pub fn stats(&self) -> &BatcherStats {
&self.stats
}
pub fn throughput(&self) -> f64 {
self.stats.throughput()
}
pub fn add_request(&mut self, request: InferenceRequest) {
let seq_group = SequenceGroup::new(request);
self.insert_waiting(seq_group);
}
fn insert_waiting(&mut self, seq_group: SequenceGroup) {
match &self.policy {
SchedulingPolicy::FCFS => {
self.waiting.push_back(seq_group);
}
SchedulingPolicy::SJF => {
let insert_idx = self
.waiting
.iter()
.position(|s| s.request.estimated_tokens > seq_group.request.estimated_tokens)
.unwrap_or(self.waiting.len());
self.waiting.insert(insert_idx, seq_group);
}
SchedulingPolicy::Priority { .. } => {
let insert_idx = self
.waiting
.iter()
.position(|s| s.request.priority < seq_group.request.priority)
.unwrap_or(self.waiting.len());
self.waiting.insert(insert_idx, seq_group);
}
SchedulingPolicy::FairShare => {
self.waiting.push_back(seq_group);
}
}
}
pub fn schedule(&mut self) -> BatchSchedule {
self.running.retain(|s| !s.is_finished);
let _available = self.max_batch_size.saturating_sub(self.running.len());
let mut prefill_count = 0;
while !self.waiting.is_empty() && self.running.len() < self.max_batch_size {
if let Some(seq_group) = self.waiting.pop_front() {
if seq_group.total_tokens() <= self.max_seq_len {
prefill_count += 1;
self.running.push(seq_group);
} else {
self.waiting.push_front(seq_group);
break;
}
}
}
while !self.swapped.is_empty() && self.running.len() < self.max_batch_size {
if let Some(seq_group) = self.swapped.pop() {
self.running.push(seq_group);
self.stats.total_swaps += 1;
}
}
let sequence_ids: Vec<SeqId> = self.running.iter().map(|s| s.request.id).collect();
let total_tokens: usize = self.running.iter().map(|s| s.total_tokens()).sum();
let decode_count = self.running.len() - prefill_count;
BatchSchedule {
batch_size: sequence_ids.len(),
sequence_ids,
total_tokens,
prefill_count,
decode_count,
}
}
pub fn process_outputs(&mut self, outputs: Vec<TokenOutput>) {
for output in outputs {
if let Some(seq_group) = self
.running
.iter_mut()
.find(|s| s.request.id == output.seq_id)
{
seq_group.add_token(output.token);
self.stats.total_tokens += 1;
if output.is_eos {
seq_group.finish();
self.stats.total_requests += 1;
}
}
}
}
pub fn preempt(&mut self, num_to_preempt: usize) -> Vec<SeqId> {
let mut preempted = Vec::new();
for _ in 0..num_to_preempt {
if self.running.is_empty() {
break;
}
let victim_idx = self
.running
.iter()
.enumerate()
.max_by_key(|(_, s)| s.total_tokens())
.map(|(i, _)| i);
if let Some(idx) = victim_idx {
let victim = self.running.remove(idx);
preempted.push(victim.request.id);
self.swapped.push(victim);
self.stats.total_preemptions += 1;
}
}
preempted
}
pub fn needs_preemption(&self, current_utilization: f64) -> bool {
current_utilization >= self.memory_threshold
&& !self.running.is_empty()
&& matches!(
self.policy,
SchedulingPolicy::Priority {
preempt_enabled: true
}
)
}
pub fn get_sequence(&self, seq_id: SeqId) -> Option<&SequenceGroup> {
self.running
.iter()
.chain(self.waiting.iter())
.chain(self.swapped.iter())
.find(|s| s.request.id == seq_id)
}
pub fn all_sequence_ids(&self) -> Vec<SeqId> {
self.running
.iter()
.chain(self.waiting.iter())
.chain(self.swapped.iter())
.map(|s| s.request.id)
.collect()
}
}
impl fmt::Display for ContinuousBatcher {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
writeln!(f, "ContinuousBatcher")?;
writeln!(
f,
" Policy: {} | Max Batch: {} | Max Seq: {}",
self.policy, self.max_batch_size, self.max_seq_len
)?;
writeln!(
f,
" Running: {} | Waiting: {} | Swapped: {}",
self.running.len(),
self.waiting.len(),
self.swapped.len()
)?;
writeln!(f, " Throughput: {:.1} tok/s", self.throughput())?;
writeln!(
f,
" Stats: tokens={}, requests={}, preemptions={}, swaps={}",
self.stats.total_tokens,
self.stats.total_requests,
self.stats.total_preemptions,
self.stats.total_swaps
)?;
Ok(())
}
}
#[cfg(test)]
mod tests;