#[derive(Debug, Clone, PartialEq, Default)]
pub struct Mat {
pub rows: usize,
pub cols: usize,
pub data: Vec<f32>,
}
fn shape_len(context: &str, rows: usize, cols: usize) -> usize {
let len = rows.checked_mul(cols);
assert!(
len.is_some(),
"{context}: rows*cols overflow ({rows} * {cols})"
);
len.unwrap_or(0)
}
impl Mat {
#[must_use]
pub fn from_vec(rows: usize, cols: usize, data: Vec<f32>) -> Self {
let len = shape_len("Mat::from_vec", rows, cols);
assert_eq!(
data.len(),
len,
"Mat::from_vec: data len {} != rows*cols {}",
data.len(),
len
);
Self { rows, cols, data }
}
#[must_use]
pub fn zeros(rows: usize, cols: usize) -> Self {
let len = shape_len("Mat::zeros", rows, cols);
Self {
rows,
cols,
data: vec![0.0f32; len],
}
}
#[must_use]
pub fn new(rows: usize, cols: usize) -> Self {
Self::zeros(rows, cols)
}
#[must_use]
pub fn len(&self) -> usize {
self.data.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.data.is_empty()
}
#[must_use]
pub fn shape(&self) -> (usize, usize) {
(self.rows, self.cols)
}
#[must_use]
pub fn get(&self, r: usize, c: usize) -> f32 {
assert!(r < self.rows && c < self.cols, "Mat::get out of bounds");
self.data[r * self.cols + c]
}
pub fn set(&mut self, r: usize, c: usize, v: f32) {
assert!(r < self.rows && c < self.cols, "Mat::set out of bounds");
self.data[r * self.cols + c] = v;
}
#[must_use]
pub fn row(&self, r: usize) -> &[f32] {
assert!(r < self.rows, "Mat::row out of bounds");
&self.data[r * self.cols..(r + 1) * self.cols]
}
pub fn row_mut(&mut self, r: usize) -> &mut [f32] {
assert!(r < self.rows, "Mat::row_mut out of bounds");
let c = self.cols;
&mut self.data[r * c..(r + 1) * c]
}
}
#[derive(Debug, Clone)]
pub enum Int8Weights {
Owned(Vec<i8>),
Shared(super::weights::SharedBytes),
}
impl std::ops::Deref for Int8Weights {
type Target = [i8];
fn deref(&self) -> &[i8] {
match self {
Int8Weights::Owned(v) => v,
Int8Weights::Shared(s) => bytemuck::cast_slice(s),
}
}
}
impl From<Vec<i8>> for Int8Weights {
fn from(v: Vec<i8>) -> Self {
Int8Weights::Owned(v)
}
}
impl PartialEq for Int8Weights {
fn eq(&self, other: &Self) -> bool {
self[..] == other[..]
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct QInt8 {
pub w: Int8Weights,
pub scales: Vec<f32>,
pub n: usize,
pub k: usize,
pub layout: WeightLayout,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum WeightLayout {
RowMajor,
SmmlaPanels,
}
impl QInt8 {
#[must_use]
pub fn new(w: Vec<i8>, scales: Vec<f32>, n: usize, k: usize) -> Self {
let len = shape_len("QInt8::new", n, k);
assert_eq!(w.len(), len, "QInt8: w len {} != n*k {}", w.len(), len);
assert_eq!(
scales.len(),
n,
"QInt8: scales len {} != n {}",
scales.len(),
n
);
Self {
w: w.into(),
scales,
n,
k,
layout: WeightLayout::RowMajor,
}
}
#[must_use]
pub fn new_shared(
w: super::weights::SharedBytes,
scales: Vec<f32>,
n: usize,
k: usize,
) -> Self {
let len = shape_len("QInt8::new_shared", n, k);
assert_eq!(w.len(), len, "QInt8: w len {} != n*k {}", w.len(), len);
assert_eq!(
scales.len(),
n,
"QInt8: scales len {} != n {}",
scales.len(),
n
);
Self {
w: Int8Weights::Shared(w),
scales,
n,
k,
layout: WeightLayout::RowMajor,
}
}
#[must_use]
pub fn new_smmla_panels(w: Vec<i8>, scales: Vec<f32>, n: usize, k: usize) -> Self {
let len = crate::simd::pack::smmla_packed_len(n, k);
assert_eq!(
w.len(),
len,
"QInt8: panel len {} != ceil(n/2)*ceil(k/8)*16 {}",
w.len(),
len
);
assert_eq!(
scales.len(),
n,
"QInt8: scales len {} != n {}",
scales.len(),
n
);
Self {
w: w.into(),
scales,
n,
k,
layout: WeightLayout::SmmlaPanels,
}
}
#[must_use]
pub fn expected_w_len(&self) -> usize {
match self.layout {
WeightLayout::RowMajor => self.n * self.k,
WeightLayout::SmmlaPanels => crate::simd::pack::smmla_packed_len(self.n, self.k),
}
}
}
#[derive(Debug, Clone)]
pub enum PackedBytes {
Owned(Vec<u8>),
Shared(super::weights::SharedBytes),
}
impl std::ops::Deref for PackedBytes {
type Target = [u8];
fn deref(&self) -> &[u8] {
match self {
PackedBytes::Owned(v) => v,
PackedBytes::Shared(s) => s,
}
}
}
impl From<Vec<u8>> for PackedBytes {
fn from(v: Vec<u8>) -> Self {
PackedBytes::Owned(v)
}
}
impl PartialEq for PackedBytes {
fn eq(&self, other: &Self) -> bool {
self[..] == other[..]
}
}
#[derive(Debug, Clone)]
pub enum GroupScales {
Owned(Vec<f32>),
RawLe(super::weights::SharedBytes),
}
impl GroupScales {
#[must_use]
pub fn len(&self) -> usize {
match self {
GroupScales::Owned(v) => v.len(),
GroupScales::RawLe(bytes) => bytes.len() / 4,
}
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
#[must_use]
pub fn slice_f32<'a>(
&'a self,
start: usize,
end: usize,
scratch: &'a mut Vec<f32>,
) -> &'a [f32] {
assert!(
start <= end && end <= self.len(),
"GroupScales::slice_f32: [{start}, {end}) out of bounds (len {})",
self.len()
);
match self {
GroupScales::Owned(v) => &v[start..end],
GroupScales::RawLe(bytes) => {
scratch.clear();
scratch.extend(
bytes[start * 4..end * 4]
.as_chunks::<4>()
.0
.iter()
.map(|c| f32::from_le_bytes(*c)),
);
scratch
}
}
}
#[must_use]
pub fn to_vec(&self) -> Vec<f32> {
let mut scratch = Vec::new();
self.slice_f32(0, self.len(), &mut scratch).to_vec()
}
}
impl From<Vec<f32>> for GroupScales {
fn from(v: Vec<f32>) -> Self {
GroupScales::Owned(v)
}
}
impl PartialEq for GroupScales {
fn eq(&self, other: &Self) -> bool {
self.len() == other.len() && self.to_vec() == other.to_vec()
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct QInt4 {
pub packed: PackedBytes,
pub scales: GroupScales,
pub n: usize,
pub k: usize,
pub group_size: usize,
pub tier: u8,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn mat_from_vec_roundtrips() {
let m = Mat::from_vec(2, 3, vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
assert_eq!(m.shape(), (2, 3));
assert_eq!(m.len(), 6);
assert!(!m.is_empty());
assert_eq!(m.get(0, 0), 1.0);
assert_eq!(m.get(0, 2), 3.0);
assert_eq!(m.get(1, 0), 4.0);
assert_eq!(m.get(1, 2), 6.0);
}
#[test]
fn mat_zeros_and_new_agree() {
let z = Mat::zeros(3, 4);
let n = Mat::new(3, 4);
assert_eq!(z, n);
assert!(z.data.iter().all(|&v| v == 0.0));
assert_eq!(z.len(), 12);
}
#[test]
fn mat_row_is_contiguous() {
let m = Mat::from_vec(2, 3, vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
assert_eq!(m.row(0), &[1.0, 2.0, 3.0]);
assert_eq!(m.row(1), &[4.0, 5.0, 6.0]);
}
#[test]
fn mat_set_and_row_mut() {
let mut m = Mat::zeros(2, 2);
m.set(0, 1, 7.0);
assert_eq!(m.get(0, 1), 7.0);
m.row_mut(1).copy_from_slice(&[8.0, 9.0]);
assert_eq!(m.row(1), &[8.0, 9.0]);
}
#[test]
#[should_panic(expected = "data len")]
fn mat_from_vec_rejects_bad_len() {
let _ = Mat::from_vec(2, 3, vec![1.0, 2.0]);
}
#[test]
#[should_panic(expected = "Mat::from_vec: rows*cols overflow")]
fn mat_from_vec_rejects_shape_overflow() {
let _ = Mat::from_vec(usize::MAX, 2, Vec::new());
}
#[test]
#[should_panic(expected = "Mat::zeros: rows*cols overflow")]
fn mat_zeros_rejects_shape_overflow_before_allocating() {
let _ = Mat::zeros(usize::MAX, 2);
}
#[test]
fn qint8_new_validates_shape() {
let q = QInt8::new(vec![1i8, 2, 3, 4, 5, 6], vec![0.1, 0.2], 2, 3);
assert_eq!(q.n, 2);
assert_eq!(q.k, 3);
assert_eq!(q.w.len(), 6);
assert_eq!(q.scales.len(), 2);
}
#[test]
#[should_panic(expected = "w len")]
fn qint8_rejects_bad_weight_len() {
let _ = QInt8::new(vec![1i8, 2, 3], vec![0.1, 0.2], 2, 3);
}
#[test]
#[should_panic(expected = "QInt8::new: rows*cols overflow")]
fn qint8_rejects_shape_overflow() {
let _ = QInt8::new(Vec::new(), Vec::new(), usize::MAX, 2);
}
#[test]
fn qint4_placeholder_constructs() {
let q = QInt4 {
packed: (0u8..16).collect::<Vec<u8>>().into(),
scales: vec![0.1, 0.2].into(),
n: 2,
k: 16,
group_size: 16,
tier: 1,
};
assert_eq!(q.packed.len(), q.n * (q.k / 2));
assert_eq!(q.scales.len(), q.n * (q.k / q.group_size));
}
#[test]
fn group_scales_raw_le_decodes_identically_to_owned() {
let values = vec![0.125f32, -3.5, 1.0e-3, 0.0, 7.25];
let bytes: Vec<u8> = values.iter().flat_map(|v| v.to_le_bytes()).collect();
let owned = GroupScales::Owned(values.clone());
assert_eq!(owned.len(), 5);
assert_eq!(owned.to_vec(), values);
let mut scratch = Vec::new();
assert_eq!(owned.slice_f32(1, 4, &mut scratch), &values[1..4]);
assert_eq!(bytes.len(), values.len() * 4);
let decoded: Vec<f32> = bytes
.as_chunks::<4>()
.0
.iter()
.map(|c| f32::from_le_bytes(*c))
.collect();
assert_eq!(decoded, values);
}
}