1use std::time::Instant;
27
28use ftts_kernels::int4::{QuantizedMatrixQ4, linear_q4};
29use ftts_kernels::int8::{Int8Tier, QuantizedMatrix, linear_q8, quantize_row_q8};
30
31fn deterministic(count: usize, seed: u64) -> Vec<f32> {
34 let mut state = seed | 1;
35 (0..count)
36 .map(|_| {
37 state ^= state << 13;
38 state ^= state >> 7;
39 state ^= state << 17;
40 ((state >> 40) as f32 / 4096.0) - 0.5
41 })
42 .collect()
43}
44
45fn main() {
46 const HIDDEN: usize = 1024;
48 const INTERMEDIATE: usize = 3072;
49 const Q_WIDTH: usize = 16 * 128;
50 const KV_WIDTH: usize = 8 * 128;
51 let projections: &[(&str, usize, usize)] = &[
52 ("q_proj", Q_WIDTH, HIDDEN),
53 ("k_proj", KV_WIDTH, HIDDEN),
54 ("v_proj", KV_WIDTH, HIDDEN),
55 ("o_proj", HIDDEN, Q_WIDTH),
56 ("gate_up", INTERMEDIATE * 2, HIDDEN),
57 ("down_proj", HIDDEN, INTERMEDIATE),
58 ];
59
60 const ROUNDS: usize = 15 * 5;
62
63 let dispatched = Int8Tier::dispatch();
64 println!("dispatched int8 route: {}", dispatched.as_str());
65 println!("shapes: microdecoder layer at m=1, {ROUNDS} rounds (15 depths x 5 layers)\n");
66 println!(
69 "{:<10} {:>6} {:>6} {:>10} {:>10} {:>10} {:>9} {:>9}",
70 "proj", "n", "k", "q8-scalar", "q4-scalar", "q8-route", "spd/q8scl", "spd/route"
71 );
72
73 let mut total_q8_scalar = 0.0_f64;
74 let mut total_q4 = 0.0_f64;
75 let mut total_q8_route = 0.0_f64;
76
77 for &(label, n, k) in projections {
78 let weight = deterministic(n * k, 0x51ED_0000 + n as u64);
79 let activation = deterministic(k, 0xA0C7_0000 + k as u64);
80 let mut x_q = vec![0_i8; k];
81 let scale = quantize_row_q8(&activation, &mut x_q);
82
83 let q8 = QuantizedMatrix::quantize(&weight, n, k);
84 let q4 = QuantizedMatrixQ4::quantize(&weight, n, k);
85 let mut out = vec![0.0_f32; n];
86
87 linear_q8(&x_q, &[scale], &q8, None, 1, &mut out, Int8Tier::Scalar);
89 linear_q4(&x_q, &[scale], &q4, None, 1, &mut out);
90
91 const REPEATS: usize = 7;
100 let mut q8_scalar = f64::MAX;
101 let mut q4_ms = f64::MAX;
102 let mut q8_route = f64::MAX;
103 for _ in 0..REPEATS {
104 let started = Instant::now();
105 for _ in 0..ROUNDS {
106 linear_q8(&x_q, &[scale], &q8, None, 1, &mut out, Int8Tier::Scalar);
107 }
108 q8_scalar = q8_scalar.min(started.elapsed().as_secs_f64() * 1000.0);
109
110 let started = Instant::now();
111 for _ in 0..ROUNDS {
112 linear_q4(&x_q, &[scale], &q4, None, 1, &mut out);
113 }
114 q4_ms = q4_ms.min(started.elapsed().as_secs_f64() * 1000.0);
115
116 let started = Instant::now();
117 for _ in 0..ROUNDS {
118 linear_q8(&x_q, &[scale], &q8, None, 1, &mut out, dispatched);
119 }
120 q8_route = q8_route.min(started.elapsed().as_secs_f64() * 1000.0);
121 }
122
123 total_q8_scalar += q8_scalar;
124 total_q4 += q4_ms;
125 total_q8_route += q8_route;
126
127 println!(
128 "{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",
129 q8_scalar / q4_ms,
130 q8_route / q4_ms
131 );
132 }
133
134 println!(
135 "\ntotals (one layer, {ROUNDS} rounds): q8-scalar {total_q8_scalar:.1} ms, \
136 q4-scalar {total_q4:.1} ms, q8-route {total_q8_route:.1} ms"
137 );
138 println!(
139 "int4 vs scalar int8 : {:.2}x ({})",
140 total_q8_scalar / total_q4,
141 if total_q4 < total_q8_scalar {
142 "FASTER - the halved-bytes thesis holds at equal implementation quality"
143 } else {
144 "SLOWER - unpack cost exceeds the traffic saving even against scalar"
145 }
146 );
147 println!(
148 "int4 vs shipping int8: {:.2}x ({})",
149 total_q8_route / total_q4,
150 if total_q4 < total_q8_route {
151 "FASTER - gate (a) PASSES; the listening gate is now worth running"
152 } else {
153 "SLOWER - gate (a) FAILS as built; int4 needs an in-register SIMD unpack to compete"
154 }
155 );
156
157 let q8_bytes: usize = projections.iter().map(|(_, n, k)| n * k).sum();
159 println!(
160 "\nweight bytes per layer: q8 {:.1} MB, q4 {:.1} MB",
161 q8_bytes as f64 / 1e6,
162 q8_bytes as f64 / 2e6
163 );
164}