use std::time::Instant;
use ftts_kernels::int4::{QuantizedMatrixQ4, linear_q4};
use ftts_kernels::int8::{Int8Tier, QuantizedMatrix, linear_q8, quantize_row_q8};
fn deterministic(count: usize, seed: u64) -> Vec<f32> {
let mut state = seed | 1;
(0..count)
.map(|_| {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
((state >> 40) as f32 / 4096.0) - 0.5
})
.collect()
}
fn main() {
const HIDDEN: usize = 1024;
const INTERMEDIATE: usize = 3072;
const Q_WIDTH: usize = 16 * 128;
const KV_WIDTH: usize = 8 * 128;
let projections: &[(&str, usize, usize)] = &[
("q_proj", Q_WIDTH, HIDDEN),
("k_proj", KV_WIDTH, HIDDEN),
("v_proj", KV_WIDTH, HIDDEN),
("o_proj", HIDDEN, Q_WIDTH),
("gate_up", INTERMEDIATE * 2, HIDDEN),
("down_proj", HIDDEN, INTERMEDIATE),
];
const ROUNDS: usize = 15 * 5;
let dispatched = Int8Tier::dispatch();
println!("dispatched int8 route: {}", dispatched.as_str());
println!("shapes: microdecoder layer at m=1, {ROUNDS} rounds (15 depths x 5 layers)\n");
println!(
"{:<10} {:>6} {:>6} {:>10} {:>10} {:>10} {:>9} {:>9}",
"proj", "n", "k", "q8-scalar", "q4-scalar", "q8-route", "spd/q8scl", "spd/route"
);
let mut total_q8_scalar = 0.0_f64;
let mut total_q4 = 0.0_f64;
let mut total_q8_route = 0.0_f64;
for &(label, n, k) in projections {
let weight = deterministic(n * k, 0x51ED_0000 + n as u64);
let activation = deterministic(k, 0xA0C7_0000 + k as u64);
let mut x_q = vec![0_i8; k];
let scale = quantize_row_q8(&activation, &mut x_q);
let q8 = QuantizedMatrix::quantize(&weight, n, k);
let q4 = QuantizedMatrixQ4::quantize(&weight, n, k);
let mut out = vec![0.0_f32; n];
linear_q8(&x_q, &[scale], &q8, None, 1, &mut out, Int8Tier::Scalar);
linear_q4(&x_q, &[scale], &q4, None, 1, &mut out);
const REPEATS: usize = 7;
let mut q8_scalar = f64::MAX;
let mut q4_ms = f64::MAX;
let mut q8_route = f64::MAX;
for _ in 0..REPEATS {
let started = Instant::now();
for _ in 0..ROUNDS {
linear_q8(&x_q, &[scale], &q8, None, 1, &mut out, Int8Tier::Scalar);
}
q8_scalar = q8_scalar.min(started.elapsed().as_secs_f64() * 1000.0);
let started = Instant::now();
for _ in 0..ROUNDS {
linear_q4(&x_q, &[scale], &q4, None, 1, &mut out);
}
q4_ms = q4_ms.min(started.elapsed().as_secs_f64() * 1000.0);
let started = Instant::now();
for _ in 0..ROUNDS {
linear_q8(&x_q, &[scale], &q8, None, 1, &mut out, dispatched);
}
q8_route = q8_route.min(started.elapsed().as_secs_f64() * 1000.0);
}
total_q8_scalar += q8_scalar;
total_q4 += q4_ms;
total_q8_route += q8_route;
println!(
"{label:<10} {n:>6} {k:>6} {q8_scalar:>9.2}m {q4_ms:>9.2}m {q8_route:>9.2}m {:>8.2}x {:>8.2}x",
q8_scalar / q4_ms,
q8_route / q4_ms
);
}
println!(
"\ntotals (one layer, {ROUNDS} rounds): q8-scalar {total_q8_scalar:.1} ms, \
q4-scalar {total_q4:.1} ms, q8-route {total_q8_route:.1} ms"
);
println!(
"int4 vs scalar int8 : {:.2}x ({})",
total_q8_scalar / total_q4,
if total_q4 < total_q8_scalar {
"FASTER - the halved-bytes thesis holds at equal implementation quality"
} else {
"SLOWER - unpack cost exceeds the traffic saving even against scalar"
}
);
println!(
"int4 vs shipping int8: {:.2}x ({})",
total_q8_route / total_q4,
if total_q4 < total_q8_route {
"FASTER - gate (a) PASSES; the listening gate is now worth running"
} else {
"SLOWER - gate (a) FAILS as built; int4 needs an in-register SIMD unpack to compete"
}
);
let q8_bytes: usize = projections.iter().map(|(_, n, k)| n * k).sum();
println!(
"\nweight bytes per layer: q8 {:.1} MB, q4 {:.1} MB",
q8_bytes as f64 / 1e6,
q8_bytes as f64 / 2e6
);
}