Skip to main content

int4_speed_gate/
int4_speed_gate.rs

1//! Doctrine #2 gate (a) for int4: is W4A8 actually FASTER than W8A8, unpack cost included?
2//!
3//! Run: `cargo run --release -p ftts-kernels --example int4_speed_gate`
4//!
5//! # Why this exists before any listening test
6//!
7//! int4 halves the microdecoder's weight bytes, and its 5-layer body is re-read fifteen times per
8//! frame, so cache residency is the whole thesis. But the unpack is not free: every nibble costs a
9//! mask and a shift that Q8 does not pay. Doctrine #2 requires this measured on the real ISA, at
10//! real shapes, INCLUDING that cost — and it comes first, because a lever that is slower does not
11//! deserve anyone's ears.
12//!
13//! # What is compared, and why two comparisons rather than one
14//!
15//! * **scalar vs scalar** isolates the memory-traffic thesis from implementation quality. Both
16//!   sides are plain Rust with no intrinsics, so a win here is the halved bytes talking.
17//! * **int4 vs the DISPATCHED int8 route** is the shipping question. On this machine that route is
18//!   NEON SDOT — a hand-written SIMD island — against a scalar int4 kernel with no SIMD unpack.
19//!   Losing that comparison does NOT kill int4; it says int4 needs its own in-register unpack
20//!   before it can compete, which is exactly NE-INH-004's documented escape clause.
21//!
22//! Shapes are the microdecoder's own: hidden 1024, intermediate 3072, 16 Q / 8 KV heads at
23//! head_dim 128, per `MicrodecoderConfig::default`. Every projection runs at m = 1, the decode
24//! geometry, because the microdecoder is fifteen sequential single-token steps.
25
26use std::time::Instant;
27
28use ftts_kernels::int4::{QuantizedMatrixQ4, linear_q4};
29use ftts_kernels::int8::{Int8Tier, QuantizedMatrix, linear_q8, quantize_row_q8};
30
31/// Deterministic operands spanning the representable range, so neither side is flattered by a
32/// sparse or small-magnitude matrix.
33fn 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    // (label, n, k) for one microdecoder layer, at decode geometry m = 1.
47    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    // Fifteen depths x five layers is how often this body is walked per frame.
61    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    // The ratio columns are SPEED ratios (q4's throughput relative to the q8 variant):
67    // t_q8 / t_q4, so 1.0 means parity and 0.05 means q4 runs at 5% of q8's speed.
68    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        // Warm the caches for whichever side runs first, so the ordering does not decide it.
88        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        // INTERLEAVED repeats, reporting the minimum of each variant.
92        //
93        // A first attempt timed the three variants in sequential blocks and produced 98.60 ms for
94        // k_proj against 5.14 ms for v_proj — identical 1024x1024 shapes, 19x apart. That is
95        // scheduler and thermal noise, and it is precisely what doctrine #8's same-thermal-window
96        // rule exists to prevent. Interleaving puts all three variants in the same window on every
97        // repeat, and the minimum is the least noise-contaminated estimator of a deterministic
98        // kernel: noise only ever adds time.
99        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    // Bytes are the thesis; report them so the ratio can be read against the traffic it saves.
158    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}