use std::collections::{BTreeSet, HashMap};
use std::path::{Path, PathBuf};
use lattice_inference::model::qwen35::qwen_required_tensor_names;
use lattice_inference::model::qwen35_config::Qwen35Config;
fn kv_cache_dtype_bytes() -> usize {
if matches!(
std::env::var("LATTICE_KV_F16").as_deref(),
Ok("1") | Ok("true")
) {
2
} else {
4
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Placement {
Cpu,
Metal,
}
impl std::fmt::Display for Placement {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Placement::Cpu => write!(f, "CPU"),
Placement::Metal => write!(f, "Metal GPU"),
}
}
}
#[derive(Debug)]
struct TensorEntry {
dtype: String,
byte_len: u64,
}
struct WeightInventory {
total_bytes: u64,
tensor_count: usize,
quantization: String,
quantization_error: Option<String>,
missing_tensors: Vec<String>,
unsupported_dtypes: Vec<String>,
has_mtp_tensors: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct MemoryPlan {
pub weight_bytes: u64,
pub kv_bytes_per_token: u64,
pub available_memory_bytes: u64,
pub max_context_by_memory: u64,
pub max_position_embeddings: usize,
pub max_context_len: usize,
pub requested_context: Option<usize>,
pub requested_fits: Option<bool>,
}
pub fn plan_memory(
weight_bytes: u64,
kv_bytes_per_token: u64,
available_memory_bytes: u64,
max_position_embeddings: usize,
requested_context: Option<usize>,
) -> MemoryPlan {
let usable = available_memory_bytes.saturating_sub(weight_bytes);
let max_context_by_memory = if kv_bytes_per_token == 0 {
u64::MAX
} else {
usable / kv_bytes_per_token
};
let max_context_len = max_context_by_memory.min(max_position_embeddings as u64) as usize;
let requested_fits = requested_context.map(|c| c <= max_context_len);
MemoryPlan {
weight_bytes,
kv_bytes_per_token,
available_memory_bytes,
max_context_by_memory,
max_position_embeddings,
max_context_len,
requested_context,
requested_fits,
}
}
pub const METAL_RUNTIME_MAX_CACHE_LEN: usize = 4096;
pub fn effective_max_position_embeddings(
placement: Placement,
max_position_embeddings: usize,
) -> usize {
if placement == Placement::Metal {
max_position_embeddings.min(METAL_RUNTIME_MAX_CACHE_LEN)
} else {
max_position_embeddings
}
}
fn human_bytes(bytes: u64) -> String {
const KIB: f64 = 1024.0;
const MIB: f64 = KIB * 1024.0;
const GIB: f64 = MIB * 1024.0;
let b = bytes as f64;
if b >= GIB {
format!("{:.2} GiB", b / GIB)
} else if b >= MIB {
format!("{:.2} MiB", b / MIB)
} else if b >= KIB {
format!("{:.2} KiB", b / KIB)
} else {
format!("{bytes} B")
}
}
pub fn detect_total_memory_bytes() -> Option<u64> {
#[cfg(target_os = "macos")]
{
std::process::Command::new("sysctl")
.args(["-n", "hw.memsize"])
.output()
.ok()
.filter(|o| o.status.success())
.and_then(|o| String::from_utf8(o.stdout).ok())
.and_then(|s| s.trim().parse::<u64>().ok())
}
#[cfg(target_os = "linux")]
{
std::fs::read_to_string("/proc/meminfo")
.ok()
.and_then(|contents| {
contents.lines().find_map(|line| {
let rest = line.strip_prefix("MemTotal:")?;
let kb_str = rest.trim().strip_suffix(" kB")?.trim();
kb_str.parse::<u64>().ok().map(|kb| kb * 1024)
})
})
}
#[cfg(not(any(target_os = "macos", target_os = "linux")))]
{
None
}
}
fn read_safetensors_header(path: &Path) -> Result<HashMap<String, TensorEntry>, String> {
use std::io::Read;
let mut file =
std::fs::File::open(path).map_err(|e| format!("failed to open {}: {e}", path.display()))?;
let file_len = file
.metadata()
.map_err(|e| format!("failed to stat {}: {e}", path.display()))?
.len();
let mut len_buf = [0u8; 8];
file.read_exact(&mut len_buf)
.map_err(|e| format!("failed to read header length from {}: {e}", path.display()))?;
let header_len = u64::from_le_bytes(len_buf);
if header_len > file_len.saturating_sub(8) {
return Err(format!(
"{}: header length {header_len} exceeds file size {file_len}",
path.display()
));
}
let mut header_buf = vec![0u8; header_len as usize];
file.read_exact(&mut header_buf)
.map_err(|e| format!("failed to read header from {}: {e}", path.display()))?;
let header_str = std::str::from_utf8(&header_buf)
.map_err(|e| format!("{} header is not valid UTF-8: {e}", path.display()))?;
let root: serde_json::Value = serde_json::from_str(header_str)
.map_err(|e| format!("{} header is not valid JSON: {e}", path.display()))?;
let obj = root
.as_object()
.ok_or_else(|| format!("{} header is not a JSON object", path.display()))?;
let mut out = HashMap::with_capacity(obj.len());
for (name, entry) in obj {
if name == "__metadata__" {
continue;
}
let dtype = entry
.get("dtype")
.and_then(|v| v.as_str())
.unwrap_or("UNKNOWN")
.to_string();
let Some(offsets) = entry.get("data_offsets").and_then(|v| v.as_array()) else {
return Err(format!(
"tensor '{name}' in {} has no data_offsets",
path.display()
));
};
if offsets.len() != 2 {
return Err(format!(
"tensor '{name}' in {} has malformed data_offsets",
path.display()
));
}
let start = offsets[0].as_u64().unwrap_or(0);
let end = offsets[1].as_u64().unwrap_or(0);
out.insert(
name.clone(),
TensorEntry {
dtype,
byte_len: end.saturating_sub(start),
},
);
}
Ok(out)
}
const SUPPORTED_DTYPES: [&str; 3] = ["F32", "F16", "BF16"];
fn cpu_resident_bytes(dtype: &str, on_disk_byte_len: u64) -> u64 {
match dtype {
"F16" | "BF16" => on_disk_byte_len * 2,
_ => on_disk_byte_len,
}
}
fn inspect_safetensors_dir(dir: &Path, cfg: &Qwen35Config) -> Result<WeightInventory, String> {
let single = dir.join("model.safetensors");
let index_path = dir.join("model.safetensors.index.json");
let all_tensors: HashMap<String, TensorEntry> = if single.exists() {
read_safetensors_header(&single)?
} else if index_path.exists() {
let index_bytes = std::fs::read(&index_path)
.map_err(|e| format!("failed to read {}: {e}", index_path.display()))?;
let index: serde_json::Value = serde_json::from_slice(&index_bytes)
.map_err(|e| format!("{} is not valid JSON: {e}", index_path.display()))?;
let weight_map = index
.get("weight_map")
.and_then(|v| v.as_object())
.ok_or_else(|| format!("{} has no weight_map object", index_path.display()))?;
let mut shard_names: BTreeSet<String> = BTreeSet::new();
for v in weight_map.values() {
if let Some(s) = v.as_str() {
shard_names.insert(s.to_string());
}
}
let mut merged = HashMap::new();
for shard_name in shard_names {
let shard_path = lattice_inference::weights::contained_shard_path(dir, &shard_name)
.map_err(|e| e.to_string())?;
if !shard_path.exists() {
return Err(format!(
"shard '{shard_name}' referenced by {} not found in {}",
index_path.display(),
dir.display()
));
}
merged.extend(read_safetensors_header(&shard_path)?);
}
merged
} else {
return Err(format!(
"no model.safetensors or model.safetensors.index.json in {}",
dir.display()
));
};
let total_bytes: u64 = all_tensors
.values()
.map(|t| cpu_resident_bytes(&t.dtype, t.byte_len))
.sum();
let dtypes: BTreeSet<&str> = all_tensors.values().map(|t| t.dtype.as_str()).collect();
let quantization = if dtypes.is_empty() {
"unknown".to_string()
} else {
dtypes.into_iter().collect::<Vec<_>>().join(", ")
};
let mut missing_tensors = Vec::new();
let mut unsupported_dtypes = Vec::new();
for name in qwen_required_tensor_names(cfg) {
match all_tensors.get(&name) {
Some(entry) if !SUPPORTED_DTYPES.contains(&entry.dtype.as_str()) => {
unsupported_dtypes.push(format!(
"tensor '{name}' has dtype {}, which is not supported (supported: F32, F16, BF16)",
entry.dtype
));
}
Some(_) => {}
None => missing_tensors.push(name),
}
}
Ok(WeightInventory {
total_bytes,
tensor_count: all_tensors.len(),
quantization,
quantization_error: None,
missing_tensors,
unsupported_dtypes,
has_mtp_tensors: false,
})
}
fn q4_resident_bytes(name_or_file: &str, on_disk_bytes: u64, tie_word_embeddings: bool) -> u64 {
let mut n = name_or_file.replace('.', "_");
if let Some(stripped) = n.strip_suffix("_q4").or_else(|| n.strip_suffix("_f16")) {
n = stripped.to_string();
}
if n.ends_with("norm_weight")
|| n.ends_with("A_log")
|| n.ends_with("dt_bias")
|| n.ends_with("conv1d_weight")
|| n.ends_with("pre_fc_norm_embedding_weight")
|| n.ends_with("pre_fc_norm_hidden_weight")
{
return on_disk_bytes.saturating_mul(2);
}
if n.ends_with("in_proj_a_weight") || n.ends_with("in_proj_b_weight") {
return (on_disk_bytes as f64 * 3.2).round() as u64;
}
if n.ends_with("in_proj_qkv_weight") || n.ends_with("in_proj_z_weight") {
return on_disk_bytes.saturating_mul(2);
}
if n.ends_with("embed_tokens_weight") {
let ratio = if tie_word_embeddings { 4.2 } else { 3.2 };
return (on_disk_bytes as f64 * ratio).round() as u64;
}
on_disk_bytes
}
fn detect_q4_quantization_label(dir: &Path) -> Result<String, String> {
let sample = std::fs::read_dir(dir)
.map_err(|e| format!("failed to read directory {}: {e}", dir.display()))?
.flatten()
.find(|e| {
e.file_name()
.to_str()
.map(|n| n.ends_with(".q4") && !n.starts_with("merged_qkvz_"))
.unwrap_or(false)
});
let Some(sample) = sample else {
return Ok("Q4 (no .q4 files found to sample)".to_string());
};
let mut file = std::fs::File::open(sample.path())
.map_err(|e| format!("failed to open {}: {e}", sample.path().display()))?;
lattice_inference::weights::q4_weights::read_q4_header(&mut file)
.map(|_| "Q4_0 (lattice native, v2 asymmetric)".to_string())
.map_err(|e| {
format!(
"unsupported quantization scheme in {}: {e}",
sample.path().display()
)
})
}
fn inspect_q4_dir(dir: &Path, cfg: &Qwen35Config) -> Result<WeightInventory, String> {
let (quantization, quantization_error) = match detect_q4_quantization_label(dir) {
Ok(label) => (label, None),
Err(e) => ("unknown (see blocking reasons)".to_string(), Some(e)),
};
let manifest = lattice_inference::quant::q4_manifest::load_manifest(dir)?;
if let Some(manifest) = manifest {
let entries = manifest.tensors;
let mut total_bytes = 0u64;
let mut missing_tensors = Vec::new();
let mut present_names: BTreeSet<String> = BTreeSet::new();
let mut has_mtp_tensors = false;
for entry in &entries {
let Ok(file_path) = lattice_inference::weights::contained_shard_path(dir, &entry.file)
else {
missing_tensors.push(format!(
"{} (listed in quantize_index.json as '{}', \
path escapes the model directory)",
entry.name, entry.file
));
continue;
};
match std::fs::metadata(&file_path) {
Ok(meta) => {
total_bytes +=
q4_resident_bytes(&entry.name, meta.len(), cfg.tie_word_embeddings);
present_names.insert(entry.name.clone());
has_mtp_tensors |= entry.name.starts_with("mtp.");
}
Err(_) => missing_tensors.push(format!(
"{} (listed in quantize_index.json as '{}', file not found)",
entry.name, entry.file
)),
}
}
for name in qwen_required_tensor_names(cfg) {
if !present_names.contains(&name) {
missing_tensors.push(name);
}
}
Ok(WeightInventory {
total_bytes,
tensor_count: entries.len(),
quantization,
quantization_error,
missing_tensors,
unsupported_dtypes: Vec::new(),
has_mtp_tensors,
})
} else {
let mut total_bytes = 0u64;
let mut tensor_count = 0usize;
let mut has_mtp_tensors = false;
let read_dir = std::fs::read_dir(dir)
.map_err(|e| format!("failed to read directory {}: {e}", dir.display()))?;
for entry in read_dir.flatten() {
let file_name = entry.file_name();
let Some(name) = file_name.to_str() else {
continue;
};
if name.starts_with("merged_qkvz_") {
continue;
}
if (name.ends_with(".q4") || name.ends_with(".f16"))
&& let Ok(meta) = entry.metadata()
{
total_bytes += q4_resident_bytes(name, meta.len(), cfg.tie_word_embeddings);
tensor_count += 1;
has_mtp_tensors |= name.starts_with("mtp_");
}
}
Ok(WeightInventory {
total_bytes,
tensor_count,
quantization,
quantization_error,
missing_tensors: Vec::new(),
unsupported_dtypes: Vec::new(),
has_mtp_tensors,
})
}
}
#[derive(Debug)]
pub struct DoctorReport {
pub model_dir: PathBuf,
pub format: crate::backend::ModelFormat,
pub placement: Placement,
pub quantization: String,
pub tensor_count: usize,
pub weight_bytes: u64,
pub kv_bytes_per_token: u64,
pub max_position_embeddings: usize,
pub metal_runtime_cache_cap: Option<usize>,
pub available_memory_bytes: Option<u64>,
pub max_context_len: Option<usize>,
pub requested_context: Option<usize>,
pub requested_fits: Option<bool>,
pub tokenizer_path: PathBuf,
pub tokenizer_present: bool,
pub missing_tensors: Vec<String>,
pub blocking_reasons: Vec<String>,
pub has_mtp_tensors: bool,
}
impl DoctorReport {
pub fn is_ready(&self) -> bool {
self.blocking_reasons.is_empty()
}
}
impl std::fmt::Display for DoctorReport {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
writeln!(f, "Model directory : {}", self.model_dir.display())?;
writeln!(f, "Format : {:?}", self.format)?;
writeln!(f, "Placement : {}", self.placement)?;
writeln!(f, "Quantization : {}", self.quantization)?;
writeln!(f, "Tensors found : {}", self.tensor_count)?;
writeln!(
f,
"Weight memory : {} ({} bytes)",
human_bytes(self.weight_bytes),
self.weight_bytes
)?;
writeln!(
f,
"KV cache : {}/token ({} bytes/token; at 4096 tokens ~= {})",
human_bytes(self.kv_bytes_per_token),
self.kv_bytes_per_token,
human_bytes(self.kv_bytes_per_token.saturating_mul(4096))
)?;
writeln!(
f,
"Model max context (max_position_embeddings): {}",
self.max_position_embeddings
)?;
if let Some(cap) = self.metal_runtime_cache_cap {
writeln!(
f,
"Metal runtime cache cap: {cap} tokens (MetalChatBackend::MAX_CACHE_LEN; \
the chat/serve binaries never allow more, even if memory allows it)"
)?;
}
match self.available_memory_bytes {
Some(avail) => writeln!(
f,
"Detected system memory: {} ({} bytes)",
human_bytes(avail),
avail
)?,
None => writeln!(
f,
"Detected system memory: unknown (unsupported OS or query failed)"
)?,
}
match self.max_context_len {
Some(max_ctx) => writeln!(
f,
"Max feasible context length: {max_ctx} tokens \
(weights + KV cache only -- activation buffers, GDN recurrent-state \
scratch, and tokenizer tables are not counted)"
)?,
None => writeln!(
f,
"Max feasible context length: unknown (system memory undetected)"
)?,
}
if self.has_mtp_tensors {
writeln!(
f,
"Note: this directory includes MTP (multi-token prediction) files -- \
the separate MTP K/V cache and session buffers `from_q4_dir` allocates \
for them are also not counted in the estimate above."
)?;
}
if let Some(requested) = self.requested_context {
match self.requested_fits {
Some(true) => writeln!(f, "Requested context {requested}: fits")?,
Some(false) => writeln!(f, "Requested context {requested}: DOES NOT FIT")?,
None => writeln!(
f,
"Requested context {requested}: unknown (system memory undetected)"
)?,
}
}
writeln!(
f,
"Tokenizer : {} ({})",
self.tokenizer_path.display(),
if self.tokenizer_present {
"found"
} else {
"MISSING"
}
)?;
if !self.missing_tensors.is_empty() {
writeln!(
f,
"Missing required tensors ({}):",
self.missing_tensors.len()
)?;
for name in self.missing_tensors.iter().take(20) {
writeln!(f, " - {name}")?;
}
if self.missing_tensors.len() > 20 {
writeln!(f, " ... and {} more", self.missing_tensors.len() - 20)?;
}
}
writeln!(f)?;
if self.is_ready() {
writeln!(f, "Result: OK -- weights + KV cache fit; ready to load")?;
} else {
writeln!(f, "Result: NOT READY")?;
for reason in &self.blocking_reasons {
writeln!(f, " - {reason}")?;
}
}
Ok(())
}
}
pub fn build_report(
model_dir: &Path,
tokenizer_dir: Option<&Path>,
requested_context: Option<usize>,
available_memory_override: Option<u64>,
) -> Result<DoctorReport, String> {
let format = crate::backend::detect_format(model_dir);
let (placement, cfg, inventory) = match format {
crate::backend::ModelFormat::Safetensors => {
let cfg = Qwen35Config::from_model_dir(model_dir)
.map_err(|e| format!("config.json load failed: {e}"))?;
let inventory = inspect_safetensors_dir(model_dir, &cfg)?;
(Placement::Cpu, cfg, inventory)
}
crate::backend::ModelFormat::Q4 => {
let cfg = Qwen35Config::from_model_dir(model_dir)
.map_err(|e| format!("config.json load failed: {e}"))?;
let inventory = inspect_q4_dir(model_dir, &cfg)?;
(Placement::Metal, cfg, inventory)
}
crate::backend::ModelFormat::Unknown => {
return Err(crate::backend::unrecognized_format_message(model_dir));
}
_ => {
return Err(crate::backend::unrecognized_format_message(model_dir));
}
};
let tokenizer_path = tokenizer_dir.unwrap_or(model_dir).join("tokenizer.json");
let tokenizer_present = tokenizer_path.exists();
let kv_bytes_per_token = cfg.kv_bytes_per_token(kv_cache_dtype_bytes()) as u64;
let available_memory_bytes = available_memory_override.or_else(detect_total_memory_bytes);
let effective_max_position_embeddings =
effective_max_position_embeddings(placement, cfg.max_position_embeddings);
let (max_context_len, requested_fits) = match available_memory_bytes {
Some(avail) => {
let plan = plan_memory(
inventory.total_bytes,
kv_bytes_per_token,
avail,
effective_max_position_embeddings,
requested_context,
);
(Some(plan.max_context_len), plan.requested_fits)
}
None => (None, None),
};
let mut blocking_reasons = Vec::new();
if let Some(e) = &inventory.quantization_error {
blocking_reasons.push(e.clone());
}
blocking_reasons.extend(inventory.unsupported_dtypes.iter().cloned());
if !inventory.missing_tensors.is_empty() {
blocking_reasons.push(format!(
"{} required tensor(s) missing (see list below), e.g. '{}'",
inventory.missing_tensors.len(),
inventory.missing_tensors[0]
));
}
if !tokenizer_present {
blocking_reasons.push(format!(
"tokenizer.json not found at {}",
tokenizer_path.display()
));
}
if format == crate::backend::ModelFormat::Q4 && !cfg!(feature = "metal-gpu") {
blocking_reasons.push(crate::backend::metal_gpu_required_message(model_dir));
}
if requested_fits == Some(false) {
let requested = requested_context.unwrap_or(0);
let max_ctx = max_context_len.unwrap_or(0);
blocking_reasons.push(format!(
"requested context {requested} does not fit: max feasible is {max_ctx} tokens"
));
}
if let Some(avail) = available_memory_bytes
&& inventory.total_bytes > avail
{
blocking_reasons.push(format!(
"estimated weight memory {} exceeds detected system memory {}: this model is \
unlikely to load on this machine (estimate of peak resident footprint from \
on-disk tensor sizes -- mmap'd weights fault into physical memory lazily on \
first touch, so this is not a measurement of memory actually held at any \
instant)",
human_bytes(inventory.total_bytes),
human_bytes(avail)
));
}
if max_context_len == Some(0) {
blocking_reasons.push(
"max feasible context length is 0 tokens: weights and/or KV cache do not fit \
in available memory (not even a single token of context)"
.to_string(),
);
}
Ok(DoctorReport {
model_dir: model_dir.to_path_buf(),
format,
placement,
quantization: inventory.quantization,
tensor_count: inventory.tensor_count,
weight_bytes: inventory.total_bytes,
kv_bytes_per_token,
max_position_embeddings: cfg.max_position_embeddings,
metal_runtime_cache_cap: (placement == Placement::Metal)
.then_some(METAL_RUNTIME_MAX_CACHE_LEN),
available_memory_bytes,
max_context_len,
requested_context,
requested_fits,
tokenizer_path,
tokenizer_present,
missing_tensors: inventory.missing_tensors,
blocking_reasons,
has_mtp_tensors: inventory.has_mtp_tensors,
})
}
#[cfg(test)]
mod tests {
use super::*;
use std::fs;
fn tempdir(name: &str) -> PathBuf {
let mut dir = std::env::temp_dir();
dir.push(format!(
"lattice-doctor-test-{name}-{}-{}",
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_nanos())
.unwrap_or(0)
));
fs::create_dir_all(&dir).expect("create tempdir");
dir
}
#[test]
fn human_bytes_formats_units() {
assert_eq!(human_bytes(0), "0 B");
assert_eq!(human_bytes(512), "512 B");
assert_eq!(human_bytes(1024), "1.00 KiB");
assert_eq!(human_bytes(1024 * 1024), "1.00 MiB");
assert_eq!(human_bytes(1024 * 1024 * 1024), "1.00 GiB");
assert_eq!(human_bytes(1536 * 1024 * 1024), "1.50 GiB");
}
#[test]
fn plan_memory_hand_computed() {
let plan = plan_memory(100, 10, 1000, 1_000_000, Some(50));
assert_eq!(plan.max_context_by_memory, 90);
assert_eq!(plan.max_context_len, 90); assert_eq!(plan.requested_fits, Some(true));
let plan2 = plan_memory(100, 10, 1000, 1_000_000, Some(95));
assert_eq!(plan2.requested_fits, Some(false)); }
#[test]
fn plan_memory_capped_by_max_position_embeddings() {
let plan = plan_memory(0, 1, 1_000_000_000_000, 4096, None);
assert_eq!(plan.max_context_by_memory, 1_000_000_000_000);
assert_eq!(plan.max_context_len, 4096);
}
#[test]
fn effective_max_position_embeddings_caps_metal_placement_at_runtime_limit() {
assert_eq!(
effective_max_position_embeddings(Placement::Metal, 131_072),
METAL_RUNTIME_MAX_CACHE_LEN
);
}
#[test]
fn effective_max_position_embeddings_leaves_smaller_model_ceiling_untouched() {
assert_eq!(
effective_max_position_embeddings(Placement::Metal, 2048),
2048
);
}
#[test]
fn effective_max_position_embeddings_does_not_cap_cpu_placement() {
assert_eq!(
effective_max_position_embeddings(Placement::Cpu, 131_072),
131_072
);
}
#[test]
fn plan_memory_weight_bytes_exceeding_available_yields_zero_context() {
let plan = plan_memory(2_000, 10, 1_000, 1_000_000, Some(1));
assert_eq!(plan.max_context_by_memory, 0);
assert_eq!(plan.max_context_len, 0);
assert_eq!(plan.requested_fits, Some(false));
}
#[test]
fn plan_memory_kv_bytes_per_token_matches_qwen35_0_8b_doc_identity() {
let cfg = Qwen35Config::qwen35_0_8b();
assert_eq!(cfg.num_full_attention_layers(), 6);
assert_eq!(cfg.full_kv_dim(), 512);
assert_eq!(cfg.kv_bytes_per_token(2), 12_288);
assert_eq!(cfg.kv_bytes_per_token(4), 24_576);
}
fn write_fake_safetensors(path: &Path, tensors: &[(&str, &str, u64, u64)]) {
let mut header = serde_json::Map::new();
for (name, dtype, start, end) in tensors {
header.insert(
(*name).to_string(),
serde_json::json!({
"dtype": dtype,
"shape": [1],
"data_offsets": [start, end],
}),
);
}
let header_json = serde_json::Value::Object(header).to_string();
let header_bytes = header_json.as_bytes();
let mut buf = Vec::new();
buf.extend_from_slice(&(header_bytes.len() as u64).to_le_bytes());
buf.extend_from_slice(header_bytes);
let payload_len = tensors.iter().map(|(_, _, _, e)| *e).max().unwrap_or(0);
buf.resize(buf.len() + payload_len as usize, 0);
fs::write(path, buf).expect("write fake safetensors file");
}
#[test]
fn read_safetensors_header_computes_exact_byte_lengths() {
let dir = tempdir("st-header");
let path = dir.join("model.safetensors");
write_fake_safetensors(
&path,
&[("tensor.a", "F32", 0, 400), ("tensor.b", "BF16", 400, 600)],
);
let tensors = read_safetensors_header(&path).unwrap();
assert_eq!(tensors.len(), 2);
assert_eq!(tensors["tensor.a"].byte_len, 400);
assert_eq!(tensors["tensor.a"].dtype, "F32");
assert_eq!(tensors["tensor.b"].byte_len, 200);
assert_eq!(tensors["tensor.b"].dtype, "BF16");
fs::remove_dir_all(&dir).ok();
}
#[test]
fn read_safetensors_header_rejects_oversized_header_length() {
let dir = tempdir("st-header-bad");
let path = dir.join("model.safetensors");
let mut buf = Vec::new();
buf.extend_from_slice(&(1_000_000_000u64).to_le_bytes());
buf.extend_from_slice(b"tiny");
fs::write(&path, buf).unwrap();
let err = read_safetensors_header(&path).unwrap_err();
assert!(err.contains("exceeds file size"));
fs::remove_dir_all(&dir).ok();
}
fn required_tensor_fixture(cfg: &Qwen35Config) -> Vec<(String, String, u64, u64)> {
let mut offset = 0u64;
qwen_required_tensor_names(cfg)
.into_iter()
.map(|name| {
let start = offset;
offset += 64;
(name, "F32".to_string(), start, offset)
})
.collect()
}
#[test]
fn inspect_safetensors_dir_all_tensors_present_no_missing() {
let dir = tempdir("st-complete");
let cfg = Qwen35Config::qwen35_0_8b();
let tensors = required_tensor_fixture(&cfg);
let refs: Vec<(&str, &str, u64, u64)> = tensors
.iter()
.map(|(n, d, s, e)| (n.as_str(), d.as_str(), *s, *e))
.collect();
write_fake_safetensors(&dir.join("model.safetensors"), &refs);
let inv = inspect_safetensors_dir(&dir, &cfg).unwrap();
assert!(inv.missing_tensors.is_empty());
assert!(inv.unsupported_dtypes.is_empty());
assert_eq!(inv.total_bytes, tensors.len() as u64 * 64);
assert_eq!(inv.quantization, "F32");
fs::remove_dir_all(&dir).ok();
}
#[test]
fn inspect_safetensors_dir_scales_f16_bf16_to_resident_f32_bytes() {
let dir = tempdir("st-mixed-dtype");
let cfg = Qwen35Config::qwen35_0_8b();
let mut tensors = required_tensor_fixture(&cfg);
assert!(
tensors.len() >= 2,
"fixture must have room to mutate two entries"
);
tensors[0].1 = "F16".to_string();
tensors[1].1 = "BF16".to_string();
let refs: Vec<(&str, &str, u64, u64)> = tensors
.iter()
.map(|(n, d, s, e)| (n.as_str(), d.as_str(), *s, *e))
.collect();
write_fake_safetensors(&dir.join("model.safetensors"), &refs);
let inv = inspect_safetensors_dir(&dir, &cfg).unwrap();
let expected = (tensors.len() as u64 - 2) * 64 + 2 * 128;
assert_eq!(inv.total_bytes, expected);
fs::remove_dir_all(&dir).ok();
}
#[test]
fn inspect_safetensors_dir_detects_missing_tensor() {
let dir = tempdir("st-missing");
let cfg = Qwen35Config::qwen35_0_8b();
let mut tensors = required_tensor_fixture(&cfg);
let (dropped_name, _, _, _) = tensors.pop().unwrap();
let refs: Vec<(&str, &str, u64, u64)> = tensors
.iter()
.map(|(n, d, s, e)| (n.as_str(), d.as_str(), *s, *e))
.collect();
write_fake_safetensors(&dir.join("model.safetensors"), &refs);
let inv = inspect_safetensors_dir(&dir, &cfg).unwrap();
assert_eq!(inv.missing_tensors, vec![dropped_name]);
fs::remove_dir_all(&dir).ok();
}
#[test]
fn inspect_safetensors_dir_flags_unsupported_dtype() {
let dir = tempdir("st-baddtype");
let cfg = Qwen35Config::qwen35_0_8b();
let mut tensors = required_tensor_fixture(&cfg);
tensors[0].1 = "I64".to_string();
let bad_name = tensors[0].0.clone();
let refs: Vec<(&str, &str, u64, u64)> = tensors
.iter()
.map(|(n, d, s, e)| (n.as_str(), d.as_str(), *s, *e))
.collect();
write_fake_safetensors(&dir.join("model.safetensors"), &refs);
let inv = inspect_safetensors_dir(&dir, &cfg).unwrap();
assert_eq!(inv.unsupported_dtypes.len(), 1);
assert!(inv.unsupported_dtypes[0].contains(&bad_name));
assert!(inv.unsupported_dtypes[0].contains("I64"));
fs::remove_dir_all(&dir).ok();
}
#[test]
fn inspect_safetensors_dir_prefers_single_file_over_index() {
let dir = tempdir("st-precedence");
let cfg = Qwen35Config::qwen35_0_8b();
let tensors = required_tensor_fixture(&cfg);
let refs: Vec<(&str, &str, u64, u64)> = tensors
.iter()
.map(|(n, d, s, e)| (n.as_str(), d.as_str(), *s, *e))
.collect();
write_fake_safetensors(&dir.join("model.safetensors"), &refs);
fs::write(dir.join("model.safetensors.index.json"), b"not valid json").unwrap();
let inv = inspect_safetensors_dir(&dir, &cfg).unwrap();
assert!(inv.missing_tensors.is_empty());
fs::remove_dir_all(&dir).ok();
}
fn write_fake_q4_file(
path: &Path,
version: u32,
shape: &[u64],
original_len: u64,
n_blocks: usize,
) {
let mut buf = Vec::new();
buf.extend_from_slice(b"KHQ4");
buf.extend_from_slice(&version.to_le_bytes());
buf.extend_from_slice(&(shape.len() as u32).to_le_bytes());
for s in shape {
buf.extend_from_slice(&s.to_le_bytes());
}
buf.extend_from_slice(&original_len.to_le_bytes());
buf.resize(buf.len() + n_blocks * 20, 0);
fs::write(path, buf).expect("write fake q4 file");
}
#[test]
fn detect_q4_quantization_label_accepts_v2_file() {
let dir = tempdir("q4-label-v2");
write_fake_q4_file(&dir.join("sample.q4"), 2, &[32], 32, 1);
let label = detect_q4_quantization_label(&dir).unwrap();
assert!(label.contains("Q4_0"));
fs::remove_dir_all(&dir).ok();
}
#[test]
fn detect_q4_quantization_label_rejects_legacy_v1_file() {
let dir = tempdir("q4-label-v1");
write_fake_q4_file(&dir.join("sample.q4"), 1, &[32], 32, 1);
let err = detect_q4_quantization_label(&dir).unwrap_err();
assert!(err.contains("legacy"));
fs::remove_dir_all(&dir).ok();
}
#[test]
fn inspect_q4_dir_uses_index_manifest_and_ignores_merged_cache_files() {
let dir = tempdir("q4-merge-safe");
write_fake_q4_file(&dir.join("layer0_qkv.q4"), 2, &[32], 32, 1);
write_fake_q4_file(&dir.join("layer0_z.q4"), 2, &[32], 32, 1);
write_fake_q4_file(&dir.join("merged_qkvz_0_100_50.q4"), 2, &[64], 64, 2);
let qkv_len = fs::metadata(dir.join("layer0_qkv.q4")).unwrap().len();
let z_len = fs::metadata(dir.join("layer0_z.q4")).unwrap().len();
let index = serde_json::json!([
{"name": "model.language_model.layers.0.linear_attn.in_proj_qkv.weight", "file": "layer0_qkv.q4", "quantized": true, "shape": [32], "numel": 32},
{"name": "model.language_model.layers.0.linear_attn.in_proj_z.weight", "file": "layer0_z.q4", "quantized": true, "shape": [32], "numel": 32},
]);
fs::write(
dir.join("quantize_index.json"),
serde_json::to_vec(&index).unwrap(),
)
.unwrap();
let cfg = Qwen35Config::qwen35_0_8b();
let inv = inspect_q4_dir(&dir, &cfg).unwrap();
assert_eq!(inv.total_bytes, 2 * (qkv_len + z_len));
fs::remove_dir_all(&dir).ok();
}
#[test]
fn inspect_q4_dir_resident_bytes_exceed_raw_on_disk_sum_for_expanded_tensors() {
let dir = tempdir("q4-resident-expansion");
write_fake_q4_file(
&dir.join("embed.q4"),
2,
&[32],
32,
10, );
write_fake_q4_file(&dir.join("proj_a.q4"), 2, &[32], 32, 10);
write_fake_q4_file(&dir.join("q_proj.q4"), 2, &[32], 32, 10);
let embed_len = fs::metadata(dir.join("embed.q4")).unwrap().len();
let proj_a_len = fs::metadata(dir.join("proj_a.q4")).unwrap().len();
let q_proj_len = fs::metadata(dir.join("q_proj.q4")).unwrap().len();
let index = serde_json::json!([
{"name": "model.language_model.embed_tokens.weight", "file": "embed.q4", "quantized": true, "shape": [32], "numel": 32},
{"name": "model.language_model.layers.0.linear_attn.in_proj_a.weight", "file": "proj_a.q4", "quantized": true, "shape": [32], "numel": 32},
{"name": "model.language_model.layers.0.self_attn.q_proj.weight", "file": "q_proj.q4", "quantized": true, "shape": [32], "numel": 32},
]);
fs::write(
dir.join("quantize_index.json"),
serde_json::to_vec(&index).unwrap(),
)
.unwrap();
let cfg = Qwen35Config::qwen35_0_8b();
let inv = inspect_q4_dir(&dir, &cfg).unwrap();
let naive_sum = embed_len + proj_a_len + q_proj_len;
let expected = (embed_len as f64 * 4.2).round() as u64
+ (proj_a_len as f64 * 3.2).round() as u64
+ q_proj_len;
assert!(
inv.total_bytes > naive_sum,
"resident estimate ({}) must exceed the naive on-disk sum ({naive_sum}) \
once embed_tokens/in_proj_a expansion is accounted for",
inv.total_bytes
);
assert_eq!(inv.total_bytes, expected);
fs::remove_dir_all(&dir).ok();
}
#[test]
fn q4_resident_bytes_classifies_every_known_category() {
assert_eq!(
q4_resident_bytes("model.language_model.norm.weight", 100, true),
200
);
assert_eq!(
q4_resident_bytes("model_language_model_norm_weight.f16", 100, true),
200
);
assert_eq!(
q4_resident_bytes("model.language_model.layers.0.linear_attn.A_log", 100, true),
200
);
assert_eq!(
q4_resident_bytes(
"model.language_model.layers.0.linear_attn.dt_bias",
100,
true
),
200
);
assert_eq!(
q4_resident_bytes(
"model.language_model.layers.0.linear_attn.conv1d.weight",
100,
true
),
200
);
assert_eq!(
q4_resident_bytes(
"model.language_model.layers.0.linear_attn.in_proj_a.weight",
100,
true
),
320
);
assert_eq!(
q4_resident_bytes(
"model.language_model.layers.0.linear_attn.in_proj_b.weight",
100,
true
),
320
);
assert_eq!(
q4_resident_bytes(
"model.language_model.layers.0.linear_attn.in_proj_qkv.weight",
100,
true
),
200
);
assert_eq!(
q4_resident_bytes(
"model.language_model.layers.0.linear_attn.in_proj_z.weight",
100,
true
),
200
);
assert_eq!(
q4_resident_bytes("model.language_model.embed_tokens.weight", 100, true),
420
);
assert_eq!(
q4_resident_bytes("mtp.pre_fc_norm_embedding.weight", 100, true),
200
);
assert_eq!(
q4_resident_bytes("mtp.pre_fc_norm_hidden.weight", 100, true),
200
);
assert_eq!(
q4_resident_bytes("mtp_pre_fc_norm_embedding_weight.f16", 100, true),
200
);
assert_eq!(
q4_resident_bytes("mtp_pre_fc_norm_hidden_weight.f16", 100, true),
200
);
assert_eq!(
q4_resident_bytes(
"model.language_model.layers.0.self_attn.q_proj.weight",
100,
true
),
100
);
assert_eq!(
q4_resident_bytes(
"model.language_model.layers.0.mlp.gate_proj.weight",
100,
false
),
100
);
assert_eq!(
q4_resident_bytes(
"model.language_model.layers.0.linear_attn.out_proj.weight",
100,
true
),
100
);
assert_eq!(q4_resident_bytes("lm_head.weight", 100, false), 100);
}
#[test]
fn q4_resident_bytes_embed_tokens_scales_by_tie_word_embeddings() {
let tensor = "model.language_model.embed_tokens.weight";
assert_eq!(q4_resident_bytes(tensor, 1000, true), 4200);
assert_eq!(q4_resident_bytes(tensor, 1000, false), 3200);
assert_eq!(q4_resident_bytes("lm_head.weight", 1000, true), 1000);
assert_eq!(q4_resident_bytes("lm_head.weight", 1000, false), 1000);
}
#[test]
fn inspect_q4_dir_manifest_flags_missing_file() {
let dir = tempdir("q4-missing-file");
let index = serde_json::json!([
{"name": "some.tensor.weight", "file": "does_not_exist.q4", "quantized": true, "shape": [32], "numel": 32},
]);
fs::write(
dir.join("quantize_index.json"),
serde_json::to_vec(&index).unwrap(),
)
.unwrap();
write_fake_q4_file(&dir.join("sample.q4"), 2, &[32], 32, 1);
let cfg = Qwen35Config::qwen35_0_8b();
let inv = inspect_q4_dir(&dir, &cfg).unwrap();
assert!(
inv.missing_tensors
.iter()
.any(|m| m.contains("some.tensor.weight"))
);
fs::remove_dir_all(&dir).ok();
}
#[test]
fn inspect_q4_dir_falls_back_to_directory_scan_without_manifest() {
let dir = tempdir("q4-no-manifest");
write_fake_q4_file(&dir.join("a.q4"), 2, &[32], 32, 1);
write_fake_q4_file(&dir.join("b.q4"), 2, &[32], 32, 1);
write_fake_q4_file(&dir.join("merged_qkvz_x.q4"), 2, &[64], 64, 2);
let a_len = fs::metadata(dir.join("a.q4")).unwrap().len();
let b_len = fs::metadata(dir.join("b.q4")).unwrap().len();
let cfg = Qwen35Config::qwen35_0_8b();
let inv = inspect_q4_dir(&dir, &cfg).unwrap();
assert_eq!(inv.total_bytes, a_len + b_len);
assert_eq!(inv.tensor_count, 2);
fs::remove_dir_all(&dir).ok();
}
#[test]
fn inspect_q4_dir_fallback_scan_applies_resident_bytes_by_sanitized_filename() {
let dir = tempdir("q4-no-manifest-embed");
write_fake_q4_file(
&dir.join("model_language_model_embed_tokens_weight.q4"),
2,
&[32],
32,
10,
);
write_fake_q4_file(
&dir.join("model_language_model_layers_0_self_attn_q_proj_weight.q4"),
2,
&[32],
32,
10,
);
let embed_len = fs::metadata(dir.join("model_language_model_embed_tokens_weight.q4"))
.unwrap()
.len();
let q_proj_len =
fs::metadata(dir.join("model_language_model_layers_0_self_attn_q_proj_weight.q4"))
.unwrap()
.len();
let cfg = Qwen35Config::qwen35_0_8b();
let inv = inspect_q4_dir(&dir, &cfg).unwrap();
let expected = (embed_len as f64 * 4.2).round() as u64 + q_proj_len;
assert_eq!(inv.total_bytes, expected);
assert!(inv.total_bytes > embed_len + q_proj_len);
fs::remove_dir_all(&dir).ok();
}
#[test]
fn inspect_q4_dir_detects_mtp_tensors_via_manifest() {
let dir = tempdir("q4-mtp-manifest");
write_fake_q4_file(&dir.join("q_proj.q4"), 2, &[32], 32, 1);
write_fake_q4_file(&dir.join("mtp_norm.f16"), 2, &[32], 32, 1);
let index = serde_json::json!([
{"name": "model.language_model.layers.0.self_attn.q_proj.weight", "file": "q_proj.q4", "quantized": true, "shape": [32], "numel": 32},
{"name": "mtp.norm.weight", "file": "mtp_norm.f16", "quantized": false, "shape": [32], "numel": 32},
]);
fs::write(
dir.join("quantize_index.json"),
serde_json::to_vec(&index).unwrap(),
)
.unwrap();
let cfg = Qwen35Config::qwen35_0_8b();
let inv = inspect_q4_dir(&dir, &cfg).unwrap();
assert!(inv.has_mtp_tensors);
fs::remove_dir_all(&dir).ok();
}
#[test]
fn inspect_q4_dir_no_mtp_tensors_via_manifest_when_absent() {
let dir = tempdir("q4-no-mtp-manifest");
write_fake_q4_file(&dir.join("q_proj.q4"), 2, &[32], 32, 1);
let index = serde_json::json!([
{"name": "model.language_model.layers.0.self_attn.q_proj.weight", "file": "q_proj.q4", "quantized": true, "shape": [32], "numel": 32},
]);
fs::write(
dir.join("quantize_index.json"),
serde_json::to_vec(&index).unwrap(),
)
.unwrap();
let cfg = Qwen35Config::qwen35_0_8b();
let inv = inspect_q4_dir(&dir, &cfg).unwrap();
assert!(!inv.has_mtp_tensors);
fs::remove_dir_all(&dir).ok();
}
#[test]
fn inspect_q4_dir_accepts_both_bare_array_and_quarot_object_manifest_shapes() {
let bare_dir = tempdir("q4-manifest-bare-array");
let object_dir = tempdir("q4-manifest-quarot-object");
for dir in [&bare_dir, &object_dir] {
write_fake_q4_file(&dir.join("q_proj.q4"), 2, &[32], 32, 10);
}
let q_proj_len = fs::metadata(bare_dir.join("q_proj.q4")).unwrap().len();
let tensors = serde_json::json!([
{"name": "model.language_model.layers.0.self_attn.q_proj.weight", "file": "q_proj.q4", "quantized": true, "shape": [32], "numel": 32},
]);
fs::write(
bare_dir.join("quantize_index.json"),
serde_json::to_vec(&tensors).unwrap(),
)
.unwrap();
let quarot_seed: u64 = 0xCAFE_BABE_DEAD_BEEF;
let object_manifest = serde_json::json!({
"quarot_seed": quarot_seed,
"tensors": [
{"name": "model.language_model.layers.0.self_attn.q_proj.weight", "file": "q_proj.q4", "quantized": true, "shape": [32], "numel": 32},
],
});
fs::write(
object_dir.join("quantize_index.json"),
serde_json::to_vec(&object_manifest).unwrap(),
)
.unwrap();
let cfg = Qwen35Config::qwen35_0_8b();
let bare_inv = inspect_q4_dir(&bare_dir, &cfg)
.expect("bare-array quantize_index.json (quantize_q4's shape) must load");
let object_inv = inspect_q4_dir(&object_dir, &cfg)
.expect("object-form quantize_index.json (quantize_quarot's shape) must load");
assert_eq!(object_inv.total_bytes, bare_inv.total_bytes);
assert_eq!(object_inv.total_bytes, q_proj_len);
assert_eq!(object_inv.tensor_count, bare_inv.tensor_count);
assert_eq!(object_inv.tensor_count, 1);
assert_eq!(object_inv.missing_tensors, bare_inv.missing_tensors);
assert_eq!(object_inv.has_mtp_tensors, bare_inv.has_mtp_tensors);
fs::remove_dir_all(&bare_dir).ok();
fs::remove_dir_all(&object_dir).ok();
}
#[test]
fn inspect_q4_dir_malformed_manifest_surfaces_precise_schema_error() {
let dir = tempdir("q4-manifest-malformed");
fs::write(dir.join("quantize_index.json"), br#"[{"name": "x"}]"#).unwrap();
let cfg = Qwen35Config::qwen35_0_8b();
let err = inspect_q4_dir(&dir, &cfg)
.err()
.expect("bare-array entry missing `file` must fail");
assert!(
err.contains("file"),
"error must name the missing `file` field; got: {err}"
);
assert!(
!err.contains("did not match any variant"),
"error must not be the generic untagged-enum fallthrough; got: {err}"
);
fs::remove_dir_all(&dir).ok();
}
#[test]
fn inspect_q4_dir_detects_mtp_tensors_via_fallback_scan() {
let dir = tempdir("q4-mtp-fallback");
write_fake_q4_file(
&dir.join("model_language_model_layers_0_self_attn_q_proj_weight.q4"),
2,
&[32],
32,
1,
);
write_fake_q4_file(
&dir.join("mtp_pre_fc_norm_embedding_weight.f16"),
2,
&[32],
32,
1,
);
let cfg = Qwen35Config::qwen35_0_8b();
let inv = inspect_q4_dir(&dir, &cfg).unwrap();
assert!(inv.has_mtp_tensors);
fs::remove_dir_all(&dir).ok();
}
#[test]
fn inspect_q4_dir_no_mtp_tensors_via_fallback_scan_when_absent() {
let dir = tempdir("q4-no-mtp-fallback");
write_fake_q4_file(
&dir.join("model_language_model_layers_0_self_attn_q_proj_weight.q4"),
2,
&[32],
32,
1,
);
let cfg = Qwen35Config::qwen35_0_8b();
let inv = inspect_q4_dir(&dir, &cfg).unwrap();
assert!(!inv.has_mtp_tensors);
fs::remove_dir_all(&dir).ok();
}
#[test]
fn detect_total_memory_bytes_is_plausible_when_known() {
if let Some(bytes) = detect_total_memory_bytes() {
assert!(bytes > 0);
assert!(bytes < (1u64 << 50)); }
}
fn write_config_json(dir: &Path, cfg: &Qwen35Config) {
let config_json = serde_json::json!({
"text_config": {
"hidden_size": cfg.hidden_size,
"num_hidden_layers": cfg.num_hidden_layers,
"vocab_size": cfg.vocab_size,
"intermediate_size": cfg.intermediate_size,
"num_attention_heads": cfg.num_attention_heads,
"num_key_value_heads": cfg.num_key_value_heads,
"head_dim": cfg.head_dim,
"rope_theta": cfg.rope_theta,
"partial_rotary_factor": cfg.partial_rotary_factor,
"linear_num_key_heads": cfg.linear_num_key_heads,
"linear_num_value_heads": cfg.linear_num_value_heads,
"linear_key_head_dim": cfg.linear_key_head_dim,
"linear_value_head_dim": cfg.linear_value_head_dim,
"linear_conv_kernel_dim": cfg.linear_conv_kernel_dim,
"tie_word_embeddings": cfg.tie_word_embeddings,
"max_position_embeddings": cfg.max_position_embeddings,
"eos_token_id": cfg.eos_token_id,
"full_attention_interval": cfg.full_attention_interval,
}
});
fs::write(
dir.join("config.json"),
serde_json::to_vec_pretty(&config_json).unwrap(),
)
.unwrap();
}
fn write_complete_safetensors_fixture(dir: &Path, cfg: &Qwen35Config) {
let tensors = required_tensor_fixture(cfg);
let refs: Vec<(&str, &str, u64, u64)> = tensors
.iter()
.map(|(n, d, s, e)| (n.as_str(), d.as_str(), *s, *e))
.collect();
write_fake_safetensors(&dir.join("model.safetensors"), &refs);
write_config_json(dir, cfg);
fs::write(dir.join("tokenizer.json"), b"{}").unwrap();
}
#[test]
fn build_report_happy_path_is_ready() {
let dir = tempdir("report-happy");
let cfg = Qwen35Config::qwen35_0_8b();
write_complete_safetensors_fixture(&dir, &cfg);
let report = build_report(&dir, None, Some(4096), Some(1u64 << 40)).unwrap();
assert!(
report.is_ready(),
"blocking reasons: {:?}",
report.blocking_reasons
);
assert_eq!(report.placement, Placement::Cpu);
assert_eq!(report.requested_fits, Some(true));
fs::remove_dir_all(&dir).ok();
}
#[test]
fn build_report_infeasible_context_via_memory_override_exits_not_ready() {
let dir = tempdir("report-infeasible");
let cfg = Qwen35Config::qwen35_0_8b();
write_complete_safetensors_fixture(&dir, &cfg);
let report = build_report(&dir, None, Some(1), Some(1024)).unwrap();
assert!(!report.is_ready());
assert_eq!(report.requested_fits, Some(false));
assert!(
report
.blocking_reasons
.iter()
.any(|r| r.contains("does not fit"))
);
fs::remove_dir_all(&dir).ok();
}
#[test]
fn build_report_context_exceeding_any_real_machine_is_not_ready() {
let dir = tempdir("report-absurd-context");
let cfg = Qwen35Config::qwen35_0_8b();
write_complete_safetensors_fixture(&dir, &cfg);
let report = build_report(&dir, None, Some(usize::MAX / 2), None).unwrap();
if report.available_memory_bytes.is_some() {
assert_eq!(report.requested_fits, Some(false));
assert!(!report.is_ready());
}
fs::remove_dir_all(&dir).ok();
}
#[test]
fn build_report_weights_exceed_ram_fails_without_explicit_context() {
let dir = tempdir("report-weights-exceed-ram-no-ctx");
let cfg = Qwen35Config::qwen35_0_8b();
write_complete_safetensors_fixture(&dir, &cfg);
let weight_bytes = required_tensor_fixture(&cfg).len() as u64 * 64;
let report = build_report(&dir, None, None, Some(weight_bytes - 1)).unwrap();
assert!(!report.is_ready());
assert_eq!(report.requested_context, None);
assert!(
report
.blocking_reasons
.iter()
.any(|r| r.contains("exceeds detected system memory")),
"blocking reasons: {:?}",
report.blocking_reasons
);
fs::remove_dir_all(&dir).ok();
}
#[test]
fn build_report_zero_feasible_context_fails_without_explicit_context() {
let dir = tempdir("report-zero-ctx-no-request");
let cfg = Qwen35Config::qwen35_0_8b();
write_complete_safetensors_fixture(&dir, &cfg);
let weight_bytes = required_tensor_fixture(&cfg).len() as u64 * 64;
let kv_bytes_per_token = cfg.kv_bytes_per_token(kv_cache_dtype_bytes()) as u64;
assert!(
kv_bytes_per_token > 1,
"fixture assumption: at least 2 bytes/token of KV cache"
);
let report = build_report(&dir, None, None, Some(weight_bytes + 1)).unwrap();
assert_eq!(report.max_context_len, Some(0));
assert!(!report.is_ready());
assert_eq!(report.requested_context, None);
assert!(
report
.blocking_reasons
.iter()
.any(|r| r.contains("max feasible context length is 0 tokens")),
"blocking reasons: {:?}",
report.blocking_reasons
);
fs::remove_dir_all(&dir).ok();
}
#[test]
fn build_report_fits_within_ram_is_ready_without_explicit_context() {
let dir = tempdir("report-fits-no-ctx");
let cfg = Qwen35Config::qwen35_0_8b();
write_complete_safetensors_fixture(&dir, &cfg);
let report = build_report(&dir, None, None, Some(1u64 << 40)).unwrap();
assert!(
report.is_ready(),
"blocking reasons: {:?}",
report.blocking_reasons
);
assert!(report.max_context_len.unwrap_or(0) > 0);
fs::remove_dir_all(&dir).ok();
}
fn write_minimal_q4_fixture(dir: &Path, cfg: &Qwen35Config) -> (u64, u64) {
write_fake_q4_file(&dir.join("embed.q4"), 2, &[32], 32, 5);
write_fake_q4_file(&dir.join("q_proj.q4"), 2, &[32], 32, 3);
let embed_bytes = fs::metadata(dir.join("embed.q4")).unwrap().len();
let q_proj_bytes = fs::metadata(dir.join("q_proj.q4")).unwrap().len();
let index = serde_json::json!([
{"name": "model.language_model.embed_tokens.weight", "file": "embed.q4", "quantized": true, "shape": [32], "numel": 32},
{"name": "model.language_model.layers.0.self_attn.q_proj.weight", "file": "q_proj.q4", "quantized": true, "shape": [32], "numel": 32},
]);
fs::write(
dir.join("quantize_index.json"),
serde_json::to_vec(&index).unwrap(),
)
.unwrap();
write_config_json(dir, cfg);
fs::write(dir.join("tokenizer.json"), b"{}").unwrap();
(embed_bytes, q_proj_bytes)
}
#[test]
fn build_report_q4_untied_embed_tokens_threshold_no_longer_overcounts() {
let dir = tempdir("report-q4-untied-threshold");
let mut cfg = Qwen35Config::qwen35_0_8b();
cfg.tie_word_embeddings = false;
let (embed_bytes, q_proj_bytes) = write_minimal_q4_fixture(&dir, &cfg);
let corrected_total = (embed_bytes as f64 * 3.2).round() as u64 + q_proj_bytes; let stale_would_have_been = (embed_bytes as f64 * 4.2).round() as u64 + q_proj_bytes; assert!(stale_would_have_been > corrected_total);
let threshold_override = corrected_total + (stale_would_have_been - corrected_total) / 2;
assert!(threshold_override > corrected_total);
assert!(threshold_override < stale_would_have_been);
let report = build_report(&dir, None, None, Some(threshold_override)).unwrap();
assert!(
!report
.blocking_reasons
.iter()
.any(|r| r.contains("exceeds detected system memory")),
"corrected Q4 estimate must fit within a budget between the corrected and \
stale flat-4.2x estimate; blocking reasons: {:?}",
report.blocking_reasons
);
fs::remove_dir_all(&dir).ok();
}
#[test]
fn build_report_q4_true_over_budget_still_fails() {
let dir = tempdir("report-q4-true-over-budget");
let mut cfg = Qwen35Config::qwen35_0_8b();
cfg.tie_word_embeddings = false;
let (embed_bytes, q_proj_bytes) = write_minimal_q4_fixture(&dir, &cfg);
let corrected_total = (embed_bytes as f64 * 3.2).round() as u64 + q_proj_bytes;
let report = build_report(&dir, None, None, Some(corrected_total - 1)).unwrap();
assert!(
report
.blocking_reasons
.iter()
.any(|r| r.contains("exceeds detected system memory")),
"a checkpoint exceeding even the corrected estimate must still fail closed; \
blocking reasons: {:?}",
report.blocking_reasons
);
fs::remove_dir_all(&dir).ok();
}
#[test]
fn build_report_missing_tokenizer_is_not_ready() {
let dir = tempdir("report-no-tokenizer");
let cfg = Qwen35Config::qwen35_0_8b();
let tensors = required_tensor_fixture(&cfg);
let refs: Vec<(&str, &str, u64, u64)> = tensors
.iter()
.map(|(n, d, s, e)| (n.as_str(), d.as_str(), *s, *e))
.collect();
write_fake_safetensors(&dir.join("model.safetensors"), &refs);
write_config_json(&dir, &cfg);
let report = build_report(&dir, None, None, Some(1u64 << 40)).unwrap();
assert!(!report.is_ready());
assert!(!report.tokenizer_present);
fs::remove_dir_all(&dir).ok();
}
#[test]
fn build_report_unknown_format_is_err() {
let dir = tempdir("report-unknown");
let err = build_report(&dir, None, None, None).unwrap_err();
assert!(err.contains("not a recognized model directory"));
fs::remove_dir_all(&dir).ok();
}
#[test]
fn build_report_safetensors_missing_config_json_is_hard_error() {
let dir = tempdir("report-safetensors-no-config");
let cfg = Qwen35Config::qwen35_0_8b();
let tensors = required_tensor_fixture(&cfg);
let refs: Vec<(&str, &str, u64, u64)> = tensors
.iter()
.map(|(n, d, s, e)| (n.as_str(), d.as_str(), *s, *e))
.collect();
write_fake_safetensors(&dir.join("model.safetensors"), &refs);
let err = build_report(&dir, None, None, Some(1u64 << 40)).unwrap_err();
assert!(
err.contains("config.json"),
"error must name the missing config.json, not a guessed preset: {err}"
);
fs::remove_dir_all(&dir).ok();
}
#[test]
fn build_report_q4_missing_config_json_is_hard_error() {
let dir = tempdir("report-q4-no-config");
write_fake_q4_file(&dir.join("embed.q4"), 2, &[32], 32, 5);
write_fake_q4_file(&dir.join("q_proj.q4"), 2, &[32], 32, 3);
let index = serde_json::json!([
{"name": "model.language_model.embed_tokens.weight", "file": "embed.q4", "quantized": true, "shape": [32], "numel": 32},
{"name": "model.language_model.layers.0.self_attn.q_proj.weight", "file": "q_proj.q4", "quantized": true, "shape": [32], "numel": 32},
]);
fs::write(
dir.join("quantize_index.json"),
serde_json::to_vec(&index).unwrap(),
)
.unwrap();
let err = build_report(&dir, None, None, Some(1u64 << 40)).unwrap_err();
assert!(
err.contains("config.json"),
"error must name the missing config.json, not a guessed preset: {err}"
);
fs::remove_dir_all(&dir).ok();
}
fn minimal_doctor_report(has_mtp_tensors: bool) -> DoctorReport {
DoctorReport {
model_dir: PathBuf::from("/fake/model"),
format: crate::backend::ModelFormat::Q4,
placement: Placement::Metal,
quantization: "Q4_0".to_string(),
tensor_count: 1,
weight_bytes: 100,
kv_bytes_per_token: 100,
max_position_embeddings: 4096,
metal_runtime_cache_cap: Some(METAL_RUNTIME_MAX_CACHE_LEN),
available_memory_bytes: Some(1u64 << 40),
max_context_len: Some(4096),
requested_context: None,
requested_fits: None,
tokenizer_path: PathBuf::from("/fake/model/tokenizer.json"),
tokenizer_present: true,
missing_tensors: Vec::new(),
blocking_reasons: Vec::new(),
has_mtp_tensors,
}
}
#[test]
fn doctor_report_display_includes_mtp_disclosure_when_mtp_tensors_present() {
let report = minimal_doctor_report(true);
let text = format!("{report}");
assert!(
text.contains("MTP"),
"report must disclose the uncounted MTP K/V cache when MTP tensors are present:\n{text}"
);
}
#[test]
fn doctor_report_display_omits_mtp_disclosure_when_no_mtp_tensors() {
let report = minimal_doctor_report(false);
let text = format!("{report}");
assert!(
!text.contains("MTP"),
"report must not mention MTP for a directory with no MTP tensors:\n{text}"
);
}
}