use crate::ltxdit::{Shared, rows};
use crate::pool::Pool;
use std::collections::HashMap;
use std::path::Path;
pub struct LoraBranch {
a: Vec<f32>, b: Vec<f32>, rank: usize,
inn: usize,
out: usize,
scale: f32,
id: usize,
resonance: std::sync::atomic::AtomicU32,
live: std::sync::atomic::AtomicBool,
}
pub fn route_threshold() -> Option<f32> {
use std::sync::atomic::Ordering::Relaxed;
let mut bits = ROUTE.load(Relaxed);
if bits == u32::MAX {
bits = std::env::var("CMF_LORA_ROUTE")
.ok()
.and_then(|v| v.parse::<f32>().ok())
.filter(|v| *v > 0.0)
.map_or(0, f32::to_bits);
ROUTE.store(bits, Relaxed);
}
(bits != 0).then(|| f32::from_bits(bits))
}
static ROUTE: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(u32::MAX);
pub fn set_route_threshold(t: Option<f32>) {
ROUTE.store(
t.filter(|v| *v > 0.0).map_or(0, f32::to_bits),
std::sync::atomic::Ordering::Relaxed,
);
}
pub fn wants_measurement() -> bool {
probe_on() || route_threshold().is_some()
}
pub fn probe_report_on() -> bool {
probe_on()
}
fn probe_on() -> bool {
static P: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*P.get_or_init(|| std::env::var("CMF_LORA_PROBE").is_ok())
}
static NEXT_ID: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(1);
impl LoraBranch {
pub fn rank(&self) -> usize {
self.rank
}
#[cfg(target_os = "macos")]
pub(crate) fn side(&self) -> crate::gpu_metal::LoraSide<'_> {
crate::gpu_metal::LoraSide {
a: &self.a,
b: &self.b,
rank: self.rank,
scale: self.scale,
id: self.id,
}
}
pub fn live(&self) -> bool {
self.live.load(std::sync::atomic::Ordering::Relaxed)
}
pub fn resonance(&self) -> f32 {
f32::from_bits(self.resonance.load(std::sync::atomic::Ordering::Relaxed))
}
pub fn add(&self, x: &[f32], n: usize, dst: &mut [f32], pool: Option<&Pool>) {
debug_assert_eq!(x.len(), n * self.inn);
debug_assert_eq!(dst.len(), n * self.out);
if n == 0 || !self.live() {
return;
}
let mut h = vec![0f32; n * self.rank];
let mut d = vec![0f32; n * self.out];
crate::gpu::cpu_scope(|| {
crate::fcd_ops::gemm_nt(x, &self.a, &mut h, n, self.inn, self.rank, pool);
crate::fcd_ops::gemm_nt(&h, &self.b, &mut d, n, self.rank, self.out, pool);
});
let scale = self.scale;
let measure = (probe_on() || route_threshold().is_some())
&& self.resonance.load(std::sync::atomic::Ordering::Relaxed) == 0;
if measure {
let (mut dd, mut bb) = (0f64, 0f64);
for (&dv, &bv) in d.iter().zip(dst.iter()) {
dd += (scale * dv) as f64 * (scale * dv) as f64;
bb += bv as f64 * bv as f64;
}
let r = if bb > 0.0 {
(dd / bb).sqrt() as f32
} else {
f32::INFINITY
};
self.resonance.store(
r.max(f32::MIN_POSITIVE).to_bits(),
std::sync::atomic::Ordering::Relaxed,
);
if let Some(t) = route_threshold() {
if r < t {
self.live.store(false, std::sync::atomic::Ordering::Relaxed);
}
}
}
let sink = Shared(dst.as_mut_ptr());
rows(pool, n, &|s, e| {
let row = unsafe { sink.at(s * self.out, (e - s) * self.out) };
for (o, v) in row.iter_mut().enumerate() {
*v += scale * d[s * self.out + o];
}
});
}
}
pub struct SlotEmbed {
freqs: Vec<f32>,
w0: Vec<f32>,
b0: Vec<f32>,
w2: Vec<f32>,
b2: Vec<f32>,
hidden: usize,
dim: usize,
}
impl SlotEmbed {
pub fn dim(&self) -> usize {
self.dim
}
pub fn embed(&self, slot_id: usize) -> Vec<f32> {
let scaled = slot_id as f32 / 16.0;
let mut feat = Vec::with_capacity(1 + 2 * self.freqs.len());
feat.push(scaled);
for f in &self.freqs {
feat.push((scaled * f).sin());
}
for f in &self.freqs {
feat.push((scaled * f).cos());
}
let win = feat.len();
let mut hid = vec![0f32; self.hidden];
for (o, hv) in hid.iter_mut().enumerate() {
let row = &self.w0[o * win..(o + 1) * win];
let mut acc = self.b0[o];
for (fv, wv) in feat.iter().zip(row) {
acc += fv * wv;
}
*hv = acc / (1.0 + (-acc).exp());
}
let mut outv = vec![0f32; self.dim];
for (o, ov) in outv.iter_mut().enumerate() {
let row = &self.w2[o * self.hidden..(o + 1) * self.hidden];
let mut acc = self.b2[o];
for (hv, wv) in hid.iter().zip(row) {
acc += hv * wv;
}
*ov = acc;
}
outv
}
}
pub struct LoraBank {
pairs: HashMap<String, (Vec<f32>, Vec<f32>, usize, usize, usize)>,
pub slot: Option<SlotEmbed>,
pub meta: HashMap<String, String>,
scale: f32,
}
fn st_read(path: &Path) -> Result<(HashMap<String, (Vec<usize>, Vec<f32>)>, HashMap<String, String>), String> {
let bytes = std::fs::read(path).map_err(|e| format!("{}: {e}", path.display()))?;
if bytes.len() < 8 {
return Err("lora: truncated safetensors header".into());
}
let hlen = u64::from_le_bytes(bytes[..8].try_into().unwrap()) as usize;
let header: serde_json::Value = serde_json::from_slice(
bytes.get(8..8 + hlen).ok_or("lora: header past end of file")?,
)
.map_err(|e| format!("lora header: {e}"))?;
let base = 8 + hlen;
let obj = header.as_object().ok_or("lora: header not an object")?;
let mut meta = HashMap::new();
let mut out = HashMap::new();
for (name, m) in obj {
if name == "__metadata__" {
if let Some(o) = m.as_object() {
for (k, v) in o {
if let Some(s) = v.as_str() {
meta.insert(k.clone(), s.to_string());
}
}
}
continue;
}
let dtype = m["dtype"].as_str().ok_or("lora: dtype")?;
let shape: Vec<usize> = m["shape"]
.as_array()
.ok_or("lora: shape")?
.iter()
.map(|v| v.as_u64().unwrap_or(0) as usize)
.collect();
let offs = m["data_offsets"].as_array().ok_or("lora: offsets")?;
let s = offs[0].as_u64().unwrap_or(0) as usize + base;
let e = offs[1].as_u64().unwrap_or(0) as usize + base;
let raw = bytes.get(s..e).ok_or("lora: tensor span past end of file")?;
let mut data = Vec::with_capacity(shape.iter().product::<usize>().max(1));
match dtype {
"F32" => {
for c in raw.chunks_exact(4) {
data.push(f32::from_le_bytes(c.try_into().unwrap()));
}
}
"F16" => {
for c in raw.chunks_exact(2) {
data.push(cortiq_core::quant::f16_to_f32(u16::from_le_bytes(
c.try_into().unwrap(),
)));
}
}
"BF16" => {
for c in raw.chunks_exact(2) {
let b = u16::from_le_bytes(c.try_into().unwrap());
data.push(f32::from_bits((b as u32) << 16));
}
}
other => return Err(format!("lora: unsupported dtype {other} on {name}")),
}
out.insert(name.clone(), (shape, data));
}
Ok((out, meta))
}
impl LoraBank {
pub fn load(path: &Path, strength: f32) -> Result<LoraBank, String> {
let (tensors, meta) = st_read(path)?;
let mut a_side: HashMap<String, (Vec<usize>, Vec<f32>)> = HashMap::new();
let mut b_side: HashMap<String, (Vec<usize>, Vec<f32>)> = HashMap::new();
let mut slot_parts: HashMap<String, (Vec<usize>, Vec<f32>)> = HashMap::new();
for (name, val) in tensors {
let short = name
.strip_prefix("diffusion_model.")
.or_else(|| name.strip_prefix("base_model.model."))
.or_else(|| name.strip_prefix("transformer."))
.unwrap_or(&name);
let short = short.strip_prefix("dit.").unwrap_or(short).to_string();
if let Some(rest) = short.strip_prefix("reference_slot_embedding.") {
slot_parts.insert(rest.to_string(), val);
} else if let Some(base) = short.strip_suffix(".lora_A.weight") {
a_side.insert(base.to_string(), val);
} else if let Some(base) = short.strip_suffix(".lora_B.weight") {
b_side.insert(base.to_string(), val);
} else if let Some(base) = short.strip_suffix(".lora_down.weight") {
a_side.insert(base.to_string(), val);
} else if let Some(base) = short.strip_suffix(".lora_up.weight") {
b_side.insert(base.to_string(), val);
}
}
let alpha = meta
.get("alpha")
.or_else(|| meta.get("lora_alpha"))
.and_then(|v| v.parse::<f32>().ok());
let mut pairs = HashMap::new();
for (base, (ashape, adata)) in a_side {
let Some((bshape, bdata)) = b_side.remove(&base) else {
return Err(format!("lora: {base} has an A side and no B side"));
};
if ashape.len() != 2 || bshape.len() != 2 {
return Err(format!("lora: {base} is not a matrix pair"));
}
let (rank, inn) = (ashape[0], ashape[1]);
let (out, rank_b) = (bshape[0], bshape[1]);
if rank != rank_b {
return Err(format!(
"lora: {base} rank mismatch — A is {rank}, B is {rank_b}"
));
}
let scale = match alpha {
Some(al) if rank > 0 => strength * al / rank as f32,
_ => strength,
};
pairs.insert(base, (adata, bdata, rank, inn, out));
let _ = scale; }
if !b_side.is_empty() {
let orphan = b_side.keys().next().cloned().unwrap_or_default();
return Err(format!("lora: {orphan} has a B side and no A side"));
}
let slot = if slot_parts.is_empty() {
None
} else {
let need = |k: &str| -> Result<&(Vec<usize>, Vec<f32>), String> {
slot_parts
.get(k)
.ok_or_else(|| format!("lora: reference_slot_embedding.{k} is missing"))
};
let freqs = need("frequencies")?.1.clone();
let (s0, w0) = need("net.0.weight").map(|t| (t.0.clone(), t.1.clone()))?;
let b0 = need("net.0.bias")?.1.clone();
let (s2, w2) = need("net.2.weight").map(|t| (t.0.clone(), t.1.clone()))?;
let b2 = need("net.2.bias")?.1.clone();
if s0.len() != 2 || s2.len() != 2 {
return Err("lora: slot embedding layers are not matrices".into());
}
if s0[1] != 1 + 2 * freqs.len() {
return Err(format!(
"lora: slot embedding takes {} features, {} frequencies imply {}",
s0[1],
freqs.len(),
1 + 2 * freqs.len()
));
}
Some(SlotEmbed {
freqs,
w0,
b0,
w2,
b2,
hidden: s0[0],
dim: s2[0],
})
};
let scale = match alpha {
Some(al) => {
let r = pairs.values().next().map(|p| p.2).unwrap_or(1).max(1);
strength * al / r as f32
}
None => strength,
};
Ok(LoraBank { pairs, slot, meta, scale })
}
pub fn len(&self) -> usize {
self.pairs.len()
}
pub fn is_empty(&self) -> bool {
self.pairs.is_empty()
}
pub fn rank(&self) -> usize {
self.pairs.values().next().map(|p| p.2).unwrap_or(0)
}
pub fn keys(&self) -> Vec<&str> {
self.pairs.keys().map(|s| s.as_str()).collect()
}
pub fn branch_for(
&self,
name: &str,
out: usize,
inn: usize,
) -> Result<Option<LoraBranch>, String> {
let Some(br) = self.branch(name) else {
return Ok(None);
};
if br.out != out || br.inn != inn {
return Err(format!(
"lora: {name} is [{}, {}] in the adapter and [{out}, {inn}] in this container \
— that adapter was trained for a different model",
br.out, br.inn
));
}
Ok(Some(br))
}
pub fn branch(&self, name: &str) -> Option<LoraBranch> {
let key = name.strip_prefix("dit.").unwrap_or(name);
let (a, b, rank, inn, out) = self.pairs.get(key)?;
Some(LoraBranch {
a: a.clone(),
b: b.clone(),
rank: *rank,
inn: *inn,
out: *out,
scale: self.scale,
id: NEXT_ID.fetch_add(1, std::sync::atomic::Ordering::Relaxed),
resonance: Default::default(),
live: std::sync::atomic::AtomicBool::new(true),
})
}
pub fn check_reference_convention(&self) -> Result<(), String> {
if let Some(order) = self.meta.get("reference_token_order") {
if order != "prepend" {
return Err(format!(
"lora: reference_token_order={order}, this build only prepends"
));
}
}
if let Some(off) = self.meta.get("reference_slot_time_offsets") {
if off != "pic1_based_negative_time" {
return Err(format!(
"lora: reference_slot_time_offsets={off}, this build only places \
references at negative latent frames"
));
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn branch_matches_dense() {
let (n, inn, rank, out) = (3usize, 5usize, 2usize, 4usize);
let a: Vec<f32> = (0..rank * inn).map(|i| (i as f32 * 0.37).sin()).collect();
let b: Vec<f32> = (0..out * rank).map(|i| (i as f32 * 0.11).cos()).collect();
let x: Vec<f32> = (0..n * inn).map(|i| (i as f32 * 0.7).sin()).collect();
let br = LoraBranch {
a: a.clone(),
b: b.clone(),
rank,
inn,
out,
scale: 0.5,
id: 0,
resonance: Default::default(),
live: std::sync::atomic::AtomicBool::new(true),
};
let mut got = vec![1.5f32; n * out];
br.add(&x, n, &mut got, None);
for t in 0..n {
for o in 0..out {
let mut acc = 0f32;
for r in 0..rank {
let h: f32 = (0..inn).map(|i| x[t * inn + i] * a[r * inn + i]).sum();
acc += h * b[o * rank + r];
}
let want = 1.5 + 0.5 * acc;
assert!(
(got[t * out + o] - want).abs() < 1e-4,
"row {t} col {o}: {} vs {want}",
got[t * out + o]
);
}
}
}
fn write_pairs(path: &std::path::Path, pairs: &[(&str, usize, usize)]) {
use std::io::Write;
let mut header = serde_json::Map::new();
let mut blob: Vec<u8> = Vec::new();
for (base, inn, rank) in pairs {
for (suffix, shape) in [
("lora_A.weight", vec![*rank, *inn]),
("lora_B.weight", vec![*inn, *rank]),
] {
let count: usize = shape.iter().product();
let start = blob.len();
for i in 0..count {
blob.extend_from_slice(&(i as f32 * 0.25).to_le_bytes());
}
header.insert(
format!("{base}.{suffix}"),
serde_json::json!({
"dtype": "F32",
"shape": shape,
"data_offsets": [start, blob.len()],
}),
);
}
}
let hdr = serde_json::to_vec(&serde_json::Value::Object(header)).unwrap();
let mut f = std::fs::File::create(path).unwrap();
f.write_all(&(hdr.len() as u64).to_le_bytes()).unwrap();
f.write_all(&hdr).unwrap();
f.write_all(&blob).unwrap();
}
#[test]
fn mmh3_names_bind_to_container_projections() {
let dir = std::env::temp_dir().join(format!("mmh3lora{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let path = dir.join("h3.safetensors");
write_pairs(
&path,
&[
("diffusion_model.blocks.0.attn.qkv_proj", 3, 2),
("base_model.model.dit.blocks.1.mlp.fc1", 3, 2),
],
);
let bank = LoraBank::load(&path, 1.0).unwrap();
assert!(bank.branch("dit.blocks.0.attn.qkv_proj").is_some());
assert!(bank.branch("dit.blocks.1.mlp.fc1").is_some());
assert!(bank.branch("dit.blocks.2.mlp.fc1").is_none());
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn a_branch_of_the_wrong_shape_is_refused() {
let dir = std::env::temp_dir().join(format!("lorashape{}", std::process::id()));
std::fs::create_dir_all(&dir).unwrap();
let path = dir.join("wrong.safetensors");
write_pairs(&path, &[("diffusion_model.blocks.0.attn.qkv_proj", 3, 2)]);
let bank = LoraBank::load(&path, 1.0).unwrap();
let key = "dit.blocks.0.attn.qkv_proj";
assert!(
matches!(bank.branch_for(key, 3, 3), Ok(Some(_))),
"its own shape binds"
);
let err = match bank.branch_for(key, 4096, 3) {
Err(e) => e,
Ok(_) => panic!("a [4096, 3] projection must not take a [3, 3] branch"),
};
assert!(err.contains("different model"), "{err}");
assert!(matches!(bank.branch_for(key, 3, 4096), Err(_)));
assert!(matches!(
bank.branch_for("dit.blocks.9.attn.qkv_proj", 3, 3),
Ok(None)
));
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn router_silences_a_quiet_branch() {
set_route_threshold(Some(0.01));
let quiet = LoraBranch {
a: vec![1.0; 2],
b: vec![1e-6; 2],
rank: 1,
inn: 2,
out: 2,
scale: 1.0,
id: 0,
resonance: Default::default(),
live: std::sync::atomic::AtomicBool::new(true),
};
let mut out = vec![1.0f32; 2];
quiet.add(&[1.0, 1.0], 1, &mut out, None);
assert!(!quiet.live(), "a 1e-6 branch on a unit base must be routed off");
let loud = LoraBranch {
a: vec![1.0; 2],
b: vec![1.0; 2],
rank: 1,
inn: 2,
out: 2,
scale: 1.0,
id: 0,
resonance: Default::default(),
live: std::sync::atomic::AtomicBool::new(true),
};
let mut out = vec![1.0f32; 2];
loud.add(&[1.0, 1.0], 1, &mut out, None);
assert!(loud.live(), "a branch the size of the base must stay");
assert!(loud.resonance() > 1.0);
set_route_threshold(None);
}
#[test]
fn zero_b_changes_nothing() {
let br = LoraBranch {
a: vec![1.0; 4],
b: vec![0.0; 6],
rank: 2,
inn: 2,
out: 3,
scale: 1.0,
id: 0,
resonance: Default::default(),
live: std::sync::atomic::AtomicBool::new(true),
};
let mut out = vec![7.0f32; 3];
br.add(&[1.0, 2.0], 1, &mut out, None);
assert_eq!(out, vec![7.0, 7.0, 7.0]);
}
#[test]
fn names_bind_to_container_projections() {
use std::io::Write;
let dir = std::env::temp_dir().join("cmf_lora_name_test");
std::fs::create_dir_all(&dir).unwrap();
let path = dir.join("tiny.safetensors");
let names = [
("diffusion_model.transformer_blocks.0.attn1.to_q.lora_A.weight", vec![2usize, 3]),
("diffusion_model.transformer_blocks.0.attn1.to_q.lora_B.weight", vec![4, 2]),
];
let mut header = serde_json::Map::new();
let mut blob: Vec<u8> = Vec::new();
for (n, shape) in &names {
let count: usize = shape.iter().product();
let start = blob.len();
for i in 0..count {
blob.extend_from_slice(&(i as f32 * 0.25).to_le_bytes());
}
header.insert(
(*n).to_string(),
serde_json::json!({
"dtype": "F32",
"shape": shape,
"data_offsets": [start, blob.len()],
}),
);
}
let hdr = serde_json::to_vec(&serde_json::Value::Object(header)).unwrap();
let mut f = std::fs::File::create(&path).unwrap();
f.write_all(&(hdr.len() as u64).to_le_bytes()).unwrap();
f.write_all(&hdr).unwrap();
f.write_all(&blob).unwrap();
drop(f);
let bank = LoraBank::load(&path, 1.0).expect("load");
assert_eq!(bank.len(), 1);
assert_eq!(bank.rank(), 2);
assert!(bank.slot.is_none());
let br = bank
.branch("dit.transformer_blocks.0.attn1.to_q")
.expect("the container's name must find the adapter's branch");
assert_eq!(br.rank(), 2);
assert!(bank.branch("dit.transformer_blocks.0.attn1.to_k").is_none());
let mut out = vec![0f32; 4];
br.add(&[1.0, 0.0, 0.0], 1, &mut out, None);
let want = [0.25, 0.75, 1.25, 1.75].map(|v: f32| v * 0.75);
for (g, w) in out.iter().zip(&want) {
assert!((g - w).abs() < 1e-5, "{g} vs {w}");
}
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn orphan_side_is_refused() {
use std::io::Write;
let dir = std::env::temp_dir().join("cmf_lora_orphan_test");
std::fs::create_dir_all(&dir).unwrap();
let path = dir.join("orphan.safetensors");
let mut header = serde_json::Map::new();
let mut blob: Vec<u8> = Vec::new();
for i in 0..6 {
blob.extend_from_slice(&(i as f32).to_le_bytes());
}
header.insert(
"diffusion_model.transformer_blocks.0.attn1.to_q.lora_A.weight".to_string(),
serde_json::json!({"dtype":"F32","shape":[2,3],"data_offsets":[0,24]}),
);
let hdr = serde_json::to_vec(&serde_json::Value::Object(header)).unwrap();
let mut f = std::fs::File::create(&path).unwrap();
f.write_all(&(hdr.len() as u64).to_le_bytes()).unwrap();
f.write_all(&hdr).unwrap();
f.write_all(&blob).unwrap();
drop(f);
let err = match LoraBank::load(&path, 1.0) {
Err(e) => e,
Ok(_) => panic!("an adapter with a lone A side must be refused"),
};
assert!(err.contains("no B side"), "{err}");
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn slot_embedding_matches_definition() {
let s = SlotEmbed {
freqs: vec![2.0],
w0: vec![1.0, 0.5, -0.25],
b0: vec![0.1],
w2: vec![2.0],
b2: vec![-0.3],
hidden: 1,
dim: 1,
};
let v = 3.0f32 / 16.0;
let feat = [v, (v * 2.0).sin(), (v * 2.0).cos()];
let pre = 0.1 + feat[0] * 1.0 + feat[1] * 0.5 + feat[2] * -0.25;
let hid = pre / (1.0 + (-pre).exp());
let want = -0.3 + hid * 2.0;
let got = s.embed(3);
assert_eq!(got.len(), 1);
assert!((got[0] - want).abs() < 1e-6, "{} vs {want}", got[0]);
}
}