#[cfg(any(feature = "bnb", feature = "gguf"))]
use std::borrow::Cow;
use std::collections::HashMap;
use std::path::{Path, PathBuf};
use crate::{AnamnesisError, Dtype, ParseLimits};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum ConvertTarget {
Safetensors,
Gguf,
BnbNf4,
}
impl ConvertTarget {
pub fn parse(raw: &str) -> crate::Result<Self> {
match raw.to_ascii_lowercase().as_str() {
"safetensors" | "bf16" => Ok(Self::Safetensors),
"gguf" => Ok(Self::Gguf),
"bnb-nf4" | "bnb_nf4" | "nf4" => Ok(Self::BnbNf4),
other => Err(AnamnesisError::Unsupported {
format: other.to_owned(),
detail: "supported convert targets: `safetensors` (alias `bf16`), \
`gguf`, `bnb-nf4`. Quantised GGUF targets need Phase 8.5"
.into(),
}),
}
}
#[must_use]
pub const fn extension(self) -> &'static str {
match self {
Self::Safetensors | Self::BnbNf4 => "safetensors",
Self::Gguf => "gguf",
}
}
#[must_use]
pub const fn suffix(self) -> &'static str {
match self {
Self::Safetensors => "bf16",
Self::Gguf => "gguf",
Self::BnbNf4 => "bnb-nf4",
}
}
}
#[derive(Debug, Clone, Default)]
#[non_exhaustive]
pub struct ConvertOptions {
pub limits: ParseLimits,
pub threads: Option<usize>,
#[cfg(feature = "gguf")]
pub gguf_metadata: HashMap<String, crate::GgufMetadataValue>,
pub output_dtype: Option<Dtype>,
}
impl ConvertOptions {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn with_limits(mut self, limits: ParseLimits) -> Self {
self.limits = limits;
self
}
#[must_use]
pub fn with_threads(mut self, n: usize) -> Self {
self.threads = Some(n.max(1));
self
}
#[must_use]
pub fn with_output_dtype(mut self, dtype: Dtype) -> Self {
self.output_dtype = Some(dtype);
self
}
#[cfg(feature = "gguf")]
#[must_use]
pub fn with_gguf_metadata(
mut self,
metadata: HashMap<String, crate::GgufMetadataValue>,
) -> Self {
self.gguf_metadata = metadata;
self
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
#[non_exhaustive]
pub struct ConvertStats {
pub tensors: usize,
pub dequantized: usize,
pub quantized: usize,
pub passthrough: usize,
}
#[derive(Debug, Clone)]
pub(crate) struct HubTensor {
pub(crate) name: String,
pub(crate) shape: Vec<usize>,
pub(crate) dtype: Dtype,
pub(crate) data: Vec<u8>,
}
#[derive(Debug, Default)]
pub(crate) struct Hub {
pub(crate) tensors: Vec<HubTensor>,
pub(crate) st_metadata: Option<HashMap<String, String>>,
pub(crate) dequantized: usize,
#[cfg(feature = "gguf")]
pub(crate) gguf_metadata: HashMap<String, crate::GgufMetadataValue>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum Format {
Safetensors,
#[cfg(feature = "pth")]
Pth,
#[cfg(feature = "npz")]
Npz,
#[cfg(feature = "gguf")]
Gguf,
}
#[cfg(not(all(feature = "pth", feature = "npz", feature = "gguf")))]
fn missing_feature_err(format_name: &str, kind: &str, feature_flag: &str) -> AnamnesisError {
AnamnesisError::Unsupported {
format: format_name.into(),
detail: format!(
"input is {kind} but the `{feature_flag}` Cargo feature is not enabled in this \
build — rebuild with `cargo install anamnesis --features cli,{feature_flag}` \
(or `cargo build --features cli,{feature_flag}`) to add support"
),
}
}
fn has_magic(path: &Path, magic: [u8; 4]) -> bool {
let mut buf = [0u8; 4];
std::fs::File::open(path)
.and_then(|mut f| {
use std::io::Read as _;
f.read_exact(&mut buf)
})
.is_ok_and(|()| buf == magic)
}
#[allow(clippy::unnecessary_wraps)]
pub(crate) fn detect_format(path: &Path) -> crate::Result<Format> {
let ext = path
.extension()
.and_then(|e| e.to_str())
.unwrap_or("")
.to_ascii_lowercase();
match ext.as_str() {
"safetensors" => Ok(Format::Safetensors),
"pth" | "pt" => {
#[cfg(feature = "pth")]
{
Ok(Format::Pth)
}
#[cfg(not(feature = "pth"))]
{
Err(missing_feature_err("PyTorch", "a .pth/.pt file", "pth"))
}
}
"npz" => {
#[cfg(feature = "npz")]
{
Ok(Format::Npz)
}
#[cfg(not(feature = "npz"))]
{
Err(missing_feature_err("NumPy NPZ", "a .npz file", "npz"))
}
}
"gguf" => {
#[cfg(feature = "gguf")]
{
Ok(Format::Gguf)
}
#[cfg(not(feature = "gguf"))]
{
Err(missing_feature_err("GGUF", "a .gguf file", "gguf"))
}
}
"bin" => {
if has_magic(path, *b"PK\x03\x04") {
#[cfg(feature = "pth")]
{
return Ok(Format::Pth);
}
#[cfg(not(feature = "pth"))]
{
return Err(missing_feature_err(
"PyTorch",
"a .bin file with ZIP magic (PyTorch pickle archive)",
"pth",
));
}
}
if has_magic(path, *b"GGUF") {
#[cfg(feature = "gguf")]
{
return Ok(Format::Gguf);
}
#[cfg(not(feature = "gguf"))]
{
return Err(missing_feature_err(
"GGUF",
"a .bin file with GGUF magic",
"gguf",
));
}
}
Ok(Format::Safetensors)
}
_ => {
if has_magic(path, *b"GGUF") {
#[cfg(feature = "gguf")]
{
return Ok(Format::Gguf);
}
#[cfg(not(feature = "gguf"))]
{
return Err(missing_feature_err(
"GGUF",
"a file whose first four bytes are the GGUF magic",
"gguf",
));
}
}
Ok(Format::Safetensors)
}
}
}
const QUANT_SUFFIXES: &[&str] = &[
"-GPTQ-Int4",
"-GPTQ-Int8",
"-gptq-int4",
"-gptq-int8",
"-gptq4",
"-gptq8",
"-GPTQ",
"-gptq",
"_gptq",
"-AWQ",
"-awq",
"_awq",
"-bnb-4bit",
"-bnb-int8",
"-bnb",
"_bnb",
"-4bit",
"-int4",
"-int8",
"-fp8",
"_fp8",
"-FP8",
];
#[must_use]
pub(crate) fn strip_quant_suffix(stem: &str) -> &str {
QUANT_SUFFIXES
.iter()
.find_map(|qs| stem.strip_suffix(qs))
.unwrap_or(stem)
}
#[must_use]
pub fn derive_output_path(input: &Path, target: ConvertTarget) -> PathBuf {
derive_output_path_for_dtype(input, target, Dtype::BF16)
}
#[must_use]
pub fn derive_output_path_for_dtype(
input: &Path,
target: ConvertTarget,
dequant_dtype: Dtype,
) -> PathBuf {
let stem = input
.file_stem()
.and_then(|s| s.to_str())
.unwrap_or("output");
#[allow(clippy::wildcard_enum_match_arm)]
let suffix = match target {
ConvertTarget::Safetensors => match dequant_dtype {
Dtype::F32 => "f32",
Dtype::F16 => "f16",
_ => "bf16",
},
other => other.suffix(),
};
let new_name = format!(
"{}-{}.{}",
strip_quant_suffix(stem),
suffix,
target.extension()
);
input
.parent()
.map_or_else(|| PathBuf::from(&new_name), |p| p.join(&new_name))
}
pub fn convert(
input: &Path,
target: ConvertTarget,
output: &Path,
options: &ConvertOptions,
) -> crate::Result<ConvertStats> {
let hub = read_hub(input, options)?;
write_hub(&hub, target, output, options)
}
fn resolve_output_dtype(options: &ConvertOptions) -> crate::Result<Dtype> {
let requested = options.output_dtype.unwrap_or(Dtype::BF16);
#[allow(clippy::wildcard_enum_match_arm)]
match requested {
Dtype::BF16 | Dtype::F32 | Dtype::F16 => Ok(requested),
other => Err(AnamnesisError::Unsupported {
format: "convert".into(),
detail: format!(
"output dtype {other} is not a dequantisation output width \
(supported: bf16, f32, f16)"
),
}),
}
}
fn read_hub(input: &Path, options: &ConvertOptions) -> crate::Result<Hub> {
let threads = crate::model::resolve_thread_budget(options.threads);
let out_dtype = resolve_output_dtype(options)?;
match detect_format(input)? {
Format::Safetensors => read_safetensors(input, &options.limits, threads, out_dtype),
#[cfg(feature = "pth")]
Format::Pth => read_pth(input, &options.limits),
#[cfg(feature = "npz")]
Format::Npz => read_npz(input, &options.limits),
#[cfg(feature = "gguf")]
Format::Gguf => read_gguf(input, &options.limits, threads, out_dtype),
}
}
#[cfg_attr(not(feature = "gguf"), allow(unused_variables))]
fn write_hub(
hub: &Hub,
target: ConvertTarget,
output: &Path,
options: &ConvertOptions,
) -> crate::Result<ConvertStats> {
match target {
ConvertTarget::Safetensors => write_safetensors(hub, output),
#[cfg(feature = "gguf")]
ConvertTarget::Gguf => write_gguf_target(hub, output, options),
#[cfg(not(feature = "gguf"))]
ConvertTarget::Gguf => Err(AnamnesisError::Unsupported {
format: "gguf".into(),
detail: "GGUF emit requires the `gguf` Cargo feature; rebuild with \
`--features cli,gguf`"
.into(),
}),
#[cfg(feature = "bnb")]
ConvertTarget::BnbNf4 => write_bnb_nf4_target(hub, output),
#[cfg(not(feature = "bnb"))]
ConvertTarget::BnbNf4 => Err(AnamnesisError::Unsupported {
format: "bnb-nf4".into(),
detail: "BnB-NF4 encode requires the `bnb` Cargo feature; rebuild with \
`--features cli,bnb`"
.into(),
}),
}
}
fn read_safetensors(
path: &Path,
limits: &ParseLimits,
threads: usize,
out_dtype: Dtype,
) -> crate::Result<Hub> {
let model = crate::parse_with_limits(path, limits)?;
#[allow(clippy::wildcard_enum_match_arm)]
let (tensors, dequantized) = match out_dtype {
Dtype::BF16 => model.hub_tensors::<crate::Bf16Out>(threads)?,
Dtype::F32 => model.hub_tensors::<crate::F32Out>(threads)?,
Dtype::F16 => model.hub_tensors::<crate::F16Out>(threads)?,
other => {
return Err(AnamnesisError::Unsupported {
format: "safetensors".into(),
detail: format!("{other} is not a dequantisation output width"),
});
}
};
Ok(Hub {
tensors,
st_metadata: model.header.metadata.clone(),
dequantized,
#[cfg(feature = "gguf")]
gguf_metadata: HashMap::new(),
})
}
#[cfg(feature = "npz")]
fn read_npz(path: &Path, limits: &ParseLimits) -> crate::Result<Hub> {
let mut map = crate::parse_npz_with_limits(path, limits)?;
let mut names: Vec<String> = map.keys().cloned().collect();
names.sort();
let mut tensors = Vec::with_capacity(names.len());
for name in names {
let t = map.remove(&name).ok_or_else(|| AnamnesisError::Parse {
reason: format!("NPZ tensor `{name}` vanished mid-iteration"),
})?;
tensors.push(HubTensor {
name,
shape: t.shape,
dtype: npz_dtype_to_hub(t.dtype),
data: t.data,
});
}
Ok(Hub {
tensors,
st_metadata: None,
dequantized: 0,
#[cfg(feature = "gguf")]
gguf_metadata: HashMap::new(),
})
}
#[cfg(feature = "pth")]
fn read_pth(path: &Path, limits: &ParseLimits) -> crate::Result<Hub> {
let parsed = crate::parse_pth_with_limits(path, limits)?;
let pth_tensors = parsed.tensors()?;
let mut tensors = Vec::with_capacity(pth_tensors.len());
for t in pth_tensors {
tensors.push(HubTensor {
name: t.name,
shape: t.shape,
dtype: pth_dtype_to_hub(t.dtype)?,
data: t.data.into_owned(),
});
}
Ok(Hub {
tensors,
st_metadata: None,
dequantized: 0,
#[cfg(feature = "gguf")]
gguf_metadata: HashMap::new(),
})
}
#[cfg(feature = "gguf")]
fn read_gguf(
path: &Path,
limits: &ParseLimits,
threads: usize,
out_dtype: Dtype,
) -> crate::Result<Hub> {
#[allow(clippy::wildcard_enum_match_arm)]
match out_dtype {
Dtype::BF16 => read_gguf_as::<crate::Bf16Out>(path, limits, threads),
Dtype::F32 => read_gguf_as::<crate::F32Out>(path, limits, threads),
Dtype::F16 => read_gguf_as::<crate::F16Out>(path, limits, threads),
other => Err(AnamnesisError::Unsupported {
format: "GGUF".into(),
detail: format!(
"output dtype {other} is not a dequantisation output width \
(supported: bf16, f32, f16)"
),
}),
}
}
#[cfg(feature = "gguf")]
fn read_gguf_as<E: crate::OutputElement>(
path: &Path,
limits: &ParseLimits,
threads: usize,
) -> crate::Result<Hub> {
let parsed = crate::parse_gguf_with_limits(path, limits)?;
let views: Vec<crate::GgufTensor<'_>> = parsed.tensors().collect();
let work_bytes: u64 = views.iter().fold(0u64, |acc, view| {
acc.saturating_add(u64::try_from(view.data.len()).unwrap_or(u64::MAX))
});
let tensors: Vec<HubTensor> = crate::parallel::map_indexed(
&views,
threads,
work_bytes,
|_, tensor| {
let mut shape: Vec<usize> = tensor.shape.to_vec();
shape.reverse();
if tensor.dtype.is_quantized() {
let n_elements = tensor
.shape
.iter()
.try_fold(1usize, |acc, &d| acc.checked_mul(d))
.ok_or_else(|| AnamnesisError::Parse {
reason: format!(
"GGUF tensor `{}` shape {:?} element count overflows usize",
tensor.name, tensor.shape
),
})?;
let data = crate::dequantize_gguf::<E>(&tensor.data, tensor.dtype, n_elements)?;
Ok(HubTensor {
name: tensor.name.to_owned(),
shape,
dtype: E::DTYPE,
data,
})
} else {
Ok(HubTensor {
name: tensor.name.to_owned(),
shape,
dtype: gguf_type_to_hub(tensor.dtype)?,
data: tensor.data.to_vec(),
})
}
},
|_| {},
)?;
let dequantized = views.iter().filter(|v| v.dtype.is_quantized()).count();
Ok(Hub {
tensors,
st_metadata: None,
dequantized,
gguf_metadata: parsed.metadata().clone(),
})
}
fn write_safetensors(hub: &Hub, output: &Path) -> crate::Result<ConvertStats> {
let mut views: Vec<(String, safetensors::tensor::TensorView<'_>)> =
Vec::with_capacity(hub.tensors.len());
for t in &hub.tensors {
let st_dtype = t.dtype.to_safetensors_dtype()?;
let view = safetensors::tensor::TensorView::new(st_dtype, t.shape.clone(), &t.data)
.map_err(|e| AnamnesisError::Parse {
reason: format!("failed to create TensorView for `{}`: {e}", t.name),
})?;
views.push((t.name.clone(), view));
}
safetensors::tensor::serialize_to_file(views, hub.st_metadata.clone(), output).map_err(
#[allow(clippy::wildcard_enum_match_arm)]
|e| match e {
safetensors::SafeTensorError::IoError(io_err) => AnamnesisError::Io(io_err),
other => AnamnesisError::Parse {
reason: format!("failed to write safetensors file: {other}"),
},
},
)?;
Ok(ConvertStats {
tensors: hub.tensors.len(),
dequantized: hub.dequantized,
quantized: 0,
passthrough: hub.tensors.len().saturating_sub(hub.dequantized),
})
}
#[cfg(feature = "gguf")]
fn write_gguf_target(
hub: &Hub,
output: &Path,
options: &ConvertOptions,
) -> crate::Result<ConvertStats> {
use crate::{GgufWriteTensor, write_gguf};
let mut owned: Vec<(String, crate::GgufType, Vec<usize>, &[u8])> =
Vec::with_capacity(hub.tensors.len());
for t in &hub.tensors {
let gguf_dtype = hub_dtype_to_gguf(t.dtype)?;
let mut msb_first = t.shape.clone();
msb_first.reverse();
owned.push((t.name.clone(), gguf_dtype, msb_first, t.data.as_slice()));
}
let tensors: Vec<GgufWriteTensor<'_>> = owned
.iter()
.map(|(name, dtype, shape, data)| GgufWriteTensor {
name: name.as_str(),
shape: shape.as_slice(),
dtype: *dtype,
data,
})
.collect();
let metadata: Cow<'_, HashMap<String, crate::GgufMetadataValue>> =
if options.gguf_metadata.is_empty() {
Cow::Borrowed(&hub.gguf_metadata)
} else {
let mut merged = hub.gguf_metadata.clone();
merged.extend(
options
.gguf_metadata
.iter()
.map(|(k, v)| (k.clone(), v.clone())),
);
Cow::Owned(merged)
};
write_gguf(output, &tensors, &metadata)?;
Ok(ConvertStats {
tensors: tensors.len(),
dequantized: hub.dequantized,
quantized: 0,
passthrough: tensors.len().saturating_sub(hub.dequantized),
})
}
#[cfg(feature = "bnb")]
fn write_bnb_nf4_target(hub: &Hub, output: &Path) -> crate::Result<ConvertStats> {
use crate::{BnbWriteInput, classify_inputs, write_bnb_nf4_safetensors};
let mut owned: Vec<(String, Vec<usize>, Cow<'_, [u8]>)> = Vec::with_capacity(hub.tensors.len());
for t in &hub.tensors {
let bf16 = to_bf16_bytes(&t.data, t.dtype, &t.name)?;
owned.push((t.name.clone(), t.shape.clone(), bf16));
}
let inputs: Vec<BnbWriteInput<'_>> = owned
.iter()
.map(|(name, shape, bf16)| BnbWriteInput {
name: name.as_str(),
shape: shape.as_slice(),
bf16_data: bf16.as_ref(),
})
.collect();
let stats = classify_inputs(&inputs);
write_bnb_nf4_safetensors(&inputs, output)?;
Ok(ConvertStats {
tensors: inputs.len(),
dequantized: hub.dequantized,
quantized: stats.quantized,
passthrough: stats.passthrough,
})
}
#[cfg(feature = "gguf")]
pub fn parse_gguf_kv_arg(arg: &str) -> crate::Result<(String, crate::GgufMetadataValue)> {
let (key, value) = arg.split_once('=').ok_or_else(|| AnamnesisError::Parse {
reason: format!("--gguf-kv `{arg}`: expected `key=value`"),
})?;
if key.is_empty() {
return Err(AnamnesisError::Parse {
reason: format!("--gguf-kv `{arg}`: empty key"),
});
}
Ok((
key.to_owned(),
crate::GgufMetadataValue::String(value.to_owned()),
))
}
#[cfg(feature = "gguf")]
pub fn parse_gguf_metadata_json(
json: &str,
) -> crate::Result<HashMap<String, crate::GgufMetadataValue>> {
let parsed: serde_json::Value =
serde_json::from_str(json).map_err(|e| AnamnesisError::Parse {
reason: format!("--gguf-metadata: invalid JSON: {e}"),
})?;
let obj = parsed.as_object().ok_or_else(|| AnamnesisError::Parse {
reason: format!(
"--gguf-metadata: expected a top-level JSON object, found {}",
json_type_name(&parsed)
),
})?;
let mut out = HashMap::with_capacity(obj.len());
for (key, value) in obj {
out.insert(key.clone(), json_to_metadata_value(key, value)?);
}
Ok(out)
}
#[cfg(feature = "gguf")]
const fn json_type_name(value: &serde_json::Value) -> &'static str {
match *value {
serde_json::Value::Null => "null",
serde_json::Value::Bool(_) => "a boolean",
serde_json::Value::Number(_) => "a number",
serde_json::Value::String(_) => "a string",
serde_json::Value::Array(_) => "an array",
serde_json::Value::Object(_) => "an object",
}
}
#[cfg(feature = "gguf")]
fn json_as_int(key: &str, value: &serde_json::Value) -> crate::Result<i128> {
let number = value.as_number().ok_or_else(|| AnamnesisError::Parse {
reason: format!(
"--gguf-metadata `{key}`: expected an integer, found {}",
json_type_name(value)
),
})?;
if let Some(u) = number.as_u64() {
return Ok(i128::from(u));
}
if let Some(i) = number.as_i64() {
return Ok(i128::from(i));
}
Err(AnamnesisError::Parse {
reason: format!("--gguf-metadata `{key}`: expected an integer, found a float"),
})
}
#[cfg(feature = "gguf")]
fn json_as_float(key: &str, value: &serde_json::Value) -> crate::Result<f64> {
value.as_f64().ok_or_else(|| AnamnesisError::Parse {
reason: format!(
"--gguf-metadata `{key}`: expected a number, found {}",
json_type_name(value)
),
})
}
#[cfg(feature = "gguf")]
fn json_as_bool(key: &str, value: &serde_json::Value) -> crate::Result<bool> {
value.as_bool().ok_or_else(|| AnamnesisError::Parse {
reason: format!(
"--gguf-metadata `{key}`: expected a boolean, found {}",
json_type_name(value)
),
})
}
#[cfg(feature = "gguf")]
fn json_as_string(key: &str, value: &serde_json::Value) -> crate::Result<String> {
value
.as_str()
.map(ToOwned::to_owned)
.ok_or_else(|| AnamnesisError::Parse {
reason: format!(
"--gguf-metadata `{key}`: expected a string, found {}",
json_type_name(value)
),
})
}
#[cfg(feature = "gguf")]
fn narrow_int<T: TryFrom<i128>>(key: &str, raw: i128, type_name: &str) -> crate::Result<T> {
T::try_from(raw).map_err(|_| AnamnesisError::Parse {
reason: format!("--gguf-metadata `{key}`: value {raw} is out of range for {type_name}"),
})
}
#[cfg(feature = "gguf")]
fn scalar_of_type(
key: &str,
type_tag: &str,
value: &serde_json::Value,
) -> crate::Result<crate::GgufMetadataValue> {
use crate::GgufMetadataValue as V;
Ok(match type_tag {
"u8" => V::U8(narrow_int(key, json_as_int(key, value)?, "u8")?),
"i8" => V::I8(narrow_int(key, json_as_int(key, value)?, "i8")?),
"u16" => V::U16(narrow_int(key, json_as_int(key, value)?, "u16")?),
"i16" => V::I16(narrow_int(key, json_as_int(key, value)?, "i16")?),
"u32" => V::U32(narrow_int(key, json_as_int(key, value)?, "u32")?),
"i32" => V::I32(narrow_int(key, json_as_int(key, value)?, "i32")?),
"u64" => V::U64(narrow_int(key, json_as_int(key, value)?, "u64")?),
"i64" => V::I64(narrow_int(key, json_as_int(key, value)?, "i64")?),
"f32" => {
#[allow(clippy::as_conversions, clippy::cast_possible_truncation)]
let narrowed = json_as_float(key, value)? as f32;
V::F32(narrowed)
}
"f64" => V::F64(json_as_float(key, value)?),
"bool" => V::Bool(json_as_bool(key, value)?),
"string" => V::String(json_as_string(key, value)?),
other => {
return Err(AnamnesisError::Parse {
reason: format!(
"--gguf-metadata `{key}`: unknown type `{other}` \
(expected u8/i8/u16/i16/u32/i32/u64/i64/f32/f64/bool/string/array)"
),
});
}
})
}
#[cfg(feature = "gguf")]
fn array_of_type(
key: &str,
item_type: &str,
items: &[serde_json::Value],
) -> crate::Result<crate::GgufMetadataArray> {
use crate::GgufMetadataArray as A;
fn collect<T, F: Fn(&serde_json::Value) -> crate::Result<T>>(
items: &[serde_json::Value],
f: F,
) -> crate::Result<Vec<T>> {
items.iter().map(f).collect()
}
Ok(match item_type {
"u8" => A::U8(collect(items, |v| {
narrow_int(key, json_as_int(key, v)?, "u8")
})?),
"i8" => A::I8(collect(items, |v| {
narrow_int(key, json_as_int(key, v)?, "i8")
})?),
"u16" => A::U16(collect(items, |v| {
narrow_int(key, json_as_int(key, v)?, "u16")
})?),
"i16" => A::I16(collect(items, |v| {
narrow_int(key, json_as_int(key, v)?, "i16")
})?),
"u32" => A::U32(collect(items, |v| {
narrow_int(key, json_as_int(key, v)?, "u32")
})?),
"i32" => A::I32(collect(items, |v| {
narrow_int(key, json_as_int(key, v)?, "i32")
})?),
"u64" => A::U64(collect(items, |v| {
narrow_int(key, json_as_int(key, v)?, "u64")
})?),
"i64" => A::I64(collect(items, |v| {
narrow_int(key, json_as_int(key, v)?, "i64")
})?),
"f32" => A::F32(collect(items, |v| {
#[allow(clippy::as_conversions, clippy::cast_possible_truncation)]
let narrowed = json_as_float(key, v)? as f32;
Ok(narrowed)
})?),
"f64" => A::F64(collect(items, |v| json_as_float(key, v))?),
"bool" => A::Bool(collect(items, |v| json_as_bool(key, v))?),
"string" => A::String(collect(items, |v| json_as_string(key, v))?),
other => {
return Err(AnamnesisError::Parse {
reason: format!(
"--gguf-metadata `{key}`: unknown array item type `{other}` \
(expected u8/i8/u16/i16/u32/i32/u64/i64/f32/f64/bool/string)"
),
});
}
})
}
#[cfg(feature = "gguf")]
fn infer_type_tag(key: &str, value: &serde_json::Value) -> crate::Result<&'static str> {
match *value {
serde_json::Value::Bool(_) => Ok("bool"),
serde_json::Value::String(_) => Ok("string"),
serde_json::Value::Number(ref n) => {
if n.is_f64() && !n.is_u64() && !n.is_i64() {
return Ok("f32");
}
let raw = json_as_int(key, value)?;
if (0..=i128::from(u32::MAX)).contains(&raw) {
Ok("u32")
} else if i64::try_from(raw).is_ok() {
Ok("i64")
} else {
Ok("u64")
}
}
serde_json::Value::Null | serde_json::Value::Array(_) | serde_json::Value::Object(_) => {
Err(AnamnesisError::Parse {
reason: format!(
"--gguf-metadata `{key}`: cannot infer a scalar type from {}",
json_type_name(value)
),
})
}
}
}
#[cfg(feature = "gguf")]
fn json_to_metadata_value(
key: &str,
value: &serde_json::Value,
) -> crate::Result<crate::GgufMetadataValue> {
if let Some(obj) = value.as_object() {
let type_tag = obj
.get("type")
.and_then(serde_json::Value::as_str)
.ok_or_else(|| AnamnesisError::Parse {
reason: format!(
"--gguf-metadata `{key}`: an object value must carry a string `type` field \
(explicit form: {{\"type\": \"u32\", \"value\": 32}})"
),
})?;
let inner = obj.get("value").ok_or_else(|| AnamnesisError::Parse {
reason: format!("--gguf-metadata `{key}`: explicit form is missing `value`"),
})?;
if type_tag == "array" {
let item_type = obj
.get("item_type")
.and_then(serde_json::Value::as_str)
.ok_or_else(|| AnamnesisError::Parse {
reason: format!(
"--gguf-metadata `{key}`: an `array` needs a string `item_type`"
),
})?;
let items = inner.as_array().ok_or_else(|| AnamnesisError::Parse {
reason: format!(
"--gguf-metadata `{key}`: `array` expects a JSON array `value`, found {}",
json_type_name(inner)
),
})?;
return Ok(crate::GgufMetadataValue::Array(Box::new(array_of_type(
key, item_type, items,
)?)));
}
return scalar_of_type(key, type_tag, inner);
}
if let Some(items) = value.as_array() {
let first = items.first().ok_or_else(|| AnamnesisError::Parse {
reason: format!(
"--gguf-metadata `{key}`: cannot infer an item type from an empty array \
(use the explicit form: {{\"type\": \"array\", \"item_type\": \"i32\", \
\"value\": []}})"
),
})?;
let item_type = infer_type_tag(key, first)?;
return Ok(crate::GgufMetadataValue::Array(Box::new(array_of_type(
key, item_type, items,
)?)));
}
let type_tag = infer_type_tag(key, value)?;
scalar_of_type(key, type_tag, value)
}
#[cfg(feature = "npz")]
const fn npz_dtype_to_hub(dtype: crate::NpzDtype) -> Dtype {
use crate::NpzDtype;
match dtype {
NpzDtype::Bool => Dtype::Bool,
NpzDtype::U8 => Dtype::U8,
NpzDtype::I8 => Dtype::I8,
NpzDtype::U16 => Dtype::U16,
NpzDtype::I16 => Dtype::I16,
NpzDtype::U32 => Dtype::U32,
NpzDtype::I32 => Dtype::I32,
NpzDtype::U64 => Dtype::U64,
NpzDtype::I64 => Dtype::I64,
NpzDtype::F16 => Dtype::F16,
NpzDtype::BF16 => Dtype::BF16,
NpzDtype::F32 => Dtype::F32,
NpzDtype::F64 => Dtype::F64,
}
}
#[cfg(feature = "pth")]
#[allow(clippy::unnecessary_wraps)]
const fn pth_dtype_to_hub(dtype: crate::PthDtype) -> crate::Result<Dtype> {
use crate::PthDtype;
Ok(match dtype {
PthDtype::F16 => Dtype::F16,
PthDtype::BF16 => Dtype::BF16,
PthDtype::F32 => Dtype::F32,
PthDtype::F64 => Dtype::F64,
PthDtype::U8 => Dtype::U8,
PthDtype::I8 => Dtype::I8,
PthDtype::I16 => Dtype::I16,
PthDtype::I32 => Dtype::I32,
PthDtype::I64 => Dtype::I64,
PthDtype::Bool => Dtype::Bool,
})
}
#[cfg(feature = "gguf")]
fn gguf_type_to_hub(dtype: crate::GgufType) -> crate::Result<Dtype> {
use crate::GgufType;
#[allow(clippy::wildcard_enum_match_arm)]
match dtype {
GgufType::F32 => Ok(Dtype::F32),
GgufType::F16 => Ok(Dtype::F16),
GgufType::BF16 => Ok(Dtype::BF16),
GgufType::F64 => Ok(Dtype::F64),
GgufType::I8 => Ok(Dtype::I8),
GgufType::I16 => Ok(Dtype::I16),
GgufType::I32 => Ok(Dtype::I32),
GgufType::I64 => Ok(Dtype::I64),
other => Err(AnamnesisError::Unsupported {
format: "GGUF".into(),
detail: format!("no scalar hub dtype for {other}"),
}),
}
}
#[cfg(feature = "gguf")]
fn hub_dtype_to_gguf(dtype: Dtype) -> crate::Result<crate::GgufType> {
use crate::GgufType;
match dtype {
Dtype::F32 => Ok(GgufType::F32),
Dtype::F16 => Ok(GgufType::F16),
Dtype::BF16 => Ok(GgufType::BF16),
Dtype::F64 => Ok(GgufType::F64),
Dtype::I8 => Ok(GgufType::I8),
Dtype::I16 => Ok(GgufType::I16),
Dtype::I32 => Ok(GgufType::I32),
Dtype::I64 => Ok(GgufType::I64),
Dtype::F8E4M3
| Dtype::F8E5M2
| Dtype::Bool
| Dtype::U8
| Dtype::U16
| Dtype::U32
| Dtype::U64 => Err(AnamnesisError::Unsupported {
format: "gguf".into(),
detail: format!(
"no GGUF dtype counterpart for {dtype} \
(Bool/unsigned-integer/FP8 are not in the GGUF scalar surface)"
),
}),
}
}
#[cfg(feature = "bnb")]
fn to_bf16_bytes<'a>(data: &'a [u8], dtype: Dtype, name: &str) -> crate::Result<Cow<'a, [u8]>> {
match dtype {
Dtype::BF16 => Ok(Cow::Borrowed(data)),
Dtype::F32 => {
if !data.len().is_multiple_of(4) {
return Err(AnamnesisError::Parse {
reason: format!(
"bnb-nf4 `{name}`: F32 byte count {} is not a multiple of 4",
data.len()
),
});
}
let mut out = vec![0u8; data.len() / 2];
for (chunk, out_pair) in data.chunks_exact(4).zip(out.chunks_exact_mut(2)) {
#[allow(clippy::indexing_slicing)]
let arr: [u8; 4] = [chunk[0], chunk[1], chunk[2], chunk[3]];
let bf16 = crate::remember::fp8::f32_bits_to_bf16_bits(u32::from_le_bytes(arr));
out_pair.copy_from_slice(&bf16.to_le_bytes());
}
Ok(Cow::Owned(out))
}
Dtype::F16 => {
if !data.len().is_multiple_of(2) {
return Err(AnamnesisError::Parse {
reason: format!(
"bnb-nf4 `{name}`: F16 byte count {} is not a multiple of 2",
data.len()
),
});
}
let mut out = vec![0u8; data.len()];
for (chunk, out_pair) in data.chunks_exact(2).zip(out.chunks_exact_mut(2)) {
#[allow(clippy::indexing_slicing)]
let arr: [u8; 2] = [chunk[0], chunk[1]];
let bits = half::f16::from_le_bytes(arr).to_f32().to_bits();
let bf16 = crate::remember::fp8::f32_bits_to_bf16_bits(bits);
out_pair.copy_from_slice(&bf16.to_le_bytes());
}
Ok(Cow::Owned(out))
}
Dtype::F8E4M3
| Dtype::F8E5M2
| Dtype::F64
| Dtype::Bool
| Dtype::U8
| Dtype::I8
| Dtype::U16
| Dtype::I16
| Dtype::U32
| Dtype::I32
| Dtype::U64
| Dtype::I64 => Err(AnamnesisError::Unsupported {
format: "bnb-nf4".into(),
detail: format!(
"tensor `{name}` has dtype {dtype}; only F32/F16/BF16 inputs are \
supported for BnB-NF4 conversion"
),
}),
}
}
#[cfg(all(test, feature = "gguf"))]
#[allow(clippy::panic, clippy::unwrap_used, clippy::expect_used)]
mod gguf_metadata_tests {
use super::{parse_gguf_kv_arg, parse_gguf_metadata_json};
use crate::{GgufMetadataArray as A, GgufMetadataValue as V};
#[test]
fn plain_scalars_are_inferred() {
let meta = parse_gguf_metadata_json(
r#"{"s": "llama", "b": true, "small": 32, "neg": -5, "f": 1.5}"#,
)
.expect("parse");
assert_eq!(meta.get("s"), Some(&V::String("llama".to_owned())));
assert_eq!(meta.get("b"), Some(&V::Bool(true)));
assert_eq!(meta.get("small"), Some(&V::U32(32)));
assert_eq!(meta.get("neg"), Some(&V::I64(-5)));
assert_eq!(meta.get("f"), Some(&V::F32(1.5)));
}
#[test]
fn plain_arrays_take_their_first_element_type() {
let meta =
parse_gguf_metadata_json(r#"{"toks": ["a", "b"], "ids": [1, 2]}"#).expect("parse");
assert_eq!(
meta.get("toks"),
Some(&V::Array(Box::new(A::String(vec![
"a".to_owned(),
"b".to_owned()
]))))
);
assert_eq!(
meta.get("ids"),
Some(&V::Array(Box::new(A::U32(vec![1, 2]))))
);
}
#[test]
fn explicit_form_pins_an_exact_width() {
let meta = parse_gguf_metadata_json(
r#"{"blocks": {"type": "u32", "value": 32}, "eps": {"type": "f32", "value": 1e-5}}"#,
)
.expect("parse");
assert_eq!(meta.get("blocks"), Some(&V::U32(32)));
assert_eq!(meta.get("eps"), Some(&V::F32(1e-5)));
}
#[test]
fn explicit_array_fixes_the_token_type_case() {
let inferred = parse_gguf_metadata_json(r#"{"tt": [1, 1, 2]}"#).expect("parse");
assert_eq!(
inferred.get("tt"),
Some(&V::Array(Box::new(A::U32(vec![1, 1, 2])))),
"inference alone yields U32 — the reason the escape hatch exists"
);
let explicit = parse_gguf_metadata_json(
r#"{"tt": {"type": "array", "item_type": "i32", "value": [1, 1, 2]}}"#,
)
.expect("parse");
assert_eq!(
explicit.get("tt"),
Some(&V::Array(Box::new(A::I32(vec![1, 1, 2]))))
);
}
#[test]
fn malformed_documents_are_rejected_with_the_key_named() {
assert!(parse_gguf_metadata_json("{not json").is_err());
assert!(parse_gguf_metadata_json("[1, 2]").is_err());
let err = parse_gguf_metadata_json(r#"{"empty": []}"#).unwrap_err();
assert!(err.to_string().contains("empty"), "got: {err}");
assert!(parse_gguf_metadata_json(r#"{"k": {"type": "u128", "value": 1}}"#).is_err());
let err = parse_gguf_metadata_json(r#"{"k": {"type": "u8", "value": 300}}"#).unwrap_err();
assert!(err.to_string().contains("out of range"), "got: {err}");
assert!(parse_gguf_metadata_json(r#"{"k": null}"#).is_err());
}
#[test]
fn kv_args_are_string_valued_and_split_on_the_first_equals() {
let (key, value) = parse_gguf_kv_arg("general.architecture=llama").expect("parse");
assert_eq!(key, "general.architecture");
assert_eq!(value, V::String("llama".to_owned()));
let (key, value) = parse_gguf_kv_arg("k=a=b").expect("parse");
assert_eq!(key, "k");
assert_eq!(value, V::String("a=b".to_owned()));
assert_eq!(
parse_gguf_kv_arg("n=32").expect("parse").1,
V::String("32".to_owned())
);
assert!(parse_gguf_kv_arg("no-equals").is_err());
assert!(parse_gguf_kv_arg("=empty-key").is_err());
}
}
#[cfg(all(test, feature = "bnb"))]
#[allow(clippy::panic, clippy::unwrap_used, clippy::expect_used)]
mod to_bf16_bytes_tests {
use super::{Dtype, to_bf16_bytes};
#[test]
fn f32_arm_rounds_to_nearest_even() {
let cases = [
(0x3F80_9000_u32, 0x3F81_u16, 0x3F80_u16),
(0x3F80_8000, 0x3F80, 0x3F80),
(0x3F81_8000, 0x3F82, 0x3F81),
];
let mut input = Vec::new();
for (bits, _, _) in cases {
input.extend_from_slice(&bits.to_le_bytes());
}
let out = to_bf16_bytes(&input, Dtype::F32, "w").expect("F32 narrowing");
for (i, (bits, want_rne, would_truncate)) in cases.into_iter().enumerate() {
#[allow(clippy::indexing_slicing)]
let got = u16::from_le_bytes([out[i * 2], out[i * 2 + 1]]);
assert_eq!(got, want_rne, "0x{bits:08X} should round to nearest even");
if want_rne != would_truncate {
assert_ne!(got, would_truncate, "0x{bits:08X} must not truncate");
}
}
}
#[test]
fn bf16_arm_borrows_unchanged() {
let input = vec![0x34_u8, 0x12, 0x78, 0x56];
let out = to_bf16_bytes(&input, Dtype::BF16, "w").expect("BF16 passthrough");
assert_eq!(&*out, &input[..]);
assert!(
matches!(out, std::borrow::Cow::Borrowed(_)),
"must not copy"
);
}
#[test]
fn f16_arm_rounds_to_nearest_even() {
let value = half::f16::from_bits(0x3C05);
let input = value.to_le_bytes();
let out = to_bf16_bytes(&input, Dtype::F16, "w").expect("F16 narrowing");
#[allow(clippy::indexing_slicing)]
let got = u16::from_le_bytes([out[0], out[1]]);
let widened = value.to_f32().to_bits();
#[allow(clippy::as_conversions, clippy::cast_possible_truncation)]
let truncated = (widened >> 16) as u16;
assert_eq!(got, crate::remember::fp8::f32_bits_to_bf16_bits(widened));
assert_ne!(got, truncated, "F16 arm must not truncate either");
}
}
#[cfg(test)]
#[allow(clippy::panic, clippy::unwrap_used, clippy::expect_used)]
mod stats_tests {
use super::{ConvertOptions, ConvertTarget, convert};
use std::path::Path;
#[test]
fn stats_partition_the_written_tensors() {
let input = Path::new("tests/fixtures/safetensors_reference/fp8.safetensors");
assert!(input.exists(), "committed FP8 fixture missing");
let dir = tempfile::tempdir().expect("tempdir");
let out = dir.path().join("out.safetensors");
let stats = convert(
input,
ConvertTarget::Safetensors,
&out,
&ConvertOptions::new(),
)
.expect("convert fp8 -> safetensors");
assert!(
stats.dequantized > 0,
"the FP8 fixture has a quantised weight: {stats:?}"
);
assert_eq!(stats.quantized, 0, "safetensors target quantises nothing");
assert_eq!(
stats.dequantized + stats.passthrough,
stats.tensors,
"dequantised + passthrough must equal the tensors written: {stats:?}"
);
}
}
#[cfg(all(test, feature = "gguf"))]
#[allow(
clippy::panic,
clippy::unwrap_used,
clippy::expect_used,
clippy::as_conversions,
clippy::cast_possible_truncation,
clippy::indexing_slicing,
// EXHAUSTIVE: `GgufType` is `#[non_exhaustive]`, so the fixture builder's
// dtype matches need a wildcard. It models only the three dtypes the
// fixture uses and panics on anything else — a test-authoring error, never
// reachable from input. Hoisted here because an arm-level `#[allow]` does
// not suppress `wildcard_enum_match_arm`.
clippy::wildcard_enum_match_arm
)]
mod quantized_gguf_tests {
use super::{
ConvertOptions, ConvertTarget, convert, derive_output_path, derive_output_path_for_dtype,
};
use crate::GgufType;
use crate::parallel::MIN_PARALLEL_BYTES;
use std::path::PathBuf;
const ALIGNMENT: usize = 32;
struct Spec {
name: &'static str,
dtype: GgufType,
shape: Vec<usize>,
}
impl Spec {
fn new(name: &'static str, dtype: GgufType, shape: &[usize]) -> Self {
Self {
name,
dtype,
shape: shape.to_vec(),
}
}
fn n_elements(&self) -> usize {
self.shape.iter().product()
}
fn byte_len(&self) -> usize {
let n = self.n_elements();
match self.dtype {
GgufType::F32 => n * 4,
GgufType::Q8_0 => (n / 32) * 34,
GgufType::Q4_K => (n / 256) * 144,
other => panic!("fixture builder does not model {other:?}"),
}
}
fn discriminant(&self) -> u32 {
match self.dtype {
GgufType::F32 => 0,
GgufType::Q8_0 => 8,
GgufType::Q4_K => 12,
other => panic!("fixture builder does not model {other:?}"),
}
}
fn data(&self) -> Vec<u8> {
let len = self.byte_len();
let mut buf: Vec<u8> = (0..len)
.map(|i| (i.wrapping_mul(2_654_435_761) & 0xFF) as u8)
.collect();
match self.dtype {
GgufType::Q8_0 => {
for block in buf.chunks_exact_mut(34) {
block[0] = 0x00;
block[1] = 0x3C; }
}
GgufType::Q4_K => {
for block in buf.chunks_exact_mut(144) {
block[0] = 0x00;
block[1] = 0x3C; block[2] = 0x00;
block[3] = 0x38; }
}
GgufType::F32 => {}
other => panic!("fixture builder does not model {other:?}"),
}
buf
}
}
fn fixture_specs() -> Vec<Spec> {
let mut specs = Vec::new();
for i in 0..8 {
let shape: Vec<usize> = if i == 0 {
vec![350, 400] } else {
vec![140_000]
};
specs.push(Spec::new(
match i {
0 => "blk.0.attn_norm.weight",
1 => "blk.1.attn_norm.weight",
2 => "blk.2.attn_norm.weight",
3 => "blk.3.attn_norm.weight",
4 => "blk.4.attn_norm.weight",
5 => "blk.5.attn_norm.weight",
6 => "blk.6.attn_norm.weight",
_ => "output_norm.weight",
},
GgufType::F32,
&shape,
));
}
specs.push(Spec::new("token_embd.weight", GgufType::Q8_0, &[512, 512]));
specs.push(Spec::new("blk.0.attn_q.weight", GgufType::Q8_0, &[32_768]));
specs.push(Spec::new("blk.1.attn_q.weight", GgufType::Q8_0, &[16_384]));
specs.push(Spec::new("blk.2.attn_q.weight", GgufType::Q8_0, &[8_192]));
specs.push(Spec::new("blk.3.attn_q.weight", GgufType::Q8_0, &[4_096]));
specs.push(Spec::new(
"blk.0.ffn_down.weight",
GgufType::Q4_K,
&[65_536],
));
specs.push(Spec::new(
"blk.1.ffn_down.weight",
GgufType::Q4_K,
&[32_768],
));
specs.push(Spec::new(
"blk.2.ffn_down.weight",
GgufType::Q4_K,
&[16_384],
));
specs.push(Spec::new("blk.3.ffn_down.weight", GgufType::Q4_K, &[8_192]));
specs
}
fn push_u32(buf: &mut Vec<u8>, v: u32) {
buf.extend_from_slice(&v.to_le_bytes());
}
fn push_u64(buf: &mut Vec<u8>, v: u64) {
buf.extend_from_slice(&v.to_le_bytes());
}
fn push_string(buf: &mut Vec<u8>, s: &str) {
push_u64(buf, s.len() as u64);
buf.extend_from_slice(s.as_bytes());
}
fn pad_to_alignment(buf: &mut Vec<u8>) {
while !buf.len().is_multiple_of(ALIGNMENT) {
buf.push(0);
}
}
fn build_quantized_gguf(specs: &[Spec]) -> Vec<u8> {
let mut buf = Vec::new();
buf.extend_from_slice(b"GGUF");
push_u32(&mut buf, 3); push_u64(&mut buf, specs.len() as u64);
push_u64(&mut buf, 2);
push_string(&mut buf, "general.architecture");
push_u32(&mut buf, 8);
push_string(&mut buf, "llama");
push_string(&mut buf, "general.alignment");
push_u32(&mut buf, 4);
push_u32(&mut buf, ALIGNMENT as u32);
let mut relative = 0usize;
let mut offsets = Vec::with_capacity(specs.len());
for spec in specs {
offsets.push(relative);
relative += spec.byte_len();
while !relative.is_multiple_of(ALIGNMENT) {
relative += 1;
}
}
for (spec, &offset) in specs.iter().zip(offsets.iter()) {
push_string(&mut buf, spec.name);
push_u32(&mut buf, spec.shape.len() as u32);
for &d in &spec.shape {
push_u64(&mut buf, d as u64);
}
push_u32(&mut buf, spec.discriminant());
push_u64(&mut buf, offset as u64);
}
pad_to_alignment(&mut buf);
let data_start = buf.len();
for (spec, &offset) in specs.iter().zip(offsets.iter()) {
debug_assert_eq!(buf.len() - data_start, offset, "offset table drift");
buf.extend_from_slice(&spec.data());
pad_to_alignment(&mut buf);
}
buf
}
fn assert_bytes_eq(actual: &[u8], expected: &[u8], context: &str) {
assert_eq!(
actual.len(),
expected.len(),
"{context}: length differs ({} vs {} bytes)",
actual.len(),
expected.len()
);
if let Some((offset, (a, e))) = actual
.iter()
.zip(expected.iter())
.enumerate()
.find(|(_, (a, e))| a != e)
.map(|(i, (a, e))| (i, (*a, *e)))
{
panic!(
"{context}: first difference at byte {offset} of {} (got {a:#04x}, expected {e:#04x})",
actual.len()
);
}
}
fn write_fixture() -> (tempfile::TempDir, PathBuf, Vec<Spec>) {
let specs = fixture_specs();
let bytes = build_quantized_gguf(&specs);
let dir = tempfile::tempdir().expect("tempdir");
let path = dir.path().join("quantized.gguf");
std::fs::write(&path, &bytes).expect("write fixture");
(dir, path, specs)
}
#[test]
fn fixture_crosses_the_parallel_threshold() {
let specs = fixture_specs();
let total: u64 = specs.iter().map(|s| s.byte_len() as u64).sum();
assert!(
total > MIN_PARALLEL_BYTES,
"fixture is {total} B but MIN_PARALLEL_BYTES is {MIN_PARALLEL_BYTES} B — \
the determinism tests would not exercise the parallel path"
);
assert_eq!(specs.len(), 17, "a prime tensor count is deliberate");
}
#[test]
fn quantized_gguf_converts_to_expected_bf16() {
let (_dir, path, specs) = write_fixture();
let dir_out = tempfile::tempdir().expect("tempdir");
let out = dir_out.path().join("out.safetensors");
let stats = convert(
&path,
ConvertTarget::Safetensors,
&out,
&ConvertOptions::new().with_threads(4),
)
.expect("convert quantised gguf -> safetensors");
assert_eq!(stats.tensors, 17);
assert_eq!(stats.dequantized, 9, "5 Q8_0 + 4 Q4_K are dequantised");
assert_eq!(stats.passthrough, 8, "the 8 F32 tensors pass through");
let written = std::fs::read(&out).expect("read output");
let tensors = safetensors::SafeTensors::deserialize(&written).expect("parse output");
for spec in &specs {
let view = tensors.tensor(spec.name).expect("tensor present in output");
let expected_shape: Vec<usize> = spec.shape.iter().copied().rev().collect();
assert_eq!(view.shape(), expected_shape.as_slice(), "{}", spec.name);
match spec.dtype {
GgufType::F32 => {
assert_eq!(view.dtype(), safetensors::Dtype::F32, "{}", spec.name);
assert_bytes_eq(view.data(), &spec.data(), spec.name);
}
dtype => {
assert_eq!(view.dtype(), safetensors::Dtype::BF16, "{}", spec.name);
let expected =
crate::dequantize_gguf_to_bf16(&spec.data(), dtype, spec.n_elements())
.expect("oracle dequant");
assert_bytes_eq(
view.data(),
&expected,
&format!("{} vs the kernel called directly", spec.name),
);
}
}
}
}
#[test]
fn gguf_to_safetensors_deterministic_across_thread_counts() {
let (_dir, path, _specs) = write_fixture();
let dir_out = tempfile::tempdir().expect("tempdir");
let baseline_path = dir_out.path().join("baseline.safetensors");
convert(
&path,
ConvertTarget::Safetensors,
&baseline_path,
&ConvertOptions::new().with_threads(1),
)
.expect("baseline convert");
let baseline = std::fs::read(&baseline_path).expect("read baseline");
for n in [1usize, 2, 4, 8, 16] {
let out = dir_out.path().join(format!("t{n}.safetensors"));
convert(
&path,
ConvertTarget::Safetensors,
&out,
&ConvertOptions::new().with_threads(n),
)
.expect("threaded convert");
assert_bytes_eq(
&std::fs::read(&out).expect("read output"),
&baseline,
&format!("safetensors output at {n} threads"),
);
}
let default_out = dir_out.path().join("default.safetensors");
convert(
&path,
ConvertTarget::Safetensors,
&default_out,
&ConvertOptions::new(),
)
.expect("default convert");
assert_bytes_eq(
&std::fs::read(&default_out).expect("read output"),
&baseline,
"the default thread budget vs the sequential baseline",
);
}
#[test]
fn gguf_to_other_targets_deterministic_across_thread_counts() {
let (_dir, path, _specs) = write_fixture();
let dir_out = tempfile::tempdir().expect("tempdir");
let targets: &[(ConvertTarget, &str)] = &[
(ConvertTarget::Gguf, "gguf"),
#[cfg(feature = "bnb")]
(ConvertTarget::BnbNf4, "safetensors"),
];
for &(target, ext) in targets {
let baseline_path = dir_out.path().join(format!("baseline-{ext}.{ext}"));
convert(
&path,
target,
&baseline_path,
&ConvertOptions::new().with_threads(1),
)
.expect("baseline convert");
let baseline = std::fs::read(&baseline_path).expect("read baseline");
for n in [2usize, 4, 8] {
let out = dir_out.path().join(format!("t{n}-{ext}.{ext}"));
convert(&path, target, &out, &ConvertOptions::new().with_threads(n))
.expect("threaded convert");
assert_bytes_eq(
&std::fs::read(&out).expect("read output"),
&baseline,
&format!("{target:?} output at {n} threads"),
);
}
}
}
#[test]
fn convert_honours_every_output_dtype_end_to_end() {
for (requested, expected_st) in [
(crate::Dtype::BF16, safetensors::Dtype::BF16),
(crate::Dtype::F32, safetensors::Dtype::F32),
(crate::Dtype::F16, safetensors::Dtype::F16),
] {
let (_dir, path, specs) = write_fixture();
let dir_out = tempfile::tempdir().expect("tempdir");
let out = dir_out.path().join("out.safetensors");
let stats = convert(
&path,
ConvertTarget::Safetensors,
&out,
&ConvertOptions::new()
.with_threads(4)
.with_output_dtype(requested),
)
.unwrap_or_else(|e| panic!("convert at {requested}: {e}"));
assert_eq!(stats.dequantized, 9, "{requested}");
let written = std::fs::read(&out).expect("read output");
let tensors = safetensors::SafeTensors::deserialize(&written).expect("parse output");
for spec in &specs {
let view = tensors.tensor(spec.name).expect("tensor present");
match spec.dtype {
GgufType::F32 => {
assert_eq!(
view.dtype(),
safetensors::Dtype::F32,
"{} passthrough must ignore --out-dtype {requested}",
spec.name
);
assert_bytes_eq(view.data(), &spec.data(), spec.name);
}
dtype => {
assert_eq!(view.dtype(), expected_st, "{}", spec.name);
let expected = match requested {
crate::Dtype::BF16 => crate::dequantize_gguf::<crate::Bf16Out>(
&spec.data(),
dtype,
spec.n_elements(),
),
crate::Dtype::F32 => crate::dequantize_gguf::<crate::F32Out>(
&spec.data(),
dtype,
spec.n_elements(),
),
_ => crate::dequantize_gguf::<crate::F16Out>(
&spec.data(),
dtype,
spec.n_elements(),
),
}
.expect("oracle dequant");
assert_bytes_eq(
view.data(),
&expected,
&format!("{} at {requested}", spec.name),
);
}
}
}
}
}
#[test]
fn output_dtype_changes_the_dequantised_payload_width() {
let mut payloads = Vec::new();
for requested in [crate::Dtype::BF16, crate::Dtype::F16, crate::Dtype::F32] {
let (_dir, path, specs) = write_fixture();
let dir_out = tempfile::tempdir().expect("tempdir");
let out = dir_out.path().join("out.safetensors");
convert(
&path,
ConvertTarget::Safetensors,
&out,
&ConvertOptions::new().with_output_dtype(requested),
)
.expect("convert");
let written = std::fs::read(&out).expect("read output");
let tensors = safetensors::SafeTensors::deserialize(&written).expect("parse output");
let dequantised: usize = specs
.iter()
.filter(|s| s.dtype != GgufType::F32)
.map(|s| tensors.tensor(s.name).expect("tensor present").data().len())
.sum();
payloads.push(dequantised);
}
assert_eq!(
payloads[0], payloads[1],
"BF16 and F16 are both 2 bytes per element"
);
assert_eq!(
payloads[2],
payloads[0] * 2,
"F32 is exactly twice BF16: {payloads:?}"
);
}
#[test]
fn output_dtype_is_deterministic_across_thread_counts() {
for requested in [crate::Dtype::BF16, crate::Dtype::F32, crate::Dtype::F16] {
let (_dir, path, _specs) = write_fixture();
let dir_out = tempfile::tempdir().expect("tempdir");
let mut baseline: Option<Vec<u8>> = None;
for threads in [1usize, 2, 4, 8] {
let out = dir_out.path().join(format!("out-{requested}-{threads}.st"));
convert(
&path,
ConvertTarget::Safetensors,
&out,
&ConvertOptions::new()
.with_threads(threads)
.with_output_dtype(requested),
)
.expect("convert");
let bytes = std::fs::read(&out).expect("read output");
match &baseline {
None => baseline = Some(bytes),
Some(expected) => assert_bytes_eq(
&bytes,
expected,
&format!("{requested} at {threads} threads vs 1 thread"),
),
}
}
}
}
#[test]
fn derived_output_path_names_the_dtype_it_holds() {
let input = std::path::Path::new("/models/smollm2.gguf");
for (dtype, expected) in [
(crate::Dtype::BF16, "smollm2-bf16.safetensors"),
(crate::Dtype::F32, "smollm2-f32.safetensors"),
(crate::Dtype::F16, "smollm2-f16.safetensors"),
] {
let path = derive_output_path_for_dtype(input, ConvertTarget::Safetensors, dtype);
assert_eq!(
path.file_name().and_then(|s| s.to_str()),
Some(expected),
"safetensors target at {dtype}"
);
}
assert_eq!(
derive_output_path(input, ConvertTarget::Safetensors)
.file_name()
.and_then(|s| s.to_str()),
Some("smollm2-bf16.safetensors"),
);
for dtype in [crate::Dtype::BF16, crate::Dtype::F32] {
assert_eq!(
derive_output_path_for_dtype(input, ConvertTarget::Gguf, dtype)
.file_name()
.and_then(|s| s.to_str()),
Some("smollm2-gguf.gguf"),
"gguf target must ignore the dequant dtype"
);
}
}
#[test]
fn unsupported_output_dtype_is_rejected() {
let (_dir, path, _specs) = write_fixture();
let dir_out = tempfile::tempdir().expect("tempdir");
let out = dir_out.path().join("out.safetensors");
let err = convert(
&path,
ConvertTarget::Safetensors,
&out,
&ConvertOptions::new().with_output_dtype(crate::Dtype::I64),
)
.expect_err("I64 is not an output width");
let msg = err.to_string();
assert!(
msg.contains("bf16, f32, f16"),
"error should list the supported widths, got: {msg}"
);
}
}
#[cfg(all(test, feature = "gguf"))]
#[allow(
clippy::panic,
clippy::unwrap_used,
clippy::expect_used,
clippy::as_conversions,
clippy::cast_precision_loss,
clippy::indexing_slicing
)]
mod hub_scaling_bench {
use super::{ConvertOptions, read_hub};
use std::path::{Path, PathBuf};
use std::time::Instant;
const SAMPLES: usize = 5;
fn model_path(file_name: &str) -> PathBuf {
Path::new(env!("CARGO_MANIFEST_DIR"))
.join("tests")
.join("fixtures")
.join("gguf_reference")
.join("models")
.join(file_name)
}
fn python_baseline(model: &str) -> Option<(f64, String, String)> {
let stem = model.strip_suffix(".gguf").unwrap_or(model);
let sidecar = Path::new(env!("CARGO_MANIFEST_DIR"))
.join("tests")
.join("fixtures")
.join("gguf_reference")
.join(format!("{stem}.dequant.timing.json"));
let raw = std::fs::read_to_string(sidecar).ok()?;
let value: serde_json::Value = serde_json::from_str(&raw).ok()?;
let seconds = value.get("py_seconds")?.as_f64()?;
let library = value.get("py_library")?.as_str()?.to_owned();
let note = value
.get("note")
.and_then(serde_json::Value::as_str)
.unwrap_or("")
.to_owned();
Some((seconds, library, note))
}
fn hub_scaling_for(model: &str) {
let input = model_path(model);
if !input.exists() {
eprintln!("SKIP {model}: fixture absent (gitignored)");
return;
}
let input_mib = std::fs::metadata(&input).expect("stat").len() as f64 / (1024.0 * 1024.0);
eprintln!("\n=== read_hub({model}) — {input_mib:.1} MiB input ===");
let python = python_baseline(model);
let mut baseline = 0.0_f64;
for threads in [1usize, 2, 4, 8, 16] {
let options = ConvertOptions::new().with_threads(threads);
drop(read_hub(&input, &options).expect("warm-up read_hub"));
let mut samples: Vec<f64> = Vec::with_capacity(SAMPLES);
for _ in 0..SAMPLES {
let start = Instant::now();
let hub = read_hub(&input, &options).expect("read_hub");
samples.push(start.elapsed().as_secs_f64() * 1000.0);
drop(hub);
}
samples.sort_by(|a, b| a.partial_cmp(b).unwrap());
let median = samples[SAMPLES / 2];
if threads == 1 {
baseline = median;
}
let vs_python = python.as_ref().map_or_else(String::new, |&(py, _, _)| {
format!(" | {:.1}x vs gguf-py", (py * 1000.0) / median)
});
eprintln!(
"{threads:>2} threads: median {median:>8.2} ms (min {:.2}, max {:.2}) -> {:.2}x{vs_python}",
samples[0],
samples[SAMPLES - 1],
baseline / median
);
}
match python {
Some((py, library, note)) => {
eprintln!("\npython baseline: {:.1} ms ({library})", py * 1000.0);
if !note.is_empty() {
eprintln!(" caveat: {note}");
}
}
None => {
eprintln!("\npython baseline: no sidecar (run generate_gguf_dequant_timings.py)");
}
}
}
#[test]
#[ignore = "ad-hoc measurement; run explicitly with --ignored"]
fn hub_scaling_smollm2_q4_k_m() {
hub_scaling_for("SmolLM2-135M-Instruct-Q4_K_M.gguf");
}
#[test]
#[ignore = "ad-hoc measurement; run explicitly with --ignored"]
fn hub_scaling_tinyllama_q5_0() {
hub_scaling_for("tinyllama-1.1b-chat-v1.0.Q5_0.gguf");
}
}