use std::ffi::{CStr, CString};
use std::os::raw::c_void;
use std::path::Path;
use std::ptr::NonNull;
use std::slice;
use llama_cpp_sys_4 as sys;
use crate::model::LlamaModel;
#[derive(Debug, thiserror::Error)]
pub enum MtmdError {
#[error("failed to create mtmd context (null return from mtmd_init_from_file)")]
ContextCreateFailed,
#[error("failed to create mtmd bitmap")]
BitmapCreateFailed,
#[error("invalid path: {0}")]
InvalidPath(#[from] std::ffi::NulError),
#[error("path is not valid UTF-8")]
PathNotUtf8,
#[error("tokenize error: code {0} (1 = bitmap count mismatch, 2 = preprocessing error)")]
TokenizeError(i32),
#[error("encode error: code {0}")]
EncodeError(i32),
#[error("chunk save error: code {0}")]
ChunkSaveFailed(i32),
#[error("failed to load an input chunk from the buffer")]
ChunkLoadFailed,
#[error("batch add error: code {0} (2 = batch full, 3 = incompatible with existing chunks)")]
BatchAddFailed(i32),
#[error("failed to create an mtmd batch")]
BatchCreateFailed,
#[error("eval error: code {0}")]
EvalError(i32),
#[error("failed to open video stream (null return from mtmd_helper_video_init)")]
VideoInitFailed,
#[error("video read error: code {0}")]
VideoReadError(i32),
}
pub type Result<T> = std::result::Result<T, MtmdError>;
pub type MtmdProgressCallback = unsafe extern "C" fn(progress: f32, user_data: *mut c_void) -> bool;
pub struct MtmdContextParams {
pub(crate) params: sys::mtmd_context_params,
}
impl std::fmt::Debug for MtmdContextParams {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MtmdContextParams")
.field("use_gpu", &self.params.use_gpu)
.field("print_timings", &self.params.print_timings)
.field("n_threads", &self.params.n_threads)
.field("warmup", &self.params.warmup)
.field("image_min_tokens", &self.params.image_min_tokens)
.field("image_max_tokens", &self.params.image_max_tokens)
.finish()
}
}
impl Default for MtmdContextParams {
fn default() -> Self {
let params = unsafe { sys::mtmd_context_params_default() };
Self { params }
}
}
impl MtmdContextParams {
#[must_use]
pub fn use_gpu(mut self, v: bool) -> Self {
self.params.use_gpu = v;
self
}
#[must_use]
pub fn print_timings(mut self, v: bool) -> Self {
self.params.print_timings = v;
self
}
#[must_use]
pub fn n_threads(mut self, n: i32) -> Self {
self.params.n_threads = n;
self
}
#[must_use]
pub fn warmup(mut self, v: bool) -> Self {
self.params.warmup = v;
self
}
#[must_use]
pub fn image_min_tokens(mut self, n: i32) -> Self {
self.params.image_min_tokens = n;
self
}
#[must_use]
pub fn image_max_tokens(mut self, n: i32) -> Self {
self.params.image_max_tokens = n;
self
}
#[must_use]
pub fn with_batch_max_tokens(mut self, n: i32) -> Self {
self.params.batch_max_tokens = n;
self
}
#[must_use]
pub fn batch_max_tokens(&self) -> i32 {
self.params.batch_max_tokens
}
#[must_use]
pub fn with_flash_attn_type(
mut self,
flash_attn_type: crate::context::params::LlamaFlashAttnType,
) -> Self {
self.params.flash_attn_type = flash_attn_type.into();
self
}
#[must_use]
pub fn flash_attn_type(&self) -> crate::context::params::LlamaFlashAttnType {
crate::context::params::LlamaFlashAttnType::from(self.params.flash_attn_type)
}
#[must_use]
pub fn with_progress_callback(
mut self,
callback: Option<MtmdProgressCallback>,
user_data: *mut c_void,
) -> Self {
self.params.progress_callback = callback;
self.params.progress_callback_user_data = user_data;
self
}
pub fn media_marker(mut self, marker: Option<&str>) -> std::result::Result<Self, MtmdError> {
match marker {
None => {
self.params.media_marker = std::ptr::null();
Ok(self)
}
Some(s) => {
let cs = CString::new(s)?;
self.params.media_marker = cs.as_ptr();
std::mem::forget(cs);
Ok(self)
}
}
}
}
pub struct MtmdContext {
ptr: NonNull<sys::mtmd_context>,
}
unsafe impl Send for MtmdContext {}
unsafe impl Sync for MtmdContext {}
impl std::fmt::Debug for MtmdContext {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MtmdContext")
.field("ptr", &self.ptr)
.finish()
}
}
impl Drop for MtmdContext {
fn drop(&mut self) {
unsafe { sys::mtmd_free(self.ptr.as_ptr()) }
}
}
impl MtmdContext {
#[must_use]
pub fn default_marker() -> &'static str {
let ptr = unsafe { sys::mtmd_default_marker() };
unsafe { CStr::from_ptr(ptr) }
.to_str()
.unwrap_or("<__media__>")
}
#[allow(clippy::needless_pass_by_value)]
pub fn init_from_file(
mmproj_path: impl AsRef<Path>,
text_model: &LlamaModel,
params: MtmdContextParams,
) -> Result<Self> {
let path = mmproj_path
.as_ref()
.to_str()
.ok_or(MtmdError::PathNotUtf8)?;
let c_path = CString::new(path)?;
let ptr = unsafe {
sys::mtmd_init_from_file(c_path.as_ptr(), text_model.model.as_ptr(), params.params)
};
let ptr = NonNull::new(ptr).ok_or(MtmdError::ContextCreateFailed)?;
Ok(Self { ptr })
}
pub fn void_logs() {
unsafe extern "C" fn noop(
_level: sys::ggml_log_level,
_text: *const ::std::os::raw::c_char,
_ud: *mut ::std::os::raw::c_void,
) {
}
unsafe { sys::mtmd_log_set(Some(noop), std::ptr::null_mut()) };
}
pub fn void_helper_logs() {
unsafe extern "C" fn noop(
_level: sys::ggml_log_level,
_text: *const ::std::os::raw::c_char,
_ud: *mut ::std::os::raw::c_void,
) {
}
unsafe { sys::mtmd_helper_log_set(Some(noop), std::ptr::null_mut()) };
}
#[must_use]
pub fn supports_vision(&self) -> bool {
unsafe { sys::mtmd_support_vision(self.ptr.as_ptr()) }
}
#[must_use]
pub fn supports_audio(&self) -> bool {
unsafe { sys::mtmd_support_audio(self.ptr.as_ptr()) }
}
#[must_use]
pub fn supports_video(&self) -> bool {
unsafe { sys::mtmd_helper_support_video(self.ptr.as_ptr()) }
}
#[must_use]
pub fn marker(&self) -> &str {
let ptr = unsafe { sys::mtmd_get_marker(self.ptr.as_ptr()) };
if ptr.is_null() {
return Self::default_marker();
}
unsafe { CStr::from_ptr(ptr) }
.to_str()
.unwrap_or_else(|_| Self::default_marker())
}
#[must_use]
pub fn audio_sample_rate(&self) -> i32 {
unsafe { sys::mtmd_get_audio_sample_rate(self.ptr.as_ptr()) }
}
#[must_use]
pub fn decode_use_non_causal(&self, chunk: &MtmdInputChunk<'_>) -> bool {
unsafe { sys::mtmd_decode_use_non_causal(self.ptr.as_ptr(), chunk.as_ptr()) }
}
#[must_use]
pub fn decode_use_mrope(&self) -> bool {
unsafe { sys::mtmd_decode_use_mrope(self.ptr.as_ptr()) }
}
pub fn tokenize(
&self,
text: &MtmdInputText<'_>,
bitmaps: &[&MtmdBitmap],
output: &mut MtmdInputChunks,
) -> Result<()> {
let mut bitmap_ptrs: Vec<*const sys::mtmd_bitmap> = bitmaps
.iter()
.map(|b| b.ptr.as_ptr().cast_const())
.collect();
let c_text = text.as_raw();
let ret = unsafe {
sys::mtmd_tokenize(
self.ptr.as_ptr(),
output.ptr.as_ptr(),
&raw const c_text,
bitmap_ptrs.as_mut_ptr(),
bitmap_ptrs.len(),
)
};
if ret != 0 {
return Err(MtmdError::TokenizeError(ret));
}
Ok(())
}
pub fn tokenize_from_parts(
&self,
parts: &[MtmdInputPart<'_>],
add_special: bool,
output: &mut MtmdInputChunks,
) -> Result<()> {
let raw_texts: Vec<sys::mtmd_input_text> = parts
.iter()
.filter_map(|part| match part {
MtmdInputPart::Text(text) => Some(text.as_raw()),
MtmdInputPart::Bitmap(_) => None,
})
.collect();
let mut next_text = 0usize;
let raw_parts: Vec<sys::mtmd_input_part> = parts
.iter()
.map(|part| match part {
MtmdInputPart::Text(_) => {
let raw = &raw_texts[next_text];
next_text += 1;
sys::mtmd_input_part {
text: std::ptr::from_ref(raw),
bitmap: std::ptr::null(),
}
}
MtmdInputPart::Bitmap(bitmap) => sys::mtmd_input_part {
text: std::ptr::null(),
bitmap: bitmap.ptr.as_ptr().cast_const(),
},
})
.collect();
let part_ptrs: Vec<*const sys::mtmd_input_part> =
raw_parts.iter().map(std::ptr::from_ref).collect();
let ret = unsafe {
sys::mtmd_tokenize_from_parts(
self.ptr.as_ptr(),
output.ptr.as_ptr(),
part_ptrs.as_ptr(),
part_ptrs.len(),
add_special,
)
};
if ret != 0 {
return Err(MtmdError::TokenizeError(ret));
}
Ok(())
}
#[must_use]
pub fn gen_audio_info(&self) -> Option<MtmdGenAudioInfo> {
let info = unsafe { sys::mtmd_gen_audio_get_info(self.ptr.as_ptr()) };
if info.type_ == sys::MTMD_GEN_AUDIO_TYPE_NONE {
return None;
}
let variant = if info.model_variant.is_null() {
None
} else {
Some(
unsafe { CStr::from_ptr(info.model_variant) }
.to_string_lossy()
.into_owned(),
)
};
Some(MtmdGenAudioInfo {
pipeline: MtmdGenAudioType::from_raw(info.type_),
sample_rate: info.sample_rate,
model_variant: variant,
})
}
#[must_use]
pub fn model_can_chat(&self, ctx: &crate::context::LlamaContext<'_>) -> bool {
unsafe { sys::mtmd_helper_model_can_chat(ctx.context.as_ptr(), self.ptr.as_ptr()) }
}
pub fn encode_chunk(&self, chunk: &MtmdInputChunk<'_>) -> Result<()> {
let ret = unsafe { sys::mtmd_encode_chunk(self.ptr.as_ptr(), chunk.ptr) };
if ret != 0 {
return Err(MtmdError::EncodeError(ret));
}
Ok(())
}
#[must_use]
pub fn output_embd(&self, n_elements: usize) -> &[f32] {
let ptr = unsafe { sys::mtmd_get_output_embd(self.ptr.as_ptr()) };
if ptr.is_null() || n_elements == 0 {
return &[];
}
unsafe { slice::from_raw_parts(ptr, n_elements) }
}
#[allow(clippy::too_many_arguments, clippy::not_unsafe_ptr_arg_deref)]
pub fn eval_chunks(
&self,
lctx: *mut sys::llama_context,
chunks: &MtmdInputChunks,
n_past: i32,
seq_id: i32,
n_batch: i32,
logits_last: bool,
new_n_past: &mut i32,
) -> Result<()> {
let ret = unsafe {
sys::mtmd_helper_eval_chunks(
self.ptr.as_ptr(),
lctx,
chunks.ptr.as_ptr(),
n_past,
seq_id,
n_batch,
logits_last,
new_n_past,
)
};
if ret != 0 {
return Err(MtmdError::EvalError(ret));
}
Ok(())
}
#[allow(clippy::too_many_arguments, clippy::not_unsafe_ptr_arg_deref)]
pub fn eval_chunk_single(
&self,
lctx: *mut sys::llama_context,
chunk: &MtmdInputChunk<'_>,
n_past: i32,
seq_id: i32,
n_batch: i32,
logits_last: bool,
new_n_past: &mut i32,
) -> Result<()> {
let ret = unsafe {
sys::mtmd_helper_eval_chunk_single(
self.ptr.as_ptr(),
lctx,
chunk.ptr,
n_past,
seq_id,
n_batch,
logits_last,
new_n_past,
)
};
if ret != 0 {
return Err(MtmdError::EvalError(ret));
}
Ok(())
}
#[allow(clippy::too_many_arguments, clippy::not_unsafe_ptr_arg_deref)]
pub fn decode_image_chunk(
&self,
lctx: *mut sys::llama_context,
chunk: &MtmdInputChunk<'_>,
encoded_embd: &[f32],
n_past: i32,
seq_id: i32,
n_batch: i32,
new_n_past: &mut i32,
) -> Result<()> {
let ret = unsafe {
sys::mtmd_helper_decode_image_chunk(
self.ptr.as_ptr(),
lctx,
chunk.ptr,
encoded_embd.as_ptr().cast_mut(),
n_past,
seq_id,
n_batch,
new_n_past,
None,
std::ptr::null_mut(),
)
};
if ret != 0 {
return Err(MtmdError::EvalError(ret));
}
Ok(())
}
#[must_use]
pub fn as_ptr(&self) -> *mut sys::mtmd_context {
self.ptr.as_ptr()
}
}
#[derive(Debug)]
pub struct MtmdInputText<'a> {
text: Vec<u8>,
text_len: usize,
add_special: bool,
parse_special: bool,
_marker: std::marker::PhantomData<&'a ()>,
}
impl<'a> MtmdInputText<'a> {
pub(crate) fn as_raw(&self) -> sys::mtmd_input_text {
sys::mtmd_input_text {
text: self.text.as_ptr().cast(),
text_len: self.text_len,
add_special: self.add_special,
parse_special: self.parse_special,
}
}
#[must_use]
pub fn new(text: &'a str, add_special: bool, parse_special: bool) -> Self {
Self::from_bytes(text.as_bytes(), add_special, parse_special)
}
#[must_use]
pub fn from_bytes(text: &'a [u8], add_special: bool, parse_special: bool) -> Self {
let text_len = text.len();
let mut buf = Vec::with_capacity(text_len + 1);
buf.extend_from_slice(text);
buf.push(0); Self {
text: buf,
text_len,
add_special,
parse_special,
_marker: std::marker::PhantomData,
}
}
pub fn try_new(
text: &'a str,
add_special: bool,
parse_special: bool,
) -> std::result::Result<Self, std::ffi::NulError> {
Ok(Self::new(text, add_special, parse_special))
}
}
pub struct MtmdBitmap {
ptr: NonNull<sys::mtmd_bitmap>,
}
unsafe impl Send for MtmdBitmap {}
unsafe impl Sync for MtmdBitmap {}
impl std::fmt::Debug for MtmdBitmap {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MtmdBitmap")
.field("nx", &self.nx())
.field("ny", &self.ny())
.field("n_bytes", &self.n_bytes())
.field("is_audio", &self.is_audio())
.finish()
}
}
impl Drop for MtmdBitmap {
fn drop(&mut self) {
unsafe { sys::mtmd_bitmap_free(self.ptr.as_ptr()) }
}
}
impl MtmdBitmap {
pub fn from_rgb(nx: u32, ny: u32, data: &[u8]) -> Result<Self> {
let ptr = unsafe { sys::mtmd_bitmap_init(nx, ny, data.as_ptr()) };
let ptr = NonNull::new(ptr).ok_or(MtmdError::BitmapCreateFailed)?;
Ok(Self { ptr })
}
pub fn from_audio(samples: &[f32]) -> Result<Self> {
let ptr = unsafe { sys::mtmd_bitmap_init_from_audio(samples.len(), samples.as_ptr()) };
let ptr = NonNull::new(ptr).ok_or(MtmdError::BitmapCreateFailed)?;
Ok(Self { ptr })
}
fn from_wrapper(wrapper: sys::mtmd_helper_bitmap_wrapper) -> Result<Self> {
if !wrapper.video_ctx.is_null() {
unsafe { sys::mtmd_helper_video_free(wrapper.video_ctx) };
}
let ptr = NonNull::new(wrapper.bitmap).ok_or(MtmdError::BitmapCreateFailed)?;
Ok(Self { ptr })
}
pub fn from_file(ctx: &MtmdContext, path: impl AsRef<Path>) -> Result<Self> {
let path = path.as_ref().to_str().ok_or(MtmdError::PathNotUtf8)?;
let c_path = CString::new(path)?;
let wrapper = unsafe {
sys::mtmd_helper_bitmap_init_from_file(
ctx.ptr.as_ptr(),
c_path.as_ptr(),
false,
sys::mtmd_helper_init_opt_default(),
)
};
Self::from_wrapper(wrapper)
}
pub fn from_buf(ctx: &MtmdContext, buf: &[u8]) -> Result<Self> {
let wrapper = unsafe {
sys::mtmd_helper_bitmap_init_from_buf(
ctx.ptr.as_ptr(),
buf.as_ptr(),
buf.len(),
false,
sys::mtmd_helper_init_opt_default(),
)
};
Self::from_wrapper(wrapper)
}
pub fn set_mergeable(&mut self, mergeable: bool) {
unsafe { sys::mtmd_bitmap_set_mergeable(self.ptr.as_ptr(), mergeable) }
}
#[must_use]
pub fn nx(&self) -> u32 {
unsafe { sys::mtmd_bitmap_get_nx(self.ptr.as_ptr()) }
}
#[must_use]
pub fn ny(&self) -> u32 {
unsafe { sys::mtmd_bitmap_get_ny(self.ptr.as_ptr()) }
}
#[must_use]
pub fn n_bytes(&self) -> usize {
unsafe { sys::mtmd_bitmap_get_n_bytes(self.ptr.as_ptr()) }
}
#[must_use]
pub fn is_audio(&self) -> bool {
unsafe { sys::mtmd_bitmap_is_audio(self.ptr.as_ptr()) }
}
#[must_use]
pub fn data(&self) -> &[u8] {
let n = self.n_bytes();
if n == 0 {
return &[];
}
let ptr = unsafe { sys::mtmd_bitmap_get_data(self.ptr.as_ptr()) };
unsafe { slice::from_raw_parts(ptr, n) }
}
#[must_use]
pub fn id(&self) -> Option<&str> {
let ptr = unsafe { sys::mtmd_bitmap_get_id(self.ptr.as_ptr()) };
if ptr.is_null() {
return None;
}
unsafe { CStr::from_ptr(ptr) }.to_str().ok()
}
pub fn set_id(&mut self, id: &str) -> std::result::Result<(), std::ffi::NulError> {
let cs = CString::new(id)?;
unsafe { sys::mtmd_bitmap_set_id(self.ptr.as_ptr(), cs.as_ptr()) };
Ok(())
}
}
extern "C" {
fn free(ptr: *mut std::os::raw::c_void);
fn strdup(s: *const std::os::raw::c_char) -> *mut std::os::raw::c_char;
}
pub struct MtmdVideoParams {
params: sys::mtmd_helper_video_init_params,
ffmpeg_bin_dir: Option<CString>,
}
impl std::fmt::Debug for MtmdVideoParams {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MtmdVideoParams")
.field("fps_target", &self.params.fps_target)
.field("timestamp_interval_ms", &self.params.timestamp_interval_ms)
.field("ffmpeg_bin_dir", &self.ffmpeg_bin_dir)
.finish()
}
}
impl Default for MtmdVideoParams {
fn default() -> Self {
let params = unsafe { sys::mtmd_helper_video_init_params_default() };
Self {
params,
ffmpeg_bin_dir: None,
}
}
}
impl MtmdVideoParams {
#[must_use]
pub fn fps_target(mut self, fps: f32) -> Self {
self.params.fps_target = fps;
self
}
#[must_use]
pub fn timestamp_interval_ms(mut self, ms: i64) -> Self {
self.params.timestamp_interval_ms = ms;
self
}
pub fn ffmpeg_bin_dir(mut self, dir: Option<&str>) -> Result<Self> {
match dir {
None => {
self.params.ffmpeg_bin_dir = std::ptr::null();
self.ffmpeg_bin_dir = None;
}
Some(d) => {
let cs = CString::new(d)?;
self.params.ffmpeg_bin_dir = cs.as_ptr();
self.ffmpeg_bin_dir = Some(cs);
}
}
Ok(self)
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct MtmdVideoInfo {
pub width: u32,
pub height: u32,
pub fps: f32,
pub n_frames: i32,
}
#[derive(Debug)]
pub enum MtmdVideoItem {
Frame(MtmdBitmap),
Text(String),
}
pub struct MtmdVideo {
ptr: NonNull<sys::mtmd_helper_video>,
}
impl std::fmt::Debug for MtmdVideo {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MtmdVideo")
.field("info", &self.info())
.finish()
}
}
impl Drop for MtmdVideo {
fn drop(&mut self) {
unsafe { sys::mtmd_helper_video_free(self.ptr.as_ptr()) }
}
}
impl MtmdVideo {
pub fn from_file(
ctx: &MtmdContext,
path: impl AsRef<Path>,
params: &MtmdVideoParams,
) -> Result<Self> {
let path = path.as_ref().to_str().ok_or(MtmdError::PathNotUtf8)?;
let c_path = CString::new(path)?;
let ptr = unsafe {
sys::mtmd_helper_video_init(ctx.ptr.as_ptr(), c_path.as_ptr(), params.params)
};
let ptr = NonNull::new(ptr).ok_or(MtmdError::VideoInitFailed)?;
Ok(Self { ptr })
}
pub fn from_buf(ctx: &MtmdContext, buf: &[u8], params: &MtmdVideoParams) -> Result<Self> {
let ptr = unsafe {
sys::mtmd_helper_video_init_from_buf(
ctx.ptr.as_ptr(),
buf.as_ptr(),
buf.len(),
params.params,
)
};
let ptr = NonNull::new(ptr).ok_or(MtmdError::VideoInitFailed)?;
Ok(Self { ptr })
}
#[must_use]
pub fn info(&self) -> MtmdVideoInfo {
let info = unsafe { sys::mtmd_helper_video_get_info(self.ptr.as_ptr()) };
MtmdVideoInfo {
width: info.width,
height: info.height,
fps: info.fps,
n_frames: info.n_frames,
}
}
pub fn read_next(&mut self) -> Result<Option<MtmdVideoItem>> {
let mut out_bitmap: *mut sys::mtmd_bitmap = std::ptr::null_mut();
let mut out_text: *mut std::os::raw::c_char = std::ptr::null_mut();
let ret = unsafe {
sys::mtmd_helper_video_read_next(
self.ptr.as_ptr(),
&raw mut out_bitmap,
&raw mut out_text,
)
};
match ret {
0 => {
if let Some(ptr) = NonNull::new(out_bitmap) {
Ok(Some(MtmdVideoItem::Frame(MtmdBitmap { ptr })))
} else if !out_text.is_null() {
let text = unsafe { CStr::from_ptr(out_text) }
.to_string_lossy()
.into_owned();
unsafe { free(out_text.cast()) };
Ok(Some(MtmdVideoItem::Text(text)))
} else {
Ok(None)
}
}
-1 => Ok(None), other => Err(MtmdError::VideoReadError(other)),
}
}
}
pub struct MtmdInputChunks {
ptr: NonNull<sys::mtmd_input_chunks>,
}
impl std::fmt::Debug for MtmdInputChunks {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MtmdInputChunks")
.field("len", &self.len())
.finish()
}
}
impl Drop for MtmdInputChunks {
fn drop(&mut self) {
unsafe { sys::mtmd_input_chunks_free(self.ptr.as_ptr()) }
}
}
impl MtmdInputChunks {
#[must_use]
pub fn new() -> Self {
let ptr = unsafe { sys::mtmd_input_chunks_init() };
let ptr = NonNull::new(ptr).expect("mtmd_input_chunks_init returned null");
Self { ptr }
}
#[must_use]
pub fn len(&self) -> usize {
unsafe { sys::mtmd_input_chunks_size(self.ptr.as_ptr()) }
}
pub fn load_chunk(buf: &[u8]) -> Result<OwnedMtmdInputChunk> {
let ptr = unsafe {
sys::mtmd_input_chunk_load(buf.as_ptr().cast::<std::os::raw::c_char>(), buf.len())
};
NonNull::new(ptr)
.map(|ptr| OwnedMtmdInputChunk { ptr })
.ok_or(MtmdError::ChunkLoadFailed)
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
#[must_use]
pub fn get(&self, idx: usize) -> Option<MtmdInputChunk<'_>> {
if idx >= self.len() {
return None;
}
let ptr = unsafe { sys::mtmd_input_chunks_get(self.ptr.as_ptr(), idx) };
if ptr.is_null() {
return None;
}
Some(MtmdInputChunk {
ptr,
_marker: std::marker::PhantomData,
})
}
pub fn iter(&self) -> impl Iterator<Item = MtmdInputChunk<'_>> {
(0..self.len()).filter_map(|i| self.get(i))
}
#[must_use]
pub fn n_tokens(&self) -> usize {
unsafe { sys::mtmd_helper_get_n_tokens(self.ptr.as_ptr()) }
}
#[must_use]
pub fn n_pos(&self) -> i32 {
unsafe { sys::mtmd_helper_get_n_pos(self.ptr.as_ptr()) }
}
}
impl Default for MtmdInputChunks {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum MtmdInputChunkType {
Text,
Image,
Audio,
}
impl From<sys::mtmd_input_chunk_type> for MtmdInputChunkType {
fn from(v: sys::mtmd_input_chunk_type) -> Self {
if v == sys::MTMD_INPUT_CHUNK_TYPE_IMAGE {
Self::Image
} else if v == sys::MTMD_INPUT_CHUNK_TYPE_AUDIO {
Self::Audio
} else {
Self::Text
}
}
}
#[derive(Debug)]
pub struct MtmdInputChunk<'chunks> {
ptr: *const sys::mtmd_input_chunk,
_marker: std::marker::PhantomData<&'chunks MtmdInputChunks>,
}
impl<'chunks> MtmdInputChunk<'chunks> {
#[must_use]
pub fn chunk_type(&self) -> MtmdInputChunkType {
let t = unsafe { sys::mtmd_input_chunk_get_type(self.ptr) };
MtmdInputChunkType::from(t)
}
#[must_use]
pub fn n_tokens(&self) -> usize {
unsafe { sys::mtmd_input_chunk_get_n_tokens(self.ptr) }
}
#[must_use]
pub fn n_pos(&self) -> i32 {
unsafe { sys::mtmd_input_chunk_get_n_pos(self.ptr) }
}
pub fn save(&self) -> Result<Vec<u8>> {
let mut needed: usize = 0;
let rc = unsafe {
sys::mtmd_input_chunk_save(self.ptr, std::ptr::null_mut(), 0, &raw mut needed)
};
if rc != 0 && needed == 0 {
return Err(MtmdError::ChunkSaveFailed(rc));
}
let mut buf = vec![0u8; needed];
let rc = unsafe {
sys::mtmd_input_chunk_save(
self.ptr,
buf.as_mut_ptr().cast::<std::os::raw::c_char>(),
buf.len(),
&raw mut needed,
)
};
if rc != 0 {
return Err(MtmdError::ChunkSaveFailed(rc));
}
buf.truncate(needed);
Ok(buf)
}
pub fn to_owned_chunk(&self) -> Result<OwnedMtmdInputChunk> {
let ptr = unsafe { sys::mtmd_input_chunk_copy(self.ptr) };
NonNull::new(ptr)
.map(|ptr| OwnedMtmdInputChunk { ptr })
.ok_or(MtmdError::ChunkLoadFailed)
}
pub fn to_placeholder(&self) -> Result<OwnedMtmdInputChunk> {
let ptr = unsafe { sys::mtmd_input_chunk_get_placeholder(self.ptr) };
NonNull::new(ptr)
.map(|ptr| OwnedMtmdInputChunk { ptr })
.ok_or(MtmdError::ChunkLoadFailed)
}
#[must_use]
pub fn text_tokens(&self) -> Option<&[i32]> {
if self.chunk_type() != MtmdInputChunkType::Text {
return None;
}
let mut n: usize = 0;
let ptr = unsafe { sys::mtmd_input_chunk_get_tokens_text(self.ptr, &raw mut n) };
if ptr.is_null() || n == 0 {
return Some(&[]);
}
Some(unsafe { slice::from_raw_parts(ptr, n) })
}
#[must_use]
pub fn image_tokens(&self) -> Option<MtmdImageTokens<'chunks>> {
match self.chunk_type() {
MtmdInputChunkType::Image | MtmdInputChunkType::Audio => {}
MtmdInputChunkType::Text => return None,
}
let ptr = unsafe { sys::mtmd_input_chunk_get_tokens_image(self.ptr) };
if ptr.is_null() {
return None;
}
Some(MtmdImageTokens {
ptr,
_marker: std::marker::PhantomData,
})
}
#[must_use]
pub fn id(&self) -> Option<&str> {
let ptr = unsafe { sys::mtmd_input_chunk_get_id(self.ptr) };
if ptr.is_null() {
return None;
}
unsafe { CStr::from_ptr(ptr) }.to_str().ok()
}
#[must_use]
pub fn as_ptr(&self) -> *const sys::mtmd_input_chunk {
self.ptr
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[repr(C)]
pub struct MtmdDecoderPos {
pub t: u32,
pub x: u32,
pub y: u32,
pub z: u32,
}
#[derive(Debug)]
pub struct MtmdImageTokens<'chunks> {
ptr: *const sys::mtmd_image_tokens,
_marker: std::marker::PhantomData<&'chunks MtmdInputChunks>,
}
impl MtmdImageTokens<'_> {
#[must_use]
pub fn n_tokens(&self) -> usize {
unsafe { sys::mtmd_image_tokens_get_n_tokens(self.ptr) }
}
#[must_use]
pub fn nx(&self) -> usize {
unsafe { sys::mtmd_image_tokens_get_nx(self.ptr) }
}
#[must_use]
pub fn ny(&self) -> usize {
unsafe { sys::mtmd_image_tokens_get_ny(self.ptr) }
}
#[must_use]
pub fn n_pos(&self) -> i32 {
unsafe { sys::mtmd_image_tokens_get_n_pos(self.ptr) }
}
#[must_use]
pub fn id(&self) -> Option<&str> {
let ptr = unsafe { sys::mtmd_image_tokens_get_id(self.ptr) };
if ptr.is_null() {
return None;
}
unsafe { CStr::from_ptr(ptr) }.to_str().ok()
}
#[must_use]
pub fn decoder_positions(&self, pos_0: i32) -> Vec<MtmdDecoderPos> {
let n = self.n_tokens();
let mut out = vec![MtmdDecoderPos::default(); n];
if n == 0 {
return out;
}
unsafe {
sys::mtmd_helper_image_get_decoder_pos(
self.ptr,
pos_0,
out.as_mut_ptr().cast::<sys::mtmd_decoder_pos>(),
);
}
out
}
}
use crate::context::LlamaContext;
impl LlamaContext<'_> {
#[must_use]
pub fn as_ptr(&self) -> *mut sys::llama_context {
self.context.as_ptr()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn decoder_pos_layout_matches_sys() {
assert_eq!(
std::mem::size_of::<MtmdDecoderPos>(),
std::mem::size_of::<sys::mtmd_decoder_pos>(),
);
assert_eq!(
std::mem::align_of::<MtmdDecoderPos>(),
std::mem::align_of::<sys::mtmd_decoder_pos>(),
);
assert_eq!(std::mem::offset_of!(MtmdDecoderPos, t), 0);
assert_eq!(std::mem::offset_of!(MtmdDecoderPos, x), 4);
assert_eq!(std::mem::offset_of!(MtmdDecoderPos, y), 8);
assert_eq!(std::mem::offset_of!(MtmdDecoderPos, z), 12);
}
#[test]
fn input_text_records_byte_length_and_nul_terminates() {
let input = MtmdInputText::new("hello", true, false);
assert_eq!(input.text_len, 5);
assert_eq!(input.text, b"hello\0");
assert!(input.add_special);
assert!(!input.parse_special);
}
#[test]
fn input_text_preserves_interior_nul() {
let input = MtmdInputText::from_bytes(b"a\0b", false, true);
assert_eq!(input.text_len, 3);
assert_eq!(input.text, b"a\0b\0");
}
#[test]
fn mmproj_caps_on_a_non_mmproj_file_reports_nothing() {
let caps = mmproj_caps("/definitely/not/a/model.gguf").expect("no NUL in path");
assert!(!caps.vision);
assert!(!caps.audio);
}
#[test]
fn mmproj_caps_rejects_interior_nul() {
assert!(mmproj_caps("a\0b").is_err());
}
#[test]
fn load_chunk_rejects_garbage() {
let err = MtmdInputChunks::load_chunk(b"not a serialized chunk").unwrap_err();
assert!(matches!(err, MtmdError::ChunkLoadFailed), "got {err:?}");
}
#[test]
fn load_chunk_rejects_empty_input() {
assert!(MtmdInputChunks::load_chunk(&[]).is_err());
}
#[test]
fn load_chunk_rejects_truncated_input() {
assert!(MtmdInputChunks::load_chunk(&[0u8; 4]).is_err());
}
#[test]
fn input_text_try_new_is_infallible() {
let input = MtmdInputText::try_new("marker \u{1} data", true, true)
.expect("try_new no longer rejects any input");
assert_eq!(input.text_len, "marker \u{1} data".len());
}
}
pub struct OwnedMtmdInputChunk {
ptr: NonNull<sys::mtmd_input_chunk>,
}
impl std::fmt::Debug for OwnedMtmdInputChunk {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OwnedMtmdInputChunk")
.field("chunk_type", &self.chunk_type())
.field("n_tokens", &self.n_tokens())
.finish()
}
}
impl Drop for OwnedMtmdInputChunk {
fn drop(&mut self) {
unsafe { sys::mtmd_input_chunk_free(self.ptr.as_ptr()) }
}
}
impl OwnedMtmdInputChunk {
#[must_use]
pub fn chunk_type(&self) -> MtmdInputChunkType {
MtmdInputChunkType::from(unsafe { sys::mtmd_input_chunk_get_type(self.ptr.as_ptr()) })
}
#[must_use]
pub fn n_tokens(&self) -> usize {
unsafe { sys::mtmd_input_chunk_get_n_tokens(self.ptr.as_ptr()) }
}
#[must_use]
pub fn n_pos(&self) -> i32 {
unsafe { sys::mtmd_input_chunk_get_n_pos(self.ptr.as_ptr()) }
}
pub fn save(&self) -> Result<Vec<u8>> {
let mut needed: usize = 0;
let rc = unsafe {
sys::mtmd_input_chunk_save(
self.ptr.as_ptr(),
std::ptr::null_mut(),
0,
&raw mut needed,
)
};
if rc != 0 && needed == 0 {
return Err(MtmdError::ChunkSaveFailed(rc));
}
let mut buf = vec![0u8; needed];
let rc = unsafe {
sys::mtmd_input_chunk_save(
self.ptr.as_ptr(),
buf.as_mut_ptr().cast::<std::os::raw::c_char>(),
buf.len(),
&raw mut needed,
)
};
if rc != 0 {
return Err(MtmdError::ChunkSaveFailed(rc));
}
buf.truncate(needed);
Ok(buf)
}
}
pub struct MtmdBatch<'ctx> {
ptr: NonNull<sys::mtmd_batch>,
_ctx: std::marker::PhantomData<&'ctx MtmdContext>,
}
impl std::fmt::Debug for MtmdBatch<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MtmdBatch").finish_non_exhaustive()
}
}
impl Drop for MtmdBatch<'_> {
fn drop(&mut self) {
unsafe { sys::mtmd_batch_free(self.ptr.as_ptr()) }
}
}
impl<'ctx> MtmdBatch<'ctx> {
pub fn new(ctx: &'ctx MtmdContext) -> Result<Self> {
let ptr = unsafe { sys::mtmd_batch_init(ctx.ptr.as_ptr()) };
NonNull::new(ptr)
.map(|ptr| Self {
ptr,
_ctx: std::marker::PhantomData,
})
.ok_or(MtmdError::BatchCreateFailed)
}
pub fn add_chunk(&mut self, chunk: &MtmdInputChunk<'_>) -> Result<()> {
let rc = unsafe { sys::mtmd_batch_add_chunk(self.ptr.as_ptr(), chunk.ptr) };
if rc == 0 {
Ok(())
} else {
Err(MtmdError::BatchAddFailed(rc))
}
}
pub fn encode(&mut self) -> Result<()> {
let rc = unsafe { sys::mtmd_batch_encode(self.ptr.as_ptr()) };
if rc == 0 {
Ok(())
} else {
Err(MtmdError::EncodeError(rc))
}
}
#[must_use]
pub fn output_embd(&self, chunk: &MtmdInputChunk<'_>, n_embd: usize) -> Option<&[f32]> {
let ptr = unsafe { sys::mtmd_batch_get_output_embd(self.ptr.as_ptr(), chunk.ptr) };
if ptr.is_null() {
return None;
}
let len = chunk.n_tokens().checked_mul(n_embd)?;
if len == 0 {
return Some(&[]);
}
Some(unsafe { slice::from_raw_parts(ptr, len) })
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum MtmdAudioOutType {
Pcm,
Wav,
}
impl MtmdAudioOutType {
fn as_raw(self) -> sys::mtmd_helper_gen_audio_outtype {
match self {
Self::Pcm => sys::MTMD_HELPER_GEN_AUDIO_OUTTYPE_PCM,
Self::Wav => sys::MTMD_HELPER_GEN_AUDIO_OUTTYPE_WAV,
}
}
}
#[derive(Debug, Clone)]
pub struct MtmdAudioRequest {
pub seq_id: i32,
pub prompt: String,
pub lang: Option<String>,
pub top_k: i32,
pub top_p: f32,
pub seed: u32,
pub out_type: MtmdAudioOutType,
}
impl MtmdAudioRequest {
#[must_use]
pub fn new(prompt: impl Into<String>) -> Self {
Self {
seq_id: 0,
prompt: prompt.into(),
lang: None,
top_k: 40,
top_p: 0.9,
seed: u32::MAX,
out_type: MtmdAudioOutType::Wav,
}
}
}
pub struct MtmdAudioGen<'ctx> {
ptr: NonNull<sys::mtmd_helper_gen_audio>,
_ctx: std::marker::PhantomData<&'ctx MtmdContext>,
}
impl std::fmt::Debug for MtmdAudioGen<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MtmdAudioGen").finish_non_exhaustive()
}
}
impl Drop for MtmdAudioGen<'_> {
fn drop(&mut self) {
unsafe { sys::mtmd_helper_gen_audio_free(self.ptr.as_ptr()) }
}
}
impl<'ctx> MtmdAudioGen<'ctx> {
pub fn new(
lctx: &mut crate::context::LlamaContext<'_>,
mctx: &'ctx MtmdContext,
) -> Result<Self> {
let ptr = unsafe {
sys::mtmd_helper_gen_audio_init(lctx.context.as_ptr(), mctx.ptr.as_ptr())
};
NonNull::new(ptr)
.map(|ptr| Self {
ptr,
_ctx: std::marker::PhantomData,
})
.ok_or(MtmdError::ContextCreateFailed)
}
pub fn reset(&mut self) {
unsafe { sys::mtmd_helper_gen_audio_reset(self.ptr.as_ptr()) }
}
pub fn set_input(
&mut self,
request: &MtmdAudioRequest,
speaker_ref: Option<&MtmdBitmap>,
) -> Result<()> {
let prompt = CString::new(request.prompt.as_str())?;
let lang = request.lang.as_deref().map(CString::new).transpose()?;
let inp = sys::mtmd_helper_gen_audio_inp {
seq_id: request.seq_id,
prompt: prompt.as_ptr(),
prompt_len: request.prompt.len(),
speaker_ref: speaker_ref.map_or(std::ptr::null_mut(), |b| b.ptr.as_ptr()),
lang: lang.as_ref().map_or(std::ptr::null(), |c| c.as_ptr()),
top_k: request.top_k,
top_p: request.top_p,
seed: request.seed,
out_type: request.out_type.as_raw(),
};
let rc = unsafe { sys::mtmd_helper_gen_audio_set_input(self.ptr.as_ptr(), &raw const inp) };
if rc == 0 {
Ok(())
} else {
Err(MtmdError::EvalError(rc))
}
}
pub fn step_prompt(&mut self, n_batch: i32) -> Result<i32> {
let rc = unsafe { sys::mtmd_helper_gen_audio_step_prompt(self.ptr.as_ptr(), n_batch) };
if rc < 0 {
return Err(MtmdError::EvalError(rc));
}
Ok(rc)
}
pub fn step_gen(
&mut self,
sampled: Option<crate::token::LlamaToken>,
h_state_in: Option<&[f32]>,
n_text_embd: usize,
) -> Result<(Option<Vec<f32>>, bool)> {
let token = sampled.map_or(sys::LLAMA_TOKEN_NULL, |t| t.0);
let in_ptr = h_state_in.map_or(std::ptr::null(), <[f32]>::as_ptr);
let mut out_ptr: *const f32 = std::ptr::null();
let mut stop = false;
let rc = unsafe {
sys::mtmd_helper_gen_audio_step_gen(
self.ptr.as_ptr(),
token,
in_ptr,
&raw mut out_ptr,
&raw mut stop,
)
};
if rc < 0 {
return Err(MtmdError::EvalError(rc));
}
let state = if out_ptr.is_null() || n_text_embd == 0 {
None
} else {
Some(unsafe { slice::from_raw_parts(out_ptr, n_text_embd) }.to_vec())
};
Ok((state, stop))
}
pub fn output(&mut self) -> Result<(i32, Vec<u8>, i64)> {
let mut sample_rate: i32 = 0;
let mut data: *const std::os::raw::c_char = std::ptr::null();
let mut data_len: usize = 0;
let mut n_samples: i64 = 0;
let rc = unsafe {
sys::mtmd_helper_gen_audio_get_output(
self.ptr.as_ptr(),
&raw mut sample_rate,
&raw mut data,
&raw mut data_len,
&raw mut n_samples,
)
};
if rc != 0 {
return Err(MtmdError::EvalError(rc));
}
let bytes = if data.is_null() || data_len == 0 {
Vec::new()
} else {
unsafe { slice::from_raw_parts(data.cast::<u8>(), data_len) }.to_vec()
};
Ok((sample_rate, bytes, n_samples))
}
}
#[derive(Debug)]
pub enum MtmdInputPart<'a> {
Text(&'a MtmdInputText<'a>),
Bitmap(&'a MtmdBitmap),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum MtmdGenAudioType {
Qwen3Tts,
PocketTts,
Unknown,
}
impl MtmdGenAudioType {
fn from_raw(raw: sys::mtmd_gen_audio_type) -> Self {
match raw {
sys::MTMD_GEN_AUDIO_TYPE_QWEN3TTS => Self::Qwen3Tts,
sys::MTMD_GEN_AUDIO_TYPE_POCKETTTS => Self::PocketTts,
_ => Self::Unknown,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct MtmdGenAudioInfo {
pub pipeline: MtmdGenAudioType,
pub sample_rate: i32,
pub model_variant: Option<String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct MtmdCaps {
pub vision: bool,
pub audio: bool,
}
pub fn mmproj_caps(path: impl AsRef<Path>) -> Result<MtmdCaps> {
let path = path.as_ref().to_str().ok_or(MtmdError::PathNotUtf8)?;
let c_path = CString::new(path)?;
let caps = unsafe { sys::mtmd_get_cap_from_file(c_path.as_ptr()) };
Ok(MtmdCaps {
vision: caps.inp_vision,
audio: caps.inp_audio,
})
}
#[derive(Debug)]
pub enum MtmdLazyChunk {
Bitmap(MtmdBitmap),
Text(String),
End,
}
pub struct MtmdLazyBitmap {
bitmap: MtmdBitmap,
_callback: Box<LazyCallbackState>,
}
impl std::fmt::Debug for MtmdLazyBitmap {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MtmdLazyBitmap").finish_non_exhaustive()
}
}
struct LazyCallbackState {
func: Box<dyn FnMut(usize) -> MtmdLazyChunk>,
}
impl MtmdLazyBitmap {
pub fn new<F>(ctx: &MtmdContext, id: &str, callback: F) -> Result<Self>
where
F: FnMut(usize) -> MtmdLazyChunk + 'static,
{
let c_id = CString::new(id)?;
let mut state = Box::new(LazyCallbackState {
func: Box::new(callback),
});
let user_data = std::ptr::from_mut(state.as_mut()).cast::<std::os::raw::c_void>();
let ptr = unsafe {
sys::mtmd_bitmap_init_lazy(
ctx.ptr.as_ptr(),
c_id.as_ptr(),
user_data,
Some(lazy_trampoline),
)
};
let bitmap = MtmdBitmap {
ptr: NonNull::new(ptr).ok_or(MtmdError::BitmapCreateFailed)?,
};
Ok(Self {
bitmap,
_callback: state,
})
}
#[must_use]
pub fn as_bitmap(&self) -> &MtmdBitmap {
&self.bitmap
}
}
extern "C" fn lazy_trampoline(
chunk_idx: usize,
user_data: *mut std::os::raw::c_void,
out_bitmap: *mut *mut sys::mtmd_bitmap,
out_text: *mut *mut std::os::raw::c_char,
) -> std::os::raw::c_int {
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
if user_data.is_null() {
return -2;
}
let state = unsafe { &mut *user_data.cast::<LazyCallbackState>() };
match (state.func)(chunk_idx) {
MtmdLazyChunk::Bitmap(bitmap) => {
let raw = bitmap.ptr.as_ptr();
std::mem::forget(bitmap);
unsafe { *out_bitmap = raw };
0
}
MtmdLazyChunk::Text(text) => {
let Ok(c_text) = CString::new(text) else {
return -2;
};
let dup = unsafe { strdup(c_text.as_ptr()) };
if dup.is_null() {
return -2;
}
unsafe { *out_text = dup };
0
}
MtmdLazyChunk::End => -1,
}
}));
result.unwrap_or(-2)
}