use std::sync::OnceLock;
pub const Q8_MAX_ABS: i8 = 127;
const TEAM_WORK_THRESHOLD_BYTES: usize = 512 * 1024;
pub fn quantize_row_q8(row: &[f32], output: &mut [i8]) -> f32 {
assert_eq!(output.len(), row.len(), "quantize output length mismatch");
let mut maximum = 0.0_f32;
for (index, &value) in row.iter().enumerate() {
assert!(
value.is_finite(),
"non-finite value {value} at index {index} reached the Q8 quantizer"
);
maximum = maximum.max(value.abs());
}
if maximum == 0.0 {
output.fill(0);
return 1.0;
}
let scale = maximum / 127.0;
if scale == 0.0 {
output.fill(0);
return 1.0;
}
for (&value, slot) in row.iter().zip(output.iter_mut()) {
let rounded = (value / scale).clamp(-127.0, 127.0).round_ties_even();
*slot = rounded as i8;
}
scale
}
#[derive(Clone, Debug)]
pub struct QuantizedMatrix {
pub data: Vec<i8>,
pub scales: Vec<f32>,
pub n: usize,
pub k: usize,
}
impl QuantizedMatrix {
#[must_use]
pub fn concat_rows(parts: &[&Self]) -> Self {
let k = parts.first().expect("at least one part").k;
assert!(parts.iter().all(|part| part.k == k), "parts must share k");
let n = parts.iter().map(|part| part.n).sum();
let mut data = Vec::with_capacity(n * k);
let mut scales = Vec::with_capacity(n);
for part in parts {
data.extend_from_slice(&part.data);
scales.extend_from_slice(&part.scales);
}
Self { data, scales, n, k }
}
#[must_use]
pub fn quantize(weight: &[f32], n: usize, k: usize) -> Self {
assert_eq!(weight.len(), n * k, "weight must be [n, k]");
let mut data = vec![0_i8; n * k];
let mut scales = vec![0.0_f32; n];
for ((weight_row, data_row), scale) in weight
.chunks_exact(k)
.zip(data.chunks_exact_mut(k))
.zip(scales.iter_mut())
{
*scale = quantize_row_q8(weight_row, data_row);
}
Self { data, scales, n, k }
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum Int8Tier {
Scalar,
Autovec,
NeonSdot,
WasmSimd128,
}
impl Int8Tier {
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::Scalar => "scalar",
Self::Autovec => "autovec",
Self::NeonSdot => "neon-sdot",
Self::WasmSimd128 => "wasm-simd128",
}
}
#[must_use]
pub fn available() -> Vec<Self> {
let mut tiers = vec![Self::Scalar, Self::Autovec];
if neon_sdot_available() {
tiers.push(Self::NeonSdot);
}
if wasm_simd128_available() {
tiers.push(Self::WasmSimd128);
}
tiers
}
#[must_use]
pub fn dispatch() -> Self {
if wasm_simd128_available() {
return Self::WasmSimd128;
}
match std::env::var("FTTS_INT8_TIER").as_deref() {
Ok("scalar") => Self::Scalar,
Ok("autovec") => Self::Autovec,
Ok("neon-sdot") if neon_sdot_available() => Self::NeonSdot,
_ if neon_sdot_available() => Self::NeonSdot,
_ => Self::Scalar,
}
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum QuantLinearMode {
W8A8(Int8Tier),
W8A16,
}
impl QuantLinearMode {
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::W8A8(_) => "w8a8",
Self::W8A16 => "w8a16",
}
}
}
pub fn linear_w8a16(
x: &[f32],
weight: &QuantizedMatrix,
bias: Option<&[f32]>,
m: usize,
out: &mut [f32],
) {
let (n, k) = (weight.n, weight.k);
assert_eq!(x.len(), m * k, "x must be [m, k]");
assert_eq!(out.len(), m * n, "out must be [m, n]");
if let Some(bias) = bias {
assert_eq!(bias.len(), n, "bias must be [n]");
}
for col in 0..n {
let w_row = &weight.data[col * k..(col + 1) * k];
let w_scale = weight.scales[col];
let bias_term = bias.map(|b| b[col]);
for row in 0..m {
let x_row = &x[row * k..(row + 1) * k];
let acc = dot_w8a16(x_row, w_row);
let value = acc * w_scale;
out[row * n + col] = bias_term.map_or(value, |b| value + b);
}
}
}
fn dot_w8a16(x: &[f32], w: &[i8]) -> f32 {
const LANES: usize = 8;
let mut lanes = [0.0_f32; LANES];
let chunks = x.len() / LANES;
for chunk in 0..chunks {
let base = chunk * LANES;
for lane in 0..LANES {
lanes[lane] = f32::from(w[base + lane]).mul_add(x[base + lane], lanes[lane]);
}
}
let mut sum: f32 = lanes.iter().sum();
for index in chunks * LANES..x.len() {
sum = f32::from(w[index]).mul_add(x[index], sum);
}
sum
}
#[must_use]
pub fn quant_mode_from_environment() -> QuantLinearMode {
match std::env::var("FTTS_INT8").as_deref() {
Ok("w8a16") => QuantLinearMode::W8A16,
_ => QuantLinearMode::W8A8(autotuned_plan().decode_gemv),
}
}
pub fn quant_linear(
mode: QuantLinearMode,
x: &[f32],
weight: &QuantizedMatrix,
bias: Option<&[f32]>,
m: usize,
out: &mut [f32],
) {
match mode {
QuantLinearMode::W8A8(tier) => linear_q8_dynamic(x, weight, bias, m, out, tier),
QuantLinearMode::W8A16 => linear_w8a16(x, weight, bias, m, out),
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct KernelPlanV0 {
pub decode_gemv: Int8Tier,
pub batch_gemm: Int8Tier,
}
pub fn autotuned_plan() -> KernelPlanV0 {
static PLAN: OnceLock<KernelPlanV0> = OnceLock::new();
*PLAN.get_or_init(|| {
#[cfg(target_arch = "wasm32")]
{
let tier = Int8Tier::dispatch();
KernelPlanV0 {
decode_gemv: tier,
batch_gemm: tier,
}
}
#[cfg(not(target_arch = "wasm32"))]
{
if std::env::var("FTTS_INT8_TIER").is_ok() {
let forced = Int8Tier::dispatch();
return KernelPlanV0 {
decode_gemv: forced,
batch_gemm: forced,
};
}
if let Some(cached) = load_persisted_plan() {
return cached;
}
let plan = KernelPlanV0 {
decode_gemv: fastest_tier(&[(1, 1024, 256), (1, 3072, 256)]),
batch_gemm: fastest_tier(&[(16, 1024, 128), (16, 3072, 64)]),
};
persist_plan(plan);
plan
}
})
}
fn plan_cache_path() -> Option<std::path::PathBuf> {
std::env::var_os("HOME")
.map(|home| std::path::PathBuf::from(home).join(".cache/franken_tts/kernel_plan_v0.txt"))
}
fn plan_cache_key() -> String {
let tiers: Vec<&str> = Int8Tier::available().iter().map(|t| t.as_str()).collect();
format!(
"v1|crate={}|arch={}-{}|cores={}|tiers={}",
env!("CARGO_PKG_VERSION"),
std::env::consts::ARCH,
std::env::consts::OS,
std::thread::available_parallelism().map_or(1, usize::from),
tiers.join(",")
)
}
fn load_persisted_plan() -> Option<KernelPlanV0> {
let text = {
use std::io::Read as _;
let mut text = String::new();
let file = std::fs::File::open(plan_cache_path()?).ok()?;
file.take(512).read_to_string(&mut text).ok()?;
text
};
let mut lines = text.lines();
if lines.next()? != plan_cache_key() {
return None;
}
let parse = |line: &str| match line {
"scalar" => Some(Int8Tier::Scalar),
"autovec" => Some(Int8Tier::Autovec),
"neon-sdot" if neon_sdot_available() => Some(Int8Tier::NeonSdot),
_ => None,
};
Some(KernelPlanV0 {
decode_gemv: parse(lines.next()?)?,
batch_gemm: parse(lines.next()?)?,
})
}
fn persist_plan(plan: KernelPlanV0) {
let Some(path) = plan_cache_path() else {
return;
};
if let Some(parent) = path.parent() {
let _ = std::fs::create_dir_all(parent);
}
let _ = std::fs::write(
path,
format!(
"{}\n{}\n{}\n",
plan_cache_key(),
plan.decode_gemv.as_str(),
plan.batch_gemm.as_str()
),
);
}
fn fastest_tier(probes: &[(usize, usize, usize)]) -> Int8Tier {
use std::time::Instant;
let tiers = Int8Tier::available();
let mut best = (tiers[0], f64::MAX);
for &tier in &tiers {
let mut total = 0.0_f64;
for &(m, k, n) in probes {
let x_q: Vec<i8> = (0..m * k)
.map(|i| (((i * 37 + 11) % 255) as i32 - 127) as i8)
.collect();
let x_scales = vec![1.0_f32; m];
let weight = QuantizedMatrix {
data: (0..n * k)
.map(|i| (((i * 29 + 5) % 255) as i32 - 127) as i8)
.collect(),
scales: vec![1.0_f32; n],
n,
k,
};
let mut out = vec![0.0_f32; m * n];
let mut rounds: Vec<f64> = (0..3)
.map(|_| {
let start = Instant::now();
crate::team::with_team_bypassed(|| {
linear_q8(&x_q, &x_scales, &weight, None, m, &mut out, tier);
});
start.elapsed().as_secs_f64()
})
.collect();
rounds.sort_by(f64::total_cmp);
total += rounds[1];
}
if total < best.1 {
best = (tier, total);
}
}
best.0
}
#[must_use]
pub fn neon_sdot_available() -> bool {
#[cfg(all(target_arch = "aarch64", feature = "neon-dotprod"))]
{
neon_dotprod::available()
}
#[cfg(not(all(target_arch = "aarch64", feature = "neon-dotprod")))]
{
false
}
}
#[must_use]
pub fn wasm_simd128_available() -> bool {
cfg!(all(target_arch = "wasm32", target_feature = "simd128"))
}
#[must_use]
pub fn dot_i32(a: &[i8], b: &[i8], tier: Int8Tier) -> i32 {
assert_eq!(a.len(), b.len(), "int8 dot inputs must match");
match tier {
Int8Tier::Scalar => dot_i32_scalar(a, b),
Int8Tier::Autovec => dot_i32_autovec(a, b),
Int8Tier::NeonSdot => dot_i32_neon_or_panic(a, b),
Int8Tier::WasmSimd128 => dot_i32_wasm_or_panic(a, b),
}
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
fn dot_i32_wasm_or_panic(a: &[i8], b: &[i8]) -> i32 {
wasm_simd128::dot_i32(a, b)
}
#[cfg(not(all(target_arch = "wasm32", target_feature = "simd128")))]
fn dot_i32_wasm_or_panic(_a: &[i8], _b: &[i8]) -> i32 {
panic!("wasm-simd128 route selected on a build without the island");
}
#[cfg(all(target_arch = "aarch64", feature = "neon-dotprod"))]
fn dot_i32_neon_or_panic(a: &[i8], b: &[i8]) -> i32 {
assert!(
neon_dotprod::available(),
"neon-sdot route selected without FEAT_DotProd"
);
neon_dotprod::dot_i32(a, b)
}
#[cfg(not(all(target_arch = "aarch64", feature = "neon-dotprod")))]
fn dot_i32_neon_or_panic(_a: &[i8], _b: &[i8]) -> i32 {
panic!("neon-sdot route selected on a build without the island");
}
fn dot_i32_scalar(a: &[i8], b: &[i8]) -> i32 {
let mut sum = 0_i32;
for index in 0..a.len() {
sum += i32::from(a[index]) * i32::from(b[index]);
}
sum
}
fn dot_i32_autovec(a: &[i8], b: &[i8]) -> i32 {
const LANES: usize = 8;
let mut lanes = [0_i32; LANES];
let chunks = a.len() / LANES;
for chunk in 0..chunks {
let base = chunk * LANES;
for lane in 0..LANES {
lanes[lane] += i32::from(a[base + lane]) * i32::from(b[base + lane]);
}
}
let mut sum: i32 = lanes.iter().sum();
for index in chunks * LANES..a.len() {
sum += i32::from(a[index]) * i32::from(b[index]);
}
sum
}
#[cfg(all(target_arch = "aarch64", feature = "neon-dotprod"))]
mod neon_dotprod {
use core::arch::aarch64::{vaddq_s32, vaddvq_s32, vdotq_s32, vdupq_n_s32, vld1q_s8};
#[must_use]
pub fn available() -> bool {
std::arch::is_aarch64_feature_detected!("dotprod")
}
#[must_use]
pub fn dot_i32(a: &[i8], b: &[i8]) -> i32 {
debug_assert!(available(), "SDOT island entered without FEAT_DotProd");
unsafe { dot_i32_sdot(a, b) }
}
#[target_feature(enable = "neon,dotprod")]
unsafe fn dot_i32_sdot(a: &[i8], b: &[i8]) -> i32 {
let len = a.len();
let a_ptr = a.as_ptr();
let b_ptr = b.as_ptr();
let mut acc0 = vdupq_n_s32(0);
let mut acc1 = vdupq_n_s32(0);
let mut acc2 = vdupq_n_s32(0);
let mut acc3 = vdupq_n_s32(0);
let mut index = 0_usize;
while index + 64 <= len {
unsafe {
acc0 = vdotq_s32(acc0, vld1q_s8(a_ptr.add(index)), vld1q_s8(b_ptr.add(index)));
acc1 = vdotq_s32(
acc1,
vld1q_s8(a_ptr.add(index + 16)),
vld1q_s8(b_ptr.add(index + 16)),
);
acc2 = vdotq_s32(
acc2,
vld1q_s8(a_ptr.add(index + 32)),
vld1q_s8(b_ptr.add(index + 32)),
);
acc3 = vdotq_s32(
acc3,
vld1q_s8(a_ptr.add(index + 48)),
vld1q_s8(b_ptr.add(index + 48)),
);
}
index += 64;
}
while index + 16 <= len {
unsafe {
acc0 = vdotq_s32(acc0, vld1q_s8(a_ptr.add(index)), vld1q_s8(b_ptr.add(index)));
}
index += 16;
}
let mut sum = vaddvq_s32(vaddq_s32(vaddq_s32(acc0, acc1), vaddq_s32(acc2, acc3)));
while index < len {
sum += i32::from(a[index]) * i32::from(b[index]);
index += 1;
}
sum
}
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
mod wasm_simd128 {
use core::arch::wasm32::{
i16x8_extend_high_i8x16, i16x8_extend_low_i8x16, i32x4_add, i32x4_dot_i16x8,
i32x4_extract_lane, i32x4_splat, v128, v128_load,
};
#[must_use]
pub fn dot_i32(a: &[i8], b: &[i8]) -> i32 {
unsafe { dot_i32_simd128(a, b) }
}
#[inline]
unsafe fn accumulate_block(acc: v128, a: *const i8, b: *const i8) -> v128 {
let (left, right) = unsafe { (v128_load(a.cast()), v128_load(b.cast())) };
let low = i32x4_dot_i16x8(i16x8_extend_low_i8x16(left), i16x8_extend_low_i8x16(right));
let high = i32x4_dot_i16x8(
i16x8_extend_high_i8x16(left),
i16x8_extend_high_i8x16(right),
);
i32x4_add(acc, i32x4_add(low, high))
}
pub fn linear_blocked(
x_q: &[i8],
x_scales: &[f32],
weight: &super::QuantizedMatrix,
bias: Option<&[f32]>,
m: usize,
out: &mut [f32],
) {
let (n, k) = (weight.n, weight.k);
let mut col = 0;
while col + 4 <= n {
for row in 0..m {
let x_row = &x_q[row * k..(row + 1) * k];
let acc = unsafe { dot4_simd128(x_row, &weight.data[col * k..], k) };
for (lane, accumulated) in acc.iter().enumerate() {
let column = col + lane;
#[allow(clippy::cast_precision_loss)]
let value = *accumulated as f32 * (x_scales[row] * weight.scales[column]);
out[row * n + column] = bias.map_or(value, |values| value + values[column]);
}
}
col += 4;
}
while col < n {
let w_row = &weight.data[col * k..(col + 1) * k];
for row in 0..m {
let x_row = &x_q[row * k..(row + 1) * k];
#[allow(clippy::cast_precision_loss)]
let value = dot_i32(x_row, w_row) as f32 * (x_scales[row] * weight.scales[col]);
out[row * n + col] = bias.map_or(value, |values| value + values[col]);
}
col += 1;
}
}
unsafe fn dot4_simd128(x: &[i8], weights: &[i8], k: usize) -> [i32; 4] {
let x_ptr = x.as_ptr();
let w_ptr = weights.as_ptr();
let mut acc = [i32x4_splat(0); 4];
let mut index = 0_usize;
while index + 16 <= k {
let (low, high) = unsafe {
let block = v128_load(x_ptr.add(index).cast());
(
i16x8_extend_low_i8x16(block),
i16x8_extend_high_i8x16(block),
)
};
for (lane, accumulator) in acc.iter_mut().enumerate() {
let w = unsafe { v128_load(w_ptr.add(lane * k + index).cast()) };
let products = i32x4_add(
i32x4_dot_i16x8(low, i16x8_extend_low_i8x16(w)),
i32x4_dot_i16x8(high, i16x8_extend_high_i8x16(w)),
);
*accumulator = i32x4_add(*accumulator, products);
}
index += 16;
}
let mut sums = [0_i32; 4];
for (lane, sum) in sums.iter_mut().enumerate() {
let total = acc[lane];
*sum = i32x4_extract_lane::<0>(total)
+ i32x4_extract_lane::<1>(total)
+ i32x4_extract_lane::<2>(total)
+ i32x4_extract_lane::<3>(total);
for tail in index..k {
let w = unsafe { *w_ptr.add(lane * k + tail) };
*sum += i32::from(x[tail]) * i32::from(w);
}
}
sums
}
unsafe fn dot_i32_simd128(a: &[i8], b: &[i8]) -> i32 {
let len = a.len();
let a_ptr = a.as_ptr();
let b_ptr = b.as_ptr();
let mut acc0 = i32x4_splat(0);
let mut acc1 = i32x4_splat(0);
let mut acc2 = i32x4_splat(0);
let mut acc3 = i32x4_splat(0);
let mut index = 0_usize;
while index + 64 <= len {
unsafe {
acc0 = accumulate_block(acc0, a_ptr.add(index), b_ptr.add(index));
acc1 = accumulate_block(acc1, a_ptr.add(index + 16), b_ptr.add(index + 16));
acc2 = accumulate_block(acc2, a_ptr.add(index + 32), b_ptr.add(index + 32));
acc3 = accumulate_block(acc3, a_ptr.add(index + 48), b_ptr.add(index + 48));
}
index += 64;
}
while index + 16 <= len {
unsafe {
acc0 = accumulate_block(acc0, a_ptr.add(index), b_ptr.add(index));
}
index += 16;
}
let total = i32x4_add(i32x4_add(acc0, acc1), i32x4_add(acc2, acc3));
let mut sum = i32x4_extract_lane::<0>(total)
+ i32x4_extract_lane::<1>(total)
+ i32x4_extract_lane::<2>(total)
+ i32x4_extract_lane::<3>(total);
while index < len {
sum += i32::from(a[index]) * i32::from(b[index]);
index += 1;
}
sum
}
}
#[allow(clippy::too_many_arguments)]
pub fn linear_q8(
x_q: &[i8],
x_scales: &[f32],
weight: &QuantizedMatrix,
bias: Option<&[f32]>,
m: usize,
out: &mut [f32],
tier: Int8Tier,
) {
let (n, k) = (weight.n, weight.k);
assert_eq!(x_q.len(), m * k, "x_q must be [m, k]");
assert_eq!(x_scales.len(), m, "x_scales must be [m]");
assert_eq!(out.len(), m * n, "out must be [m, n]");
if let Some(bias) = bias {
assert_eq!(bias.len(), n, "bias must be [n]");
}
if n * k >= TEAM_WORK_THRESHOLD_BYTES
&& !crate::team::thread_bypassed()
&& let Some(team) = crate::team::armed()
{
team.linear_q8(x_q, x_scales, weight, bias, m, out, tier);
return;
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
if matches!(tier, Int8Tier::WasmSimd128) {
wasm_simd128::linear_blocked(x_q, x_scales, weight, bias, m, out);
return;
}
for col in 0..n {
let w_row = &weight.data[col * k..(col + 1) * k];
let w_scale = weight.scales[col];
let bias_term = bias.map(|b| b[col]);
for row in 0..m {
let x_row = &x_q[row * k..(row + 1) * k];
let acc = dot_i32(x_row, w_row, tier);
let value = acc as f32 * (x_scales[row] * w_scale);
out[row * n + col] = bias_term.map_or(value, |b| value + b);
}
}
}
pub fn linear_q8_dynamic(
x: &[f32],
weight: &QuantizedMatrix,
bias: Option<&[f32]>,
m: usize,
out: &mut [f32],
tier: Int8Tier,
) {
let k = weight.k;
assert_eq!(x.len(), m * k, "x must be [m, k]");
let mut x_q = vec![0_i8; m * k];
let mut x_scales = vec![0.0_f32; m];
for ((x_row, q_row), scale) in x
.chunks_exact(k)
.zip(x_q.chunks_exact_mut(k))
.zip(x_scales.iter_mut())
{
*scale = quantize_row_q8(x_row, q_row);
}
linear_q8(&x_q, &x_scales, weight, bias, m, out, tier);
}
#[cfg(test)]
mod tests {
use super::*;
fn pseudo_random_q8(len: usize, seed: u64) -> Vec<i8> {
let mut state = seed;
(0..len)
.map(|_| {
state = state.wrapping_add(0x9e37_79b9_7f4a_7c15);
let mut z = state;
z = (z ^ (z >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9);
z = (z ^ (z >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb);
z ^= z >> 31;
((z % 255) as i32 - 127) as i8
})
.collect()
}
const MODEL_SHAPES: &[(usize, usize)] = &[
(2048, 1024), (1024, 1024), (1024, 2048), (3072, 1024), (1024, 3072), ];
#[test]
fn every_tier_is_exactly_equal_in_i32_at_every_model_shape() {
for &(n, k) in MODEL_SHAPES {
let a = pseudo_random_q8(k, 0x5eed_0001 ^ (n as u64) << 20 ^ k as u64);
let w = pseudo_random_q8(n * k, 0x5eed_0002 ^ (n as u64) << 20 ^ k as u64);
for row in [0, n / 2, n - 1] {
let w_row = &w[row * k..(row + 1) * k];
let reference = dot_i32(&a, w_row, Int8Tier::Scalar);
for tier in Int8Tier::available() {
assert_eq!(
dot_i32(&a, w_row, tier),
reference,
"tier {} diverged at shape {n}x{k} row {row}",
tier.as_str()
);
}
}
}
}
#[test]
fn every_tier_survives_the_all_extreme_reduction_at_the_binding_census_k() {
for k in [2048_usize, 3072, 4608, 7168, 8192] {
let a = vec![127_i8; k];
let b = vec![127_i8; k];
let negative = vec![-127_i8; k];
let expected = 127_i64 * 127 * k as i64;
for tier in Int8Tier::available() {
assert_eq!(
i64::from(dot_i32(&a, &b, tier)),
expected,
"positive all-extreme diverged on {} at K={k}",
tier.as_str()
);
assert_eq!(
i64::from(dot_i32(&a, &negative, tier)),
-expected,
"negative all-extreme diverged on {} at K={k}",
tier.as_str()
);
}
}
}
#[test]
fn tail_lengths_that_defeat_block_boundaries_stay_exact() {
for len in [1_usize, 7, 15, 16, 17, 63, 64, 65, 100, 129] {
let a = pseudo_random_q8(len, tail_seed(len));
let b = pseudo_random_q8(len, tail_seed(len) ^ 1);
let reference = dot_i32(&a, &b, Int8Tier::Scalar);
for tier in Int8Tier::available() {
assert_eq!(
dot_i32(&a, &b, tier),
reference,
"len={len} {}",
tier.as_str()
);
}
}
}
#[test]
fn quantizer_matches_the_canonical_converter_semantics() {
let row = [
-127.0_f32, -126.5, -125.5, -1.5, -0.5, 0.5, 1.5, 125.5, 126.5, 127.0,
];
let mut q = [0_i8; 10];
let scale = quantize_row_q8(&row, &mut q);
assert_eq!(scale.to_bits(), 1.0_f32.to_bits());
assert_eq!(q, [-127, -126, -126, -2, 0, 0, 2, 126, 126, 127]);
let zeros = [0.0_f32; 4];
let mut qz = [1_i8; 4];
assert_eq!(
quantize_row_q8(&zeros, &mut qz).to_bits(),
1.0_f32.to_bits()
);
assert_eq!(qz, [0, 0, 0, 0]);
let matrix = QuantizedMatrix::quantize(&[2.0, -1.0, 0.0, 3.0], 2, 2);
assert_eq!(matrix.scales[0].to_bits(), (2.0_f32 / 127.0).to_bits());
assert_eq!(matrix.scales[1].to_bits(), (3.0_f32 / 127.0).to_bits());
assert!(matrix.data.iter().all(|&b| b != -128));
}
#[test]
fn dynamic_w8a8_linear_tracks_the_f32_reference_within_quant_error() {
let (n, k) = (64_usize, 128_usize);
let mut weight = vec![0.0_f32; n * k];
let mut x = vec![0.0_f32; k];
let mut state = 0x1234_5678_u64;
let mut next = || {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1);
((state >> 33) as f32 / (1u64 << 31) as f32) - 1.0
};
for value in weight.iter_mut() {
*value = next();
}
for value in x.iter_mut() {
*value = next();
}
let quantized = QuantizedMatrix::quantize(&weight, n, k);
let mut out_q8 = vec![0.0_f32; n];
linear_q8_dynamic(&x, &quantized, None, 1, &mut out_q8, Int8Tier::Autovec);
let mut out_f32 = vec![0.0_f32; n];
crate::f32ref::linear(&x, &weight, None, 1, k, n, &mut out_f32);
let dot = |a: &[f32], b: &[f32]| a.iter().zip(b).map(|(x, y)| x * y).sum::<f32>();
let cosine = dot(&out_q8, &out_f32)
/ (dot(&out_q8, &out_q8).sqrt() * dot(&out_f32, &out_f32).sqrt());
assert!(
cosine > 0.999,
"W8A8 dequant plumbing is broken: cosine {cosine}"
);
}
#[test]
fn tiers_produce_bit_identical_f32_output_not_merely_close() {
let (n, k) = (256_usize, 1024_usize);
let weight: Vec<f32> = pseudo_random_q8(n * k, 77)
.iter()
.map(|&b| f32::from(b) / 64.0)
.collect();
let x: Vec<f32> = pseudo_random_q8(k, 78)
.iter()
.map(|&b| f32::from(b) / 64.0)
.collect();
let quantized = QuantizedMatrix::quantize(&weight, n, k);
let mut reference = vec![0.0_f32; n];
linear_q8_dynamic(&x, &quantized, None, 1, &mut reference, Int8Tier::Scalar);
for tier in Int8Tier::available() {
let mut out = vec![0.0_f32; n];
linear_q8_dynamic(&x, &quantized, None, 1, &mut out, tier);
for (index, (a, b)) in reference.iter().zip(&out).enumerate() {
assert_eq!(
a.to_bits(),
b.to_bits(),
"tier {} f32 output differs at {index}",
tier.as_str()
);
}
}
}
fn tail_seed(len: usize) -> u64 {
0x7a11_0000 ^ len as u64
}
}