use std::ffi::CString;
use std::ptr::NonNull;
use llama_cpp_sys_4 as sys;
use crate::shim::{check_status, last_error, read_tokens, ShimError};
use crate::token::LlamaToken;
pub type NgramError = ShimError;
type Result<T> = std::result::Result<T, NgramError>;
#[allow(clippy::similar_names)]
pub fn ngram_simple_draft(
size_ngram: u16,
size_mgram: u16,
tokens: &[LlamaToken],
sampled: LlamaToken,
) -> Result<Vec<LlamaToken>> {
let raw: Vec<i32> = tokens.iter().map(|t| t.0).collect();
read_tokens(|out, cap, len| unsafe {
sys::common_shim_ngram_simple_draft(
size_ngram,
size_mgram,
raw.as_ptr(),
raw.len(),
sampled.0,
out,
cap,
len,
)
})
}
pub struct NgramCache {
raw: NonNull<sys::common_shim_ngram_cache>,
}
impl std::fmt::Debug for NgramCache {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("NgramCache").field("len", &self.len()).finish()
}
}
unsafe impl Send for NgramCache {}
impl Drop for NgramCache {
fn drop(&mut self) {
unsafe { sys::common_shim_ngram_cache_free(self.raw.as_ptr()) }
}
}
impl Default for NgramCache {
fn default() -> Self {
Self::new()
}
}
impl NgramCache {
#[must_use]
pub fn new() -> Self {
let raw = unsafe { sys::common_shim_ngram_cache_init() };
Self {
raw: NonNull::new(raw).expect("common_shim_ngram_cache_init returned null"),
}
}
pub fn load(path: &str) -> Result<Self> {
let c_path = CString::new(path)?;
let raw = unsafe { sys::common_shim_ngram_cache_load(c_path.as_ptr()) };
NonNull::new(raw)
.map(|raw| Self { raw })
.ok_or_else(|| NgramError::Failed(last_error()))
}
pub fn save(&mut self, path: &str) -> Result<()> {
let c_path = CString::new(path)?;
let status =
unsafe { sys::common_shim_ngram_cache_save(self.raw.as_ptr(), c_path.as_ptr()) };
check_status(status)
}
pub fn merge(&mut self, other: &mut NgramCache) -> Result<()> {
let status =
unsafe { sys::common_shim_ngram_cache_merge(self.raw.as_ptr(), other.raw.as_ptr()) };
check_status(status)
}
#[must_use]
pub fn len(&self) -> usize {
unsafe { sys::common_shim_ngram_cache_size(self.raw.as_ptr()) }
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn update(
&mut self,
ngram_min: i32,
ngram_max: i32,
tokens: &[LlamaToken],
nnew: i32,
print_progress: bool,
) -> Result<()> {
let raw: Vec<i32> = tokens.iter().map(|t| t.0).collect();
let status = unsafe {
sys::common_shim_ngram_cache_update(
self.raw.as_ptr(),
ngram_min,
ngram_max,
raw.as_ptr(),
raw.len(),
nnew,
print_progress,
)
};
check_status(status)
}
}
pub fn ngram_cache_draft(
tokens: &[LlamaToken],
n_draft: i32,
ngram_min: i32,
ngram_max: i32,
context: Option<&mut NgramCache>,
dynamic: Option<&mut NgramCache>,
statik: Option<&mut NgramCache>,
) -> Result<Vec<LlamaToken>> {
if tokens.is_empty() {
return Err(NgramError::InvalidArg);
}
let raw: Vec<i32> = tokens.iter().map(|t| t.0).collect();
let ctx_ptr = context.map_or(std::ptr::null_mut(), |c| c.raw.as_ptr());
let dyn_ptr = dynamic.map_or(std::ptr::null_mut(), |c| c.raw.as_ptr());
let sta_ptr = statik.map_or(std::ptr::null_mut(), |c| c.raw.as_ptr());
read_tokens(|out, cap, len| unsafe {
sys::common_shim_ngram_cache_draft(
raw.as_ptr(),
raw.len(),
n_draft,
ngram_min,
ngram_max,
ctx_ptr,
dyn_ptr,
sta_ptr,
out,
cap,
len,
)
})
}
pub struct NgramMap {
raw: NonNull<sys::common_shim_ngram_map>,
}
impl std::fmt::Debug for NgramMap {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("NgramMap").finish_non_exhaustive()
}
}
unsafe impl Send for NgramMap {}
impl Drop for NgramMap {
fn drop(&mut self) {
unsafe { sys::common_shim_ngram_map_free(self.raw.as_ptr()) }
}
}
impl NgramMap {
pub fn new(size_key: u16, size_value: u16, key_only: bool, min_hits: u16) -> Result<Self> {
let raw =
unsafe { sys::common_shim_ngram_map_init(size_key, size_value, key_only, min_hits) };
NonNull::new(raw)
.map(|raw| Self { raw })
.ok_or_else(|| NgramError::Failed(last_error()))
}
pub fn begin(&mut self, tokens: &[LlamaToken]) -> Result<()> {
let raw: Vec<i32> = tokens.iter().map(|t| t.0).collect();
let status = unsafe {
sys::common_shim_ngram_map_begin(self.raw.as_ptr(), raw.as_ptr(), raw.len())
};
check_status(status)
}
pub fn draft(&mut self, tokens: &[LlamaToken], sampled: LlamaToken) -> Result<Vec<LlamaToken>> {
let raw: Vec<i32> = tokens.iter().map(|t| t.0).collect();
read_tokens(|out, cap, len| unsafe {
sys::common_shim_ngram_map_draft(
self.raw.as_ptr(),
raw.as_ptr(),
raw.len(),
sampled.0,
out,
cap,
len,
)
})
}
pub fn accept(&mut self, n_accepted: u16) {
unsafe { sys::common_shim_ngram_map_accept(self.raw.as_ptr(), n_accepted) }
}
}
#[cfg(test)]
mod tests {
use super::*;
fn toks(v: &[i32]) -> Vec<LlamaToken> {
v.iter().copied().map(LlamaToken).collect()
}
#[test]
fn simple_draft_predicts_a_repeat() {
let history = toks(&[9, 1, 2, 3, 4, 5, 1]);
let draft = ngram_simple_draft(2, 2, &history, LlamaToken(2)).unwrap();
assert_eq!(
draft,
toks(&[3, 4]),
"expected the continuation of the earlier `1 2`"
);
}
#[test]
fn simple_draft_with_sampled_already_in_history_finds_nothing() {
let history = toks(&[9, 1, 2, 3, 4, 5, 1, 2]);
let draft = ngram_simple_draft(2, 2, &history, LlamaToken(2)).unwrap();
assert!(draft.is_empty(), "got {draft:?}");
}
#[test]
fn simple_draft_is_empty_below_the_length_floor() {
let history = toks(&[1, 2, 1, 2, 1]);
assert!(ngram_simple_draft(2, 2, &history, LlamaToken(2))
.unwrap()
.is_empty());
}
#[test]
fn simple_draft_is_empty_without_a_repeat() {
let history = toks(&[1, 2, 3, 4, 5]);
let draft = ngram_simple_draft(2, 2, &history, LlamaToken(5)).unwrap();
assert!(draft.is_empty(), "expected no draft, got {draft:?}");
}
#[test]
fn simple_draft_handles_empty_history() {
assert!(ngram_simple_draft(2, 2, &[], LlamaToken(1)).unwrap().is_empty());
}
#[test]
fn cache_starts_empty_and_learns() {
let mut cache = NgramCache::new();
assert!(cache.is_empty());
let tokens = toks(&[1, 2, 3, 1, 2, 3, 1, 2, 3]);
cache
.update(1, 4, &tokens, i32::try_from(tokens.len()).unwrap(), false)
.unwrap();
assert!(!cache.is_empty(), "update recorded nothing");
}
#[test]
fn cache_round_trips_through_disk() {
let dir = std::env::temp_dir();
let path = dir.join("llama_cpp_rs_ngram_test.bin");
let path_str = path.to_str().unwrap();
let mut cache = NgramCache::new();
let tokens = toks(&[7, 8, 9, 7, 8, 9, 7, 8, 9]);
cache
.update(1, 4, &tokens, i32::try_from(tokens.len()).unwrap(), false)
.unwrap();
let saved_len = cache.len();
cache.save(path_str).unwrap();
let loaded = NgramCache::load(path_str).unwrap();
assert_eq!(loaded.len(), saved_len, "cache changed size across disk");
let _ = std::fs::remove_file(&path);
}
#[test]
fn cache_load_rejects_a_missing_file() {
assert!(NgramCache::load("/definitely/not/a/cache.bin").is_err());
}
#[test]
fn cache_rejects_interior_nul_in_path() {
let mut cache = NgramCache::new();
assert!(matches!(cache.save("a\0b"), Err(NgramError::Nul(_))));
assert!(matches!(NgramCache::load("a\0b"), Err(NgramError::Nul(_))));
}
#[test]
fn cache_merge_is_additive() {
let tokens = toks(&[4, 5, 6, 4, 5, 6, 4, 5, 6]);
let mut a = NgramCache::new();
a.update(1, 4, &tokens, i32::try_from(tokens.len()).unwrap(), false).unwrap();
let before = a.len();
let mut b = NgramCache::new();
b.update(1, 4, &tokens, i32::try_from(tokens.len()).unwrap(), false).unwrap();
a.merge(&mut b).unwrap();
assert!(a.len() >= before, "merge lost entries");
}
#[test]
fn cache_draft_requires_tokens() {
assert!(matches!(
ngram_cache_draft(&[], 4, 1, 4, None, None, None),
Err(NgramError::InvalidArg)
));
}
#[test]
fn cache_draft_with_no_caches_is_empty() {
let tokens = toks(&[1, 2, 3]);
let draft = ngram_cache_draft(&tokens, 4, 1, 4, None, None, None).unwrap();
assert!(draft.is_empty(), "got {draft:?}");
}
#[test]
fn cache_draft_predicts_a_learned_repeat() {
let tokens = toks(&[1, 2, 3, 1, 2, 3, 1, 2, 3, 1, 2]);
let mut cache = NgramCache::new();
cache
.update(1, 4, &tokens, i32::try_from(tokens.len()).unwrap(), false)
.unwrap();
let draft = ngram_cache_draft(&tokens, 4, 1, 4, Some(&mut cache), None, None).unwrap();
assert!(
draft.contains(&LlamaToken(3)),
"expected 3 after 1,2; got {draft:?}"
);
}
#[test]
fn map_builds_and_drafts() {
let mut map = NgramMap::new(2, 2, false, 1).expect("map");
let tokens = toks(&[1, 2, 3, 4, 1, 2]);
map.begin(&tokens).unwrap();
let _ = map.draft(&tokens, LlamaToken(2)).unwrap();
map.accept(0);
}
#[test]
fn map_begin_accepts_an_empty_prompt() {
let mut map = NgramMap::new(2, 2, false, 1).expect("map");
map.begin(&[]).unwrap();
}
}