1use crate::context::params::LlamaContextParams;
4use crate::model::params::kv_overrides::KvOverrides;
5use crate::LlamaCppError;
6use std::ffi::{c_char, c_void, CStr};
7use std::fmt::{Debug, Formatter};
8use std::pin::Pin;
9use std::ptr::null;
10
11pub mod kv_overrides;
12
13#[cfg(feature = "common")]
15#[derive(Debug, Clone)]
16pub struct FitResult {
17 pub n_ctx: u32,
19}
20
21#[cfg(feature = "common")]
23#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
24pub enum FitError {
25 #[error("could not find allocations that fit available memory")]
27 Failure,
28 #[error("hard error during parameter fitting")]
30 Error,
31}
32
33#[allow(clippy::cast_possible_wrap)]
34#[allow(clippy::cast_possible_truncation)]
35const LLAMA_SPLIT_MODE_NONE: i8 = llama_cpp_sys_2::LLAMA_SPLIT_MODE_NONE as i8;
36#[allow(clippy::cast_possible_wrap)]
37#[allow(clippy::cast_possible_truncation)]
38const LLAMA_SPLIT_MODE_LAYER: i8 = llama_cpp_sys_2::LLAMA_SPLIT_MODE_LAYER as i8;
39#[allow(clippy::cast_possible_wrap)]
40#[allow(clippy::cast_possible_truncation)]
41const LLAMA_SPLIT_MODE_ROW: i8 = llama_cpp_sys_2::LLAMA_SPLIT_MODE_ROW as i8;
42#[allow(clippy::cast_possible_wrap)]
43#[allow(clippy::cast_possible_truncation)]
44const LLAMA_SPLIT_MODE_TENSOR: i8 = llama_cpp_sys_2::LLAMA_SPLIT_MODE_TENSOR as i8;
45
46#[repr(i8)]
48#[derive(Copy, Clone, Debug, PartialEq, Eq)]
49pub enum LlamaSplitMode {
50 None = LLAMA_SPLIT_MODE_NONE,
52 Layer = LLAMA_SPLIT_MODE_LAYER,
54 Row = LLAMA_SPLIT_MODE_ROW,
56 Tensor = LLAMA_SPLIT_MODE_TENSOR,
58}
59
60#[derive(Debug, Clone, Copy, PartialEq, Eq)]
62pub struct LlamaSplitModeParseError(pub i32);
63
64impl TryFrom<i32> for LlamaSplitMode {
69 type Error = LlamaSplitModeParseError;
70
71 fn try_from(value: i32) -> Result<Self, Self::Error> {
72 let i8_value = value
73 .try_into()
74 .map_err(|_| LlamaSplitModeParseError(value))?;
75 match i8_value {
76 LLAMA_SPLIT_MODE_NONE => Ok(Self::None),
77 LLAMA_SPLIT_MODE_LAYER => Ok(Self::Layer),
78 LLAMA_SPLIT_MODE_ROW => Ok(Self::Row),
79 LLAMA_SPLIT_MODE_TENSOR => Ok(Self::Tensor),
80 _ => Err(LlamaSplitModeParseError(value)),
81 }
82 }
83}
84
85impl TryFrom<u32> for LlamaSplitMode {
90 type Error = LlamaSplitModeParseError;
91
92 fn try_from(value: u32) -> Result<Self, Self::Error> {
93 let i8_value = value
94 .try_into()
95 .map_err(|_| LlamaSplitModeParseError(value.try_into().unwrap_or(i32::MAX)))?;
96 match i8_value {
97 LLAMA_SPLIT_MODE_NONE => Ok(Self::None),
98 LLAMA_SPLIT_MODE_LAYER => Ok(Self::Layer),
99 LLAMA_SPLIT_MODE_ROW => Ok(Self::Row),
100 LLAMA_SPLIT_MODE_TENSOR => Ok(Self::Tensor),
101 _ => Err(LlamaSplitModeParseError(
102 value.try_into().unwrap_or(i32::MAX),
103 )),
104 }
105 }
106}
107
108impl From<LlamaSplitMode> for i32 {
110 fn from(value: LlamaSplitMode) -> Self {
111 match value {
112 LlamaSplitMode::None => LLAMA_SPLIT_MODE_NONE.into(),
113 LlamaSplitMode::Layer => LLAMA_SPLIT_MODE_LAYER.into(),
114 LlamaSplitMode::Row => LLAMA_SPLIT_MODE_ROW.into(),
115 LlamaSplitMode::Tensor => LLAMA_SPLIT_MODE_TENSOR.into(),
116 }
117 }
118}
119
120impl From<LlamaSplitMode> for u32 {
122 fn from(value: LlamaSplitMode) -> Self {
123 match value {
124 LlamaSplitMode::None => LLAMA_SPLIT_MODE_NONE as u32,
125 LlamaSplitMode::Layer => LLAMA_SPLIT_MODE_LAYER as u32,
126 LlamaSplitMode::Row => LLAMA_SPLIT_MODE_ROW as u32,
127 LlamaSplitMode::Tensor => LLAMA_SPLIT_MODE_TENSOR as u32,
128 }
129 }
130}
131
132impl Default for LlamaSplitMode {
134 fn default() -> Self {
135 LlamaSplitMode::Layer
136 }
137}
138
139pub const LLAMA_CPP_MAX_DEVICES: usize = 16;
144
145fn load_mode_from_flags(use_mmap: bool, use_mlock: bool) -> llama_cpp_sys_2::llama_load_mode {
148 match (use_mmap, use_mlock) {
149 (false, false) => llama_cpp_sys_2::LLAMA_LOAD_MODE_NONE,
150 (true, false) => llama_cpp_sys_2::LLAMA_LOAD_MODE_MMAP,
151 (false, true) => llama_cpp_sys_2::LLAMA_LOAD_MODE_MLOCK,
152 (true, true) => llama_cpp_sys_2::LLAMA_LOAD_MODE_MMAP_MLOCK,
153 }
154}
155
156#[allow(clippy::module_name_repetitions)]
158pub struct LlamaModelParams {
159 pub(crate) params: llama_cpp_sys_2::llama_model_params,
160 kv_overrides: Vec<llama_cpp_sys_2::llama_model_kv_override>,
161 buft_overrides: Vec<llama_cpp_sys_2::llama_model_tensor_buft_override>,
162 devices: Pin<Box<[llama_cpp_sys_2::ggml_backend_dev_t; LLAMA_CPP_MAX_DEVICES]>>,
163 tensor_split: Vec<f32>,
164 progress_callback: Option<Box<dyn FnMut(f32) -> bool>>,
165}
166
167impl Debug for LlamaModelParams {
168 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
169 f.debug_struct("LlamaModelParams")
170 .field("n_gpu_layers", &self.params.n_gpu_layers)
171 .field("main_gpu", &self.params.main_gpu)
172 .field("vocab_only", &self.params.vocab_only)
173 .field("use_mmap", &self.use_mmap())
174 .field("use_mlock", &self.use_mlock())
175 .field("split_mode", &self.split_mode())
176 .field("devices", &self.devices)
177 .field("kv_overrides", &"vec of kv_overrides")
178 .finish()
179 }
180}
181
182impl LlamaModelParams {
183 #[must_use]
195 pub fn kv_overrides<'a>(&'a self) -> KvOverrides<'a> {
196 KvOverrides::new(self)
197 }
198
199 #[allow(clippy::missing_panics_doc)] pub fn append_kv_override(
222 mut self: Pin<&mut Self>,
223 key: &CStr,
224 value: kv_overrides::ParamOverrideValue,
225 ) {
226 let kv_override = self
227 .kv_overrides
228 .get_mut(0)
229 .expect("kv_overrides did not have a next allocated");
230
231 assert_eq!(kv_override.key[0], 0, "last kv_override was not empty");
232
233 for (i, &c) in key.to_bytes_with_nul().iter().enumerate() {
235 kv_override.key[i] = c_char::try_from(c).expect("invalid character in key");
236 }
237
238 kv_override.tag = value.tag();
239 kv_override.__bindgen_anon_1 = value.value();
240
241 self.params.kv_overrides = null();
243
244 self.kv_overrides
246 .push(llama_cpp_sys_2::llama_model_kv_override {
247 key: [0; 128],
248 tag: 0,
249 __bindgen_anon_1: llama_cpp_sys_2::llama_model_kv_override__bindgen_ty_1 {
250 val_i64: 0,
251 },
252 });
253
254 self.params.kv_overrides = self.kv_overrides.as_ptr();
256
257 eprintln!("saved ptr: {:?}", self.params.kv_overrides);
258 }
259}
260
261impl LlamaModelParams {
262 pub fn add_cpu_moe_override(self: Pin<&mut Self>) {
264 self.add_cpu_buft_override(c"\\.ffn_(up|down|gate)_(ch|)exps");
265 }
266
267 pub fn add_cpu_buft_override(mut self: Pin<&mut Self>, key: &CStr) {
270 let buft_override = self
271 .buft_overrides
272 .get_mut(0)
273 .expect("buft_overrides did not have a next allocated");
274
275 assert!(
276 buft_override.pattern.is_null(),
277 "last buft_override was not empty"
278 );
279
280 for &c in key.to_bytes_with_nul().iter() {
282 c_char::try_from(c).expect("invalid character in key");
283 }
284
285 buft_override.pattern = key.as_ptr();
286 buft_override.buft = unsafe { llama_cpp_sys_2::ggml_backend_cpu_buffer_type() };
287
288 self.params.tensor_buft_overrides = null();
290
291 self.buft_overrides
293 .push(llama_cpp_sys_2::llama_model_tensor_buft_override {
294 pattern: std::ptr::null(),
295 buft: std::ptr::null_mut(),
296 });
297
298 self.params.tensor_buft_overrides = self.buft_overrides.as_ptr();
300 }
301
302 #[must_use]
313 pub fn tensor_buft_override_patterns(&self) -> Vec<String> {
314 self.buft_overrides
315 .iter()
316 .filter(|o| !o.pattern.is_null())
317 .map(|o| {
318 unsafe { CStr::from_ptr(o.pattern) }
326 .to_string_lossy()
327 .into_owned()
328 })
329 .collect()
330 }
331}
332
333#[cfg(feature = "common")]
334impl LlamaModelParams {
335 pub fn fit_params(
372 mut self: Pin<&mut Self>,
373 model_path: &CStr,
374 cparams: &mut LlamaContextParams,
375 margins: &mut [usize],
376 n_ctx_min: u32,
377 log_level: llama_cpp_sys_2::ggml_log_level,
378 ) -> Result<FitResult, FitError> {
379 let max_devices = unsafe { llama_cpp_sys_2::llama_max_devices() };
380 let max_buft = unsafe { llama_cpp_sys_2::llama_max_tensor_buft_overrides() };
381
382 self.tensor_split.clear();
384 self.tensor_split.resize(max_devices, 0.0);
385
386 self.buft_overrides.clear();
388 self.buft_overrides.resize(
389 max_buft + 1,
390 llama_cpp_sys_2::llama_model_tensor_buft_override {
391 pattern: std::ptr::null(),
392 buft: std::ptr::null_mut(),
393 },
394 );
395
396 self.params.tensor_split = null::<f32>();
398 self.params.tensor_buft_overrides = null();
399
400 let status = unsafe {
401 llama_cpp_sys_2::llama_rs_fit_params(
402 model_path.as_ptr(),
403 &raw mut self.params,
404 &raw mut cparams.context_params,
405 self.tensor_split.as_mut_ptr(),
406 self.buft_overrides.as_mut_ptr(),
407 margins.as_mut_ptr(),
408 n_ctx_min,
409 log_level,
410 )
411 };
412
413 match status {
415 0 => {}
416 1 => return Err(FitError::Failure),
417 _ => return Err(FitError::Error),
418 }
419
420 self.params.tensor_split = self.tensor_split.as_ptr();
422 self.params.tensor_buft_overrides = self.buft_overrides.as_ptr();
423
424 Ok(FitResult {
425 n_ctx: cparams.context_params.n_ctx,
426 })
427 }
428}
429
430impl LlamaModelParams {
431 #[must_use]
433 pub fn n_gpu_layers(&self) -> i32 {
434 self.params.n_gpu_layers
435 }
436
437 #[must_use]
439 pub fn main_gpu(&self) -> i32 {
440 self.params.main_gpu
441 }
442
443 #[must_use]
445 pub fn vocab_only(&self) -> bool {
446 self.params.vocab_only
447 }
448
449 #[must_use]
456 pub fn use_mmap(&self) -> bool {
457 matches!(
458 self.params.load_mode,
459 llama_cpp_sys_2::LLAMA_LOAD_MODE_MMAP | llama_cpp_sys_2::LLAMA_LOAD_MODE_MMAP_MLOCK
460 )
461 }
462
463 #[must_use]
465 pub fn use_mlock(&self) -> bool {
466 matches!(
467 self.params.load_mode,
468 llama_cpp_sys_2::LLAMA_LOAD_MODE_MLOCK | llama_cpp_sys_2::LLAMA_LOAD_MODE_MMAP_MLOCK
469 )
470 }
471
472 pub fn split_mode(&self) -> Result<LlamaSplitMode, LlamaSplitModeParseError> {
477 LlamaSplitMode::try_from(self.params.split_mode)
478 }
479
480 #[must_use]
482 pub fn devices(&self) -> Vec<usize> {
483 let mut backend_devices = Vec::new();
484 for i in 0..unsafe { llama_cpp_sys_2::ggml_backend_dev_count() } {
485 let dev = unsafe { llama_cpp_sys_2::ggml_backend_dev_get(i) };
486 backend_devices.push(dev);
487 }
488 let mut devices = Vec::new();
489 for &dev in self.devices.iter() {
490 if dev.is_null() {
491 break;
492 }
493 if let Some((index, _)) = backend_devices
494 .iter()
495 .enumerate()
496 .find(|&(_i, &d)| d == dev)
497 {
498 devices.push(index);
499 }
500 }
501 devices
502 }
503
504 #[must_use]
512 pub fn with_n_gpu_layers(mut self, n_gpu_layers: u32) -> Self {
513 let n_gpu_layers = i32::try_from(n_gpu_layers).unwrap_or(i32::MAX);
516 self.params.n_gpu_layers = n_gpu_layers;
517 self
518 }
519
520 #[must_use]
524 pub fn with_main_gpu(mut self, main_gpu: i32) -> Self {
525 self.params.main_gpu = main_gpu;
526 self
527 }
528
529 #[must_use]
531 pub fn with_vocab_only(mut self, vocab_only: bool) -> Self {
532 self.params.vocab_only = vocab_only;
533 self
534 }
535
536 #[must_use]
538 pub fn with_use_mmap(mut self, use_mmap: bool) -> Self {
539 self.params.load_mode = load_mode_from_flags(use_mmap, self.use_mlock());
540 self
541 }
542
543 #[must_use]
545 pub fn with_use_mlock(mut self, use_mlock: bool) -> Self {
546 self.params.load_mode = load_mode_from_flags(self.use_mmap(), use_mlock);
547 self
548 }
549
550 #[must_use]
552 pub fn with_split_mode(mut self, split_mode: LlamaSplitMode) -> Self {
553 self.params.split_mode = split_mode.into();
554 self
555 }
556
557 pub fn with_devices(mut self, devices: &[usize]) -> Result<Self, LlamaCppError> {
568 for dev in self.devices.iter_mut() {
569 *dev = std::ptr::null_mut();
570 }
571 let max_devices = crate::max_devices().min(LLAMA_CPP_MAX_DEVICES);
573 if devices.len() > max_devices {
574 return Err(LlamaCppError::MaxDevicesExceeded(max_devices));
575 }
576 for (i, &dev) in devices.iter().enumerate() {
577 if dev >= unsafe { llama_cpp_sys_2::ggml_backend_dev_count() } {
578 return Err(LlamaCppError::BackendDeviceNotFound(dev));
579 }
580 let backend_dev = unsafe { llama_cpp_sys_2::ggml_backend_dev_get(dev) };
581 self.devices[i] = backend_dev;
582 }
583 if self.devices.is_empty() {
584 self.params.devices = std::ptr::null_mut();
585 } else {
586 self.params.devices = self.devices.as_mut_ptr();
587 }
588 Ok(self)
589 }
590
591 #[must_use]
597 pub fn with_no_alloc(mut self, no_alloc: bool) -> Self {
598 self.params.no_alloc = no_alloc;
599 if no_alloc {
600 self = self.with_use_mmap(false);
601 }
602 self
603 }
604
605 #[must_use]
609 pub fn no_alloc(&self) -> bool {
610 self.params.no_alloc
611 }
612
613 #[must_use]
616 pub fn with_progress_callback<F: FnMut(f32) -> bool + 'static>(mut self, callback: F) -> Self {
617 unsafe extern "C" fn trampoline<F: FnMut(f32) -> bool>(
618 progress: f32,
619 user_data: *mut c_void,
620 ) -> bool {
621 let callback = unsafe { &mut *user_data.cast::<F>() };
622 callback(progress)
623 }
624
625 let mut callback = Box::new(callback);
626 self.params.progress_callback_user_data =
627 std::ptr::from_mut(&mut *callback).cast::<c_void>();
628 self.params.progress_callback = Some(trampoline::<F>);
629 self.progress_callback = Some(callback);
630 self
631 }
632}
633
634impl Default for LlamaModelParams {
649 fn default() -> Self {
650 let default_params = unsafe { llama_cpp_sys_2::llama_model_default_params() };
651 LlamaModelParams {
652 params: default_params,
653 kv_overrides: vec![llama_cpp_sys_2::llama_model_kv_override {
655 key: [0; 128],
656 tag: 0,
657 __bindgen_anon_1: llama_cpp_sys_2::llama_model_kv_override__bindgen_ty_1 {
658 val_i64: 0,
659 },
660 }],
661 buft_overrides: vec![llama_cpp_sys_2::llama_model_tensor_buft_override {
662 pattern: std::ptr::null(),
663 buft: std::ptr::null_mut(),
664 }],
665 devices: Box::pin([std::ptr::null_mut(); 16]),
666 tensor_split: Vec::new(),
667 progress_callback: None,
668 }
669 }
670}
671
672#[cfg(test)]
673mod tests {
674 use super::{LlamaModelParams, LlamaSplitMode};
675 use std::pin::pin;
676
677 #[test]
678 fn tensor_buft_override_patterns_empty_by_default() {
679 assert!(LlamaModelParams::default()
681 .tensor_buft_override_patterns()
682 .is_empty());
683 }
684
685 #[test]
686 fn tensor_buft_override_patterns_reads_back_added_override() {
687 let mut params = pin!(LlamaModelParams::default());
692 params.as_mut().add_cpu_moe_override();
693 assert_eq!(
694 params.tensor_buft_override_patterns(),
695 vec!["\\.ffn_(up|down|gate)_(ch|)exps".to_owned()],
696 );
697 }
698
699 #[test]
700 fn tensor_split_mode_round_trips() {
701 assert_eq!(
702 LlamaSplitMode::try_from(llama_cpp_sys_2::LLAMA_SPLIT_MODE_TENSOR),
703 Ok(LlamaSplitMode::Tensor)
704 );
705 assert_eq!(
706 u32::from(LlamaSplitMode::Tensor),
707 llama_cpp_sys_2::LLAMA_SPLIT_MODE_TENSOR as u32
708 );
709 assert_eq!(
710 i32::from(LlamaSplitMode::Tensor),
711 llama_cpp_sys_2::LLAMA_SPLIT_MODE_TENSOR as i32
712 );
713 }
714
715 #[test]
716 fn progress_callback_round_trips_and_can_abort() {
717 use super::LlamaModelParams;
718 use std::cell::Cell;
719 use std::rc::Rc;
720
721 let calls = Rc::new(Cell::new(0_u32));
722 let counter = Rc::clone(&calls);
723 let params = LlamaModelParams::default().with_progress_callback(move |_progress| {
724 counter.set(counter.get() + 1);
725 false
726 });
727
728 assert!(params.params.progress_callback.is_some());
729 assert!(!params.params.progress_callback_user_data.is_null());
730
731 let trampoline = params.params.progress_callback.unwrap();
732 let user_data = params.params.progress_callback_user_data;
733 let first = unsafe { trampoline(0.5, user_data) };
734 let second = unsafe { trampoline(1.0, user_data) };
735
736 assert!(!first && !second, "returning false signals an abort");
737 assert_eq!(calls.get(), 2);
738 }
739}