use super::error::QuantizeError;
use super::ggml_type::GgmlType;
use super::llama_ftype::LlamaFtype;
use super::tensor_ref::{ArchName, TensorRef};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum TensorCategory {
Output,
TokenEmbd,
AttnV,
AttnQ,
AttnK,
AttnQkv,
AttnKvB,
AttnOutput,
FfnUp,
FfnGate,
FfnDown,
Other,
}
impl TensorCategory {
pub fn classify(name: &str) -> Self {
if name == "output.weight" {
return TensorCategory::Output;
}
if name == "token_embd.weight" || name == "per_layer_token_embd.weight" {
return TensorCategory::TokenEmbd;
}
if name.contains("attn_qkv.weight") {
return TensorCategory::AttnQkv;
}
if name.contains("attn_kv_b.weight") {
return TensorCategory::AttnKvB;
}
if name.contains("attn_v.weight") {
return TensorCategory::AttnV;
}
if name.contains("attn_k.weight") {
return TensorCategory::AttnK;
}
if name.contains("attn_q.weight") {
return TensorCategory::AttnQ;
}
if name.contains("attn_output.weight") {
return TensorCategory::AttnOutput;
}
if name.contains("ffn_up") {
return TensorCategory::FfnUp;
}
if name.contains("ffn_gate") {
return TensorCategory::FfnGate;
}
if name.contains("ffn_down") {
return TensorCategory::FfnDown;
}
TensorCategory::Other
}
pub const fn is_attn_v(self) -> bool {
matches!(
self,
TensorCategory::AttnV | TensorCategory::AttnQkv | TensorCategory::AttnKvB
)
}
}
pub fn tensor_type_fallback(target: GgmlType, n_per_row: usize) -> Result<GgmlType, QuantizeError> {
let target_block = target.block_size();
if n_per_row % target_block == 0 {
return Ok(target);
}
let downshift = match target {
GgmlType::IQ1_S
| GgmlType::IQ1_M
| GgmlType::IQ2_XXS
| GgmlType::IQ2_XS
| GgmlType::IQ2_S
| GgmlType::IQ3_XXS
| GgmlType::IQ3_S
| GgmlType::IQ4_XS => GgmlType::IQ4_NL,
GgmlType::Q2_K | GgmlType::Q3_K | GgmlType::TQ1_0 | GgmlType::TQ2_0 => GgmlType::Q4_0,
GgmlType::Q4_K => GgmlType::Q5_0,
GgmlType::Q5_K => GgmlType::Q5_1,
GgmlType::Q6_K => GgmlType::Q8_0,
_ => {
return Err(QuantizeError::NotBlockAligned {
ggml_type: target,
n_per_row,
block_size: target_block,
});
}
};
if n_per_row % downshift.block_size() != 0 {
return Err(QuantizeError::NotBlockAligned {
ggml_type: downshift,
n_per_row,
block_size: downshift.block_size(),
});
}
Ok(downshift)
}
#[derive(Debug, Clone, Copy)]
pub struct HParams {
pub n_expert: u32,
pub n_head: u32,
pub n_head_kv: u32,
pub n_layer: u32,
pub n_mtp_layers: u32,
}
impl HParams {
pub const fn n_gqa(&self) -> u32 {
if self.n_head_kv == 0 {
return 0;
}
self.n_head / self.n_head_kv
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum LlmType {
M70B,
Other,
}
#[derive(Debug, Clone, Copy, Default)]
pub struct QuantizeParams {
pub output_tensor_type: Option<GgmlType>,
pub token_embedding_type: Option<GgmlType>,
}
#[derive(Debug, Clone)]
pub struct QsState {
pub ftype: LlamaFtype,
pub arch: ArchName,
pub model_type: LlmType,
pub hparams: HParams,
pub params: QuantizeParams,
pub has_tied_embeddings: bool,
pub has_imatrix: bool,
pub n_attention_wv: i32,
pub i_attention_wv: i32,
pub n_ffn_down: i32,
pub i_ffn_down: i32,
pub n_ffn_gate: i32,
pub i_ffn_gate: i32,
pub n_ffn_up: i32,
pub i_ffn_up: i32,
}
impl QsState {
pub const fn new(
ftype: LlamaFtype,
arch: ArchName,
model_type: LlmType,
hparams: HParams,
) -> Self {
Self {
ftype,
arch,
model_type,
hparams,
params: QuantizeParams {
output_tensor_type: None,
token_embedding_type: None,
},
has_tied_embeddings: true,
has_imatrix: false,
n_attention_wv: 0,
i_attention_wv: 0,
n_ffn_down: 0,
i_ffn_down: 0,
n_ffn_gate: 0,
i_ffn_gate: 0,
n_ffn_up: 0,
i_ffn_up: 0,
}
}
}
#[inline]
const fn use_more_bits(i_layer: i32, n_layers: i32) -> bool {
i_layer < n_layers / 8 || i_layer >= 7 * n_layers / 8 || (i_layer - n_layers / 8) % 3 == 2
}
fn layer_info(
i_layer: i32,
n_layer: i32,
name: &str,
n_expert: i32,
) -> Result<(i32, i32), QuantizeError> {
if n_expert > 1 {
let parsed = name.strip_prefix("blk.").and_then(|rest| {
let dot = rest.find('.')?;
rest[..dot].parse::<i32>().ok()
});
let parsed = match parsed {
Some(v) => v,
None => {
return Err(QuantizeError::BadLayerForTensor {
name: name.to_string(),
n_layer,
});
}
};
if parsed < 0 || parsed >= n_layer {
return Err(QuantizeError::BadLayerForTensor {
name: name.to_string(),
n_layer,
});
}
Ok((parsed, n_layer))
} else {
Ok((i_layer, n_layer))
}
}
pub struct StandardPolicy;
impl StandardPolicy {
pub const fn new() -> Self {
Self
}
pub fn target_for(
&self,
qs: &mut QsState,
tensor: &TensorRef,
category: TensorCategory,
) -> Result<GgmlType, QuantizeError> {
let name = tensor.name;
let arch = qs.arch;
let ftype = qs.ftype;
let mut new_type = ftype.primary_type();
let n_expert = (qs.hparams.n_expert as i32).max(1);
if category == TensorCategory::Output
|| (qs.has_tied_embeddings && category == TensorCategory::TokenEmbd)
{
if let Some(t) = qs.params.output_tensor_type {
new_type = t;
} else {
let nx = tensor.n_per_row() as i64;
let qk_k = new_type.block_size() as i64;
if ftype == LlamaFtype::MostlyMXFP4_MOE {
new_type = GgmlType::Q8_0;
} else if arch == ArchName::Falcon || nx % qk_k != 0 {
new_type = GgmlType::Q8_0;
} else if matches!(
ftype,
LlamaFtype::MostlyIQ2_XXS
| LlamaFtype::MostlyIQ2_XS
| LlamaFtype::MostlyIQ3_XXS
| LlamaFtype::MostlyIQ1_S
| LlamaFtype::MostlyIQ2_S
| LlamaFtype::MostlyIQ2_M
| LlamaFtype::MostlyIQ1_M
) {
new_type = GgmlType::Q5_K;
} else if new_type != GgmlType::Q8_0 {
new_type = GgmlType::Q6_K;
}
}
} else if ftype == LlamaFtype::MostlyMXFP4_MOE {
let is_3d = tensor.shape.len() > 2 && tensor.shape[2] > 1;
new_type = if is_3d {
GgmlType::MXFP4
} else {
GgmlType::Q8_0
};
} else if category == TensorCategory::TokenEmbd {
if let Some(t) = qs.params.token_embedding_type {
new_type = t;
} else if matches!(
ftype,
LlamaFtype::MostlyIQ2_XXS
| LlamaFtype::MostlyIQ2_XS
| LlamaFtype::MostlyIQ1_S
| LlamaFtype::MostlyIQ1_M
) {
new_type = GgmlType::Q2_K;
} else if matches!(ftype, LlamaFtype::MostlyIQ2_S | LlamaFtype::MostlyIQ2_M) {
new_type = GgmlType::IQ3_S;
} else if ftype == LlamaFtype::MostlyIQ3_XXS {
new_type = GgmlType::IQ3_S;
} else if matches!(ftype, LlamaFtype::MostlyTQ1_0 | LlamaFtype::MostlyTQ2_0) {
new_type = GgmlType::Q4_K;
}
} else if matches!(
ftype,
LlamaFtype::MostlyIQ2_XXS
| LlamaFtype::MostlyIQ2_XS
| LlamaFtype::MostlyIQ1_S
| LlamaFtype::MostlyIQ2_S
| LlamaFtype::MostlyIQ2_M
| LlamaFtype::MostlyIQ1_M
) {
if category.is_attn_v() {
if qs.hparams.n_gqa() >= 4 || qs.hparams.n_expert >= 4 {
new_type = GgmlType::Q4_K;
} else if matches!(ftype, LlamaFtype::MostlyIQ2_S | LlamaFtype::MostlyIQ2_M) {
new_type = GgmlType::IQ3_S;
} else {
new_type = GgmlType::Q2_K;
}
qs.i_attention_wv += 1;
} else if qs.hparams.n_expert == 8 && category == TensorCategory::AttnK {
new_type = GgmlType::Q4_K;
} else if category == TensorCategory::FfnDown {
if qs.i_ffn_down < qs.n_ffn_down / 8 {
new_type = if matches!(ftype, LlamaFtype::MostlyIQ2_S | LlamaFtype::MostlyIQ2_M)
{
GgmlType::IQ3_S
} else {
GgmlType::Q2_K
};
}
qs.i_ffn_down += 1;
} else if category == TensorCategory::AttnOutput {
if qs.hparams.n_expert == 8 {
new_type = GgmlType::Q5_K;
} else if matches!(ftype, LlamaFtype::MostlyIQ1_S | LlamaFtype::MostlyIQ1_M) {
new_type = GgmlType::IQ2_XXS;
} else if matches!(ftype, LlamaFtype::MostlyIQ2_S | LlamaFtype::MostlyIQ2_M) {
new_type = GgmlType::IQ3_S;
}
}
} else if category.is_attn_v() {
if ftype == LlamaFtype::MostlyQ2_K {
new_type = if qs.hparams.n_gqa() >= 4 {
GgmlType::Q4_K
} else {
GgmlType::Q3_K
};
} else if ftype == LlamaFtype::MostlyQ2_K_S && qs.hparams.n_gqa() >= 4 {
new_type = GgmlType::Q4_K;
} else if ftype == LlamaFtype::MostlyIQ3_XXS {
new_type = if qs.hparams.n_gqa() >= 4 {
GgmlType::Q4_K
} else if !qs.has_imatrix {
GgmlType::IQ3_S
} else {
GgmlType::IQ3_XXS
};
} else if matches!(ftype, LlamaFtype::MostlyIQ3_XS | LlamaFtype::MostlyIQ3_S)
&& qs.hparams.n_gqa() >= 4
{
new_type = GgmlType::Q4_K;
} else if ftype == LlamaFtype::MostlyIQ3_M {
new_type = GgmlType::Q4_K;
} else if ftype == LlamaFtype::MostlyQ3_K_M {
new_type = if qs.i_attention_wv < 2 {
GgmlType::Q5_K
} else {
GgmlType::Q4_K
};
} else if ftype == LlamaFtype::MostlyQ3_K_L {
new_type = GgmlType::Q5_K;
} else if matches!(ftype, LlamaFtype::MostlyIQ4_NL | LlamaFtype::MostlyIQ4_XS)
&& qs.hparams.n_gqa() >= 4
{
new_type = GgmlType::Q5_K;
} else if matches!(ftype, LlamaFtype::MostlyQ4_K_M | LlamaFtype::MostlyQ5_K_M)
&& use_more_bits(qs.i_attention_wv, qs.n_attention_wv)
{
new_type = GgmlType::Q6_K;
} else if ftype == LlamaFtype::MostlyQ4_K_S && qs.i_attention_wv < 4 {
new_type = GgmlType::Q5_K;
}
if qs.model_type == LlmType::M70B && matches!(new_type, GgmlType::Q3_K | GgmlType::Q4_K)
{
new_type = GgmlType::Q5_K;
}
if qs.hparams.n_expert == 8 {
new_type = GgmlType::Q8_0;
}
qs.i_attention_wv += 1;
} else if category == TensorCategory::AttnK {
if qs.hparams.n_expert == 8 {
new_type = GgmlType::Q8_0;
} else if ftype == LlamaFtype::MostlyIQ3_XS {
new_type = GgmlType::IQ3_XXS;
} else if ftype == LlamaFtype::MostlyIQ3_XXS {
new_type = GgmlType::IQ2_S;
}
} else if category == TensorCategory::AttnQ {
if ftype == LlamaFtype::MostlyIQ3_XS {
new_type = GgmlType::IQ3_XXS;
} else if ftype == LlamaFtype::MostlyIQ3_XXS {
new_type = GgmlType::IQ2_S;
}
} else if category == TensorCategory::FfnDown {
let (i_layer, n_layer) = layer_info(qs.i_ffn_down, qs.n_ffn_down, name, n_expert)?;
if ftype == LlamaFtype::MostlyQ2_K {
new_type = GgmlType::Q3_K;
} else if ftype == LlamaFtype::MostlyQ2_K_S {
if i_layer < n_layer / 8 {
new_type = GgmlType::Q4_K;
}
} else if ftype == LlamaFtype::MostlyIQ3_XXS && !qs.has_imatrix {
new_type = if i_layer < n_layer / 8 {
GgmlType::Q4_K
} else {
GgmlType::Q3_K
};
} else if ftype == LlamaFtype::MostlyQ3_K_M {
new_type = if i_layer < n_layer / 16 {
GgmlType::Q5_K
} else if arch != ArchName::Falcon || use_more_bits(i_layer, n_layer) {
GgmlType::Q4_K
} else {
GgmlType::Q3_K
};
} else if ftype == LlamaFtype::MostlyIQ3_M
&& (i_layer < n_layer / 8
|| (qs.hparams.n_expert == 8 && use_more_bits(i_layer, n_layer)))
{
new_type = GgmlType::Q4_K;
} else if ftype == LlamaFtype::MostlyQ3_K_L {
new_type = if arch == ArchName::Falcon {
GgmlType::Q4_K
} else {
GgmlType::Q5_K
};
} else if ftype == LlamaFtype::MostlyQ4_K_M {
if arch == ArchName::Falcon {
new_type = if i_layer < n_layer / 16 {
GgmlType::Q6_K
} else if use_more_bits(i_layer, n_layer) {
GgmlType::Q5_K
} else {
GgmlType::Q4_K
};
} else if use_more_bits(i_layer, n_layer) {
new_type = GgmlType::Q6_K;
}
} else if i_layer < n_layer / 8
&& matches!(ftype, LlamaFtype::MostlyIQ4_NL | LlamaFtype::MostlyIQ4_XS)
&& !qs.has_imatrix
{
new_type = GgmlType::Q5_K;
} else if ftype == LlamaFtype::MostlyQ5_K_M && use_more_bits(i_layer, n_layer) {
new_type = GgmlType::Q6_K;
} else if ftype == LlamaFtype::MostlyQ4_K_S
&& arch != ArchName::Falcon
&& i_layer < n_layer / 8
{
new_type = GgmlType::Q5_K;
} else if matches!(ftype, LlamaFtype::MostlyQ4_0 | LlamaFtype::MostlyQ5_0)
&& qs.has_imatrix
&& i_layer < n_layer / 8
{
new_type = if ftype == LlamaFtype::MostlyQ4_0 {
GgmlType::Q4_1
} else {
GgmlType::Q5_1
};
}
qs.i_ffn_down += 1;
} else if category == TensorCategory::AttnOutput {
if arch != ArchName::Falcon {
if qs.hparams.n_expert == 8 {
if matches!(
ftype,
LlamaFtype::MostlyQ2_K
| LlamaFtype::MostlyIQ3_XS
| LlamaFtype::MostlyIQ3_XXS
| LlamaFtype::MostlyQ3_K_S
| LlamaFtype::MostlyQ3_K_M
| LlamaFtype::MostlyIQ4_NL
| LlamaFtype::MostlyQ4_K_S
| LlamaFtype::MostlyQ4_K_M
| LlamaFtype::MostlyIQ3_S
| LlamaFtype::MostlyIQ3_M
| LlamaFtype::MostlyIQ4_XS
) {
new_type = GgmlType::Q5_K;
}
} else {
if ftype == LlamaFtype::MostlyQ2_K {
new_type = GgmlType::Q3_K;
} else if ftype == LlamaFtype::MostlyIQ3_XXS {
new_type = GgmlType::IQ3_S;
} else if ftype == LlamaFtype::MostlyQ3_K_M {
new_type = GgmlType::Q4_K;
} else if ftype == LlamaFtype::MostlyQ3_K_L {
new_type = GgmlType::Q5_K;
} else if ftype == LlamaFtype::MostlyIQ3_M {
new_type = GgmlType::Q4_K;
}
}
} else if ftype == LlamaFtype::MostlyQ3_K_L {
new_type = GgmlType::Q4_K;
}
} else if category == TensorCategory::AttnQkv {
if matches!(
ftype,
LlamaFtype::MostlyQ3_K_M | LlamaFtype::MostlyQ3_K_L | LlamaFtype::MostlyIQ3_M
) {
new_type = GgmlType::Q4_K;
} else if ftype == LlamaFtype::MostlyQ4_K_M {
new_type = GgmlType::Q5_K;
} else if ftype == LlamaFtype::MostlyQ5_K_M {
new_type = GgmlType::Q6_K;
}
} else if category == TensorCategory::FfnGate {
let (i_layer, n_layer) = layer_info(qs.i_ffn_gate, qs.n_ffn_gate, name, n_expert)?;
if ftype == LlamaFtype::MostlyIQ3_XS
&& i_layer >= n_layer / 8
&& i_layer < 7 * n_layer / 8
{
new_type = GgmlType::IQ3_XXS;
}
qs.i_ffn_gate += 1;
} else if category == TensorCategory::FfnUp {
let (i_layer, n_layer) = layer_info(qs.i_ffn_up, qs.n_ffn_up, name, n_expert)?;
if ftype == LlamaFtype::MostlyIQ3_XS
&& i_layer >= n_layer / 8
&& i_layer < 7 * n_layer / 8
{
new_type = GgmlType::IQ3_XXS;
}
qs.i_ffn_up += 1;
}
tensor_type_fallback(new_type, tensor.n_per_row())
}
}
impl Default for StandardPolicy {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::super::tensor_ref::{ArchName, SourceDtype};
use super::*;
fn mk_qs(ftype: LlamaFtype, arch: ArchName) -> QsState {
let hparams = HParams {
n_expert: 0,
n_head: 32,
n_head_kv: 8,
n_layer: 32,
n_mtp_layers: 0,
};
QsState::new(ftype, arch, LlmType::Other, hparams)
}
fn mk_tensor<'a>(name: &'a str, shape: &'a [usize], arch: ArchName) -> TensorRef<'a> {
TensorRef {
name,
shape,
source_dtype: SourceDtype::BF16,
arch,
layer_index: None,
}
}
#[test]
fn passthrough_f32() {
let mut qs = mk_qs(LlamaFtype::AllF32, ArchName::Llama3);
let shape = [4096, 1];
let t = mk_tensor("blk.0.attn_q.weight", &shape, ArchName::Llama3);
let pol = StandardPolicy::new();
let cat = TensorCategory::classify(t.name);
assert_eq!(pol.target_for(&mut qs, &t, cat).unwrap(), GgmlType::F32);
}
#[test]
fn passthrough_f16() {
let mut qs = mk_qs(LlamaFtype::MostlyF16, ArchName::Llama3);
let shape = [4096, 1];
let t = mk_tensor("blk.0.attn_q.weight", &shape, ArchName::Llama3);
let pol = StandardPolicy::new();
let cat = TensorCategory::classify(t.name);
assert_eq!(pol.target_for(&mut qs, &t, cat).unwrap(), GgmlType::F16);
}
#[test]
fn passthrough_bf16() {
let mut qs = mk_qs(LlamaFtype::BF16, ArchName::Llama3);
let shape = [4096, 1];
let t = mk_tensor("blk.0.attn_q.weight", &shape, ArchName::Llama3);
let pol = StandardPolicy::new();
let cat = TensorCategory::classify(t.name);
assert_eq!(pol.target_for(&mut qs, &t, cat).unwrap(), GgmlType::BF16);
}
#[test]
fn output_q5_k_m_bumps_to_q6_k() {
let mut qs = mk_qs(LlamaFtype::MostlyQ5_K_M, ArchName::Llama3);
let shape = [4096, 1];
let t = mk_tensor("output.weight", &shape, ArchName::Llama3);
let pol = StandardPolicy::new();
assert_eq!(
pol.target_for(&mut qs, &t, TensorCategory::Output).unwrap(),
GgmlType::Q6_K
);
}
#[test]
fn token_embd_q4_k_s_stays_q4_k() {
let mut qs = mk_qs(LlamaFtype::MostlyQ4_K_S, ArchName::Llama3);
qs.has_tied_embeddings = false;
let shape = [4096, 1];
let t = mk_tensor("token_embd.weight", &shape, ArchName::Llama3);
let pol = StandardPolicy::new();
assert_eq!(
pol.target_for(&mut qs, &t, TensorCategory::TokenEmbd)
.unwrap(),
GgmlType::Q4_K
);
}
#[test]
fn token_embd_tied_routes_to_output_branch() {
let mut qs = mk_qs(LlamaFtype::MostlyQ5_K_M, ArchName::Llama3);
let shape = [4096, 1];
let t = mk_tensor("token_embd.weight", &shape, ArchName::Llama3);
let pol = StandardPolicy::new();
assert_eq!(
pol.target_for(&mut qs, &t, TensorCategory::TokenEmbd)
.unwrap(),
GgmlType::Q6_K
);
}
#[test]
fn attn_v_q5_k_m_first_layers_bump_q6_k() {
let mut qs = mk_qs(LlamaFtype::MostlyQ5_K_M, ArchName::Llama3);
qs.n_attention_wv = 32;
qs.i_attention_wv = 0;
let shape = [1024, 1];
let t = mk_tensor("blk.0.attn_v.weight", &shape, ArchName::Llama3);
let pol = StandardPolicy::new();
assert_eq!(
pol.target_for(&mut qs, &t, TensorCategory::AttnV).unwrap(),
GgmlType::Q6_K
);
assert_eq!(
qs.i_attention_wv, 1,
"counter must be incremented per C:548"
);
}
#[test]
fn ffn_down_q4_k_m_first_layers_bump_q6_k() {
let mut qs = mk_qs(LlamaFtype::MostlyQ4_K_M, ArchName::Llama3);
qs.n_ffn_down = 32;
qs.i_ffn_down = 0; let shape = [4096, 1];
let t = mk_tensor("blk.0.ffn_down.weight", &shape, ArchName::Llama3);
let pol = StandardPolicy::new();
assert_eq!(
pol.target_for(&mut qs, &t, TensorCategory::FfnDown)
.unwrap(),
GgmlType::Q6_K
);
assert_eq!(qs.i_ffn_down, 1);
}
#[test]
fn all_v1_arches_q5_k_m_attn_q_basic() {
let arches = [
ArchName::Gemma4,
ArchName::Gemma4Mmproj,
ArchName::Qwen35Moe,
ArchName::Qwen3VlText,
ArchName::Bert,
ArchName::NomicBert,
ArchName::Llama3,
ArchName::MiniMaxM2,
];
for arch in arches {
let mut qs = mk_qs(LlamaFtype::MostlyQ5_K_M, arch);
let shape = [4096, 1];
let t = mk_tensor("blk.0.attn_q.weight", &shape, arch);
let pol = StandardPolicy::new();
let got = pol.target_for(&mut qs, &t, TensorCategory::AttnQ).unwrap();
assert_eq!(got, GgmlType::Q5_K, "arch {} should pick Q5_K", arch.name());
}
}
#[test]
fn tensor_type_fallback_chains_on_misaligned_k_quant() {
let mut qs = mk_qs(LlamaFtype::MostlyQ5_K_M, ArchName::Llama3);
let shape = [17, 1];
let t = mk_tensor("blk.0.attn_q.weight", &shape, ArchName::Llama3);
let pol = StandardPolicy::new();
let err = pol
.target_for(&mut qs, &t, TensorCategory::AttnQ)
.unwrap_err();
assert!(matches!(err, QuantizeError::NotBlockAligned { .. }));
}
#[test]
fn tensor_type_fallback_first_shift_q5_k_to_q5_1() {
let mut qs = mk_qs(LlamaFtype::MostlyQ5_K_M, ArchName::Llama3);
let shape = [128, 1];
let t = mk_tensor("blk.0.attn_q.weight", &shape, ArchName::Llama3);
let pol = StandardPolicy::new();
assert_eq!(
pol.target_for(&mut qs, &t, TensorCategory::AttnQ).unwrap(),
GgmlType::Q5_1
);
}
#[test]
fn category_classify_basic() {
assert_eq!(
TensorCategory::classify("token_embd.weight"),
TensorCategory::TokenEmbd
);
assert_eq!(
TensorCategory::classify("output.weight"),
TensorCategory::Output
);
assert_eq!(
TensorCategory::classify("blk.0.attn_q.weight"),
TensorCategory::AttnQ
);
assert_eq!(
TensorCategory::classify("blk.10.attn_v.weight"),
TensorCategory::AttnV
);
assert_eq!(
TensorCategory::classify("blk.5.attn_k.weight"),
TensorCategory::AttnK
);
assert_eq!(
TensorCategory::classify("blk.3.attn_output.weight"),
TensorCategory::AttnOutput
);
assert_eq!(
TensorCategory::classify("blk.7.ffn_down.weight"),
TensorCategory::FfnDown
);
assert_eq!(
TensorCategory::classify("blk.7.ffn_down_exps.weight"),
TensorCategory::FfnDown
);
assert_eq!(
TensorCategory::classify("blk.7.ffn_up.weight"),
TensorCategory::FfnUp
);
assert_eq!(
TensorCategory::classify("blk.7.ffn_up_exps.weight"),
TensorCategory::FfnUp
);
assert_eq!(
TensorCategory::classify("blk.7.ffn_gate.weight"),
TensorCategory::FfnGate
);
assert_eq!(
TensorCategory::classify("blk.7.ffn_gate_exps.weight"),
TensorCategory::FfnGate
);
assert_eq!(
TensorCategory::classify("blk.7.ffn_gate_inp.weight"),
TensorCategory::FfnGate
);
assert_eq!(
TensorCategory::classify("blk.0.attn_qkv.weight"),
TensorCategory::AttnQkv
);
assert_eq!(
TensorCategory::classify("blk.0.attn_kv_b.weight"),
TensorCategory::AttnKvB
);
assert_eq!(
TensorCategory::classify("blk.0.attn_norm.weight"),
TensorCategory::Other
);
assert_eq!(
TensorCategory::classify("per_layer_token_embd.weight"),
TensorCategory::TokenEmbd
);
}
#[test]
fn fallback_passthrough_when_aligned() {
assert_eq!(
tensor_type_fallback(GgmlType::Q5_K, 512).unwrap(),
GgmlType::Q5_K
);
assert_eq!(
tensor_type_fallback(GgmlType::Q4_K, 256).unwrap(),
GgmlType::Q4_K
);
assert_eq!(
tensor_type_fallback(GgmlType::Q4_0, 32).unwrap(),
GgmlType::Q4_0
);
}
#[test]
fn fallback_q4_k_to_q5_0() {
assert_eq!(
tensor_type_fallback(GgmlType::Q4_K, 128).unwrap(),
GgmlType::Q5_0
);
}
#[test]
fn fallback_q5_k_to_q5_1() {
assert_eq!(
tensor_type_fallback(GgmlType::Q5_K, 160).unwrap(),
GgmlType::Q5_1
);
}
#[test]
fn fallback_q6_k_to_q8_0() {
assert_eq!(
tensor_type_fallback(GgmlType::Q6_K, 96).unwrap(),
GgmlType::Q8_0
);
}
#[test]
fn fallback_q2_q3_to_q4_0() {
assert_eq!(
tensor_type_fallback(GgmlType::Q2_K, 128).unwrap(),
GgmlType::Q4_0
);
assert_eq!(
tensor_type_fallback(GgmlType::Q3_K, 64).unwrap(),
GgmlType::Q4_0
);
}
#[test]
fn fallback_second_misalignment_is_typed_error() {
let err = tensor_type_fallback(GgmlType::Q5_K, 15).unwrap_err();
assert!(matches!(err, QuantizeError::NotBlockAligned { .. }));
}
#[test]
fn fallback_no_path_for_unaligned_legacy() {
let err = tensor_type_fallback(GgmlType::Q4_0, 17).unwrap_err();
assert!(matches!(err, QuantizeError::NotBlockAligned { .. }));
}
}