use std::collections::{HashSet, VecDeque};
use std::sync::atomic::Ordering;
use std::sync::mpsc::Receiver;
use std::sync::Arc;
use ferrox_core::cache::{KvCache, PagedKvCache};
use ferrox_models::Decoder;
use ferrox_models::MultiSeqKv;
use crate::generate::{acquire_paged_caches, DecodeError, FinishReason, PagedKvConfig, Usage};
use crate::stop::StopStep;
use super::block_budget::BlockBudget;
use super::clock::RowClock;
use super::config::{decode_log_interval_from_env, send_finished, BatcherConfig, DecodeFn};
use super::counters::Counters;
use super::prefill::{Prefill, PrefillState};
use super::queue::{AbortId, AbortInbox, QueueGate};
use super::row::{Job, RowKv, Rows, Slot, Uid};
use super::status::{BatchStatus, PoolUsage, PrefillSnapshot, StatusReporter};
pub(super) fn drain_channel(rx: &Receiver<Job>, waiting: &mut VecDeque<Job>) {
while let Ok(job) = rx.try_recv() {
waiting.push_back(job);
}
}
#[allow(clippy::too_many_arguments)] pub(super) fn admit(
decoder: &Arc<Decoder>,
waiting: &mut VecDeque<Job>,
prefills: &mut VecDeque<Prefill>,
decoding: usize,
config: &BatcherConfig,
queue: &QueueGate,
budget: &BlockBudget,
paged: Option<&PagedKvConfig>,
) -> PrefillSnapshot {
let mut snapshot = PrefillSnapshot::default();
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, paged) {
Some(prefill) => {
snapshot.new_seqs += 1;
snapshot.new_tokens += prefill.prompt_tokens;
prefills.push_back(prefill);
}
None => budget.release(blocks),
}
}
snapshot
}
pub(super) 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();
send_finished(
&job.reply,
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);
send_finished(
&prefill.reply,
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);
}
}
pub(super) fn batch_status(
rows: &Rows,
prefills: &VecDeque<Prefill>,
waiting: &VecDeque<Job>,
budget: &BlockBudget,
) -> BatchStatus {
let kv_pages = match budget.total {
Some(total) => PoolUsage::from_available(total, budget.free.load(Ordering::Relaxed)),
None => PoolUsage::default(),
};
BatchStatus {
running_reqs: rows.len() + prefills.len(),
queue_reqs: waiting.len(),
kv_pages,
page_size: budget.block_size,
window: None,
recurrent: None,
}
}
#[allow(clippy::too_many_arguments)]
pub(super) 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>,
paged: Option<PagedKvConfig>,
) {
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();
let started = std::time::Instant::now();
let mut reporter = StatusReporter::new(
decode_log_interval_from_env(),
started.elapsed().as_secs_f64(),
);
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);
let admitted = admit(
&decoder,
&mut waiting,
&mut prefills,
rows.len(),
&config,
&queue,
&budget,
paged.as_ref(),
);
if admitted.new_seqs > 0 {
let status = batch_status(&rows, &prefills, &waiting, &budget);
tracing::info!(
"{}",
reporter.report_prefill(started.elapsed().as_secs_f64(), &admitted, &status)
);
}
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 = match crate::sample_step::sample_next(
&mut slot.sample,
&slot.logits,
&slot.params,
&slot.prompt_ids,
&slot.generated_ids,
&slot.stop_tokens,
&|id| decode(&[id]),
) {
Ok(crate::sample_step::Step::Token(next)) => next,
Ok(crate::sample_step::Step::GrammarComplete) => {
slot.finish = Some(FinishReason::Stop);
continue;
}
Err(e) => {
slot.fail(e);
continue;
}
};
let model_eos = !slot.params.ignore_eos && slot.stop_tokens.contains(next);
if model_eos || slot.stops.is_stop_token(next) {
slot.finish = Some(FinishReason::Stop);
continue;
}
slot.generated_ids.push(next);
slot.clock.token();
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 paged = matches!(rows.get(active[0]).map(|s| &s.kv), Some(RowKv::Paged(_)));
let mut contiguous_refs: Vec<Vec<KvCache>> = Vec::new();
let mut paged_refs: Vec<Vec<PagedKvCache>> = Vec::new();
for &uid in &active {
let slot = rows.get_mut(uid).expect("active row exists");
let pos = slot.pos;
let token = *slot
.generated_ids
.last()
.expect("an active row has a token");
match &mut slot.kv {
RowKv::Contiguous(c) => contiguous_refs.push(std::mem::take(c)),
RowKv::Paged(lease) => {
lease.observe_sampled(token, pos + 1, false);
lease.before_step(pos);
paged_refs.push(std::mem::take(lease.caches_mut()));
}
}
}
debug_assert!(
contiguous_refs.len() == active.len() || paged_refs.len() == active.len(),
"every row in a batch must share one KV backing"
);
let logits_batch = if paged {
let store = Arc::clone(match &rows.get(active[0]).expect("active row exists").kv {
RowKv::Paged(lease) => lease.store(),
RowKv::Contiguous(_) => unreachable!("checked above"),
});
decoder.forward_multi_seq_kv(
&tokens,
&positions,
&mut MultiSeqKv::Paged {
caches: &mut paged_refs,
stores: &store,
},
)
} else {
decoder.forward_multi_seq(&tokens, &positions, &mut contiguous_refs)
};
counters.decode_steps.fetch_add(1, Ordering::Relaxed);
if let Some(line) = reporter.report_decode(
started.elapsed().as_secs_f64(),
active.len(),
&batch_status(&rows, &prefills, &waiting, &budget),
) {
tracing::info!("{line}");
}
for (j, &uid) in active.iter().enumerate() {
let slot = rows
.get_mut(uid)
.expect("an active row cannot vanish mid-step");
match &mut slot.kv {
RowKv::Contiguous(c) => *c = std::mem::take(&mut contiguous_refs[j]),
RowKv::Paged(lease) => *lease.caches_mut() = std::mem::take(&mut paged_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);
}
}
pub(super) fn apply_stop_buffer(slot: &mut Slot, piece: &str) -> bool {
match slot.stops.push(piece) {
StopStep::Emit(text) => {
if !text.is_empty() {
slot.visible.push_str(&text);
let _ = slot.reply.send(super::config::BatcherEvent::Chunk(text));
}
false
}
StopStep::Matched { text, stop } => {
if !text.is_empty() {
slot.visible.push_str(&text);
let _ = slot.reply.send(super::config::BatcherEvent::Chunk(text));
}
slot.finish = Some(FinishReason::StopSequence(stop));
true
}
}
}
pub(super) fn accept(
decoder: &Arc<Decoder>,
job: Job,
chunk_size: usize,
paged: Option<&PagedKvConfig>,
) -> Option<Prefill> {
let vocab_size = decoder.config.vocab_size;
if let Some(&bad) = job.prompt_tokens.iter().find(|&&t| t >= vocab_size) {
send_finished(
&job.reply,
Err(DecodeError::TokenOutOfVocab {
token: bad,
vocab_size,
}),
);
return None;
}
let state = match paged {
Some(config) => {
let max_seq_len = job.prompt_tokens.len() + job.params.max_tokens;
match acquire_paged_caches(decoder, config, &job.prompt_tokens, max_seq_len) {
Ok(lease) => PrefillState::new_paged(
Arc::clone(decoder),
&job.prompt_tokens,
chunk_size,
lease,
),
Err(_) => {
send_finished(&job.reply, Err(DecodeError::KvPoolExhausted));
return None;
}
}
}
None => PrefillState::new(Arc::clone(decoder), &job.prompt_tokens, chunk_size),
};
Some(Prefill {
state,
clock: RowClock::start(),
prompt_tokens: job.prompt_tokens.len(),
params: job.params,
stop_tokens: job.stop_tokens,
reply: job.reply,
abort: job.abort,
blocks: job.blocks,
})
}