use std::path::PathBuf;
use std::process::Command;
fn nvcc_version(p: &std::path::Path) -> Option<(u32, u32)> {
let out = Command::new(p).arg("--version").output().ok()?;
if !out.status.success() {
return None;
}
let s = String::from_utf8_lossy(&out.stdout);
let rel = s.split("release ").nth(1)?;
let rel = rel.split(',').next()?.trim();
let mut it = rel.split('.');
let maj = it.next()?.parse::<u32>().ok()?;
let min = it.next().and_then(|m| m.parse::<u32>().ok()).unwrap_or(0);
Some((maj, min))
}
fn resolve_nvcc() -> String {
println!("cargo:rerun-if-env-changed=MEMRA_NVCC");
println!("cargo:rerun-if-env-changed=CUDA_HOME");
println!("cargo:rerun-if-env-changed=CUDA_PATH");
println!("cargo:rerun-if-env-changed=CUDA_ROOT");
if let Ok(p) = std::env::var("MEMRA_NVCC") {
println!("cargo:warning=nvcc from MEMRA_NVCC: {p}");
return p;
}
for var in ["CUDA_HOME", "CUDA_PATH", "CUDA_ROOT"] {
if let Ok(root) = std::env::var(var) {
let cand = PathBuf::from(&root).join("bin/nvcc");
if cand.is_file() {
println!("cargo:warning=nvcc from ${var}: {}", cand.display());
return cand.to_string_lossy().into_owned();
}
}
}
let mut cands: Vec<PathBuf> = Vec::new();
if let Ok(path) = std::env::var("PATH") {
cands.extend(
path.split(':')
.filter(|d| !d.is_empty())
.map(|d| PathBuf::from(d).join("nvcc")),
);
}
cands.push(PathBuf::from("/usr/local/cuda/bin/nvcc"));
if let Ok(rd) = std::fs::read_dir("/usr/local") {
for ent in rd.flatten() {
if ent.file_name().to_string_lossy().starts_with("cuda-") {
cands.push(ent.path().join("bin/nvcc"));
}
}
}
let mut seen: Vec<PathBuf> = Vec::new();
let mut ranked: Vec<((u32, u32), PathBuf)> = Vec::new();
for c in cands {
if !c.is_file() {
continue;
}
let abs = c.canonicalize().unwrap_or(c);
if seen.contains(&abs) {
continue;
}
seen.push(abs.clone());
if let Some(v) = nvcc_version(&abs) {
ranked.push((v, abs));
}
}
ranked.sort_by(|a, b| b.0.cmp(&a.0));
if let Some(((maj, min), p)) = ranked.first() {
let others: Vec<String> = ranked[1..]
.iter()
.map(|((a, b), q)| format!("{a}.{b} @ {}", q.display()))
.collect();
println!(
"cargo:warning=nvcc auto-detected CUDA {maj}.{min} at {}{}",
p.display(),
if others.is_empty() {
String::new()
} else {
format!(
" (newest of {}; also saw {})",
ranked.len(),
others.join(", ")
)
}
);
return p.to_string_lossy().into_owned();
}
let pin = "/usr/local/cuda-13.1/bin/nvcc";
println!(
"cargo:warning=no runnable nvcc found via MEMRA_NVCC / CUDA_HOME / PATH / \
/usr/local/cuda*; falling back to {pin} — set MEMRA_NVCC=<path/to/nvcc>"
);
pin.to_string()
}
fn detect_arch() -> String {
let cap = Command::new("nvidia-smi")
.args(["--query-gpu=compute_cap", "--format=csv,noheader"])
.output()
.ok()
.filter(|o| o.status.success())
.and_then(|o| String::from_utf8(o.stdout).ok())
.and_then(|s| s.lines().next().map(|l| l.trim().to_string()));
let arch = match cap.as_deref() {
Some("12.0") | Some("12.1") => "120a",
Some("10.0") => "100a",
Some("9.0") => "90a",
Some("8.9") => "89",
_ => "120a",
};
match &cap {
Some(c) => println!(
"cargo:warning=MEMRA_CUDA_ARCH auto-detected {arch} (compute_cap {c}); set MEMRA_CUDA_ARCH to override"
),
None => {
println!("cargo:warning=no GPU visible; defaulting MEMRA_CUDA_ARCH=120a (compile-only)")
}
}
arch.to_string()
}
fn main() {
let out = PathBuf::from(std::env::var("OUT_DIR").unwrap());
if std::env::var_os("DOCS_RS").is_some() {
for stem in [
"kernels",
"hybrid",
"qmatvec",
"flash_attn",
"qmatvec_gemm",
"moe_router",
"spec_sample",
"flash_attn_vq4",
"flash_attn_vf8",
"flash_attn_kf8",
"flash_attn_kf8vq4",
"flash_attn_kf8vf8",
] {
std::fs::write(out.join(format!("{stem}.fatbin")), []).unwrap();
}
for (env, stem) in [
("MEMRA_ENGINE_FATBIN", "kernels"),
("MEMRA_HYBRID_FATBIN", "hybrid"),
("MEMRA_QMATVEC_FATBIN", "qmatvec"),
("MEMRA_FLASH_FATBIN", "flash_attn"),
("MEMRA_GEMM_FATBIN", "qmatvec_gemm"),
("MEMRA_ROUTER_FATBIN", "moe_router"),
("MEMRA_SAMPLE_FATBIN", "spec_sample"),
("MEMRA_FLASH_FATBIN_VQ4", "flash_attn_vq4"),
("MEMRA_FLASH_FATBIN_VF8", "flash_attn_vf8"),
("MEMRA_FLASH_FATBIN_KF8", "flash_attn_kf8"),
("MEMRA_FLASH_FATBIN_KF8VQ4", "flash_attn_kf8vq4"),
("MEMRA_FLASH_FATBIN_KF8VF8", "flash_attn_kf8vf8"),
] {
println!(
"cargo:rustc-env={env}={}",
out.join(format!("{stem}.fatbin")).display()
);
}
println!("cargo:rustc-check-cfg=cfg(memra_portable_cuda)");
println!("cargo:rustc-check-cfg=cfg(memra_hopper_mma)");
println!("cargo:rustc-check-cfg=cfg(memra_cutlass)");
println!("cargo:rustc-env=MEMRA_BUILT_CUDA_ARCH=120a");
return;
}
let nvcc = resolve_nvcc();
println!("cargo:rerun-if-env-changed=MEMRA_CUDA_ARCH");
println!("cargo:rerun-if-env-changed=MEMRA_CUTLASS");
println!("cargo:rustc-check-cfg=cfg(memra_portable_cuda)");
println!("cargo:rustc-check-cfg=cfg(memra_hopper_mma)");
println!("cargo:rustc-check-cfg=cfg(memra_cutlass)");
let cuda_arch = std::env::var("MEMRA_CUDA_ARCH").unwrap_or_else(|_| detect_arch());
assert!(
matches!(cuda_arch.as_str(), "120a" | "100a" | "90a" | "89"),
"MEMRA_CUDA_ARCH must be 120a (default), 100a (B200), 90a (Hopper), or 89 (portable eval)"
);
let portable = matches!(cuda_arch.as_str(), "89" | "90a");
let hopper_mma = cuda_arch == "90a";
assert!(
!(cuda_arch != "120a" && std::env::var_os("MEMRA_CUTLASS").is_some()),
"MEMRA_CUTLASS is sm_120a-only and cannot be enabled for this CUDA architecture"
);
let gencode = format!("arch=compute_{cuda_arch},code=sm_{cuda_arch}");
if portable {
println!("cargo:rustc-cfg=memra_portable_cuda");
}
if hopper_mma {
println!("cargo:rustc-cfg=memra_hopper_mma");
}
println!("cargo:rustc-env=MEMRA_BUILT_CUDA_ARCH={cuda_arch}");
for (src, env) in [
("cu/kernels.cu", "MEMRA_ENGINE_FATBIN"),
("cu/hybrid.cu", "MEMRA_HYBRID_FATBIN"),
("cu/qmatvec.cu", "MEMRA_QMATVEC_FATBIN"),
("cu/flash_attn.cu", "MEMRA_FLASH_FATBIN"),
("cu/qmatvec_gemm.cu", "MEMRA_GEMM_FATBIN"),
("cu/moe_router.cu", "MEMRA_ROUTER_FATBIN"),
("cu/spec_sample.cu", "MEMRA_SAMPLE_FATBIN"),
] {
println!("cargo:rerun-if-changed={src}");
println!("cargo:rerun-if-changed=cu/wgmma_common.cuh");
let stem = src.split('/').last().unwrap().trim_end_matches(".cu");
let fatbin = out.join(format!("{stem}.fatbin"));
let mut args = vec!["-gencode", &gencode, "-O3", "--fatbin"];
if portable {
args.push("-DMEMRA_PORTABLE_CUDA=1");
}
if hopper_mma {
args.push("-DMEMRA_HOPPER_MMA=1");
}
if cuda_arch == "100a" && src == "cu/qmatvec_gemm.cu" {
args.push("-DMEMRA_DISABLE_NATIVE_FP4=1");
}
println!("cargo:rerun-if-env-changed=MEMRA_FA_PP_MINBLOCKS");
let fa_mb = std::env::var("MEMRA_FA_PP_MINBLOCKS").ok();
let fa_mb_arg;
if let (Some(mb), "cu/flash_attn.cu") = (&fa_mb, src) {
fa_mb_arg = format!("-DFA_PP_MINBLOCKS={mb}");
args.push(&fa_mb_arg);
}
args.extend(["-o", fatbin.to_str().unwrap(), src]);
let status = Command::new(&nvcc).args(args).status().expect("spawn nvcc");
assert!(status.success(), "nvcc fatbin build failed for {src}");
println!("cargo:rustc-env={env}={}", fatbin.display());
}
for (suffix, kfmt, vfmt) in [
("VQ4", 0, 1),
("VF8", 0, 2),
("KF8", 1, 0),
("KF8VQ4", 1, 1),
("KF8VF8", 1, 2),
] {
let fatbin = out.join(format!("flash_attn_{}.fatbin", suffix.to_lowercase()));
let mut args = vec![
"-gencode".to_string(),
gencode.clone(),
"-O3".to_string(),
"--fatbin".to_string(),
];
if portable {
args.push("-DMEMRA_PORTABLE_CUDA=1".to_string());
}
if hopper_mma {
args.push("-DMEMRA_HOPPER_MMA=1".to_string());
}
args.extend([
format!("-DMEMRA_KV_KFMT={kfmt}"),
format!("-DMEMRA_KV_VFMT={vfmt}"),
"-o".to_string(),
fatbin.to_string_lossy().into_owned(),
"cu/flash_attn.cu".to_string(),
]);
let status = Command::new(&nvcc)
.args(args)
.status()
.expect("spawn nvcc (flash_attn kv-format variant)");
assert!(
status.success(),
"nvcc fatbin build failed for flash_attn kv variant {suffix}"
);
println!(
"cargo:rustc-env=MEMRA_FLASH_FATBIN_{suffix}={}",
fatbin.display()
);
}
{
let mut objs: Vec<PathBuf> = Vec::new();
println!("cargo:rerun-if-env-changed=MEMRA_MMQ_X_Q45K");
let q45k_x = std::env::var("MEMRA_MMQ_X_Q45K").ok();
println!("cargo:rerun-if-env-changed=MEMRA_MMQ_X_Q4");
let q4_x = std::env::var("MEMRA_MMQ_X_Q4").ok();
println!("cargo:rerun-if-env-changed=MEMRA_MMQ_X_W4A8");
let w4a8_x = std::env::var("MEMRA_MMQ_X_W4A8").ok();
println!("cargo:rerun-if-env-changed=MEMRA_MMQ_X_IQEXP");
let iqexp_x = std::env::var("MEMRA_MMQ_X_IQEXP").ok();
println!("cargo:rerun-if-env-changed=MEMRA_IQEXP_K16");
let iqexp_k16 = std::env::var("MEMRA_IQEXP_K16").ok();
println!("cargo:rerun-if-env-changed=MEMRA_MMQ_Y_W4A8");
let w4a8_y = std::env::var("MEMRA_MMQ_Y_W4A8").ok();
println!("cargo:rerun-if-env-changed=MEMRA_MMQ_FOLD_CEILING");
let w4a8_fold_ceiling = std::env::var("MEMRA_MMQ_FOLD_CEILING").ok();
println!("cargo:rerun-if-env-changed=MEMRA_MMQ_F8F4_PLAIN");
let f8f4_plain = std::env::var("MEMRA_MMQ_F8F4_PLAIN").ok();
println!("cargo:rerun-if-env-changed=MEMRA_MMQ_X_FP8");
let fp8_x = std::env::var("MEMRA_MMQ_X_FP8").ok();
println!("cargo:rerun-if-env-changed=MEMRA_MMQ_FP8BLK_PLAIN");
let fp8blk_plain = std::env::var("MEMRA_MMQ_FP8BLK_PLAIN").ok();
println!("cargo:rerun-if-env-changed=ACCPROBE_F32_PLAIN");
let accprobe_plain = std::env::var("ACCPROBE_F32_PLAIN").ok();
for mmq_src in [
"cu/mmq_fp4.cu",
"cu/mmq_q45k.cu",
"cu/mmq_nvfp4_w4a8.cu",
"cu/mmq_iq_experts.cu",
"cu/mmq_q8_0.cu",
"cu/mmq_q4_0.cu",
"cu/fp8_prefill.cu",
"cu/f16_prefill.cu",
"cu/mmq_nvfp4_f8f4.cu",
"cu/fa3_prefill.cu",
"cu/moe_f16_grouped.cu",
"cu/fp8_blk_dequant.cu",
"cu/mmq_fp8_blk.cu",
"cu/mmq_q8_0_f32acc.cu",
] {
println!("cargo:rerun-if-changed={mmq_src}");
let compile_src = if cuda_arch != "120a" && mmq_src == "cu/mmq_fp4.cu" {
"cu/mmq_fp4_stub.cu"
} else if portable && mmq_src == "cu/mmq_nvfp4_w4a8.cu" {
"cu/mmq_nvfp4_w4a8_stub.cu"
} else if portable && mmq_src == "cu/mmq_fp8_blk.cu" {
"cu/mmq_fp8_blk_stub.cu"
} else {
mmq_src
};
println!("cargo:rerun-if-changed={compile_src}");
let stem = mmq_src.split('/').last().unwrap().trim_end_matches(".cu");
let obj = out.join(format!("{stem}.o"));
let mut args: Vec<String> = vec![
"-gencode".into(),
gencode.clone(),
"-O3".into(),
"-std=c++17".into(),
"--expt-relaxed-constexpr".into(),
];
if mmq_src.ends_with("mmq_q45k.cu") {
if let Some(x) = &q45k_x {
args.push(format!("-DMMQ_X={x}"));
}
}
if mmq_src.ends_with("mmq_q4_0.cu") {
if let Some(x) = &q4_x {
args.push(format!("-DMMQ_X={x}"));
}
}
if mmq_src.ends_with("mmq_nvfp4_w4a8.cu") {
if let Some(x) = &w4a8_x {
args.push(format!("-DMMQ_X={x}"));
}
if let Some(y) = &w4a8_y {
args.push(format!("-DMMQ_Y={y}"));
}
if let Some(v) = &w4a8_fold_ceiling {
args.push(format!("-DMEMRA_MMQ_FOLD_CEILING={v}"));
}
if f8f4_plain.as_deref() == Some("1") {
args.push("-DMEMRA_F8F4_PLAIN_MMA".into());
}
}
if mmq_src.ends_with("mmq_iq_experts.cu") {
if let Some(x) = &iqexp_x {
args.push(format!("-DMMQ_X={x}"));
}
if iqexp_k16.as_deref() == Some("1") {
args.push("-DMEMRA_IQEXP_K16_MMA".into());
}
}
if mmq_src.ends_with("mmq_fp8_blk.cu") {
if let Some(x) = &fp8_x {
args.push(format!("-DFP8_MMQ_X={x}"));
}
if fp8blk_plain.as_deref() == Some("1") {
args.push("-DMEMRA_FP8BLK_PLAIN_MMA".into());
}
}
if mmq_src.ends_with("mmq_q8_0_f32acc.cu") {
if accprobe_plain.as_deref() == Some("1") {
args.push("-DMEMRA_ACCPROBE_PLAIN_MMA".into());
}
}
if mmq_src.ends_with("fa3_prefill.cu") && cuda_arch != "90a" {
args.push("-DMEMRA_FA3_STUB".into());
}
args.extend([
"-c".into(),
compile_src.into(),
"-o".into(),
obj.to_str().unwrap().into(),
]);
let status = Command::new(&nvcc)
.args(&args)
.status()
.expect("spawn nvcc (mmq)");
assert!(
status.success(),
"nvcc static-lib build failed for {mmq_src}"
);
objs.push(obj);
}
let lib = out.join("libmemra_mmq.a");
let _ = std::fs::remove_file(&lib);
let mut ar_args = vec!["crus".to_string(), lib.to_str().unwrap().to_string()];
ar_args.extend(objs.iter().map(|o| o.to_str().unwrap().to_string()));
let status = Command::new("ar")
.args(&ar_args)
.status()
.expect("spawn ar (mmq)");
assert!(status.success(), "ar failed for {}", lib.display());
println!("cargo:rustc-link-search=native={}", out.display());
println!("cargo:rustc-link-lib=static:+whole-archive=memra_mmq");
let cuda_lib = std::path::Path::new(&nvcc)
.parent()
.and_then(|p| p.parent())
.map(|p| p.join("lib64"))
.unwrap_or_else(|| std::path::PathBuf::from("/usr/local/cuda-13.1/lib64"));
println!("cargo:rustc-link-search=native={}", cuda_lib.display());
println!("cargo:rustc-link-lib=dylib=cudart");
println!("cargo:rustc-link-lib=dylib=stdc++");
println!("cargo:rustc-link-lib=dylib=cublasLt");
println!("cargo:rustc-link-lib=dylib=cublas");
println!(
"cargo:rustc-link-search=native={}/stubs",
cuda_lib.display()
);
println!("cargo:rustc-link-lib=dylib=cuda");
}
if std::env::var("MEMRA_CUTLASS").is_ok() {
let cutlass_src = "cu/cutlass_fp4_sm120.cu";
println!("cargo:rerun-if-changed={cutlass_src}");
let cutlass_root = std::env::var("MEMRA_CUTLASS_ROOT").unwrap_or_else(|_| {
"/home/avifenesh/.venvs/torch/lib/python3.12/site-packages/flashinfer/data/cutlass"
.into()
});
let cutlass_inc = format!("{cutlass_root}/include");
let cutlass_util = format!("{cutlass_root}/tools/util/include");
let obj = out.join("cutlass_fp4_sm120.o");
let lib = out.join("libmemra_cutlass.a");
let status = Command::new(&nvcc)
.args([
"-gencode",
"arch=compute_120a,code=sm_120a",
"-O3",
"-std=c++17",
"--expt-relaxed-constexpr",
"-DENABLE_BF16",
"-DENABLE_FP4",
"-DCUTLASS_ENABLE_GDC_FOR_SM100=1",
"-I",
&cutlass_inc,
"-I",
&cutlass_util,
"-c",
cutlass_src,
"-o",
obj.to_str().unwrap(),
])
.status()
.expect("spawn nvcc (cutlass)");
assert!(
status.success(),
"nvcc static-lib build failed for {cutlass_src}"
);
let _ = std::fs::remove_file(&lib);
let status = Command::new("ar")
.args(["crus", lib.to_str().unwrap(), obj.to_str().unwrap()])
.status()
.expect("spawn ar");
assert!(status.success(), "ar failed for {}", lib.display());
println!("cargo:rustc-link-search=native={}", out.display());
println!("cargo:rustc-link-arg=-Wl,--whole-archive");
println!("cargo:rustc-link-arg={}", lib.display());
println!("cargo:rustc-link-arg=-Wl,--no-whole-archive");
let cuda_lib = std::path::Path::new(&nvcc)
.parent()
.and_then(|p| p.parent())
.map(|p| p.join("lib64"))
.unwrap_or_else(|| std::path::PathBuf::from("/usr/local/cuda-13.1/lib64"));
println!("cargo:rustc-link-search=native={}", cuda_lib.display());
println!("cargo:rustc-link-lib=dylib=cudart");
println!("cargo:rustc-link-lib=dylib=stdc++");
println!("cargo:rustc-cfg=memra_cutlass");
}
}