#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum LogitScaleUse {
#[default]
NotApplied,
Reciprocal,
AsIs,
AsIsOptional,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum AttentionScaleKey {
#[default]
NotRead,
Scale,
OutputScale,
}
impl AttentionScaleKey {
pub fn suffix(self) -> Option<&'static str> {
match self {
Self::NotRead => None,
Self::Scale => Some("attention.scale"),
Self::OutputScale => Some("attention.output_scale"),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum MultiplierDefaults {
#[default]
FromFileOnly,
MiniCpm,
Grok,
}
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,
},
Self::Grok => DeclaredMultipliers {
logit: Some(0.577_350_3),
residual: None,
embedding: Some(78.383_67),
attention: Some(0.088_388_35),
},
}
}
pub fn attn_logit_softcap(self) -> Option<f32> {
match self {
Self::FromFileOnly | Self::MiniCpm => None,
Self::Grok => Some(30.0),
}
}
pub fn yarn_beta_fast(self) -> Option<f32> {
match self {
Self::FromFileOnly | Self::MiniCpm => None,
Self::Grok => Some(8.0),
}
}
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 enum ResidualScaleUse {
#[default]
NotRead,
BranchOutput,
NormedInputRequired,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct MultiplierSupport {
pub embedding: bool,
pub residual: ResidualScaleUse,
pub logit: LogitScaleUse,
pub attention: AttentionScaleKey,
pub defaults: MultiplierDefaults,
}
impl MultiplierSupport {
pub const NONE: Self = Self {
embedding: false,
residual: ResidualScaleUse::NotRead,
logit: LogitScaleUse::NotApplied,
attention: AttentionScaleKey::NotRead,
defaults: MultiplierDefaults::FromFileOnly,
};
pub const GRANITE: Self = Self {
embedding: true,
residual: ResidualScaleUse::BranchOutput,
logit: LogitScaleUse::Reciprocal,
attention: AttentionScaleKey::Scale,
defaults: MultiplierDefaults::FromFileOnly,
};
pub const MINICPM: Self = Self {
embedding: true,
residual: ResidualScaleUse::BranchOutput,
logit: LogitScaleUse::Reciprocal,
attention: AttentionScaleKey::NotRead,
defaults: MultiplierDefaults::MiniCpm,
};
pub const GROK: Self = Self {
embedding: true,
residual: ResidualScaleUse::NotRead,
logit: LogitScaleUse::AsIs,
attention: AttentionScaleKey::OutputScale,
defaults: MultiplierDefaults::Grok,
};
pub const TALKIE: Self = Self {
embedding: false,
residual: ResidualScaleUse::NotRead,
logit: LogitScaleUse::AsIs,
attention: AttentionScaleKey::NotRead,
defaults: MultiplierDefaults::FromFileOnly,
};
pub const COMMAND_R: Self = Self {
embedding: false,
residual: ResidualScaleUse::NotRead,
logit: LogitScaleUse::AsIsOptional,
attention: AttentionScaleKey::NotRead,
defaults: MultiplierDefaults::FromFileOnly,
};
pub const EMBEDDING_ONLY: Self = Self {
embedding: true,
residual: ResidualScaleUse::NotRead,
logit: LogitScaleUse::NotApplied,
attention: AttentionScaleKey::NotRead,
defaults: MultiplierDefaults::FromFileOnly,
};
pub const MINIMAX_01: Self = Self {
embedding: false,
residual: ResidualScaleUse::NormedInputRequired,
logit: LogitScaleUse::NotApplied,
attention: AttentionScaleKey::NotRead,
defaults: MultiplierDefaults::FromFileOnly,
};
pub const COHERE2: Self = Self {
embedding: false,
residual: ResidualScaleUse::NotRead,
logit: LogitScaleUse::AsIs,
attention: AttentionScaleKey::NotRead,
defaults: MultiplierDefaults::FromFileOnly,
};
}
const MULTIPLIER_ARCHITECTURES: &[(&str, MultiplierSupport)] = &[
("granite", MultiplierSupport::GRANITE),
("granitemoe", MultiplierSupport::GRANITE),
("granite-moe", MultiplierSupport::GRANITE),
("granitehybrid", MultiplierSupport::GRANITE),
("granite-hybrid", MultiplierSupport::GRANITE),
("granite_swa", MultiplierSupport::GRANITE),
("minicpm", MultiplierSupport::MINICPM),
("grok", MultiplierSupport::GROK),
("talkie", MultiplierSupport::TALKIE),
("command-r", MultiplierSupport::COMMAND_R),
("cohere2", MultiplierSupport::COHERE2),
("minimax-01", MultiplierSupport::MINIMAX_01),
("muse-glimmer", MultiplierSupport::COHERE2),
("hrm_text", MultiplierSupport::EMBEDDING_ONLY),
("cohere2moe", MultiplierSupport::COHERE2),
];
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 normed_residual_scale: Option<f32>,
pub logit_multiplier: Option<f32>,
pub attention_scale: Option<f32>,
}
#[derive(Debug, Clone, PartialEq)]
pub enum MultiplierError {
MissingRequiredLogitScale,
NonPositiveLogitScale(f32),
MissingRequiredResidualScale,
}
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::MissingRequiredResidualScale => format!(
"`{arch}.residual_scale` is REQUIRED for this architecture \
(src/models/minimax-01.cpp:6 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 scales every logit by it \
(src/models/granite.cpp:180 divides, src/models/grok.cpp:211 multiplies), so \
zero blanks or divides away the whole vocabulary and a negative value \
reorders it"
),
}
}
}
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))
}
LogitScaleUse::AsIs => {
let v = logit.ok_or(MultiplierError::MissingRequiredLogitScale)?;
if v <= 0.0 {
return Err(MultiplierError::NonPositiveLogitScale(v));
}
scale_or_none(Some(v))
}
LogitScaleUse::AsIsOptional => match logit {
None | Some(0.0) => None,
Some(v) if v < 0.0 => return Err(MultiplierError::NonPositiveLogitScale(v)),
Some(v) => scale_or_none(Some(v)),
},
};
let restates_kernel = |v: &f32| {
let kernel = 1.0 / (dims.head_dim as f32).sqrt();
(v - kernel).abs() > f32::EPSILON * kernel.max(1.0)
};
let attention_scale = match support.attention {
AttentionScaleKey::NotRead => None,
AttentionScaleKey::Scale => attention.filter(|&v| v != 0.0).filter(restates_kernel),
AttentionScaleKey::OutputScale => attention.filter(restates_kernel),
};
Ok(ResolvedMultipliers {
embedding_scale: support
.embedding
.then(|| scale_or_none(embedding))
.flatten(),
residual_scale: match support.residual {
ResidualScaleUse::BranchOutput => scale_or_none(residual),
ResidualScaleUse::NotRead | ResidualScaleUse::NormedInputRequired => None,
},
normed_residual_scale: match support.residual {
ResidualScaleUse::NormedInputRequired => {
Some(residual.ok_or(MultiplierError::MissingRequiredResidualScale)?)
}
ResidualScaleUse::NotRead | ResidualScaleUse::BranchOutput => None,
},
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"
);
}
#[test]
fn a_grok_file_declaring_nothing_is_scaled_by_all_of_grok_cpps_defaults() {
let got = resolve(
MultiplierSupport::GROK,
DeclaredMultipliers::default(),
dims(6),
)
.expect("defaults resolve");
assert_eq!(got.embedding_scale, Some(78.383_67));
assert_eq!(
got.logit_multiplier,
Some(0.577_350_3),
"grok.cpp:211 MULTIPLIES by f_logit_scale; a reciprocal here would be 1.732"
);
assert_eq!(got.attention_scale, Some(0.088_388_35));
assert_eq!(got.residual_scale, None, "grok has no residual multiplier");
let real = resolve(
MultiplierSupport::GROK,
DeclaredMultipliers::default(),
MultiplierDims {
head_dim: 128,
n_layer: 64,
n_embd: 6144,
},
)
.expect("resolves");
assert_eq!(real.attention_scale, None);
assert_eq!(MultiplierDefaults::Grok.attn_logit_softcap(), Some(30.0));
assert_eq!(MultiplierDefaults::Grok.yarn_beta_fast(), Some(8.0));
assert_eq!(MultiplierDefaults::MiniCpm.attn_logit_softcap(), None);
assert_eq!(MultiplierDefaults::FromFileOnly.yarn_beta_fast(), None);
}
#[test]
fn command_rs_logit_scale_is_optional_and_talkies_is_not() {
let with = |logit: Option<f32>| DeclaredMultipliers {
logit,
..Default::default()
};
let cr = |logit| resolve(MultiplierSupport::COMMAND_R, with(logit), dims(6));
assert_eq!(cr(None).expect("absent is no scale").logit_multiplier, None);
assert_eq!(
cr(Some(0.0)).expect("zero is no scale").logit_multiplier,
None
);
assert_eq!(
cr(Some(0.0625)).expect("resolves").logit_multiplier,
Some(0.0625)
);
assert!(matches!(
cr(Some(-1.0)),
Err(MultiplierError::NonPositiveLogitScale(_))
));
assert!(matches!(
resolve(MultiplierSupport::TALKIE, with(None), dims(6)),
Err(MultiplierError::MissingRequiredLogitScale)
));
let cr_keys = crate::capability::unsupported_scaling_keys("command-r");
assert!(cr_keys
.iter()
.any(|(k, _, _)| k == "command-r.residual_scale"));
assert!(!cr_keys.iter().any(|(k, _, _)| k == "command-r.logit_scale"));
}
#[test]
fn a_grok_file_declaring_its_keys_overrides_every_default() {
let got = resolve(
MultiplierSupport::GROK,
DeclaredMultipliers {
logit: Some(2.5),
residual: None,
embedding: Some(3.0),
attention: Some(0.25),
},
dims(6),
)
.expect("resolves");
assert_eq!(got.logit_multiplier, Some(2.5), "multiplied as-is");
assert_eq!(got.embedding_scale, Some(3.0));
assert_eq!(got.attention_scale, Some(0.25));
}
#[test]
fn the_output_scale_key_has_no_sentinel_and_the_scale_key_does() {
let grok = resolve(
MultiplierSupport::GROK,
DeclaredMultipliers {
attention: Some(0.0),
..Default::default()
},
dims(6),
)
.expect("resolves");
assert_eq!(grok.attention_scale, Some(0.0));
let granite = resolve(
MultiplierSupport::GRANITE,
DeclaredMultipliers {
logit: Some(2.0),
attention: Some(0.0),
..Default::default()
},
dims(6),
)
.expect("resolves");
assert_eq!(granite.attention_scale, None);
assert_eq!(
AttentionScaleKey::OutputScale.suffix(),
Some("attention.output_scale")
);
assert_eq!(AttentionScaleKey::Scale.suffix(), Some("attention.scale"));
assert_eq!(AttentionScaleKey::NotRead.suffix(), None);
}
#[test]
fn a_non_positive_as_is_logit_scale_is_refused() {
for bad in [0.0f32, -0.5] {
assert_eq!(
resolve(
MultiplierSupport::GROK,
DeclaredMultipliers {
logit: Some(bad),
..Default::default()
},
dims(6),
),
Err(MultiplierError::NonPositiveLogitScale(bad))
);
}
}
}