llama-cpp-4 0.6.1

llama.cpp bindings for Rust
Documentation
//! Tests for sampler creation and introspection (no model needed for most).

// These assertions check exact literals that were just written into token data,
// so bit-exact float comparison is the property under test.
#![allow(clippy::float_cmp)]

use llama_cpp_4::sampling::LlamaSampler;
use llama_cpp_4::token::data::LlamaTokenData;
use llama_cpp_4::token::data_array::LlamaTokenDataArray;
use llama_cpp_4::token::LlamaToken;

#[test]
fn test_greedy_sampler() {
    let sampler = LlamaSampler::greedy();
    assert_eq!(sampler.name(), "greedy");
}

#[test]
fn test_dist_sampler() {
    let sampler = LlamaSampler::dist(42);
    assert_eq!(sampler.name(), "dist");
    assert_eq!(sampler.get_seed(), 42);
}

#[test]
fn test_temp_sampler() {
    let sampler = LlamaSampler::temp(0.8);
    assert_eq!(sampler.name(), "temp");
}

#[test]
fn test_temp_ext_sampler() {
    let sampler = LlamaSampler::temp_ext(0.8, 0.1, 1.0);
    assert_eq!(sampler.name(), "temp-ext");
}

#[test]
fn test_top_k_sampler() {
    let sampler = LlamaSampler::top_k(40);
    assert_eq!(sampler.name(), "top-k");
}

#[test]
fn test_top_p_sampler() {
    let sampler = LlamaSampler::top_p(0.9, 1);
    assert_eq!(sampler.name(), "top-p");
}

#[test]
fn test_min_p_sampler() {
    let sampler = LlamaSampler::min_p(0.05, 1);
    assert_eq!(sampler.name(), "min-p");
}

#[test]
fn test_typical_sampler() {
    let sampler = LlamaSampler::typical(1.0, 1);
    let name = sampler.name();
    assert!(
        name.contains("typical"),
        "expected 'typical' in name, got: {name}"
    );
}

#[test]
fn test_xtc_sampler() {
    let sampler = LlamaSampler::xtc(0.5, 0.1, 1, 42);
    assert_eq!(sampler.name(), "xtc");
}

#[test]
fn test_top_n_sigma_sampler() {
    let sampler = LlamaSampler::top_n_sigma(2.0);
    assert_eq!(sampler.name(), "top-n-sigma");
}

#[test]
fn test_adaptive_p_sampler() {
    let sampler = LlamaSampler::adaptive_p(0.9, 0.95, 42);
    assert_eq!(sampler.name(), "adaptive-p");
    // get_seed may return the seed or LLAMA_DEFAULT_SEED depending on implementation
    let _ = sampler.get_seed();
}

#[test]
fn test_mirostat_sampler() {
    let sampler = LlamaSampler::mirostat(32000, 42, 5.0, 0.1, 100);
    assert_eq!(sampler.name(), "mirostat");
}

#[test]
fn test_mirostat_v2_sampler() {
    let sampler = LlamaSampler::mirostat_v2(42, 5.0, 0.1);
    assert_eq!(sampler.name(), "mirostat-v2");
}

#[test]
fn test_logit_bias_sampler() {
    let biases = vec![(LlamaToken(0), -10.0), (LlamaToken(1), 5.0)];
    let sampler = LlamaSampler::logit_bias(32000, &biases);
    assert_eq!(sampler.name(), "logit-bias");
}

/// llama.cpp `b10470` moved `n_vocab` out of `llama_sampler_data` and into the
/// penalty sampler, so `penalties` gained a leading `n_vocab`. Constructing one
/// pins the argument order — passing `penalty_last_n` where `n_vocab` belongs
/// still compiles (both are `i32`) and would only show up at sample time.
#[test]
fn test_penalties_sampler() {
    let sampler = LlamaSampler::penalties(32000, 64, 1.1, 0.0, 0.0);
    assert_eq!(sampler.name(), "penalties");
}

#[test]
fn test_penalties_simple_sampler() {
    let sampler = LlamaSampler::penalties_simple(32000, 64, 1.1);
    assert_eq!(sampler.name(), "penalties");
}

/// When every penalty is at its "disabled" value llama.cpp does not build a
/// penalty sampler at all — it substitutes an identity sampler named with a
/// `?` prefix. Pinned because it is surprising: the call succeeds, but the
/// returned sampler does nothing.
#[test]
fn test_penalties_sampler_disabled_returns_noop() {
    // penalty_last_n = 0 disables regardless of the other values.
    assert_eq!(
        LlamaSampler::penalties(32000, 0, 1.1, 0.5, 0.5).name(),
        "?penalties"
    );
    // ...as does every penalty sitting at its neutral value.
    assert_eq!(
        LlamaSampler::penalties(32000, 64, 1.0, 0.0, 0.0).name(),
        "?penalties"
    );
}

/// A penalties sampler actually applies inside a chain: token 0 is repeated in
/// the accepted history, so its logit must end up below its untouched peer.
#[test]
fn test_penalties_sampler_penalizes_repeats() {
    let n_vocab = 4;
    let mut chain =
        LlamaSampler::chain_simple([LlamaSampler::penalties(n_vocab, 64, 2.0, 0.0, 0.0)]);
    for _ in 0..3 {
        chain.accept(LlamaToken(0));
    }

    let mut data = LlamaTokenDataArray::from_iter(
        (0..n_vocab).map(|id| LlamaTokenData::new(LlamaToken(id), 1.0, 0.0)),
        false,
    );
    chain.apply(&mut data);

    let repeated = data.data[0].logit();
    let untouched = data.data[1].logit();
    assert!(
        repeated < untouched,
        "repeated token logit {repeated} should be penalized below {untouched}"
    );
}

#[test]
fn test_chain_creation() {
    let chain = LlamaSampler::chain_simple([
        LlamaSampler::top_k(40),
        LlamaSampler::top_p(0.9, 1),
        LlamaSampler::temp(0.8),
        LlamaSampler::greedy(),
    ]);
    assert_eq!(chain.name(), "chain");
    assert_eq!(chain.chain_n(), 4);
}

#[test]
fn test_chain_remove() {
    let mut chain = LlamaSampler::chain_simple([LlamaSampler::top_k(40), LlamaSampler::greedy()]);
    assert_eq!(chain.chain_n(), 2);
    let removed = chain.chain_remove(0);
    assert_eq!(removed.name(), "top-k");
    assert_eq!(chain.chain_n(), 1);
}

#[test]
fn test_sampler_clone() {
    let original = LlamaSampler::chain_simple([
        LlamaSampler::top_k(40),
        LlamaSampler::temp(0.8),
        LlamaSampler::greedy(),
    ]);
    let cloned = original.clone_sampler();
    assert_eq!(original.chain_n(), cloned.chain_n());
    assert_eq!(original.name(), cloned.name());
}

/// `copy_state_from` rewinds a sampler in place instead of allocating a new
/// one. Uses a penalties sampler because it carries observable state: the token
/// history it penalizes.
#[test]
fn test_sampler_copy_state_restores_history() {
    let n_vocab = 4;
    let make = || LlamaSampler::chain_simple([LlamaSampler::penalties(n_vocab, 64, 2.0, 0.0, 0.0)]);
    let logits = || {
        LlamaTokenDataArray::from_iter(
            (0..n_vocab).map(|id| LlamaTokenData::new(LlamaToken(id), 1.0, 0.0)),
            false,
        )
    };

    // A pristine checkpoint, and a sampler that has seen token 0 repeatedly.
    let pristine = make();
    let mut advanced = make();
    for _ in 0..3 {
        advanced.accept(LlamaToken(0));
    }

    let mut penalized = logits();
    advanced.apply(&mut penalized);
    assert!(
        penalized.data[0].logit() < penalized.data[1].logit(),
        "precondition: the advanced sampler penalizes the repeated token"
    );

    // Rewinding to the pristine state must drop that history.
    advanced.copy_state_from(&pristine);
    let mut after_restore = logits();
    advanced.apply(&mut after_restore);
    assert_eq!(
        after_restore.data[0].logit(),
        after_restore.data[1].logit(),
        "after copy_state_from the accepted history should be gone"
    );
}

#[test]
fn test_sampler_reset() {
    let mut sampler = LlamaSampler::chain_simple([LlamaSampler::greedy()]);
    sampler.reset();
    // Should not panic
    assert_eq!(sampler.chain_n(), 1);
}

#[test]
fn test_sampler_perf_data() {
    let sampler = LlamaSampler::chain_simple([LlamaSampler::greedy()]);
    let perf = sampler.perf_data();
    assert_eq!(perf.n_sample, 0);
    assert_eq!(perf.t_sample_ms, 0.0);
}

#[test]
fn test_sampler_perf_reset() {
    let mut sampler = LlamaSampler::chain_simple([LlamaSampler::greedy()]);
    sampler.perf_reset();
    let perf = sampler.perf_data();
    assert_eq!(perf.n_sample, 0);
}

#[test]
fn test_greedy_selects_max() {
    let mut data_array = LlamaTokenDataArray::new(
        vec![
            LlamaTokenData::new(LlamaToken(0), 1.0, 0.0),
            LlamaTokenData::new(LlamaToken(1), 5.0, 0.0),
            LlamaTokenData::new(LlamaToken(2), 3.0, 0.0),
        ],
        false,
    );
    data_array.apply_sampler(&mut LlamaSampler::greedy());
    assert_eq!(data_array.selected_token(), Some(LlamaToken(1)));
}

#[test]
fn test_top_k_filters() {
    let mut data_array = LlamaTokenDataArray::new(
        vec![
            LlamaTokenData::new(LlamaToken(0), 1.0, 0.0),
            LlamaTokenData::new(LlamaToken(1), 5.0, 0.0),
            LlamaTokenData::new(LlamaToken(2), 3.0, 0.0),
            LlamaTokenData::new(LlamaToken(3), 2.0, 0.0),
        ],
        false,
    );
    data_array.apply_sampler(&mut LlamaSampler::top_k(2));
    assert_eq!(data_array.data.len(), 2);
}

#[test]
fn test_temp_scales_logits() {
    let mut data_array = LlamaTokenDataArray::new(
        vec![
            LlamaTokenData::new(LlamaToken(0), 2.0, 0.0),
            LlamaTokenData::new(LlamaToken(1), 4.0, 0.0),
        ],
        false,
    );
    data_array.apply_sampler(&mut LlamaSampler::temp(0.5));
    assert_eq!(data_array.data[0].logit(), 4.0);
    assert_eq!(data_array.data[1].logit(), 8.0);
}

#[test]
fn test_chain_simple_applies_all() {
    let mut data_array = LlamaTokenDataArray::new(
        vec![
            LlamaTokenData::new(LlamaToken(0), 1.0, 0.0),
            LlamaTokenData::new(LlamaToken(1), 5.0, 0.0),
            LlamaTokenData::new(LlamaToken(2), 3.0, 0.0),
        ],
        false,
    );
    data_array.apply_sampler(&mut LlamaSampler::chain_simple([
        LlamaSampler::top_k(2),
        LlamaSampler::greedy(),
    ]));
    assert_eq!(data_array.data.len(), 2);
    assert_eq!(data_array.selected_token(), Some(LlamaToken(1)));
}

#[test]
fn test_accept_tokens() {
    let mut sampler = LlamaSampler::chain_simple([LlamaSampler::greedy()]);
    sampler.accept(LlamaToken(0));
    sampler.accept_many([LlamaToken(1), LlamaToken(2)]);
    // Should not panic
}

#[test]
fn test_with_tokens() {
    let sampler = LlamaSampler::chain_simple([LlamaSampler::greedy()])
        .with_tokens([LlamaToken(0), LlamaToken(1)]);
    assert_eq!(sampler.name(), "chain");
}