use crate::backend;
#[cfg(feature = "metal-gpu")]
pub(crate) fn load_q4_config(
dir: &std::path::Path,
) -> Result<lattice_inference::model::qwen35_config::Qwen35Config, String> {
lattice_inference::model::qwen35_config::Qwen35Config::from_model_dir(dir)
.map_err(|e| format!("config.json load failed: {e}"))
}
#[cfg(feature = "metal-gpu")]
struct MetalChatBackend {
state: lattice_inference::forward::metal_qwen35::MetalQwen35State,
tokenizer: lattice_inference::tokenizer::bpe::BpeTokenizer,
}
#[cfg(feature = "metal-gpu")]
impl MetalChatBackend {
const MAX_CACHE_LEN: usize = 4096;
fn load(
dir: &std::path::Path,
tokenizer_dir: Option<&std::path::Path>,
) -> Result<Self, String> {
let tokenizer_path = tokenizer_dir.unwrap_or(dir).join("tokenizer.json");
let tokenizer =
lattice_inference::tokenizer::bpe::BpeTokenizer::from_tokenizer_json(&tokenizer_path)
.map_err(|e| format!("tokenizer load failed ({}): {e}", tokenizer_path.display()))?;
let cfg = load_q4_config(dir)?;
let state = lattice_inference::forward::metal_qwen35::MetalQwen35State::from_q4_dir(
dir,
&tokenizer_path,
&cfg,
Self::MAX_CACHE_LEN,
)
.map_err(|e| format!("Q4 model load failed: {e}"))?;
Ok(Self { state, tokenizer })
}
fn generate(
&mut self,
prompt: &str,
gen_cfg: &lattice_inference::model::qwen35_config::GenerateConfig,
) -> Result<
lattice_inference::model::qwen35_config::GenerateOutput,
lattice_inference::error::InferenceError,
> {
self.state.generate(prompt, &self.tokenizer, gen_cfg)
}
}
#[cfg(feature = "metal-gpu")]
pub(crate) fn chat_max_cache_len() -> usize {
MetalChatBackend::MAX_CACHE_LEN
}
pub(crate) fn run_chat(
model_path: &str,
max_tokens: usize,
temperature: f32,
tokenizer_dir: Option<&str>,
) {
use std::io::{BufRead, Write};
use std::path::Path;
let path = Path::new(model_path);
let format = backend::detect_format(path);
#[cfg(feature = "metal-gpu")]
let tokenizer_dir_path = tokenizer_dir.map(Path::new);
#[cfg(not(feature = "metal-gpu"))]
let _ = tokenizer_dir;
eprintln!("Loading model from {model_path}...");
enum Backend {
Cpu(Box<lattice_inference::model::qwen35::Qwen35Model>),
#[cfg(feature = "metal-gpu")]
Metal(Box<MetalChatBackend>),
}
let mut model = match format {
backend::ModelFormat::Safetensors => {
match lattice_inference::model::qwen35::Qwen35Model::from_safetensors(path) {
Ok(m) => Backend::Cpu(Box::new(m)),
Err(e) => {
eprintln!("Error: failed to load model: {e}");
std::process::exit(1);
}
}
}
backend::ModelFormat::Q4 => {
#[cfg(feature = "metal-gpu")]
{
match MetalChatBackend::load(path, tokenizer_dir_path) {
Ok(m) => Backend::Metal(Box::new(m)),
Err(e) => {
eprintln!("Error: failed to load Q4 model: {e}");
std::process::exit(1);
}
}
}
#[cfg(not(feature = "metal-gpu"))]
{
eprintln!("Error: {}", backend::metal_gpu_required_message(path));
std::process::exit(1);
}
}
backend::ModelFormat::Unknown => {
eprintln!("Error: {}", backend::unrecognized_format_message(path));
std::process::exit(1);
}
_ => {
eprintln!("Error: {}", backend::unrecognized_format_message(path));
std::process::exit(1);
}
};
eprintln!("Model loaded. Type 'exit' or 'quit' to stop.\n");
let gen_cfg = lattice_inference::model::qwen35_config::GenerateConfig {
max_new_tokens: max_tokens,
temperature,
..Default::default()
};
let stdin = std::io::stdin();
let mut stdout = std::io::stdout();
for line in stdin.lock().lines() {
let prompt = match line {
Ok(l) => l,
Err(e) => {
eprintln!("Error reading input: {e}");
break;
}
};
let trimmed = prompt.trim();
if trimmed.is_empty() {
continue;
}
if trimmed.eq_ignore_ascii_case("exit") || trimmed.eq_ignore_ascii_case("quit") {
break;
}
match &mut model {
Backend::Cpu(m) => match m.generate(trimmed, &gen_cfg) {
Ok(output) => {
let _ = writeln!(stdout, "{}", output.text);
let _ = writeln!(
stdout,
"[{} prompt tokens, {} generated]",
output.prompt_tokens, output.generated_tokens
);
}
Err(e) => {
eprintln!("Generation error: {e}");
}
},
#[cfg(feature = "metal-gpu")]
Backend::Metal(m) => match m.generate(trimmed, &gen_cfg) {
Ok(output) => {
let _ = writeln!(stdout, "{}", output.text);
let _ = writeln!(
stdout,
"[{} prompt tokens, {} generated]",
output.prompt_tokens, output.generated_tokens
);
}
Err(e) => {
eprintln!("Generation error: {e}");
}
},
}
}
}