use std::collections::{HashMap, HashSet, VecDeque};
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::mpsc::{self, Receiver, Sender};
use std::sync::{Arc, Mutex};
use std::thread::{self, JoinHandle};
use ferrox_core::cache::KvCache;
use ferrox_models::sampling::Sampler;
use ferrox_models::tokenizer::StopTokens;
use ferrox_models::Ceiling;
use ferrox_models::Decoder;
use crate::budget::ContextCeiling;
use crate::generate::{DecodeError, FinishReason, GenerationParams, Usage};
use crate::stop::{StopMatcher, StopStep};
type DecodeFn = Arc<dyn Fn(&[usize]) -> String + Send + Sync>;
type JobResult = Result<(FinishReason, Vec<usize>, String, Usage), DecodeError>;
pub const DEFAULT_PREFILL_CHUNK: usize = 128;
pub const DEFAULT_KV_BLOCK_SIZE: usize = 256;
pub const DEFAULT_MAX_QUEUE: usize = 512;
#[derive(Clone, Copy, Debug)]
pub struct BatcherConfig {
pub max_seqs: usize,
pub prefill_chunk: usize,
pub max_queue: usize,
pub kv_block_size: usize,
pub kv_blocks: Option<usize>,
pub max_context: Option<usize>,
}
impl Default for BatcherConfig {
fn default() -> Self {
BatcherConfig {
max_seqs: usize::MAX,
prefill_chunk: DEFAULT_PREFILL_CHUNK,
max_queue: DEFAULT_MAX_QUEUE,
kv_block_size: DEFAULT_KV_BLOCK_SIZE,
kv_blocks: None,
max_context: None,
}
}
}
impl BatcherConfig {
pub fn from_env() -> Self {
BatcherConfig {
max_seqs: env_positive("FERROX_CB_MAX_SEQS").unwrap_or(usize::MAX),
prefill_chunk: env_positive("FERROX_CB_PREFILL_CHUNK").unwrap_or(DEFAULT_PREFILL_CHUNK),
max_queue: env_positive("FERROX_CB_MAX_QUEUE").unwrap_or(DEFAULT_MAX_QUEUE),
kv_block_size: env_positive("FERROX_CB_KV_BLOCK_SIZE").unwrap_or(DEFAULT_KV_BLOCK_SIZE),
kv_blocks: env_positive("FERROX_CB_KV_BLOCKS"),
max_context: env_positive("FERROX_CB_MAX_CONTEXT"),
}
}
}
struct QueueGate {
depth: AtomicUsize,
cap: usize,
rejected: AtomicU64,
}
impl QueueGate {
fn new(cap: usize) -> Self {
QueueGate {
depth: AtomicUsize::new(0),
cap,
rejected: AtomicU64::new(0),
}
}
fn try_reserve(&self) -> Result<(), usize> {
let mut current = self.depth.load(Ordering::Acquire);
loop {
if current >= self.cap {
self.rejected.fetch_add(1, Ordering::Relaxed);
return Err(current);
}
match self.depth.compare_exchange_weak(
current,
current + 1,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => return Ok(()),
Err(actual) => current = actual,
}
}
}
fn release(&self) {
let previous = self.depth.fetch_sub(1, Ordering::AcqRel);
debug_assert!(previous > 0, "queue depth underflow");
}
fn depth(&self) -> usize {
self.depth.load(Ordering::Relaxed)
}
fn rejected(&self) -> u64 {
self.rejected.load(Ordering::Relaxed)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)]
struct AbortId(u64);
#[derive(Default)]
struct AbortInbox {
pending: Mutex<HashSet<AbortId>>,
next_id: AtomicU64,
aborted: AtomicU64,
}
impl AbortInbox {
fn next_id(&self) -> AbortId {
AbortId(self.next_id.fetch_add(1, Ordering::Relaxed))
}
fn enqueue(&self, id: AbortId) {
self.pending
.lock()
.unwrap_or_else(|p| p.into_inner())
.insert(id);
}
fn drain(&self) -> HashSet<AbortId> {
std::mem::take(&mut *self.pending.lock().unwrap_or_else(|p| p.into_inner()))
}
fn aborted(&self) -> u64 {
self.aborted.load(Ordering::Relaxed)
}
}
struct BlockBudget {
block_size: usize,
ceiling: Arc<ContextCeiling>,
total: Option<usize>,
free: AtomicUsize,
rejected_too_large: AtomicU64,
}
impl BlockBudget {
fn new(block_size: usize, total: Option<usize>, ceiling: Arc<ContextCeiling>) -> Self {
assert!(block_size > 0, "kv block size must be positive");
BlockBudget {
block_size,
ceiling,
total,
free: AtomicUsize::new(total.unwrap_or(0)),
rejected_too_large: AtomicU64::new(0),
}
}
fn bytes_for(&self, positions: usize) -> u64 {
self.ceiling.bytes_for(positions)
}
fn immovable_refusal(&self, positions: usize) -> Option<DecodeError> {
if let Some(err) = self.ceiling.refusal(positions) {
return Some(err);
}
let total = self.total?;
let blocks = self.blocks_for(positions);
if blocks <= total {
return None;
}
self.rejected_too_large.fetch_add(1, Ordering::Relaxed);
let limit_positions = total * self.block_size;
Some(DecodeError::KvBudgetExceeded {
binding: Ceiling::DeviceMemory.code(),
estimated_bytes: self.bytes_for(positions),
limit_bytes: self.bytes_for(limit_positions),
positions,
positions_limit: limit_positions,
detail: format!(
"request needs {blocks} KV blocks ({positions} token positions at {} per \
block) but this server's whole KV budget is {total} blocks; an idle server \
would refuse it identically",
self.block_size
),
})
}
fn blocks_for(&self, positions: usize) -> usize {
positions.div_ceil(self.block_size).max(1)
}
fn try_reserve(&self, blocks: usize) -> bool {
if self.total.is_none() {
return true;
}
let free = self.free.load(Ordering::Relaxed);
if blocks > free {
return false;
}
self.free.store(free - blocks, Ordering::Relaxed);
true
}
fn release(&self, blocks: usize) {
let Some(total) = self.total else {
return;
};
let free = self.free.load(Ordering::Relaxed);
debug_assert!(
free + blocks <= total,
"released more blocks than were ever reserved"
);
self.free
.store((free + blocks).min(total), Ordering::Relaxed);
}
fn free(&self) -> usize {
self.total
.map(|_| self.free.load(Ordering::Relaxed))
.unwrap_or(0)
}
}
fn env_positive(name: &str) -> Option<usize> {
let raw = std::env::var(name).ok()?;
let value: usize = raw
.parse()
.unwrap_or_else(|_| panic!("{name} must be a positive integer"));
assert!(value > 0, "{name} must be a positive integer");
Some(value)
}
#[derive(Default)]
struct Counters {
prefill_chunks: AtomicU64,
prefill_tokens: AtomicU64,
decode_steps: AtomicU64,
peak_blocks: AtomicUsize,
}
#[derive(Debug, Clone, Copy, Default, serde::Serialize)]
pub struct BatcherStats {
pub prefill_chunks: u64,
pub prefill_tokens: u64,
pub decode_steps: u64,
pub queue_depth: usize,
pub queue_rejected: u64,
pub kv_blocks_total: usize,
pub kv_blocks_free: usize,
pub kv_block_size: usize,
pub kv_rejected_too_large: u64,
pub kv_rejected_context_length: u64,
pub kv_blocks_peak: usize,
pub aborted: u64,
}
pub struct PrefillState {
decoder: Arc<Decoder>,
caches: Vec<KvCache>,
tokens: Vec<usize>,
tokens_processed: usize,
logits: Vec<f32>,
chunk_size: usize,
}
impl PrefillState {
pub fn new(decoder: Arc<Decoder>, prompt_tokens: &[usize], chunk_size: usize) -> Self {
assert!(chunk_size > 0, "prefill chunk size must be positive");
let caches: Vec<KvCache> = decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
let tokens = if prompt_tokens.is_empty() {
vec![0]
} else {
prompt_tokens.to_vec()
};
PrefillState {
decoder,
caches,
tokens,
tokens_processed: 0,
logits: Vec::new(),
chunk_size,
}
}
pub fn tokens_processed(&self) -> usize {
self.tokens_processed
}
pub fn tokens_remaining(&self) -> usize {
self.tokens.len() - self.tokens_processed
}
pub fn is_done(&self) -> bool {
self.tokens_remaining() == 0
}
pub fn step_chunk(&mut self) -> bool {
let end = (self.tokens_processed + self.chunk_size).min(self.tokens.len());
for pos in self.tokens_processed..end {
self.logits = self
.decoder
.forward_token(self.tokens[pos], pos, &mut self.caches);
}
self.tokens_processed = end;
self.is_done()
}
fn into_decode_start(self) -> (Vec<KvCache>, Vec<f32>, usize) {
debug_assert!(self.is_done(), "prefill must finish before decoding");
(self.caches, self.logits, self.tokens_processed)
}
}
struct Job {
prompt_tokens: Vec<usize>,
params: GenerationParams,
stop_tokens: StopTokens,
reply: Sender<JobResult>,
abort: AbortId,
blocks: usize,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)]
struct Uid(u64);
#[derive(Default)]
struct Rows {
state: HashMap<Uid, Slot>,
order: Vec<Uid>,
next_uid: u64,
}
impl Rows {
fn insert(&mut self, slot: Slot) -> Uid {
let uid = Uid(self.next_uid);
self.next_uid += 1;
self.state.insert(uid, slot);
self.order.push(uid);
uid
}
fn len(&self) -> usize {
self.order.len()
}
fn blocks_held(&self) -> usize {
self.state.values().map(|slot| slot.blocks).sum()
}
fn mark_cancelled(&mut self, ids: &HashSet<AbortId>) -> Vec<AbortId> {
let mut consumed = Vec::new();
for uid in &self.order {
if let Some(slot) = self.state.get_mut(uid) {
if ids.contains(&slot.abort) && slot.finish.is_none() {
slot.finish = Some(FinishReason::Cancelled);
consumed.push(slot.abort);
}
}
}
consumed
}
fn is_empty(&self) -> bool {
self.order.is_empty()
}
fn get(&self, uid: Uid) -> Option<&Slot> {
self.state.get(&uid)
}
fn get_mut(&mut self, uid: Uid) -> Option<&mut Slot> {
self.state.get_mut(&uid)
}
fn remove(&mut self, uid: Uid) -> Option<Slot> {
self.order.retain(|&u| u != uid);
self.state.remove(&uid)
}
fn ready(&self) -> Vec<Uid> {
self.order
.iter()
.copied()
.filter(|uid| {
self.state
.get(uid)
.is_some_and(|s| s.finish.is_none() && s.generated_ids.len() < s.max_tokens)
})
.collect()
}
fn flush_finished(&mut self, budget: &BlockBudget) {
let finished: Vec<Uid> = self
.order
.iter()
.copied()
.filter(|uid| self.state.get(uid).is_some_and(|s| s.finish.is_some()))
.collect();
for uid in finished {
if let Some(slot) = self.remove(uid) {
budget.release(slot.blocks);
reply_finished(slot);
}
}
}
}
struct Slot {
caches: Vec<KvCache>,
pos: usize,
logits: Vec<f32>,
sampler: Sampler,
generated_ids: Vec<usize>,
visible: String,
stops: StopMatcher,
prompt_tokens: usize,
max_tokens: usize,
stop_tokens: StopTokens,
params: GenerationParams,
reply: Sender<JobResult>,
finish: Option<FinishReason>,
abort: AbortId,
blocks: usize,
}
#[derive(Clone)]
pub struct ContinuousBatcher {
tx: Sender<Job>,
counters: Arc<Counters>,
queue: Arc<QueueGate>,
budget: Arc<BlockBudget>,
aborts: Arc<AbortInbox>,
}
struct WorkerGuard {
_join: JoinHandle<()>,
}
impl ContinuousBatcher {
#[cfg(test)]
pub fn spawn_with_config(
decoder: Arc<Decoder>,
decode: DecodeFn,
config: BatcherConfig,
) -> Self {
let ceiling = Arc::new(ContextCeiling::new(
config.max_context,
ferrox_models::KvShape::from_config(&decoder.config, ferrox_models::KvElem::F32, 1),
));
Self::spawn_with_ceiling(decoder, decode, config, ceiling)
}
pub fn spawn_with_ceiling(
decoder: Arc<Decoder>,
decode: DecodeFn,
config: BatcherConfig,
ceiling: Arc<ContextCeiling>,
) -> Self {
let (tx, rx) = mpsc::channel::<Job>();
let counters = Arc::new(Counters::default());
let queue = Arc::new(QueueGate::new(config.max_queue));
let budget = Arc::new(BlockBudget::new(
config.kv_block_size,
config.kv_blocks,
ceiling,
));
let aborts = Arc::new(AbortInbox::default());
let worker_counters = Arc::clone(&counters);
let worker_queue = Arc::clone(&queue);
let worker_budget = Arc::clone(&budget);
let worker_aborts = Arc::clone(&aborts);
let _join = thread::Builder::new()
.name("ferrox-continuous-batch".into())
.spawn(move || {
worker_loop(
decoder,
decode,
rx,
config,
worker_counters,
worker_queue,
worker_budget,
worker_aborts,
)
})
.expect("spawn continuous-batch worker");
let _guard: &'static WorkerGuard = Box::leak(Box::new(WorkerGuard { _join }));
ContinuousBatcher {
tx,
counters,
queue,
budget,
aborts,
}
}
pub fn stats(&self) -> BatcherStats {
BatcherStats {
prefill_chunks: self.counters.prefill_chunks.load(Ordering::Relaxed),
prefill_tokens: self.counters.prefill_tokens.load(Ordering::Relaxed),
decode_steps: self.counters.decode_steps.load(Ordering::Relaxed),
queue_depth: self.queue.depth(),
queue_rejected: self.queue.rejected(),
kv_blocks_total: self.budget.total.unwrap_or(0),
kv_blocks_free: self.budget.free(),
kv_block_size: self.budget.block_size,
kv_rejected_too_large: self.budget.rejected_too_large.load(Ordering::Relaxed),
kv_rejected_context_length: self.budget.ceiling.refused(),
kv_blocks_peak: self.counters.peak_blocks.load(Ordering::Relaxed),
aborted: self.aborts.aborted(),
}
}
pub fn generate(
&self,
prompt_tokens: Vec<usize>,
params: GenerationParams,
stop_tokens: StopTokens,
) -> Result<(FinishReason, Vec<usize>, String, Usage), DecodeError> {
let positions = prompt_tokens.len().saturating_add(params.max_tokens);
let blocks = self.budget.blocks_for(positions);
if let Some(refusal) = self.budget.immovable_refusal(positions) {
return Err(refusal);
}
self.queue
.try_reserve()
.map_err(|queued| DecodeError::QueueFull {
queued,
cap: self.queue.cap,
})?;
let (reply_tx, reply_rx) = mpsc::channel();
let abort = self.aborts.next_id();
let cancel = params.cancel.clone();
if self
.tx
.send(Job {
prompt_tokens,
params,
stop_tokens,
reply: reply_tx,
abort,
blocks,
})
.is_err()
{
self.queue.release();
return Err(DecodeError::KvPoolExhausted);
}
if let Some(token) = cancel {
let inbox = Arc::clone(&self.aborts);
token.on_cancel(move || inbox.enqueue(abort));
}
reply_rx.recv().unwrap_or(Err(DecodeError::KvPoolExhausted))
}
}
fn drain_channel(rx: &Receiver<Job>, waiting: &mut VecDeque<Job>) {
while let Ok(job) = rx.try_recv() {
waiting.push_back(job);
}
}
fn admit(
decoder: &Arc<Decoder>,
waiting: &mut VecDeque<Job>,
prefills: &mut VecDeque<Prefill>,
decoding: usize,
config: &BatcherConfig,
queue: &QueueGate,
budget: &BlockBudget,
) {
while let Some(job) = waiting.front() {
if decoding + prefills.len() >= config.max_seqs {
break;
}
let blocks = job.blocks;
if !budget.try_reserve(blocks) {
break;
}
let job = waiting.pop_front().expect("front() just succeeded");
queue.release();
match accept(decoder, job, config.prefill_chunk) {
Some(prefill) => prefills.push_back(prefill),
None => budget.release(blocks),
}
}
}
fn apply_aborts(
inbox: &AbortInbox,
carried: &mut HashSet<AbortId>,
waiting: &mut VecDeque<Job>,
prefills: &mut VecDeque<Prefill>,
rows: &mut Rows,
queue: &QueueGate,
budget: &BlockBudget,
) {
carried.extend(inbox.drain());
if carried.is_empty() {
return;
}
let mut stopped = 0u64;
let mut still_waiting = VecDeque::with_capacity(waiting.len());
while let Some(job) = waiting.pop_front() {
if carried.remove(&job.abort) {
queue.release();
let _ = job.reply.send(Ok((
FinishReason::Cancelled,
Vec::new(),
String::new(),
Usage::new(job.prompt_tokens.len(), 0),
)));
stopped += 1;
} else {
still_waiting.push_back(job);
}
}
*waiting = still_waiting;
let mut still_prefilling = VecDeque::with_capacity(prefills.len());
while let Some(prefill) = prefills.pop_front() {
if carried.remove(&prefill.abort) {
budget.release(prefill.blocks);
let _ = prefill.reply.send(Ok((
FinishReason::Cancelled,
Vec::new(),
String::new(),
Usage::new(prefill.prompt_tokens, 0),
)));
stopped += 1;
} else {
still_prefilling.push_back(prefill);
}
}
*prefills = still_prefilling;
for id in rows.mark_cancelled(carried) {
carried.remove(&id);
stopped += 1;
}
if stopped > 0 {
inbox.aborted.fetch_add(stopped, Ordering::Relaxed);
}
}
#[allow(clippy::too_many_arguments)]
fn worker_loop(
decoder: Arc<Decoder>,
decode: DecodeFn,
rx: Receiver<Job>,
config: BatcherConfig,
counters: Arc<Counters>,
queue: Arc<QueueGate>,
budget: Arc<BlockBudget>,
aborts: Arc<AbortInbox>,
) {
let mut rows = Rows::default();
let mut prefills: VecDeque<Prefill> = VecDeque::new();
let mut waiting: VecDeque<Job> = VecDeque::new();
let mut carried_aborts: HashSet<AbortId> = HashSet::new();
loop {
if rows.is_empty() && prefills.is_empty() && waiting.is_empty() {
match rx.recv() {
Ok(job) => waiting.push_back(job),
Err(_) => break,
}
}
drain_channel(&rx, &mut waiting);
apply_aborts(
&aborts,
&mut carried_aborts,
&mut waiting,
&mut prefills,
&mut rows,
&queue,
&budget,
);
rows.flush_finished(&budget);
admit(
&decoder,
&mut waiting,
&mut prefills,
rows.len(),
&config,
&queue,
&budget,
);
let held = rows.blocks_held() + prefills.iter().map(|p| p.blocks).sum::<usize>();
counters.peak_blocks.fetch_max(held, Ordering::Relaxed);
debug_assert!(
!(rows.is_empty() && prefills.is_empty() && !waiting.is_empty()),
"idle worker cannot admit its queue: {} jobs stuck",
waiting.len()
);
if let Some(mut prefill) = prefills.pop_front() {
let before = prefill.state.tokens_processed();
let done = prefill.state.step_chunk();
counters.prefill_chunks.fetch_add(1, Ordering::Relaxed);
counters.prefill_tokens.fetch_add(
(prefill.state.tokens_processed() - before) as u64,
Ordering::Relaxed,
);
if done {
rows.insert(prefill.into_slot());
} else {
prefills.push_back(prefill);
}
}
if rows.is_empty() {
continue;
}
let ready = rows.ready();
if ready.is_empty() {
rows.flush_finished(&budget);
continue;
}
let mut active: Vec<Uid> = Vec::with_capacity(ready.len());
for uid in ready {
let Some(slot) = rows.get_mut(uid) else {
continue;
};
let next =
slot.sampler
.sample(&slot.logits, &slot.params.sampling, &slot.generated_ids);
if slot.stop_tokens.contains(next) || slot.stops.is_stop_token(next) {
slot.finish = Some(FinishReason::Stop);
continue;
}
slot.generated_ids.push(next);
let piece = decode(&[next]);
if apply_stop_buffer(slot, &piece) {
continue;
}
active.push(uid);
}
if !active.is_empty() {
let tokens: Vec<usize> = active
.iter()
.map(|&uid| *rows.get(uid).unwrap().generated_ids.last().unwrap())
.collect();
let positions: Vec<usize> = active
.iter()
.map(|&uid| rows.get(uid).unwrap().pos)
.collect();
let mut cache_refs: Vec<Vec<KvCache>> = active
.iter()
.map(|&uid| std::mem::take(&mut rows.get_mut(uid).unwrap().caches))
.collect();
let logits_batch = decoder.forward_multi_seq(&tokens, &positions, &mut cache_refs);
counters.decode_steps.fetch_add(1, Ordering::Relaxed);
for (j, &uid) in active.iter().enumerate() {
let slot = rows
.get_mut(uid)
.expect("an active row cannot vanish mid-step");
slot.caches = std::mem::take(&mut cache_refs[j]);
slot.logits = logits_batch[j].clone();
slot.pos += 1;
if slot.generated_ids.len() >= slot.max_tokens {
slot.finish = Some(FinishReason::Length);
}
}
}
rows.flush_finished(&budget);
}
}
fn apply_stop_buffer(slot: &mut Slot, piece: &str) -> bool {
match slot.stops.push(piece) {
StopStep::Emit(text) => {
slot.visible.push_str(&text);
false
}
StopStep::Matched(text) => {
slot.visible.push_str(&text);
slot.finish = Some(FinishReason::Stop);
true
}
}
}
struct Prefill {
state: PrefillState,
prompt_tokens: usize,
params: GenerationParams,
stop_tokens: StopTokens,
reply: Sender<JobResult>,
abort: AbortId,
blocks: usize,
}
impl Prefill {
fn into_slot(self) -> Slot {
let Prefill {
state,
prompt_tokens,
params,
stop_tokens,
reply,
abort,
blocks,
} = self;
let (caches, logits, pos) = state.into_decode_start();
Slot {
caches,
pos,
logits,
sampler: Sampler::new(params.seed),
generated_ids: Vec::with_capacity(params.max_tokens),
visible: String::new(),
stops: StopMatcher::new(¶ms.stop, ¶ms.stop_token_ids),
prompt_tokens,
max_tokens: params.max_tokens,
stop_tokens,
params,
reply,
finish: None,
abort,
blocks,
}
}
}
fn accept(decoder: &Arc<Decoder>, job: Job, chunk_size: usize) -> Option<Prefill> {
let vocab_size = decoder.config.vocab_size;
if let Some(&bad) = job.prompt_tokens.iter().find(|&&t| t >= vocab_size) {
let _ = job.reply.send(Err(DecodeError::TokenOutOfVocab {
token: bad,
vocab_size,
}));
return None;
}
Some(Prefill {
state: PrefillState::new(Arc::clone(decoder), &job.prompt_tokens, chunk_size),
prompt_tokens: job.prompt_tokens.len(),
params: job.params,
stop_tokens: job.stop_tokens,
reply: job.reply,
abort: job.abort,
blocks: job.blocks,
})
}
fn reply_finished(mut slot: Slot) {
let finish = slot.finish.expect("only a finished row is replied to");
let tail = slot.stops.flush();
slot.visible.push_str(&tail);
let usage = Usage::new(slot.prompt_tokens, slot.generated_ids.len());
let _ = slot
.reply
.send(Ok((finish, slot.generated_ids, slot.visible, usage)));
}
#[cfg(test)]
mod tests {
use super::*;
use crate::cancel::CancelToken;
use ferrox_models::config::test_dense_fixture;
use ferrox_models::sampling::SamplingParams;
use std::sync::{Barrier, Mutex};
fn tiny_decoder() -> Arc<Decoder> {
let cfg = test_dense_fixture();
let vocab = cfg.vocab_size;
Arc::new(Decoder::new_random_small(cfg, 2, vocab))
}
fn greedy_params(max_tokens: usize, seed: u64) -> GenerationParams {
GenerationParams {
max_tokens,
sampling: SamplingParams {
temperature: 0.0,
top_p: 1.0,
top_k: 0,
repetition_penalty: 1.0,
presence_penalty: 0.0,
frequency_penalty: 0.0,
},
seed,
stop: vec![],
stop_token_ids: Vec::new(),
json_object: false,
cancel: None,
}
}
fn identity_decode() -> DecodeFn {
Arc::new(|ids: &[usize]| {
ids.iter()
.map(|id| char::from_u32(65 + (*id as u32 % 26)).unwrap_or('?'))
.collect()
})
}
fn sequential_ids(
decoder: &Decoder,
prompt: &[usize],
params: &GenerationParams,
) -> Vec<usize> {
let mut caches: Vec<KvCache> = decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
let mut pos = 0;
let mut logits = Vec::new();
for &tok in prompt {
logits = decoder.forward_token(tok, pos, &mut caches);
pos += 1;
}
let mut sampler = Sampler::new(params.seed);
let mut generated = Vec::new();
for _ in 0..params.max_tokens {
let next = sampler.sample(&logits, ¶ms.sampling, &generated);
generated.push(next);
logits = decoder.forward_token(next, pos, &mut caches);
pos += 1;
}
generated
}
#[test]
fn continuous_batch_matches_sequential_generate_token_ids() {
let decoder = tiny_decoder();
let prompts: [Vec<usize>; 2] = [vec![1, 2, 3], vec![4, 5]];
let params = [greedy_params(8, 7), greedy_params(5, 11)];
let sequential: Vec<Vec<usize>> = prompts
.iter()
.zip(params.iter())
.map(|(p, par)| sequential_ids(&decoder, p, par))
.collect();
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
identity_decode(),
BatcherConfig {
prefill_chunk: 1,
..BatcherConfig::default()
},
);
let barrier = Arc::new(Barrier::new(3));
let results = Arc::new(Mutex::new(vec![None, None]));
let mut threads = Vec::new();
for i in 0..2 {
let batcher = batcher.clone();
let barrier = Arc::clone(&barrier);
let results = Arc::clone(&results);
let prompt = prompts[i].clone();
let par = GenerationParams {
max_tokens: params[i].max_tokens,
sampling: SamplingParams {
temperature: params[i].sampling.temperature,
top_p: params[i].sampling.top_p,
top_k: params[i].sampling.top_k,
repetition_penalty: params[i].sampling.repetition_penalty,
presence_penalty: params[i].sampling.presence_penalty,
frequency_penalty: params[i].sampling.frequency_penalty,
},
seed: params[i].seed,
stop: vec![],
stop_token_ids: Vec::new(),
json_object: params[i].json_object,
cancel: params[i].cancel.clone(),
};
threads.push(thread::spawn(move || {
barrier.wait();
let out = batcher
.generate(prompt, par, StopTokens::default())
.expect("batch generate");
results.lock().unwrap()[i] = Some(out.1);
}));
}
barrier.wait();
for t in threads {
t.join().unwrap();
}
let got = results.lock().unwrap();
assert_eq!(got[0].as_ref().unwrap(), &sequential[0]);
assert_eq!(got[1].as_ref().unwrap(), &sequential[1]);
}
#[test]
fn continuous_batch_honors_stop_sequence_in_decoded_text() {
let decoder = tiny_decoder();
let decode: DecodeFn = Arc::new(|ids: &[usize]| {
ids.iter()
.map(|id| match id % 3 {
0 => 'X',
1 => 'Y',
_ => 'Z',
})
.collect()
});
let prompt = vec![1usize, 2, 3];
let mut params = greedy_params(32, 3);
let ids = sequential_ids(&decoder, &prompt, ¶ms);
let full: String = ids
.iter()
.map(|id| match id % 3 {
0 => 'X',
1 => 'Y',
_ => 'Z',
})
.collect();
assert!(
full.len() >= 4,
"need enough tokens to place a mid-stream stop"
);
let stop = full[2..4].to_string();
params.stop = vec![stop.clone()];
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
decode,
BatcherConfig {
prefill_chunk: 2,
..BatcherConfig::default()
},
);
let (finish, _ids, text, _usage) = batcher
.generate(prompt, params, StopTokens::default())
.expect("batch generate");
assert_eq!(finish, FinishReason::Stop);
assert!(
!text.contains(&stop),
"stop string must be trimmed from visible text: text={text:?} stop={stop:?}"
);
assert_eq!(&full[..full.find(&stop).unwrap()], text);
}
#[test]
fn continuous_batch_stops_on_any_member_of_the_stop_set() {
let decoder = tiny_decoder();
let decode: DecodeFn = Arc::new(|_: &[usize]| String::new());
let prompt = vec![1usize, 2, 3];
let params = greedy_params(32, 3);
let ids = sequential_ids(&decoder, &prompt, ¶ms);
assert!(ids.len() > 3, "need a mid-stream token to stop on");
let turn_ender = ids[2];
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
decode,
BatcherConfig {
prefill_chunk: 2,
..BatcherConfig::default()
},
);
let (finish, got, _text, usage) = batcher
.generate(prompt, params, StopTokens::from_eos(Some(turn_ender)))
.expect("batch generate");
assert_eq!(finish, FinishReason::Stop);
assert_eq!(got, ids[..2].to_vec());
assert_eq!(usage.completion_tokens, 2);
}
#[test]
fn prefill_step_chunk_is_bounded_and_resumable() {
let decoder = tiny_decoder();
let prompt: Vec<usize> = (1..=7).collect();
let mut state = PrefillState::new(Arc::clone(&decoder), &prompt, 3);
assert_eq!(state.tokens_remaining(), 7);
assert_eq!(state.tokens_processed(), 0);
assert!(!state.step_chunk());
assert_eq!(state.tokens_processed(), 3, "a chunk may not overrun");
assert_eq!(state.tokens_remaining(), 4);
assert!(!state.step_chunk());
assert_eq!(state.tokens_processed(), 6);
assert!(state.step_chunk(), "final short chunk finishes the prompt");
assert_eq!(state.tokens_processed(), 7);
assert_eq!(state.tokens_remaining(), 0);
assert!(state.is_done());
assert!(state.step_chunk(), "stepping a finished prefill is a no-op");
assert_eq!(state.tokens_processed(), 7);
}
#[test]
fn empty_prompt_prefills_one_stand_in_token() {
let decoder = tiny_decoder();
let mut state = PrefillState::new(Arc::clone(&decoder), &[], 4);
assert_eq!(state.tokens_remaining(), 1);
assert!(state.step_chunk());
let (_caches, logits, pos) = state.into_decode_start();
assert_eq!(pos, 1);
assert_eq!(logits.len(), decoder.config.vocab_size);
}
#[test]
fn prefill_chunking_does_not_change_logits() {
let decoder = tiny_decoder();
let prompt: Vec<usize> = (0..11).map(|i| (i * 3 + 1) % 16).collect();
let mut sequential: Vec<f32> = Vec::new();
let mut caches: Vec<KvCache> = decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect();
for (pos, &tok) in prompt.iter().enumerate() {
sequential = decoder.forward_token(tok, pos, &mut caches);
}
for chunk in [1usize, 2, 5, 11, 64] {
let mut state = PrefillState::new(Arc::clone(&decoder), &prompt, chunk);
while !state.step_chunk() {}
let (_caches, logits, pos) = state.into_decode_start();
assert_eq!(pos, prompt.len());
assert_eq!(
logits, sequential,
"chunk size {chunk} changed the prefill logits"
);
}
}
#[test]
fn long_prefill_does_not_freeze_an_in_flight_decode() {
let decoder = tiny_decoder();
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
identity_decode(),
BatcherConfig {
prefill_chunk: 1,
..BatcherConfig::default()
},
);
let decode_job = {
let batcher = batcher.clone();
thread::spawn(move || {
batcher.generate(vec![1, 2], greedy_params(90, 5), StopTokens::default())
})
};
let deadline = std::time::Instant::now() + std::time::Duration::from_secs(60);
while batcher.stats().decode_steps < 2 {
assert!(std::time::Instant::now() < deadline, "decode never started");
thread::yield_now();
}
let long_prompt: Vec<usize> = (0..40).map(|i| (i % 16) + 1).collect();
let total = long_prompt.len() as u64;
let prefill_at_submit = batcher.stats().prefill_tokens;
let prefill_job = {
let batcher = batcher.clone();
thread::spawn(move || {
batcher.generate(long_prompt, greedy_params(1, 9), StopTokens::default())
})
};
let decode_before = loop {
assert!(
std::time::Instant::now() < deadline,
"never observed the long prompt mid-prefill"
);
let st = batcher.stats();
let progressed = st.prefill_tokens - prefill_at_submit;
assert!(
progressed < total,
"the whole prompt was prefilled without ever being observed \
partially done: prefill ran as one unbounded unit of work"
);
if progressed > 0 {
break st.decode_steps;
}
thread::yield_now();
};
loop {
assert!(
std::time::Instant::now() < deadline,
"decode stalled while a long prompt prefilled"
);
let st = batcher.stats();
if st.decode_steps > decode_before {
break;
}
assert!(
st.prefill_tokens - prefill_at_submit < total,
"the prompt finished prefilling before the in-flight decode \
took a single step: prefill froze decode"
);
thread::yield_now();
}
let (_finish, ids, _text, _usage) = prefill_job.join().unwrap().expect("prefill job");
assert_eq!(ids.len(), 1);
let (_finish, ids, _text, _usage) = decode_job.join().unwrap().expect("decode job");
assert_eq!(ids.len(), 90);
}
#[test]
fn max_seqs_cap_counts_prefilling_prompts_and_still_serves_both() {
let decoder = tiny_decoder();
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
identity_decode(),
BatcherConfig {
max_seqs: 1,
prefill_chunk: 1,
..BatcherConfig::default()
},
);
let expected: Vec<Vec<usize>> = [(vec![1usize, 2, 3], 6u64), (vec![4usize, 5], 6)]
.iter()
.map(|(p, seed)| sequential_ids(&decoder, p, &greedy_params(6, *seed)))
.collect();
let handles: Vec<_> = [(vec![1usize, 2, 3], 6u64), (vec![4usize, 5], 6)]
.into_iter()
.map(|(prompt, seed)| {
let batcher = batcher.clone();
thread::spawn(move || {
batcher
.generate(prompt, greedy_params(6, seed), StopTokens::default())
.expect("generate")
.1
})
})
.collect();
let got: Vec<Vec<usize>> = handles.into_iter().map(|h| h.join().unwrap()).collect();
assert_eq!(got[0], expected[0]);
assert_eq!(got[1], expected[1]);
}
fn test_shape() -> ferrox_models::KvShape {
ferrox_models::KvShape::from_config(&test_dense_fixture(), ferrox_models::KvElem::F32, 1)
}
fn no_budget() -> BlockBudget {
BlockBudget::new(
DEFAULT_KV_BLOCK_SIZE,
None,
Arc::new(ContextCeiling::new(None, test_shape())),
)
}
fn budget(block_size: usize, total: Option<usize>) -> BlockBudget {
BlockBudget::new(
block_size,
total,
Arc::new(ContextCeiling::new(None, test_shape())),
)
}
fn test_slot(max_tokens: usize, seed: u64) -> (Slot, mpsc::Receiver<JobResult>) {
let (tx, rx) = mpsc::channel();
let params = greedy_params(max_tokens, seed);
(
Slot {
caches: Vec::new(),
pos: 0,
logits: Vec::new(),
sampler: Sampler::new(seed),
generated_ids: Vec::new(),
visible: String::new(),
stops: StopMatcher::new(¶ms.stop, ¶ms.stop_token_ids),
prompt_tokens: 0,
max_tokens,
stop_tokens: StopTokens::default(),
params,
reply: tx,
abort: AbortId(0),
blocks: 1,
finish: None,
},
rx,
)
}
#[test]
fn removing_a_row_never_reassigns_another_rows_state() {
let mut rows = Rows::default();
let (a, _ra) = test_slot(3, 11);
let (b, _rb) = test_slot(5, 22);
let (c, _rc) = test_slot(7, 33);
let a = rows.insert(a);
let b = rows.insert(b);
let c = rows.insert(c);
assert_eq!(rows.order, vec![a, b, c]);
let removed = rows.remove(b).expect("b was present");
assert_eq!(removed.max_tokens, 5);
assert!(
rows.get(b).is_none(),
"a stale uid must resolve to nothing, never to another request's row"
);
assert_eq!(rows.get(a).expect("a still in flight").max_tokens, 3);
assert_eq!(
rows.get(c).expect("c still in flight").max_tokens,
7,
"c must still be c after b left"
);
assert_eq!(rows.order, vec![a, c], "admission order is preserved");
assert_eq!(rows.len(), 2);
let mut positional = vec![3usize, 5, 7];
let c_index = 2;
positional.swap_remove(1);
assert_eq!(positional[1], 7, "C moved into B's index");
assert!(
positional.get(c_index).is_none(),
"C's index now names nothing"
);
}
#[test]
fn uids_are_unique_and_insertion_does_not_disturb_existing_rows() {
let mut rows = Rows::default();
let (a, _ra) = test_slot(3, 11);
let a = rows.insert(a);
let (b, _rb) = test_slot(5, 22);
let b = rows.insert(b);
rows.remove(a);
let (c, _rc) = test_slot(7, 33);
let c = rows.insert(c);
assert_ne!(c, a, "a uid is never reused after its row leaves");
assert_ne!(c, b);
assert_eq!(rows.get(b).expect("b untouched").max_tokens, 5);
assert_eq!(rows.get(c).expect("c inserted").max_tokens, 7);
}
#[test]
fn flush_replies_on_each_rows_own_channel() {
let mut rows = Rows::default();
let (a, ra) = test_slot(3, 11);
let (mut b, rb) = test_slot(5, 22);
b.finish = Some(FinishReason::Stop);
b.visible.push_str("bee");
b.generated_ids.push(7);
let a = rows.insert(a);
let b = rows.insert(b);
let (c, _rc) = test_slot(7, 33);
let c = rows.insert(c);
assert_eq!(rows.ready(), vec![a, c], "a finished row takes no step");
rows.flush_finished(&no_budget());
assert!(rows.get(b).is_none());
assert_eq!(rows.order, vec![a, c]);
let (finish, ids, text, usage) =
rb.try_recv().expect("b's caller got a reply").expect("ok");
assert_eq!(finish, FinishReason::Stop);
assert_eq!(ids, vec![7]);
assert_eq!(text, "bee");
assert_eq!(usage.completion_tokens, 1);
assert!(
ra.try_recv().is_err(),
"an unfinished row's caller must not be replied to"
);
}
#[test]
fn a_row_leaving_mid_batch_does_not_shift_its_neighbours_output() {
let decoder = tiny_decoder();
let prompts = [vec![1usize, 2, 3], vec![4usize, 5], vec![6usize]];
let budgets = [25usize, 25, 20];
let refs: Vec<Vec<usize>> = prompts
.iter()
.zip(budgets.iter())
.map(|(p, &n)| sequential_ids(&decoder, p, &greedy_params(n, 4)))
.collect();
let letter = |id: &usize| char::from_u32(65 + (*id as u32 % 26)).unwrap_or('?');
let middle_text: String = refs[1].iter().map(letter).collect();
assert!(middle_text.len() >= 4);
let stop = middle_text[2..4].to_string();
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
identity_decode(),
BatcherConfig {
prefill_chunk: 1,
..BatcherConfig::default()
},
);
let barrier = Arc::new(Barrier::new(prompts.len()));
let handles: Vec<_> = (0..prompts.len())
.map(|i| {
let batcher = batcher.clone();
let barrier = Arc::clone(&barrier);
let prompt = prompts[i].clone();
let mut params = greedy_params(budgets[i], 4);
if i == 1 {
params.stop = vec![stop.clone()];
}
thread::spawn(move || {
barrier.wait();
batcher
.generate(prompt, params, StopTokens::default())
.expect("generate")
.1
})
})
.collect();
let got: Vec<Vec<usize>> = handles.into_iter().map(|h| h.join().unwrap()).collect();
assert_eq!(got[0], refs[0], "row 0 received another row's output");
assert_eq!(got[2], refs[2], "row 2 received another row's output");
assert!(
got[1].len() < refs[1].len() && refs[1].starts_with(&got[1]),
"the stopped row must be a strict prefix of its own stream"
);
}
#[test]
fn queue_gate_admits_up_to_its_cap_and_frees_slots_on_release() {
let gate = QueueGate::new(2);
assert!(gate.try_reserve().is_ok());
assert!(gate.try_reserve().is_ok());
assert_eq!(gate.depth(), 2);
assert_eq!(gate.try_reserve(), Err(2), "the refusal reports the depth");
assert_eq!(gate.rejected(), 1);
gate.release();
assert_eq!(gate.depth(), 1);
assert!(
gate.try_reserve().is_ok(),
"a released slot must be reusable"
);
assert_eq!(gate.depth(), 2);
}
#[test]
fn queue_gate_never_exceeds_its_cap_under_concurrent_submitters() {
const THREADS: usize = 32;
const CAP: usize = 4;
for round in 0..64 {
let gate = Arc::new(QueueGate::new(CAP));
let barrier = Arc::new(Barrier::new(THREADS));
let admitted = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let handles: Vec<_> = (0..THREADS)
.map(|_| {
let gate = Arc::clone(&gate);
let barrier = Arc::clone(&barrier);
let admitted = Arc::clone(&admitted);
thread::spawn(move || {
barrier.wait();
if gate.try_reserve().is_ok() {
admitted.fetch_add(1, Ordering::Relaxed);
}
})
})
.collect();
for h in handles {
h.join().unwrap();
}
assert_eq!(
admitted.load(Ordering::Relaxed),
CAP,
"round {round}: exactly the cap may be admitted"
);
assert_eq!(gate.depth(), CAP, "round {round}: depth matches admissions");
assert_eq!(gate.rejected(), (THREADS - CAP) as u64);
}
}
#[test]
fn a_full_queue_refuses_new_jobs_with_queue_full() {
let decoder = tiny_decoder();
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
identity_decode(),
BatcherConfig {
max_queue: 0,
..BatcherConfig::default()
},
);
let err = batcher
.generate(vec![1, 2, 3], greedy_params(4, 1), StopTokens::default())
.expect_err("a full queue must refuse");
assert!(
matches!(err, DecodeError::QueueFull { queued: 0, cap: 0 }),
"expected QueueFull, got {err:?}"
);
assert_eq!(err.retry_after_secs(), Some(1), "a queue drains; say so");
let stats = batcher.stats();
assert_eq!(stats.queue_rejected, 1);
assert_eq!(stats.queue_depth, 0, "a refused job holds nothing");
}
#[test]
fn queue_depth_returns_to_zero_after_a_served_request() {
let decoder = tiny_decoder();
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
identity_decode(),
BatcherConfig {
max_queue: 1,
prefill_chunk: 1,
..BatcherConfig::default()
},
);
for _ in 0..3 {
batcher
.generate(vec![1, 2, 3], greedy_params(2, 1), StopTokens::default())
.expect("a cap of 1 still serves requests one after another");
}
assert_eq!(batcher.stats().queue_depth, 0);
assert_eq!(batcher.stats().queue_rejected, 0);
}
fn budget_config(block_size: usize, blocks: usize) -> BatcherConfig {
BatcherConfig {
prefill_chunk: 1,
kv_block_size: block_size,
kv_blocks: Some(blocks),
..BatcherConfig::default()
}
}
#[test]
fn blocks_are_counted_in_positions_and_always_round_up() {
let budget = budget(4, Some(10));
assert_eq!(budget.blocks_for(0), 1);
assert_eq!(budget.blocks_for(1), 1);
assert_eq!(budget.blocks_for(4), 1);
assert_eq!(budget.blocks_for(5), 2);
assert_eq!(budget.blocks_for(8), 2);
assert_eq!(budget.blocks_for(9), 3);
}
#[test]
fn an_unconfigured_budget_admits_everything() {
let budget = budget(4, None);
assert!(budget.immovable_refusal(usize::MAX).is_none());
assert!(budget.try_reserve(1_000_000));
budget.release(1_000_000);
}
#[test]
fn a_request_larger_than_the_whole_budget_is_refused_rather_than_queued() {
let decoder = tiny_decoder();
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
identity_decode(),
budget_config(4, 2),
);
let err = batcher
.generate(
vec![1, 2, 3, 4, 5, 6],
greedy_params(8, 1),
StopTokens::from_eos(None),
)
.expect_err("14 positions cannot fit an 8-position server");
let shape = test_shape();
match &err {
DecodeError::KvBudgetExceeded {
binding,
estimated_bytes,
limit_bytes,
positions,
positions_limit,
detail,
} => {
assert_eq!(*binding, "device_memory_budget_exceeded");
assert_eq!(*positions, 14);
assert_eq!(*positions_limit, 8, "2 blocks x 4 positions");
assert_eq!(*estimated_bytes, shape.kv_bytes_for_tokens(14));
assert_eq!(*limit_bytes, shape.kv_bytes_for_tokens(8));
assert!(estimated_bytes > limit_bytes);
assert!(detail.contains("14"), "{detail}");
}
other => panic!("expected KvBudgetExceeded, got {other:?}"),
}
assert_eq!(err.retry_after_secs(), None);
let stats = batcher.stats();
assert_eq!(stats.kv_rejected_too_large, 1);
assert_eq!(
stats.kv_rejected_context_length, 0,
"no per-request context ceiling is configured here"
);
assert_eq!(
stats.queue_rejected, 0,
"too-big and under-pressure are different counters"
);
assert_eq!(
stats.queue_depth, 0,
"an impossible request must not occupy a queue slot"
);
assert_eq!(stats.kv_blocks_free, 2, "nothing was reserved");
batcher
.generate(vec![1, 2], greedy_params(2, 1), StopTokens::from_eos(None))
.expect("4 positions fit");
}
#[test]
fn a_request_longer_than_the_context_ceiling_names_that_ceiling() {
let decoder = tiny_decoder();
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
identity_decode(),
BatcherConfig {
prefill_chunk: 1,
max_context: Some(6),
kv_block_size: 4,
kv_blocks: Some(1024),
..BatcherConfig::default()
},
);
let err = batcher
.generate(
vec![1, 2, 3, 4],
greedy_params(4, 1),
StopTokens::from_eos(None),
)
.expect_err("8 positions against a 6-position ceiling");
let shape = test_shape();
match &err {
DecodeError::KvBudgetExceeded {
binding,
estimated_bytes,
limit_bytes,
positions,
positions_limit,
detail,
} => {
assert_eq!(*binding, "context_length_exceeded");
assert_eq!(*positions, 8);
assert_eq!(*positions_limit, 6);
assert_eq!(*estimated_bytes, shape.kv_bytes_for_tokens(8));
assert_eq!(*limit_bytes, shape.kv_bytes_for_tokens(6));
assert!(detail.contains("max_tokens"), "{detail}");
}
other => panic!("expected KvBudgetExceeded, got {other:?}"),
}
assert_eq!(err.retry_after_secs(), None);
let stats = batcher.stats();
assert_eq!(stats.kv_rejected_context_length, 1);
assert_eq!(
stats.kv_rejected_too_large, 0,
"the machine's budget was never the binding ceiling"
);
assert_eq!(stats.queue_rejected, 0);
batcher
.generate(
vec![1, 2, 3, 4],
greedy_params(2, 1),
StopTokens::from_eos(None),
)
.expect("6 positions is 6 positions");
}
#[test]
fn the_context_ceiling_is_reported_before_the_device_ceiling() {
let decoder = tiny_decoder();
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
identity_decode(),
BatcherConfig {
prefill_chunk: 1,
max_context: Some(6),
kv_block_size: 4,
kv_blocks: Some(2),
..BatcherConfig::default()
},
);
let err = batcher
.generate(
vec![1, 2, 3, 4, 5, 6],
greedy_params(8, 1),
StopTokens::from_eos(None),
)
.expect_err("14 positions breaks both ceilings");
assert!(
matches!(
&err,
DecodeError::KvBudgetExceeded { binding, .. }
if *binding == "context_length_exceeded"
),
"got {err:?}"
);
let stats = batcher.stats();
assert_eq!(stats.kv_rejected_context_length, 1);
assert_eq!(stats.kv_rejected_too_large, 0);
}
#[test]
fn without_ceilings_nothing_is_refused_as_too_large() {
let budget = BlockBudget::new(4, None, Arc::new(ContextCeiling::new(None, test_shape())));
assert!(budget.immovable_refusal(1_000_000).is_none());
assert_eq!(budget.rejected_too_large.load(Ordering::Relaxed), 0);
assert_eq!(budget.ceiling.refused(), 0);
}
#[test]
fn concurrent_requests_never_hold_more_blocks_than_the_budget() {
let decoder = tiny_decoder();
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
identity_decode(),
budget_config(4, 4),
);
let start = Arc::new(Barrier::new(6));
let handles: Vec<_> = (0..6)
.map(|i| {
let batcher = batcher.clone();
let start = Arc::clone(&start);
thread::spawn(move || {
start.wait();
batcher
.generate(
vec![1, 2, 3],
greedy_params(5, i as u64),
StopTokens::from_eos(None),
)
.expect("every request fits the budget on its own")
})
})
.collect();
for handle in handles {
handle.join().expect("no submitter panicked");
}
let stats = batcher.stats();
assert!(
stats.kv_blocks_peak <= stats.kv_blocks_total,
"admission handed out {} blocks from a budget of {}",
stats.kv_blocks_peak,
stats.kv_blocks_total
);
assert_eq!(stats.kv_rejected_too_large, 0, "all six fit individually");
}
#[test]
fn every_admitted_request_gives_its_blocks_back() {
let decoder = tiny_decoder();
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
identity_decode(),
budget_config(4, 4),
);
for i in 0..8 {
batcher
.generate(
vec![1, 2, 3],
greedy_params(4, i),
StopTokens::from_eos(None),
)
.expect("generate");
}
for _ in 0..200 {
if batcher.stats().kv_blocks_free == 4 {
break;
}
thread::sleep(std::time::Duration::from_millis(5));
}
let stats = batcher.stats();
assert_eq!(
stats.kv_blocks_free, stats.kv_blocks_total,
"an idle server must own its whole budget again"
);
assert!(stats.kv_blocks_peak > 0, "something was actually reserved");
}
#[test]
fn a_job_rejected_at_validation_gives_its_blocks_back() {
let decoder = tiny_decoder();
let vocab = decoder.config.vocab_size;
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
identity_decode(),
budget_config(4, 4),
);
assert!(matches!(
batcher.generate(
vec![vocab + 1],
greedy_params(2, 1),
StopTokens::from_eos(None)
),
Err(DecodeError::TokenOutOfVocab { .. })
));
for _ in 0..200 {
if batcher.stats().kv_blocks_free == 4 {
break;
}
thread::sleep(std::time::Duration::from_millis(5));
}
assert_eq!(batcher.stats().kv_blocks_free, 4);
batcher
.generate(vec![1, 2], greedy_params(2, 1), StopTokens::from_eos(None))
.expect("a bad request must not poison the budget");
}
#[test]
fn a_head_job_that_does_not_fit_holds_the_line() {
let decoder = tiny_decoder();
let config = budget_config(4, 4);
let budget = BlockBudget::new(
config.kv_block_size,
config.kv_blocks,
Arc::new(ContextCeiling::new(config.max_context, test_shape())),
);
let queue = QueueGate::new(config.max_queue);
assert!(budget.try_reserve(2));
let mut waiting: VecDeque<Job> = VecDeque::new();
let (big_tx, _big_rx) = mpsc::channel();
let (small_tx, _small_rx) = mpsc::channel();
waiting.push_back(Job {
prompt_tokens: vec![1, 2, 3],
params: greedy_params(2, 1),
stop_tokens: StopTokens::from_eos(None),
reply: big_tx,
abort: AbortId(0),
blocks: 3,
});
waiting.push_back(Job {
prompt_tokens: vec![1],
params: greedy_params(2, 2),
stop_tokens: StopTokens::from_eos(None),
reply: small_tx,
abort: AbortId(1),
blocks: 1,
});
queue.try_reserve().expect("cap 512");
queue.try_reserve().expect("cap 512");
let mut prefills: VecDeque<Prefill> = VecDeque::new();
admit(
&decoder,
&mut waiting,
&mut prefills,
0,
&config,
&queue,
&budget,
);
assert!(
prefills.is_empty(),
"the 1-block job must not jump the 3-block job that cannot fit"
);
assert_eq!(waiting.len(), 2);
assert_eq!(queue.depth(), 2, "neither job has stopped waiting");
budget.release(2);
admit(
&decoder,
&mut waiting,
&mut prefills,
0,
&config,
&queue,
&budget,
);
assert_eq!(prefills.len(), 2);
assert!(waiting.is_empty());
assert_eq!(queue.depth(), 0);
assert_eq!(budget.free(), 0, "3 + 1 blocks are now out");
}
#[test]
fn the_sequence_cap_and_the_block_cap_compose() {
let decoder = tiny_decoder();
let config = BatcherConfig {
max_seqs: 1,
..budget_config(4, 8)
};
let budget = BlockBudget::new(
config.kv_block_size,
config.kv_blocks,
Arc::new(ContextCeiling::new(config.max_context, test_shape())),
);
let queue = QueueGate::new(config.max_queue);
let mut waiting: VecDeque<Job> = VecDeque::new();
let mut receivers = Vec::new();
for i in 0..3 {
let (tx, rx) = mpsc::channel();
receivers.push(rx);
waiting.push_back(Job {
prompt_tokens: vec![1, 2],
params: greedy_params(2, i),
stop_tokens: StopTokens::from_eos(None),
reply: tx,
abort: AbortId(i),
blocks: 1,
});
queue.try_reserve().expect("cap 512");
}
let mut prefills: VecDeque<Prefill> = VecDeque::new();
admit(
&decoder,
&mut waiting,
&mut prefills,
0,
&config,
&queue,
&budget,
);
assert_eq!(prefills.len(), 1, "max_seqs still binds");
assert_eq!(budget.free(), 7, "only the admitted job reserved");
assert_eq!(receivers.len(), 3);
}
fn cancellable_params(max_tokens: usize, seed: u64) -> (GenerationParams, CancelToken) {
let token = CancelToken::new();
let mut params = greedy_params(max_tokens, seed);
params.cancel = Some(token.clone());
(params, token)
}
fn abortable_job(abort: AbortId, prompt: Vec<usize>) -> (Job, mpsc::Receiver<JobResult>) {
let (tx, rx) = mpsc::channel();
(
Job {
prompt_tokens: prompt,
params: greedy_params(4, 1),
stop_tokens: StopTokens::from_eos(None),
reply: tx,
abort,
blocks: 1,
},
rx,
)
}
#[test]
fn a_decoding_request_stops_when_it_is_cancelled() {
let decoder = tiny_decoder();
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
identity_decode(),
BatcherConfig {
prefill_chunk: 1,
..BatcherConfig::default()
},
);
let (params, token) = cancellable_params(4000, 9);
let worker = {
let batcher = batcher.clone();
thread::spawn(move || {
batcher.generate(vec![1, 2, 3], params, StopTokens::from_eos(None))
})
};
let deadline = std::time::Instant::now() + std::time::Duration::from_secs(10);
while batcher.stats().decode_steps < 3 {
assert!(
std::time::Instant::now() < deadline,
"never started decoding"
);
thread::sleep(std::time::Duration::from_millis(1));
}
token.cancel();
let (finish, ids, _text, usage) = worker.join().expect("no panic").expect("generate");
assert_eq!(finish, FinishReason::Cancelled);
assert!(
ids.len() < 4000,
"cancelling did not shorten the decode: {} tokens",
ids.len()
);
assert!(
!ids.is_empty(),
"the tokens produced before the cancel must survive it"
);
assert_eq!(usage.completion_tokens, ids.len());
assert_eq!(batcher.stats().aborted, 1);
}
#[test]
fn a_cancelled_row_leaves_through_the_one_exit_every_row_uses() {
let mut rows = Rows::default();
let budget = budget(4, Some(4));
assert!(budget.try_reserve(2));
let (mut a, ra) = test_slot(9, 11);
a.abort = AbortId(7);
a.blocks = 2;
a.generated_ids.push(3);
a.visible.push_str("hi");
let (b, rb) = test_slot(9, 22);
let a_uid = rows.insert(a);
let b_uid = rows.insert(b);
let consumed = rows.mark_cancelled(&HashSet::from([AbortId(7)]));
assert_eq!(consumed, vec![AbortId(7)]);
assert!(
rows.get(a_uid).is_some(),
"the row must still be in the table until the flush: its KV \
buffers are what the batch is built from"
);
assert_eq!(
budget.free(),
2,
"marking must not release blocks -- the flush does"
);
assert!(ra.try_recv().is_err(), "no reply until the row leaves");
rows.flush_finished(&budget);
assert!(rows.get(a_uid).is_none());
assert!(rows.get(b_uid).is_some(), "only the cancelled row left");
assert_eq!(budget.free(), 4, "blocks came back exactly once");
let (finish, ids, text, _usage) = ra.try_recv().expect("one reply").expect("ok");
assert_eq!(finish, FinishReason::Cancelled);
assert_eq!(ids, vec![3], "partial output survives the cancel");
assert_eq!(text, "hi");
assert!(
ra.try_recv().is_err(),
"a cancelled row must be replied to exactly once"
);
assert!(rb.try_recv().is_err(), "the other row is untouched");
}
#[test]
fn a_cancel_that_arrives_before_its_job_is_not_lost() {
let config = budget_config(4, 4);
let inbox = AbortInbox::default();
let queue = QueueGate::new(config.max_queue);
let budget = BlockBudget::new(
config.kv_block_size,
config.kv_blocks,
Arc::new(ContextCeiling::new(config.max_context, test_shape())),
);
let mut carried = HashSet::new();
let mut waiting = VecDeque::new();
let mut prefills = VecDeque::new();
let mut rows = Rows::default();
inbox.enqueue(AbortId(42));
apply_aborts(
&inbox,
&mut carried,
&mut waiting,
&mut prefills,
&mut rows,
&queue,
&budget,
);
assert_eq!(inbox.aborted(), 0, "nothing to stop yet");
let (job, rx) = abortable_job(AbortId(42), vec![1, 2]);
waiting.push_back(job);
queue.try_reserve().expect("cap");
apply_aborts(
&inbox,
&mut carried,
&mut waiting,
&mut prefills,
&mut rows,
&queue,
&budget,
);
assert!(waiting.is_empty(), "the late job must still be cancelled");
assert_eq!(inbox.aborted(), 1);
assert_eq!(queue.depth(), 0, "its queue slot came back");
let (finish, ids, _, usage) = rx.try_recv().expect("reply").expect("ok");
assert_eq!(finish, FinishReason::Cancelled);
assert!(ids.is_empty(), "it never ran a token");
assert_eq!(usage.prompt_tokens, 2);
}
#[test]
fn a_cancelled_prefill_is_abandoned_and_gives_its_blocks_back() {
let decoder = tiny_decoder();
let config = budget_config(4, 4);
let inbox = AbortInbox::default();
let queue = QueueGate::new(config.max_queue);
let budget = BlockBudget::new(
config.kv_block_size,
config.kv_blocks,
Arc::new(ContextCeiling::new(config.max_context, test_shape())),
);
let mut carried = HashSet::new();
let mut waiting = VecDeque::new();
let mut prefills = VecDeque::new();
let mut rows = Rows::default();
let (job, rx) = abortable_job(AbortId(5), vec![1, 2, 3, 4]);
waiting.push_back(job);
queue.try_reserve().expect("cap");
admit(
&decoder,
&mut waiting,
&mut prefills,
0,
&config,
&queue,
&budget,
);
assert_eq!(prefills.len(), 1);
assert_eq!(budget.free(), 3, "the prefill holds its reservation");
inbox.enqueue(AbortId(5));
apply_aborts(
&inbox,
&mut carried,
&mut waiting,
&mut prefills,
&mut rows,
&queue,
&budget,
);
assert!(prefills.is_empty(), "the remaining chunks never run");
assert_eq!(budget.free(), 4, "an abandoned prefill releases its blocks");
assert_eq!(inbox.aborted(), 1);
assert_eq!(
rx.try_recv().expect("reply").expect("ok").0,
FinishReason::Cancelled
);
}
#[test]
fn cancelling_one_request_leaves_its_neighbours_running() {
let decoder = tiny_decoder();
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
identity_decode(),
BatcherConfig {
prefill_chunk: 1,
..BatcherConfig::default()
},
);
let expected = sequential_ids(&decoder, &[4, 5], &greedy_params(6, 3));
let (doomed_params, token) = cancellable_params(4000, 9);
let doomed = {
let batcher = batcher.clone();
thread::spawn(move || {
batcher.generate(vec![1, 2, 3], doomed_params, StopTokens::from_eos(None))
})
};
let survivor = {
let batcher = batcher.clone();
thread::spawn(move || {
batcher.generate(vec![4, 5], greedy_params(6, 3), StopTokens::from_eos(None))
})
};
let deadline = std::time::Instant::now() + std::time::Duration::from_secs(10);
while batcher.stats().decode_steps < 3 {
assert!(
std::time::Instant::now() < deadline,
"never started decoding"
);
thread::sleep(std::time::Duration::from_millis(1));
}
token.cancel();
let (finish, _, _, _) = doomed.join().expect("no panic").expect("generate");
assert_eq!(finish, FinishReason::Cancelled);
let (finish, ids, _, _) = survivor.join().expect("no panic").expect("generate");
assert_eq!(finish, FinishReason::Length);
assert_eq!(
ids, expected,
"an uncancelled request must produce exactly what it would have alone"
);
}
#[test]
fn a_token_level_stop_ends_a_batched_row() {
let decoder = tiny_decoder();
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
identity_decode(),
BatcherConfig {
prefill_chunk: 1,
..BatcherConfig::default()
},
);
let baseline = sequential_ids(&decoder, &[1, 2, 3], &greedy_params(8, 4));
assert!(baseline.len() > 1, "need something to stop before the end");
let stop_token = baseline[1];
let (finish, ids, _text, usage) = batcher
.generate(
vec![1, 2, 3],
GenerationParams {
stop_token_ids: vec![stop_token],
..greedy_params(8, 4)
},
StopTokens::from_eos(None),
)
.expect("generate");
assert_eq!(finish, FinishReason::Stop);
assert_eq!(
ids,
baseline[..1].to_vec(),
"the stop token itself is not part of the answer"
);
assert_eq!(usage.completion_tokens, 1);
}
#[test]
fn a_multi_token_stop_string_truncates_a_batched_row() {
let decoder = tiny_decoder();
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
identity_decode(),
BatcherConfig {
prefill_chunk: 1,
..BatcherConfig::default()
},
);
let baseline = sequential_ids(&decoder, &[1, 2, 3], &greedy_params(8, 4));
let decode = identity_decode();
let full = decode(&baseline);
if full.chars().count() < 4 {
return;
}
let cut: Vec<char> = full.chars().collect();
let stop_str: String = cut[1..3].iter().collect();
let expected: String = cut[..1].iter().collect();
let (finish, _ids, text, _usage) = batcher
.generate(
vec![1, 2, 3],
GenerationParams {
stop: vec![stop_str.clone()],
..greedy_params(8, 4)
},
StopTokens::from_eos(None),
)
.expect("generate");
assert_eq!(finish, FinishReason::Stop);
assert_eq!(
text, expected,
"a stop spanning two tokens must cut where it starts"
);
assert!(!text.contains(&stop_str));
}
#[test]
fn a_batched_row_that_never_matches_loses_no_output() {
let decoder = tiny_decoder();
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
identity_decode(),
BatcherConfig {
prefill_chunk: 1,
..BatcherConfig::default()
},
);
let baseline = sequential_ids(&decoder, &[1, 2, 3], &greedy_params(8, 4));
let expected = identity_decode()(&baseline);
let (finish, ids, text, _usage) = batcher
.generate(
vec![1, 2, 3],
GenerationParams {
stop: vec!["ZZ_NEVER_MATCHES_ZZ".to_string()],
..greedy_params(8, 4)
},
StopTokens::from_eos(None),
)
.expect("generate");
assert_eq!(finish, FinishReason::Length);
assert_eq!(ids, baseline);
assert_eq!(
text, expected,
"buffering is about when text is released, never whether"
);
}
#[test]
fn an_uncancelled_request_is_unaffected_by_the_abort_path() {
let decoder = tiny_decoder();
let batcher = ContinuousBatcher::spawn_with_config(
Arc::clone(&decoder),
identity_decode(),
BatcherConfig {
prefill_chunk: 1,
..BatcherConfig::default()
},
);
let (params, _token) = cancellable_params(6, 3);
let (finish, ids, _, _) = batcher
.generate(vec![4, 5], params, StopTokens::from_eos(None))
.expect("generate");
assert_eq!(finish, FinishReason::Length);
assert_eq!(ids, sequential_ids(&decoder, &[4, 5], &greedy_params(6, 3)));
assert_eq!(batcher.stats().aborted, 0);
}
}