#[cfg(any(feature = "mmap", not(target_arch = "wasm32")))]
use std::path::Path;
use std::sync::Arc;
use anyhow::{Context, Result, bail, ensure};
use crate::gguf::GgufFile;
pub const LORA_TARGET_COUNT: usize = 13;
#[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Debug)]
pub enum LoraTarget {
AttnQ,
AttnK,
AttnV,
AttnOutput,
FfnGate,
FfnUp,
FfnDown,
ShortconvInProj,
ShortconvOutProj,
FfnGateInp,
FfnGateExps,
FfnUpExps,
FfnDownExps,
}
impl LoraTarget {
pub const ALL: [LoraTarget; LORA_TARGET_COUNT] = [
LoraTarget::AttnQ,
LoraTarget::AttnK,
LoraTarget::AttnV,
LoraTarget::AttnOutput,
LoraTarget::FfnGate,
LoraTarget::FfnUp,
LoraTarget::FfnDown,
LoraTarget::ShortconvInProj,
LoraTarget::ShortconvOutProj,
LoraTarget::FfnGateInp,
LoraTarget::FfnGateExps,
LoraTarget::FfnUpExps,
LoraTarget::FfnDownExps,
];
pub fn index(self) -> usize {
match self {
LoraTarget::AttnQ => 0,
LoraTarget::AttnK => 1,
LoraTarget::AttnV => 2,
LoraTarget::AttnOutput => 3,
LoraTarget::FfnGate => 4,
LoraTarget::FfnUp => 5,
LoraTarget::FfnDown => 6,
LoraTarget::ShortconvInProj => 7,
LoraTarget::ShortconvOutProj => 8,
LoraTarget::FfnGateInp => 9,
LoraTarget::FfnGateExps => 10,
LoraTarget::FfnUpExps => 11,
LoraTarget::FfnDownExps => 12,
}
}
pub fn is_expert(self) -> bool {
matches!(
self,
LoraTarget::FfnGateExps | LoraTarget::FfnUpExps | LoraTarget::FfnDownExps
)
}
fn is_routed_ffn(self) -> bool {
self.is_expert() || self == LoraTarget::FfnGateInp
}
fn gguf_stem(self) -> &'static str {
match self {
LoraTarget::AttnQ => "attn_q",
LoraTarget::AttnK => "attn_k",
LoraTarget::AttnV => "attn_v",
LoraTarget::AttnOutput => "attn_output",
LoraTarget::FfnGate => "ffn_gate",
LoraTarget::FfnUp => "ffn_up",
LoraTarget::FfnDown => "ffn_down",
LoraTarget::ShortconvInProj => "shortconv.in_proj",
LoraTarget::ShortconvOutProj => "shortconv.out_proj",
LoraTarget::FfnGateInp => "ffn_gate_inp",
LoraTarget::FfnGateExps => "ffn_gate_exps",
LoraTarget::FfnUpExps => "ffn_up_exps",
LoraTarget::FfnDownExps => "ffn_down_exps",
}
}
fn from_gguf_stem(stem: &str) -> Option<LoraTarget> {
LoraTarget::ALL.into_iter().find(|t| t.gguf_stem() == stem)
}
fn from_peft_module(module: &str) -> Option<LoraTarget> {
match module {
"self_attn.q_proj" => Some(LoraTarget::AttnQ),
"self_attn.k_proj" => Some(LoraTarget::AttnK),
"self_attn.v_proj" => Some(LoraTarget::AttnV),
"self_attn.o_proj" => Some(LoraTarget::AttnOutput),
"mlp.gate_proj" => Some(LoraTarget::FfnGate),
"mlp.up_proj" => Some(LoraTarget::FfnUp),
"mlp.down_proj" => Some(LoraTarget::FfnDown),
"conv.in_proj" => Some(LoraTarget::ShortconvInProj),
"conv.out_proj" => Some(LoraTarget::ShortconvOutProj),
_ => None,
}
}
}
#[derive(Clone)]
pub struct LoraTargetWeights {
pub a: Vec<f32>,
pub b: Vec<f32>,
pub rank: usize,
pub k: usize,
pub d: usize,
pub scale: f32,
}
impl LoraTargetWeights {
fn new(
a: Vec<f32>,
rank_a: usize,
k: usize,
b: Vec<f32>,
d: usize,
rank_b: usize,
alpha: f32,
) -> Result<Self> {
ensure!(
rank_a == rank_b,
"LoRA rank mismatch between A ({rank_a}) and B ({rank_b})"
);
ensure!(rank_a > 0 && k > 0 && d > 0, "LoRA dims must be non-zero");
ensure!(
rank_a <= MAX_LORA_RANK,
"LoRA rank {rank_a} exceeds the supported maximum ({MAX_LORA_RANK})"
);
let ak = rank_a.checked_mul(k).context("LoRA A dims overflow")?;
let dr = d.checked_mul(rank_a).context("LoRA B dims overflow")?;
ensure!(a.len() == ak, "LoRA A size {} != rank*k {ak}", a.len());
ensure!(b.len() == dr, "LoRA B size {} != d*rank {dr}", b.len());
Ok(Self {
a,
b,
rank: rank_a,
k,
d,
scale: alpha / rank_a as f32,
})
}
}
#[derive(Clone)]
enum TargetDelta {
Dense(LoraTargetWeights),
Experts(Vec<LoraTargetWeights>),
}
#[derive(Default, Clone)]
pub struct LoraLayer {
targets: [Option<TargetDelta>; LORA_TARGET_COUNT],
}
pub struct LoraAdapterWeights {
layers: Vec<LoraLayer>,
default_scale: f32,
pub classifier_weight: Option<Vec<f32>>,
pub classifier_bias: Option<Vec<f32>>,
pub class_labels: Vec<String>,
pub num_classes: usize,
}
impl LoraAdapterWeights {
pub fn is_classifier(&self) -> bool {
self.classifier_weight.is_some()
}
pub fn num_classes(&self) -> usize {
self.num_classes
}
pub fn with_class_labels(mut self: Arc<Self>, labels: Vec<String>) -> Arc<Self> {
let n_cls = labels.len();
if let Some(s) = Arc::get_mut(&mut self) {
s.num_classes = n_cls;
s.class_labels = labels;
self
} else {
Arc::new(Self {
layers: self.layers.clone(),
default_scale: self.default_scale,
classifier_weight: self.classifier_weight.clone(),
classifier_bias: self.classifier_bias.clone(),
class_labels: labels,
num_classes: n_cls,
})
}
}
#[cfg(test)]
pub fn new_classifier_for_testing(
classifier_weight: Vec<f32>,
classifier_bias: Option<Vec<f32>>,
class_labels: Vec<String>,
) -> Arc<Self> {
let num_classes = if !class_labels.is_empty() {
class_labels.len()
} else if let Some(ref b) = classifier_bias {
b.len()
} else {
0
};
Arc::new(Self {
layers: Vec::new(),
default_scale: 1.0,
classifier_weight: Some(classifier_weight),
classifier_bias,
class_labels,
num_classes,
})
}
pub fn get(&self, layer: usize, target: LoraTarget) -> Option<&LoraTargetWeights> {
match self.layers.get(layer)?.targets[target.index()].as_ref()? {
TargetDelta::Dense(t) => Some(t),
TargetDelta::Experts(_) => None,
}
}
pub fn get_expert(
&self,
layer: usize,
target: LoraTarget,
expert: usize,
) -> Option<&LoraTargetWeights> {
match self.layers.get(layer)?.targets[target.index()].as_ref()? {
TargetDelta::Experts(per_expert) => per_expert.get(expert),
TargetDelta::Dense(_) => None,
}
}
pub fn has_moe_deltas(&self) -> bool {
self.layers.iter().any(|l| {
LoraTarget::ALL
.into_iter()
.filter(|t| t.is_routed_ffn())
.any(|t| l.targets[t.index()].is_some())
})
}
pub fn n_layers(&self) -> usize {
self.layers.len()
}
pub fn default_scale(&self) -> f32 {
self.default_scale
}
pub fn target_count(&self) -> usize {
self.layers
.iter()
.map(|l| l.targets.iter().filter(|t| t.is_some()).count())
.sum()
}
pub fn validate_dims(&self, config: &crate::model::ModelConfig) -> Result<()> {
let n_layers = config
.block_types
.len()
.max(config.kv_heads_per_layer.len());
let q_dim = config.n_heads * config.head_dim;
for (layer, l) in self.layers.iter().enumerate() {
for target in LoraTarget::ALL {
let Some(delta) = l.targets[target.index()].as_ref() else {
continue;
};
ensure!(
layer < n_layers,
"LoRA references layer {layer} but the model has {n_layers} layers"
);
let moe = config
.moe
.as_ref()
.filter(|m| m.is_moe_layer.get(layer).copied().unwrap_or(false));
match (target, moe.is_some()) {
(LoraTarget::FfnGate | LoraTarget::FfnUp | LoraTarget::FfnDown, true) => bail!(
"LoRA target {target:?} on layer {layer} adapts a dense feed-forward \
block, but that layer is mixture-of-experts; it needs the stacked \
per-expert tensors (`ffn_*_exps.weight.lora_{{a,b}}`)"
),
(t, false) if t.is_routed_ffn() => bail!(
"LoRA target {target:?} on layer {layer} adapts a mixture-of-experts \
feed-forward block, but that layer is dense"
),
_ => {}
}
let kv_heads = config
.kv_heads_per_layer
.get(layer)
.copied()
.filter(|&h| h > 0)
.unwrap_or(config.n_kv_heads);
let kv_dim = kv_heads * config.head_dim;
let (want_k, want_d) = match target {
LoraTarget::AttnQ => (config.hidden_size, q_dim),
LoraTarget::AttnK | LoraTarget::AttnV => (config.hidden_size, kv_dim),
LoraTarget::AttnOutput => (q_dim, config.hidden_size),
LoraTarget::FfnGate | LoraTarget::FfnUp => {
(config.hidden_size, config.intermediate_size)
}
LoraTarget::FfnDown => (config.intermediate_size, config.hidden_size),
LoraTarget::ShortconvInProj => (config.hidden_size, 3 * config.hidden_size),
LoraTarget::ShortconvOutProj => (config.hidden_size, config.hidden_size),
LoraTarget::FfnGateInp => (config.hidden_size, moe.map_or(0, |m| m.n_expert)),
LoraTarget::FfnGateExps | LoraTarget::FfnUpExps => {
(config.hidden_size, moe.map_or(0, |m| m.expert_ff_len))
}
LoraTarget::FfnDownExps => {
(moe.map_or(0, |m| m.expert_ff_len), config.hidden_size)
}
};
let pairs = match delta {
TargetDelta::Dense(t) => std::slice::from_ref(t),
TargetDelta::Experts(per_expert) => {
let want_experts = moe.map_or(0, |m| m.n_expert);
ensure!(
per_expert.len() == want_experts,
"LoRA target {target:?} on layer {layer} carries {} expert deltas, \
but the model has {want_experts} experts",
per_expert.len()
);
per_expert.as_slice()
}
};
for t in pairs {
ensure!(
t.k == want_k && t.d == want_d,
"LoRA target {target:?} on layer {layer} has dims (in={}, out={}), \
but the model expects (in={want_k}, out={want_d}). Adapter built for a \
different model?",
t.k,
t.d
);
}
}
}
if let Some(ref w) = self.classifier_weight {
let hs = config.hidden_size;
ensure!(
hs > 0 && w.len().is_multiple_of(hs),
"classifier weight length {} is not a multiple of hidden_size {}",
w.len(),
hs
);
let num_classes = w.len() / hs;
if let Some(ref b) = self.classifier_bias {
ensure!(
b.len() == num_classes,
"classifier bias length {} does not match num_classes {}",
b.len(),
num_classes
);
}
if !self.class_labels.is_empty() {
ensure!(
self.class_labels.len() == num_classes,
"class_labels count {} does not match num_classes {}",
self.class_labels.len(),
num_classes
);
}
}
Ok(())
}
#[cfg(feature = "mmap")]
pub fn from_gguf(path: &Path) -> Result<Arc<Self>> {
let gguf = GgufFile::open(path).with_context(|| format!("open adapter {path:?}"))?;
Self::from_gguf_file(&gguf)
}
pub fn from_gguf_bytes(bytes: Arc<[u8]>) -> Result<Arc<Self>> {
let gguf = GgufFile::from_bytes(bytes).context("parse adapter GGUF bytes")?;
Self::from_gguf_file(&gguf)
}
fn from_gguf_file(gguf: &GgufFile) -> Result<Arc<Self>> {
let alpha_meta = gguf.get_f32("adapter.lora.alpha");
let mut builder = AdapterBuilder::new();
for (name, info) in &gguf.tensors {
let Some((layer, target, is_a)) = parse_gguf_lora_name(name) else {
continue;
};
let (rows, cols, n_slices) = match info.shape[..] {
[cols, rows] => (rows, cols, 1),
[cols, rows, n_slices] => (rows, cols, n_slices),
_ => bail!(
"LoRA tensor {name} has rank {}, expected 2 (dense) or 3 (per-expert)",
info.shape.len()
),
};
let data = gguf.get_tensor(name)?.to_f32_vec();
let factor = Factor::new(data, rows, cols, n_slices)
.with_context(|| format!("LoRA tensor {name}"))?;
builder.add_factor(layer, target, is_a, factor);
}
let (classifier_weight, classifier_bias, num_classes) =
if let Ok(tensor) = gguf.get_tensor("classifier.weight") {
let w = tensor.to_f32_vec();
let b = gguf
.get_tensor("classifier.bias")
.ok()
.map(|t| t.to_f32_vec());
let n_cls = b
.as_ref()
.map(|v| v.len())
.unwrap_or_else(|| tensor.shape().get(1).copied().unwrap_or(0));
(Some(w), b, n_cls)
} else {
(None, None, 0)
};
let class_labels = gguf
.get_string_array("token_classifier.labels")
.map(|arr| arr.into_iter().map(|s| s.to_string()).collect())
.unwrap_or_default();
builder.finish(
alpha_meta,
classifier_weight,
classifier_bias,
class_labels,
num_classes,
)
}
#[cfg(not(target_arch = "wasm32"))]
pub fn from_safetensors(path: &Path, alpha: Option<f32>) -> Result<Arc<Self>> {
let bytes = std::fs::read(path).with_context(|| format!("read adapter {path:?}"))?;
let alpha = alpha.or_else(|| {
path.parent().and_then(|dir| {
let p = dir.join("adapter_config.json");
let content = std::fs::read_to_string(p).ok()?;
let val: serde_json::Value = serde_json::from_str(&content).ok()?;
val.get("lora_alpha")
.and_then(|v| v.as_f64())
.map(|f| f as f32)
})
});
let mut adapter = Self::from_safetensors_bytes(&bytes, alpha)?;
if let Some(parent) = path
.parent()
.filter(|_| adapter.is_classifier() && adapter.class_labels.is_empty())
{
let labels = load_labels_from_dir(parent);
if !labels.is_empty() {
adapter = adapter.with_class_labels(labels);
}
}
Ok(adapter)
}
#[cfg(not(target_arch = "wasm32"))]
pub fn load_from_path(path: &Path) -> Result<Arc<Self>> {
let is_gguf = if let Ok(mut f) = std::fs::File::open(path) {
use std::io::Read;
let mut magic = [0u8; 4];
f.read_exact(&mut magic).is_ok() && &magic == b"GGUF"
} else {
false
};
if is_gguf {
Self::from_gguf(path)
} else {
Self::from_safetensors(path, None)
}
}
pub fn from_safetensors_bytes(bytes: &[u8], alpha: Option<f32>) -> Result<Arc<Self>> {
let st = SafeTensors::parse(bytes)?;
let mut builder = AdapterBuilder::new();
let mut classifier_weight = None;
let mut classifier_bias = None;
let mut num_classes = 0;
for (name, entry) in st.tensors() {
if name.ends_with("classifier.weight") {
if let Ok((rows, _cols)) = entry.shape2() {
num_classes = rows;
}
classifier_weight = Some(st.dequantize(entry, bytes)?);
continue;
}
if name.ends_with("classifier.bias") {
let b = st.dequantize(entry, bytes)?;
num_classes = b.len();
classifier_bias = Some(b);
continue;
}
if name.contains("experts") && (name.contains("lora_A") || name.contains("lora_B")) {
bail!(
"PEFT adapter tensor {name} targets a mixture-of-experts projection, which \
cera only loads from GGUF. Convert the adapter with llama.cpp's \
`convert_lora_to_gguf.py` and load the result."
);
}
let Some((layer, target, is_a)) = parse_peft_lora_name(name) else {
continue;
};
let (rows, cols) = entry
.shape2()
.with_context(|| format!("tensor {name} not 2-D"))?;
let data = st.dequantize(entry, bytes)?;
let factor =
Factor::new(data, rows, cols, 1).with_context(|| format!("LoRA tensor {name}"))?;
builder.add_factor(layer, target, is_a, factor);
}
builder.finish(
alpha,
classifier_weight,
classifier_bias,
Vec::new(),
num_classes,
)
}
}
#[cfg(not(target_arch = "wasm32"))]
fn load_labels_from_dir(dir: &Path) -> Vec<String> {
let schema_path = dir.join("label_schema.json");
if let Ok(content) = std::fs::read_to_string(&schema_path) {
let val_opt = serde_json::from_str::<serde_json::Value>(&content).ok();
if let Some(types) = val_opt
.as_ref()
.and_then(|v| v.get("types_in_order")?.as_array())
{
let mut labels = vec!["O".to_string()];
for item in types {
if let Some(t_str) = item.as_str() {
labels.push(format!("B-{t_str}"));
labels.push(format!("I-{t_str}"));
labels.push(format!("E-{t_str}"));
labels.push(format!("S-{t_str}"));
}
}
return labels;
}
}
for filename in &["config.json", "adapter_config.json"] {
let p = dir.join(filename);
if let Ok(content) = std::fs::read_to_string(&p) {
let val_opt = serde_json::from_str::<serde_json::Value>(&content).ok();
if let Some(id2label) = val_opt
.as_ref()
.and_then(|v| v.get("id2label")?.as_object())
{
let mut pairs = Vec::new();
for (k, v) in id2label {
if let (Ok(idx), Some(lbl)) = (k.parse::<usize>(), v.as_str()) {
pairs.push((idx, lbl.to_string()));
}
}
if !pairs.is_empty() {
pairs.sort_by_key(|&(idx, _)| idx);
return pairs.into_iter().map(|(_, l)| l).collect();
}
}
}
}
Vec::new()
}
const MAX_LORA_LAYERS: usize = 8192;
const MAX_LORA_EXPERTS: usize = 4096;
pub const MAX_LORA_RANK: usize = 512;
#[derive(Default)]
struct AdapterBuilder {
factors: std::collections::HashMap<(usize, LoraTarget), FactorPair>,
max_layer: usize,
}
struct Factor {
data: Vec<f32>,
rows: usize,
cols: usize,
n_slices: usize,
}
impl Factor {
fn new(data: Vec<f32>, rows: usize, cols: usize, n_slices: usize) -> Result<Self> {
let want = rows
.checked_mul(cols)
.and_then(|m| m.checked_mul(n_slices))
.context("LoRA factor dims overflow")?;
ensure!(n_slices > 0, "LoRA factor has no slices");
ensure!(
n_slices <= MAX_LORA_EXPERTS,
"LoRA factor is stacked {n_slices} deep, over the sane maximum of {MAX_LORA_EXPERTS} experts"
);
ensure!(
data.len() == want,
"LoRA factor has {} elements, expected {n_slices}×{rows}×{cols} = {want}",
data.len()
);
Ok(Self {
data,
rows,
cols,
n_slices,
})
}
fn slice(&self, i: usize) -> Vec<f32> {
let n = self.rows * self.cols;
self.data[i * n..(i + 1) * n].to_vec()
}
}
#[derive(Default)]
struct FactorPair {
a: Option<Factor>,
b: Option<Factor>,
}
impl AdapterBuilder {
fn new() -> Self {
Self::default()
}
fn add_factor(&mut self, layer: usize, target: LoraTarget, is_a: bool, factor: Factor) {
self.max_layer = self.max_layer.max(layer);
let slot = self.factors.entry((layer, target)).or_default();
if is_a {
slot.a = Some(factor);
} else {
slot.b = Some(factor);
}
}
fn finish(
self,
alpha: Option<f32>,
classifier_weight: Option<Vec<f32>>,
classifier_bias: Option<Vec<f32>>,
class_labels: Vec<String>,
num_classes: usize,
) -> Result<Arc<LoraAdapterWeights>> {
ensure!(
!self.factors.is_empty() || classifier_weight.is_some(),
"adapter contains no LoRA or classifier tensors"
);
let n_cls = if !class_labels.is_empty() {
class_labels.len()
} else if let Some(ref b) = classifier_bias {
b.len()
} else {
num_classes
};
if self.factors.is_empty() {
return Ok(Arc::new(LoraAdapterWeights {
layers: Vec::new(),
default_scale: 1.0,
classifier_weight,
classifier_bias,
class_labels,
num_classes: n_cls,
}));
}
ensure!(
self.max_layer < MAX_LORA_LAYERS,
"adapter layer index {} exceeds the sane maximum ({MAX_LORA_LAYERS})",
self.max_layer
);
let n_layers = self.max_layer + 1;
let mut layers: Vec<LoraLayer> = (0..n_layers).map(|_| LoraLayer::default()).collect();
let mut factors: Vec<_> = self.factors.into_iter().collect();
factors.sort_by_key(|&(key, _)| key);
let mut default_scale = 1.0f32;
let mut scale_set = false;
for ((layer, target), pair) in factors {
let a = pair
.a
.with_context(|| format!("layer {layer} target {target:?}: missing lora_a"))?;
let b = pair
.b
.with_context(|| format!("layer {layer} target {target:?}: missing lora_b"))?;
ensure!(
a.n_slices == b.n_slices,
"layer {layer} target {target:?}: lora_a has {} slices but lora_b has {}",
a.n_slices,
b.n_slices
);
let alpha = alpha.unwrap_or(a.rows as f32);
let delta = if target.is_expert() {
let per_expert = (0..a.n_slices)
.map(|e| {
LoraTargetWeights::new(
a.slice(e),
a.rows,
a.cols,
b.slice(e),
b.rows,
b.cols,
alpha,
)
.with_context(|| format!("layer {layer} target {target:?} expert {e}"))
})
.collect::<Result<Vec<_>>>()?;
TargetDelta::Experts(per_expert)
} else {
ensure!(
a.n_slices == 1,
"layer {layer} target {target:?} is a single projection carrying one \
low-rank pair, but its factors are stacked {} deep",
a.n_slices
);
TargetDelta::Dense(
LoraTargetWeights::new(a.data, a.rows, a.cols, b.data, b.rows, b.cols, alpha)
.with_context(|| format!("layer {layer} target {target:?}"))?,
)
};
if !scale_set {
default_scale = match &delta {
TargetDelta::Dense(t) => t.scale,
TargetDelta::Experts(per_expert) => per_expert.first().map_or(1.0, |t| t.scale),
};
scale_set = true;
}
layers[layer].targets[target.index()] = Some(delta);
}
Ok(Arc::new(LoraAdapterWeights {
layers,
default_scale,
classifier_weight,
classifier_bias,
class_labels,
num_classes: n_cls,
}))
}
}
fn parse_gguf_lora_name(name: &str) -> Option<(usize, LoraTarget, bool)> {
let rest = name.strip_prefix("blk.")?;
let (layer_str, rest) = rest.split_once('.')?;
let layer: usize = layer_str.parse().ok()?;
let (stem, suffix) = rest.split_once(".weight.")?;
let is_a = match suffix {
"lora_a" => true,
"lora_b" => false,
_ => return None,
};
let target = LoraTarget::from_gguf_stem(stem)?;
Some((layer, target, is_a))
}
fn parse_peft_lora_name(name: &str) -> Option<(usize, LoraTarget, bool)> {
let idx = name.find("layers.")?;
let after = &name[idx + "layers.".len()..];
let (layer_str, rest) = after.split_once('.')?;
let layer: usize = layer_str.parse().ok()?;
let rest = rest.strip_suffix(".weight")?;
let (module, ab) = rest.rsplit_once('.')?;
let is_a = match ab {
"lora_A" => true,
"lora_B" => false,
_ => return None,
};
let target = LoraTarget::from_peft_module(module)?;
Some((layer, target, is_a))
}
pub fn apply_decode(t: &LoraTargetWeights, x: &[f32], y: &mut [f32], tmp: &mut Vec<f32>) {
debug_assert_eq!(x.len(), t.k);
debug_assert_eq!(y.len(), t.d);
if t.scale == 0.0 {
return;
}
tmp.clear();
tmp.resize(t.rank, 0.0);
for (row, tmp_r) in t.a.chunks_exact(t.k).zip(tmp.iter_mut()) {
let acc: f32 = row.iter().zip(x).map(|(w, &xi)| w * xi).sum();
*tmp_r = acc * t.scale;
}
for (row, yi) in t.b.chunks_exact(t.rank).zip(y.iter_mut()) {
let acc: f32 = row.iter().zip(tmp.iter()).map(|(w, &ti)| w * ti).sum();
*yi += acc;
}
}
pub fn apply_prefill(
t: &LoraTargetWeights,
x: &[f32],
y: &mut [f32],
n: usize,
tmp: &mut Vec<f32>,
) {
debug_assert_eq!(x.len(), t.k * n);
debug_assert_eq!(y.len(), t.d * n);
if n == 0 || t.scale == 0.0 {
return;
}
tmp.clear();
tmp.resize(t.rank * n + n, 0.0);
let (tmp_rank, acc_row) = tmp.split_at_mut(t.rank * n);
for (r, tmp_row) in tmp_rank.chunks_exact_mut(n).enumerate() {
let a_row = &t.a[r * t.k..(r + 1) * t.k];
for (kk, &a_val) in a_row.iter().enumerate() {
let x_row = &x[kk * n..(kk + 1) * n];
for (t_j, &x_j) in tmp_row.iter_mut().zip(x_row) {
*t_j += a_val * x_j;
}
}
for t_j in tmp_row.iter_mut() {
*t_j *= t.scale;
}
}
for (o, y_row) in y.chunks_exact_mut(n).enumerate() {
let b_row = &t.b[o * t.rank..(o + 1) * t.rank];
acc_row.fill(0.0);
for (r, &b_val) in b_row.iter().enumerate() {
let tmp_row = &tmp_rank[r * n..(r + 1) * n];
for (a_j, &t_j) in acc_row.iter_mut().zip(tmp_row) {
*a_j += b_val * t_j;
}
}
for (y_j, &a_j) in y_row.iter_mut().zip(acc_row.iter()) {
*y_j += a_j;
}
}
}
pub fn apply_attn_qkv(
lora: &LoraAdapterWeights,
layer: usize,
x: &[f32],
q: &mut [f32],
k: &mut [f32],
v: &mut [f32],
tmp: &mut Vec<f32>,
) {
if let Some(t) = lora.get(layer, LoraTarget::AttnQ) {
apply_decode(t, x, q, tmp);
}
if let Some(t) = lora.get(layer, LoraTarget::AttnK) {
apply_decode(t, x, k, tmp);
}
if let Some(t) = lora.get(layer, LoraTarget::AttnV) {
apply_decode(t, x, v, tmp);
}
}
struct StEntry {
dtype: String,
shape: Vec<usize>,
begin: usize,
end: usize,
}
impl StEntry {
fn shape2(&self) -> Result<(usize, usize)> {
ensure!(self.shape.len() == 2, "expected 2-D, got {:?}", self.shape);
Ok((self.shape[0], self.shape[1]))
}
}
struct SafeTensors {
entries: Vec<(String, StEntry)>,
data_start: usize,
}
impl SafeTensors {
fn parse(bytes: &[u8]) -> Result<Self> {
ensure!(bytes.len() >= 8, "safetensors: truncated header length");
let header_len = usize::try_from(u64::from_le_bytes(bytes[0..8].try_into().unwrap()))
.context("safetensors: header length too large for this platform")?;
let header_end = 8usize
.checked_add(header_len)
.context("safetensors: header length overflow")?;
ensure!(
header_end <= bytes.len(),
"safetensors: header exceeds file"
);
let header: serde_json::Value = serde_json::from_slice(&bytes[8..header_end])
.context("safetensors: bad JSON header")?;
let obj = header
.as_object()
.context("safetensors: header is not an object")?;
let mut entries = Vec::new();
for (name, v) in obj {
if name == "__metadata__" {
continue;
}
let dtype = v
.get("dtype")
.and_then(|d| d.as_str())
.with_context(|| format!("{name}: missing dtype"))?
.to_string();
let shape = v
.get("shape")
.and_then(|s| s.as_array())
.with_context(|| format!("{name}: missing shape"))?
.iter()
.map(|n| n.as_u64().and_then(|u| usize::try_from(u).ok()))
.collect::<Option<Vec<_>>>()
.with_context(|| format!("{name}: bad shape (or a dim too large)"))?;
let offsets = v
.get("data_offsets")
.and_then(|o| o.as_array())
.with_context(|| format!("{name}: missing data_offsets"))?;
ensure!(
offsets.len() == 2,
"{name}: data_offsets must be [begin, end]"
);
let to_usize = |v: &serde_json::Value| -> Result<usize> {
usize::try_from(v.as_u64().context("bad data_offset")?)
.context("data_offset too large for this platform")
};
let begin = to_usize(&offsets[0])?;
let end = to_usize(&offsets[1])?;
entries.push((
name.clone(),
StEntry {
dtype,
shape,
begin,
end,
},
));
}
Ok(Self {
entries,
data_start: header_end,
})
}
fn tensors(&self) -> impl Iterator<Item = (&str, &StEntry)> {
self.entries.iter().map(|(n, e)| (n.as_str(), e))
}
fn dequantize(&self, e: &StEntry, bytes: &[u8]) -> Result<Vec<f32>> {
let start = self
.data_start
.checked_add(e.begin)
.context("safetensors: offset overflow")?;
let end = self
.data_start
.checked_add(e.end)
.context("safetensors: offset overflow")?;
ensure!(
end <= bytes.len() && start <= end,
"safetensors: tensor slice out of range"
);
let raw = &bytes[start..end];
let n = e
.shape
.iter()
.try_fold(1usize, |acc, &d| acc.checked_mul(d))
.context("safetensors: shape product overflows usize")?;
let expect_bytes = |elt: usize| n.checked_mul(elt).context("safetensors: size overflow");
match e.dtype.as_str() {
"F32" => {
ensure!(raw.len() == expect_bytes(4)?, "F32 byte count mismatch");
Ok(raw
.as_chunks::<4>()
.0
.iter()
.map(|c| f32::from_le_bytes(*c))
.collect())
}
"F16" => {
ensure!(raw.len() == expect_bytes(2)?, "F16 byte count mismatch");
Ok(raw
.as_chunks::<2>()
.0
.iter()
.map(|c| crate::quant::f16_to_f32(u16::from_le_bytes(*c)))
.collect())
}
"BF16" => {
ensure!(raw.len() == expect_bytes(2)?, "BF16 byte count mismatch");
Ok(raw
.as_chunks::<2>()
.0
.iter()
.map(|c| crate::quant::bf16_to_f32(u16::from_le_bytes(*c)))
.collect())
}
other => bail!("unsupported safetensors dtype for LoRA: {other}"),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn gguf_name_parse() {
assert_eq!(
parse_gguf_lora_name("blk.12.attn_q.weight.lora_a"),
Some((12, LoraTarget::AttnQ, true))
);
assert_eq!(
parse_gguf_lora_name("blk.0.ffn_down.weight.lora_b"),
Some((0, LoraTarget::FfnDown, false))
);
assert_eq!(
parse_gguf_lora_name("blk.4.shortconv.in_proj.weight.lora_a"),
Some((4, LoraTarget::ShortconvInProj, true))
);
assert_eq!(
parse_gguf_lora_name("blk.15.shortconv.out_proj.weight.lora_b"),
Some((15, LoraTarget::ShortconvOutProj, false))
);
assert_eq!(
parse_gguf_lora_name("blk.2.ffn_gate_exps.weight.lora_a"),
Some((2, LoraTarget::FfnGateExps, true))
);
assert_eq!(
parse_gguf_lora_name("blk.2.ffn_up_exps.weight.lora_b"),
Some((2, LoraTarget::FfnUpExps, false))
);
assert_eq!(
parse_gguf_lora_name("blk.23.ffn_down_exps.weight.lora_a"),
Some((23, LoraTarget::FfnDownExps, true))
);
assert_eq!(
parse_gguf_lora_name("blk.2.ffn_gate_inp.weight.lora_a"),
Some((2, LoraTarget::FfnGateInp, true))
);
assert_eq!(
parse_gguf_lora_name("blk.2.ffn_gate.weight.lora_a"),
Some((2, LoraTarget::FfnGate, true))
);
assert_eq!(parse_gguf_lora_name("blk.3.attn_q.weight"), None);
assert_eq!(parse_gguf_lora_name("blk.3.attn_norm.weight.lora_a"), None);
assert_eq!(parse_gguf_lora_name("token_embd.weight"), None);
}
#[test]
fn peft_name_parse() {
assert_eq!(
parse_peft_lora_name("base_model.model.model.layers.7.self_attn.q_proj.lora_A.weight"),
Some((7, LoraTarget::AttnQ, true))
);
assert_eq!(
parse_peft_lora_name("base_model.model.model.layers.31.mlp.up_proj.lora_B.weight"),
Some((31, LoraTarget::FfnUp, false))
);
assert_eq!(
parse_peft_lora_name("model.layers.2.self_attn.o_proj.lora_A.weight"),
Some((2, LoraTarget::AttnOutput, true))
);
assert_eq!(
parse_peft_lora_name("base_model.model.model.layers.0.conv.in_proj.lora_A.weight"),
Some((0, LoraTarget::ShortconvInProj, true))
);
assert_eq!(
parse_peft_lora_name("model.layers.3.conv.out_proj.lora_B.weight"),
Some((3, LoraTarget::ShortconvOutProj, false))
);
assert_eq!(
parse_peft_lora_name("base_model.model.model.layers.0.input_layernorm.weight"),
None
);
}
fn synth_safetensors(rank: usize, k: usize, d: usize, a_val: f32, b_val: f32) -> Vec<u8> {
let a: Vec<f32> = vec![a_val; rank * k];
let b: Vec<f32> = vec![b_val; d * rank];
let a_bytes: Vec<u8> = a.iter().flat_map(|x| x.to_le_bytes()).collect();
let b_bytes: Vec<u8> = b.iter().flat_map(|x| x.to_le_bytes()).collect();
let a_name = "base_model.model.model.layers.0.self_attn.q_proj.lora_A.weight";
let b_name = "base_model.model.model.layers.0.self_attn.q_proj.lora_B.weight";
let header = serde_json::json!({
a_name: { "dtype": "F32", "shape": [rank, k], "data_offsets": [0, a_bytes.len()] },
b_name: { "dtype": "F32", "shape": [d, rank], "data_offsets": [a_bytes.len(), a_bytes.len() + b_bytes.len()] },
});
let header_str = serde_json::to_vec(&header).unwrap();
let mut out = Vec::new();
out.extend_from_slice(&(header_str.len() as u64).to_le_bytes());
out.extend_from_slice(&header_str);
out.extend_from_slice(&a_bytes);
out.extend_from_slice(&b_bytes);
out
}
#[test]
fn safetensors_load_shapes_and_scale() {
let (rank, k, d) = (8, 64, 128);
let buf = synth_safetensors(rank, k, d, 0.5, 0.25);
let adapter = LoraAdapterWeights::from_safetensors_bytes(&buf, Some(16.0)).unwrap();
assert_eq!(adapter.n_layers(), 1);
assert_eq!(adapter.target_count(), 1);
let t = adapter.get(0, LoraTarget::AttnQ).expect("q_proj present");
assert_eq!((t.rank, t.k, t.d), (rank, k, d));
assert_eq!(t.a.len(), rank * k);
assert_eq!(t.b.len(), d * rank);
assert!((t.scale - 2.0).abs() < 1e-6, "scale {}", t.scale);
assert!(adapter.get(0, LoraTarget::FfnDown).is_none());
let a2 = LoraAdapterWeights::from_safetensors_bytes(&buf, None).unwrap();
assert!((a2.get(0, LoraTarget::AttnQ).unwrap().scale - 1.0).abs() < 1e-6);
}
#[test]
fn apply_math_and_noop() {
let (rank, k, d) = (2, 3, 4);
let buf = synth_safetensors(rank, k, d, 1.0, 0.0);
let adapter = LoraAdapterWeights::from_safetensors_bytes(&buf, Some(4.0)).unwrap();
let t = adapter.get(0, LoraTarget::AttnQ).unwrap();
let x = vec![1.0, 2.0, 3.0];
let mut y = vec![10.0, 20.0, 30.0, 40.0];
let before = y.clone();
let mut tmp = Vec::new();
apply_decode(t, &x, &mut y, &mut tmp);
assert_eq!(y, before, "B=0 must be a no-op");
let buf = synth_safetensors(rank, k, d, 1.0, 1.0);
let adapter = LoraAdapterWeights::from_safetensors_bytes(&buf, Some(4.0)).unwrap();
let t = adapter.get(0, LoraTarget::AttnQ).unwrap();
let mut y = vec![0.0; d];
apply_decode(t, &x, &mut y, &mut tmp);
let sum_x: f32 = x.iter().sum();
let expected = rank as f32 * (t.scale * sum_x); for &yi in &y {
assert!((yi - expected).abs() < 1e-5, "{yi} != {expected}");
}
}
#[test]
fn apply_prefill_matches_per_column_decode() {
let (rank, k, d, n) = (3, 4, 5, 3);
let buf = synth_safetensors(rank, k, d, 0.5, 0.25);
let adapter = LoraAdapterWeights::from_safetensors_bytes(&buf, Some(6.0)).unwrap();
let t = adapter.get(0, LoraTarget::AttnQ).unwrap();
let mut x = vec![0.0f32; k * n];
for i in 0..k {
for j in 0..n {
x[i * n + j] = (i as f32 + 1.0) * (j as f32 + 1.0) * 0.1;
}
}
let mut y_batched: Vec<f32> = (0..d * n).map(|i| i as f32 * 0.01).collect();
let mut y_ref = y_batched.clone();
let mut tmp = Vec::new();
apply_prefill(t, &x, &mut y_batched, n, &mut tmp);
for j in 0..n {
let x_col: Vec<f32> = (0..k).map(|i| x[i * n + j]).collect();
let mut y_col: Vec<f32> = (0..d).map(|o| y_ref[o * n + j]).collect();
apply_decode(t, &x_col, &mut y_col, &mut tmp);
for (o, &v) in y_col.iter().enumerate() {
y_ref[o * n + j] = v;
}
}
for (i, (a, b)) in y_batched.iter().zip(&y_ref).enumerate() {
assert_eq!(
a.to_bits(),
b.to_bits(),
"element {i}: batched {a} != per-column {b}"
);
}
}
#[test]
fn empty_adapter_errors() {
let header = serde_json::json!({
"some.other.weight": { "dtype": "F32", "shape": [2, 2], "data_offsets": [0, 16] },
});
let hs = serde_json::to_vec(&header).unwrap();
let mut buf = Vec::new();
buf.extend_from_slice(&(hs.len() as u64).to_le_bytes());
buf.extend_from_slice(&hs);
buf.extend_from_slice(&[0u8; 16]);
assert!(LoraAdapterWeights::from_safetensors_bytes(&buf, None).is_err());
}
fn st_buf(header: serde_json::Value, data_len: usize) -> Vec<u8> {
let hs = serde_json::to_vec(&header).unwrap();
let mut buf = Vec::new();
buf.extend_from_slice(&(hs.len() as u64).to_le_bytes());
buf.extend_from_slice(&hs);
buf.extend_from_slice(&vec![0u8; data_len]);
buf
}
#[test]
fn rejects_absurd_layer_index() {
let a = "base_model.model.model.layers.9999999999.self_attn.q_proj.lora_A.weight";
let b = "base_model.model.model.layers.9999999999.self_attn.q_proj.lora_B.weight";
let buf = st_buf(
serde_json::json!({
a: { "dtype": "F32", "shape": [1, 1], "data_offsets": [0, 4] },
b: { "dtype": "F32", "shape": [1, 1], "data_offsets": [4, 8] },
}),
8,
);
assert!(LoraAdapterWeights::from_safetensors_bytes(&buf, None).is_err());
}
#[test]
fn all_targets_are_in_index_order() {
for (i, target) in LoraTarget::ALL.into_iter().enumerate() {
assert_eq!(target.index(), i, "{target:?}");
}
let mut stems: Vec<&str> = LoraTarget::ALL.iter().map(|t| t.gguf_stem()).collect();
stems.sort_unstable();
stems.dedup();
assert_eq!(
stems.len(),
LORA_TARGET_COUNT,
"duplicate gguf stems: {stems:?}"
);
}
fn load_err(r: Result<Arc<LoraAdapterWeights>>) -> String {
match r {
Ok(_) => panic!("expected the adapter to be rejected"),
Err(e) => format!("{e:?}"),
}
}
fn push_gguf_string(out: &mut Vec<u8>, s: &str) {
out.extend_from_slice(&(s.len() as u64).to_le_bytes());
out.extend_from_slice(s.as_bytes());
}
fn synth_gguf(tensors: &[(&str, Vec<usize>, Vec<f32>)], alpha: Option<f32>) -> Vec<u8> {
let mut out = Vec::new();
out.extend_from_slice(b"GGUF");
out.extend_from_slice(&3u32.to_le_bytes());
out.extend_from_slice(&(tensors.len() as u64).to_le_bytes());
out.extend_from_slice(&(alpha.is_some() as u64).to_le_bytes());
if let Some(a) = alpha {
push_gguf_string(&mut out, "adapter.lora.alpha");
out.extend_from_slice(&6u32.to_le_bytes()); out.extend_from_slice(&a.to_le_bytes());
}
let mut offset = 0u64;
for (name, ne, data) in tensors {
push_gguf_string(&mut out, name);
out.extend_from_slice(&(ne.len() as u32).to_le_bytes());
ne.iter()
.for_each(|&d| out.extend_from_slice(&(d as u64).to_le_bytes()));
out.extend_from_slice(&0u32.to_le_bytes()); out.extend_from_slice(&offset.to_le_bytes());
offset += (data.len() * 4) as u64;
}
while !out.len().is_multiple_of(32) {
out.push(0);
}
for (_, _, data) in tensors {
out.extend(data.iter().flat_map(|x| x.to_le_bytes()));
}
out
}
fn synth_expert_gguf(n_expert: usize, rank: usize, k: usize, d: usize) -> Vec<u8> {
let a: Vec<f32> = (0..n_expert)
.flat_map(|e| std::iter::repeat_n(e as f32 + 1.0, rank * k))
.collect();
let b: Vec<f32> = (0..n_expert)
.flat_map(|e| std::iter::repeat_n(1.0 / (e as f32 + 1.0), d * rank))
.collect();
synth_gguf(
&[
(
"blk.0.ffn_gate_exps.weight.lora_a",
vec![k, rank, n_expert],
a,
),
(
"blk.0.ffn_gate_exps.weight.lora_b",
vec![rank, d, n_expert],
b,
),
],
Some(rank as f32),
)
}
#[test]
fn gguf_expert_adapter_splits_per_expert() {
let (n_expert, rank, k, d) = (4, 2, 3, 5);
let buf = synth_expert_gguf(n_expert, rank, k, d);
let adapter = LoraAdapterWeights::from_gguf_bytes(Arc::from(buf.into_boxed_slice()))
.expect("expert adapter loads");
assert_eq!(adapter.target_count(), 1);
assert!(adapter.get(0, LoraTarget::FfnGateExps).is_none());
for e in 0..n_expert {
let t = adapter
.get_expert(0, LoraTarget::FfnGateExps, e)
.unwrap_or_else(|| panic!("expert {e} present"));
assert_eq!((t.rank, t.k, t.d), (rank, k, d));
assert!(
t.a.iter().all(|&x| x == e as f32 + 1.0),
"expert {e} got A = {:?}",
&t.a[..t.a.len().min(4)]
);
assert!(
t.b.iter().all(|&x| x == 1.0 / (e as f32 + 1.0)),
"expert {e} got B = {:?}",
&t.b[..t.b.len().min(4)]
);
}
assert!(
adapter
.get_expert(0, LoraTarget::FfnGateExps, n_expert)
.is_none()
);
}
#[test]
fn expert_deltas_apply_independently() {
let (n_expert, rank, k, d) = (3, 2, 4, 3);
let buf = synth_expert_gguf(n_expert, rank, k, d);
let adapter =
LoraAdapterWeights::from_gguf_bytes(Arc::from(buf.into_boxed_slice())).unwrap();
let x = vec![1.0f32, 2.0, 3.0, 4.0];
let sum_x: f32 = x.iter().sum();
let mut tmp = Vec::new();
for e in 0..n_expert {
let t = adapter.get_expert(0, LoraTarget::FfnGateExps, e).unwrap();
let mut y = vec![0.0f32; d];
apply_decode(t, &x, &mut y, &mut tmp);
let expected = rank as f32 * sum_x;
assert!(
y.iter().all(|&v| (v - expected).abs() < 1e-4),
"expert {e}: {y:?} != {expected}"
);
}
}
#[test]
fn an_unstacked_expert_factor_is_rejected_against_the_model() {
let (rank, k, d) = (2, 8, 16);
let buf = synth_gguf(
&[
(
"blk.2.ffn_gate_exps.weight.lora_a",
vec![k, rank],
vec![1.0; rank * k],
),
(
"blk.2.ffn_gate_exps.weight.lora_b",
vec![rank, d],
vec![1.0; d * rank],
),
],
Some(rank as f32),
);
let adapter = LoraAdapterWeights::from_gguf_bytes(Arc::from(buf.into_boxed_slice()))
.expect("a rank-2 expert factor loads as a one-expert stack");
let err = adapter
.validate_dims(&moe_config())
.unwrap_err()
.to_string();
assert!(err.contains("1 expert deltas"), "{err}");
}
#[test]
fn rejects_stacked_factors_on_a_dense_target() {
let (n_expert, rank, k, d) = (2, 2, 3, 5);
let buf = synth_gguf(
&[
(
"blk.0.ffn_gate.weight.lora_a",
vec![k, rank, n_expert],
vec![1.0; n_expert * rank * k],
),
(
"blk.0.ffn_gate.weight.lora_b",
vec![rank, d, n_expert],
vec![1.0; n_expert * d * rank],
),
],
Some(rank as f32),
);
let err = load_err(LoraAdapterWeights::from_gguf_bytes(Arc::from(
buf.into_boxed_slice(),
)));
assert!(err.contains("stacked"), "{err}");
}
#[test]
fn rejects_absurd_expert_count() {
let n = MAX_LORA_EXPERTS + 1;
let buf = synth_gguf(
&[
(
"blk.0.ffn_gate_exps.weight.lora_a",
vec![1, 1, n],
vec![1.0; n],
),
(
"blk.0.ffn_gate_exps.weight.lora_b",
vec![1, 1, n],
vec![1.0; n],
),
],
Some(1.0),
);
let err = load_err(LoraAdapterWeights::from_gguf_bytes(Arc::from(
buf.into_boxed_slice(),
)));
assert!(err.contains("over the sane maximum"), "{err}");
}
#[test]
fn rejects_mismatched_expert_counts() {
let (rank, k, d) = (2, 3, 5);
let buf = synth_gguf(
&[
(
"blk.0.ffn_gate_exps.weight.lora_a",
vec![k, rank, 4],
vec![1.0; 4 * rank * k],
),
(
"blk.0.ffn_gate_exps.weight.lora_b",
vec![rank, d, 2],
vec![1.0; 2 * d * rank],
),
],
Some(rank as f32),
);
let err = load_err(LoraAdapterWeights::from_gguf_bytes(Arc::from(
buf.into_boxed_slice(),
)));
assert!(err.contains("slices"), "{err}");
}
#[test]
fn peft_expert_adapter_is_refused_not_dropped() {
let name = "base_model.model.model.layers.0.mlp.experts.3.gate_proj.lora_A.weight";
let buf = st_buf(
serde_json::json!({
name: { "dtype": "F32", "shape": [2, 2], "data_offsets": [0, 16] },
}),
16,
);
let err = load_err(LoraAdapterWeights::from_safetensors_bytes(&buf, None));
assert!(err.contains("convert_lora_to_gguf"), "{err}");
}
fn moe_config() -> crate::model::ModelConfig {
crate::model::ModelConfig {
architecture: "lfm2moe".to_string(),
n_layers: 4,
hidden_size: 8,
intermediate_size: 32,
n_heads: 2,
n_kv_heads: 2,
head_dim: 4,
vocab_size: 16,
max_seq_len: 32,
rope_theta: 10000.0,
rms_norm_eps: 1e-5,
block_types: vec![crate::model::BlockType::Attention; 4],
conv_kernel_size: None,
kv_heads_per_layer: vec![2; 4],
scalars: crate::model::ScalarMultipliers::default(),
moe: Some(crate::model::MoeConfig {
n_expert: 3,
n_expert_used: 2,
expert_ff_len: 16,
is_moe_layer: vec![false, false, true, true],
}),
is_causal: true,
class_labels: Vec::new(),
}
}
#[test]
fn has_moe_deltas_covers_the_router_and_the_experts() {
assert!(
adapter_for(0, LoraTarget::FfnGateInp, 8, 4, 1).has_moe_deltas(),
"a router delta is a routed-FFN delta"
);
assert!(
adapter_for(0, LoraTarget::FfnGateExps, 8, 16, 3).has_moe_deltas(),
"a per-expert delta is a routed-FFN delta"
);
assert!(
!adapter_for(0, LoraTarget::FfnGate, 8, 16, 1).has_moe_deltas(),
"a dense FFN delta must not trip the gate; every GPU backend applies it"
);
assert!(
!adapter_for(0, LoraTarget::AttnQ, 8, 8, 1).has_moe_deltas(),
"an attention delta must not trip the gate"
);
}
fn adapter_for(
layer: usize,
target: LoraTarget,
k: usize,
d: usize,
n_slices: usize,
) -> Arc<LoraAdapterWeights> {
let rank = 2;
let stem = target.gguf_stem();
let (ne_a, ne_b) = if n_slices > 1 {
(vec![k, rank, n_slices], vec![rank, d, n_slices])
} else {
(vec![k, rank], vec![rank, d])
};
let buf = synth_gguf(
&[
(
&format!("blk.{layer}.{stem}.weight.lora_a"),
ne_a,
vec![0.5; n_slices * rank * k],
),
(
&format!("blk.{layer}.{stem}.weight.lora_b"),
ne_b,
vec![0.5; n_slices * d * rank],
),
],
Some(rank as f32),
);
LoraAdapterWeights::from_gguf_bytes(Arc::from(buf.into_boxed_slice()))
.expect("synthetic adapter loads")
}
#[test]
fn validate_dims_rejects_a_dense_ffn_adapter_on_a_routed_layer() {
let cfg = moe_config();
adapter_for(1, LoraTarget::FfnGate, 8, 32, 1)
.validate_dims(&cfg)
.expect("dense adapter on a dense layer");
let err = adapter_for(2, LoraTarget::FfnGate, 8, 32, 1)
.validate_dims(&cfg)
.unwrap_err()
.to_string();
assert!(err.contains("mixture-of-experts"), "{err}");
}
#[test]
fn validate_dims_rejects_an_expert_adapter_on_a_dense_layer() {
let cfg = moe_config();
adapter_for(3, LoraTarget::FfnGateExps, 8, 16, 3)
.validate_dims(&cfg)
.expect("expert adapter on a routed layer");
let err = adapter_for(0, LoraTarget::FfnGateExps, 8, 16, 3)
.validate_dims(&cfg)
.unwrap_err()
.to_string();
assert!(err.contains("is dense"), "{err}");
}
#[test]
fn validate_dims_uses_the_expert_width_not_the_dense_one() {
let cfg = moe_config();
let err = adapter_for(2, LoraTarget::FfnGateExps, 8, 32, 3)
.validate_dims(&cfg)
.unwrap_err()
.to_string();
assert!(err.contains("out=16"), "{err}");
}
#[test]
fn validate_dims_accepts_the_router_and_checks_its_width() {
let cfg = moe_config();
adapter_for(2, LoraTarget::FfnGateInp, 8, 3, 1)
.validate_dims(&cfg)
.expect("router adapter validates");
let err = adapter_for(2, LoraTarget::FfnGateInp, 8, 4, 1)
.validate_dims(&cfg)
.unwrap_err()
.to_string();
assert!(err.contains("out=3"), "{err}");
}
#[test]
fn validate_dims_rejects_wrong_expert_count() {
let cfg = moe_config();
let err = adapter_for(2, LoraTarget::FfnGateExps, 8, 16, 2)
.validate_dims(&cfg)
.unwrap_err()
.to_string();
assert!(err.contains("2 expert deltas"), "{err}");
}
#[test]
fn rejects_overflow_shape() {
let big = (u64::MAX / 2) as usize;
let name = "base_model.model.model.layers.0.self_attn.q_proj.lora_A.weight";
let buf = st_buf(
serde_json::json!({
name: { "dtype": "F32", "shape": [big, big], "data_offsets": [0, 4] },
}),
4,
);
assert!(LoraAdapterWeights::from_safetensors_bytes(&buf, None).is_err());
}
#[test]
fn peft_classifier_adapter_loads_and_validates() {
let w_bytes = 24 * 4;
let b_bytes = 3 * 4;
let total_bytes = w_bytes + b_bytes;
let buf = st_buf(
serde_json::json!({
"base_model.model.classifier.weight": {
"dtype": "F32",
"shape": [3, 8],
"data_offsets": [0, w_bytes]
},
"base_model.model.classifier.bias": {
"dtype": "F32",
"shape": [3],
"data_offsets": [w_bytes, total_bytes]
}
}),
total_bytes,
);
let adapter = LoraAdapterWeights::from_safetensors_bytes(&buf, None).unwrap();
assert!(adapter.is_classifier());
assert_eq!(adapter.num_classes(), 3);
let labels = vec![
"O".to_string(),
"B-EMAIL".to_string(),
"I-EMAIL".to_string(),
];
let adapter = adapter.with_class_labels(labels);
assert_eq!(adapter.num_classes(), 3);
let cfg = crate::model::ModelConfig {
architecture: "lfm2".to_string(),
n_layers: 2,
hidden_size: 8,
intermediate_size: 32,
n_heads: 2,
n_kv_heads: 2,
head_dim: 4,
vocab_size: 16,
max_seq_len: 32,
rope_theta: 10000.0,
rms_norm_eps: 1e-5,
block_types: vec![crate::model::BlockType::Attention; 2],
conv_kernel_size: None,
kv_heads_per_layer: vec![2; 2],
scalars: crate::model::ScalarMultipliers::default(),
moe: None,
is_causal: true,
class_labels: Vec::new(),
};
assert!(adapter.validate_dims(&cfg).is_ok());
let bad_cfg = crate::model::ModelConfig {
hidden_size: 16,
..cfg
};
assert!(adapter.validate_dims(&bad_cfg).is_err());
}
}