use std::sync::Arc;
use frink_gguf::{GgufValue, TensorSource};
use frink_moe::{ClampForm, GluAct, XieluParams};
use crate::config::{FfnActivation, ModelConfig};
use crate::loader::LoadError;
#[derive(Debug, Clone, PartialEq)]
pub struct XieluLayers(Arc<[XieluParams]>);
impl XieluLayers {
pub fn new(layers: Vec<XieluParams>) -> Self {
Self(layers.into())
}
pub fn layer(&self, il: usize) -> XieluParams {
self.0[il]
}
pub fn len(&self) -> usize {
self.0.len()
}
pub fn is_empty(&self) -> bool {
self.0.is_empty()
}
}
pub const XIELU_KEYS: [&str; 4] = ["xielu.alpha_n", "xielu.alpha_p", "xielu.beta", "xielu.eps"];
fn read_f32_per_layer(
file: &impl TensorSource,
key: &str,
n_layers: usize,
) -> Result<Vec<f32>, LoadError> {
let Some(value) = file.metadata(key) else {
return Err(LoadError::MissingHparam(key.to_string()));
};
match value {
GgufValue::Array(items) => {
if items.len() != n_layers {
return Err(LoadError::UnsupportedFeature(
key.to_string(),
format!(
"array of {} entries for {n_layers} layers; llama.cpp refuses this too \
(`key has wrong array length`, llama-model-loader.cpp:464-465)",
items.len()
),
));
}
let mut out = Vec::with_capacity(n_layers);
for (il, item) in items.iter().enumerate() {
out.push(item.as_f32().ok_or_else(|| {
LoadError::UnsupportedFeature(
key.to_string(),
format!("entry {il} is not a float: {item:?}"),
)
})?);
}
Ok(out)
}
scalar => scalar
.as_f32()
.map(|v| vec![v; n_layers])
.ok_or_else(|| LoadError::MissingHparam(key.to_string())),
}
}
pub fn read_xielu_layers(
file: &impl TensorSource,
n_layers: usize,
) -> Result<XieluLayers, LoadError> {
let [alpha_n, alpha_p, beta, eps] = XIELU_KEYS;
let alpha_n = read_f32_per_layer(file, alpha_n, n_layers)?;
let alpha_p = read_f32_per_layer(file, alpha_p, n_layers)?;
let beta = read_f32_per_layer(file, beta, n_layers)?;
let eps = read_f32_per_layer(file, eps, n_layers)?;
Ok(XieluLayers::new(
(0..n_layers)
.map(|il| XieluParams::from_gguf(alpha_n[il], alpha_p[il], beta[il], eps[il]))
.collect(),
))
}
#[derive(Debug, Clone, PartialEq)]
pub struct SwigluClamps {
routed: Arc<[f32]>,
dense: Arc<[f32]>,
form: ClampForm,
}
pub const CLAMP_BEFORE_SILU: &[(&str, &str)] = &[
("maple", "src/llama-graph.cpp:2228 (routed only)"),
("deepseek4", "src/llama-graph.cpp:1834,2228 (own engine)"),
("hy_v4", "src/llama-graph.cpp:2228 (refused: dedicated)"),
(
"dflash",
"src/llama-graph.cpp:1834,2228 when dsv4_hc_mult > 0 (deferred)",
),
];
pub fn clamp_form(arch: &str) -> ClampForm {
if CLAMP_BEFORE_SILU.iter().any(|(a, _)| *a == arch) {
ClampForm::BeforeSilu
} else {
ClampForm::AfterSilu
}
}
const CLAMP_EPS: f32 = 1e-6;
impl SwigluClamps {
pub fn new(routed: Vec<f32>, dense: Vec<f32>, form: ClampForm) -> Self {
assert_eq!(routed.len(), dense.len(), "one entry per layer in both");
Self {
routed: routed.into(),
dense: dense.into(),
form,
}
}
pub fn routed(&self, il: usize) -> GluAct {
self.act(self.routed[il])
}
pub fn dense(&self, il: usize) -> GluAct {
self.act(self.dense[il])
}
fn act(&self, limit: f32) -> GluAct {
if limit > CLAMP_EPS {
GluAct::SwigluClamped {
limit,
form: self.form,
}
} else {
GluAct::Swiglu
}
}
pub fn len(&self) -> usize {
self.routed.len()
}
pub fn is_empty(&self) -> bool {
self.routed.is_empty()
}
}
pub const SWIGLU_CLAMP_READERS: &[&str] = &["step35", "maple"];
pub fn reads_swiglu_clamps(arch: &str) -> bool {
SWIGLU_CLAMP_READERS.contains(&arch)
}
pub fn read_swiglu_clamps(
file: &impl TensorSource,
arch: &str,
trunk: &crate::mtp_blocks::TrunkLayers,
) -> Result<Option<SwigluClamps>, LoadError> {
let key = |k: &str| format!("{arch}.{k}");
let (exp_key, shexp_key) = (key("swiglu_clamp_exp"), key("swiglu_clamp_shexp"));
if file.metadata(&exp_key).is_none() && file.metadata(&shexp_key).is_none() {
return Ok(None);
}
let read = |k: &str| -> Result<Vec<f32>, LoadError> {
if file.metadata(k).is_none() {
return Ok(vec![0.0; trunk.n_layers]);
}
let mut v = read_f32_per_layer(file, k, trunk.block_count)?;
v.truncate(trunk.n_layers);
Ok(v)
};
Ok(Some(SwigluClamps::new(
read(&exp_key)?,
read(&shexp_key)?,
clamp_form(arch),
)))
}
pub const XIELU_ARCHITECTURES: &[&str] = &["apertus"];
pub fn uses_xielu(arch: &str) -> bool {
XIELU_ARCHITECTURES.contains(&arch)
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct LayerFfnActs {
pub routed: GluAct,
pub dense: GluAct,
}
impl LayerFfnActs {
fn same(act: GluAct) -> Self {
Self {
routed: act,
dense: act,
}
}
pub fn all_swiglu(self) -> bool {
self.routed.is_swiglu() && self.dense.is_swiglu()
}
}
impl ModelConfig {
pub fn layer_ffn_acts(&self, il: usize) -> LayerFfnActs {
match &self.ffn_activation {
FfnActivation::Swiglu | FfnActivation::SwigluFused => {
LayerFfnActs::same(GluAct::Swiglu)
}
FfnActivation::Gelu => LayerFfnActs::same(GluAct::Geglu),
FfnActivation::ReluSqr => LayerFfnActs::same(GluAct::ReluSqr),
FfnActivation::GeluUngated => LayerFfnActs::same(GluAct::GeluUngated),
FfnActivation::Reglu => LayerFfnActs::same(GluAct::Reglu),
FfnActivation::Xielu(layers) => LayerFfnActs::same(GluAct::Xielu(layers.layer(il))),
FfnActivation::SwigluClamped(clamps) => LayerFfnActs {
routed: clamps.routed(il),
dense: clamps.dense(il),
},
}
}
pub fn model_ffn_act(&self) -> Option<GluAct> {
match &self.ffn_activation {
FfnActivation::Xielu(_) | FfnActivation::SwigluClamped(_) => None,
FfnActivation::Swiglu
| FfnActivation::SwigluFused
| FfnActivation::Gelu
| FfnActivation::ReluSqr
| FfnActivation::GeluUngated
| FfnActivation::Reglu => Some(self.layer_ffn_acts(0).dense),
}
}
pub fn ffn_is_ungated(&self) -> bool {
match &self.ffn_activation {
FfnActivation::ReluSqr | FfnActivation::GeluUngated | FfnActivation::Xielu(_) => true,
FfnActivation::Swiglu
| FfnActivation::SwigluFused
| FfnActivation::SwigluClamped(_)
| FfnActivation::Gelu
| FfnActivation::Reglu => false,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use frink_moe::GluAct;
#[derive(Clone, Default)]
struct Meta(Vec<(String, GgufValue)>);
impl Meta {
fn insert(&mut self, key: &str, value: GgufValue) {
self.remove(key);
self.0.push((key.to_string(), value));
}
fn remove(&mut self, key: &str) {
self.0.retain(|(k, _)| k != key);
}
}
impl TensorSource for Meta {
fn metadata(&self, key: &str) -> Option<&GgufValue> {
self.0.iter().find(|(k, _)| k == key).map(|(_, v)| v)
}
fn find_tensor(&self, _: &str) -> Option<&frink_gguf::TensorInfo> {
None
}
fn tensor_bytes(&self, name: &str) -> Result<&[u8], frink_gguf::GgufError> {
Err(frink_gguf::GgufError::TensorNotFound(name.to_string()))
}
fn tensor_mapped_range(
&self,
name: &str,
) -> Result<(Arc<frink_gguf::MmapHandle>, std::ops::Range<usize>), frink_gguf::GgufError>
{
Err(frink_gguf::GgufError::TensorNotFound(name.to_string()))
}
}
fn base_config() -> ModelConfig {
let mut cfg = crate::config::glm_5_2();
cfg.n_layers = 2;
cfg
}
fn xielu_config(layers: Vec<XieluParams>) -> ModelConfig {
let mut cfg = base_config();
cfg.n_layers = layers.len();
cfg.ffn_activation = FfnActivation::Xielu(XieluLayers::new(layers));
cfg
}
#[test]
fn layer_ffn_act_answers_each_layer_s_own_parameters() {
let p0 = XieluParams::from_gguf(0.8, 0.8, 0.5, -1e-6);
let p1 = XieluParams::from_gguf(0.2, 1.5, 0.75, -0.3);
let cfg = xielu_config(vec![p0, p1]);
assert_eq!(cfg.layer_ffn_acts(0), LayerFfnActs::same(GluAct::Xielu(p0)));
assert_eq!(cfg.layer_ffn_acts(1), LayerFfnActs::same(GluAct::Xielu(p1)));
assert_ne!(p0, p1, "the test needs two different parameter sets");
assert!(cfg.ffn_is_ungated());
assert_eq!(
cfg.model_ffn_act(),
None,
"a parameterised activation has no whole-model answer, even with equal parameters"
);
assert_eq!(xielu_config(vec![p0, p0]).model_ffn_act(), None);
}
#[test]
fn uniform_activations_answer_the_same_on_every_layer() {
for (kind, want, ungated) in [
(FfnActivation::Swiglu, GluAct::Swiglu, false),
(FfnActivation::SwigluFused, GluAct::Swiglu, false),
(FfnActivation::Gelu, GluAct::Geglu, false),
(FfnActivation::ReluSqr, GluAct::ReluSqr, true),
(FfnActivation::GeluUngated, GluAct::GeluUngated, true),
(FfnActivation::Reglu, GluAct::Reglu, false),
] {
let mut cfg = base_config();
cfg.ffn_activation = kind.clone();
for il in 0..cfg.n_layers {
assert_eq!(
cfg.layer_ffn_acts(il),
LayerFfnActs::same(want),
"{kind:?} layer {il}"
);
assert!(cfg.layer_ffn_acts(il).all_swiglu() == (want == GluAct::Swiglu));
}
assert_eq!(cfg.model_ffn_act(), Some(want), "{kind:?}");
assert_eq!(cfg.ffn_is_ungated(), ungated, "{kind:?}");
}
}
#[test]
fn the_four_keys_are_read_as_llama_cpp_reads_them() {
let arr = |v: &[f32]| GgufValue::Array(v.iter().map(|&x| GgufValue::F32(x)).collect());
let mut md = Meta::default();
md.insert("xielu.alpha_n", arr(&[0.8, 0.2]));
md.insert("xielu.alpha_p", arr(&[0.8, 1.5]));
md.insert("xielu.beta", GgufValue::F32(0.5)); md.insert("xielu.eps", arr(&[-1e-6, -0.3]));
let layers = read_xielu_layers(&md, 2).expect("reads");
assert_eq!(layers.len(), 2);
assert_eq!(
layers.layer(0),
XieluParams::from_gguf(0.8, 0.8, 0.5, -1e-6)
);
assert_eq!(layers.layer(1), XieluParams::from_gguf(0.2, 1.5, 0.5, -0.3));
let mut short = md.clone();
short.insert("xielu.eps", arr(&[-1e-6]));
let err = read_xielu_layers(&short, 2).expect_err("wrong length refuses");
let msg = err.to_string();
assert!(
msg.contains("xielu.eps") && msg.contains("1 entries"),
"{msg}"
);
let mut missing = md.clone();
missing.remove("xielu.alpha_p");
let err = read_xielu_layers(&missing, 2).expect_err("a missing key refuses");
assert!(err.to_string().contains("xielu.alpha_p"), "{err}");
}
#[test]
fn the_clamp_form_is_per_architecture_and_the_two_forms_differ() {
assert_eq!(clamp_form("maple"), ClampForm::BeforeSilu);
assert_eq!(clamp_form("step35"), ClampForm::AfterSilu);
assert_eq!(clamp_form("llama"), ClampForm::AfterSilu);
for (arch, line) in CLAMP_BEFORE_SILU {
assert!(line.contains("llama-graph.cpp:"), "`{arch}` cites no line");
assert_eq!(clamp_form(arch), ClampForm::BeforeSilu);
}
let before = GluAct::SwigluClamped {
limit: 2.0,
form: ClampForm::BeforeSilu,
};
let after = GluAct::SwigluClamped {
limit: 2.0,
form: ClampForm::AfterSilu,
};
assert!((before.combine(6.0, 1.0) - 1.761_594).abs() < 1e-5);
assert!((after.combine(6.0, 1.0) - 2.0).abs() < 1e-5);
assert!((before.combine(0.5, 1.0) - after.combine(0.5, 1.0)).abs() < 1e-7);
}
#[test]
fn the_clamp_arrays_are_read_per_site_and_zero_means_plain_swiglu() {
let mut cfg = base_config();
cfg.n_layers = 3;
cfg.ffn_activation = FfnActivation::SwigluClamped(SwigluClamps::new(
vec![0.0, 1.5, 0.0],
vec![2.0, 0.0, 1e-7],
ClampForm::AfterSilu,
));
assert_eq!(
cfg.layer_ffn_acts(0),
LayerFfnActs {
routed: GluAct::Swiglu,
dense: GluAct::SwigluClamped {
limit: 2.0,
form: ClampForm::AfterSilu
},
}
);
assert_eq!(
cfg.layer_ffn_acts(1),
LayerFfnActs {
routed: GluAct::SwigluClamped {
limit: 1.5,
form: ClampForm::AfterSilu
},
dense: GluAct::Swiglu,
}
);
assert_eq!(cfg.layer_ffn_acts(2), LayerFfnActs::same(GluAct::Swiglu));
assert!(cfg.layer_ffn_acts(2).all_swiglu() && !cfg.layer_ffn_acts(0).all_swiglu());
assert_eq!(cfg.model_ffn_act(), None);
assert!(!cfg.ffn_is_ungated());
}
#[test]
fn the_clamp_keys_are_read_as_llama_cpp_reads_them() {
let arr = |v: &[f32]| GgufValue::Array(v.iter().map(|&x| GgufValue::F32(x)).collect());
let trunk = crate::mtp_blocks::TrunkLayers {
block_count: 3,
n_layers: 2,
n_mtp_blocks: 1,
};
let mut md = Meta::default();
assert_eq!(
read_swiglu_clamps(&md, "step35", &trunk).expect("reads"),
None
);
md.insert("step35.swiglu_clamp_exp", arr(&[0.0, 7.0, 0.0]));
let clamps = read_swiglu_clamps(&md, "step35", &trunk)
.expect("reads")
.expect("one key is a table");
assert_eq!(clamps.len(), 2, "trunk entries only");
assert_eq!(
clamps.routed(1),
GluAct::SwigluClamped {
limit: 7.0,
form: ClampForm::AfterSilu
}
);
assert_eq!(clamps.dense(1), GluAct::Swiglu, "the absent key is zeros");
md.insert("step35.swiglu_clamp_shexp", GgufValue::F32(16.0));
let clamps = read_swiglu_clamps(&md, "step35", &trunk)
.expect("reads")
.expect("table");
assert_eq!(
clamps.dense(0),
GluAct::SwigluClamped {
limit: 16.0,
form: ClampForm::AfterSilu
}
);
assert_eq!(
clamps.dense(1),
GluAct::SwigluClamped {
limit: 16.0,
form: ClampForm::AfterSilu
}
);
md.insert("step35.swiglu_clamp_exp", arr(&[0.0, 7.0]));
let err = read_swiglu_clamps(&md, "step35", &trunk).expect_err("wrong length refuses");
assert!(err.to_string().contains("swiglu_clamp_exp"), "{err}");
assert!(reads_swiglu_clamps("step35"));
for arch in ["llama", "apertus", "laguna", "deepseek2", "gpt-oss"] {
assert!(!reads_swiglu_clamps(arch), "{arch}");
}
}
#[test]
fn only_the_graphs_that_call_ggml_xielu_use_it() {
assert!(uses_xielu("apertus"));
for arch in ["llama", "arcee", "plm", "step35", "gemma3"] {
assert!(!uses_xielu(arch), "{arch}");
}
assert_eq!(XIELU_KEYS[0], "xielu.alpha_n", "no architecture prefix");
}
}