use ik_llama_cpp_sys as sys;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum MtpOpType {
None,
Warmup,
UpdateAccepted,
DraftGen,
}
impl MtpOpType {
#[must_use]
pub fn to_raw(self) -> sys::llama_mtp_op_type {
match self {
MtpOpType::None => sys::MTP_OP_NONE,
MtpOpType::Warmup => sys::MTP_OP_WARMUP,
MtpOpType::UpdateAccepted => sys::MTP_OP_UPDATE_ACCEPTED,
MtpOpType::DraftGen => sys::MTP_OP_DRAFT_GEN,
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct MtpSpeculativeParams {
pub n_max: i32,
pub n_min: i32,
pub p_min: f32,
pub mtp_heads: i32,
pub temp: f32,
}
impl Default for MtpSpeculativeParams {
fn default() -> Self {
Self {
n_max: 1,
n_min: 0,
p_min: 0.0,
mtp_heads: 1,
temp: 0.0,
}
}
}
#[cfg(feature = "common")]
pub use driver::{MtpSpeculative, MtpStep};
#[cfg(feature = "common")]
mod driver {
use super::MtpSpeculativeParams;
use ik_llama_cpp_sys as sys;
use std::ptr::NonNull;
use crate::context::LlamaContext;
use crate::model::LlamaModel;
use crate::token::LlamaToken;
use crate::LlamaError;
#[derive(Debug, Clone)]
pub struct MtpStep {
pub tokens: Vec<LlamaToken>,
pub n_accepted: i32,
}
#[derive(Debug)]
pub struct MtpSpeculative<'ctx, 'model> {
raw: NonNull<sys::ik_llama_rs_mtp>,
ctx: &'ctx mut LlamaContext<'model>,
params: MtpSpeculativeParams,
}
impl<'ctx, 'model> MtpSpeculative<'ctx, 'model> {
pub fn new(
model: &LlamaModel,
ctx: &'ctx mut LlamaContext<'model>,
params: MtpSpeculativeParams,
) -> Result<Self, LlamaError> {
let cparams = ctx.raw_params; let raw = unsafe {
sys::ik_llama_rs_mtp_init(
model.model.as_ptr(),
ctx.context.as_ptr(),
&cparams,
params.n_max,
params.n_min,
params.p_min,
params.mtp_heads,
params.temp,
)
};
NonNull::new(raw)
.map(|raw| Self { raw, ctx, params })
.ok_or(LlamaError::MtpInit)
}
pub fn begin(&mut self, prompt: &[LlamaToken]) -> Result<usize, LlamaError> {
let toks: Vec<sys::llama_token> = prompt.iter().map(|t| t.0).collect();
let n =
unsafe { sys::ik_llama_rs_mtp_begin(self.raw.as_ptr(), toks.as_ptr(), toks.len()) };
if n < 0 {
return Err(LlamaError::MtpBegin(n as i32));
}
Ok(n as usize)
}
pub fn step(&mut self) -> Result<MtpStep, LlamaError> {
let cap = self.params.n_max.max(1) as usize + 2;
let mut out = vec![0 as sys::llama_token; cap];
let mut n_out: usize = 0;
let mut n_accepted: i32 = 0;
let st = unsafe {
sys::ik_llama_rs_mtp_step(
self.raw.as_ptr(),
out.as_mut_ptr(),
cap,
&mut n_out,
&mut n_accepted,
)
};
if st as i32 != 0 {
return Err(LlamaError::MtpStep(st as i32));
}
out.truncate(n_out);
Ok(MtpStep {
tokens: out.into_iter().map(LlamaToken).collect(),
n_accepted,
})
}
#[must_use]
pub fn params(&self) -> MtpSpeculativeParams {
self.params
}
pub fn context_mut(&mut self) -> &mut LlamaContext<'model> {
self.ctx
}
}
impl Drop for MtpSpeculative<'_, '_> {
fn drop(&mut self) {
unsafe { sys::ik_llama_rs_mtp_free(self.raw.as_ptr()) };
}
}
}