use std::collections::{HashMap, VecDeque};
use crate::batch::config::BatchConfig;
use crate::batch::scheduler::{FifoScheduler, Scheduler};
use crate::batch::sequence::{
AdapterKey, FinishReason, SeqId, Sequence, SequenceManager, SequenceState,
};
use crate::kv_cache::{PagePool, PagedKVCacheConfig};
use crate::sampling::{Sampler, SamplingConfig};
#[derive(Debug)]
pub struct GdnStatePool {
s_matrices: Vec<Vec<f32>>,
conv_buffers: Vec<Vec<f32>>,
free_slots: VecDeque<usize>,
seq_to_slot: HashMap<SeqId, usize>,
capacity: usize,
}
impl GdnStatePool {
pub fn new(capacity: usize, s_floats_per_slot: usize, conv_floats_per_slot: usize) -> Self {
let s_matrices = (0..capacity)
.map(|_| vec![0.0f32; s_floats_per_slot])
.collect();
let conv_buffers = (0..capacity)
.map(|_| vec![0.0f32; conv_floats_per_slot])
.collect();
let free_slots = (0..capacity).collect();
Self {
s_matrices,
conv_buffers,
free_slots,
seq_to_slot: HashMap::new(),
capacity,
}
}
pub fn alloc(&mut self, seq_id: SeqId) -> Option<usize> {
let slot = self.free_slots.pop_front()?;
self.seq_to_slot.insert(seq_id, slot);
Some(slot)
}
pub fn free(&mut self, seq_id: SeqId) {
if let Some(slot) = self.seq_to_slot.remove(&seq_id) {
self.s_matrices[slot].fill(0.0);
self.conv_buffers[slot].fill(0.0);
self.free_slots.push_back(slot);
}
}
#[inline]
pub fn slot_of(&self, seq_id: SeqId) -> Option<usize> {
self.seq_to_slot.get(&seq_id).copied()
}
#[inline]
pub fn s_matrix_mut(&mut self, slot: usize) -> &mut [f32] {
&mut self.s_matrices[slot]
}
#[inline]
pub fn conv_buffer_mut(&mut self, slot: usize) -> &mut [f32] {
&mut self.conv_buffers[slot]
}
#[inline]
pub fn free_count(&self) -> usize {
self.free_slots.len()
}
#[inline]
pub fn capacity(&self) -> usize {
self.capacity
}
}
#[derive(Debug)]
pub struct BatchStepInput<'a> {
pub seq_id: SeqId,
pub token_ids: &'a [u32],
pub start_pos: usize,
pub gdn_slot: usize,
pub adapter_id: Option<&'a AdapterKey>,
}
#[derive(Debug)]
pub struct BatchStepOutput {
pub seq_id: SeqId,
pub logits: Vec<f32>,
}
#[derive(Debug)]
pub struct InferenceRequest {
pub prompt_ids: Vec<u32>,
pub sampling: SamplingConfig,
pub lora_adapter: Option<AdapterKey>,
pub max_new_tokens: usize,
}
#[derive(Debug, Clone)]
pub struct InferenceToken {
pub seq_id: SeqId,
pub token_id: u32,
pub finished: bool,
pub finish_reason: Option<FinishReason>,
}
pub struct BatchWorker {
config: BatchConfig,
seq_manager: SequenceManager,
gdn_pool: GdnStatePool,
kv_pool: PagePool,
scheduler: FifoScheduler,
samplers: HashMap<SeqId, Sampler>,
eos_token_id: Option<u32>,
output_buffer: VecDeque<InferenceToken>,
}
impl BatchWorker {
pub fn new(
config: BatchConfig,
kv_pool_config: PagedKVCacheConfig,
s_floats_per_slot: usize,
conv_floats_per_slot: usize,
eos_token_id: Option<u32>,
) -> Self {
let floats_per_page = kv_pool_config.floats_per_page_pub();
let kv_pool = PagePool::new(kv_pool_config.max_pages, floats_per_page);
let gdn_pool = GdnStatePool::new(
config.max_batch_size,
s_floats_per_slot,
conv_floats_per_slot,
);
let scheduler = FifoScheduler::new(config.clone());
Self {
config,
seq_manager: SequenceManager::new(),
gdn_pool,
kv_pool,
scheduler,
samplers: HashMap::new(),
eos_token_id,
output_buffer: VecDeque::new(),
}
}
pub fn submit(&mut self, request: InferenceRequest) -> Option<SeqId> {
if request.prompt_ids.is_empty() {
return None;
}
let total_len = request.prompt_ids.len() + request.max_new_tokens;
if total_len > self.config.max_seq_len {
return None;
}
let id = self.seq_manager.next_id();
let mut sampler = Sampler::new(request.sampling.clone());
sampler.seed_history(&request.prompt_ids);
let seq = Sequence::new(
id,
request.prompt_ids,
request.sampling,
request.lora_adapter,
request.max_new_tokens,
);
let page_size = self.kv_pool_page_size();
self.seq_manager.add(seq, page_size);
self.samplers.insert(id, sampler);
self.scheduler.enqueue(id);
Some(id)
}
pub fn step(
&mut self,
mut forward_fn: impl FnMut(BatchStepInput<'_>, &mut GdnStatePool) -> Vec<f32>,
) -> Vec<InferenceToken> {
let mut waiting = self.seq_manager.ids_in_state_prefilling();
let mut running = self.seq_manager.ids_decoding();
waiting.sort();
running.sort();
let kv_free = self.kv_pool.free_count();
let gdn_free = self.gdn_pool.free_count();
let decision = self
.scheduler
.select_batch(&waiting, &running, kv_free, gdn_free);
for seq_id in &decision.evict {
self.evict_sequence(*seq_id, FinishReason::Preempted);
self.scheduler.on_preempt(*seq_id);
}
for (seq_id, sentinel_start, chunk_len) in &decision.prefill {
let seq_id = *seq_id;
if self.gdn_pool.slot_of(seq_id).is_none() && self.gdn_pool.alloc(seq_id).is_none() {
continue;
}
let (real_start, real_len) = {
let Some(seq) = self.seq_manager.get(seq_id) else {
continue;
};
let start = if *sentinel_start == usize::MAX {
match seq.state {
SequenceState::Prefilling { chunk_start } => chunk_start,
_ => continue, }
} else {
*sentinel_start
};
let prompt_len = seq.prompt_ids.len();
let remaining = prompt_len.saturating_sub(start);
let len = remaining.min(*chunk_len);
(start, len)
};
if real_len == 0 {
continue;
}
let Some(gdn_slot) = self.gdn_pool.slot_of(seq_id) else {
continue;
};
let logits = {
let Some(seq) = self.seq_manager.get(seq_id) else {
continue;
};
let chunk = &seq.prompt_ids[real_start..real_start + real_len];
let adapter = seq.adapter_id.as_ref();
let input = BatchStepInput {
seq_id,
token_ids: chunk,
start_pos: real_start,
gdn_slot,
adapter_id: adapter,
};
forward_fn(input, &mut self.gdn_pool)
};
if let Some(seq) = self.seq_manager.get_mut(seq_id) {
seq.advance_prefill(real_len);
}
let just_finished_prefill = self
.seq_manager
.get(seq_id)
.is_some_and(|s| s.state == SequenceState::Decoding);
if just_finished_prefill {
let token_id = {
let Some(sampler) = self.samplers.get_mut(&seq_id) else {
continue;
};
sampler.sample(&logits)
};
let done = self
.seq_manager
.get_mut(seq_id)
.is_some_and(|s| s.push_token(token_id, self.eos_token_id));
let finish_reason = if done {
self.seq_manager.get(seq_id).and_then(|s| match s.state {
SequenceState::Finished(r) => Some(r),
_ => None,
})
} else {
None
};
self.output_buffer.push_back(InferenceToken {
seq_id,
token_id,
finished: done,
finish_reason,
});
}
}
for seq_id in &decision.decode {
let seq_id = *seq_id;
let Some(gdn_slot) = self.gdn_pool.slot_of(seq_id) else {
continue;
};
let logits = {
let Some(seq) = self.seq_manager.get(seq_id) else {
continue;
};
let Some(&last_token) = seq.generated_ids.last() else {
continue; };
let start_pos = seq.position().saturating_sub(1);
let token_buf = std::slice::from_ref(&last_token);
let adapter = seq.adapter_id.as_ref();
let input = BatchStepInput {
seq_id,
token_ids: token_buf,
start_pos,
gdn_slot,
adapter_id: adapter,
};
forward_fn(input, &mut self.gdn_pool)
};
let token_id = {
let Some(sampler) = self.samplers.get_mut(&seq_id) else {
continue;
};
sampler.sample(&logits)
};
let done = self
.seq_manager
.get_mut(seq_id)
.is_some_and(|s| s.push_token(token_id, self.eos_token_id));
let finish_reason = if done {
self.seq_manager.get(seq_id).and_then(|s| match s.state {
SequenceState::Finished(r) => Some(r),
_ => None,
})
} else {
None
};
self.output_buffer.push_back(InferenceToken {
seq_id,
token_id,
finished: done,
finish_reason,
});
}
let finished_ids: Vec<SeqId> = self.seq_manager.ids_finished();
for seq_id in finished_ids {
self.evict_sequence(seq_id, FinishReason::Eos); }
self.output_buffer.drain(..).collect()
}
#[inline]
pub fn active_count(&self) -> usize {
self.seq_manager.len()
}
#[inline]
pub fn is_idle(&self) -> bool {
self.seq_manager.is_empty() && self.scheduler.waiting_count() == 0
}
fn evict_sequence(&mut self, seq_id: SeqId, _reason: FinishReason) {
if let Some((_, table)) = self.seq_manager.remove(seq_id) {
for &phys in table.physical_pages() {
self.kv_pool.free(phys);
}
}
self.gdn_pool.free(seq_id);
self.samplers.remove(&seq_id);
}
fn kv_pool_page_size(&self) -> usize {
256
}
}
pub trait PagedKVCacheConfigExt {
fn floats_per_page_pub(&self) -> usize;
}
impl PagedKVCacheConfigExt for crate::kv_cache::PagedKVCacheConfig {
fn floats_per_page_pub(&self) -> usize {
self.num_layers * 2 * self.page_size * self.kv_dim()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::batch::config::BatchConfig;
use crate::kv_cache::{EvictionPolicy, PagedKVCacheConfig};
use crate::sampling::SamplingConfig;
fn test_kv_config() -> PagedKVCacheConfig {
PagedKVCacheConfig {
page_size: 4,
max_pages: 64,
num_layers: 2,
num_kv_heads: 2,
head_dim: 4,
eviction: EvictionPolicy::None,
}
}
fn test_worker() -> BatchWorker {
let config = BatchConfig {
max_batch_size: 4,
max_seq_len: 64,
chunk_size: 8,
prefill_reserve_pages: 2,
};
BatchWorker::new(
config,
test_kv_config(),
16, 8, Some(2), )
}
fn dummy_logits(winner: u32, vocab: usize) -> Vec<f32> {
let mut v = vec![0.0f32; vocab];
if (winner as usize) < vocab {
v[winner as usize] = 10.0;
}
v
}
#[test]
fn gdn_pool_alloc_and_free() {
let mut pool = GdnStatePool::new(2, 4, 2);
assert_eq!(pool.free_count(), 2);
let s0 = pool.alloc(SeqId(0));
assert!(s0.is_some());
assert_eq!(pool.free_count(), 1);
let s1 = pool.alloc(SeqId(1));
assert!(s1.is_some());
assert_eq!(pool.free_count(), 0);
assert!(pool.alloc(SeqId(2)).is_none());
pool.free(SeqId(0));
assert_eq!(pool.free_count(), 1);
}
#[test]
fn gdn_pool_free_zeroes_buffers() {
let mut pool = GdnStatePool::new(1, 4, 2);
let slot = pool.alloc(SeqId(0)).unwrap();
pool.s_matrix_mut(slot).fill(1.0);
pool.conv_buffer_mut(slot).fill(2.0);
pool.free(SeqId(0));
let slot2 = pool.alloc(SeqId(1)).unwrap();
assert_eq!(slot2, slot);
assert!(pool.s_matrix_mut(slot2).iter().all(|&x| x == 0.0));
assert!(pool.conv_buffer_mut(slot2).iter().all(|&x| x == 0.0));
}
#[test]
fn gdn_pool_slot_of() {
let mut pool = GdnStatePool::new(2, 4, 2);
pool.alloc(SeqId(5)).unwrap();
assert!(pool.slot_of(SeqId(5)).is_some());
assert!(pool.slot_of(SeqId(99)).is_none());
}
#[test]
fn submit_valid_request_returns_seq_id() {
let mut worker = test_worker();
let req = InferenceRequest {
prompt_ids: vec![1, 2, 3],
sampling: SamplingConfig::greedy(),
lora_adapter: None,
max_new_tokens: 4,
};
let id = worker.submit(req);
assert!(id.is_some());
}
#[test]
fn submit_empty_prompt_returns_none() {
let mut worker = test_worker();
let req = InferenceRequest {
prompt_ids: vec![],
sampling: SamplingConfig::greedy(),
lora_adapter: None,
max_new_tokens: 4,
};
assert!(worker.submit(req).is_none());
}
#[test]
fn submit_too_long_returns_none() {
let mut worker = test_worker(); let req = InferenceRequest {
prompt_ids: vec![0u32; 60],
sampling: SamplingConfig::greedy(),
lora_adapter: None,
max_new_tokens: 10, };
assert!(worker.submit(req).is_none());
}
#[test]
fn submit_seeds_prompt_history_for_repetition_penalty() {
let mut worker = test_worker();
worker.submit(InferenceRequest {
prompt_ids: vec![3], sampling: SamplingConfig {
temperature: 0.0, top_k: 1,
top_p: 1.0,
repetition_penalty: 2.0,
},
lora_adapter: None,
max_new_tokens: 1,
});
let mut logits = vec![0.0f32; 8];
logits[3] = 9.0; logits[4] = 5.5;
let out = worker.step(|_input, _pool| logits.clone());
assert_eq!(out.len(), 1);
assert_eq!(
out[0].token_id, 4,
"prompt token 3 must be penalized (9.0/2.0=4.5 < 5.5) so token 4 wins; \
without seed_history in submit(), token 3 is not penalized and wins raw \
(mutation: remove seed_history call from submit())"
);
}
#[test]
fn step_no_requests_returns_empty() {
let mut worker = test_worker();
let out = worker.step(|_input, _pool| vec![0.0f32; 8]);
assert!(out.is_empty());
}
#[test]
fn step_single_request_prefill_then_decode() {
let mut worker = test_worker(); worker.submit(InferenceRequest {
prompt_ids: vec![10, 11, 12, 13],
sampling: SamplingConfig::greedy(),
lora_adapter: None,
max_new_tokens: 2,
});
let step1 = worker.step(|_input, _pool| dummy_logits(5, 8));
assert_eq!(step1.len(), 1);
assert_eq!(step1[0].token_id, 5);
assert!(!step1[0].finished);
let step2 = worker.step(|_input, _pool| dummy_logits(6, 8));
assert_eq!(step2.len(), 1);
assert_eq!(step2[0].token_id, 6);
assert!(step2[0].finished);
assert_eq!(step2[0].finish_reason, Some(FinishReason::MaxLength));
}
#[test]
fn step_eos_finishes_sequence() {
let mut worker = test_worker(); worker.submit(InferenceRequest {
prompt_ids: vec![10],
sampling: SamplingConfig::greedy(),
lora_adapter: None,
max_new_tokens: 10,
});
let out = worker.step(|_input, _pool| dummy_logits(2, 8));
assert_eq!(out.len(), 1);
assert!(out[0].finished);
assert_eq!(out[0].finish_reason, Some(FinishReason::Eos));
}
#[test]
fn step_evicts_finished_sequence() {
let mut worker = test_worker();
worker.submit(InferenceRequest {
prompt_ids: vec![10],
sampling: SamplingConfig::greedy(),
lora_adapter: None,
max_new_tokens: 1,
});
worker.step(|_input, _pool| dummy_logits(5, 8));
assert!(worker.is_idle());
}
#[test]
fn step_chunked_prefill_multiple_chunks() {
let mut worker = BatchWorker::new(
BatchConfig {
max_batch_size: 4,
max_seq_len: 128,
chunk_size: 3, prefill_reserve_pages: 0,
},
test_kv_config(),
16,
8,
Some(99), );
worker.submit(InferenceRequest {
prompt_ids: vec![1, 2, 3, 4, 5, 6, 7],
sampling: SamplingConfig::greedy(),
lora_adapter: None,
max_new_tokens: 2,
});
let s1 = worker.step(|_input, _pool| dummy_logits(10, 12));
assert!(s1.is_empty(), "no output yet during prefill");
let s2 = worker.step(|_input, _pool| dummy_logits(10, 12));
assert!(s2.is_empty());
let s3 = worker.step(|_input, _pool| dummy_logits(10, 12));
assert_eq!(s3.len(), 1);
assert_eq!(s3[0].token_id, 10);
}
#[test]
fn step_multiple_concurrent_sequences() {
let mut worker = test_worker();
worker.submit(InferenceRequest {
prompt_ids: vec![1],
sampling: SamplingConfig::greedy(),
lora_adapter: None,
max_new_tokens: 1,
});
worker.submit(InferenceRequest {
prompt_ids: vec![2],
sampling: SamplingConfig::greedy(),
lora_adapter: None,
max_new_tokens: 1,
});
let out = worker.step(|_input, _pool| dummy_logits(7, 8));
assert_eq!(out.len(), 2);
assert!(out.iter().all(|t| t.token_id == 7));
assert!(out.iter().all(|t| t.finished));
assert!(worker.is_idle());
}
#[test]
fn floats_per_page_ext() {
let cfg = test_kv_config(); assert_eq!(cfg.floats_per_page_pub(), 128);
}
}