use crate::loader::load_f32_vec;
use crate::norm::{norm_function, NormFunction, NormOp, NormParam};
use crate::LoadError;
use frink_gguf::{GgufError, TensorSource};
pub const PRE_FFN_NORM_IS_POST_ATTENTION_NORM: &[&str] = &[
"gpt-oss",
"seed_oss",
"glm4moe",
"qwen35",
"qwen35moe",
"qwen3next",
];
pub const PRE_FFN_NORM_IS_ATTN_OUTPUT_NORM: &[&str] = &["dbrx"];
pub const POST_NORMS_UNDER_GROK_NAMES: &[&str] = &["grok"];
pub const ATTN_NORM_2_FEEDS_ATTENTION: &[&str] = &["falcon"];
pub const EMBEDDING_NORM_ARCHITECTURES: &[&str] = &["bloom"];
pub const WEIGHTLESS_EMBEDDING_NORM: &[&str] = &["muse-glimmer"];
pub const NO_OUTPUT_NORM: &[&str] = &["hrm_text"];
pub fn weightless_embedding_norm(arch: &str) -> bool {
WEIGHTLESS_EMBEDDING_NORM.contains(&arch)
}
pub const OUTPUT_NORM_UNDER_EMBEDDING_NAME: &[&str] = &["lfm2", "lfm2moe"];
pub const ONE_NORM_PER_LAYER: &[&str] = &["nemotron_h", "nemotron_h_moe"];
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct StoredNorm {
pub names: &'static [&'static str],
pub required: bool,
}
impl StoredNorm {
const fn required(names: &'static [&'static str]) -> Self {
Self {
names,
required: true,
}
}
const fn optional(names: &'static [&'static str]) -> Self {
Self {
names,
required: false,
}
}
fn candidates(&self, layer: Option<usize>) -> Vec<String> {
self.candidates_for(layer, NormParam::Weight)
}
fn candidates_for(&self, layer: Option<usize>, part: NormParam) -> Vec<String> {
self.names
.iter()
.flat_map(|base| {
let full = match layer {
Some(l) => format!("blk.{l}.{base}"),
None => (*base).to_string(),
};
match part {
NormParam::Weight => vec![format!("{full}.weight"), full],
NormParam::Bias => vec![format!("{full}.bias")],
}
})
.collect()
}
pub fn load(
&self,
file: &impl TensorSource,
layer: Option<usize>,
) -> Result<Option<Vec<f32>>, LoadError> {
let candidates = self.candidates(layer);
if let Some(name) = candidates.iter().find(|n| file.find_tensor(n).is_some()) {
return load_f32_vec(file, name).map(Some);
}
if self.required {
return Err(LoadError::Gguf(GgufError::TensorNotFound(
candidates.join(" | "),
)));
}
Ok(None)
}
fn load_required_part(
&self,
file: &impl TensorSource,
layer: Option<usize>,
part: NormParam,
) -> Result<Vec<f32>, LoadError> {
debug_assert!(self.required, "load_required_part on an optional site");
let candidates = self.candidates_for(layer, part);
match candidates.iter().find(|n| file.find_tensor(n).is_some()) {
Some(name) => load_f32_vec(file, name),
None => Err(LoadError::Gguf(GgufError::TensorNotFound(
candidates.join(" | "),
))),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct NormSites {
pub function: NormFunction,
pub attn: Option<StoredNorm>,
pub ffn: Option<StoredNorm>,
pub post_attn: Option<StoredNorm>,
pub post_ffn: Option<StoredNorm>,
pub output: Option<StoredNorm>,
pub embedding: Option<StoredNorm>,
}
impl NormSites {
pub fn for_arch(arch: &str) -> Self {
Self::with_function(arch, norm_function(arch))
}
pub fn with_function(arch: &str, function: NormFunction) -> Self {
let mut sites = Self {
function,
attn: Some(StoredNorm::required(&["attn_norm"])),
ffn: Some(StoredNorm::required(&["ffn_norm"])),
post_attn: Some(StoredNorm::optional(&["post_attention_norm"])),
post_ffn: Some(StoredNorm::optional(&["post_ffw_norm"])),
output: if NO_OUTPUT_NORM.contains(&arch) {
None
} else if OUTPUT_NORM_UNDER_EMBEDDING_NAME.contains(&arch) {
Some(StoredNorm::required(&["token_embd_norm"]))
} else {
Some(StoredNorm::required(&["output_norm"]))
},
embedding: EMBEDDING_NORM_ARCHITECTURES
.contains(&arch)
.then_some(StoredNorm::required(&["token_embd_norm"])),
};
if crate::capability::is_post_norm_only(arch) {
sites.attn = None;
sites.ffn = None;
}
if PRE_FFN_NORM_IS_POST_ATTENTION_NORM.contains(&arch) {
sites.ffn = Some(StoredNorm::required(&["post_attention_norm"]));
sites.post_attn = None;
}
if PRE_FFN_NORM_IS_ATTN_OUTPUT_NORM.contains(&arch) {
sites.ffn = Some(StoredNorm::required(&["attn_output_norm"]));
}
if ONE_NORM_PER_LAYER.contains(&arch) {
sites.ffn = Some(StoredNorm::required(&["attn_norm"]));
}
if POST_NORMS_UNDER_GROK_NAMES.contains(&arch) {
sites.post_attn = Some(StoredNorm::required(&["attn_output_norm"]));
sites.post_ffn = Some(StoredNorm::required(&[
"layer_output_norm",
"post_ffw_norm",
]));
}
sites
}
pub fn for_layer(&self, arch: &str, file: &impl TensorSource, layer: usize) -> Self {
if !ATTN_NORM_2_FEEDS_ATTENTION.contains(&arch)
|| file
.find_tensor(&format!("blk.{layer}.attn_norm_2.weight"))
.is_none()
{
return *self;
}
Self {
attn: Some(StoredNorm::required(&["attn_norm_2"])),
ffn: Some(StoredNorm::required(&["attn_norm"])),
..*self
}
}
pub fn load_pre_norm(
&self,
site: Option<StoredNorm>,
file: &impl TensorSource,
layer: Option<usize>,
) -> Result<NormOp, LoadError> {
match site {
None => Ok(NormOp::None),
Some(stored) => self
.function
.resolve(|part| stored.load_required_part(file, layer, part)),
}
}
pub fn load_post_norm(
site: Option<StoredNorm>,
file: &impl TensorSource,
layer: usize,
) -> Result<Option<Vec<f32>>, LoadError> {
match site {
None => Ok(None),
Some(stored) => stored.load(file, Some(layer)),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn attn_norm_2_crosses_the_two_pre_norm_slots_per_layer() {
use crate::test_source::StubSource;
let row = NormSites::for_arch("falcon");
assert_eq!(row.function, NormFunction::LayerNormBias);
let seven_b = StubSource::with_tensors(&["blk.0.attn_norm.weight"]);
assert_eq!(row.for_layer("falcon", &seven_b, 0), row);
let forty_b = StubSource::with_tensors(&["blk.1.attn_norm_2.weight"]);
assert_eq!(
row.for_layer("falcon", &forty_b, 0),
row,
"layer 0 has none"
);
let crossed = row.for_layer("falcon", &forty_b, 1);
assert_eq!(crossed.attn, Some(StoredNorm::required(&["attn_norm_2"])));
assert_eq!(crossed.ffn, Some(StoredNorm::required(&["attn_norm"])));
assert_eq!(crossed.output, row.output);
let llama = NormSites::for_arch("llama");
assert_eq!(llama.for_layer("llama", &forty_b, 1), llama);
}
#[test]
fn the_default_row_is_the_plain_pre_norm_layer() {
let s = NormSites::for_arch("llama");
assert_eq!(s.function, NormFunction::Rms);
assert_eq!(s.attn, Some(StoredNorm::required(&["attn_norm"])));
assert_eq!(s.ffn, Some(StoredNorm::required(&["ffn_norm"])));
assert_eq!(
s.post_attn,
Some(StoredNorm::optional(&["post_attention_norm"]))
);
assert_eq!(s.post_ffn, Some(StoredNorm::optional(&["post_ffw_norm"])));
assert_eq!(s.output, Some(StoredNorm::required(&["output_norm"])));
assert_eq!(s.embedding, None);
assert_eq!(
NormSites::for_arch("bloom").embedding,
Some(StoredNorm::required(&["token_embd_norm"]))
);
}
#[test]
fn nemotron_h_has_one_norm_per_layer() {
let s = NormSites::for_arch("nemotron_h");
assert_eq!(s.attn, Some(StoredNorm::required(&["attn_norm"])));
assert_eq!(s.ffn, Some(StoredNorm::required(&["attn_norm"])));
}
#[test]
fn token_embd_norm_is_the_output_norm_on_lfm2_alone() {
let s = NormSites::for_arch("lfm2");
assert_eq!(s.output, Some(StoredNorm::required(&["token_embd_norm"])));
assert_eq!(s.embedding, None);
let b = NormSites::for_arch("bloom");
assert_eq!(b.output, Some(StoredNorm::required(&["output_norm"])));
for arch in OUTPUT_NORM_UNDER_EMBEDDING_NAME {
assert!(
!EMBEDDING_NORM_ARCHITECTURES.contains(arch),
"{arch}: one tensor cannot feed both sites"
);
}
}
#[test]
fn attn_output_norm_is_dbrxs_pre_ffn_norm_and_groks_post_attention_norm() {
let dbrx = NormSites::for_arch("dbrx");
assert_eq!(dbrx.function, NormFunction::LayerNorm);
assert_eq!(
dbrx.ffn,
Some(StoredNorm::required(&["attn_output_norm"])),
"dbrx.cpp:110-113 norms ffn_inp with attn_out_norm"
);
assert_eq!(
dbrx.post_attn,
Some(StoredNorm::optional(&["post_attention_norm"])),
"no post-attention norm of its own"
);
let grok = NormSites::for_arch("grok");
assert_eq!(grok.function, NormFunction::Rms);
assert_eq!(grok.ffn, Some(StoredNorm::required(&["ffn_norm"])));
assert_eq!(
grok.post_attn,
Some(StoredNorm::required(&["attn_output_norm"])),
"grok.cpp:143-146 norms the attention output before the residual add"
);
assert_eq!(
grok.post_ffn,
Some(StoredNorm::required(&[
"layer_output_norm",
"post_ffw_norm"
])),
"grok.cpp:75-78: LAYER_OUT_NORM first, FFN_POST_NORM as the fallback"
);
}
#[test]
fn the_two_no_tensor_shapes_are_kept_apart() {
let olmo2 = NormSites::for_arch("olmo2");
assert_eq!(olmo2.attn, None);
assert_eq!(olmo2.ffn, None);
assert_eq!(olmo2.function, NormFunction::Rms);
let olmo = NormSites::for_arch("olmo");
assert_eq!(olmo.function, NormFunction::LayerNormNoParams);
assert!(
olmo.attn.is_some(),
"the site exists; the function skips the read"
);
}
#[test]
fn the_gpt_oss_slot_reads_the_tensor_once() {
let s = NormSites::for_arch("gpt-oss");
assert_eq!(s.ffn, Some(StoredNorm::required(&["post_attention_norm"])));
assert_eq!(s.post_attn, None);
}
#[test]
fn candidates_try_the_suffixed_spelling_first_for_each_name() {
let s = StoredNorm::required(&["layer_output_norm", "post_ffw_norm"]);
assert_eq!(
s.candidates(Some(3)),
vec![
"blk.3.layer_output_norm.weight",
"blk.3.layer_output_norm",
"blk.3.post_ffw_norm.weight",
"blk.3.post_ffw_norm",
]
);
assert_eq!(
StoredNorm::required(&["output_norm"]).candidates(None),
vec!["output_norm.weight", "output_norm"]
);
}
}