use std::collections::{HashMap, HashSet};
use std::sync::mpsc::Sender;
use ferrox_core::cache::KvCache;
use ferrox_models::tokenizer::StopTokens;
use crate::generate::{DecodeError, FinishReason, GenerationParams, PagedLease};
use crate::sample_step::SampleState;
use crate::stop::StopMatcher;
use super::block_budget::BlockBudget;
use super::clock::RowClock;
use super::config::{send_finished, BatcherEvent};
use super::queue::AbortId;
use crate::utf8_stream::Utf8Stream;
pub(super) enum RowKv {
Contiguous(Vec<KvCache>),
Paged(PagedLease),
}
impl RowKv {
pub(super) fn positions_written(&mut self) -> usize {
match self {
RowKv::Contiguous(caches) => caches.first().map_or(0, |c| c.positions()),
RowKv::Paged(lease) => lease.caches_mut().first().map_or(0, |c| c.seq_len()),
}
}
}
pub(super) struct Job {
pub(super) prompt_tokens: Vec<usize>,
pub(super) params: GenerationParams,
pub(super) stop_tokens: StopTokens,
pub(super) reply: Sender<BatcherEvent>,
pub(super) abort: AbortId,
pub(super) blocks: usize,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub(super) struct Uid(u64);
#[derive(Default)]
pub(super) struct Rows {
pub(super) state: HashMap<Uid, Slot>,
pub(super) order: Vec<Uid>,
pub(super) next_uid: u64,
}
impl Rows {
pub(super) 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
}
pub(super) fn len(&self) -> usize {
self.order.len()
}
pub(super) fn blocks_held(&self) -> usize {
self.state.values().map(|slot| slot.blocks).sum()
}
pub(super) 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
}
pub(super) fn is_empty(&self) -> bool {
self.order.is_empty()
}
pub(super) fn get(&self, uid: Uid) -> Option<&Slot> {
self.state.get(&uid)
}
pub(super) fn get_mut(&mut self, uid: Uid) -> Option<&mut Slot> {
self.state.get_mut(&uid)
}
pub(super) fn remove(&mut self, uid: Uid) -> Option<Slot> {
self.order.retain(|&u| u != uid);
self.state.remove(&uid)
}
pub(super) 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()
}
pub(super) 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(mut slot) = self.remove(uid) {
budget.release(slot.blocks);
if let RowKv::Paged(lease) = &mut slot.kv {
let mut seq = std::mem::take(&mut slot.prompt_ids);
seq.extend_from_slice(&slot.generated_ids);
let bs = lease.block_size();
crate::generate::publish_to_radix(lease, &seq, bs);
}
reply_finished(slot);
}
}
}
}
pub(super) struct Slot {
pub(super) kv: RowKv,
pub(super) pos: usize,
pub(super) logits: Vec<f32>,
pub(super) sample: SampleState,
pub(super) generated_ids: Vec<usize>,
pub(super) prompt_ids: Vec<usize>,
pub(super) visible: String,
pub(super) stops: StopMatcher,
pub(super) prompt_tokens: usize,
pub(super) max_tokens: usize,
pub(super) stop_tokens: StopTokens,
pub(super) params: GenerationParams,
pub(super) reply: Sender<BatcherEvent>,
pub(super) finish: Option<FinishReason>,
pub(super) error: Option<DecodeError>,
pub(super) abort: AbortId,
pub(super) blocks: usize,
pub(super) clock: RowClock,
pub(super) utf8: Utf8Stream,
}
impl Slot {
pub(super) fn fail(&mut self, error: DecodeError) {
self.finish = Some(FinishReason::Stop);
self.error = Some(error);
}
}
pub(super) fn reply_finished(mut slot: Slot) {
if let Some(error) = slot.error {
send_finished(&slot.reply, Err(error));
return;
}
let finish = slot.finish.expect("only a finished row is replied to");
let partial = slot.utf8.flush();
if !partial.is_empty() {
slot.stops.push(&partial);
}
let tail = slot.stops.flush();
if !tail.is_empty() {
slot.visible.push_str(&tail);
let _ = slot.reply.send(BatcherEvent::Chunk(tail));
}
let usage = slot
.clock
.usage(slot.prompt_tokens, slot.generated_ids.len());
send_finished(
&slot.reply,
Ok((finish, slot.generated_ids, slot.visible, usage)),
);
}