#[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;
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum LoraTarget {
AttnQ,
AttnK,
AttnV,
AttnOutput,
FfnGate,
FfnUp,
FfnDown,
}
impl LoraTarget {
pub const ALL: [LoraTarget; 7] = [
LoraTarget::AttnQ,
LoraTarget::AttnK,
LoraTarget::AttnV,
LoraTarget::AttnOutput,
LoraTarget::FfnGate,
LoraTarget::FfnUp,
LoraTarget::FfnDown,
];
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,
}
}
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",
}
}
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),
_ => None,
}
}
}
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(Default)]
pub struct LoraLayer {
targets: [Option<LoraTargetWeights>; 7],
}
pub struct LoraAdapterWeights {
layers: Vec<LoraLayer>,
default_scale: f32,
}
impl LoraAdapterWeights {
pub fn get(&self, layer: usize, target: LoraTarget) -> Option<&LoraTargetWeights> {
self.layers.get(layer)?.targets[target.index()].as_ref()
}
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(t) = l.targets[target.index()].as_ref() else {
continue;
};
ensure!(
layer < n_layers,
"LoRA references layer {layer} but the model has {n_layers} layers"
);
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),
};
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
);
}
}
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 in gguf.tensors.keys() {
let Some((layer, target, is_a)) = parse_gguf_lora_name(name) else {
continue;
};
let (_, rows, cols, _) = gguf.tensor_meta(name)?;
let data = gguf.get_tensor(name)?.to_f32_vec();
builder.add_factor(layer, target, is_a, data, rows, cols);
}
builder.finish(alpha_meta)
}
#[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:?}"))?;
Self::from_safetensors_bytes(&bytes, alpha)
}
pub fn from_safetensors_bytes(bytes: &[u8], alpha: Option<f32>) -> Result<Arc<Self>> {
let st = SafeTensors::parse(bytes)?;
let mut builder = AdapterBuilder::new();
for (name, entry) in st.tensors() {
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)?;
builder.add_factor(layer, target, is_a, data, rows, cols);
}
builder.finish(alpha)
}
}
const MAX_LORA_LAYERS: usize = 8192;
pub const MAX_LORA_RANK: usize = 512;
#[derive(Default)]
struct AdapterBuilder {
factors: std::collections::HashMap<(usize, usize), FactorPair>,
max_layer: usize,
}
#[derive(Default)]
struct FactorPair {
a: Option<(Vec<f32>, usize, usize)>,
b: Option<(Vec<f32>, usize, usize)>,
}
impl AdapterBuilder {
fn new() -> Self {
Self::default()
}
fn add_factor(
&mut self,
layer: usize,
target: LoraTarget,
is_a: bool,
data: Vec<f32>,
rows: usize,
cols: usize,
) {
self.max_layer = self.max_layer.max(layer);
let slot = self.factors.entry((layer, target.index())).or_default();
if is_a {
slot.a = Some((data, rows, cols));
} else {
slot.b = Some((data, rows, cols));
}
}
fn finish(self, alpha: Option<f32>) -> Result<Arc<LoraAdapterWeights>> {
ensure!(!self.factors.is_empty(), "adapter contains no LoRA tensors");
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_idx), pair) in factors {
let (a, rank_a, k) = pair
.a
.with_context(|| format!("layer {layer} target {target_idx}: missing lora_a"))?;
let (b, d, rank_b) = pair
.b
.with_context(|| format!("layer {layer} target {target_idx}: missing lora_b"))?;
let alpha = alpha.unwrap_or(rank_a as f32);
let tw = LoraTargetWeights::new(a, rank_a, k, b, d, rank_b, alpha)
.with_context(|| format!("layer {layer} target {target_idx}"))?;
if !scale_set {
default_scale = tw.scale;
scale_set = true;
}
layers[layer].targets[target_idx] = Some(tw);
}
Ok(Arc::new(LoraAdapterWeights {
layers,
default_scale,
}))
}
}
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, 0.0);
for (r, tmp_row) in tmp.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];
for (r, &b_val) in b_row.iter().enumerate() {
let tmp_row = &tmp[r * n..(r + 1) * n];
for (y_j, &t_j) in y_row.iter_mut().zip(tmp_row) {
*y_j += b_val * t_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
.chunks_exact(4)
.map(|c| f32::from_le_bytes(c.try_into().unwrap()))
.collect())
}
"F16" => {
ensure!(raw.len() == expect_bytes(2)?, "F16 byte count mismatch");
Ok(raw
.chunks_exact(2)
.map(|c| half::f16::from_le_bytes(c.try_into().unwrap()).to_f32())
.collect())
}
"BF16" => {
ensure!(raw.len() == expect_bytes(2)?, "BF16 byte count mismatch");
Ok(raw
.chunks_exact(2)
.map(|c| half::bf16::from_le_bytes(c.try_into().unwrap()).to_f32())
.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.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.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 (a, b) in y_batched.iter().zip(&y_ref) {
assert!((a - b).abs() < 1e-5, "{a} != {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 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());
}
}