#![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, ComputeStorage};
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()
}
#[inline(always)]
fn scale_row_q2f(wide_dot: ComputeStorage, s_q32: u64) -> ComputeStorage {
#[cfg(table_format = "q16_16")]
{
let scaled = (wide_dot as i128 * s_q32 as i128) >> 32;
if scaled > i64::MAX as i128 || scaled < i64::MIN as i128 {
panic!("RowScaledTQ19: q2f scaled output exceeds compute range");
}
scaled as i64
}
#[cfg(table_format = "q32_32")]
{
let h = wide_dot >> 32;
let l = (wide_dot & 0xFFFF_FFFF) as i128;
let low = (l * s_q32 as i128) >> 32;
match h.checked_mul(s_q32 as i128).and_then(|hs| hs.checked_add(low)) {
Some(v) => v,
None => panic!("RowScaledTQ19: q2f scaled output exceeds compute range"),
}
}
}
pub fn matvec_q2f(&self, activations: &[BinaryStorage]) -> Vec<ComputeStorage> {
assert_eq!(activations.len(), self.cols, "RowScaledTQ19::matvec_q2f: activation length mismatch");
(0..self.rows)
.map(|r| {
let row = &self.data[r * self.cols..(r + 1) * self.cols];
Self::scale_row_q2f(ops::tq19_dot_q2f(row, activations), self.scales_q32[r])
})
.collect()
}
pub fn matvec_q2f_par(&self, activations: &[BinaryStorage]) -> Vec<ComputeStorage> {
assert_eq!(activations.len(), self.cols, "RowScaledTQ19::matvec_q2f_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_q2f(ops::tq19_dot_q2f(row, activations), self.scales_q32[r])
})
.collect()
}
pub fn matvec_q2f_batch_par(&self, batch: &[&[BinaryStorage]]) -> Vec<Vec<ComputeStorage>> {
for (i, v) in batch.iter().enumerate() {
assert_eq!(v.len(), self.cols, "RowScaledTQ19::matvec_q2f_batch_par: activation[{i}] length mismatch");
}
let per_row: Vec<Vec<ComputeStorage>> = (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_q2f(ops::tq19_dot_q2f(row, x), s))
.collect()
})
.collect();
(0..batch.len())
.map(|b| per_row.iter().map(|row| row[b]).collect())
.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));
}
fn narrow_q2f(v: ComputeStorage) -> BinaryStorage {
#[cfg(table_format = "q16_16")]
{ (v / (1i64 << crate::fixed_point::frac_config::FRAC_BITS)) as i32 }
#[cfg(table_format = "q32_32")]
{ (v / (1i128 << 32)) as i64 }
}
#[test]
fn q2f_unit_scale_exact_general_within_one_lsb_and_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 x: Vec<BinaryStorage> = (0..cols)
.map(|i| ((i as i64 * 69621 % 4001) - 2000) as BinaryStorage)
.collect();
let unit = RowScaledTQ19::from_parts(rows, cols, data.clone(), vec![1u64 << 32; rows]);
let narrow = unit.matvec(&x);
let wide = unit.matvec_q2f(&x);
for r in 0..rows {
assert_eq!(narrow_q2f(wide[r]), narrow[r], "unit-scale row {r}");
}
let scales: Vec<u64> = vec![
1u64 << 31,
1u64 << 32,
(1u64 << 32) + (1u64 << 30) + 12345,
1u64 << 20,
(1u64 << 32) + 1,
];
let rs = RowScaledTQ19::from_parts(rows, cols, data.clone(), scales.clone());
let narrow = rs.matvec(&x);
let wide = rs.matvec_q2f(&x);
#[cfg(table_format = "q16_16")]
let f = crate::fixed_point::frac_config::FRAC_BITS;
#[cfg(table_format = "q32_32")]
let f = 32u32;
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 wide_dot = (acc << f) / SCALE as i128;
let expected = (wide_dot * scales[r] as i128) >> 32;
assert_eq!(wide[r] as i128, expected, "row {r} diverged from i128 oracle");
let diff = (narrow_q2f(wide[r]) as i128 - narrow[r] as i128).abs();
assert!(diff <= 1, "row {r}: narrow(q2f) off by {diff} LSB");
}
assert_eq!(wide, rs.matvec_q2f_par(&x));
let batch: Vec<&[BinaryStorage]> = vec![&x];
assert_eq!(rs.matvec_q2f_batch_par(&batch)[0], wide);
}
#[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);
}
}