use super::nam_json::{NamModelData, WeightsLayout};
use super::namb::{FLAG_HAS_CRC32, NambHeader, crc32_ieee, crc32_ieee_update};
use super::transpose::{lstm::transpose_lstm_gate_major, wavenet::transpose_wavenet_interleaved4};
use anyhow::Result;
use std::io::Write;
pub fn encode_namb(
data: &NamModelData,
version: u16,
target_layout: WeightsLayout,
) -> Result<Vec<u8>> {
let mut buffer = Vec::new();
let transposed_weights = if target_layout != WeightsLayout::Original {
transpose_weights(data, target_layout)?
} else {
data.weights.clone()
};
let weights_bytes: Vec<u8> = transposed_weights
.iter()
.flat_map(|&f| f.to_le_bytes().to_vec())
.collect();
let mut metadata_only = data.clone();
metadata_only.weights = Vec::new();
let json_str = serde_json::to_string(&metadata_only)?;
let json_bytes = json_str.as_bytes();
let header_flags = if version >= 2 { FLAG_HAS_CRC32 } else { 0 };
let mut header = NambHeader {
magic: 0x4E414D42,
version,
layout_type: target_layout as u8,
flags: header_flags,
reserved_v2: [0; 4],
weights_offset: (std::mem::size_of::<NambHeader>() + json_bytes.len()) as u32,
reserved1: [0; 2],
crc32: if version < 2 {
crc32_ieee(&weights_bytes)
} else {
0
},
reserved2: 0,
version_str: [0; 32],
sample_rate: data.sample_rate.unwrap_or(48000.0),
input_level_dbu: data
.metadata
.as_ref()
.and_then(|m| m.input_level_dbu)
.unwrap_or(12.0),
output_level_dbu: data
.metadata
.as_ref()
.and_then(|m| m.output_level_dbu)
.unwrap_or(-6.0),
reserved3: [0; 1],
};
let ver_str = format!("NAMB v{} ({:?})", version, target_layout);
let ver_bytes = ver_str.as_bytes();
let copy_len = ver_bytes.len().min(32);
header.version_str[..copy_len].copy_from_slice(&ver_bytes[..copy_len]);
let header_bytes = unsafe {
std::slice::from_raw_parts(
(&header as *const NambHeader) as *const u8,
std::mem::size_of::<NambHeader>(),
)
};
buffer.write_all(header_bytes)?;
buffer.write_all(json_bytes)?;
buffer.write_all(&weights_bytes)?;
if version >= 2 {
let crc = crc32_ieee_update(0xFFFFFFFFu32, &buffer[..24]);
let crc = crc32_ieee_update(crc, &buffer[28..]);
let final_crc = crc ^ 0xFFFFFFFFu32;
buffer[24..28].copy_from_slice(&final_crc.to_le_bytes());
}
Ok(buffer)
}
fn transpose_weights(data: &NamModelData, layout: WeightsLayout) -> Result<Vec<f32>> {
match layout {
WeightsLayout::GateMajorLstm => transpose_lstm_gate_major(data),
WeightsLayout::Interleaved4WaveNet => transpose_wavenet_interleaved4(data),
_ => Ok(data.weights.clone()),
}
}
pub(crate) fn ensure_capacity(
weights: &[f32],
cursor: usize,
needed: usize,
label: String,
) -> Result<()> {
if cursor + needed > weights.len() {
anyhow::bail!(
"Insufficient weights for {}: needs index {}..{}, but total length is {}",
label,
cursor,
cursor + needed,
weights.len()
);
}
Ok(())
}