pub mod embeddings;
pub mod kv_cache;
pub mod params;
pub mod sampling_state;
pub mod session;
use std::ptr::NonNull;
use llama_crab_sys as sys;
use crate::batch::LlamaBatch;
use crate::error::{LlamaError, Result};
use crate::model::LlamaModel;
#[derive(Debug)]
pub struct LlamaContext {
pub(crate) handle: NonNull<sys::llama_context>,
pub(crate) model: NonNull<LlamaModel>,
}
impl LlamaContext {
pub(crate) fn from_raw(
handle: NonNull<sys::llama_context>,
model: NonNull<LlamaModel>,
) -> Self {
Self { handle, model }
}
#[must_use]
pub fn n_ctx(&self) -> u32 {
unsafe { sys::llama_n_ctx(self.handle.as_ptr()) as u32 }
}
#[must_use]
pub fn n_batch(&self) -> u32 {
unsafe { sys::llama_n_batch(self.handle.as_ptr()) as u32 }
}
#[must_use]
pub fn n_ubatch(&self) -> u32 {
unsafe { sys::llama_n_ubatch(self.handle.as_ptr()) as u32 }
}
#[must_use]
pub fn n_seq_max(&self) -> u32 {
unsafe { sys::llama_n_seq_max(self.handle.as_ptr()) as u32 }
}
#[must_use]
pub fn raw_handle(&self) -> *mut sys::llama_context {
self.handle.as_ptr()
}
pub fn decode(&mut self, batch: &LlamaBatch) -> Result<()> {
let rc = unsafe { sys::llama_decode(self.handle.as_ptr(), *batch.raw()) };
if rc != 0 {
return Err(LlamaError::Decode(rc));
}
Ok(())
}
pub fn encode(&mut self, batch: &LlamaBatch) -> Result<()> {
let rc = unsafe { sys::llama_encode(self.handle.as_ptr(), *batch.raw()) };
if rc != 0 {
return Err(LlamaError::Encode(rc));
}
Ok(())
}
#[must_use]
pub fn model(&self) -> &LlamaModel {
unsafe { &*self.model.as_ptr() }
}
pub(crate) fn raw(&self) -> *mut sys::llama_context {
self.handle.as_ptr()
}
}
unsafe impl Send for LlamaContext {}
unsafe impl Sync for LlamaContext {}
impl Drop for LlamaContext {
fn drop(&mut self) {
unsafe { sys::llama_free(self.handle.as_ptr()) };
}
}
pub use self::params::LlamaContextParams;