use std::process::Command;
include!("build_support/msl_postpass.rs");
fn main() {
let sha = std::env::var("CERA_GIT_SHA").ok().unwrap_or_else(|| {
Command::new("git")
.args(["rev-parse", "--short=12", "HEAD"])
.output()
.ok()
.filter(|o| o.status.success())
.and_then(|o| String::from_utf8(o.stdout).ok())
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.unwrap_or_else(|| "unknown".to_string())
});
println!("cargo:rustc-env=CERA_GIT_SHA={sha}");
println!("cargo:rerun-if-env-changed=CERA_GIT_SHA");
println!("cargo:rerun-if-changed=build_support/msl_postpass.rs");
if std::env::var_os("CARGO_FEATURE_GPU").is_some() {
compile_slang_kernels();
}
let want_wgsl = std::env::var_os("CARGO_FEATURE_GPU").is_some();
let target_os = std::env::var("CARGO_CFG_TARGET_OS").unwrap_or_default();
let is_apple = target_os == "macos" || target_os == "ios";
let want_msl = std::env::var_os("CARGO_FEATURE_METAL").is_some() && is_apple;
if want_wgsl || want_msl {
compile_slang_multitarget(want_wgsl, want_msl);
}
println!("cargo:rustc-check-cfg=cfg(has_blas)");
if std::env::var_os("CARGO_FEATURE_BLAS").is_some() && is_apple {
println!("cargo:rustc-cfg=has_blas");
}
}
const SLANG_KERNELS: &[&str] = &[
"mul_mat_reg_tile_q4_0",
"mul_mat_reg_tile_q8_0",
"mul_mat_reg_tile_q4_k",
"mul_mat_reg_tile_q5_k",
"mul_mat_reg_tile_q6_k",
];
fn compile_slang_kernels() {
let out_dir = std::env::var("OUT_DIR").expect("OUT_DIR unset");
let dir = "src/backend/shaders/spirv";
let slangc = find_slangc();
for name in SLANG_KERNELS {
let src = format!("{dir}/{name}.slang");
let committed = format!("{dir}/{name}.spv");
let out = format!("{out_dir}/{name}.spv");
println!("cargo:rerun-if-changed={src}");
println!("cargo:rerun-if-changed={committed}");
let compiled = slangc.as_ref().is_some_and(|sc| {
Command::new(sc)
.args([
src.as_str(),
"-target",
"spirv",
"-O3",
"-entry",
"main",
"-stage",
"compute",
"-o",
out.as_str(),
])
.status()
.map(|s| s.success())
.unwrap_or(false)
});
if !compiled {
if slangc.is_none() {
println!(
"cargo:warning=slangc not found; using committed {name}.spv. Set SLANGC or add ~/.local/slang/bin to PATH to recompile from {name}.slang."
);
} else {
println!(
"cargo:warning=slangc failed on {name}.slang; using committed {name}.spv."
);
}
std::fs::copy(&committed, &out)
.unwrap_or_else(|e| panic!("no compiled or committed SPIR-V for {name}: {e}"));
}
}
println!("cargo:rerun-if-env-changed=SLANGC");
}
const SLANG_MULTI_KERNELS: &[&str] = &[
"softmax",
"coopmat_probe",
"gemm_q8_0",
"bias_add",
"gelu",
"elementwise",
"rope",
"per_head_rmsnorm",
"layernorm_batch",
"rmsnorm_batch",
"argmax_f32",
"rmsnorm",
"conv1d",
"conv1d_fused",
"conv1d_fused_batch",
"exp_polar",
"overlap_add",
"activations",
"conv2d_direct",
"transpose_blocked",
"glu_split",
"chan_affine_silu",
"audio_xl_attention",
"stft_frame",
"power_spec",
"mel_project",
"mel_norm",
"moe_route",
"moe_gemv_q4_0",
"moe_combine",
"ffn_swiglu_q4_0",
"gemv_q4_0_fast",
"gemv_q4_0_qkv",
"bert_flash_attention",
];
fn compile_slang_multitarget(want_wgsl: bool, want_msl: bool) {
let out_dir = std::env::var("OUT_DIR").expect("OUT_DIR unset");
let dir = "src/backend/shaders/slang";
let slangc = find_slangc();
let mut targets: Vec<(&str, &str)> = Vec::with_capacity(2);
if want_wgsl {
targets.push(("wgsl", "wgsl"));
}
if want_msl {
targets.push(("metal", "metal"));
}
for name in SLANG_MULTI_KERNELS {
let src = format!("{dir}/{name}.slang");
println!("cargo:rerun-if-changed={src}");
let entries = slang_entry_points(&src, name);
for (target, ext) in &targets {
let committed = format!("{dir}/{name}.{ext}");
let out = format!("{out_dir}/{name}.{ext}");
println!("cargo:rerun-if-changed={committed}");
let mut args: Vec<&str> = vec![src.as_str(), "-target", target, "-O3"];
for e in &entries {
args.push("-entry");
args.push(e);
}
args.extend(["-stage", "compute", "-o", out.as_str()]);
let compiled = slangc.as_ref().is_some_and(|sc| {
Command::new(sc)
.args(&args)
.status()
.map(|s| s.success())
.unwrap_or(false)
});
if !compiled {
if slangc.is_none() {
println!(
"cargo:warning=slangc not found; using committed {name}.{ext}. Set SLANGC or add ~/.local/slang/bin to PATH to recompile from {name}.slang."
);
} else {
println!(
"cargo:warning=slangc failed on {name}.slang for target {target}; using committed {name}.{ext}."
);
}
std::fs::copy(&committed, &out)
.unwrap_or_else(|e| panic!("no compiled or committed {ext} for {name}: {e}"));
}
if *name == "gemm_q8_0" && *target == "metal" {
apply_msl_postpass(&out, name);
}
}
}
println!("cargo:rerun-if-env-changed=SLANGC");
}
fn apply_msl_postpass(path: &str, name: &str) {
let src = match std::fs::read_to_string(path) {
Ok(s) => s,
Err(e) => {
println!("cargo:warning=cannot read generated {name}.metal to post-process: {e}");
return;
}
};
match postpass_gemm_msl(&src) {
Ok(patched) => {
let tmp = format!("{path}.tmp");
let staged = std::fs::write(&tmp, patched).and_then(|()| std::fs::rename(&tmp, path));
if let Err(e) = staged {
let _ = std::fs::remove_file(&tmp);
println!("cargo:warning=cannot write post-processed {name}.metal: {e}");
}
}
Err(why) => println!(
"cargo:warning=MSL post-pass declined on {name}.metal ({why}); shipping unpatched Slang output, which is correct but ~5% slower than the patched shader on the simdgroup GEMM."
),
}
}
fn slang_entry_points(src_path: &str, basename: &str) -> Vec<String> {
let text = std::fs::read_to_string(src_path)
.unwrap_or_else(|e| panic!("failed to read Slang source {src_path}: {e}"));
let mut names: Vec<String> = Vec::new();
for line in text.lines() {
let line = line.trim_start();
if let Some(rest) = line.strip_prefix("//")
&& let Some(list) = rest.trim_start().strip_prefix("slang-entries:")
{
names.extend(list.split_whitespace().map(str::to_string));
}
}
if names.is_empty() {
vec![basename.to_string()]
} else {
names
}
}
fn find_slangc() -> Option<String> {
let home = std::env::var("HOME").unwrap_or_default();
let mut candidates: Vec<String> = Vec::new();
if let Some(p) = std::env::var_os("SLANGC") {
candidates.push(p.to_string_lossy().into_owned());
}
candidates.push("slangc".to_string());
candidates.push(format!("{home}/.local/slang/bin/slangc"));
candidates
.into_iter()
.find(|cand| Command::new(cand).arg("-v").output().is_ok())
}