use ik_llama_cpp_sys as sys;
use crate::token::LlamaToken;
use crate::LlamaError;
#[derive(Debug)]
pub struct LlamaBatch {
batch: sys::llama_batch,
capacity: usize,
n_seq_max: usize,
}
impl LlamaBatch {
#[must_use]
pub fn new(capacity: usize, n_seq_max: usize) -> Self {
assert!(
capacity <= i32::MAX as usize && n_seq_max <= i32::MAX as usize,
"batch dims exceed i32::MAX (capacity={capacity}, n_seq_max={n_seq_max})"
);
let batch = unsafe { sys::llama_batch_init(capacity as i32, 0, n_seq_max as i32) };
Self {
batch,
capacity,
n_seq_max,
}
}
pub fn clear(&mut self) {
self.batch.n_tokens = 0;
}
#[must_use]
pub fn n_tokens(&self) -> i32 {
self.batch.n_tokens
}
pub fn add(
&mut self,
token: LlamaToken,
pos: i32,
seq_ids: &[i32],
logits: bool,
) -> Result<(), LlamaError> {
let i = self.batch.n_tokens as usize;
if i >= self.capacity {
return Err(LlamaError::BatchOverflow {
capacity: self.capacity,
index: i,
});
}
if seq_ids.len() > self.n_seq_max {
return Err(LlamaError::TooManySeqIds {
got: seq_ids.len(),
max: self.n_seq_max,
});
}
unsafe {
*self.batch.token.add(i) = token.0;
*self.batch.pos.add(i) = pos;
*self.batch.n_seq_id.add(i) = seq_ids.len() as i32;
let seq_row = *self.batch.seq_id.add(i);
for (j, &s) in seq_ids.iter().enumerate() {
*seq_row.add(j) = s;
}
*self.batch.logits.add(i) = i8::from(logits);
}
self.batch.n_tokens += 1;
Ok(())
}
pub fn add_sequence(
&mut self,
tokens: &[LlamaToken],
seq_id: i32,
logits_all: bool,
) -> Result<(), LlamaError> {
let n = tokens.len();
for (i, &t) in tokens.iter().enumerate() {
let want_logits = logits_all || i + 1 == n;
self.add(t, i as i32, &[seq_id], want_logits)?;
}
Ok(())
}
pub(crate) fn as_raw(&self) -> sys::llama_batch {
self.batch
}
}
impl Drop for LlamaBatch {
fn drop(&mut self) {
unsafe { sys::llama_batch_free(self.batch) };
}
}