use super::model::MlxModelWeights;
use crate::debug::INVESTIGATION_ENV;
pub(super) fn profiling_enabled() -> bool {
INVESTIGATION_ENV.mlx_profile
}
#[derive(Default, Clone)]
pub struct KernelTypeProfile {
pub qkv_matmuls_us: Vec<f64>,
pub head_norms_rope_us: Vec<f64>,
pub kv_cache_copy_us: Vec<f64>,
pub sdpa_us: Vec<f64>,
pub o_proj_us: Vec<f64>,
pub mlp_matmuls_us: Vec<f64>,
pub moe_us: Vec<f64>,
pub norms_adds_us: Vec<f64>,
pub lm_head_us: f64,
}
#[derive(Default, Clone)]
pub struct TokenProfile {
pub layer_s1_us: Vec<f64>, pub layer_cpu1_us: Vec<f64>, pub layer_s2_us: Vec<f64>, pub layer_cpu2_us: Vec<f64>, pub layer_s3_us: Vec<f64>, pub layer_cpu3_us: Vec<f64>, pub layer_s4_us: Vec<f64>, pub layer_cpu4_us: Vec<f64>, pub head_session_us: f64, pub head_cpu_us: f64, pub total_us: f64,
pub s1_dispatches: Vec<usize>,
pub s2_dispatches: Vec<usize>,
pub s3_dispatches: Vec<usize>,
pub s4_dispatches: Vec<usize>,
pub head_dispatches: usize,
}
pub(super) fn merge_profiles(main: &mut Option<TokenProfile>, worker: Option<TokenProfile>) {
let (Some(m), Some(w)) = (main.as_mut(), worker) else {
return; };
for (mv, wv) in [
(&mut m.layer_s1_us, &w.layer_s1_us),
(&mut m.layer_cpu1_us, &w.layer_cpu1_us),
(&mut m.layer_s2_us, &w.layer_s2_us),
(&mut m.layer_cpu2_us, &w.layer_cpu2_us),
(&mut m.layer_s3_us, &w.layer_s3_us),
(&mut m.layer_cpu3_us, &w.layer_cpu3_us),
(&mut m.layer_s4_us, &w.layer_s4_us),
(&mut m.layer_cpu4_us, &w.layer_cpu4_us),
] {
for (mi, wi) in mv.iter_mut().zip(wv.iter()) {
*mi += *wi;
}
}
for (mv, wv) in [
(&mut m.s1_dispatches, &w.s1_dispatches),
(&mut m.s2_dispatches, &w.s2_dispatches),
(&mut m.s3_dispatches, &w.s3_dispatches),
(&mut m.s4_dispatches, &w.s4_dispatches),
] {
for (mi, wi) in mv.iter_mut().zip(wv.iter()) {
*mi += *wi;
}
}
}
pub struct ProfileAccumulator {
pub tokens: Vec<TokenProfile>,
pub warmup_count: usize,
pub enabled: bool,
}
impl ProfileAccumulator {
pub fn new(warmup: usize) -> Self {
Self {
tokens: Vec::new(),
warmup_count: warmup,
enabled: profiling_enabled(),
}
}
pub fn start_token(&self) -> Option<TokenProfile> {
if self.enabled {
Some(TokenProfile::default())
} else {
None
}
}
pub fn finish_token(&mut self, profile: Option<TokenProfile>) {
if let Some(p) = profile {
self.tokens.push(p);
}
}
pub fn print_summary(&self) {
if !self.enabled || self.tokens.is_empty() {
return;
}
let skip = self.warmup_count.min(self.tokens.len().saturating_sub(1));
let measured: Vec<&TokenProfile> = self.tokens.iter().skip(skip).collect();
if measured.is_empty() {
eprintln!("[PROFILE] No tokens after warmup to report.");
return;
}
let n = measured.len();
let num_layers = measured[0].layer_s1_us.len();
eprintln!("\n╔══════════════════════════════════════════════════════════╗");
eprintln!("║ MLX-NATIVE FORWARD PASS PROFILE ({n} tokens, {skip} warmup skipped) ║");
eprintln!("╠══════════════════════════════════════════════════════════╣");
let avg = |getter: &dyn Fn(&TokenProfile) -> &Vec<f64>| -> f64 {
let total: f64 = measured.iter().map(|t| getter(t).iter().sum::<f64>()).sum();
total / n as f64
};
let s1_avg = avg(&|t| &t.layer_s1_us);
let cpu1_avg = avg(&|t| &t.layer_cpu1_us);
let s2_avg = avg(&|t| &t.layer_s2_us);
let cpu2_avg = avg(&|t| &t.layer_cpu2_us);
let s3_avg = avg(&|t| &t.layer_s3_us);
let cpu3_avg = avg(&|t| &t.layer_cpu3_us);
let s4_avg = avg(&|t| &t.layer_s4_us);
let cpu4_avg = avg(&|t| &t.layer_cpu4_us);
let head_gpu_avg: f64 = measured.iter().map(|t| t.head_session_us).sum::<f64>() / n as f64;
let head_cpu_avg: f64 = measured.iter().map(|t| t.head_cpu_us).sum::<f64>() / n as f64;
let total_avg: f64 = measured.iter().map(|t| t.total_us).sum::<f64>() / n as f64;
let gpu_total = s1_avg + s2_avg + s3_avg + s4_avg + head_gpu_avg;
let cpu_total = cpu1_avg + cpu2_avg + cpu3_avg + cpu4_avg + head_cpu_avg;
let actual_sessions = if s2_avg + s3_avg + s4_avg + head_gpu_avg < 1.0 {
1 } else {
num_layers * 2 + 1
};
eprintln!(
"║ {} session(s)/token (single-session mode)",
actual_sessions
);
eprintln!("║");
eprintln!("║ Session breakdown (avg across {num_layers} layers, {n} tokens):");
eprintln!(
"║ S1 (QKV+attn+MLP): {:8.1} us ({:5.2} ms total)",
s1_avg / num_layers as f64,
s1_avg / 1000.0
);
eprintln!(
"║ CPU1 (eliminated): {:8.1} us ({:5.2} ms total)",
cpu1_avg / num_layers as f64,
cpu1_avg / 1000.0
);
eprintln!(
"║ S2 (SDPA+MLP): {:8.1} us ({:5.2} ms total)",
s2_avg / num_layers as f64,
s2_avg / 1000.0
);
eprintln!(
"║ CPU2 (post-FF): {:8.1} us ({:5.2} ms total)",
cpu2_avg / num_layers as f64,
cpu2_avg / 1000.0
);
eprintln!(
"║ S3 (router proj): {:8.1} us ({:5.2} ms total)",
s3_avg / num_layers as f64,
s3_avg / 1000.0
);
eprintln!(
"║ CPU3 (softmax+topk): {:8.1} us ({:5.2} ms total)",
cpu3_avg / num_layers as f64,
cpu3_avg / 1000.0
);
eprintln!(
"║ S4 (MoE experts): {:8.1} us ({:5.2} ms total)",
s4_avg / num_layers as f64,
s4_avg / 1000.0
);
eprintln!(
"║ CPU4 (post-MoE): {:8.1} us ({:5.2} ms total)",
cpu4_avg / num_layers as f64,
cpu4_avg / 1000.0
);
eprintln!(
"║ Head GPU: {:8.1} us ({:5.2} ms)",
head_gpu_avg,
head_gpu_avg / 1000.0
);
eprintln!(
"║ Head CPU: {:8.1} us ({:5.2} ms)",
head_cpu_avg,
head_cpu_avg / 1000.0
);
eprintln!("║");
eprintln!(
"║ Total: {:8.1} us ({:5.2} ms)",
total_avg,
total_avg / 1000.0
);
eprintln!(
"║ GPU sessions: {:8.1} us ({:5.1}%)",
gpu_total,
gpu_total / total_avg * 100.0
);
eprintln!(
"║ CPU ops: {:8.1} us ({:5.1}%)",
cpu_total,
cpu_total / total_avg * 100.0
);
let overhead = total_avg - gpu_total - cpu_total;
if overhead.abs() > 10.0 {
eprintln!(
"║ Unaccounted: {:8.1} us ({:5.1}%)",
overhead,
overhead / total_avg * 100.0
);
}
let last_layer_dispatch_avg = |getter: &dyn Fn(&TokenProfile) -> &Vec<usize>| -> f64 {
let total: usize = measured
.iter()
.map(|t| getter(t).last().copied().unwrap_or(0))
.sum();
total as f64 / n as f64
};
let s1_disp = last_layer_dispatch_avg(&|t| &t.s1_dispatches);
let s2_disp = last_layer_dispatch_avg(&|t| &t.s2_dispatches);
let s3_disp = last_layer_dispatch_avg(&|t| &t.s3_dispatches);
let s4_disp = last_layer_dispatch_avg(&|t| &t.s4_dispatches);
let total_token_disp: f64 = measured
.iter()
.map(|t| t.head_dispatches as f64)
.sum::<f64>()
/ n as f64;
let body_cum = s1_disp + s2_disp + s3_disp + s4_disp;
let head_disp = (total_token_disp - body_cum).max(0.0);
let total_disp = total_token_disp;
eprintln!("║");
eprintln!("║ Dispatch counts per token:");
eprintln!("║ S1: {s1_disp:.0} S2: {s2_disp:.0} S3: {s3_disp:.0} S4: {s4_disp:.0} Head: {head_disp:.0}");
eprintln!("║ Total: {total_disp:.0} dispatches/token");
eprintln!("║ (candle Phase 0 baseline: ~105 dispatches/token)");
eprintln!("║ Ratio: {:.1}x more dispatches", total_disp / 105.0);
eprintln!("║");
eprintln!("║ Per-layer detail (avg over {n} tokens, us):");
eprintln!("║ Layer | S1 | CPU1 | S2 | CPU2 | S3 | CPU3 | S4 | CPU4 | Total");
eprintln!("║ ------|--------|--------|--------|--------|--------|--------|--------|--------|------");
let detail_layers: Vec<usize> = {
let mut v: Vec<usize> = (0..3.min(num_layers)).collect();
if num_layers > 3 {
v.push(num_layers - 1);
}
v
};
for &li in &detail_layers {
let s1: f64 = measured.iter().map(|t| t.layer_s1_us[li]).sum::<f64>() / n as f64;
let c1: f64 = measured.iter().map(|t| t.layer_cpu1_us[li]).sum::<f64>() / n as f64;
let s2: f64 = measured.iter().map(|t| t.layer_s2_us[li]).sum::<f64>() / n as f64;
let c2: f64 = measured.iter().map(|t| t.layer_cpu2_us[li]).sum::<f64>() / n as f64;
let s3: f64 = measured.iter().map(|t| t.layer_s3_us[li]).sum::<f64>() / n as f64;
let c3: f64 = measured.iter().map(|t| t.layer_cpu3_us[li]).sum::<f64>() / n as f64;
let s4: f64 = measured.iter().map(|t| t.layer_s4_us[li]).sum::<f64>() / n as f64;
let c4: f64 = measured.iter().map(|t| t.layer_cpu4_us[li]).sum::<f64>() / n as f64;
let layer_type = if (li + 1) % 6 == 0 { "G" } else { "S" };
eprintln!("║ {:>2} ({}) | {:6.0} | {:6.0} | {:6.0} | {:6.0} | {:6.0} | {:6.0} | {:6.0} | {:6.0} | {:6.0}",
li, layer_type, s1, c1, s2, c2, s3, c3, s4, c4,
s1 + c1 + s2 + c2 + s3 + c3 + s4 + c4);
}
eprintln!("║");
eprintln!("║ Avg time per dispatch (session_time / dispatches):");
if s1_disp > 0.0 {
eprintln!("║ S1 (QKV+attn+MLP): {:.1} us/dispatch", s1_avg / s1_disp);
}
if s2_disp > 0.0 {
eprintln!("║ S2 (unused): {:.1} us/dispatch", s2_avg / s2_disp);
}
if s3_disp > 0.0 {
eprintln!("║ S3 (router): {:.1} us/dispatch", s3_avg / s3_disp);
}
if s4_disp > 0.0 {
eprintln!("║ S4 (MoE): {:.1} us/dispatch", s4_avg / s4_disp);
}
eprintln!("╚══════════════════════════════════════════════════════════╝");
}
}
impl MlxModelWeights {
pub fn print_kernel_profile_report(profiles: &[KernelTypeProfile]) {
if profiles.is_empty() {
eprintln!("[KERNEL_PROFILE] No tokens to report.");
return;
}
let n = profiles.len();
let num_layers = profiles[0].qkv_matmuls_us.len();
let median_sum = |getter: &dyn Fn(&KernelTypeProfile) -> &Vec<f64>| -> f64 {
let mut sums: Vec<f64> = profiles
.iter()
.map(|p| getter(p).iter().sum::<f64>())
.collect();
sums.sort_by(|a, b| a.partial_cmp(b).unwrap());
sums[sums.len() / 2]
};
let qkv_total = median_sum(&|p| &p.qkv_matmuls_us);
let norms_rope_total = median_sum(&|p| &p.head_norms_rope_us);
let kv_cache_total = median_sum(&|p| &p.kv_cache_copy_us);
let sdpa_total = median_sum(&|p| &p.sdpa_us);
let o_proj_total = median_sum(&|p| &p.o_proj_us);
let mlp_total = median_sum(&|p| &p.mlp_matmuls_us);
let moe_total = median_sum(&|p| &p.moe_us);
let norms_adds_total = median_sum(&|p| &p.norms_adds_us);
let mut head_vals: Vec<f64> = profiles.iter().map(|p| p.lm_head_us).collect();
head_vals.sort_by(|a, b| a.partial_cmp(b).unwrap());
let head_total = head_vals[head_vals.len() / 2];
let gpu_total = qkv_total
+ norms_rope_total
+ kv_cache_total
+ sdpa_total
+ o_proj_total
+ mlp_total
+ moe_total
+ norms_adds_total
+ head_total;
let qkv_per_layer = qkv_total / num_layers as f64;
let norms_rope_per_layer = norms_rope_total / num_layers as f64;
let kv_cache_per_layer = kv_cache_total / num_layers as f64;
let sdpa_per_layer = sdpa_total / num_layers as f64;
let o_proj_per_layer = o_proj_total / num_layers as f64;
let mlp_per_layer = mlp_total / num_layers as f64;
let moe_per_layer = moe_total / num_layers as f64;
let norms_adds_per_layer = norms_adds_total / num_layers as f64;
let candle_qkv_per_layer = 37.0; let candle_norms_rope_per_layer = 24.0; let candle_kv_cache_per_layer = 4.0; let candle_sdpa_per_layer = 17.0; let candle_o_proj_per_layer = 11.0; let candle_mlp_per_layer = 43.0; let candle_moe_per_layer = 127.0; let candle_norms_adds_per_layer = 20.0; let candle_lm_head = 185.0;
let candle_per_layer_total = candle_qkv_per_layer
+ candle_norms_rope_per_layer
+ candle_kv_cache_per_layer
+ candle_sdpa_per_layer
+ candle_o_proj_per_layer
+ candle_mlp_per_layer
+ candle_moe_per_layer
+ candle_norms_adds_per_layer;
let candle_layers_total = candle_per_layer_total * num_layers as f64;
let candle_total_reconstructed = candle_layers_total + candle_lm_head;
eprintln!("\n=== PER-KERNEL-TYPE PROFILING (median over {n} tokens) ===");
eprintln!("Per layer ({num_layers} layers):");
eprintln!(
" QKV matmuls (norm+3 proj): {:7.0} us [candle: ~{:.0} us] ratio: {:.1}x",
qkv_per_layer,
candle_qkv_per_layer,
qkv_per_layer / candle_qkv_per_layer
);
eprintln!(
" Head norms + RoPE (3 dispatches): {:7.0} us [candle: ~{:.0} us] ratio: {:.1}x",
norms_rope_per_layer,
candle_norms_rope_per_layer,
norms_rope_per_layer / candle_norms_rope_per_layer
);
eprintln!(
" KV cache copy (2 dispatches): {:7.0} us [candle: ~{:.0} us] ratio: {:.1}x",
kv_cache_per_layer,
candle_kv_cache_per_layer,
kv_cache_per_layer / candle_kv_cache_per_layer
);
eprintln!(
" SDPA (1 dispatch): {:7.0} us [candle: ~{:.0} us] ratio: {:.1}x",
sdpa_per_layer,
candle_sdpa_per_layer,
sdpa_per_layer / candle_sdpa_per_layer
);
eprintln!(
" O-proj matmul (1 dispatch): {:7.0} us [candle: ~{:.0} us] ratio: {:.1}x",
o_proj_per_layer,
candle_o_proj_per_layer,
o_proj_per_layer / candle_o_proj_per_layer
);
eprintln!(
" MLP matmuls (norm+3proj+gelu): {:7.0} us [candle: ~{:.0} us] ratio: {:.1}x",
mlp_per_layer,
candle_mlp_per_layer,
mlp_per_layer / candle_mlp_per_layer
);
eprintln!(
" MoE (routing+4 expert): {:7.0} us [candle: ~{:.0} us] ratio: {:.1}x",
moe_per_layer,
candle_moe_per_layer,
moe_per_layer / candle_moe_per_layer
);
eprintln!(
" Fused norms/adds (2 dispatches): {:7.0} us [candle: ~{:.0} us] ratio: {:.1}x",
norms_adds_per_layer,
candle_norms_adds_per_layer,
norms_adds_per_layer / candle_norms_adds_per_layer
);
eprintln!();
eprintln!("Head:");
eprintln!(
" lm_head GEMM (F16): {:7.0} us [candle: ~{:.0} us] ratio: {:.1}x",
head_total,
candle_lm_head,
head_total / candle_lm_head
);
eprintln!();
eprintln!(
"Total GPU per token: {:7.0} us [candle: ~{:.0} us] ratio: {:.1}x",
gpu_total,
candle_total_reconstructed,
gpu_total / candle_total_reconstructed
);
eprintln!(
" Layers total: {:7.0} us [candle: ~{:.0} us]",
gpu_total - head_total,
candle_layers_total
);
eprintln!(
" Head total: {:7.0} us [candle: ~{:.0} us]",
head_total, candle_lm_head
);
eprintln!();
eprintln!("Per-layer detail (median token, us):");
eprintln!(" Layer | Type | QKV | Nrm+RoPE | KV$ | SDPA | O-proj | MLP | MoE | Norms | Total");
eprintln!(" ------|------|--------|----------|------|-------|--------|--------|--------|-------|------");
let mid = profiles.len() / 2;
let median_p = &profiles[mid]; for li in 0..num_layers {
let lt = if (li + 1) % 6 == 0 { "G" } else { "S" };
let layer_total = median_p.qkv_matmuls_us[li]
+ median_p.head_norms_rope_us[li]
+ median_p.kv_cache_copy_us[li]
+ median_p.sdpa_us[li]
+ median_p.o_proj_us[li]
+ median_p.mlp_matmuls_us[li]
+ median_p.moe_us[li]
+ median_p.norms_adds_us[li];
eprintln!(" {:>2} | {} | {:6.0} | {:5.0} | {:4.0} | {:5.0} | {:5.0} | {:5.0} | {:5.0} | {:5.0} | {:5.0}",
li, lt,
median_p.qkv_matmuls_us[li], median_p.head_norms_rope_us[li],
median_p.kv_cache_copy_us[li], median_p.sdpa_us[li],
median_p.o_proj_us[li], median_p.mlp_matmuls_us[li],
median_p.moe_us[li], median_p.norms_adds_us[li],
layer_total);
}
let mut ratios = vec![
(
"QKV matmuls",
qkv_per_layer,
candle_qkv_per_layer,
qkv_per_layer / candle_qkv_per_layer,
),
(
"Head norms + RoPE",
norms_rope_per_layer,
candle_norms_rope_per_layer,
norms_rope_per_layer / candle_norms_rope_per_layer,
),
(
"KV cache copy",
kv_cache_per_layer,
candle_kv_cache_per_layer,
kv_cache_per_layer / candle_kv_cache_per_layer,
),
(
"SDPA",
sdpa_per_layer,
candle_sdpa_per_layer,
sdpa_per_layer / candle_sdpa_per_layer,
),
(
"O-proj matmul",
o_proj_per_layer,
candle_o_proj_per_layer,
o_proj_per_layer / candle_o_proj_per_layer,
),
(
"MLP matmuls",
mlp_per_layer,
candle_mlp_per_layer,
mlp_per_layer / candle_mlp_per_layer,
),
(
"MoE",
moe_per_layer,
candle_moe_per_layer,
moe_per_layer / candle_moe_per_layer,
),
(
"Fused norms/adds",
norms_adds_per_layer,
candle_norms_adds_per_layer,
norms_adds_per_layer / candle_norms_adds_per_layer,
),
(
"lm_head GEMM",
head_total,
candle_lm_head,
head_total / candle_lm_head,
),
];
ratios.sort_by(|a, b| b.3.partial_cmp(&a.3).unwrap());
eprintln!();
eprintln!("TOP 3 SLOWEST (highest mlx-native/candle ratio):");
for (i, (name, mlx_us, candle_us, ratio)) in ratios.iter().take(3).enumerate() {
let overhead_per_token = (mlx_us - candle_us)
* if *name != "lm_head GEMM" {
num_layers as f64
} else {
1.0
};
eprintln!(
" {}. {} — {:.1}x slower ({:.0} vs {:.0} us/layer) — {:.0} us/token overhead",
i + 1,
name,
ratio,
mlx_us,
candle_us,
overhead_per_token
);
}
eprintln!();
eprintln!("NOTE: Per-session overhead (~30-50 us/session) inflates all groups.");
eprintln!(" The ratio shows relative slowness, not absolute kernel time.");
eprintln!(
" {} sessions/token vs 1 in production mode.",
8 * num_layers + 2
);
}
}