use crate::InferenceError;
#[cfg(car_mlxlm_swift_built)]
use std::ffi::CString;
use std::path::Path;
#[cfg(car_mlxlm_swift_built)]
mod ffi {
use std::ffi::c_char;
extern "C" {
pub fn car_mlxlm_load(path: *const c_char) -> i32;
pub fn car_mlxlm_free(handle: i32);
pub fn car_mlxlm_clear_cache(handle: i32) -> i32;
pub fn car_mlxlm_forward(
handle: i32,
tokens: *const i32,
count: i32,
pos: i32,
out: *mut f32,
out_cap: i32,
) -> i32;
pub fn car_mlxlm_begin_prompt(handle: i32, tokens: *const i32, count: i32) -> i32;
pub fn car_mlxlm_last_error(buf: *mut c_char, cap: i32) -> i32;
pub fn car_mlxlm_supports(model_type: *const c_char) -> i32;
pub fn car_mlxlm_embed_load(path: *const c_char) -> i32;
pub fn car_mlxlm_embed_free(handle: i32);
pub fn car_mlxlm_embed(
handle: i32,
tokens: *const i32,
count: i32,
out: *mut f32,
out_cap: i32,
) -> i32;
}
}
pub fn is_available() -> bool {
cfg!(car_mlxlm_swift_built)
}
pub fn supports_model_type(model_type: &str) -> bool {
#[cfg(car_mlxlm_swift_built)]
{
let Ok(c) = CString::new(model_type.to_ascii_lowercase()) else {
return false;
};
unsafe { ffi::car_mlxlm_supports(c.as_ptr()) == 1 }
}
#[cfg(not(car_mlxlm_swift_built))]
{
let _ = model_type;
false
}
}
#[cfg(car_mlxlm_swift_built)]
fn last_error() -> String {
let mut buf = vec![0i8; 512];
let n = unsafe { ffi::car_mlxlm_last_error(buf.as_mut_ptr(), buf.len() as i32) };
if n <= 0 {
return "unknown error".to_string();
}
let bytes: Vec<u8> = buf[..n as usize].iter().map(|c| *c as u8).collect();
String::from_utf8_lossy(&bytes).into_owned()
}
pub struct SwiftLmBackend {
#[cfg(car_mlxlm_swift_built)]
handle: i32,
#[cfg(car_mlxlm_swift_built)]
vocab_size: usize,
tokenizer: tokenizers::Tokenizer,
context_length: usize,
eos: Vec<u32>,
chat_template: Option<crate::tasks::chat_template::ChatTemplate>,
model_type: String,
#[cfg(car_mlxlm_swift_built)]
model_dir: std::path::PathBuf,
role: BackendRole,
#[cfg(car_mlxlm_swift_built)]
embed_handle: std::cell::Cell<i32>,
}
fn derive_eos_ids(
generation_config: Option<&serde_json::Value>,
config: Option<&serde_json::Value>,
tokenizer_config: Option<&serde_json::Value>,
lookup: impl Fn(&str) -> Option<u32>,
) -> Vec<u32> {
use serde_json::Value;
let mut eos: Vec<u32> = Vec::new();
let push = |id: u32, eos: &mut Vec<u32>| {
if !eos.contains(&id) {
eos.push(id);
}
};
for source in [generation_config, config].into_iter().flatten() {
match source.get("eos_token_id") {
Some(Value::Number(n)) => {
if let Some(id) = n.as_u64().and_then(|v| u32::try_from(v).ok()) {
push(id, &mut eos);
}
}
Some(Value::Array(items)) => {
for id in items
.iter()
.filter_map(Value::as_u64)
.filter_map(|v| u32::try_from(v).ok())
{
push(id, &mut eos);
}
}
_ => {}
}
}
let named = tokenizer_config
.and_then(|c| c.get("eos_token"))
.and_then(|v| match v {
Value::String(s) => Some(s.as_str()),
Value::Object(o) => o.get("content").and_then(Value::as_str),
_ => None,
});
if let Some(id) = named.and_then(&lookup) {
push(id, &mut eos);
}
for token in [
"<|endoftext|>",
"<|im_end|>",
"</s>",
"<end_of_turn>",
"<|eot_id|>",
"<|end|>",
"<|return|>",
] {
if let Some(id) = lookup(token) {
push(id, &mut eos);
}
}
eos
}
#[cfg(test)]
mod eos_derivation_tests {
use super::derive_eos_ids;
use serde_json::json;
fn vocab(known: &'static [(&'static str, u32)]) -> impl Fn(&str) -> Option<u32> {
move |t: &str| known.iter().find(|(name, _)| *name == t).map(|(_, id)| *id)
}
#[test]
fn non_qwen_architectures_get_a_stop_token() {
let gemma = derive_eos_ids(
Some(&json!({ "eos_token_id": [1, 106] })),
Some(&json!({ "eos_token_id": 1 })),
Some(&json!({ "eos_token": { "content": "<end_of_turn>" } })),
vocab(&[("<end_of_turn>", 106), ("<eos>", 1)]),
);
assert!(gemma.contains(&106), "gemma turn ender missing: {gemma:?}");
assert!(gemma.contains(&1), "gemma sequence eos missing: {gemma:?}");
let llama = derive_eos_ids(
Some(&json!({ "eos_token_id": 128009 })),
Some(&json!({})),
Some(&json!({ "eos_token": "<|eot_id|>" })),
vocab(&[("<|eot_id|>", 128009)]),
);
assert_eq!(llama, vec![128009], "llama: {llama:?}");
let mistral = derive_eos_ids(None, Some(&json!({})), None, vocab(&[("</s>", 2)]));
assert_eq!(mistral, vec![2], "mistral: {mistral:?}");
}
#[test]
fn qwen_derivation_is_unchanged() {
let qwen = derive_eos_ids(
Some(&json!({ "eos_token_id": 151645 })),
Some(&json!({ "eos_token_id": 151643 })),
Some(&json!({ "eos_token": "<|im_end|>" })),
vocab(&[("<|endoftext|>", 151643), ("<|im_end|>", 151645)]),
);
assert!(qwen.contains(&151643) && qwen.contains(&151645), "{qwen:?}");
}
#[test]
fn sources_union_without_duplicates() {
let ids = derive_eos_ids(
Some(&json!({ "eos_token_id": [7, 7] })),
Some(&json!({ "eos_token_id": 7 })),
Some(&json!({ "eos_token": "</s>" })),
vocab(&[("</s>", 7)]),
);
assert_eq!(ids, vec![7], "duplicates leaked: {ids:?}");
}
#[test]
fn a_turn_ender_is_kept_alongside_the_sequence_eos() {
let ids = derive_eos_ids(
Some(&json!({ "eos_token_id": 7 })),
None,
None,
vocab(&[("</s>", 7), ("<end_of_turn>", 9)]),
);
assert!(ids.contains(&7) && ids.contains(&9), "{ids:?}");
}
#[test]
fn nothing_declared_and_nothing_known_yields_empty() {
let ids = derive_eos_ids(None, None, None, vocab(&[("some_unrelated_token", 3)]));
assert!(ids.is_empty(), "expected no stop tokens, got {ids:?}");
}
#[test]
fn malformed_declarations_are_ignored() {
let ids = derive_eos_ids(
Some(&json!({ "eos_token_id": "not-a-number" })),
Some(&json!({ "eos_token_id": [ -1, 4.5, null ] })),
Some(&json!({ "eos_token": { "no_content_field": true } })),
vocab(&[]),
);
assert!(ids.is_empty(), "{ids:?}");
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BackendRole {
Text,
Embedding,
}
use crate::resource_policy::NESTED_TEXT_CONFIGS;
fn text_config_u64(config: &serde_json::Value, key: &str) -> Option<u64> {
config[key].as_u64().or_else(|| {
NESTED_TEXT_CONFIGS
.iter()
.find_map(|nested| config.get(*nested)?.get(key)?.as_u64())
})
}
fn config_vocab_size(config: &serde_json::Value) -> usize {
text_config_u64(config, "vocab_size").unwrap_or(0) as usize
}
#[cfg(test)]
mod vocab_size_tests {
use super::config_vocab_size;
use serde_json::json;
#[test]
fn a_multimodal_checkpoint_declares_its_vocabulary_under_text_config() {
let gemma4 = json!({
"model_type": "gemma4_unified",
"text_config": {"model_type": "gemma4_unified_text", "vocab_size": 262144},
"vision_config": {},
});
assert_eq!(config_vocab_size(&gemma4), 262144);
assert_eq!(config_vocab_size(&json!({"vocab_size": 151936})), 151936);
assert_eq!(
config_vocab_size(&json!({"vocab_size": 7, "text_config": {"vocab_size": 9}})),
7
);
assert_eq!(config_vocab_size(&json!({})), 0);
assert_eq!(
config_vocab_size(&json!({"llm_config": {"vocab_size": 92553}})),
92553
);
assert_eq!(
config_vocab_size(&json!({"language_config": {"vocab_size": 102400}})),
102400
);
}
}
impl SwiftLmBackend {
pub fn load(model_dir: &Path) -> Result<Self, InferenceError> {
Self::load_with_role(model_dir, BackendRole::Text)
}
pub fn load_with_role(model_dir: &Path, role: BackendRole) -> Result<Self, InferenceError> {
let tokenizer = tokenizers::Tokenizer::from_file(model_dir.join("tokenizer.json"))
.map_err(|e| InferenceError::TokenizationError(format!("tokenizer.json: {e}")))?;
let config: serde_json::Value = serde_json::from_str(
&std::fs::read_to_string(model_dir.join("config.json"))
.map_err(|e| InferenceError::InferenceFailed(format!("read config.json: {e}")))?,
)
.map_err(|e| InferenceError::InferenceFailed(format!("parse config.json: {e}")))?;
let vocab_size = config_vocab_size(&config);
if vocab_size == 0 {
return Err(InferenceError::InferenceFailed(
"config.json has no vocab_size".to_string(),
));
}
let context_length = ["max_position_embeddings", "seq_length", "n_positions"]
.iter()
.find_map(|k| text_config_u64(&config, k))
.map(|v| v as usize)
.unwrap_or_else(|| {
tracing::warn!(
dir = %model_dir.display(),
"config.json declares no context window under any known key; \
assuming 32768, which may truncate prompts or overrun the model"
);
32768
});
let model_type = config["model_type"]
.as_str()
.unwrap_or_default()
.to_ascii_lowercase();
let chat_template = crate::tasks::chat_template::ChatTemplate::load(model_dir)?;
let read_json = |name: &str| -> Option<serde_json::Value> {
serde_json::from_str(&std::fs::read_to_string(model_dir.join(name)).ok()?).ok()
};
let eos = derive_eos_ids(
read_json("generation_config.json").as_ref(),
Some(&config),
read_json("tokenizer_config.json").as_ref(),
|t| tokenizer.token_to_id(t),
);
if eos.is_empty() {
tracing::warn!(
model_type = %model_type,
dir = %model_dir.display(),
"no EOS token could be derived; generation will stop only at max_tokens"
);
}
Self::load_inner(
model_dir,
role,
vocab_size,
tokenizer,
context_length,
eos,
chat_template,
model_type,
)
}
#[allow(clippy::too_many_arguments)]
fn load_inner(
model_dir: &Path,
role: BackendRole,
vocab_size: usize,
tokenizer: tokenizers::Tokenizer,
context_length: usize,
eos: Vec<u32>,
chat_template: Option<crate::tasks::chat_template::ChatTemplate>,
model_type: String,
) -> Result<Self, InferenceError> {
#[cfg(not(car_mlxlm_swift_built))]
{
let _ = (
model_dir,
role,
vocab_size,
tokenizer,
context_length,
eos,
chat_template,
model_type,
);
Err(InferenceError::InferenceFailed(
"this build has no mlx-swift-lm backend (needs a full Swift toolchain and \
CAR_BUILD_SWIFT_LM at build time)"
.to_string(),
))
}
#[cfg(car_mlxlm_swift_built)]
{
let path = CString::new(model_dir.to_string_lossy().as_bytes()).map_err(|e| {
InferenceError::InferenceFailed(format!("model path is not a C string: {e}"))
})?;
let (handle, embed_handle) = match role {
BackendRole::Text => {
let h = unsafe { ffi::car_mlxlm_load(path.as_ptr()) };
if h < 0 {
return Err(InferenceError::InferenceFailed(format!(
"mlx-swift-lm failed to load {}: {}",
model_dir.display(),
last_error()
)));
}
(h, -1)
}
BackendRole::Embedding => {
let h = unsafe { ffi::car_mlxlm_embed_load(path.as_ptr()) };
if h < 0 {
return Err(InferenceError::InferenceFailed(format!(
"mlx-swift-lm could not load {} as an embedding model: {}",
model_dir.display(),
last_error()
)));
}
(-1, h)
}
};
Ok(Self {
handle,
vocab_size,
tokenizer,
context_length,
eos,
chat_template,
model_type,
model_dir: model_dir.to_path_buf(),
embed_handle: std::cell::Cell::new(embed_handle),
role,
})
}
}
pub fn serves_text(&self) -> bool {
self.role == BackendRole::Text
}
fn require_text(&self) -> Result<(), InferenceError> {
if self.role != BackendRole::Text {
return Err(InferenceError::UnsupportedMode {
mode: "text-generation",
backend: "mlx-swift-lm-embedder",
reason: "this backend was loaded for embeddings only, so it holds no \
language-model weights; route generation to a text model",
});
}
Ok(())
}
pub fn forward(&mut self, tokens: &[u32], pos: usize) -> Result<Vec<f32>, InferenceError> {
self.require_text()?;
#[cfg(not(car_mlxlm_swift_built))]
{
let _ = (tokens, pos);
Err(InferenceError::InferenceFailed(
"this build has no mlx-swift-lm backend".to_string(),
))
}
#[cfg(car_mlxlm_swift_built)]
{
const PREFILL_STEP: usize = 512;
if tokens.len() > PREFILL_STEP {
let mut logits = Vec::new();
for (index, chunk) in tokens.chunks(PREFILL_STEP).enumerate() {
logits = self.forward(chunk, pos + index * PREFILL_STEP)?;
}
return Ok(logits);
}
let ids: Vec<i32> = tokens.iter().map(|t| *t as i32).collect();
let mut out = vec![0f32; self.vocab_size];
let n = unsafe {
ffi::car_mlxlm_forward(
self.handle,
ids.as_ptr(),
ids.len() as i32,
pos as i32,
out.as_mut_ptr(),
out.len() as i32,
)
};
if n < 0 {
return Err(InferenceError::InferenceFailed(format!(
"mlx-swift-lm forward failed: {}",
last_error()
)));
}
out.truncate(n as usize);
Ok(out)
}
}
pub fn begin_prompt(&mut self, tokens: &[u32]) -> usize {
if self.role != BackendRole::Text {
return 0;
}
#[cfg(not(car_mlxlm_swift_built))]
{
let _ = tokens;
0
}
#[cfg(car_mlxlm_swift_built)]
{
let ids: Vec<i32> = tokens.iter().map(|t| *t as i32).collect();
let n =
unsafe { ffi::car_mlxlm_begin_prompt(self.handle, ids.as_ptr(), ids.len() as i32) };
n.max(0) as usize
}
}
pub fn tokenize_raw(&self, text: &str) -> Result<Vec<u32>, InferenceError> {
self.tokenizer
.encode(text, false)
.map(|e| e.get_ids().to_vec())
.map_err(|e| InferenceError::TokenizationError(e.to_string()))
}
pub fn detokenize_raw(&self, tokens: &[u32]) -> Result<String, InferenceError> {
self.tokenizer
.decode(tokens, false)
.map_err(|e| InferenceError::TokenizationError(e.to_string()))
}
pub fn encode(&self, text: &str) -> Result<Vec<u32>, InferenceError> {
self.tokenizer
.encode(text, true)
.map(|e| e.get_ids().to_vec())
.map_err(|e| InferenceError::TokenizationError(e.to_string()))
}
pub fn decode(&self, tokens: &[u32]) -> Result<String, InferenceError> {
self.tokenizer
.decode(tokens, true)
.map_err(|e| InferenceError::TokenizationError(e.to_string()))
}
pub fn eos_token_id(&self) -> Option<u32> {
self.eos.first().copied()
}
pub fn token_id(&self, token: &str) -> Option<u32> {
self.tokenizer.token_to_id(token)
}
pub fn context_length(&self) -> usize {
self.context_length
}
pub fn supports_capability(&self, cap: crate::schema::ModelCapability) -> bool {
<Self as crate::backend::local::LocalInferenceBackend>::supports_capability(self, cap)
}
pub fn embed_one(&mut self, text: &str) -> Result<Vec<f32>, InferenceError> {
#[cfg(not(car_mlxlm_swift_built))]
{
let _ = text;
Err(InferenceError::InferenceFailed(
"this build has no mlx-swift-lm backend, so it serves no embeddings".to_string(),
))
}
#[cfg(car_mlxlm_swift_built)]
{
let handle = self.ensure_embedder()?;
let ids = self.tokenize_raw(text)?;
if ids.is_empty() {
return Ok(vec![0.0; self.embed_dimensions()?]);
}
let tokens: Vec<i32> = ids.iter().map(|&t| t as i32).collect();
let mut out = vec![0.0f32; self.embed_dimensions()?];
let n = unsafe {
ffi::car_mlxlm_embed(
handle,
tokens.as_ptr(),
tokens.len() as i32,
out.as_mut_ptr(),
out.len() as i32,
)
};
if n < 0 {
return Err(InferenceError::InferenceFailed(format!(
"mlx-swift-lm embedding failed: {}",
last_error()
)));
}
out.truncate(n as usize);
Ok(out)
}
}
pub fn embed_query(
&mut self,
text: &str,
instruction: &str,
) -> Result<Vec<f32>, InferenceError> {
self.embed_one(&format!("Instruct: {instruction}\nQuery: {text}"))
}
#[cfg(car_mlxlm_swift_built)]
fn embed_dimensions(&self) -> Result<usize, InferenceError> {
let config: serde_json::Value = serde_json::from_str(
&std::fs::read_to_string(self.model_dir.join("config.json"))
.map_err(|e| InferenceError::InferenceFailed(format!("read config.json: {e}")))?,
)
.map_err(|e| InferenceError::InferenceFailed(format!("parse config.json: {e}")))?;
config
.get("hidden_size")
.and_then(serde_json::Value::as_u64)
.map(|v| v as usize)
.ok_or_else(|| {
InferenceError::InferenceFailed(
"config.json declares no hidden_size, so the embedding width is unknown"
.to_string(),
)
})
}
#[cfg(car_mlxlm_swift_built)]
fn ensure_embedder(&self) -> Result<i32, InferenceError> {
let existing = self.embed_handle.get();
if existing > 0 {
return Ok(existing);
}
let path = CString::new(self.model_dir.to_string_lossy().as_bytes()).map_err(|e| {
InferenceError::InferenceFailed(format!("model path is not a C string: {e}"))
})?;
let handle = unsafe { ffi::car_mlxlm_embed_load(path.as_ptr()) };
if handle < 0 {
return Err(InferenceError::InferenceFailed(format!(
"mlx-swift-lm could not load {} as an embedding model: {}",
self.model_dir.display(),
last_error()
)));
}
self.embed_handle.set(handle);
Ok(handle)
}
pub fn clear_kv_cache(&mut self) {
#[cfg(car_mlxlm_swift_built)]
{
if !self.serves_text() {
return;
}
unsafe {
ffi::car_mlxlm_clear_cache(self.handle);
}
}
}
}
impl Drop for SwiftLmBackend {
fn drop(&mut self) {
#[cfg(car_mlxlm_swift_built)]
unsafe {
if self.handle > 0 {
ffi::car_mlxlm_free(self.handle);
}
let embed = self.embed_handle.get();
if embed > 0 {
ffi::car_mlxlm_embed_free(embed);
}
}
}
}
#[cfg(all(test, car_mlxlm_swift_built))]
mod swift_lm_prefix_reuse_tests {
use super::*;
fn argmax(v: &[f32]) -> usize {
v.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.map(|(i, _)| i)
.unwrap()
}
#[test]
fn batched_prefill_preserves_positions_after_prefix_reuse() {
let Some(dir) = std::env::var_os("HOME")
.map(std::path::PathBuf::from)
.map(|home| home.join(".car/models/Qwen3-0.6B-MLX"))
.filter(|dir| dir.join("config.json").exists())
else {
eprintln!("SKIP: Qwen3-0.6B-MLX is not installed");
return;
};
let mut backend = SwiftLmBackend::load(&dir).unwrap();
let prefix = backend
.tokenize_raw("The following is background context: ")
.unwrap();
let mut primer = prefix.clone();
primer.extend(
backend
.tokenize_raw("a previous unrelated question.")
.unwrap(),
);
backend.forward(&primer, 0).unwrap();
let mut target = prefix;
target.extend(
backend
.tokenize_raw(&"Background data. ".repeat(400))
.unwrap(),
);
target.extend(backend.tokenize_raw("The capital of France is").unwrap());
let reused = backend.begin_prompt(&target);
assert!(reused > 0 && target.len() - reused > 1024);
let warm = backend.forward(&target[reused..], reused).unwrap();
backend.clear_kv_cache();
let cold = backend.forward(&target, 0).unwrap();
assert_eq!(warm.len(), cold.len());
assert!(warm.iter().chain(&cold).all(|value| value.is_finite()));
assert_eq!(
argmax(&warm),
argmax(&cold),
"batched suffix changed the prediction"
);
eprintln!(
"batched prefix test: {} tokens, {reused} reused, matching top-1",
target.len()
);
}
#[test]
fn measures_whether_a_second_turn_reuses_its_prefix() {
let Some(dir) = std::env::var_os("HOME")
.map(std::path::PathBuf::from)
.map(|h| h.join(".car/models/Qwen3-0.6B-MLX"))
.filter(|d| d.join("config.json").exists())
else {
eprintln!("SKIP swift_lm_prefix_reuse_tests: Qwen3-0.6B-MLX not in ~/.car/models");
return;
};
let mut backend = SwiftLmBackend::load(&dir).expect("load");
let turn1: Vec<u32> = backend.tokenize_raw("The capital of France is").unwrap();
let start = backend.begin_prompt(&turn1);
assert_eq!(start, 0, "a cold cache must prefill from zero");
let mut pos = 0usize;
let logits = backend.forward(&turn1, pos).unwrap();
pos += turn1.len();
let mut generated = Vec::new();
for _ in 0..4 {
let next = logits
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.map(|(i, _)| i as u32)
.unwrap();
generated.push(next);
let _ = backend.forward(&[next], pos).unwrap();
pos += 1;
}
let mut turn2 = turn1.clone();
turn2.extend_from_slice(&generated);
turn2.extend(backend.tokenize_raw("\nUser: and of Italy?").unwrap());
let reuse_realistic = backend.begin_prompt(&turn2);
let mut backend2 = SwiftLmBackend::load(&dir).expect("load");
let base: Vec<u32> = backend2.tokenize_raw("The capital of France is").unwrap();
backend2.begin_prompt(&base);
backend2.forward(&base, 0).unwrap();
let mut extended = base.clone();
extended.extend(backend2.tokenize_raw(" Paris and").unwrap());
let reuse_strict = backend2.begin_prompt(&extended);
eprintln!(
"prefix reuse: realistic second turn reused {reuse_realistic} of {} tokens; \
strict extension reused {reuse_strict} of {}",
turn2.len(),
extended.len()
);
assert!(
reuse_strict > 0,
"prefix reuse does not fire even on a strict extension: the mechanism is inert"
);
let mut warm = SwiftLmBackend::load(&dir).expect("load");
let shared: Vec<u32> = warm.tokenize_raw("The capital city of").unwrap();
let mut first = shared.clone();
first.extend(warm.tokenize_raw(" Germany is Berlin").unwrap());
warm.begin_prompt(&first);
warm.forward(&first, 0).unwrap();
let mut edited = shared.clone();
edited.extend(warm.tokenize_raw(" France is").unwrap());
let reused = warm.begin_prompt(&edited);
assert!(
reused >= shared.len(),
"a divergent prompt reused {reused} tokens; the shared prefix is {} long",
shared.len()
);
let warm_logits = warm.forward(&edited[reused..], reused).unwrap();
let mut cold = SwiftLmBackend::load(&dir).expect("load");
cold.begin_prompt(&edited);
let cold_logits = cold.forward(&edited, 0).unwrap();
assert_eq!(warm_logits.len(), cold_logits.len());
let warm_top = argmax(&warm_logits);
let cold_top = argmax(&cold_logits);
assert_eq!(
warm_top, cold_top,
"trimmed-cache reuse changed the prediction: warm {warm_top} vs cold {cold_top}"
);
eprintln!(
"divergent prompt: reused {reused} of {} tokens, top-1 unchanged",
edited.len()
);
assert!(
reuse_realistic >= turn1.len(),
"a realistic second turn reused only {reuse_realistic} of {} tokens; \
prefix reuse across turns has regressed",
turn2.len()
);
}
}
#[cfg(all(test, car_mlxlm_swift_built))]
mod swift_lm_role_tests {
use super::*;
fn embedding_checkpoint() -> Option<std::path::PathBuf> {
let dir = std::env::var_os("HOME")
.map(std::path::PathBuf::from)?
.join(".car/models/Qwen3-Embedding-0.6B-MLX");
dir.join("config.json").exists().then_some(dir)
}
#[test]
fn an_embedding_backend_holds_no_language_model() {
let Some(dir) = embedding_checkpoint() else {
eprintln!("SKIP swift_lm_role_tests: Qwen3-Embedding-0.6B-MLX not in ~/.car/models");
return;
};
let mut backend =
SwiftLmBackend::load_with_role(&dir, BackendRole::Embedding).expect("load embedder");
assert!(!backend.serves_text(), "role should be Embedding");
let v = backend.embed_one("a domestic cat on a windowsill").unwrap();
assert!(v.iter().any(|x| x.abs() > 1e-6), "embedding is all zeros");
let err = backend.forward(&[1, 2, 3], 0).unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("embeddings only") || msg.contains("text-generation"),
"expected a clear refusal, got: {msg}"
);
assert_eq!(backend.begin_prompt(&[1, 2, 3]), 0, "no KV cache to reuse");
}
#[test]
fn a_text_backend_on_the_same_checkpoint_still_generates() {
let Some(dir) = embedding_checkpoint() else {
eprintln!("SKIP swift_lm_role_tests: checkpoint absent");
return;
};
let mut backend =
SwiftLmBackend::load_with_role(&dir, BackendRole::Text).expect("load as text");
assert!(backend.serves_text());
let logits = backend.forward(&[1, 2, 3], 0).expect("text role generates");
assert!(!logits.is_empty(), "no logits from a text-role backend");
}
}
#[cfg(all(test, car_mlxlm_swift_built))]
mod swift_lm_embed_tests {
use super::*;
fn embedding_model() -> Option<std::path::PathBuf> {
let dir = dirs_home()?.join(".car/models/Qwen3-Embedding-0.6B-MLX");
dir.join("config.json").exists().then_some(dir)
}
fn dirs_home() -> Option<std::path::PathBuf> {
std::env::var_os("HOME").map(std::path::PathBuf::from)
}
#[test]
fn embeds_and_the_vector_means_something() {
let Some(dir) = embedding_model() else {
eprintln!("SKIP swift_lm_embed_tests: Qwen3-Embedding-0.6B-MLX not in ~/.car/models");
return;
};
let mut backend = SwiftLmBackend::load(&dir).expect("load embedding checkpoint");
let cat = backend
.embed_one("a domestic cat sitting on a windowsill")
.unwrap();
let kitten = backend.embed_one("a small kitten by the window").unwrap();
let finance = backend
.embed_one("quarterly earnings guidance for the fiscal year")
.unwrap();
assert_eq!(cat.len(), kitten.len());
assert_eq!(cat.len(), finance.len());
assert!(cat.len() >= 512, "unexpected width {}", cat.len());
assert!(
cat.iter().any(|v| v.abs() > 1e-6),
"embedding is all zeros: the pooling step produced nothing"
);
assert!(cat.iter().all(|v| v.is_finite()), "embedding has NaN/inf");
let norm: f32 = cat.iter().map(|v| v * v).sum::<f32>().sqrt();
assert!((norm - 1.0).abs() < 1e-2, "not L2-normalized: {norm}");
let dot = |a: &[f32], b: &[f32]| a.iter().zip(b).map(|(x, y)| x * y).sum::<f32>();
let near = dot(&cat, &kitten);
let far = dot(&cat, &finance);
assert!(
near > far,
"semantic ordering lost: cat~kitten {near:.4} should exceed cat~finance {far:.4}"
);
}
#[test]
fn the_same_text_embeds_identically() {
let Some(dir) = embedding_model() else {
eprintln!("SKIP swift_lm_embed_tests: checkpoint absent");
return;
};
let mut backend = SwiftLmBackend::load(&dir).expect("load");
let a = backend.embed_one("determinism matters").unwrap();
let b = backend.embed_one("determinism matters").unwrap();
assert_eq!(a.len(), b.len());
for (x, y) in a.iter().zip(&b) {
assert!((x - y).abs() < 1e-5, "nondeterministic embedding");
}
}
}
#[cfg(all(test, car_mlxlm_swift_built))]
mod swift_lm_parity_tests {
use super::SwiftLmBackend;
use std::path::PathBuf;
const FIXTURE: &str = include_str!("../../tests/fixtures/qwen3-0.6b-reference-logits.json");
fn bf16_ulp(v: f32) -> f32 {
if v == 0.0 {
return f32::MIN_POSITIVE;
}
(v.abs().log2().floor() - 7.0).exp2()
}
fn tokenize(prompt: &str, dir: &std::path::Path) -> Option<Vec<u32>> {
let tk = tokenizers::Tokenizer::from_file(dir.join("tokenizer.json")).ok()?;
Some(tk.encode(prompt, false).ok()?.get_ids().to_vec())
}
#[test]
fn swift_backend_matches_the_mlx_lm_reference() {
let dir = PathBuf::from(std::env::var("HOME").unwrap_or_default())
.join(".car/models/Qwen3-0.6B-MLX");
if !dir.join("config.json").exists() {
eprintln!(
"SKIPPED swift_backend_matches_the_mlx_lm_reference: {} not installed. \
This check is not running — it protects nothing on this machine.",
dir.display()
);
return;
}
let fixture: serde_json::Value = serde_json::from_str(FIXTURE).expect("fixture parses");
let cases = fixture["cases"].as_array().expect("cases");
let vocab = 151_936usize;
let mut backend = SwiftLmBackend::load(&dir).expect("load via mlx-swift-lm");
for case in cases {
let prompt = case["prompt"].as_str().expect("prompt");
let Some(ids) = tokenize(prompt, &dir) else {
eprintln!("SKIPPED: no tokenizer.json at {}", dir.display());
return;
};
assert_eq!(
ids.len() as u64,
case["prompt_tokens"].as_u64().expect("prompt_tokens"),
"tokenizer disagrees with the reference on {prompt:?}"
);
backend.clear_kv_cache();
let logits = backend.forward(&ids, 0).expect("forward");
assert_eq!(logits.len(), vocab, "unexpected logit count");
let expected = case["top5"].as_array().expect("top5");
let want_top1 = expected[0][0].as_u64().expect("id") as usize;
let got_top1 = logits
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.map(|(i, _)| i)
.expect("argmax");
if got_top1 != want_top1 {
let gap = (logits[got_top1] - logits[want_top1]).abs();
assert!(
gap <= bf16_ulp(logits[want_top1]),
"{prompt:?}: Swift chose {got_top1} ({}) where the reference chose \
{want_top1} ({}) — a gap of {gap}, too large to be a tie",
logits[got_top1],
logits[want_top1]
);
}
for entry in expected {
let id = entry[0].as_u64().expect("id") as usize;
let want = entry[1].as_f64().expect("logit") as f32;
let ulps = (logits[id] - want).abs() / bf16_ulp(want);
assert!(
ulps <= 3.0,
"{prompt:?}: token {id} is {} against the reference's {want} — \
{ulps:.1} bf16 ulps, beyond cross-implementation rounding",
logits[id]
);
}
}
}
}
impl crate::backend::local::TextDecoder for SwiftLmBackend {
fn encode(&self, text: &str) -> Result<Vec<u32>, InferenceError> {
self.tokenizer
.encode(text, true)
.map(|e| e.get_ids().to_vec())
.map_err(|e| InferenceError::TokenizationError(e.to_string()))
}
fn decode(&self, tokens: &[u32]) -> Result<String, InferenceError> {
self.tokenizer
.decode(tokens, true)
.map_err(|e| InferenceError::TokenizationError(e.to_string()))
}
fn forward(&mut self, tokens: &[u32], pos: usize) -> Result<Vec<f32>, InferenceError> {
SwiftLmBackend::forward(self, tokens, pos)
}
fn eos_ids(&self) -> Vec<u32> {
self.eos.clone()
}
fn context_length(&self) -> usize {
self.context_length
}
fn clear_kv_cache(&mut self) {
SwiftLmBackend::clear_kv_cache(self)
}
fn begin_prompt(&mut self, prompt_tokens: &[u32]) -> usize {
SwiftLmBackend::begin_prompt(self, prompt_tokens)
}
}
#[cfg(all(test, car_mlxlm_swift_built))]
mod swift_lm_decode_tests {
use super::SwiftLmBackend;
use crate::backend::local::TextDecoder;
use std::path::PathBuf;
#[test]
fn generates_through_the_shared_text_decoder_surface() {
let dir = PathBuf::from(std::env::var("HOME").unwrap_or_default())
.join(".car/models/Qwen3-0.6B-MLX");
if !dir.join("config.json").exists() {
eprintln!("SKIPPED generates_through_the_shared_text_decoder_surface: no checkpoint");
return;
}
let mut backend = SwiftLmBackend::load(&dir).expect("load");
assert!(
backend.context_length() >= 4096,
"context window looks wrong"
);
assert!(!backend.eos_ids().is_empty(), "no stop tokens");
let prompt = "The capital of France is";
let ids = backend.encode(prompt).expect("encode");
assert!(!ids.is_empty());
let offset = backend.begin_prompt(&ids);
assert!(offset <= ids.len());
let mut logits = backend.forward(&ids[offset..], offset).expect("prefill");
let eos = backend.eos_ids();
let mut out = Vec::new();
for pos in (ids.len()..).take(6) {
let tok = logits
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.map(|(i, _)| i as u32)
.unwrap_or(0);
if eos.contains(&tok) {
break;
}
out.push(tok);
logits = backend.forward(&[tok], pos).expect("decode step");
}
let text = backend.decode(&out).expect("decode");
assert!(!out.is_empty(), "generated nothing");
assert!(
text.chars().any(|c| c.is_alphabetic()),
"generated no letters: {text:?} — the decode loop is producing junk"
);
eprintln!("swift decode loop produced: {text:?}");
}
}
impl crate::backend::local::LocalInferenceBackend for SwiftLmBackend {
fn backend_name(&self) -> &'static str {
"mlx-swift-lm"
}
fn supports_capability(&self, cap: crate::schema::ModelCapability) -> bool {
use crate::schema::ModelCapability as C;
match cap {
C::Generate
| C::ToolUse
| C::MultiToolCall
| C::Reasoning
| C::Summarize
| C::Code
| C::Classify
| C::Embed
| C::Rerank => true,
C::Grounding
| C::Vision
| C::VideoUnderstanding
| C::AudioUnderstanding
| C::SpeechToText
| C::TextToSpeech
| C::ImageGeneration
| C::VideoGeneration => false,
}
}
fn render_prompt(&self, req: &crate::GenerateRequest) -> Result<String, InferenceError> {
match &self.chat_template {
Some(t) => t.render_request(req),
None => Ok(crate::tasks::generate::render_chat_prompt(req)),
}
}
fn parse_tool_calls(&self, text: &str) -> (String, Vec<crate::ToolCall>) {
if self.model_type.starts_with("gemma") {
crate::tasks::generate::parse_gemma4_tool_calls(text)
} else {
crate::tasks::generate::parse_tool_calls(text)
}
}
}