#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum LogitScaleUse {
#[default]
NotApplied,
Reciprocal,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum MultiplierDefaults {
#[default]
FromFileOnly,
MiniCpm,
}
impl MultiplierDefaults {
pub fn values(self, dims: MultiplierDims) -> DeclaredMultipliers {
match self {
Self::FromFileOnly => DeclaredMultipliers::default(),
Self::MiniCpm => DeclaredMultipliers {
logit: Some(if dims.n_embd == 0 {
1.0
} else {
256.0 / dims.n_embd as f32
}),
residual: Some(1.4 / (dims.n_layer as f32).sqrt()),
embedding: Some(12.0),
attention: None,
},
}
}
fn merge(self, declared: DeclaredMultipliers, dims: MultiplierDims) -> DeclaredMultipliers {
let DeclaredMultipliers {
logit,
residual,
embedding,
attention,
} = declared;
let d = self.values(dims);
DeclaredMultipliers {
logit: logit.or(d.logit),
residual: residual.or(d.residual),
embedding: embedding.or(d.embedding),
attention: attention.or(d.attention),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct MultiplierDims {
pub head_dim: usize,
pub n_layer: usize,
pub n_embd: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct MultiplierSupport {
pub embedding: bool,
pub residual: bool,
pub logit: LogitScaleUse,
pub attention: bool,
pub defaults: MultiplierDefaults,
}
impl MultiplierSupport {
pub const NONE: Self = Self {
embedding: false,
residual: false,
logit: LogitScaleUse::NotApplied,
attention: false,
defaults: MultiplierDefaults::FromFileOnly,
};
pub const GRANITE: Self = Self {
embedding: true,
residual: true,
logit: LogitScaleUse::Reciprocal,
attention: true,
defaults: MultiplierDefaults::FromFileOnly,
};
pub const MINICPM: Self = Self {
embedding: true,
residual: true,
logit: LogitScaleUse::Reciprocal,
attention: false,
defaults: MultiplierDefaults::MiniCpm,
};
}
const MULTIPLIER_ARCHITECTURES: &[(&str, MultiplierSupport)] = &[
("granite", MultiplierSupport::GRANITE),
("granitemoe", MultiplierSupport::GRANITE),
("granite-moe", MultiplierSupport::GRANITE),
("minicpm", MultiplierSupport::MINICPM),
];
pub fn multiplier_support(arch: &str) -> MultiplierSupport {
MULTIPLIER_ARCHITECTURES
.iter()
.find(|(n, _)| *n == arch)
.map_or(MultiplierSupport::NONE, |(_, s)| *s)
}
#[derive(Debug, Clone, Copy, PartialEq, Default)]
pub struct DeclaredMultipliers {
pub logit: Option<f32>,
pub residual: Option<f32>,
pub embedding: Option<f32>,
pub attention: Option<f32>,
}
#[derive(Debug, Clone, Copy, PartialEq, Default)]
pub struct ResolvedMultipliers {
pub embedding_scale: Option<f32>,
pub residual_scale: Option<f32>,
pub logit_multiplier: Option<f32>,
pub attention_scale: Option<f32>,
}
#[derive(Debug, Clone, PartialEq)]
pub enum MultiplierError {
MissingRequiredLogitScale,
NonPositiveLogitScale(f32),
}
impl MultiplierError {
pub fn message(&self, arch: &str) -> String {
match self {
MultiplierError::MissingRequiredLogitScale => format!(
"`{arch}.logit_scale` is REQUIRED for this architecture (src/models/granite.cpp:7 \
reads it with no default) and the file does not declare it; llama.cpp refuses \
the same file"
),
MultiplierError::NonPositiveLogitScale(v) => format!(
"`{arch}.logit_scale` = {v}: the graph divides every logit by it \
(src/models/granite.cpp:180), so zero is a division by zero and a negative \
value reorders the vocabulary"
),
}
}
}
fn scale_or_none(v: Option<f32>) -> Option<f32> {
v.filter(|&v| v != 0.0 && v != 1.0)
}
pub fn resolve(
support: MultiplierSupport,
declared: DeclaredMultipliers,
dims: MultiplierDims,
) -> Result<ResolvedMultipliers, MultiplierError> {
let DeclaredMultipliers {
logit,
residual,
embedding,
attention,
} = support.defaults.merge(declared, dims);
let logit_multiplier = match support.logit {
LogitScaleUse::NotApplied => None,
LogitScaleUse::Reciprocal => {
let v = logit.ok_or(MultiplierError::MissingRequiredLogitScale)?;
if v <= 0.0 {
return Err(MultiplierError::NonPositiveLogitScale(v));
}
scale_or_none(Some(1.0 / v))
}
};
let attention_scale = if support.attention {
attention.filter(|&v| v != 0.0).filter(|&v| {
let kernel = 1.0 / (dims.head_dim as f32).sqrt();
(v - kernel).abs() > f32::EPSILON * kernel.max(1.0)
})
} else {
None
};
Ok(ResolvedMultipliers {
embedding_scale: support
.embedding
.then(|| scale_or_none(embedding))
.flatten(),
residual_scale: support.residual.then(|| scale_or_none(residual)).flatten(),
logit_multiplier,
attention_scale,
})
}
#[inline]
pub fn residual_add(hidden: &mut [f32], branch: &[f32], scale: Option<f32>) {
debug_assert_eq!(hidden.len(), branch.len());
match scale {
None => {
for (h, b) in hidden.iter_mut().zip(branch.iter()) {
*h += *b;
}
}
Some(s) => {
for (h, b) in hidden.iter_mut().zip(branch.iter()) {
*h += s * *b;
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn dims(head_dim: usize) -> MultiplierDims {
MultiplierDims {
head_dim,
n_layer: 2,
n_embd: 24,
}
}
#[test]
fn the_three_granite_rows_have_identical_multiplier_support() {
let dense = multiplier_support("granite");
assert_eq!(dense, MultiplierSupport::GRANITE);
for alias in ["granitemoe", "granite-moe"] {
assert_eq!(
multiplier_support(alias),
dense,
"`{alias}` must scale exactly like `granite`"
);
}
}
#[test]
fn an_architecture_that_was_not_read_applies_no_multipliers() {
for arch in ["llama", "qwen3", "deepseek", "not-an-architecture"] {
assert_eq!(
multiplier_support(arch),
MultiplierSupport::NONE,
"`{arch}` must not scale"
);
}
}
#[test]
fn the_gemma_family_reads_none_of_the_four_keys() {
for arch in ["gemma", "gemma2", "gemma3"] {
assert_eq!(
multiplier_support(arch),
MultiplierSupport::NONE,
"`{arch}` computes its scales; it does not read them"
);
}
}
#[test]
fn granites_logit_scale_is_inverted_at_load_time() {
let got = resolve(
MultiplierSupport::GRANITE,
DeclaredMultipliers {
logit: Some(8.0),
..Default::default()
},
dims(64),
)
.expect("8.0 resolves");
assert_eq!(got.logit_multiplier, Some(0.125));
}
#[test]
fn a_granite_file_with_no_logit_scale_is_refused_rather_than_defaulted() {
assert_eq!(
resolve(
MultiplierSupport::GRANITE,
DeclaredMultipliers::default(),
dims(64),
),
Err(MultiplierError::MissingRequiredLogitScale)
);
assert!(
MultiplierError::MissingRequiredLogitScale
.message("granite")
.contains("granite.logit_scale"),
"the message must name the key"
);
}
#[test]
fn a_non_positive_logit_scale_is_refused() {
for bad in [0.0f32, -2.0] {
assert_eq!(
resolve(
MultiplierSupport::GRANITE,
DeclaredMultipliers {
logit: Some(bad),
..Default::default()
},
dims(64),
),
Err(MultiplierError::NonPositiveLogitScale(bad)),
"logit_scale {bad} must be refused"
);
}
}
#[test]
fn each_keys_own_no_op_value_is_what_switches_it_off() {
let got = resolve(
MultiplierSupport::GRANITE,
DeclaredMultipliers {
logit: Some(1.0),
residual: Some(1.0),
embedding: Some(1.0),
attention: Some(0.0),
},
dims(64),
)
.expect("all no-ops resolve");
assert_eq!(got, ResolvedMultipliers::default(), "{got:?}");
let got = resolve(
MultiplierSupport::GRANITE,
DeclaredMultipliers {
logit: Some(2.0),
residual: Some(0.0),
embedding: Some(0.0),
attention: Some(1.0),
},
dims(64),
)
.expect("resolves");
assert_eq!(
got.residual_scale, None,
"llama.cpp's `if (f_residual_scale)` guard makes 0.0 mean off"
);
assert_eq!(got.embedding_scale, None, "llama-graph.cpp:2337 likewise");
assert_eq!(
got.attention_scale,
Some(1.0),
"1.0 is a real attention-scale override, not its sentinel"
);
}
#[test]
fn an_attention_scale_equal_to_the_kernels_own_resolves_to_none() {
let head_dim = 64;
let kernel = 1.0 / (head_dim as f32).sqrt();
let got = resolve(
MultiplierSupport::GRANITE,
DeclaredMultipliers {
logit: Some(2.0),
attention: Some(kernel),
..Default::default()
},
dims(head_dim),
)
.expect("resolves");
assert_eq!(got.attention_scale, None);
let got = resolve(
MultiplierSupport::GRANITE,
DeclaredMultipliers {
logit: Some(2.0),
attention: Some(0.015_625),
..Default::default()
},
dims(head_dim),
)
.expect("resolves");
assert_eq!(got.attention_scale, Some(0.015_625));
}
#[test]
fn support_gates_the_value_rather_than_the_value_gating_itself() {
let got = resolve(
MultiplierSupport::NONE,
DeclaredMultipliers {
logit: Some(8.0),
residual: Some(0.22),
embedding: Some(12.0),
attention: Some(0.015_625),
},
dims(64),
)
.expect("an unsupported logit_scale is not even read");
assert_eq!(got, ResolvedMultipliers::default());
}
#[test]
fn the_residual_add_scales_the_branch_and_not_the_stream() {
let mut hidden = vec![1.0f32, 2.0, 3.0];
residual_add(&mut hidden, &[10.0, 20.0, 30.0], Some(0.5));
assert_eq!(hidden, vec![6.0, 12.0, 18.0]);
let mut hidden = vec![1.0f32, 2.0, 3.0];
residual_add(&mut hidden, &[10.0, 20.0, 30.0], None);
assert_eq!(
hidden,
vec![11.0, 22.0, 33.0],
"no scale must be exactly the unscaled add, not a multiply by 1.0"
);
}
}