use ferrox_core::cache::{KvCache, PagedKvCache, PagedStoreExhausted, SharedPagedKv};
use ferrox_core::par;
use super::{Decoder, MultiSeqKv};
impl Decoder {
pub fn forward_token(
&self,
token_id: usize,
pos: usize,
kv_caches: &mut [KvCache],
) -> Vec<f32> {
par::on_workers(move || self.forward_token_on_worker(token_id, pos, kv_caches))
}
pub fn forward_token_paged(
&self,
token_id: usize,
pos: usize,
kv_caches: &mut [PagedKvCache],
stores: &SharedPagedKv,
) -> Result<Vec<f32>, PagedStoreExhausted> {
par::on_workers(move || {
self.forward_token_paged_on_worker(token_id, pos, kv_caches, stores)
})
}
pub fn forward_batch(
&self,
tokens: &[usize],
start_pos: usize,
kv_caches: &mut [KvCache],
) -> Vec<Vec<f32>> {
par::on_workers(move || {
let hiddens = self.forward_hidden_batch(tokens, start_pos, kv_caches);
if hiddens.is_empty() {
return Vec::new();
}
let batch_size = hiddens.len();
let flat: Vec<f32> = hiddens.into_iter().flatten().collect();
self.logits_from_flat_hidden(flat, batch_size)
})
}
pub fn forward_batch_with_hidden(
&self,
tokens: &[usize],
start_pos: usize,
kv_caches: &mut [KvCache],
) -> (Vec<Vec<f32>>, Vec<Vec<f32>>) {
par::on_workers(move || {
let hiddens = self.forward_hidden_batch(tokens, start_pos, kv_caches);
if hiddens.is_empty() {
return (Vec::new(), Vec::new());
}
let batch_size = hiddens.len();
let flat: Vec<f32> = hiddens.iter().flatten().copied().collect();
(self.logits_from_flat_hidden(flat, batch_size), hiddens)
})
}
pub fn forward_batch_last(
&self,
tokens: &[usize],
start_pos: usize,
kv_caches: &mut [KvCache],
) -> Vec<f32> {
par::on_workers(move || self.forward_batch_last_inner(tokens, start_pos, kv_caches, false))
}
pub fn forward_batch_last_host_kv(
&self,
tokens: &[usize],
start_pos: usize,
kv_caches: &mut [KvCache],
) -> Vec<f32> {
par::on_workers(move || self.forward_batch_last_inner(tokens, start_pos, kv_caches, true))
}
pub fn forward_batch_last_paged(
&self,
tokens: &[usize],
start_pos: usize,
kv_caches: &mut [PagedKvCache],
stores: &SharedPagedKv,
) -> Result<Vec<f32>, PagedStoreExhausted> {
par::on_workers(move || {
self.forward_batch_last_paged_on_worker(tokens, start_pos, kv_caches, stores)
})
}
pub fn forward_hidden_batch(
&self,
tokens: &[usize],
start_pos: usize,
kv_caches: &mut [KvCache],
) -> Vec<Vec<f32>> {
par::on_workers(move || {
self.forward_hidden_batch_inner(tokens, start_pos, kv_caches, false)
})
}
pub fn forward_multi_seq(
&self,
tokens: &[usize],
positions: &[usize],
kv_caches: &mut [Vec<KvCache>],
) -> Vec<Vec<f32>> {
par::on_workers(move || {
self.forward_multi_seq_kv(tokens, positions, &mut MultiSeqKv::Contiguous(kv_caches))
})
}
pub fn forward_multi_seq_kv(
&self,
tokens: &[usize],
positions: &[usize],
kv: &mut MultiSeqKv<'_>,
) -> Vec<Vec<f32>> {
par::on_workers(move || self.forward_multi_seq_kv_on_worker(tokens, positions, kv))
}
}
#[cfg(test)]
mod tests {
#[test]
fn a_public_forward_may_not_be_declared_outside_this_file() {
let body = include_str!("../decoder.rs");
let stray: Vec<&str> = body
.lines()
.map(str::trim)
.filter(|l| l.starts_with("pub fn forward"))
.collect();
assert!(
stray.is_empty(),
"these forward entry points bypass the pool wrapper in entry.rs: {stray:?}"
);
}
}