use serde::{Deserialize, Serialize};
use tracing::{debug, info, warn};
use super::fingerprint::ModelFingerprint;
use crate::core::hardware::HardwareProfile;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AutoQuantPlan {
pub base_bits: u8,
pub group_size: usize,
pub component_overrides: Vec<ComponentOverride>,
pub quant_method: String,
pub estimated_size_bytes: u64,
pub estimated_tok_per_sec: f64,
pub reasoning: String,
pub confidence: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ComponentOverride {
pub pattern: String,
pub bits: u8,
pub reason: String,
}
#[derive(Debug, Clone)]
pub struct AutoQuantConstraints {
pub min_tok_per_sec: f64,
pub quality_preference: QualityPreference,
pub override_bandwidth_gbps: Option<f64>,
pub forced_bits: Option<u8>,
}
impl Default for AutoQuantConstraints {
fn default() -> Self {
Self {
min_tok_per_sec: 80.0,
quality_preference: QualityPreference::Balanced,
override_bandwidth_gbps: None,
forced_bits: None,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum QualityPreference {
Speed,
Balanced,
Quality,
}
struct BandwidthProfile {
effective_bandwidth_gbps: f64,
}
fn estimate_bandwidth(hardware: &HardwareProfile) -> f64 {
let chip = hardware.chip_model.to_lowercase();
let _mem_gb = hardware.total_memory_gb();
let profile = if chip.contains("m5 ultra") {
BandwidthProfile {
effective_bandwidth_gbps: 780.0,
}
} else if chip.contains("m5 max") {
BandwidthProfile {
effective_bandwidth_gbps: 401.0,
}
} else if chip.contains("m5 pro") {
BandwidthProfile {
effective_bandwidth_gbps: 250.0,
}
} else if chip.contains("m5") {
BandwidthProfile {
effective_bandwidth_gbps: 100.0,
}
} else if chip.contains("m4 ultra") {
BandwidthProfile {
effective_bandwidth_gbps: 700.0,
}
} else if chip.contains("m4 max") {
BandwidthProfile {
effective_bandwidth_gbps: 370.0,
}
} else if chip.contains("m4 pro") {
BandwidthProfile {
effective_bandwidth_gbps: 230.0,
}
} else if chip.contains("m4") {
BandwidthProfile {
effective_bandwidth_gbps: 100.0,
}
} else if chip.contains("m3 ultra") {
BandwidthProfile {
effective_bandwidth_gbps: 600.0,
}
} else if chip.contains("m3 max") {
BandwidthProfile {
effective_bandwidth_gbps: 300.0,
}
} else if chip.contains("m3 pro") {
BandwidthProfile {
effective_bandwidth_gbps: 150.0,
}
} else if chip.contains("m3") {
BandwidthProfile {
effective_bandwidth_gbps: 100.0,
}
} else if chip.contains("m2 ultra") {
BandwidthProfile {
effective_bandwidth_gbps: 600.0,
}
} else if chip.contains("m2 max") {
BandwidthProfile {
effective_bandwidth_gbps: 300.0,
}
} else if chip.contains("m2 pro") {
BandwidthProfile {
effective_bandwidth_gbps: 150.0,
}
} else if chip.contains("m2") {
BandwidthProfile {
effective_bandwidth_gbps: 100.0,
}
} else if chip.contains("m1 ultra") {
BandwidthProfile {
effective_bandwidth_gbps: 500.0,
}
} else if chip.contains("m1 max") {
BandwidthProfile {
effective_bandwidth_gbps: 250.0,
}
} else if chip.contains("m1 pro") {
BandwidthProfile {
effective_bandwidth_gbps: 150.0,
}
} else if chip.contains("m1") {
BandwidthProfile {
effective_bandwidth_gbps: 60.0,
}
} else {
warn!(
chip = %hardware.chip_model,
"Unknown hardware — using conservative 100 GB/s bandwidth estimate. \
Pass --bandwidth to override."
);
BandwidthProfile {
effective_bandwidth_gbps: 100.0,
}
};
let gbps = profile.effective_bandwidth_gbps;
debug!(
chip = %hardware.chip_model,
bandwidth_gbps = gbps,
"Estimated effective memory bandwidth"
);
gbps * 1e9 }
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum ArchFamily {
GemmaMoE,
GenericMoE,
DenseDecoder,
Qwen35Dense,
Qwen35MoE,
Unknown,
}
fn classify_architecture(fingerprint: &ModelFingerprint) -> ArchFamily {
let arch = fingerprint.architecture.to_lowercase();
if arch.contains("qwen3_5") {
return if fingerprint.is_moe() {
ArchFamily::Qwen35MoE
} else {
ArchFamily::Qwen35Dense
};
}
if fingerprint.is_moe() {
if arch.contains("gemma") {
ArchFamily::GemmaMoE
} else {
ArchFamily::GenericMoE
}
} else if arch.contains("llama")
|| arch.contains("mistral")
|| arch.contains("qwen")
|| arch.contains("phi")
|| arch.contains("gemma")
|| arch.contains("starcoder")
|| arch.contains("codellama")
|| arch.contains("deepseek")
|| arch.contains("internlm")
|| arch.contains("yi")
|| arch.contains("command")
{
ArchFamily::DenseDecoder
} else {
ArchFamily::Unknown
}
}
fn estimate_bytes_per_token(
fingerprint: &ModelFingerprint,
plan_base_bits: u8,
component_overrides: &[ComponentOverride],
) -> u64 {
if !fingerprint.is_moe() {
return estimate_total_model_bytes(fingerprint, plan_base_bits, component_overrides);
}
let h = fingerprint.hidden_size as f64;
let i = fingerprint
.intermediate_size
.unwrap_or(fingerprint.hidden_size * 4) as f64;
let n_experts = fingerprint.expert_count as f64;
let n_layers = fingerprint.layer_count as f64;
let expert_params_per_layer = n_experts * 3.0 * h * i;
let num_kv_heads = fingerprint
.num_kv_heads
.unwrap_or(fingerprint.num_attention_heads) as f64;
let num_heads = fingerprint.num_attention_heads as f64;
let head_dim = h / num_heads;
let shared_attn_per_layer = h * (num_heads * head_dim) + h * (num_kv_heads * head_dim) + h * (num_kv_heads * head_dim) + (num_heads * head_dim) * h;
let total_expert_params = expert_params_per_layer * n_layers;
let total_shared_attn_params = shared_attn_per_layer * n_layers;
let embed_params = fingerprint.vocab_size as f64 * h * 2.0;
let total_shared_params = total_shared_attn_params + embed_params;
let top_k = if fingerprint.expert_count >= 64 {
8.0
} else {
2.0
};
let expert_activation_ratio = top_k / n_experts;
let shared_bytes = (total_shared_params * plan_base_bits as f64 / 8.0) as u64;
let expert_bytes_active =
(total_expert_params * expert_activation_ratio * plan_base_bits as f64 / 8.0) as u64;
debug!(
shared_gb = shared_bytes as f64 / 1e9,
expert_active_gb = expert_bytes_active as f64 / 1e9,
activation_ratio = expert_activation_ratio,
"MoE bytes-per-token breakdown"
);
shared_bytes + expert_bytes_active
}
fn estimate_total_model_bytes(
fingerprint: &ModelFingerprint,
base_bits: u8,
_component_overrides: &[ComponentOverride],
) -> u64 {
(fingerprint.total_params as f64 * base_bits as f64 / 8.0) as u64
}
pub fn resolve_auto_plan(
hardware: &HardwareProfile,
fingerprint: &ModelFingerprint,
constraints: &AutoQuantConstraints,
) -> Result<AutoQuantPlan, AutoQuantError> {
let arch_family = classify_architecture(fingerprint);
info!(
arch_family = ?arch_family,
is_moe = fingerprint.is_moe(),
params_b = fingerprint.total_params as f64 / 1e9,
"Auto-quant: classifying model"
);
let bandwidth_bps = match constraints.override_bandwidth_gbps {
Some(bw) => bw * 1e9,
None => estimate_bandwidth(hardware),
};
if let Some(forced) = constraints.forced_bits {
let overrides = build_component_overrides(arch_family, fingerprint, forced);
let total_bytes = estimate_total_model_bytes(fingerprint, forced, &overrides);
let bytes_per_token = estimate_bytes_per_token(fingerprint, forced, &overrides);
let est_tok_s = bandwidth_bps / bytes_per_token as f64;
return Ok(AutoQuantPlan {
base_bits: forced,
group_size: 64,
component_overrides: overrides,
quant_method: plan_to_quant_method(forced, arch_family, fingerprint, hardware),
estimated_size_bytes: total_bytes,
estimated_tok_per_sec: est_tok_s,
reasoning: format!(
"Forced to {}-bit by user. Estimated {:.0} tok/s.",
forced, est_tok_s
),
confidence: 0.6,
});
}
let candidates: &[u8] = match constraints.quality_preference {
QualityPreference::Speed => &[4, 3, 2],
QualityPreference::Balanced => &[8, 4, 3, 2],
QualityPreference::Quality => &[8, 6, 4, 3, 2],
};
let target_tok_s = constraints.min_tok_per_sec;
let memory_budget = (hardware.total_memory_bytes as f64 * 0.85) as u64;
let mut best_plan: Option<(u8, f64, u64)> = None;
for &bits in candidates {
let overrides = build_component_overrides(arch_family, fingerprint, bits);
let total_bytes = estimate_total_model_bytes(fingerprint, bits, &overrides);
let bytes_per_token = estimate_bytes_per_token(fingerprint, bits, &overrides);
let est_tok_s = bandwidth_bps / bytes_per_token as f64;
debug!(
bits = bits,
total_gb = total_bytes as f64 / 1e9,
bpt_gb = bytes_per_token as f64 / 1e9,
est_tok_s = est_tok_s,
"Auto-quant: evaluating {}-bit",
bits
);
if total_bytes > memory_budget {
debug!(bits = bits, "Skipping: model does not fit in memory");
continue;
}
if est_tok_s >= target_tok_s {
match constraints.quality_preference {
QualityPreference::Quality => {
best_plan = Some((bits, est_tok_s, total_bytes));
break;
}
QualityPreference::Balanced => {
best_plan = Some((bits, est_tok_s, total_bytes));
break;
}
QualityPreference::Speed => {
best_plan = Some((bits, est_tok_s, total_bytes));
break;
}
}
}
if best_plan.is_none() || bits < best_plan.unwrap().0 {
best_plan = Some((bits, est_tok_s, total_bytes));
}
}
let (base_bits, est_tok_s, total_bytes) =
best_plan.ok_or_else(|| AutoQuantError::ModelTooLarge {
reason: format!(
"Model ({:.1}B params) does not fit in {:.0} GB memory even at 2-bit quantization.",
fingerprint.total_params as f64 / 1e9,
hardware.total_memory_gb()
),
})?;
let overrides = build_component_overrides(arch_family, fingerprint, base_bits);
let quant_method = plan_to_quant_method(base_bits, arch_family, fingerprint, hardware);
let confidence = if est_tok_s >= target_tok_s * 1.2 {
0.85 } else if est_tok_s >= target_tok_s {
0.75 } else if est_tok_s >= target_tok_s * 0.8 {
0.6 } else {
0.45 };
let quality_note = match base_bits {
8 => "~97.7% token accuracy",
6 => "~96.9% token accuracy",
4 => "~90.5% token accuracy",
3 => "~85% token accuracy (estimated)",
2 => "~75% token accuracy (estimated, significant quality loss)",
_ => "unknown quality profile",
};
let reasoning = if est_tok_s >= target_tok_s {
format!(
"{}-bit base ({}) meets {:.0} tok/s target (est. {:.0} tok/s) on {} with {:.0} GB. \
Model size: {:.1} GB. {}{}",
base_bits,
quant_method,
target_tok_s,
est_tok_s,
hardware.chip_model,
hardware.total_memory_gb(),
total_bytes as f64 / 1e9,
quality_note,
if !overrides.is_empty() {
format!(
". {} component overrides applied for quality-critical layers.",
overrides.len()
)
} else {
String::new()
},
)
} else {
format!(
"Best achievable: {}-bit base ({}) yields ~{:.0} tok/s (target: {:.0}) on {} with {:.0} GB. \
Model size: {:.1} GB. {}. Consider a smaller model for the target throughput.",
base_bits,
quant_method,
est_tok_s,
target_tok_s,
hardware.chip_model,
hardware.total_memory_gb(),
total_bytes as f64 / 1e9,
quality_note,
)
};
info!(
base_bits = base_bits,
method = %quant_method,
est_tok_s = est_tok_s,
model_gb = total_bytes as f64 / 1e9,
overrides = overrides.len(),
"Auto-quant plan resolved"
);
Ok(AutoQuantPlan {
base_bits,
group_size: 64,
component_overrides: overrides,
quant_method,
estimated_size_bytes: total_bytes,
estimated_tok_per_sec: est_tok_s,
reasoning,
confidence,
})
}
fn build_component_overrides(
arch_family: ArchFamily,
fingerprint: &ModelFingerprint,
base_bits: u8,
) -> Vec<ComponentOverride> {
let mut overrides = Vec::new();
match arch_family {
ArchFamily::Qwen35Dense | ArchFamily::Qwen35MoE => {
for ssm_pattern in &[
".A_log",
".dt_bias",
".dt_proj.weight",
".dt_proj.bias",
".conv1d.weight",
] {
overrides.push(ComponentOverride {
pattern: ssm_pattern.to_string(),
bits: 8u8.max(base_bits),
reason: format!(
"ADR-012 D12 cohort prior: SSM state tensor '{ssm_pattern}' is numerically \
load-bearing (A_log exponentiated, dt drives time-step gate). \
Promoted to sensitive_bits unconditionally."
),
});
}
}
_ => {}
}
if arch_family == ArchFamily::Qwen35MoE {
for router_pattern in &[
".mlp.gate.weight", ".shared_expert.gate_proj", ".shared_expert.up_proj", ".shared_expert.down_proj", ] {
overrides.push(ComponentOverride {
pattern: router_pattern.to_string(),
bits: 8u8.max(base_bits),
reason: format!(
"ADR-012 D12 cohort prior (qwen35moe): '{router_pattern}' is a router or \
shared expert — always active, misrouting is unrecoverable. \
Promoted to sensitive_bits unconditionally."
),
});
}
let routed_bits = next_valid_bits(base_bits, 2).min(8);
if routed_bits > base_bits {
let exp_pattern = ".experts.";
overrides.push(ComponentOverride {
pattern: exp_pattern.to_string(),
bits: routed_bits,
reason: format!(
"ADR-012 D12 cohort prior (qwen35moe): routed expert tensor '{exp_pattern}' \
uses activation-score-driven heuristic with elevated threshold. \
Default: +2 bits above base ({base_bits}-bit \u{2192} {routed_bits}-bit). \
Tunable by future calibration (see ADR-012 D12)."
),
});
}
}
if fingerprint.is_moe() && !matches!(arch_family, ArchFamily::Qwen35MoE) {
overrides.push(ComponentOverride {
pattern: "router.proj".to_string(),
bits: 8.max(base_bits),
reason: "Router misrouting is catastrophic in MoE models".to_string(),
});
}
if base_bits <= 4 {
for component in &["mlp.gate_proj", "mlp.up_proj", "mlp.down_proj"] {
overrides.push(ComponentOverride {
pattern: component.to_string(),
bits: 8,
reason: format!(
"MLP {} at 8-bit eliminates FFN-induced generation artifacts ({}-bit base)",
component, base_bits
),
});
}
}
let v_proj_bits = next_valid_bits(base_bits, 2);
if v_proj_bits > base_bits {
overrides.push(ComponentOverride {
pattern: "v_proj".to_string(),
bits: v_proj_bits,
reason: format!(
"v_proj has the highest per-bit quality impact ({}-bit vs {}-bit base)",
v_proj_bits, base_bits
),
});
}
if base_bits <= 3 {
let elevated = next_valid_bits(base_bits, 2);
overrides.push(ComponentOverride {
pattern: "layers.0.".to_string(),
bits: elevated,
reason: "First transformer layer is disproportionately sensitive".to_string(),
});
let last_layer = fingerprint.layer_count.saturating_sub(1);
overrides.push(ComponentOverride {
pattern: format!("layers.{}.", last_layer),
bits: elevated,
reason: "Last transformer layer directly feeds lm_head".to_string(),
});
}
overrides
}
fn next_valid_bits(base: u8, steps: u8) -> u8 {
const VALID: [u8; 5] = [2, 3, 4, 6, 8];
let target = base + steps;
for &v in &VALID {
if v >= target {
return v;
}
}
8
}
fn plan_to_quant_method(
base_bits: u8,
_arch_family: ArchFamily,
fingerprint: &ModelFingerprint,
hardware: &HardwareProfile,
) -> String {
if let Some(entry_override) =
crate::arch::registry::lookup_auto_override(&fingerprint.architecture)
{
return entry_override;
}
let total_gb = hardware.total_memory_gb();
if fingerprint.is_moe() {
if total_gb >= 96.0 {
return "dynamic-quant-4-8".to_string();
}
return "dynamic-quant-4-6".to_string();
}
let params_b = fingerprint.total_params as f64 / 1e9;
if params_b > 30.0 && total_gb >= 64.0 {
return "imatrix-q5_k_m".to_string();
}
if params_b <= 30.0 || total_gb < 64.0 {
match base_bits {
16 => return "f16".to_string(),
8 => return "q8".to_string(),
_ => return "imatrix-q4_k_m".to_string(),
}
}
"imatrix-q4_k_m".to_string()
}
#[derive(Debug, thiserror::Error)]
pub enum AutoQuantError {
#[error("Model too large for available hardware: {reason}")]
ModelTooLarge { reason: String },
#[error("Hardware detection failed: {reason}")]
#[allow(dead_code)]
HardwareDetection { reason: String },
}
#[allow(dead_code)] pub fn plan_to_config_json(plan: &AutoQuantPlan, hardware: &HardwareProfile) -> serde_json::Value {
let mut overrides_map = serde_json::Map::new();
for ov in &plan.component_overrides {
overrides_map.insert(
ov.pattern.clone(),
serde_json::json!({
"bits": ov.bits,
"reason": ov.reason,
}),
);
}
serde_json::json!({
"quant_method": plan.quant_method,
"bits": plan.base_bits,
"group_size": plan.group_size,
"component_overrides": overrides_map,
"auto_resolved": true,
"estimated_tok_per_sec": (plan.estimated_tok_per_sec * 10.0).round() / 10.0,
"estimated_size_bytes": plan.estimated_size_bytes,
"hardware": hardware.chip_model,
"confidence": (plan.confidence * 100.0).round() / 100.0,
"reasoning": plan.reasoning,
})
}
#[cfg(test)]
mod tests {
use super::*;
fn make_hardware(chip: &str, total_gb: u64) -> HardwareProfile {
HardwareProfile {
chip_model: chip.to_string(),
total_memory_bytes: total_gb * 1024 * 1024 * 1024,
available_memory_bytes: (total_gb as f64 * 0.8) as u64 * 1024 * 1024 * 1024,
performance_cores: 14,
efficiency_cores: 4,
total_cores: 18,
memory_bandwidth_gbs: crate::core::hardware::lookup_memory_bandwidth_gbs(chip),
}
}
fn make_dense_fingerprint(param_billions: f64) -> ModelFingerprint {
ModelFingerprint {
architecture: "LlamaForCausalLM".to_string(),
total_params: (param_billions * 1e9) as u64,
layer_count: 32,
expert_count: 0,
attention_types: vec!["attention".to_string()],
hidden_size: 4096,
dtype: "bfloat16".to_string(),
intermediate_size: Some(14336),
num_attention_heads: 32,
num_kv_heads: Some(8),
vocab_size: 128256,
}
}
fn make_moe_fingerprint() -> ModelFingerprint {
ModelFingerprint {
architecture: "Gemma4ForConditionalGeneration".to_string(),
total_params: 27_000_000_000,
layer_count: 30,
expert_count: 128,
attention_types: vec![
"full_attention".to_string(),
"sliding_attention".to_string(),
],
hidden_size: 2816,
dtype: "bfloat16".to_string(),
intermediate_size: Some(2112),
num_attention_heads: 16,
num_kv_heads: Some(8),
vocab_size: 262144,
}
}
#[test]
fn test_classify_architecture() {
let dense = make_dense_fingerprint(8.0);
assert_eq!(classify_architecture(&dense), ArchFamily::DenseDecoder);
let moe = make_moe_fingerprint();
assert_eq!(classify_architecture(&moe), ArchFamily::GemmaMoE);
}
fn make_qwen35_dense_fingerprint() -> ModelFingerprint {
ModelFingerprint {
architecture: "Qwen3_5ForCausalLM".to_string(),
total_params: 32_000_000_000,
layer_count: 64,
expert_count: 0,
attention_types: vec!["full_attention".to_string(), "linear_attention".to_string()],
hidden_size: 7168,
dtype: "bfloat16".to_string(),
intermediate_size: Some(18432),
num_attention_heads: 64,
num_kv_heads: Some(8),
vocab_size: 152064,
}
}
fn make_qwen35moe_fingerprint() -> ModelFingerprint {
ModelFingerprint {
architecture: "Qwen3_5MoeForCausalLM".to_string(),
total_params: 235_000_000_000,
layer_count: 94,
expert_count: 128,
attention_types: vec!["full_attention".to_string(), "linear_attention".to_string()],
hidden_size: 7168,
dtype: "bfloat16".to_string(),
intermediate_size: Some(2048),
num_attention_heads: 64,
num_kv_heads: Some(4),
vocab_size: 152064,
}
}
#[test]
fn test_classify_qwen35_dense() {
let fp = make_qwen35_dense_fingerprint();
assert_eq!(classify_architecture(&fp), ArchFamily::Qwen35Dense);
}
#[test]
fn test_classify_qwen35moe() {
let fp = make_qwen35moe_fingerprint();
assert_eq!(classify_architecture(&fp), ArchFamily::Qwen35MoE);
}
#[test]
fn test_qwen35_dense_ssm_cohort_priors_always_promoted() {
let fp = make_qwen35_dense_fingerprint();
let overrides = build_component_overrides(ArchFamily::Qwen35Dense, &fp, 4);
for ssm in &[
".A_log",
".dt_bias",
".dt_proj.weight",
".dt_proj.bias",
".conv1d.weight",
] {
let found = overrides.iter().find(|o| o.pattern == *ssm);
assert!(found.is_some(), "qwen35 dense cohort prior missing: {ssm}");
assert!(
found.unwrap().bits >= 4u8,
"SSM prior must be >= base_bits for {ssm}"
);
assert_eq!(
found.unwrap().bits,
8,
"SSM prior for {ssm} must be 8-bit (sensitive_bits)"
);
}
}
#[test]
fn test_qwen35moe_ssm_cohort_priors_always_promoted() {
let fp = make_qwen35moe_fingerprint();
let overrides = build_component_overrides(ArchFamily::Qwen35MoE, &fp, 4);
for ssm in &[
".A_log",
".dt_bias",
".dt_proj.weight",
".dt_proj.bias",
".conv1d.weight",
] {
let found = overrides.iter().find(|o| o.pattern == *ssm);
assert!(found.is_some(), "qwen35moe SSM cohort prior missing: {ssm}");
assert_eq!(found.unwrap().bits, 8, "SSM prior for {ssm} must be 8-bit");
}
}
#[test]
fn test_qwen35moe_router_and_shared_expert_promoted() {
let fp = make_qwen35moe_fingerprint();
let overrides = build_component_overrides(ArchFamily::Qwen35MoE, &fp, 4);
for moe_pattern in &[
".mlp.gate.weight",
".shared_expert.gate_proj",
".shared_expert.up_proj",
".shared_expert.down_proj",
] {
let found = overrides.iter().find(|o| o.pattern == *moe_pattern);
assert!(
found.is_some(),
"qwen35moe router/shared-expert prior missing: {moe_pattern}"
);
assert_eq!(
found.unwrap().bits,
8,
"Router/shared-expert prior for {moe_pattern} must be 8-bit"
);
}
}
#[test]
fn test_qwen35moe_routed_experts_elevated_threshold() {
let fp = make_qwen35moe_fingerprint();
let overrides = build_component_overrides(ArchFamily::Qwen35MoE, &fp, 4);
let found = overrides.iter().find(|o| o.pattern == ".experts.");
assert!(
found.is_some(),
"qwen35moe routed-expert elevated prior missing"
);
assert_eq!(
found.unwrap().bits,
6,
"Routed expert prior must be base+2 = 6-bit at 4-bit base"
);
}
#[test]
fn test_qwen35_dense_no_moe_priors() {
let fp = make_qwen35_dense_fingerprint();
let overrides = build_component_overrides(ArchFamily::Qwen35Dense, &fp, 4);
for moe_only in &[".mlp.gate.weight", ".shared_expert.", ".experts."] {
assert!(
!overrides.iter().any(|o| o.pattern == *moe_only),
"qwen35 dense must not have MoE-only prior: {moe_only}"
);
}
}
#[test]
fn test_gemma4_regression_cohort_priors_not_added() {
let fp = make_moe_fingerprint(); assert_eq!(classify_architecture(&fp), ArchFamily::GemmaMoE);
let overrides_before = vec![
"router.proj".to_string(),
"mlp.gate_proj".to_string(),
"mlp.up_proj".to_string(),
"mlp.down_proj".to_string(),
"v_proj".to_string(),
];
let overrides = build_component_overrides(ArchFamily::GemmaMoE, &fp, 4);
let patterns: Vec<&str> = overrides.iter().map(|o| o.pattern.as_str()).collect();
for qwen35_only in &[
".A_log",
".dt_bias",
".dt_proj.weight",
".dt_proj.bias",
".conv1d.weight",
".mlp.gate.weight",
".shared_expert.",
".experts.",
] {
assert!(
!patterns.contains(qwen35_only),
"Gemma regression: qwen35-only pattern '{qwen35_only}' appeared in Gemma overrides"
);
}
for expected_pat in &overrides_before {
assert!(
patterns.contains(&expected_pat.as_str()),
"Gemma regression: expected pattern '{expected_pat}' missing from overrides"
);
}
assert_eq!(
overrides.len(),
overrides_before.len(),
"Gemma regression: override count changed — pre-P6={}, post-P6={}",
overrides_before.len(),
overrides.len()
);
}
#[test]
fn test_sensitive_layers_additive_with_cohort_priors() {
let fp = make_qwen35moe_fingerprint();
let overrides = build_component_overrides(ArchFamily::Qwen35MoE, &fp, 4);
assert!(
overrides.iter().any(|o| o.pattern == ".A_log"),
"SSM cohort prior must be present"
);
assert!(
overrides.iter().any(|o| o.pattern == "v_proj"),
"v_proj override must still be present alongside cohort priors"
);
}
#[test]
fn test_non_qwen35_dense_no_ssm_priors() {
let fp = make_dense_fingerprint(8.0); let overrides = build_component_overrides(ArchFamily::DenseDecoder, &fp, 4);
for ssm in &[
".A_log",
".dt_bias",
".dt_proj.weight",
".dt_proj.bias",
".conv1d.weight",
] {
assert!(
!overrides.iter().any(|o| o.pattern == *ssm),
"Non-qwen35 dense must not have SSM cohort prior: {ssm}"
);
}
}
#[test]
fn test_bandwidth_estimation_known_chip() {
let hw = make_hardware("Apple M5 Max", 128);
let bw = estimate_bandwidth(&hw);
assert!((bw / 1e9 - 401.0).abs() < 1.0);
}
#[test]
fn test_bandwidth_estimation_unknown_chip() {
let hw = make_hardware("Unknown Chip XYZ", 64);
let bw = estimate_bandwidth(&hw);
assert!((bw / 1e9 - 100.0).abs() < 1.0); }
#[test]
fn test_small_dense_model_on_large_machine() {
let hw = make_hardware("Apple M5 Max", 128);
let fp = make_dense_fingerprint(8.0);
let constraints = AutoQuantConstraints::default();
let plan = resolve_auto_plan(&hw, &fp, &constraints).unwrap();
assert!(plan.base_bits <= 8);
assert!(plan.estimated_tok_per_sec >= 50.0);
assert!(!plan.reasoning.is_empty());
}
#[test]
fn test_moe_model_has_router_override() {
let hw = make_hardware("Apple M5 Max", 128);
let fp = make_moe_fingerprint();
let constraints = AutoQuantConstraints::default();
let plan = resolve_auto_plan(&hw, &fp, &constraints).unwrap();
let router_override = plan
.component_overrides
.iter()
.find(|o| o.pattern == "router.proj");
assert!(
router_override.is_some(),
"MoE plan must include router.proj override"
);
assert_eq!(router_override.unwrap().bits, 8);
}
#[test]
fn test_moe_bytes_per_token_less_than_total() {
let fp = make_moe_fingerprint();
let overrides = build_component_overrides(ArchFamily::GemmaMoE, &fp, 4);
let bpt = estimate_bytes_per_token(&fp, 4, &overrides);
let total = estimate_total_model_bytes(&fp, 4, &overrides);
assert!(
bpt < total,
"MoE bytes_per_token ({}) should be less than total ({})",
bpt,
total
);
}
#[test]
fn test_dense_bytes_per_token_equals_total() {
let fp = make_dense_fingerprint(8.0);
let overrides = build_component_overrides(ArchFamily::DenseDecoder, &fp, 4);
let bpt = estimate_bytes_per_token(&fp, 4, &overrides);
let total = estimate_total_model_bytes(&fp, 4, &overrides);
assert_eq!(bpt, total);
}
#[test]
fn test_quality_preference_affects_bits() {
let hw = make_hardware("Apple M5 Max", 128);
let fp = make_dense_fingerprint(8.0);
let speed_plan = resolve_auto_plan(
&hw,
&fp,
&AutoQuantConstraints {
quality_preference: QualityPreference::Speed,
..Default::default()
},
)
.unwrap();
let quality_plan = resolve_auto_plan(
&hw,
&fp,
&AutoQuantConstraints {
quality_preference: QualityPreference::Quality,
..Default::default()
},
)
.unwrap();
assert!(quality_plan.base_bits >= speed_plan.base_bits);
}
#[test]
fn test_huge_model_too_large_for_small_machine() {
let hw = make_hardware("Apple M4", 16);
let fp = make_dense_fingerprint(405.0);
let constraints = AutoQuantConstraints::default();
let result = resolve_auto_plan(&hw, &fp, &constraints);
assert!(result.is_err());
}
#[test]
fn test_forced_bits_respected() {
let hw = make_hardware("Apple M5 Max", 128);
let fp = make_dense_fingerprint(8.0);
let constraints = AutoQuantConstraints {
forced_bits: Some(3),
..Default::default()
};
let plan = resolve_auto_plan(&hw, &fp, &constraints).unwrap();
assert_eq!(plan.base_bits, 3);
}
#[test]
fn test_v_proj_override_present_at_4bit() {
let fp = make_dense_fingerprint(8.0);
let overrides = build_component_overrides(ArchFamily::DenseDecoder, &fp, 4);
let v_proj = overrides.iter().find(|o| o.pattern == "v_proj");
assert!(v_proj.is_some(), "4-bit plan should elevate v_proj");
assert_eq!(v_proj.unwrap().bits, 6); }
#[test]
fn test_mlp_8bit_overrides_at_4bit_base() {
let fp = make_dense_fingerprint(8.0);
let overrides = build_component_overrides(ArchFamily::DenseDecoder, &fp, 4);
for component in &["mlp.gate_proj", "mlp.up_proj", "mlp.down_proj"] {
let found = overrides.iter().find(|o| o.pattern == *component);
assert!(
found.is_some(),
"4-bit plan should elevate {} to 8-bit",
component
);
assert_eq!(found.unwrap().bits, 8, "{} should be 8-bit", component);
}
}
#[test]
fn test_no_mlp_override_at_8bit_base() {
let fp = make_dense_fingerprint(8.0);
let overrides = build_component_overrides(ArchFamily::DenseDecoder, &fp, 8);
let mlp = overrides.iter().find(|o| o.pattern.contains("mlp."));
assert!(mlp.is_none(), "8-bit base should not need MLP overrides");
}
#[test]
fn test_aggressive_quant_protects_first_last_layers() {
let fp = make_dense_fingerprint(70.0);
let overrides = build_component_overrides(ArchFamily::DenseDecoder, &fp, 2);
let first = overrides.iter().find(|o| o.pattern == "layers.0.");
let last = overrides.iter().find(|o| o.pattern == "layers.31.");
assert!(first.is_some(), "2-bit should protect first layer");
assert!(last.is_some(), "2-bit should protect last layer");
}
#[test]
fn test_plan_to_config_json_structure() {
let hw = make_hardware("Apple M5 Max", 128);
let fp = make_moe_fingerprint();
let constraints = AutoQuantConstraints::default();
let plan = resolve_auto_plan(&hw, &fp, &constraints).unwrap();
let json = plan_to_config_json(&plan, &hw);
assert!(json.get("quant_method").is_some());
assert!(json.get("bits").is_some());
assert!(json.get("group_size").is_some());
assert!(json.get("component_overrides").is_some());
assert_eq!(json["auto_resolved"], true);
assert!(json["hardware"].as_str().unwrap().contains("M5 Max"));
}
#[test]
fn test_plan_to_quant_method() {
let hw = make_hardware("Apple M5 Max", 128);
let fp_dense = make_dense_fingerprint(8.0);
let fp_moe = make_moe_fingerprint();
assert_eq!(
plan_to_quant_method(4, ArchFamily::DenseDecoder, &fp_dense, &hw),
"imatrix-q4_k_m"
);
assert_eq!(
plan_to_quant_method(8, ArchFamily::DenseDecoder, &fp_dense, &hw),
"q8"
);
assert_eq!(
plan_to_quant_method(16, ArchFamily::DenseDecoder, &fp_dense, &hw),
"f16"
);
assert_eq!(
plan_to_quant_method(2, ArchFamily::GemmaMoE, &fp_moe, &hw),
"dynamic-quant-4-8"
);
}
#[test]
fn test_decision18_dense_27b_resolves_to_imatrix_q4_k_m() {
let hw = make_hardware("Apple M5 Max", 128);
let mut fp = make_dense_fingerprint(27.0);
fp.architecture = "LlamaForCausalLM".to_string();
let result = plan_to_quant_method(4, ArchFamily::DenseDecoder, &fp, &hw);
assert_eq!(
result, "imatrix-q4_k_m",
"Dense 27B (≤30B) on any RAM → imatrix-q4_k_m, got: {result}"
);
}
#[test]
fn test_decision18_dense_70b_64gb_resolves_to_imatrix_q5_k_m() {
let hw = make_hardware("Apple M5 Max", 64);
let mut fp = make_dense_fingerprint(70.0);
fp.architecture = "LlamaForCausalLM".to_string();
let result = plan_to_quant_method(4, ArchFamily::DenseDecoder, &fp, &hw);
assert_eq!(
result, "imatrix-q5_k_m",
"Dense 70B with 64 GB RAM → imatrix-q5_k_m, got: {result}"
);
}
#[test]
fn test_decision18_dense_70b_below_64gb_resolves_to_imatrix_q4_k_m() {
let hw = make_hardware("Apple M4", 32);
let mut fp = make_dense_fingerprint(70.0);
fp.architecture = "LlamaForCausalLM".to_string();
let result = plan_to_quant_method(4, ArchFamily::DenseDecoder, &fp, &hw);
assert_eq!(
result, "imatrix-q4_k_m",
"Dense 70B with 32 GB RAM → imatrix-q4_k_m, got: {result}"
);
}
#[test]
fn test_decision18_moe_apex_64gb_resolves_to_dwq_4_6() {
let hw = make_hardware("Apple M5 Max", 64);
let fp = make_moe_fingerprint();
let result = plan_to_quant_method(4, ArchFamily::GemmaMoE, &fp, &hw);
assert_eq!(
result, "dynamic-quant-4-6",
"MoE with 64 GB RAM (<96 GB) → dynamic-quant-4-6, got: {result}"
);
}
#[test]
fn test_decision18_moe_apex_128gb_resolves_to_dwq_4_8() {
let hw = make_hardware("Apple M5 Max", 128);
let fp = make_moe_fingerprint();
let result = plan_to_quant_method(4, ArchFamily::GemmaMoE, &fp, &hw);
assert_eq!(
result, "dynamic-quant-4-8",
"MoE with 128 GB RAM (≥96 GB) → dynamic-quant-4-8, got: {result}"
);
}
#[test]
fn test_decision18_arch_override_lookup_falls_through_when_none() {
let lookup = crate::arch::registry::lookup_auto_override("Qwen3_5ForCausalLM");
assert!(
lookup.is_none(),
"qwen35 ArchEntry has auto_override = None → lookup must return None, \
got: {lookup:?}"
);
let lookup = crate::arch::registry::lookup_auto_override("CompletelyUnknownArch");
assert!(
lookup.is_none(),
"Unknown arch must return None (falls through to table)"
);
let hw = make_hardware("Apple M5 Max", 128);
let mut fp = make_qwen35_dense_fingerprint();
fp.total_params = 27_000_000_000; let result = plan_to_quant_method(4, ArchFamily::Qwen35Dense, &fp, &hw);
assert_eq!(
result, "imatrix-q4_k_m",
"qwen35 dense 27B with no override → imatrix-q4_k_m"
);
}
}