use kopitiam_core::{Error, Result, Shape};
use super::Tensor;
impl Tensor {
pub fn tessdata_int8_to_f32(
weights: &[i8],
scales: &[f32],
rows: usize,
cols: usize,
) -> Result<Tensor> {
if weights.len() != rows * cols {
return Err(Error::ShapeMismatch {
expected: Shape::new([rows, cols]),
actual: Shape::new([weights.len()]),
});
}
if scales.len() != rows {
return Err(Error::ShapeMismatch {
expected: Shape::new([rows]),
actual: Shape::new([scales.len()]),
});
}
let mut out = vec![0f32; rows * cols];
for i in 0..rows {
let scale = scales[i];
for j in 0..cols {
out[i * cols + j] = weights[i * cols + j] as f32 * scale;
}
}
Tensor::from_f32(out, Shape::new([rows, cols]))
}
}
#[cfg(test)]
mod tests {
use super::*;
const INT8_MAX: f32 = 127.0;
fn encode_row(row: &[f32]) -> (Vec<i8>, f32) {
let max_abs = row.iter().fold(0f32, |m, &v| m.max(v.abs()));
let scale = max_abs / INT8_MAX; let div = if scale == 0.0 { 1.0 } else { scale };
let q: Vec<i8> = row
.iter()
.map(|&w| (w / div).round().clamp(-INT8_MAX, INT8_MAX) as i8)
.collect();
(q, scale)
}
#[test]
fn decode_round_trips_a_known_row_within_quantization_error() {
let row = vec![0.10, -0.40, 0.25, 0.80, -0.80, 0.05, 0.55, -0.30];
let (q, scale) = encode_row(&row);
let decoded = Tensor::tessdata_int8_to_f32(&q, &[scale], 1, row.len())
.unwrap()
.to_vec_f32()
.unwrap();
let half_step = scale / 2.0;
for (got, want) in decoded.iter().zip(&row) {
assert!(
(got - want).abs() <= half_step + 1e-7,
"decoded {got} too far from {want} (half-step {half_step})"
);
}
assert!((decoded[3] - 0.80).abs() < 1e-6);
assert!((decoded[4] + 0.80).abs() < 1e-6);
}
#[test]
fn decode_recovers_an_exact_multiple_row_bit_for_bit() {
let scale = 1.0 / INT8_MAX;
let q: Vec<i8> = vec![-127, -64, 0, 33, 127];
let expected: Vec<f32> = q.iter().map(|&k| k as f32 * scale).collect();
let decoded = Tensor::tessdata_int8_to_f32(&q, &[scale], 1, q.len())
.unwrap()
.to_vec_f32()
.unwrap();
assert_eq!(decoded, expected);
}
#[test]
fn decode_uses_an_independent_scale_per_row() {
let weights: Vec<i8> = vec![10, -20, 30, 1, -2, 3];
let scales = [0.5f32, 2.0f32];
let decoded = Tensor::tessdata_int8_to_f32(&weights, &scales, 2, 3)
.unwrap()
.to_vec_f32()
.unwrap();
assert_eq!(decoded, vec![5.0, -10.0, 15.0, 2.0, -4.0, 6.0]);
}
#[test]
fn decoded_weight_matmul_matches_a_known_small_matrix() {
let weights: Vec<i8> = vec![10, -20, 30, 2, -4, 6];
let scales = [0.5f32, 1.0f32];
let w = Tensor::tessdata_int8_to_f32(&weights, &scales, 2, 3).unwrap();
let x = Tensor::from_f32(vec![1.0, 2.0, 0.0, 1.0, 1.0, 0.0], [3, 2]).unwrap();
let out = w.matmul(&x).unwrap();
assert_eq!(out.shape().dims(), &[2, 2]);
assert_eq!(out.to_vec_f32().unwrap(), vec![20.0, 0.0, 8.0, 0.0]);
}
#[test]
fn decode_rejects_a_weight_length_mismatch() {
let weights = vec![1i8, 2, 3];
assert!(matches!(
Tensor::tessdata_int8_to_f32(&weights, &[1.0, 1.0], 2, 3),
Err(Error::ShapeMismatch { .. })
));
}
#[test]
fn decode_rejects_a_scale_count_mismatch() {
let weights = vec![1i8, 2, 3, 4, 5, 6];
assert!(matches!(
Tensor::tessdata_int8_to_f32(&weights, &[1.0], 2, 3),
Err(Error::ShapeMismatch { .. })
));
}
#[test]
fn decode_handles_an_all_zero_row() {
let row = vec![0.0f32; 4];
let (q, scale) = encode_row(&row);
assert_eq!(scale, 0.0);
let decoded = Tensor::tessdata_int8_to_f32(&q, &[scale], 1, 4)
.unwrap()
.to_vec_f32()
.unwrap();
assert_eq!(decoded, vec![0.0; 4]);
}
}