use super::*;
#[test]
fn continuous_batching_composes_with_paged_kv() {
let decoder = tiny_decoder();
let prompts: [Vec<usize>; 2] = [vec![1, 2, 3], vec![4, 5]];
let params = [greedy_params(6, 7), greedy_params(4, 11)];
let sequential: Vec<Vec<usize>> = prompts
.iter()
.zip(params.iter())
.map(|(p, par)| sequential_ids(&decoder, p, par))
.collect();
let paged = PagedKvConfig {
store: Arc::new(ferrox_core::cache::SharedPagedKv::new(
decoder.layers.len(),
4,
256,
decoder.config.n_kv_heads,
decoder.config.head_dim,
)),
queue_wait: std::time::Duration::ZERO,
radix: None,
anchor_token: None,
slide_interval: crate::policy::pool_budget::DEFAULT_SWA_EVICTION_INTERVAL,
};
let store = Arc::clone(&paged.store);
let free_before = store.free_groups();
let batcher = ContinuousBatcher::spawn_with_config_paged(
Arc::clone(&decoder),
identity_decode(),
BatcherConfig {
prefill_chunk: 1,
..BatcherConfig::default()
},
Some(paged),
);
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(),
ignore_eos: false,
};
threads.push(thread::spawn(move || {
barrier.wait();
let out = batcher
.generate(prompt, par, StopTokens::default())
.expect("a paged batched row must serve");
results.lock().unwrap()[i] = Some(out.1);
}));
}
barrier.wait();
for t in threads {
t.join().unwrap();
}
let got = results.lock().unwrap().clone();
for (i, want) in sequential.iter().enumerate() {
assert_eq!(
got[i].as_ref().expect("both rows replied"),
want,
"row {i}: batching over paged KV changed the ids"
);
}
drop(batcher);
std::thread::sleep(std::time::Duration::from_millis(200));
assert_eq!(
store.free_groups(),
free_before,
"a finished batched row must return its pages"
);
}
#[test]
fn a_window_model_slides_while_continuously_batched() {
let window = 8;
let block_size = 4;
let mut cfg = test_dense_fixture();
cfg.sliding_window = Some(window);
cfg.swa_pattern = None;
let vocab = cfg.vocab_size;
let decoder = Arc::new(Decoder::new_random_small(cfg, 2, vocab));
assert_eq!(decoder.config.uniform_sliding_window(), Some(window));
let max_tokens = 400;
let prompts: [Vec<usize>; 2] = [vec![1, 2, 3], vec![4, 5, 6]];
let long = |seed: u64| GenerationParams {
ignore_eos: true,
..greedy_params(max_tokens, seed)
};
let params = [long(7), long(11)];
let sequential: Vec<Vec<usize>> = prompts
.iter()
.zip(params.iter())
.map(|(p, par)| sequential_ids(&decoder, p, par))
.collect();
let paged = PagedKvConfig {
store: Arc::new(ferrox_core::cache::SharedPagedKv::new(
decoder.layers.len(),
block_size,
120,
decoder.config.n_kv_heads,
decoder.config.head_dim,
)),
queue_wait: std::time::Duration::from_secs(5),
radix: None,
anchor_token: None,
slide_interval: crate::policy::pool_budget::DEFAULT_SWA_EVICTION_INTERVAL,
};
let store = Arc::clone(&paged.store);
let free_before = store.free_groups();
let batcher = ContinuousBatcher::spawn_with_config_paged(
Arc::clone(&decoder),
identity_decode(),
BatcherConfig {
prefill_chunk: 8,
kv_block_size: block_size,
kv_blocks: Some(100),
..BatcherConfig::default()
},
Some(paged),
);
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 = params[i].clone();
threads.push(thread::spawn(move || {
barrier.wait();
let out = batcher
.generate(prompt, par, StopTokens::default())
.expect("a windowed batched row must serve");
results.lock().unwrap()[i] = Some(out.1);
}));
}
barrier.wait();
for t in threads {
t.join().unwrap();
}
let got = results.lock().unwrap().clone();
for (i, want) in sequential.iter().enumerate() {
let ids = got[i].as_ref().expect("both rows replied");
assert_eq!(ids.len(), max_tokens, "row {i} stopped early");
assert_eq!(ids, want, "row {i}: the window slide changed the ids");
}
drop(batcher);
std::thread::sleep(std::time::Duration::from_millis(200));
assert_eq!(
store.free_groups(),
free_before,
"a finished slid row must return its recycled pages too"
);
}
#[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(),
ignore_eos: false,
};
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::StopSequence(stop.clone()),
"a batched row must name the stop it hit, like an unbatched one"
);
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, _ids) = 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, _ids) = 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]);
}
#[test]
fn a_finished_batched_row_publishes_its_prefix() {
let decoder = tiny_decoder();
let block_size = 4;
let radix = Arc::new(std::sync::Mutex::new(
crate::policy::radix::RadixCache::new(block_size),
));
let paged = PagedKvConfig {
store: Arc::new(ferrox_core::cache::SharedPagedKv::new(
decoder.layers.len(),
block_size,
256,
decoder.config.n_kv_heads,
decoder.config.head_dim,
)),
queue_wait: std::time::Duration::ZERO,
radix: Some(Arc::clone(&radix)),
anchor_token: None,
slide_interval: crate::policy::pool_budget::DEFAULT_SWA_EVICTION_INTERVAL,
};
assert_eq!(
radix.lock().unwrap().total_size(),
0,
"the tree starts empty"
);
let batcher = ContinuousBatcher::spawn_with_config_paged(
Arc::clone(&decoder),
identity_decode(),
BatcherConfig {
prefill_chunk: 1,
..BatcherConfig::default()
},
Some(paged),
);
let prompt: Vec<usize> = vec![1, 2, 3, 4, 5, 6, 7, 8];
batcher
.generate(prompt, greedy_params(4, 7), StopTokens::default())
.expect("a paged batched row must serve");
drop(batcher);
std::thread::sleep(std::time::Duration::from_millis(200));
assert!(
radix.lock().unwrap().total_size() > 0,
"a finished paged row must leave its prefix in the tree for the next request"
);
}