use kopitiam_core::{DType, Error, Result};
use kopitiam_loader::{LoadedModel, TensorEntry};
use kopitiam_tensor::Tensor;
pub fn tensor_from_entry(model: &LoadedModel, entry: &TensorEntry) -> Result<Tensor> {
let bytes = model.tensor_bytes(&entry.name)?;
match entry.dtype {
DType::F32 => Tensor::from_f32(read_f32_le(bytes), entry.shape.clone()),
DType::F16 => Tensor::from_f16(read_u16_le(bytes), entry.shape.clone()),
DType::BF16 => Tensor::from_bf16(read_u16_le(bytes), entry.shape.clone()),
DType::I8 => Tensor::from_i8(bytes.iter().map(|&b| b as i8).collect(), entry.shape.clone()),
DType::I32 => Tensor::from_i32(read_i32_le(bytes), entry.shape.clone()),
DType::Q4_0
| DType::Q4_1
| DType::Q5_0
| DType::Q5_1
| DType::Q8_0
| DType::Q2_K
| DType::Q3_K
| DType::Q4_K
| DType::Q5_K
| DType::Q6_K
| DType::Q8_K => Tensor::from_quantized(entry.dtype, bytes.to_vec(), entry.shape.clone()),
}
}
pub fn load_tensor_f32(model: &LoadedModel, name: &str) -> Result<Tensor> {
let entry = model
.tensor(name)
.ok_or_else(|| Error::MissingTensor { name: name.to_string() })?;
tensor_from_entry(model, entry)?.to_dtype(DType::F32)
}
pub fn load_tensor_f32_opt(model: &LoadedModel, name: &str) -> Result<Option<Tensor>> {
match model.tensor(name) {
Some(entry) => Ok(Some(tensor_from_entry(model, entry)?.to_dtype(DType::F32)?)),
None => Ok(None),
}
}
pub fn load_matmul_weight(model: &LoadedModel, name: &str) -> Result<Tensor> {
let entry = model
.tensor(name)
.ok_or_else(|| Error::MissingTensor { name: name.to_string() })?;
let tensor = tensor_from_entry(model, entry)?;
keep_quantized_or_dequantize(tensor)
}
pub fn load_matmul_weight_opt(model: &LoadedModel, name: &str) -> Result<Option<Tensor>> {
match model.tensor(name) {
Some(entry) => {
let tensor = tensor_from_entry(model, entry)?;
Ok(Some(keep_quantized_or_dequantize(tensor)?))
}
None => Ok(None),
}
}
fn keep_quantized_or_dequantize(tensor: Tensor) -> Result<Tensor> {
if kopitiam_tensor::has_fused_matmul_kernel(tensor.dtype()) {
Ok(tensor)
} else {
tensor.to_dtype(DType::F32)
}
}
fn read_f32_le(bytes: &[u8]) -> Vec<f32> {
bytes.chunks_exact(4).map(|c| f32::from_le_bytes(c.try_into().expect("chunks_exact(4)"))).collect()
}
fn read_u16_le(bytes: &[u8]) -> Vec<u16> {
bytes.chunks_exact(2).map(|c| u16::from_le_bytes(c.try_into().expect("chunks_exact(2)"))).collect()
}
fn read_i32_le(bytes: &[u8]) -> Vec<i32> {
bytes.chunks_exact(4).map(|c| i32::from_le_bytes(c.try_into().expect("chunks_exact(4)"))).collect()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_support::synthetic_gguf::{tiny_model_bytes, write_temp_gguf};
#[test]
fn tensor_from_entry_round_trips_f32_bytes() {
let bytes = tiny_model_bytes();
let path = write_temp_gguf(&bytes, "bridge-round-trip");
let model = kopitiam_loader::load_model(&path).unwrap();
let entry = model.tensor("token_embd.weight").unwrap().clone();
let t = tensor_from_entry(&model, &entry).unwrap();
assert_eq!(t.dtype(), DType::F32);
assert_eq!(t.shape(), &entry.shape);
}
#[test]
fn load_tensor_f32_opt_is_none_for_a_missing_name() {
let bytes = tiny_model_bytes();
let path = write_temp_gguf(&bytes, "bridge-opt-missing");
let model = kopitiam_loader::load_model(&path).unwrap();
assert!(load_tensor_f32_opt(&model, "does.not.exist").unwrap().is_none());
}
#[test]
fn load_tensor_f32_errors_on_a_missing_required_name() {
let bytes = tiny_model_bytes();
let path = write_temp_gguf(&bytes, "bridge-required-missing");
let model = kopitiam_loader::load_model(&path).unwrap();
assert!(matches!(
load_tensor_f32(&model, "does.not.exist"),
Err(Error::MissingTensor { .. })
));
}
#[test]
fn a_q4_k_matmul_weight_is_dequantized_at_load_and_completes_a_forward_pass() {
use crate::test_support::synthetic_gguf::{
arbitrary_q4_k_blocks, single_quantized_weight_gguf, GGML_TYPE_Q4_K,
};
let (out_features, in_features) = (4, 256);
let bytes = arbitrary_q4_k_blocks(out_features, 0xBEEF);
let gguf = single_quantized_weight_gguf("test.weight", &[out_features, in_features], GGML_TYPE_Q4_K, &bytes);
let path = write_temp_gguf(&gguf, "q4k-load");
let model = kopitiam_loader::load_model(&path).unwrap();
let w = load_matmul_weight(&model, "test.weight").unwrap();
assert_eq!(w.dtype(), DType::F32);
assert!(!w.dtype().is_quantized());
assert!(!kopitiam_tensor::has_fused_matmul_kernel(DType::Q4_K));
let entry = model.tensor("test.weight").unwrap().clone();
let reference_w = tensor_from_entry(&model, &entry).unwrap().to_dtype(DType::F32).unwrap();
assert_eq!(w.to_vec_f32().unwrap(), reference_w.to_vec_f32().unwrap());
let x_vals: Vec<f32> = (0..in_features).map(|i| (i as f32 % 7.0 - 3.0) * 0.1).collect();
let x = Tensor::from_f32(x_vals, [1, in_features]).unwrap();
let y = crate::linear::linear(&x, &w, None).unwrap();
assert_eq!(y.dtype(), DType::F32);
let y_vals = y.to_vec_f32().unwrap();
assert_eq!(y_vals.len(), out_features);
assert!(y_vals.iter().all(|v| v.is_finite()), "forward-pass output must be finite: {y_vals:?}");
let reference = x.matmul(&reference_w.transpose(0, 1).unwrap()).unwrap().to_vec_f32().unwrap();
assert_eq!(y_vals, reference);
}
#[test]
fn a_q4_0_matmul_weight_stays_quantized_at_load_and_uses_the_fused_kernel() {
use crate::test_support::synthetic_gguf::{
quantize_q4_0_blocks, single_quantized_weight_gguf, GGML_TYPE_Q4_0,
};
let (out_features, in_features) = (3, 64); let data: Vec<f32> = (0..out_features * in_features).map(|i| (i as f32 * 0.037).sin() * 0.5).collect();
let bytes = quantize_q4_0_blocks(&data);
let gguf = single_quantized_weight_gguf("test.weight", &[out_features, in_features], GGML_TYPE_Q4_0, &bytes);
let path = write_temp_gguf(&gguf, "q40-load");
let model = kopitiam_loader::load_model(&path).unwrap();
let w = load_matmul_weight(&model, "test.weight").unwrap();
assert_eq!(w.dtype(), DType::Q4_0, "Q4_0 must stay quantized for the fused kernel");
assert!(w.dtype().is_quantized());
assert!(kopitiam_tensor::has_fused_matmul_kernel(w.dtype()));
let x_vals: Vec<f32> = (0..in_features).map(|i| (i as f32 % 5.0 - 2.0) * 0.1).collect();
let x = Tensor::from_f32(x_vals, [1, in_features]).unwrap();
let y = crate::linear::linear(&x, &w, None).unwrap();
let fused = x.quantized_matmul(&w).unwrap();
assert_eq!(y.to_vec_f32().unwrap(), fused.to_vec_f32().unwrap());
assert!(y.to_vec_f32().unwrap().iter().all(|v| v.is_finite()));
}
}