Skip to main content

llama_cpp_bindings/
llguidance_sampler.rs

1use std::ffi::c_void;
2use std::sync::Arc;
3
4use llama_cpp_error_recorder::RecordedError;
5use llama_cpp_error_recorder::record;
6use toktrie::ApproximateTokEnv;
7
8use crate::GrammarError;
9use crate::grammar_matcher::GrammarMatcher;
10use crate::mask_outcome::MaskOutcome;
11use crate::model::LlamaModel;
12use crate::sampling::LlamaSampler;
13
14struct LlgContext {
15    grammar: GrammarMatcher,
16    tok_env: Arc<ApproximateTokEnv>,
17    grammar_kind: String,
18    grammar_data: String,
19}
20
21const unsafe extern "C" fn llg_name(
22    _smpl: *const llama_cpp_bindings_sys::llama_sampler,
23) -> *const std::os::raw::c_char {
24    c"llguidance".as_ptr()
25}
26
27unsafe extern "C" fn llg_accept(
28    smpl: *mut llama_cpp_bindings_sys::llama_sampler,
29    token: llama_cpp_bindings_sys::llama_token,
30) {
31    let ctx = unsafe { &mut *(*smpl).ctx.cast::<LlgContext>() };
32
33    if let Err(grammar_error) = ctx.grammar.consume_token(token.cast_unsigned()) {
34        record(RecordedError::new(grammar_error));
35    }
36}
37
38unsafe extern "C" fn llg_apply(
39    smpl: *mut llama_cpp_bindings_sys::llama_sampler,
40    cur_p: *mut llama_cpp_bindings_sys::llama_token_data_array,
41) {
42    let ctx = unsafe { &mut *(*smpl).ctx.cast::<LlgContext>() };
43    let cur_p = unsafe { &mut *cur_p };
44
45    let mask = match ctx.grammar.compute_mask() {
46        Ok(MaskOutcome::Constrained(mask)) => mask,
47        Ok(MaskOutcome::GrammarComplete) => return,
48        Err(grammar_error) => {
49            record(RecordedError::new(grammar_error));
50
51            return;
52        }
53    };
54
55    let data = unsafe { std::slice::from_raw_parts_mut(cur_p.data, cur_p.size) };
56    for item in data.iter_mut() {
57        if !mask.is_allowed(item.id.cast_unsigned()) {
58            item.logit = f32::NEG_INFINITY;
59        }
60    }
61}
62
63unsafe extern "C" fn llg_reset(smpl: *mut llama_cpp_bindings_sys::llama_sampler) {
64    let ctx = unsafe { &mut *(*smpl).ctx.cast::<LlgContext>() };
65
66    if let Err(grammar_error) = ctx.grammar.reset() {
67        record(RecordedError::new(grammar_error));
68    }
69}
70
71unsafe extern "C" fn llg_clone(
72    smpl: *const llama_cpp_bindings_sys::llama_sampler,
73) -> *mut llama_cpp_bindings_sys::llama_sampler {
74    let ctx = unsafe { &*(*smpl).ctx.cast::<LlgContext>() };
75    let new_ctx = Box::new(LlgContext {
76        grammar: ctx.grammar.deep_clone(),
77        tok_env: Arc::clone(&ctx.tok_env),
78        grammar_kind: ctx.grammar_kind.clone(),
79        grammar_data: ctx.grammar_data.clone(),
80    });
81    unsafe {
82        llama_cpp_bindings_sys::llama_sampler_init(
83            &raw mut LLG_SAMPLER_I,
84            Box::into_raw(new_ctx).cast::<c_void>(),
85        )
86    }
87}
88
89unsafe extern "C" fn llg_free(smpl: *mut llama_cpp_bindings_sys::llama_sampler) {
90    let ctx_ptr = unsafe { (*smpl).ctx.cast::<LlgContext>() };
91    if !ctx_ptr.is_null() {
92        drop(unsafe { Box::from_raw(ctx_ptr) });
93    }
94}
95
96static mut LLG_SAMPLER_I: llama_cpp_bindings_sys::llama_sampler_i =
97    llama_cpp_bindings_sys::llama_sampler_i {
98        name: Some(llg_name),
99        accept: Some(llg_accept),
100        apply: Some(llg_apply),
101        reset: Some(llg_reset),
102        clone: Some(llg_clone),
103        free: Some(llg_free),
104        backend_init: None,
105        backend_accept: None,
106        backend_apply: None,
107        backend_set_input: None,
108    };
109
110/// # Errors
111///
112/// Returns `GrammarError` if the parser factory, grammar, or parser cannot be created.
113pub fn create_llg_sampler(
114    model: &LlamaModel,
115    grammar_kind: &str,
116    grammar_data: &str,
117) -> Result<LlamaSampler, GrammarError> {
118    let tok_env = model.approximate_tok_env()?;
119    let tok_env_dyn: Arc<dyn toktrie::TokenizerEnv + Sync> = tok_env.clone();
120
121    let factory = llguidance::ParserFactory::new_simple(&tok_env_dyn)
122        .map_err(|factory_error| GrammarError::LlguidanceError(factory_error.to_string()))?;
123
124    let grammar = llguidance::api::TopLevelGrammar::from_tagged_str(grammar_kind, grammar_data)
125        .map_err(|parse_error| GrammarError::LlguidanceError(parse_error.to_string()))?;
126
127    let parser = factory
128        .create_parser(grammar)
129        .map_err(|parser_error| GrammarError::LlguidanceError(parser_error.to_string()))?;
130
131    let ctx = Box::new(LlgContext {
132        grammar: GrammarMatcher::new(parser),
133        tok_env,
134        grammar_kind: grammar_kind.to_string(),
135        grammar_data: grammar_data.to_string(),
136    });
137
138    let sampler = unsafe {
139        llama_cpp_bindings_sys::llama_sampler_init(
140            &raw mut LLG_SAMPLER_I,
141            Box::into_raw(ctx).cast::<c_void>(),
142        )
143    };
144
145    if sampler.is_null() {
146        Err(GrammarError::LlguidanceSamplerUnavailable)
147    } else {
148        Ok(LlamaSampler { sampler })
149    }
150}