1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
//! fp8-mmq-bench — GEMM-ONLY throughput for the per-block FP8 MMQ prefill kernel vs the Q8_0 MMQ
//! floor it replaces, at the real 27B projection shapes.
//!
//! WHY THIS EXISTS separately from the model-level pp battery: on the 27B the kernel's operand (raw
//! e4m3 bytes + the f32 block grid) is made resident by the MEMRA_PP_FP8 stash, which DUPLICATES
//! every F8-origin projection on top of the resident Q8_0 and therefore has to run under a VRAM
//! budget. At the measured ~355 MiB of e4m3 per 27B layer, the 3072 MB that fits alongside a 27 GB
//! model covers a PREFIX of ~8.6 of 64 layers, so an end-to-end pp512 number is ~13% kernel and
//! ~87% floor — it cannot separate "the kernel is slower" from "the kernel barely ran". ARM A does
//! not have this problem because its scale-fold makes the operand blk=None, which admits it to the
//! MEMRA_ST_E4M3 one-copy arm (no duplicate, no budget, all 64 layers).
//!
//! So the pp comparison at equal coverage is not available on this rig, and the honest way to
//! measure the KERNEL is to measure the kernel: same shapes, same m, same device, both launchers
//! back to back, interleaved, medians. That is a GEMM-level claim and is labeled as one — it is not
//! an end-to-end speedup.
//!
//! Shapes are the 27B's own (from the safetensors headers, /root/models/qwen36-27b-fp8):
//! q_proj 5120->12288, k/v_proj 5120->1024, o_proj 6144->5120,
//! gate/up_proj 5120->17408, down_proj 17408->5120.
//!
//! The 1.7B block-128 checkpoint's shapes are ALSO available (`1p7b`), because that is the model the
//! local pp battery runs: its projections are much narrower (out_f 1024-6144 vs the 27B's
//! 1024-17408), so the token-tile selection rule that the 27B shapes calibrate is not automatically
//! right there. A pp ratio measured on the 1.7B has to be explainable by the 1.7B's own GEMM ratios.
//!
//! usage: fp8-mmq-bench [m] [reps] [27b|1p7b|step37] (default m=512, reps=9, 27b)
use memra_engine::Engine;
fn main() -> Result<(), Box<dyn std::error::Error>> {
let m: usize = std::env::args()
.nth(1)
.and_then(|s| s.parse().ok())
.unwrap_or(512);
let reps: usize = std::env::args()
.nth(2)
.and_then(|s| s.parse().ok())
.unwrap_or(9);
let set = std::env::args().nth(3).unwrap_or_else(|| "27b".to_string());
let e = Engine::new(0)?;
println!(
"GPU: {} m={m} reps={reps} shapes={set} (interleaved fp8blk,q8_0 per rep; median of reps)",
e.ctx().name()?
);
println!(
"{:<28} {:>12} {:>12} {:>10} {:>12} {:>12}",
"shape in->out", "fp8blk_ms", "q8_0_ms", "ratio", "fp8blk_TFLOP", "q8_0_TFLOP"
);
// (in_f, out_f, label)
let shapes_27b: [(usize, usize, &str); 6] = [
(5120, 12288, "q_proj"),
(5120, 1024, "k/v_proj"),
(6144, 5120, "o_proj"),
(5120, 17408, "gate/up_proj"),
(17408, 5120, "down_proj"),
(5120, 5120, "square-ref"),
];
// qwen3-1.7B: hidden 2048, 16 q heads x 128, 8 kv heads x 128, ffn 6144.
let shapes_1p7b: [(usize, usize, &str); 5] = [
(2048, 2048, "q_proj"),
(2048, 1024, "k/v_proj"),
(2048, 2048, "o_proj"),
(2048, 6144, "gate/up_proj"),
(6144, 2048, "down_proj"),
];
// Official Step-3.7 expert projections. Gate and up have the same shape and run as distinct
// checkpoint tensors; one row is sufficient for the GEMM-level tactic comparison.
let shapes_step37: [(usize, usize, &str); 2] = [
(4096, 1280, "expert_gate_or_up"),
(1280, 4096, "expert_down"),
];
let shapes: Vec<(usize, usize, &str)> = match set.as_str() {
"1p7b" => shapes_1p7b.to_vec(),
"step37" => shapes_step37.to_vec(),
_ => shapes_27b.to_vec(),
};
for (in_f, out_f, label) in shapes {
// Weight operands. Both arms get the SAME logical weight: the e4m3 codes are the source of
// truth, and the Q8_0 slab is produced from them by the merged ARM B' device dequant — i.e.
// exactly the floor path this kernel competes with, not a synthetic Q8_0.
let mut codes = vec![0u8; out_f * in_f];
let mut s: u32 = 0x1234_5678;
for c in codes.iter_mut() {
s = s.wrapping_mul(1664525).wrapping_add(1013904223);
let v = ((s >> 16) & 0x7F) as u8;
// avoid magnitude 0x7F: the hardware MMA reads it as NaN (host decodes 0.0), and the
// dispatch path refuses any tensor containing one, so a bench must not contain one.
let v = if v == 0x7F { 0x30 } else { v };
*c = v | (((s >> 8) & 1) as u8) << 7;
}
let (rows, cols) = (out_f.div_ceil(128), in_f.div_ceil(128));
let grid: Vec<f32> = (0..rows * cols)
.map(|i| 0.5f32 + (i % 7) as f32 * 0.125)
.collect();
let w_f8 = e.htod_bytes(&codes)?;
let g_d = e.htod(&grid)?;
let w_q8 = e.fp8_blk_dequant_q8_0(&codes, &grid, out_f, in_f)?;
let x: Vec<f32> = (0..m * in_f)
.map(|i| ((i % 251) as f32 - 125.0) / 64.0)
.collect();
let x_d = e.htod(&x)?;
// warmup both arms (allocation, autotune)
let _ = e.qmatvec_mmq_fp8_blk(&w_f8, &g_d, &x_d, m, in_f, out_f)?;
let _ = e.qmatvec_mmq_q8_0_raw(&w_q8, &x_d, m, in_f, out_f)?;
e.stream().synchronize()?;
let mut t_f8: Vec<f64> = Vec::with_capacity(reps);
let mut t_q8: Vec<f64> = Vec::with_capacity(reps);
for _ in 0..reps {
// INTERLEAVED inside the rep loop: the two arms then share one clock/thermal regime,
// which back-to-back blocks of N would not (clock drift is not a valid denominator).
let t0 = std::time::Instant::now();
let _ = e.qmatvec_mmq_fp8_blk(&w_f8, &g_d, &x_d, m, in_f, out_f)?;
e.stream().synchronize()?;
t_f8.push(t0.elapsed().as_secs_f64());
let t1 = std::time::Instant::now();
let _ = e.qmatvec_mmq_q8_0_raw(&w_q8, &x_d, m, in_f, out_f)?;
e.stream().synchronize()?;
t_q8.push(t1.elapsed().as_secs_f64());
}
t_f8.sort_by(f64::total_cmp);
t_q8.sort_by(f64::total_cmp);
let (a, b) = (t_f8[reps / 2], t_q8[reps / 2]);
let flop = 2.0 * m as f64 * in_f as f64 * out_f as f64;
println!(
"{:<28} {:>12.4} {:>12.4} {:>9.3}x {:>12.1} {:>12.1}",
format!("{label} {in_f}->{out_f}"),
a * 1e3,
b * 1e3,
b / a,
flop / a / 1e12,
flop / b / 1e12
);
}
Ok(())
}