use frink_gguf::TensorSource;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ParallelNorm {
SharedNorm,
TwoNorms,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ParallelWhen {
Always,
FfnNormAbsent,
ParallelResidualKey,
}
#[derive(Debug, Clone, Copy)]
pub struct ParallelResidual {
pub arch: &'static str,
pub norm: ParallelNorm,
pub when: ParallelWhen,
pub second_norm: Option<&'static str>,
pub lines: &'static str,
}
pub const PARALLEL_RESIDUAL_GRAPHS: &[ParallelResidual] = &[
ParallelResidual {
arch: "stablelm",
norm: ParallelNorm::SharedNorm,
when: ParallelWhen::FfnNormAbsent,
second_norm: None,
lines: "src/models/stablelm.cpp:38-39,129-138,147",
},
ParallelResidual {
arch: "gptneox",
norm: ParallelNorm::TwoNorms,
when: ParallelWhen::ParallelResidualKey,
second_norm: None,
lines: "src/models/gptneox.cpp:5,143-166",
},
ParallelResidual {
arch: "phi2",
norm: ParallelNorm::SharedNorm,
when: ParallelWhen::Always,
second_norm: None,
lines: "src/models/phi2.cpp:67,108,116-117",
},
ParallelResidual {
arch: "falcon",
norm: ParallelNorm::SharedNorm,
when: ParallelWhen::Always,
second_norm: Some("attn_norm_2"),
lines: "src/models/falcon.cpp:35-36,79-85,124-135",
},
ParallelResidual {
arch: "command-r",
norm: ParallelNorm::SharedNorm,
when: ParallelWhen::Always,
second_norm: None,
lines: "src/models/command-r.cpp:68,106-119",
},
ParallelResidual {
arch: "cohere2",
norm: ParallelNorm::SharedNorm,
when: ParallelWhen::Always,
second_norm: None,
lines: "src/models/cohere2.cpp:120-134",
},
ParallelResidual {
arch: "cohere2moe",
norm: ParallelNorm::SharedNorm,
when: ParallelWhen::Always,
second_norm: None,
lines: "src/models/cohere2moe.cpp:222-266",
},
ParallelResidual {
arch: "plamo",
norm: ParallelNorm::SharedNorm,
when: ParallelWhen::Always,
second_norm: None,
lines: "src/models/plamo.cpp:59-64,97-98,111-112",
},
];
pub fn parallel_residual(arch: &str) -> Option<&'static ParallelResidual> {
PARALLEL_RESIDUAL_GRAPHS.iter().find(|row| row.arch == arch)
}
pub fn layer_is_parallel(file: &impl TensorSource, arch: &str, l: usize) -> bool {
let Some(row) = parallel_residual(arch) else {
return false;
};
match row.when {
ParallelWhen::Always => true,
ParallelWhen::FfnNormAbsent => file
.find_tensor(&format!("blk.{l}.ffn_norm.weight"))
.is_none(),
ParallelWhen::ParallelResidualKey => file
.metadata_bool(&format!("{arch}.use_parallel_residual"))
.unwrap_or(false),
}
}
pub fn layer_parallel_norm(file: &impl TensorSource, arch: &str, l: usize) -> Option<ParallelNorm> {
let row = parallel_residual(arch)?;
if !layer_is_parallel(file, arch, l) {
return None;
}
let has_second = row.second_norm.is_some_and(|name| {
file.find_tensor(&format!("blk.{l}.{name}.weight"))
.is_some()
});
Some(if has_second {
ParallelNorm::TwoNorms
} else {
row.norm
})
}
pub fn model_has_parallel_layer(file: &impl TensorSource, arch: &str, n_layers: usize) -> bool {
(0..n_layers).any(|l| layer_is_parallel(file, arch, l))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn every_row_is_a_registered_architecture_and_the_generic_ones_are_audited() {
for row in PARALLEL_RESIDUAL_GRAPHS {
assert!(
crate::capability::resolve_profile(row.arch).is_some(),
"`{}` ({}) is not a registered architecture",
row.arch,
row.lines
);
let generic = matches!(
crate::capability::resolve_architecture(row.arch),
Some(crate::capability::ArchPath::GenericGqa { .. })
);
assert_eq!(
generic,
matches!(
row.arch,
"stablelm"
| "gptneox"
| "plamo"
| "command-r"
| "falcon"
| "phi2"
| "cohere2"
| "cohere2moe"
),
"`{}`: a generic-path row here must have a golden in \
tests/parallel_residual_graphs.rs or its own graph test",
row.arch
);
}
}
#[test]
fn no_architecture_has_two_rows() {
let mut names: Vec<&str> = PARALLEL_RESIDUAL_GRAPHS.iter().map(|r| r.arch).collect();
names.sort_unstable();
names.dedup();
assert_eq!(names.len(), PARALLEL_RESIDUAL_GRAPHS.len());
assert_eq!(PARALLEL_RESIDUAL_GRAPHS.len(), 8, "the measured reach");
}
#[test]
fn a_sequential_graph_is_never_parallel() {
let file = crate::test_source::StubSource::with_tensors(&[]);
assert!(!layer_is_parallel(&file, "llama", 0));
assert_eq!(layer_parallel_norm(&file, "llama", 0), None);
assert!(!model_has_parallel_layer(&file, "llama", 4));
}
#[test]
fn stablelm_is_decided_by_ffn_norm_presence_per_layer() {
use crate::test_source::StubSource;
let sequential =
StubSource::with_tensors(&["blk.0.ffn_norm.weight", "blk.1.ffn_norm.weight"]);
assert_eq!(layer_parallel_norm(&sequential, "stablelm", 0), None);
assert_eq!(layer_parallel_norm(&sequential, "stablelm", 1), None);
assert!(!model_has_parallel_layer(&sequential, "stablelm", 2));
let mixed = StubSource::with_tensors(&["blk.0.ffn_norm.weight"]);
assert_eq!(layer_parallel_norm(&mixed, "stablelm", 0), None);
assert_eq!(
layer_parallel_norm(&mixed, "stablelm", 1),
Some(ParallelNorm::SharedNorm)
);
assert!(model_has_parallel_layer(&mixed, "stablelm", 2));
assert!(!model_has_parallel_layer(&mixed, "stablelm", 1));
}
#[test]
fn the_key_the_second_norm_and_the_unconditional_rules() {
use crate::test_source::StubSource;
use frink_gguf::GgufValue;
let neox_seq = StubSource::with_tensors(&["blk.0.ffn_norm.weight"])
.with_key("gptneox.use_parallel_residual", GgufValue::Bool(false));
assert!(!layer_is_parallel(&neox_seq, "gptneox", 0));
let neox_par = StubSource::with_tensors(&["blk.0.ffn_norm.weight"])
.with_key("gptneox.use_parallel_residual", GgufValue::Bool(true));
assert_eq!(
layer_parallel_norm(&neox_par, "gptneox", 0),
Some(ParallelNorm::TwoNorms)
);
let falcon_7b = StubSource::with_tensors(&["blk.0.attn_norm.weight"]);
assert_eq!(
layer_parallel_norm(&falcon_7b, "falcon", 0),
Some(ParallelNorm::SharedNorm)
);
let falcon_40b = StubSource::with_tensors(&["blk.0.attn_norm_2.weight"]);
let phi2 = StubSource::with_tensors(&["blk.0.ffn_norm.weight"]);
assert_eq!(
layer_parallel_norm(&phi2, "phi2", 0),
Some(ParallelNorm::SharedNorm)
);
assert_eq!(
layer_parallel_norm(&falcon_40b, "falcon", 0),
Some(ParallelNorm::TwoNorms)
);
}
}