use std::io::Seek;
use std::path::Path;
use candle_core::quantized::gguf_file;
use candle_core::{Device, Tensor};
use candle_transformers::generation::{LogitsProcessor, Sampling};
use candle_transformers::models::{quantized_qwen2, quantized_qwen3};
use tokenizers::Tokenizer;
use super::LocalModelError;
enum Weights {
Qwen2(quantized_qwen2::ModelWeights),
Qwen3(quantized_qwen3::ModelWeights),
}
impl Weights {
fn forward(&mut self, x: &Tensor, offset: usize) -> candle_core::Result<Tensor> {
match self {
Self::Qwen2(m) => m.forward(x, offset),
Self::Qwen3(m) => m.forward(x, offset),
}
}
fn clear_kv_cache(&mut self) {
match self {
Self::Qwen2(m) => m.clear_kv_cache(),
Self::Qwen3(m) => m.clear_kv_cache(),
}
}
}
const GGUF_FILE: &str = "model.gguf";
const TOKENIZER_FILE: &str = "tokenizer.json";
#[derive(Debug, Clone, Copy)]
pub struct GenConfig {
pub max_new_tokens: usize,
pub temperature: Option<f64>,
pub top_p: Option<f64>,
pub seed: u64,
}
impl Default for GenConfig {
fn default() -> Self {
Self {
max_new_tokens: 320,
temperature: None,
top_p: None,
seed: 299_792_458,
}
}
}
pub struct LocalGenerator {
model: Weights,
tokenizer: Tokenizer,
device: Device,
eos: Vec<u32>,
suppress_think: bool,
}
impl LocalGenerator {
pub fn load(model_dir: &Path) -> Result<Self, LocalModelError> {
let device = Device::Cpu;
let tokenizer = Tokenizer::from_file(model_dir.join(TOKENIZER_FILE))
.map_err(|e| LocalModelError::Tokenizer(e.to_string()))?;
let mut file = std::fs::File::open(model_dir.join(GGUF_FILE))?;
let content = gguf_file::Content::read(&mut file)?;
let arch = match content
.metadata
.get("general.architecture")
.and_then(|v| v.to_string().ok())
{
Some(a) => a.clone(),
None => "qwen2".to_owned(),
};
file.rewind()?;
let model = match arch.as_str() {
"qwen2" => Weights::Qwen2(quantized_qwen2::ModelWeights::from_gguf(
content, &mut file, &device,
)?),
"qwen3" => Weights::Qwen3(quantized_qwen3::ModelWeights::from_gguf(
content, &mut file, &device,
)?),
other => return Err(LocalModelError::UnsupportedArch(other.to_owned())),
};
let suppress_think = arch == "qwen3";
let eos: Vec<u32> = ["<|im_end|>", "<|endoftext|>"]
.iter()
.filter_map(|t| tokenizer.token_to_id(t))
.collect();
if eos.is_empty() {
return Err(LocalModelError::Tokenizer(
"tokenizer has no end-of-turn token (`<|im_end|>`/`<|endoftext|>`)".to_owned(),
));
}
Ok(Self {
model,
tokenizer,
device,
eos,
suppress_think,
})
}
pub fn generate(
&mut self,
system: Option<&str>,
user: &str,
cfg: &GenConfig,
) -> Result<String, LocalModelError> {
self.model.clear_kv_cache();
let prompt = chatml(system, user, self.suppress_think);
let encoding = self
.tokenizer
.encode(prompt, true)
.map_err(|e| LocalModelError::Tokenizer(e.to_string()))?;
let prompt_tokens = encoding.get_ids();
if prompt_tokens.is_empty() {
return Ok(String::new());
}
let sampling = match cfg.temperature {
None => Sampling::ArgMax,
Some(t) => match cfg.top_p {
None => Sampling::All { temperature: t },
Some(p) => Sampling::TopP { p, temperature: t },
},
};
let mut sampler = LogitsProcessor::from_sampling(cfg.seed, sampling);
let input = Tensor::new(prompt_tokens, &self.device)?.unsqueeze(0)?;
let mut next = sampler.sample(&self.model.forward(&input, 0)?.squeeze(0)?)?;
let mut generated: Vec<u32> = Vec::new();
for index_pos in (prompt_tokens.len()..).take(cfg.max_new_tokens) {
if self.eos.contains(&next) {
break;
}
generated.push(next);
let input = Tensor::new(&[next], &self.device)?.unsqueeze(0)?;
let logits = self.model.forward(&input, index_pos)?.squeeze(0)?;
next = sampler.sample(&logits)?;
}
self.tokenizer
.decode(&generated, true)
.map(|s| s.trim().to_owned())
.map_err(|e| LocalModelError::Tokenizer(e.to_string()))
}
}
fn chatml(system: Option<&str>, user: &str, suppress_think: bool) -> String {
let system = system.unwrap_or("You are a precise technical writer.");
let think = if suppress_think {
"<think>\n\n</think>\n\n"
} else {
""
};
format!(
"<|im_start|>system\n{system}<|im_end|>\n\
<|im_start|>user\n{user}<|im_end|>\n\
<|im_start|>assistant\n{think}"
)
}
#[cfg(test)]
mod tests {
use super::chatml;
#[test]
fn chatml_wraps_system_and_user() {
let p = chatml(Some("Be terse."), "Hello", false);
assert!(p.starts_with("<|im_start|>system\nBe terse.<|im_end|>"));
assert!(p.contains("<|im_start|>user\nHello<|im_end|>"));
assert!(p.ends_with("<|im_start|>assistant\n"));
}
#[test]
fn chatml_suppresses_qwen3_thinking() {
let p = chatml(None, "Hi", true);
assert!(p.ends_with("<|im_start|>assistant\n<think>\n\n</think>\n\n"));
let q = chatml(None, "Hi", false);
assert!(q.ends_with("<|im_start|>assistant\n"));
}
}