use crate::EMBEDDINGS_TENSOR;
use cyberbrain_core::{Error, Result};
use safetensors::{Dtype, SafeTensors};
pub(crate) struct Matrix {
pub rows: usize,
pub dim: usize,
pub data: Vec<f32>,
pub dtype: &'static str,
}
pub(crate) fn load_matrix(bytes: &[u8]) -> Result<Matrix> {
let st = SafeTensors::deserialize(bytes)
.map_err(|e| Error::Embed(format!("weights are not a valid safetensors file: {e}")))?;
let names = st.names();
let name = if names.contains(&EMBEDDINGS_TENSOR) {
EMBEDDINGS_TENSOR
} else if names.len() == 1 {
names[0]
} else {
return Err(Error::Embed(format!(
"weights hold no tensor named {EMBEDDINGS_TENSOR:?} and are not a single-tensor \
file; found {:?}",
names
)));
};
let view = st
.tensor(name)
.map_err(|e| Error::Embed(format!("cannot read tensor {name:?}: {e}")))?;
let shape = view.shape();
let [rows, dim] = shape else {
return Err(Error::Embed(format!(
"tensor {name:?} must be 2-D [vocab, dim], has shape {shape:?}"
)));
};
let (rows, dim) = (*rows, *dim);
if rows == 0 || dim == 0 {
return Err(Error::Embed(format!(
"tensor {name:?} has a zero dimension: shape {shape:?}"
)));
}
let raw = view.data();
let n = rows * dim;
let (data, dtype) = match view.dtype() {
Dtype::F32 => (
widen(raw, 4, n, |b| f32::from_le_bytes([b[0], b[1], b[2], b[3]])),
"f32",
),
Dtype::F16 => (
widen(raw, 2, n, |b| f16_to_f32(u16::from_le_bytes([b[0], b[1]]))),
"f16",
),
Dtype::BF16 => (
widen(raw, 2, n, |b| bf16_to_f32(u16::from_le_bytes([b[0], b[1]]))),
"bf16",
),
Dtype::I8 => (widen(raw, 1, n, |b| (b[0] as i8) as f32), "int8"),
other => {
return Err(Error::Embed(format!(
"tensor {name:?} has dtype {other:?}; supported: F32, F16, BF16, I8"
)));
}
};
if data.len() != n {
return Err(Error::Embed(format!(
"tensor {name:?} declares {n} elements but carries {}",
data.len()
)));
}
if let Some(pos) = data.iter().position(|x| !x.is_finite()) {
return Err(Error::Embed(format!(
"tensor {name:?} contains a non-finite value at row {} column {}; refusing to \
load a matrix that would poison cosine similarity",
pos / dim,
pos % dim
)));
}
Ok(Matrix {
rows,
dim,
data,
dtype,
})
}
fn widen(raw: &[u8], width: usize, n: usize, f: impl Fn(&[u8]) -> f32) -> Vec<f32> {
raw.chunks_exact(width).take(n).map(f).collect()
}
pub(crate) fn f16_to_f32(bits: u16) -> f32 {
let sign = ((bits >> 15) & 1) as u32;
let exp = ((bits >> 10) & 0x1f) as u32;
let frac = (bits & 0x3ff) as u32;
let out = match exp {
0 => {
if frac == 0 {
sign << 31
} else {
let shift = frac.leading_zeros() - 21; let mant = (frac << shift) & 0x3ff;
let e = 127 - 15 + 1 - shift;
(sign << 31) | (e << 23) | (mant << 13)
}
}
31 => (sign << 31) | 0x7f80_0000 | (frac << 13),
_ => (sign << 31) | ((exp + 127 - 15) << 23) | (frac << 13),
};
f32::from_bits(out)
}
pub(crate) fn bf16_to_f32(bits: u16) -> f32 {
f32::from_bits((bits as u32) << 16)
}