#![cfg(any(table_format = "q16_16", table_format = "q32_32"))]
use super::ops;
use super::MAX_RAW;
use crate::fixed_point::universal::fasc::stack_evaluator::BinaryStorage;
use rayon::prelude::*;
#[derive(Clone)]
pub struct RowScaledTQ19 {
rows: usize,
cols: usize,
data: Vec<i16>,
scales_q32: Vec<u64>,
}
impl RowScaledTQ19 {
pub fn from_parts(rows: usize, cols: usize, data: Vec<i16>, scales_q32: Vec<u64>) -> Self {
assert_eq!(data.len(), rows * cols, "RowScaledTQ19: data length mismatch");
assert_eq!(scales_q32.len(), rows, "RowScaledTQ19: scales length mismatch");
debug_assert!(data.iter().all(|&w| (w as i32).abs() <= MAX_RAW as i32));
Self { rows, cols, data, scales_q32 }
}
pub fn rows(&self) -> usize { self.rows }
pub fn cols(&self) -> usize { self.cols }
pub fn data(&self) -> &[i16] { &self.data }
pub fn scales_q32(&self) -> &[u64] { &self.scales_q32 }
pub fn size_bytes(&self) -> usize {
self.data.len() * 2 + self.scales_q32.len() * 8
}
#[inline(always)]
fn scale_row(dot: BinaryStorage, s_q32: u64) -> BinaryStorage {
let scaled = (dot as i128 * s_q32 as i128) >> 32;
if scaled > BinaryStorage::MAX as i128 || scaled < BinaryStorage::MIN as i128 {
panic!("RowScaledTQ19: scaled output exceeds storage range");
}
scaled as BinaryStorage
}
pub fn matvec(&self, activations: &[BinaryStorage]) -> Vec<BinaryStorage> {
assert_eq!(activations.len(), self.cols, "RowScaledTQ19::matvec: activation length mismatch");
(0..self.rows)
.map(|r| {
let row = &self.data[r * self.cols..(r + 1) * self.cols];
Self::scale_row(ops::tq19_dot(row, activations), self.scales_q32[r])
})
.collect()
}
pub fn matvec_par(&self, activations: &[BinaryStorage]) -> Vec<BinaryStorage> {
assert_eq!(activations.len(), self.cols, "RowScaledTQ19::matvec_par: activation length mismatch");
(0..self.rows)
.into_par_iter()
.map(|r| {
let row = &self.data[r * self.cols..(r + 1) * self.cols];
Self::scale_row(ops::tq19_dot(row, activations), self.scales_q32[r])
})
.collect()
}
pub fn matvec_batch_par(&self, batch: &[&[BinaryStorage]]) -> Vec<Vec<BinaryStorage>> {
for (i, v) in batch.iter().enumerate() {
assert_eq!(v.len(), self.cols, "RowScaledTQ19::matvec_batch_par: activation[{i}] length mismatch");
}
let per_row: Vec<Vec<BinaryStorage>> = (0..self.rows)
.into_par_iter()
.map(|r| {
let row = &self.data[r * self.cols..(r + 1) * self.cols];
let s = self.scales_q32[r];
batch
.iter()
.map(|x| Self::scale_row(ops::tq19_dot(row, x), s))
.collect()
})
.collect();
(0..batch.len())
.map(|b| per_row.iter().map(|row| row[b]).collect())
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tq19::{TQ19Matrix, SCALE};
#[test]
fn unit_scale_matches_plain_tq19() {
let rows = 7;
let cols = 64;
let data: Vec<i16> = (0..rows * cols)
.map(|i| ((i as i64 * 2654435761 % 59049) - 29524) as i16)
.collect();
let tq = TQ19Matrix::new(rows, cols, data.clone());
let rs = RowScaledTQ19::from_parts(rows, cols, data, vec![1u64 << 32; rows]);
let x: Vec<BinaryStorage> = (0..cols)
.map(|i| ((i as i64 * 40503 % 2049) - 1024) as BinaryStorage)
.collect();
assert_eq!(tq.matvec(&x), rs.matvec(&x));
}
#[test]
fn rowscaled_matvec_matches_i128_oracle() {
let rows = 5;
let cols = 97; let data: Vec<i16> = (0..rows * cols)
.map(|i| ((i as i64 * 48271 % 59049) - 29524) as i16)
.collect();
let scales: Vec<u64> = vec![
1u64 << 31,
1u64 << 32,
(1u64 << 32) + (1u64 << 30) + 12345,
1u64 << 20,
(1u64 << 32) + 1,
];
let x: Vec<BinaryStorage> = (0..cols)
.map(|i| ((i as i64 * 69621 % 4001) - 2000) as BinaryStorage)
.collect();
let rs = RowScaledTQ19::from_parts(rows, cols, data.clone(), scales.clone());
let got = rs.matvec(&x);
for r in 0..rows {
let mut acc: i128 = 0;
for c in 0..cols {
acc += data[r * cols + c] as i128 * x[c] as i128;
}
let dot = acc / SCALE as i128;
let expected = ((dot * scales[r] as i128) >> 32) as BinaryStorage;
assert_eq!(got[r], expected, "row {} diverged from i128 oracle", r);
}
assert_eq!(got, rs.matvec_par(&x));
}
#[test]
fn halved_scale_doubles_resolution() {
let cols = 32;
let data: Vec<i16> = vec![6; cols];
let rs = RowScaledTQ19::from_parts(1, cols, data, vec![1u64 << 31]);
let x: Vec<BinaryStorage> = vec![SCALE as BinaryStorage; cols];
assert_eq!(rs.matvec(&x)[0], 3 * cols as BinaryStorage);
}
}