1- /// Vulkan backend benchmarks — GPU GEMV vs CPU matmul at model-realistic sizes,
2- /// dispatch overhead, upload/download costs, and elementwise ops.
1+ /// Vulkan backend benchmarks — GPU matmul, elementwise ops, and full MLP.
2+ ///
3+ /// All "vulkan_*" benchmarks exercise the GPU compute path.
4+ /// All "cpu_*" benchmarks use CPU-only candle ops for comparison.
35///
46/// Run on Steam Deck: `cargo bench -p cake-core --features vulkan -- vulkan`
57
@@ -25,27 +27,47 @@ fn vulkan_dispatch_overhead(bencher: divan::Bencher) {
2527 bencher. bench_local ( || backend. silu_mul ( & a, & b) . unwrap ( ) ) ;
2628}
2729
28- // ── GPU GEMV vs CPU matmul at model sizes ────────────────────────────
30+ // ── GPU matmul at model sizes ──────────── ────────────────────────────
2931// Qwen3-0.6B: hidden=1024, intermediate=3072, head_dim=128
30- // QKV: (1,1024) × (1024,4096), O: (1,1024) × (1024,1024)
31- // gate_up: (1,1024) × (1024,6144), down: (1,3072) × (3072,1024)
32+ // M=2 is the smallest prefill batch that goes through GPU (M>1).
3233
33- #[ divan:: bench( args = [ 1024 , 4096 , 6144 ] ) ]
34- fn vulkan_gemv_1024xN ( bencher : divan:: Bencher , n : usize ) {
34+ #[ divan:: bench( args = [ 2 , 8 , 32 , 64 ] ) ]
35+ fn vulkan_gemm_Mx1024x4096 ( bencher : divan:: Bencher , m : usize ) {
3536 let backend = vk ( ) ;
36- let a = cpu_tensor ( & [ 1 , 1024 ] , 1100 ) ;
37- let b = cpu_tensor ( & [ 1024 , n ] , 1101 ) ;
37+ let a = cpu_tensor ( & [ m , 1024 ] , 1300 ) ;
38+ let b = cpu_tensor ( & [ 1024 , 4096 ] , 1301 ) ;
3839 bencher. bench_local ( || backend. matmul ( & a, & b) . unwrap ( ) ) ;
3940}
4041
41- #[ divan:: bench]
42- fn vulkan_gemv_3072x1024 ( bencher : divan:: Bencher ) {
42+ #[ divan:: bench( args = [ 2 , 8 , 32 , 64 ] ) ]
43+ fn cpu_gemm_Mx1024x4096 ( bencher : divan:: Bencher , m : usize ) {
44+ let a = cpu_tensor ( & [ m, 1024 ] , 1300 ) ;
45+ let b = cpu_tensor ( & [ 1024 , 4096 ] , 1301 ) ;
46+ bencher. bench_local ( || a. matmul ( & b) . unwrap ( ) ) ;
47+ }
48+
49+ // ── GPU matmul at other model shapes ─────────────────────────────────
50+ // gate_up: Mx1024x6144, down: Mx3072x1024
51+
52+ #[ divan:: bench( args = [ 2 , 8 , 32 ] ) ]
53+ fn vulkan_gemm_Mx1024x6144 ( bencher : divan:: Bencher , m : usize ) {
4354 let backend = vk ( ) ;
44- let a = cpu_tensor ( & [ 1 , 3072 ] , 1200 ) ;
55+ let a = cpu_tensor ( & [ m, 1024 ] , 1100 ) ;
56+ let b = cpu_tensor ( & [ 1024 , 6144 ] , 1101 ) ;
57+ bencher. bench_local ( || backend. matmul ( & a, & b) . unwrap ( ) ) ;
58+ }
59+
60+ #[ divan:: bench( args = [ 2 , 8 , 32 ] ) ]
61+ fn vulkan_gemm_Mx3072x1024 ( bencher : divan:: Bencher , m : usize ) {
62+ let backend = vk ( ) ;
63+ let a = cpu_tensor ( & [ m, 3072 ] , 1200 ) ;
4564 let b = cpu_tensor ( & [ 3072 , 1024 ] , 1201 ) ;
4665 bencher. bench_local ( || backend. matmul ( & a, & b) . unwrap ( ) ) ;
4766}
4867
68+ // ── CPU generation (M=1) baseline ────────────────────────────────────
69+ // M=1 uses CPU fallback (dispatch overhead > compute gain).
70+
4971#[ divan:: bench( args = [ 1024 , 4096 , 6144 ] ) ]
5072fn cpu_gemv_1024xN ( bencher : divan:: Bencher , n : usize ) {
5173 let a = cpu_tensor ( & [ 1 , 1024 ] , 1100 ) ;
@@ -60,34 +82,18 @@ fn cpu_gemv_3072x1024(bencher: divan::Bencher) {
6082 bencher. bench_local ( || a. matmul ( & b) . unwrap ( ) ) ;
6183}
6284
63- // ── GPU GEMM (prefill) at model sizes ────────────────────────────────
64-
65- #[ divan:: bench( args = [ 8 , 32 , 64 ] ) ]
66- fn vulkan_gemm_Mx1024x4096 ( bencher : divan:: Bencher , m : usize ) {
67- let backend = vk ( ) ;
68- let a = cpu_tensor ( & [ m, 1024 ] , 1300 ) ;
69- let b = cpu_tensor ( & [ 1024 , 4096 ] , 1301 ) ;
70- bencher. bench_local ( || backend. matmul ( & a, & b) . unwrap ( ) ) ;
71- }
85+ // ── Elementwise ops — GPU path (large tensors) ──────────────────────
86+ // Above 8192 element threshold to ensure GPU dispatch.
7287
73- #[ divan:: bench( args = [ 8 , 32 , 64 ] ) ]
74- fn cpu_gemm_Mx1024x4096 ( bencher : divan:: Bencher , m : usize ) {
75- let a = cpu_tensor ( & [ m, 1024 ] , 1300 ) ;
76- let b = cpu_tensor ( & [ 1024 , 4096 ] , 1301 ) ;
77- bencher. bench_local ( || a. matmul ( & b) . unwrap ( ) ) ;
78- }
79-
80- // ── Elementwise ops at model sizes ───────────────────────────────────
81-
82- #[ divan:: bench( args = [ 1024 , 3072 ] ) ]
88+ #[ divan:: bench( args = [ 16384 , 32768 ] ) ]
8389fn vulkan_silu_mul ( bencher : divan:: Bencher , size : usize ) {
8490 let backend = vk ( ) ;
8591 let gate = cpu_tensor ( & [ 1 , 1 , size] , 1400 ) ;
8692 let up = cpu_tensor ( & [ 1 , 1 , size] , 1401 ) ;
8793 bencher. bench_local ( || backend. silu_mul ( & gate, & up) . unwrap ( ) ) ;
8894}
8995
90- #[ divan:: bench( args = [ 1024 , 3072 ] ) ]
96+ #[ divan:: bench( args = [ 16384 , 32768 ] ) ]
9197fn cpu_silu_mul ( bencher : divan:: Bencher , size : usize ) {
9298 let gate = cpu_tensor ( & [ 1 , 1 , size] , 1400 ) ;
9399 let up = cpu_tensor ( & [ 1 , 1 , size] , 1401 ) ;
@@ -96,7 +102,7 @@ fn cpu_silu_mul(bencher: divan::Bencher, size: usize) {
96102 } ) ;
97103}
98104
99- #[ divan:: bench( args = [ 1024 , 3072 ] ) ]
105+ #[ divan:: bench( args = [ 16384 , 32768 ] ) ]
100106fn vulkan_add3 ( bencher : divan:: Bencher , size : usize ) {
101107 let backend = vk ( ) ;
102108 let a = cpu_tensor ( & [ 1 , 1 , size] , 1500 ) ;
@@ -105,7 +111,7 @@ fn vulkan_add3(bencher: divan::Bencher, size: usize) {
105111 bencher. bench_local ( || backend. add3 ( & a, & b, & c) . unwrap ( ) ) ;
106112}
107113
108- // ── RMS norm (CPU-only in current backend) ───── ──────────────────────
114+ // ── RMS norm (CPU fallback in current backend) ──────────────────────
109115
110116#[ divan:: bench( args = [ 1024 , 3072 ] ) ]
111117fn vulkan_rms_norm_gated ( bencher : divan:: Bencher , size : usize ) {
@@ -126,32 +132,34 @@ fn vulkan_add_rms_norm(bencher: divan::Bencher, size: usize) {
126132}
127133
128134// ── Full MLP pass (gate_up + silu_mul + down) ────────────────────────
135+ // Prefill MLP at M=8: all ops go through GPU.
129136
130137#[ divan:: bench]
131138fn vulkan_mlp_full ( bencher : divan:: Bencher ) {
132139 let backend = vk ( ) ;
133- let x = cpu_tensor ( & [ 1 , 1024 ] , 1800 ) ;
134- let gate_up_w = cpu_tensor ( & [ 6144 , 1024 ] , 1801 ) ;
135- let down_w = cpu_tensor ( & [ 1024 , 3072 ] , 1802 ) ;
140+ let x = cpu_tensor ( & [ 8 , 1024 ] , 1800 ) ;
141+ // Pre-transpose weights so they get cached on GPU across iterations
142+ let gate_up_wt = cpu_tensor ( & [ 6144 , 1024 ] , 1801 ) . t ( ) . unwrap ( ) . contiguous ( ) . unwrap ( ) ;
143+ let down_wt = cpu_tensor ( & [ 1024 , 3072 ] , 1802 ) . t ( ) . unwrap ( ) . contiguous ( ) . unwrap ( ) ;
136144 bencher. bench_local ( || {
137- let fused = backend. matmul ( & x, & gate_up_w . t ( ) . unwrap ( ) ) . unwrap ( ) ;
145+ let fused = backend. matmul ( & x, & gate_up_wt ) . unwrap ( ) ;
138146 let gate = fused. narrow ( 1 , 0 , 3072 ) . unwrap ( ) . contiguous ( ) . unwrap ( ) ;
139147 let up = fused. narrow ( 1 , 3072 , 3072 ) . unwrap ( ) . contiguous ( ) . unwrap ( ) ;
140148 let act = backend. silu_mul ( & gate, & up) . unwrap ( ) ;
141- backend. matmul ( & act, & down_w . t ( ) . unwrap ( ) ) . unwrap ( )
149+ backend. matmul ( & act, & down_wt ) . unwrap ( )
142150 } ) ;
143151}
144152
145153#[ divan:: bench]
146154fn cpu_mlp_full ( bencher : divan:: Bencher ) {
147- let x = cpu_tensor ( & [ 1 , 1024 ] , 1800 ) ;
148- let gate_up_w = cpu_tensor ( & [ 6144 , 1024 ] , 1801 ) ;
149- let down_w = cpu_tensor ( & [ 1024 , 3072 ] , 1802 ) ;
155+ let x = cpu_tensor ( & [ 8 , 1024 ] , 1800 ) ;
156+ let gate_up_wt = cpu_tensor ( & [ 6144 , 1024 ] , 1801 ) . t ( ) . unwrap ( ) . contiguous ( ) . unwrap ( ) ;
157+ let down_wt = cpu_tensor ( & [ 1024 , 3072 ] , 1802 ) . t ( ) . unwrap ( ) . contiguous ( ) . unwrap ( ) ;
150158 bencher. bench_local ( || {
151- let fused = x. matmul ( & gate_up_w . t ( ) . unwrap ( ) ) . unwrap ( ) ;
159+ let fused = x. matmul ( & gate_up_wt ) . unwrap ( ) ;
152160 let gate = fused. narrow ( 1 , 0 , 3072 ) . unwrap ( ) ;
153161 let up = fused. narrow ( 1 , 3072 , 3072 ) . unwrap ( ) ;
154162 let act = ( candle_nn:: ops:: silu ( & gate) . unwrap ( ) * & up) . unwrap ( ) ;
155- act. matmul ( & down_w . t ( ) . unwrap ( ) ) . unwrap ( )
163+ act. matmul ( & down_wt ) . unwrap ( )
156164 } ) ;
157165}
0 commit comments