use super::super::ggml_type::GgmlType;
use super::super::tensor_ref::{ArchName, TensorRef};
use super::arches::{classify_moe_tensor, is_apex_supported_arch, MoeTensorRole};
use super::error::ApexError;
use super::mudler_config::MudlerConfig;
use super::rules::{
attn_region, exp_region, shared_region, tier_rules, ApexTier, AttnRegion, ExpRegion,
SharedRegion,
};
pub const SUPPORTED_FOR_IMATRIX: &[&str] = &["qwen3moe", "qwen35moe", "gemma4"];
#[derive(Debug, Clone, Copy)]
pub struct ApexPolicy {
pub tier: ApexTier,
pub n_layers: u32,
pub n_expert: u32,
pub arch: ArchName,
pub mudler_override: Option<&'static MudlerConfig>,
}
impl ApexPolicy {
pub fn new(
tier: ApexTier,
arch: ArchName,
n_layers: u32,
n_expert: u32,
) -> Result<Self, ApexError> {
if !is_apex_supported_arch(arch) {
return Err(ApexError::unsupported_arch(arch));
}
if n_layers == 0 {
return Err(ApexError::MissingHParam {
hparam: "num_hidden_layers",
});
}
if n_expert <= 1 {
return Err(ApexError::DenseModelNotSupported {
arch: arch.name(),
n_expert,
});
}
if tier.requires_imatrix() {
return Err(ApexError::ImatrixRequiresInference {
tier: tier.cli_name(),
arch: arch.name(),
supported_for_imatrix: SUPPORTED_FOR_IMATRIX,
});
}
Ok(Self {
tier,
n_layers,
n_expert,
arch,
mudler_override: None,
})
}
pub fn new_with_imatrix(
tier: ApexTier,
arch: ArchName,
n_layers: u32,
n_expert: u32,
) -> Result<Self, ApexError> {
if !is_apex_supported_arch(arch) {
return Err(ApexError::unsupported_arch(arch));
}
if n_layers == 0 {
return Err(ApexError::MissingHParam {
hparam: "num_hidden_layers",
});
}
if n_expert <= 1 {
return Err(ApexError::DenseModelNotSupported {
arch: arch.name(),
n_expert,
});
}
if tier.requires_imatrix() && !SUPPORTED_FOR_IMATRIX.contains(&arch.name()) {
return Err(ApexError::ImatrixRequiresInference {
tier: tier.cli_name(),
arch: arch.name(),
supported_for_imatrix: SUPPORTED_FOR_IMATRIX,
});
}
Ok(Self {
tier,
n_layers,
n_expert,
arch,
mudler_override: None,
})
}
pub fn with_mudler_override(mut self, mudler: &'static MudlerConfig) -> Self {
self.mudler_override = Some(mudler);
self
}
pub fn target_for(&self, tensor: &TensorRef) -> Result<GgmlType, ApexError> {
if let Some(mudler) = self.mudler_override {
if mudler.contains_match(tensor.name) {
return mudler.target_for(tensor.name);
}
let role = classify_moe_tensor(self.arch, tensor.name);
match role {
MoeTensorRole::TokenEmbd
| MoeTensorRole::Output
| MoeTensorRole::RouterGate
| MoeTensorRole::Norm => {
}
MoeTensorRole::RoutedExpert
| MoeTensorRole::SharedExpert
| MoeTensorRole::Attention
| MoeTensorRole::Ssm
| MoeTensorRole::Other => {
return Err(ApexError::TensorNotInMudlerConfig {
source_path: mudler.source_path.to_string(),
tensor_name: tensor.name.to_string(),
});
}
}
}
let role = classify_moe_tensor(self.arch, tensor.name);
let rules = tier_rules(self.tier);
match role {
MoeTensorRole::TokenEmbd => Ok(GgmlType::Q6_K),
MoeTensorRole::Output => Ok(GgmlType::Q6_K),
MoeTensorRole::RouterGate => Ok(GgmlType::Q5_0),
MoeTensorRole::Norm => Ok(GgmlType::F32),
MoeTensorRole::RoutedExpert => {
let layer = self.require_layer_index(tensor)?;
Ok(match exp_region(layer, self.n_layers) {
ExpRegion::Edge => rules.edge_exp,
ExpRegion::Near => rules.near_exp,
ExpRegion::Mid => rules.mid_exp,
})
}
MoeTensorRole::SharedExpert => {
let layer = self.require_layer_index(tensor)?;
Ok(match shared_region(layer, self.n_layers) {
SharedRegion::Edge => rules.edge_shared,
SharedRegion::Mid => rules.mid_shared,
})
}
MoeTensorRole::Attention => {
let layer = self.require_layer_index(tensor)?;
Ok(match attn_region(layer, self.n_layers) {
AttnRegion::Edge => rules.edge_attn,
AttnRegion::Mid => rules.mid_attn,
})
}
MoeTensorRole::Ssm => {
let layer = self.require_layer_index(tensor)?;
Ok(match attn_region(layer, self.n_layers) {
AttnRegion::Edge => rules.edge_attn,
AttnRegion::Mid => rules.mid_attn,
})
}
MoeTensorRole::Other => {
let region = match tensor.layer_index {
Some(l) => attn_region(l, self.n_layers),
None => AttnRegion::Mid,
};
Ok(match region {
AttnRegion::Edge => rules.edge_attn,
AttnRegion::Mid => rules.mid_attn,
})
}
}
}
fn require_layer_index(&self, tensor: &TensorRef) -> Result<usize, ApexError> {
let layer = tensor
.layer_index
.ok_or_else(|| ApexError::MissingLayerIndex {
name: tensor.name.to_string(),
})?;
if (layer as u64) >= (self.n_layers as u64) {
return Err(ApexError::LayerIndexOutOfRange {
name: tensor.name.to_string(),
layer_index: layer,
n_layers: self.n_layers,
});
}
Ok(layer)
}
}
#[cfg(test)]
mod tests {
use super::super::super::tensor_ref::SourceDtype;
use super::*;
fn tref<'a>(name: &'a str, arch: ArchName, layer: Option<usize>) -> TensorRef<'a> {
static SHAPE: [usize; 2] = [4096, 4096];
TensorRef {
name,
shape: &SHAPE,
source_dtype: SourceDtype::BF16,
arch,
layer_index: layer,
}
}
#[test]
fn apex_policy_unsupported_arch_errors() {
let err = ApexPolicy::new(ApexTier::Quality, ArchName::Llama3, 32, 0).unwrap_err();
match err {
ApexError::UnsupportedArch { arch, supported } => {
assert_eq!(arch, "llama");
assert!(supported.contains(&"qwen3moe"));
assert!(supported.contains(&"gemma4"));
assert!(supported.contains(&"minimax-m2"));
}
other => panic!("expected UnsupportedArch, got {other:?}"),
}
let err = ApexPolicy::new(ApexTier::Quality, ArchName::Bert, 24, 0).unwrap_err();
assert!(matches!(err, ApexError::UnsupportedArch { .. }));
}
#[test]
fn apex_policy_dense_model_errors() {
let err = ApexPolicy::new(ApexTier::Quality, ArchName::Gemma4, 30, 0).unwrap_err();
assert!(matches!(err, ApexError::DenseModelNotSupported { .. }));
let err = ApexPolicy::new(ApexTier::Quality, ArchName::Gemma4, 30, 1).unwrap_err();
assert!(matches!(err, ApexError::DenseModelNotSupported { .. }));
}
#[test]
fn apex_policy_routed_expert_edge_layer_0() {
let p = ApexPolicy::new(ApexTier::Quality, ArchName::Qwen35Moe, 40, 128).unwrap();
let t = tref("blk.0.ffn_gate_exps.weight", ArchName::Qwen35Moe, Some(0));
assert_eq!(p.target_for(&t).unwrap(), GgmlType::Q6_K);
}
#[test]
fn apex_policy_routed_expert_mid_layer_20() {
let p = ApexPolicy::new(ApexTier::Quality, ArchName::Qwen35Moe, 40, 128).unwrap();
let t = tref("blk.20.ffn_gate_exps.weight", ArchName::Qwen35Moe, Some(20));
assert_eq!(p.target_for(&t).unwrap(), GgmlType::IQ4_XS);
}
#[test]
fn apex_policy_attention_edge_vs_mid() {
let p = ApexPolicy::new(ApexTier::Mini, ArchName::Gemma4, 30, 8).unwrap();
let t_edge = tref("blk.2.attn_q.weight", ArchName::Gemma4, Some(2));
assert_eq!(p.target_for(&t_edge).unwrap(), GgmlType::Q4_K);
let t_mid = tref("blk.3.attn_q.weight", ArchName::Gemma4, Some(3));
assert_eq!(p.target_for(&t_mid).unwrap(), GgmlType::Q3_K);
}
#[test]
fn apex_policy_router_gate_q5_0() {
for tier in [
ApexTier::Quality,
ApexTier::Balanced,
ApexTier::Compact,
ApexTier::Mini,
] {
let p = ApexPolicy::new(tier, ArchName::Qwen35Moe, 40, 128).unwrap();
let t = tref("blk.5.ffn_gate_inp.weight", ArchName::Qwen35Moe, Some(5));
assert_eq!(p.target_for(&t).unwrap(), GgmlType::Q5_0, "tier {tier:?}");
}
}
#[test]
fn apex_policy_norm_f32() {
for tier in [
ApexTier::Quality,
ApexTier::Balanced,
ApexTier::Compact,
ApexTier::Mini,
] {
let p = ApexPolicy::new(tier, ArchName::Qwen35Moe, 40, 128).unwrap();
for name in [
"blk.5.attn_norm.weight",
"blk.5.ffn_norm.weight",
"blk.5.attn_q_norm.weight",
"output_norm.weight",
] {
let layer = if name.starts_with("blk.") {
Some(5)
} else {
None
};
let t = tref(name, ArchName::Qwen35Moe, layer);
assert_eq!(
p.target_for(&t).unwrap(),
GgmlType::F32,
"tier {tier:?} name {name}"
);
}
}
}
#[test]
fn apex_policy_token_embd_q6_k() {
for tier in [
ApexTier::Quality,
ApexTier::Balanced,
ApexTier::Compact,
ApexTier::Mini,
] {
let p = ApexPolicy::new(tier, ArchName::Qwen35Moe, 40, 128).unwrap();
let t = tref("token_embd.weight", ArchName::Qwen35Moe, None);
assert_eq!(p.target_for(&t).unwrap(), GgmlType::Q6_K, "tier {tier:?}");
}
}
#[test]
fn apex_policy_missing_layer_index_errors() {
let p = ApexPolicy::new(ApexTier::Quality, ArchName::Qwen35Moe, 40, 128).unwrap();
let t = tref("blk.0.ffn_gate_exps.weight", ArchName::Qwen35Moe, None);
let err = p.target_for(&t).unwrap_err();
assert!(matches!(err, ApexError::MissingLayerIndex { .. }));
}
#[test]
fn apex_policy_layer_index_out_of_range_errors() {
let p = ApexPolicy::new(ApexTier::Quality, ArchName::Qwen35Moe, 40, 128).unwrap();
let t = tref("blk.40.ffn_gate_exps.weight", ArchName::Qwen35Moe, Some(40));
let err = p.target_for(&t).unwrap_err();
assert!(matches!(err, ApexError::LayerIndexOutOfRange { .. }));
}
#[test]
fn apex_policy_matches_carnice_qwen36_layer_5() {
let p = ApexPolicy::new(ApexTier::Quality, ArchName::Qwen35Moe, 41, 128).unwrap();
let exp = tref("blk.5.ffn_gate_exps.weight", ArchName::Qwen35Moe, Some(5));
assert_eq!(p.target_for(&exp).unwrap(), GgmlType::Q5_K);
let shexp = tref("blk.5.ffn_gate_shexp.weight", ArchName::Qwen35Moe, Some(5));
assert_eq!(p.target_for(&shexp).unwrap(), GgmlType::Q8_0);
let attn = tref("blk.5.attn_q.weight", ArchName::Qwen35Moe, Some(5));
assert_eq!(p.target_for(&attn).unwrap(), GgmlType::Q6_K);
}
#[test]
fn apex_policy_matches_gemma4_mini_layer_10() {
let p = ApexPolicy::new(ApexTier::Mini, ArchName::Gemma4, 30, 8).unwrap();
let exp = tref("blk.10.ffn_gate_exps.weight", ArchName::Gemma4, Some(10));
assert_eq!(p.target_for(&exp).unwrap(), GgmlType::IQ2_S);
let shexp = tref("blk.10.ffn_gate_shexp.weight", ArchName::Gemma4, Some(10));
assert_eq!(p.target_for(&shexp).unwrap(), GgmlType::Q4_K);
let attn = tref("blk.10.attn_q.weight", ArchName::Gemma4, Some(10));
assert_eq!(p.target_for(&attn).unwrap(), GgmlType::Q3_K);
}
#[test]
fn apex_policy_with_mudler_override_wins_for_enumerated_tensors() {
use super::super::fingerprint::vendor_config_content;
use super::super::mudler_config::MudlerConfig;
let content =
vendor_config_content("vendor/apex-quant/configs/gemma4_26b_balanced.txt").unwrap();
let mudler: &'static MudlerConfig = Box::leak(Box::new(
MudlerConfig::parse(content, "test/gemma4_26b_balanced.txt").unwrap(),
));
let p = ApexPolicy::new(ApexTier::Balanced, ArchName::Gemma4, 30, 128)
.unwrap()
.with_mudler_override(mudler);
let exp_0 = tref("blk.0.ffn_gate_exps.weight", ArchName::Gemma4, Some(0));
assert_eq!(p.target_for(&exp_0).unwrap(), GgmlType::Q6_K);
let exp_5 = tref("blk.5.ffn_gate_exps.weight", ArchName::Gemma4, Some(5));
assert_eq!(p.target_for(&exp_5).unwrap(), GgmlType::Q5_K);
let rg = tref("blk.5.ffn_gate_inp.weight", ArchName::Gemma4, Some(5));
assert_eq!(p.target_for(&rg).unwrap(), GgmlType::Q5_0);
let te = tref("token_embd.weight", ArchName::Gemma4, None);
assert_eq!(p.target_for(&te).unwrap(), GgmlType::Q6_K);
let nm = tref("blk.5.attn_norm.weight", ArchName::Gemma4, Some(5));
assert_eq!(p.target_for(&nm).unwrap(), GgmlType::F32);
}
#[test]
fn p4b_mudler_override_is_tier_independent_for_i_and_non_i_siblings() {
use super::super::fingerprint::vendor_config_content;
use super::super::mudler_config::MudlerConfig;
let content = vendor_config_content("vendor/apex-quant/configs/gemma4_26b_balanced.txt")
.expect("vendored balanced config must be baked in");
let mudler: &'static MudlerConfig = Box::leak(Box::new(
MudlerConfig::parse(content, "test/gemma4_26b_balanced.txt:p4b")
.expect("vendored config must parse"),
));
let non_i = ApexPolicy::new(ApexTier::Balanced, ArchName::Gemma4, 30, 128)
.expect("non-I policy must construct")
.with_mudler_override(mudler);
let i = ApexPolicy::new_with_imatrix(ApexTier::IBalanced, ArchName::Gemma4, 30, 128)
.expect("I-tier policy must construct on Gemma4")
.with_mudler_override(mudler);
const EXPECTED_ENUMERATED_COUNT: usize = 450;
assert_eq!(
mudler.map.len(),
EXPECTED_ENUMERATED_COUNT,
"§P4b override-equivalence fixture drift: gemma4_26b_balanced.txt \
now enumerates {} tensors (was {EXPECTED_ENUMERATED_COUNT}). If \
the upstream vendor sync intentionally changed the surface, update \
this constant + ADR-033 §P4b at the same commit.",
mudler.map.len(),
);
for (name, _expected) in mudler.map.iter() {
let canonical_name = format!("{name}.weight");
let layer_index = canonical_name
.strip_prefix("blk.")
.and_then(|r| r.find('.').map(|i| (r, i)))
.and_then(|(r, dot)| r[..dot].parse::<usize>().ok());
let shape = [4096usize, 1];
let tref = TensorRef {
name: &canonical_name,
shape: &shape,
source_dtype: SourceDtype::BF16,
arch: ArchName::Gemma4,
layer_index,
};
let a = non_i.target_for(&tref).expect("non-I override path");
let b = i.target_for(&tref).expect("I override path");
assert_eq!(
a, b,
"§P4b override-equivalence: tensor `{canonical_name}` \
produced non-I={a:?}, I={b:?} — the override path must \
be tier-independent (see policy.rs:277-310). Update both \
the override branch and this test at the same commit.",
);
}
for (name, layer) in [
("token_embd.weight", None),
("output.weight", None),
("output_norm.weight", None),
("blk.5.attn_norm.weight", Some(5)),
("blk.5.ffn_norm.weight", Some(5)),
("blk.5.ffn_gate_inp.weight", Some(5)),
] {
let shape = [4096usize, 1];
let tref = TensorRef {
name,
shape: &shape,
source_dtype: SourceDtype::BF16,
arch: ArchName::Gemma4,
layer_index: layer,
};
let a = non_i.target_for(&tref).expect("non-I structural arm");
let b = i.target_for(&tref).expect("I structural arm");
assert_eq!(
a, b,
"§P4b override-equivalence: structural tensor `{name}` \
produced non-I={a:?}, I={b:?} — fall-through arms in \
policy.rs must stay tier-independent.",
);
}
for name in [
"blk.99.ffn_gate_exps.weight", "blk.99.ffn_gate_shexp.weight", "blk.99.attn_q.weight", ] {
let shape = [4096usize, 1];
let tref = TensorRef {
name,
shape: &shape,
source_dtype: SourceDtype::BF16,
arch: ArchName::Gemma4,
layer_index: Some(99),
};
let a_res = non_i.target_for(&tref);
let b_res = i.target_for(&tref);
match (&a_res, &b_res) {
(
Err(ApexError::TensorNotInMudlerConfig {
source_path: a_path,
tensor_name: a_name,
}),
Err(ApexError::TensorNotInMudlerConfig {
source_path: b_path,
tensor_name: b_name,
}),
) => {
assert_eq!(
a_path, b_path,
"§P4b override-miss source_path drift: non-I={a_path}, I={b_path}",
);
assert_eq!(
a_name, b_name,
"§P4b override-miss tensor_name drift: non-I={a_name}, I={b_name}",
);
}
_ => panic!(
"§P4b override-miss expected TensorNotInMudlerConfig from both \
siblings for `{name}` — non-I={a_res:?}, I={b_res:?}. If the \
enumerated arms changed semantics, update both this assertion \
and ADR-033 §P4b at the same commit."
),
}
}
}
#[test]
fn apex_policy_new_rejects_i_tier() {
for tier in [ApexTier::IQuality, ApexTier::IBalanced, ApexTier::ICompact] {
let err = ApexPolicy::new(tier, ArchName::Gemma4, 30, 128).unwrap_err();
match err {
ApexError::ImatrixRequiresInference {
tier: t,
arch,
supported_for_imatrix,
} => {
assert_eq!(t, tier.cli_name());
assert_eq!(arch, "gemma4");
assert_eq!(supported_for_imatrix, &["qwen3moe", "qwen35moe", "gemma4"]);
}
other => panic!("expected ImatrixRequiresInference for {tier:?}, got {other:?}"),
}
}
}
#[test]
fn apex_policy_new_with_imatrix_accepts_i_tier_for_supported_arches() {
for tier in [ApexTier::IQuality, ApexTier::IBalanced, ApexTier::ICompact] {
let p = ApexPolicy::new_with_imatrix(tier, ArchName::Gemma4, 30, 128).unwrap();
assert_eq!(p.tier, tier);
let p2 = ApexPolicy::new_with_imatrix(tier, ArchName::Qwen35Moe, 40, 128).unwrap();
assert_eq!(p2.tier, tier);
}
}
#[test]
fn apex_policy_new_with_imatrix_rejects_i_tier_for_unsupported_arch() {
let err = ApexPolicy::new_with_imatrix(ApexTier::IBalanced, ArchName::MiniMaxM2, 32, 128)
.unwrap_err();
assert!(matches!(err, ApexError::ImatrixRequiresInference { .. }));
}
#[test]
fn apex_policy_new_with_imatrix_accepts_non_i_tiers() {
for tier in [
ApexTier::Quality,
ApexTier::Balanced,
ApexTier::Compact,
ApexTier::Mini,
] {
ApexPolicy::new_with_imatrix(tier, ArchName::Gemma4, 30, 128)
.unwrap_or_else(|e| panic!("non-I tier {tier:?} rejected: {e}"));
}
}
#[test]
fn apex_policy_accepts_non_i_tiers() {
for tier in [
ApexTier::Quality,
ApexTier::Balanced,
ApexTier::Compact,
ApexTier::Mini,
] {
ApexPolicy::new(tier, ArchName::Gemma4, 30, 128)
.unwrap_or_else(|e| panic!("non-I tier {tier:?} rejected: {e}"));
}
}
}