Skip to content

Commit cb206f3

Browse files
committed
Merge branch 'main' of github.com:evilsocket/cake
2 parents b56dc63 + 7e73384 commit cb206f3

9 files changed

Lines changed: 1561 additions & 637 deletions

File tree

Cargo.lock

Lines changed: 4 additions & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

cake-core/Cargo.toml

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -42,10 +42,13 @@ candle-transformers = { version = "0.9" }
4242
candle-flash-attn = { version = "0.9", optional = true }
4343
candle-metal-kernels = { version = "0.9", optional = true }
4444
objc2-metal = { version = "0.3", optional = true }
45+
ash = { version = "0.38", optional = true, default-features = false, features = ["linked", "debug", "std"] }
46+
gpu-allocator = { version = "0.27", optional = true, default-features = false, features = ["vulkan"] }
4547
wgpu = { version = "24", optional = true }
4648
libloading = { version = "0.8", optional = true }
4749
pollster = { version = "0.4", optional = true }
4850
bytemuck = { version = "1", optional = true, features = ["derive"] }
51+
naga = { version = "24", optional = true, features = ["wgsl-in", "spv-out"] }
4952
half = "2"
5053
image = "0.25.2"
5154
hf-hub = "0.5"
@@ -69,7 +72,7 @@ default = ["master", "llama", "qwen2", "qwen3_5", "qwen3", "qwen3_moe", "qwen3_5
6972
metal = ["candle-core/metal", "candle-nn/metal", "candle-transformers/metal", "dep:candle-metal-kernels", "dep:objc2-metal"]
7073
cuda = ["candle-core/cuda", "candle-nn/cuda", "candle-transformers/cuda", "dep:bindgen_cuda"]
7174
flash-attn = ["cuda", "dep:candle-flash-attn"]
72-
vulkan = ["dep:wgpu", "dep:pollster", "dep:bytemuck"]
75+
vulkan = ["dep:ash", "dep:gpu-allocator", "dep:bytemuck"]
7376
rocm = ["dep:libloading"]
7477

7578
master = ["dep:actix-web", "dep:async-stream", "dep:uuid"]
@@ -92,6 +95,7 @@ luxtts = ["dep:rustfft"]
9295

9396
[build-dependencies]
9497
bindgen_cuda = { version = "0.1.6", optional = true }
98+
naga = { version = "24", features = ["wgsl-in", "spv-out"] }
9599

96100
[dev-dependencies]
97101
tempfile = "3"

cake-core/benches/bench_vulkan.rs

Lines changed: 52 additions & 44 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,7 @@
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])]
5072
fn 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])]
8389
fn 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])]
9197
fn 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])]
100106
fn 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])]
111117
fn 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]
131138
fn 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]
146154
fn 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
}

cake-core/build.rs

Lines changed: 60 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,64 @@
11
fn main() {
2+
#[cfg(feature = "vulkan")]
3+
{
4+
println!("cargo::rerun-if-changed=src/backends/vulkan/ops.wgsl");
5+
6+
let wgsl_src = std::fs::read_to_string("src/backends/vulkan/ops.wgsl")
7+
.expect("failed to read ops.wgsl");
8+
let module = naga::front::wgsl::parse_str(&wgsl_src)
9+
.expect("failed to parse WGSL");
10+
let info = naga::valid::Validator::new(
11+
naga::valid::ValidationFlags::all(),
12+
naga::valid::Capabilities::empty(),
13+
)
14+
.validate(&module)
15+
.expect("WGSL validation failed");
16+
17+
let options = naga::back::spv::Options {
18+
lang_version: (1, 3),
19+
..Default::default()
20+
};
21+
// Generate one SPIR-V module per entry point
22+
let entry_points: Vec<String> = module
23+
.entry_points
24+
.iter()
25+
.map(|ep| ep.name.clone())
26+
.collect();
27+
28+
let out_dir = std::path::PathBuf::from(std::env::var("OUT_DIR").unwrap());
29+
let mut includes = String::new();
30+
31+
for ep_name in &entry_points {
32+
let pipeline_options = naga::back::spv::PipelineOptions {
33+
shader_stage: naga::ShaderStage::Compute,
34+
entry_point: ep_name.clone(),
35+
};
36+
let spv_words = naga::back::spv::write_vec(
37+
&module,
38+
&info,
39+
&options,
40+
Some(&pipeline_options),
41+
)
42+
.unwrap_or_else(|e| panic!("SPIR-V generation failed for {ep_name}: {e}"));
43+
44+
// Write SPIR-V binary
45+
let spv_path = out_dir.join(format!("ops_{ep_name}.spv"));
46+
let bytes: Vec<u8> = spv_words.iter().flat_map(|w| w.to_le_bytes()).collect();
47+
std::fs::write(&spv_path, &bytes).unwrap();
48+
49+
includes.push_str(&format!(
50+
"(\"{ep_name}\", include_bytes!(concat!(env!(\"OUT_DIR\"), \"/ops_{ep_name}.spv\"))),\n"
51+
));
52+
}
53+
54+
// Write a Rust file with all SPIR-V modules
55+
let rs_path = out_dir.join("spirv_ops.rs");
56+
let code = format!(
57+
"static SPIRV_MODULES: &[(&str, &[u8])] = &[\n{includes}];\n"
58+
);
59+
std::fs::write(&rs_path, &code).expect("failed to write SPIR-V module list");
60+
}
61+
262
#[cfg(feature = "cuda")]
363
{
464
println!("cargo::rerun-if-changed=src/backends/cuda/ops.cu");

0 commit comments

Comments
 (0)