use std::collections::BTreeMap;
use frink_gguf::{GgmlType, GgufFile};
use super::policy::Target;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Role {
Output,
TokenEmbd,
AttnV,
AttnK,
AttnQ,
FfnDown,
AttnOutput,
AttnQkv,
FfnGate,
FfnUp,
Other,
}
impl Role {
pub fn of(name: &str, has_output: bool) -> Role {
if name == "output.weight" {
return Role::Output;
}
let is_token_embd = name == "token_embd.weight" || name == "per_layer_token_embd.weight";
if is_token_embd {
return if has_output {
Role::TokenEmbd
} else {
Role::Output
};
}
for (needle, role) in [
("attn_v.weight", Role::AttnV),
("attn_k.weight", Role::AttnK),
("attn_q.weight", Role::AttnQ),
("ffn_down", Role::FfnDown),
("attn_output.weight", Role::AttnOutput),
("attn_qkv.weight", Role::AttnQkv),
("ffn_gate", Role::FfnGate),
("ffn_up", Role::FfnUp),
] {
if name.contains(needle) {
return role;
}
}
Role::Other
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum When {
Always,
UseMoreBits,
FirstEighth,
FirstFour,
}
impl When {
fn holds(self, i: usize, n: usize) -> bool {
match self {
When::Always => true,
When::UseMoreBits => {
if i < n / 8 {
return true;
}
if n != 0 && i >= 7 * n / 8 {
return true;
}
(i - n / 8) % 3 == 2
}
When::FirstEighth => i < n / 8,
When::FirstFour => i < 4,
}
}
}
struct Promotion {
role: Role,
ftype: Target,
not_falcon: bool,
when: When,
to: GgmlType,
}
const PROMOTIONS: &[Promotion] = &[
Promotion {
role: Role::Output,
ftype: Target::Q4_K_S,
not_falcon: true,
when: When::Always,
to: GgmlType::Q6K,
},
Promotion {
role: Role::Output,
ftype: Target::Q4_K_M,
not_falcon: true,
when: When::Always,
to: GgmlType::Q6K,
},
Promotion {
role: Role::Output,
ftype: Target::Q5_K_S,
not_falcon: true,
when: When::Always,
to: GgmlType::Q6K,
},
Promotion {
role: Role::Output,
ftype: Target::Q5_K_M,
not_falcon: true,
when: When::Always,
to: GgmlType::Q6K,
},
Promotion {
role: Role::AttnV,
ftype: Target::Q4_K_M,
not_falcon: false,
when: When::UseMoreBits,
to: GgmlType::Q6K,
},
Promotion {
role: Role::AttnV,
ftype: Target::Q5_K_M,
not_falcon: false,
when: When::UseMoreBits,
to: GgmlType::Q6K,
},
Promotion {
role: Role::AttnV,
ftype: Target::Q4_K_S,
not_falcon: false,
when: When::FirstFour,
to: GgmlType::Q5K,
},
Promotion {
role: Role::FfnDown,
ftype: Target::Q4_K_M,
not_falcon: true,
when: When::UseMoreBits,
to: GgmlType::Q6K,
},
Promotion {
role: Role::FfnDown,
ftype: Target::Q5_K_M,
not_falcon: false,
when: When::UseMoreBits,
to: GgmlType::Q6K,
},
Promotion {
role: Role::FfnDown,
ftype: Target::Q4_K_S,
not_falcon: true,
when: When::FirstEighth,
to: GgmlType::Q5K,
},
Promotion {
role: Role::AttnQkv,
ftype: Target::Q4_K_M,
not_falcon: false,
when: When::Always,
to: GgmlType::Q5K,
},
Promotion {
role: Role::AttnQkv,
ftype: Target::Q5_K_M,
not_falcon: false,
when: When::Always,
to: GgmlType::Q6K,
},
];
const FALCON_Q4KM_FFN_DOWN_SIXTEENTH: GgmlType = GgmlType::Q6K;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ModelShape {
pub n_layer: usize,
pub n_expert: usize,
pub is_falcon: bool,
pub is_70b: bool,
}
impl ModelShape {
pub fn from_header(file: &GgufFile) -> ModelShape {
let arch = file.metadata_str("general.architecture").unwrap_or("");
let key =
|suffix: &str| file.metadata_u64(&format!("{arch}.{suffix}")).unwrap_or(0) as usize;
let n_layer = key("block_count");
let n_head = key("attention.head_count");
let n_head_kv = key("attention.head_count_kv");
ModelShape {
n_layer,
n_expert: key("expert_count"),
is_falcon: arch == "falcon",
is_70b: match arch {
"llama" | "llama-embed" => n_layer == 80 && n_head != n_head_kv,
"deci" | "qwen2" | "olmo" => n_layer == 80,
_ => false,
},
}
}
}
pub struct Recipe {
target: Target,
shape: ModelShape,
n_attention_wv: usize,
has_output: bool,
i_attention_wv: usize,
i_ffn_down: usize,
i_ffn_gate: usize,
i_ffn_up: usize,
}
impl Recipe {
pub fn new(target: Target, shape: ModelShape, tensor_names: &[String]) -> Recipe {
let n_attention_wv = tensor_names
.iter()
.filter(|n| {
n.contains("attn_v.weight")
|| n.contains("attn_qkv.weight")
|| n.contains("attn_kv_b.weight")
})
.count();
Recipe {
target,
shape,
n_attention_wv,
has_output: tensor_names.iter().any(|n| n == "output.weight"),
i_attention_wv: 0,
i_ffn_down: 0,
i_ffn_gate: 0,
i_ffn_up: 0,
}
}
pub fn tensor_type(&mut self, name: &str, shape: &[u64]) -> GgmlType {
let role = Role::of(name, self.has_output);
let default = self.target.ggml_type();
let ty = match role {
Role::Output => self.output_head_type(shape, default),
Role::AttnV => {
let t = self.chain(role, self.i_attention_wv, self.n_attention_wv, default);
let t = if self.shape.is_70b && matches!(t, GgmlType::Q3K | GgmlType::Q4K) {
GgmlType::Q5K
} else {
t
};
let t = if self.shape.n_expert == 8 {
GgmlType::Q8_0
} else {
t
};
self.i_attention_wv += 1;
t
}
Role::AttnK => {
if self.shape.n_expert == 8 {
GgmlType::Q8_0
} else {
default
}
}
Role::FfnDown => {
let (i, n) = self.ffn_layer(name, self.i_ffn_down);
let t = if self.shape.is_falcon && self.target == Target::Q4_K_M {
if i < n / 16 {
FALCON_Q4KM_FFN_DOWN_SIXTEENTH
} else if When::UseMoreBits.holds(i, n) {
GgmlType::Q5K
} else {
default
}
} else {
self.chain(role, i, n, default)
};
self.i_ffn_down += 1;
t
}
Role::AttnOutput => {
if !self.shape.is_falcon
&& self.shape.n_expert == 8
&& matches!(self.target, Target::Q4_K_S | Target::Q4_K_M)
{
GgmlType::Q5K
} else {
default
}
}
Role::FfnGate => {
let (i, n) = self.ffn_layer(name, self.i_ffn_gate);
let t = self.chain(role, i, n, default);
self.i_ffn_gate += 1;
t
}
Role::FfnUp => {
let (i, n) = self.ffn_layer(name, self.i_ffn_up);
let t = self.chain(role, i, n, default);
self.i_ffn_up += 1;
t
}
Role::AttnQkv | Role::AttnQ | Role::TokenEmbd | Role::Other => {
self.chain(role, 0, 0, default)
}
};
ty
}
fn output_head_type(&self, shape: &[u64], default: GgmlType) -> GgmlType {
let nx = shape.first().copied().unwrap_or(0) as usize;
let qk_k = default.block_layout().1;
if self.shape.is_falcon || qk_k == 0 || !nx.is_multiple_of(qk_k) {
return GgmlType::Q8_0;
}
self.chain(Role::Output, 0, 0, default)
}
fn chain(&self, role: Role, i: usize, n: usize, default: GgmlType) -> GgmlType {
for p in PROMOTIONS {
if p.role != role || p.ftype != self.target {
continue;
}
if p.not_falcon && self.shape.is_falcon {
continue;
}
if p.when.holds(i, n) {
return p.to;
}
}
default
}
pub fn resolve_all<F>(
target: Target,
shape: ModelShape,
tensors: &[(String, Vec<u64>)],
eligible: F,
) -> BTreeMap<String, GgmlType>
where
F: Fn(&str, &[u64]) -> bool,
{
let names: Vec<String> = tensors.iter().map(|(n, _)| n.clone()).collect();
let mut recipe = Recipe::new(target, shape, &names);
let mut order: Vec<usize> = (0..tensors.len()).collect();
order.sort_by(|&a, &b| llama_cpp_weight_order(&tensors[a].0, &tensors[b].0));
let mut out = BTreeMap::new();
for i in order {
let (name, shape) = &tensors[i];
if eligible(name, shape) {
out.insert(name.clone(), recipe.tensor_type(name, shape));
}
}
out
}
fn ffn_layer(&self, name: &str, counter: usize) -> (usize, usize) {
let n = self.shape.n_layer;
if self.shape.n_expert.max(1) > 1 {
if let Some(rest) = name.strip_prefix("blk.") {
if let Some((digits, _)) = rest.split_once('.') {
if let Ok(i) = digits.parse::<usize>() {
return (i, n);
}
}
}
}
(counter, n)
}
}
fn llama_cpp_weight_order(a: &str, b: &str) -> std::cmp::Ordering {
fn layer(name: &str) -> i64 {
name.strip_prefix("blk.")
.and_then(|rest| rest.split_once('.'))
.and_then(|(digits, _)| digits.parse::<i64>().ok())
.unwrap_or(-1)
}
layer(a).cmp(&layer(b)).then_with(|| a.cmp(b))
}
#[cfg(test)]
mod tests {
use super::*;
fn shape(n_layer: usize) -> ModelShape {
ModelShape {
n_layer,
n_expert: 0,
is_falcon: false,
is_70b: false,
}
}
fn names(n_layer: usize, fused_qkv: bool) -> Vec<String> {
let mut v = vec!["token_embd.weight".to_string()];
for i in 0..n_layer {
if fused_qkv {
v.push(format!("blk.{i}.attn_qkv.weight"));
} else {
v.push(format!("blk.{i}.attn_q.weight"));
v.push(format!("blk.{i}.attn_k.weight"));
v.push(format!("blk.{i}.attn_v.weight"));
}
v.push(format!("blk.{i}.attn_output.weight"));
v.push(format!("blk.{i}.ffn_gate.weight"));
v.push(format!("blk.{i}.ffn_up.weight"));
v.push(format!("blk.{i}.ffn_down.weight"));
}
v.push("output.weight".to_string());
v
}
fn walk(target: Target, sh: ModelShape, names: &[String]) -> Vec<(String, GgmlType)> {
let mut r = Recipe::new(target, sh, names);
names
.iter()
.map(|n| (n.clone(), r.tensor_type(n, &[4096, 4096])))
.collect()
}
#[test]
fn use_more_bits_is_llama_cpps_integer_arithmetic() {
let n = 32;
let got: Vec<usize> = (0..n).filter(|&i| When::UseMoreBits.holds(i, n)).collect();
assert_eq!(
got,
vec![0, 1, 2, 3, 6, 9, 12, 15, 18, 21, 24, 27, 28, 29, 30, 31]
);
assert!(!When::UseMoreBits.holds(0, 0));
}
#[test]
fn every_k_quant_mix_sends_the_output_head_to_q6_k() {
for t in [
Target::Q4_K_S,
Target::Q4_K_M,
Target::Q5_K_S,
Target::Q5_K_M,
] {
let ns = names(4, false);
let got = walk(t, shape(4), &ns);
let out = got.iter().find(|(n, _)| n == "output.weight").unwrap();
assert_eq!(out.1, GgmlType::Q6K, "{t:?}");
let emb = got.iter().find(|(n, _)| n == "token_embd.weight").unwrap();
assert_eq!(emb.1, t.ggml_type(), "{t:?} token_embd");
}
let ns = names(4, false);
let got = walk(Target::Q6_K, shape(4), &ns);
assert_eq!(
got.iter().find(|(n, _)| n == "output.weight").unwrap().1,
GgmlType::Q6K
);
}
#[test]
fn a_tied_embedding_table_gets_the_output_heads_type() {
let ns: Vec<String> = names(4, false)
.into_iter()
.filter(|n| n != "output.weight")
.collect();
let got = walk(Target::Q4_K_M, shape(4), &ns);
assert_eq!(
got.iter()
.find(|(n, _)| n == "token_embd.weight")
.unwrap()
.1,
GgmlType::Q6K
);
}
#[test]
fn q4_k_m_promotes_attn_v_and_ffn_down_on_exactly_the_use_more_bits_layers() {
let n = 16;
let ns = names(n, false);
let got = walk(Target::Q4_K_M, shape(n), &ns);
for (role, want_hits) in [("attn_v", true), ("ffn_down", true)] {
let hits: Vec<usize> = got
.iter()
.filter(|(name, _)| name.contains(role))
.enumerate()
.filter(|(_, (_, t))| *t == GgmlType::Q6K)
.map(|(i, _)| i)
.collect();
let want: Vec<usize> = (0..n).filter(|&i| When::UseMoreBits.holds(i, n)).collect();
assert_eq!(hits, want, "{role}");
assert!(want_hits && !hits.is_empty());
}
for (name, t) in &got {
if name.contains("attn_q.weight") || name.contains("attn_k.weight") {
assert_eq!(*t, GgmlType::Q4K, "{name}");
}
}
}
#[test]
fn q4_k_s_promotes_a_prefix_where_q4_k_m_promotes_a_pattern() {
let n = 16;
let ns = names(n, false);
let got = walk(Target::Q4_K_S, shape(n), &ns);
let v: Vec<GgmlType> = got
.iter()
.filter(|(name, _)| name.contains("attn_v"))
.map(|(_, t)| *t)
.collect();
assert_eq!(&v[..4], &[GgmlType::Q5K; 4]);
assert!(v[4..].iter().all(|t| *t == GgmlType::Q4K));
let d: Vec<GgmlType> = got
.iter()
.filter(|(name, _)| name.contains("ffn_down"))
.map(|(_, t)| *t)
.collect();
assert_eq!(&d[..2], &[GgmlType::Q5K; 2]); assert!(d[2..].iter().all(|t| *t == GgmlType::Q4K));
}
#[test]
fn a_fused_qkv_checkpoint_promotes_every_layer_and_still_counts_them() {
let n = 8;
let ns = names(n, true);
let got = walk(Target::Q4_K_M, shape(n), &ns);
let q: Vec<GgmlType> = got
.iter()
.filter(|(name, _)| name.contains("attn_qkv"))
.map(|(_, t)| *t)
.collect();
assert_eq!(q, vec![GgmlType::Q5K; n]);
let got5 = walk(Target::Q5_K_M, shape(n), &ns);
assert!(got5
.iter()
.filter(|(name, _)| name.contains("attn_qkv"))
.all(|(_, t)| *t == GgmlType::Q6K));
let r = Recipe::new(Target::Q4_K_M, shape(n), &ns);
assert_eq!(r.n_attention_wv, n);
}
#[test]
fn the_post_chain_overrides_beat_the_chain_row_they_follow() {
let n = 16;
let ns = names(n, false);
let moe = ModelShape {
n_expert: 8,
..shape(n)
};
let got = walk(Target::Q4_K_M, moe, &ns);
assert!(got
.iter()
.filter(|(name, _)| name.contains("attn_v") || name.contains("attn_k"))
.all(|(_, t)| *t == GgmlType::Q8_0));
assert!(got
.iter()
.filter(|(name, _)| name.contains("attn_output"))
.all(|(_, t)| *t == GgmlType::Q5K));
let big = ModelShape {
is_70b: true,
..shape(n)
};
let got = walk(Target::Q4_K_M, big, &ns);
let v: Vec<GgmlType> = got
.iter()
.filter(|(name, _)| name.contains("attn_v"))
.map(|(_, t)| *t)
.collect();
for (i, t) in v.iter().enumerate() {
let want = if When::UseMoreBits.holds(i, n) {
GgmlType::Q6K
} else {
GgmlType::Q5K
};
assert_eq!(*t, want, "layer {i}");
}
let falcon = ModelShape {
is_falcon: true,
..shape(n)
};
let got = walk(Target::Q4_K_M, falcon, &ns);
assert_eq!(
got.iter().find(|(n, _)| n == "output.weight").unwrap().1,
GgmlType::Q8_0
);
}
#[test]
fn an_output_head_with_an_awkward_row_length_goes_to_q8_0() {
let ns = names(2, false);
let mut r = Recipe::new(Target::Q4_K_M, shape(2), &ns);
assert_eq!(
r.tensor_type("output.weight", &[4096 + 32, 100]),
GgmlType::Q8_0
);
let mut r = Recipe::new(Target::Q4_K_M, shape(2), &ns);
assert_eq!(r.tensor_type("output.weight", &[4096, 100]), GgmlType::Q6K);
}
#[test]
fn the_recipe_only_promotes_to_types_frink_can_encode() {
let writable = [GgmlType::Q8_0, GgmlType::Q4K, GgmlType::Q5K, GgmlType::Q6K];
for p in PROMOTIONS {
assert!(
writable.contains(&p.to),
"{:?}/{:?} promotes to {:?}, which has no frink encoder",
p.role,
p.ftype,
p.to
);
}
assert!(writable.contains(&FALCON_Q4KM_FFN_DOWN_SIXTEENTH));
for t in [GgmlType::Q5K, GgmlType::Q8_0, GgmlType::Q6K] {
assert!(writable.contains(&t));
}
}
#[test]
fn the_table_and_the_target_list_cover_each_other() {
for p in PROMOTIONS {
assert!(Target::ALL.contains(&p.ftype), "{:?}", p.ftype);
}
let no_rows: Vec<&'static str> = Target::ALL
.iter()
.filter(|t| !PROMOTIONS.iter().any(|p| p.ftype == **t))
.map(|t| t.name())
.collect();
assert_eq!(no_rows, vec!["Q8_0", "Q6_K"]);
}
#[test]
fn tensor_names_land_in_llama_cpps_arms() {
let cases: &[(&str, Role)] = &[
("output.weight", Role::Output),
("token_embd.weight", Role::TokenEmbd),
("blk.0.attn_v.weight", Role::AttnV),
("blk.0.attn_k.weight", Role::AttnK),
("blk.0.attn_q.weight", Role::AttnQ),
("blk.0.attn_qkv.weight", Role::AttnQkv),
("blk.0.attn_output.weight", Role::AttnOutput),
("blk.0.ffn_down.weight", Role::FfnDown),
("blk.0.ffn_down_exps.weight", Role::FfnDown),
("blk.0.ffn_down_shexp.weight", Role::FfnDown),
("blk.0.ffn_gate.weight", Role::FfnGate),
("blk.0.ffn_up.weight", Role::FfnUp),
("blk.0.attn_norm.weight", Role::Other),
];
for (name, want) in cases {
assert_eq!(Role::of(name, true), *want, "{name}");
}
assert_eq!(Role::of("token_embd.weight", false), Role::Output);
}
#[test]
fn tensors_are_walked_in_llama_cpps_weights_map_order_not_the_files() {
let mut names = vec![
"blk.10.ffn_down.weight",
"token_embd.weight",
"blk.2.attn_v.weight",
"output_norm.weight",
"blk.2.attn_k.weight",
"blk.1.ffn_down.weight",
];
names.sort_by(|a, b| llama_cpp_weight_order(a, b));
assert_eq!(
names,
vec![
"output_norm.weight",
"token_embd.weight",
"blk.1.ffn_down.weight",
"blk.2.attn_k.weight",
"blk.2.attn_v.weight",
"blk.10.ffn_down.weight",
]
);
}
#[test]
fn resolve_all_is_independent_of_the_order_it_is_handed_the_tensors() {
let n = 16;
let tensors: Vec<(String, Vec<u64>)> = names(n, false)
.into_iter()
.map(|s| (s, vec![4096u64, 4096]))
.collect();
let forward = Recipe::resolve_all(Target::Q4_K_M, shape(n), &tensors, |_, _| true);
let mut reversed = tensors.clone();
reversed.reverse();
assert_eq!(
forward,
Recipe::resolve_all(Target::Q4_K_M, shape(n), &reversed, |_, _| true)
);
let mut lexicographic = tensors.clone();
lexicographic.sort_by(|a, b| a.0.cmp(&b.0));
assert_eq!(
forward,
Recipe::resolve_all(Target::Q4_K_M, shape(n), &lexicographic, |_, _| true)
);
let distinct: std::collections::BTreeSet<_> =
forward.values().map(|t| format!("{t:?}")).collect();
assert!(distinct.len() > 1, "{distinct:?}");
}
#[test]
fn a_moe_checkpoints_ffn_layer_comes_from_the_name_not_the_counter() {
let n = 16;
let mut ns = names(n, false);
ns.retain(|s| !s.contains("ffn_down"));
let moe = ModelShape {
n_expert: 4,
..shape(n)
};
let mut r = Recipe::new(Target::Q4_K_M, moe, &ns);
assert_eq!(
r.tensor_type("blk.5.ffn_down.weight", &[4096, 4096]),
GgmlType::Q4K
);
assert_eq!(
r.tensor_type("blk.0.ffn_down.weight", &[4096, 4096]),
GgmlType::Q6K
);
}
}