use std::sync::mpsc::Sender;
use ferrox_core::cache::KvCache;
use ferrox_models::tokenizer::StopTokens;
use ferrox_models::Decoder;
use std::sync::Arc;
use crate::generate::{GenerationParams, PagedLease};
use super::config::BatcherEvent;
use super::queue::AbortId;
use super::row::{RowKv, Slot};
use crate::sample_step::SampleState;
use crate::stop::StopMatcher;
pub struct PrefillState {
decoder: Arc<Decoder>,
kv: RowKv,
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 {
let kv = RowKv::Contiguous(
decoder
.layers
.iter()
.map(|_| KvCache::new(decoder.config.n_kv_heads, decoder.config.head_dim))
.collect(),
);
PrefillState::over(decoder, kv, prompt_tokens, chunk_size)
}
pub fn new_paged(
decoder: Arc<Decoder>,
prompt_tokens: &[usize],
chunk_size: usize,
lease: PagedLease,
) -> Self {
PrefillState::over(decoder, RowKv::Paged(lease), prompt_tokens, chunk_size)
}
fn over(
decoder: Arc<Decoder>,
mut kv: RowKv,
prompt_tokens: &[usize],
chunk_size: usize,
) -> Self {
assert!(chunk_size > 0, "prefill chunk size must be positive");
let tokens = if prompt_tokens.is_empty() {
vec![0]
} else {
prompt_tokens.to_vec()
};
let tokens_processed = kv.positions_written();
debug_assert!(
tokens_processed < tokens.len(),
"a prefill must have at least one token left to run, \
got {tokens_processed} already done of {} -- the adopted \
prefix was not backed off",
tokens.len()
);
PrefillState {
decoder,
kv,
tokens,
tokens_processed,
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());
if self.tokens_processed >= end {
return self.is_done();
}
let pos = self.kv.positions_written();
debug_assert_eq!(
pos, self.tokens_processed,
"the KV write cursor and the prompt cursor diverged"
);
let chunk = &self.tokens[self.tokens_processed..end];
self.logits = match &mut self.kv {
RowKv::Contiguous(caches) => {
self.decoder.forward_batch_last_host_kv(chunk, pos, caches)
}
RowKv::Paged(lease) => {
let store = Arc::clone(lease.store());
self.decoder
.forward_batch_last_paged(chunk, pos, lease.caches_mut(), &store)
.expect("the row's pages were reserved at admission")
}
};
self.tokens_processed = end;
self.is_done()
}
pub(super) fn into_decode_start(self) -> (RowKv, Vec<f32>, usize, Vec<usize>) {
debug_assert!(self.is_done(), "prefill must finish before decoding");
(self.kv, self.logits, self.tokens_processed, self.tokens)
}
}
pub(super) struct Prefill {
pub(super) state: PrefillState,
pub(super) prompt_tokens: usize,
pub(super) params: GenerationParams,
pub(super) stop_tokens: StopTokens,
pub(super) reply: Sender<BatcherEvent>,
pub(super) abort: AbortId,
pub(super) blocks: usize,
}
impl Prefill {
pub(super) fn into_slot(self) -> Slot {
let Prefill {
state,
prompt_tokens,
params,
stop_tokens,
reply,
abort,
blocks,
} = self;
let (kv, logits, pos, prompt_ids) = state.into_decode_start();
Slot {
kv,
pos,
logits,
sample: SampleState::new(params.seed),
generated_ids: Vec::with_capacity(params.max_tokens),
prompt_ids,
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,
error: None,
abort,
blocks,
}
}
}