use std::path::{Path, PathBuf};
use memmap2::Mmap;
use safetensors::{Dtype, SafeTensors};
use crate::vae::error::{VaeError, VaeResult};
use crate::vae::weights::Tensor;
const TO_OUT: &str = ".to_out.";
const TO_OUT_NESTED: &str = ".to_out.0.";
#[inline]
fn bf16_to_f32(bits: u16) -> f32 {
f32::from_bits((bits as u32) << 16)
}
pub struct VaeSafetensors {
path: PathBuf,
mmap: Mmap,
}
impl VaeSafetensors {
pub fn open(path: &Path) -> VaeResult<Self> {
let file = std::fs::File::open(path).map_err(|source| VaeError::Io {
path: path.display().to_string(),
source,
})?;
let mmap = unsafe { Mmap::map(&file) }.map_err(|source| VaeError::Io {
path: path.display().to_string(),
source,
})?;
SafeTensors::deserialize(&mmap).map_err(|e| VaeError::Npy {
path: path.display().to_string(),
reason: format!("safetensors header: {e}"),
})?;
Ok(Self {
path: path.to_path_buf(),
mmap,
})
}
fn view(&self) -> VaeResult<SafeTensors<'_>> {
SafeTensors::deserialize(&self.mmap).map_err(|e| VaeError::Npy {
path: self.path.display().to_string(),
reason: format!("safetensors header: {e}"),
})
}
fn on_disk_name(key: &str) -> String {
if key.contains(TO_OUT) && !key.contains(TO_OUT_NESTED) {
key.replacen(TO_OUT, TO_OUT_NESTED, 1)
} else {
key.to_string()
}
}
pub fn load_tensor(&self, key: &str) -> VaeResult<Tensor> {
let st = self.view()?;
let name = Self::on_disk_name(key);
let view = match st.tensor(&name) {
Ok(v) => v,
Err(_) => {
return Err(VaeError::MissingWeight {
name: key.to_string(),
})
}
};
let shape: Vec<usize> = view.shape().to_vec();
let numel: usize = shape.iter().product();
let raw = view.data();
let flat = match view.dtype() {
Dtype::BF16 => {
if raw.len() != numel * 2 {
return Err(VaeError::Shape(format!(
"{key}: bf16 byte len {} != 2*{numel}",
raw.len()
)));
}
raw.chunks_exact(2)
.map(|c| bf16_to_f32(u16::from_le_bytes([c[0], c[1]])))
.collect::<Vec<f32>>()
}
Dtype::F32 => {
if raw.len() != numel * 4 {
return Err(VaeError::Shape(format!(
"{key}: f32 byte len {} != 4*{numel}",
raw.len()
)));
}
raw.chunks_exact(4)
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect::<Vec<f32>>()
}
other => {
return Err(VaeError::Npy {
path: self.path.display().to_string(),
reason: format!("{key}: unsupported dtype {other:?} (expected BF16 or F32)"),
})
}
};
if shape.len() == 4 {
let (out_c, in_c, kh, kw) = (shape[0], shape[1], shape[2], shape[3]);
let data = transpose_oikk_to_okki(&flat, out_c, in_c, kh, kw);
return Ok(Tensor {
data,
shape: vec![out_c, kh, kw, in_c],
});
}
Ok(Tensor { data: flat, shape })
}
}
fn transpose_oikk_to_okki(
src: &[f32],
out_c: usize,
in_c: usize,
kh: usize,
kw: usize,
) -> Vec<f32> {
let mut dst = vec![0.0f32; out_c * kh * kw * in_c];
for o in 0..out_c {
for i in 0..in_c {
for a in 0..kh {
for b in 0..kw {
let s = ((o * in_c + i) * kh + a) * kw + b;
let d = ((o * kh + a) * kw + b) * in_c + i;
dst[d] = src[s];
}
}
}
}
dst
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn name_map_rewrites_to_out_only() {
assert_eq!(
VaeSafetensors::on_disk_name("decoder.mid_block.attentions.0.to_out.weight"),
"decoder.mid_block.attentions.0.to_out.0.weight"
);
assert_eq!(
VaeSafetensors::on_disk_name("decoder.mid_block.attentions.0.to_out.bias"),
"decoder.mid_block.attentions.0.to_out.0.bias"
);
assert_eq!(
VaeSafetensors::on_disk_name("decoder.mid_block.attentions.0.to_out.0.weight"),
"decoder.mid_block.attentions.0.to_out.0.weight"
);
for k in [
"decoder.conv_in.weight",
"decoder.mid_block.attentions.0.to_q.weight",
"bn.running_mean",
"post_quant_conv.bias",
] {
assert_eq!(VaeSafetensors::on_disk_name(k), k);
}
}
#[test]
fn bf16_decode_is_exact_left_shift() {
for &v in &[1.5f32, -2.25, 0.0, -0.0, 65536.0, 0.5] {
let bits = (v.to_bits() >> 16) as u16;
assert_eq!(bf16_to_f32(bits), f32::from_bits((v.to_bits() >> 16) << 16));
}
assert_eq!(bf16_to_f32(0x3FC0), 1.5);
}
#[test]
fn conv_transpose_matches_manual_index() {
let (o_c, i_c, kh, kw) = (2usize, 3usize, 2usize, 2usize);
let mut src = vec![0.0f32; o_c * i_c * kh * kw];
for o in 0..o_c {
for i in 0..i_c {
for a in 0..kh {
for b in 0..kw {
let s = ((o * i_c + i) * kh + a) * kw + b;
src[s] = (s + 1) as f32;
}
}
}
}
let dst = transpose_oikk_to_okki(&src, o_c, i_c, kh, kw);
assert_eq!(dst.len(), o_c * kh * kw * i_c);
for o in 0..o_c {
for i in 0..i_c {
for a in 0..kh {
for b in 0..kw {
let s = ((o * i_c + i) * kh + a) * kw + b;
let d = ((o * kh + a) * kw + b) * i_c + i;
assert_eq!(dst[d], src[s], "mismatch at o{o} i{i} a{a} b{b}");
}
}
}
}
}
}