use crate::error::{RealizarError, Result};
use crate::format::{detect_format, ModelFormat};
use std::path::PathBuf;
use std::time::Instant;
pub(crate) fn qtype_to_dtype_str(qtype: u32) -> &'static str {
crate::gguf::admitted_from_id(qtype).map_or("Unknown", crate::gguf::GgmlQuantType::as_str)
}
pub(crate) fn body_quant_label(body: &[u32], lm_head: u32) -> String {
let name = |q: u32| {
trueno_quant::GgmlType::from_id(q)
.map_or_else(|| format!("ggml type {q}"), |t| t.as_str().to_string())
};
let mut counts: std::collections::BTreeMap<u32, usize> = std::collections::BTreeMap::new();
for &q in body {
*counts.entry(q).or_insert(0) += 1;
}
let mut ranked: Vec<(String, usize, u32)> =
counts.into_iter().map(|(q, n)| (name(q), n, q)).collect();
ranked.sort_by(|a, b| b.1.cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
let Some((_, _, dominant)) = ranked.first().cloned() else {
return name(lm_head);
};
let body_label = if ranked.len() == 1 {
ranked[0].0.clone()
} else {
let parts: Vec<String> = ranked.iter().map(|(n, c, _)| format!("{n}×{c}")).collect();
format!("mixed({})", parts.join(","))
};
if lm_head == dominant {
body_label
} else {
format!("{body_label} lm_head={}", name(lm_head))
}
}
pub(crate) fn body_qtypes(gguf: &crate::gguf::GGUFModel) -> Vec<u32> {
gguf.tensors
.iter()
.filter(|t| t.name.starts_with("blk.") && t.n_dims >= 2)
.map(|t| t.qtype)
.collect()
}
pub(crate) fn model_body_qtypes(model: &crate::gguf::OwnedQuantizedModel) -> Vec<u32> {
use crate::gguf::OwnedQKVWeights;
let mut out = Vec::new();
for layer in model.layers() {
match &layer.qkv_weight {
OwnedQKVWeights::Fused(t) => out.push(t.qtype),
OwnedQKVWeights::Separate { q, k, v } => out.extend([q.qtype, k.qtype, v.qtype]),
}
out.push(layer.attn_output_weight.qtype);
out.push(layer.ffn_up_weight.qtype);
out.push(layer.ffn_down_weight.qtype);
if let Some(gate) = layer.ffn_gate_weight.as_ref() {
out.push(gate.qtype);
}
}
out
}
pub(crate) fn safetensors_dtype_ggml_id(
dtype: &crate::safetensors::SafetensorsDtype,
) -> Option<u32> {
use crate::safetensors::SafetensorsDtype as D;
match dtype {
D::F32 => Some(0),
D::F16 => Some(1),
D::BF16 => Some(30),
_ => None,
}
}
pub(crate) fn safetensors_quant_label(path: &std::path::Path) -> String {
let model = match crate::safetensors::MappedSafeTensorsModel::load(path) {
Ok(m) => m,
Err(e) => return format!("unknown (header unreadable: {e})"),
};
let mut body = Vec::new();
let mut head = None;
for name in model.tensor_names() {
let Some(info) = model.get_tensor_info(name) else {
continue;
};
let Some(id) = safetensors_dtype_ggml_id(&info.dtype) else {
continue;
};
if name.contains(".layers.") && info.shape.len() >= 2 {
body.push(id);
} else if name == "lm_head.weight"
|| (head.is_none() && name.ends_with("embed_tokens.weight"))
{
head = Some(id);
}
}
match head {
Some(h) => body_quant_label(&body, h),
None if body.is_empty() => "unknown (no float weights in the header)".to_string(),
None => {
let dominant = body_quant_label(&body, u32::MAX);
dominant
.split(" lm_head=")
.next()
.unwrap_or(&dominant)
.to_string()
},
}
}
#[derive(Debug, Clone)]
pub struct InferenceConfig {
pub model_path: PathBuf,
pub prompt: Option<String>,
pub input_tokens: Option<Vec<u32>>,
pub max_tokens: usize,
pub temperature: f32,
pub top_k: usize,
pub top_p: Option<f32>,
pub seed: u64,
pub repeat_penalty: f32,
pub repeat_last_n: usize,
pub no_gpu: bool,
pub accel_forced: bool,
pub trace: bool,
pub trace_verbose: bool,
pub trace_output: Option<PathBuf>,
pub trace_steps: Option<Vec<String>>,
pub verbose: bool,
pub stop_tokens: Vec<u32>,
#[doc(hidden)]
pub use_mock_backend: bool,
pub force_chat_template: bool,
pub thinking: Option<bool>,
}
pub const DEFAULT_TOP_K: usize = 40;
#[must_use]
pub fn sampling_top_k(temperature: f32, requested: Option<usize>) -> usize {
if temperature == 0.0 {
1
} else {
requested.unwrap_or(DEFAULT_TOP_K)
}
}
impl InferenceConfig {
#[must_use]
pub fn new(model_path: impl Into<PathBuf>) -> Self {
Self {
model_path: model_path.into(),
prompt: None,
input_tokens: None,
max_tokens: 32,
temperature: 0.0, top_k: 1,
top_p: None,
seed: 42,
repeat_penalty: 1.0,
repeat_last_n: 64,
no_gpu: false,
accel_forced: false,
trace: false,
trace_verbose: false,
trace_output: None,
trace_steps: None,
verbose: false,
stop_tokens: Vec::new(),
use_mock_backend: false,
force_chat_template: false,
thinking: None,
}
}
#[must_use]
pub fn with_prompt(mut self, prompt: impl Into<String>) -> Self {
self.prompt = Some(prompt.into());
self
}
#[must_use]
pub fn with_input_tokens(mut self, tokens: Vec<u32>) -> Self {
self.input_tokens = Some(tokens);
self
}
#[must_use]
pub fn with_max_tokens(mut self, max_tokens: usize) -> Self {
self.max_tokens = max_tokens;
self
}
#[must_use]
pub fn with_temperature(mut self, temperature: f32) -> Self {
self.temperature = temperature;
self
}
#[must_use]
pub fn with_top_k(mut self, top_k: usize) -> Self {
self.top_k = top_k;
self
}
#[must_use]
pub fn with_top_p(mut self, top_p: Option<f32>) -> Self {
self.top_p = top_p;
self
}
#[must_use]
pub fn with_seed(mut self, seed: u64) -> Self {
self.seed = seed;
self
}
#[must_use]
pub fn with_repeat_penalty(mut self, repeat_penalty: f32) -> Self {
self.repeat_penalty = repeat_penalty;
self
}
#[must_use]
pub fn with_repeat_last_n(mut self, repeat_last_n: usize) -> Self {
self.repeat_last_n = repeat_last_n;
self
}
#[must_use]
pub fn without_gpu(mut self) -> Self {
self.no_gpu = true;
self
}
#[must_use]
pub fn with_accel_forced(mut self, accel_forced: bool) -> Self {
self.accel_forced = accel_forced;
self
}
#[must_use]
pub fn with_verbose(mut self, verbose: bool) -> Self {
self.verbose = verbose;
self
}
#[must_use]
pub fn with_force_chat_template(mut self, force: bool) -> Self {
self.force_chat_template = force;
self
}
#[must_use]
pub fn with_thinking(mut self, thinking: Option<bool>) -> Self {
self.thinking = thinking;
self
}
#[must_use]
pub fn with_trace(mut self, trace: bool) -> Self {
self.trace = trace;
self
}
#[must_use]
pub fn with_trace_output(mut self, path: impl Into<PathBuf>) -> Self {
self.trace_output = Some(path.into());
self
}
#[must_use]
pub fn with_stop_tokens(mut self, stop_tokens: Vec<u32>) -> Self {
self.stop_tokens = stop_tokens;
self
}
pub(crate) fn apply_sampling_to(&self, gen_config: &mut crate::gguf::QuantizedGenerateConfig) {
gen_config.temperature = self.temperature;
gen_config.top_k = self.top_k;
gen_config.top_p = self.top_p.unwrap_or(1.0);
gen_config.seed = self.seed;
gen_config.repeat_penalty = self.repeat_penalty;
gen_config.repeat_last_n = self.repeat_last_n;
}
}
#[derive(Debug, Clone)]
pub struct PreparedTokens {
tokens: Vec<u32>,
input_count: usize,
}
impl PreparedTokens {
#[must_use]
pub fn tokens(&self) -> &[u32] {
&self.tokens
}
#[must_use]
pub fn input_count(&self) -> usize {
self.input_count
}
}
pub fn prepare_tokens(config: &InferenceConfig, format: &ModelFormat) -> Result<PreparedTokens> {
if let Some(ref tokens) = config.input_tokens {
return Ok(PreparedTokens {
input_count: tokens.len(),
tokens: tokens.clone(),
});
}
let prompt = match config.prompt {
Some(ref p) => p.clone(),
None => {
return Ok(PreparedTokens {
tokens: vec![1u32],
input_count: 1,
})
},
};
match format {
ModelFormat::Gguf => prepare_tokens_gguf(config, &prompt),
ModelFormat::SafeTensors => prepare_tokens_safetensors(config, &prompt),
ModelFormat::Apr => prepare_tokens_apr(config, &prompt),
}
}
fn thinking_mode(config: &InferenceConfig, formatted: String) -> Result<String> {
if config.thinking.is_none() {
return Ok(formatted);
}
crate::chat_template::apply_thinking_mode(&formatted, config.thinking)
}
fn sibling_tokenizer_config(model_path: &std::path::Path) -> Option<String> {
let text = std::fs::read_to_string(model_path.with_file_name("tokenizer_config.json")).ok()?;
let v: serde_json::Value = serde_json::from_str(&text).ok()?;
let declared = match v.get("chat_template")? {
serde_json::Value::String(s) => !s.is_empty(),
serde_json::Value::Array(list) => list
.iter()
.any(|t| t.get("name").and_then(serde_json::Value::as_str) == Some("default")),
_ => false,
};
declared.then_some(text)
}
fn prepare_tokens_gguf(config: &InferenceConfig, prompt: &str) -> Result<PreparedTokens> {
use crate::chat_template::{format_messages, ChatMessage};
use crate::gguf::{GGUFValue, MappedGGUFModel};
let mapped = MappedGGUFModel::from_path(&config.model_path)?;
let gguf_arch = mapped.model.architecture().unwrap_or("transformer");
let has_chat_template = mapped
.model
.metadata
.get("tokenizer.chat_template")
.is_some_and(|v| matches!(v, GGUFValue::String(s) if !s.is_empty()));
let model_name = config
.model_path
.file_name()
.and_then(|n| n.to_str())
.unwrap_or("");
let filename_instruct = model_name.to_lowercase().contains("instruct")
|| model_name.to_lowercase().contains("-chat");
let formatted_prompt = if config.force_chat_template || has_chat_template || filename_instruct {
let template_hint = apr_arch_to_template_hint(gguf_arch, model_name);
let messages = vec![ChatMessage::user(prompt)];
let own = has_chat_template.then_some(|t: Option<bool>| {
crate::chat_template::render_official_for_model(&mapped.model, &messages, t)
});
crate::chat_template::official_or_legacy(
own,
|| {
format_messages(&messages, Some(template_hint))
.unwrap_or_else(|_| prompt.to_string())
},
config.thinking,
)?
} else {
thinking_mode(config, prompt.to_string())?
};
if config.verbose {
eprintln!(
"[DEBUG] has_chat_template={}, filename_instruct={}",
has_chat_template, filename_instruct
);
eprintln!(
"[DEBUG] formatted_prompt={:?}",
log_head(&formatted_prompt, 200)
);
}
let mut tokens = mapped.model.encode(&formatted_prompt).ok_or_else(|| {
RealizarError::InferenceError(format!(
"Tokenizer encode failed for GGUF model (no tokenizer data in GGUF file?). \
Prompt length: {} chars",
formatted_prompt.len()
))
})?;
let add_bos = match mapped
.model
.metadata
.get(crate::gguf::keys::TOKENIZER_ADD_BOS)
{
Some(GGUFValue::Bool(b)) => *b,
_ => {
let arch = mapped
.model
.metadata
.get(crate::gguf::keys::GENERAL_ARCHITECTURE)
.and_then(|v| {
if let GGUFValue::String(s) = v {
Some(s.as_str())
} else {
None
}
})
.unwrap_or("unknown");
let constraints = crate::gguf::ArchConstraints::from_architecture(arch);
constraints.positional_encoding != crate::gguf::PositionalEncoding::Absolute
},
};
if add_bos {
if let Some(bos_id) = mapped.model.bos_token_id() {
if tokens.first() != Some(&bos_id) {
tokens.insert(0, bos_id);
}
}
}
if config.verbose {
eprintln!(
"[DEBUG] add_bos={}, encoded {} tokens: {:?}",
add_bos,
tokens.len(),
&tokens[..tokens.len().min(30)]
);
}
Ok(PreparedTokens {
input_count: tokens.len(),
tokens,
})
}
fn prepare_tokens_safetensors(config: &InferenceConfig, prompt: &str) -> Result<PreparedTokens> {
use crate::apr::AprV2Model;
use crate::chat_template::{format_messages, ChatMessage};
use crate::safetensors::SafetensorsConfig;
let st_config = SafetensorsConfig::load_from_sibling(&config.model_path);
let architecture = st_config
.as_ref()
.map(SafetensorsConfig::architecture)
.unwrap_or_default();
let model_name = config
.model_path
.file_name()
.and_then(|n| n.to_str())
.unwrap_or("");
let arch_lower = architecture.to_lowercase();
let is_instruct = config.force_chat_template
|| arch_lower.contains("instruct")
|| model_name.to_lowercase().contains("instruct")
|| matches!(
arch_lower.as_str(),
"qwen2forcausallm" | "llamaforcausallm" | "mistralforcausallm" | "phiforcausallm"
);
let formatted_prompt = if is_instruct {
let template_hint = safetensors_arch_to_template_hint(&architecture, model_name);
let messages = vec![ChatMessage::user(prompt)];
let tc = sibling_tokenizer_config(&config.model_path);
let msgs = &messages;
let own = tc.as_deref().map(|json| {
move |t: Option<bool>| {
crate::chat_template::render_official_from_tokenizer_config(json, msgs, t)
}
});
crate::chat_template::official_or_legacy(
own,
|| {
format_messages(&messages, Some(template_hint))
.unwrap_or_else(|_| prompt.to_string())
},
config.thinking,
)?
} else {
thinking_mode(config, prompt.to_string())?
};
let tokens =
AprV2Model::encode_text(&config.model_path, &formatted_prompt).ok_or_else(|| {
RealizarError::InferenceError(format!(
"Tokenizer encode failed for SafeTensors model (no tokenizer.json sibling?). \
Prompt length: {} chars",
formatted_prompt.len()
))
})?;
Ok(PreparedTokens {
input_count: tokens.len(),
tokens,
})
}
fn prepare_tokens_apr(config: &InferenceConfig, prompt: &str) -> Result<PreparedTokens> {
use crate::apr::AprV2Model;
use crate::chat_template::{format_messages, ChatMessage};
let model_name = config
.model_path
.file_name()
.and_then(|n| n.to_str())
.unwrap_or("");
let (apr_arch, own_template) = if config.model_path.extension().is_some_and(|e| e == "apr") {
match AprV2Model::load(&config.model_path) {
Ok(model) => {
let meta = model.metadata();
let arch = meta.architecture.clone().unwrap_or_default();
let tmpl = meta
.extra
.get("tokenizer.chat_template")
.and_then(|v| v.as_str())
.filter(|s| !s.is_empty())
.map(str::to_string);
(arch, tmpl)
},
Err(_) => (String::new(), None),
}
} else {
(String::new(), None)
};
let has_chat_template = own_template.is_some();
let filename_instruct = model_name.to_lowercase().contains("instruct")
|| model_name.to_lowercase().contains("-chat");
let is_instruct = config.force_chat_template || has_chat_template || filename_instruct;
let formatted_prompt = if is_instruct {
let template_hint = apr_arch_to_template_hint(&apr_arch, model_name);
let messages = vec![ChatMessage::user(prompt)];
let msgs = &messages;
let own = own_template.as_deref().map(|tpl| {
move |t: Option<bool>| {
crate::chat_template::render_official(tpl, None, None, msgs, true, t)
}
});
crate::chat_template::official_or_legacy(
own,
|| {
format_messages(&messages, Some(template_hint))
.unwrap_or_else(|_| prompt.to_string())
},
config.thinking,
)?
} else {
thinking_mode(config, prompt.to_string())?
};
let tokens =
AprV2Model::encode_text(&config.model_path, &formatted_prompt).ok_or_else(|| {
RealizarError::InferenceError(format!(
"Tokenizer encode failed for APR model (no tokenizer in APR metadata?). \
Prompt length: {} chars",
formatted_prompt.len()
))
})?;
Ok(PreparedTokens {
input_count: tokens.len(),
tokens,
})
}
fn safetensors_arch_to_template_hint(architecture: &str, _model_name: &str) -> &'static str {
crate::tensor_names::normalize_architecture(architecture)
}
pub(crate) fn log_head(s: &str, max: usize) -> &str {
let mut end = s.len().min(max);
while !s.is_char_boundary(end) {
end -= 1;
}
&s[..end]
}
#[cfg(test)]
mod log_head_4018 {
use super::log_head;
#[test]
fn a_non_ascii_prompt_cut_mid_char_does_not_panic() {
let p = format!(
"<|im_start|>user\nx{}<|im_end|>\n<|im_start|>assistant\n",
"\u{6c34}".repeat(80)
);
assert!(
!p.is_char_boundary(200),
"the fixture must put byte 200 mid-char"
);
let head = log_head(&p, 200);
assert!(head.len() <= 200 && head.len() >= 197, "{}", head.len());
assert!(p.starts_with(head));
}
#[test]
fn ascii_and_short_inputs_are_unchanged() {
assert_eq!(log_head("What is 2+2?", 200), "What is 2+2?");
let a = "a".repeat(300);
assert_eq!(log_head(&a, 200).len(), 200);
assert_eq!(log_head("", 200), "");
}
#[test]
fn no_log_site_byte_slices_text_any_more() {
for (f, old) in [
(
"src/infer/mod.rs",
"&formatted_prompt[..formatted_prompt.len().min(200)]",
),
(
"src/infer/inference_result.rs",
"&raw_text[..raw_text.len().min(200)]",
),
] {
let src = std::fs::read_to_string(format!("{}/{f}", env!("CARGO_MANIFEST_DIR")))
.expect("source");
let code = src.split("mod log_head_4018").next().unwrap_or(&src);
assert!(
!code.contains(old),
"{f} still slices text at a fixed byte length: {old}"
);
}
}
}
include!("inference_result.rs");
include!("gguf_gpu_generate.rs");
include!("mod_log_transformer_eos.rs");
include!("mod_05.rs");
include!("batch.rs");
pub mod qwen3_moe_dispatch;
pub mod qwen3_moe_generate;
pub mod run_report;
#[cfg(test)]
#[path = "tests_quant_label_4006.rs"]
mod tests_quant_label_4006;
#[cfg(test)]
#[path = "tests_sampling_3760.rs"]
mod tests_sampling_3760;
#[cfg(test)]
#[path = "tests_sampling_default_3754.rs"]
mod tests_sampling_default_3754;
#[cfg(test)]
#[path = "tests_dense_session_4268.rs"]
mod tests_dense_session_4268;
#[cfg(test)]
mod sibling_tokenizer_config_3990 {
use super::sibling_tokenizer_config;
fn with_config(json: Option<&str>) -> Option<String> {
let dir = tempfile::tempdir().expect("tempdir");
if let Some(j) = json {
std::fs::write(dir.path().join("tokenizer_config.json"), j).expect("write");
}
sibling_tokenizer_config(&dir.path().join("model.safetensors"))
}
#[test]
fn a_declared_template_is_found_in_both_forms_and_nothing_else_is() {
assert!(with_config(Some(r#"{"chat_template": "{{ messages }}"}"#)).is_some());
assert!(with_config(Some(
r#"{"chat_template": [{"name": "default", "template": "x"}]}"#
))
.is_some());
for j in [
None,
Some("{}"),
Some(r#"{"chat_template": ""}"#),
Some(r#"{"chat_template": [{"name": "tool_use", "template": "x"}]}"#),
Some("not json"),
] {
assert!(with_config(j).is_none(), "{j:?}");
}
}
}