Files
foxhunt/crates/ml-alpha/build.rs

116 lines
3.9 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! Pre-compile all ml-alpha CUDA kernels into arch-specific cubins.
//!
//! Per `feedback_no_nvrtc.md`: no runtime kernel compilation.
//! Per `pearl_build_rs_rerun_if_env_changed.md`: every `std::env::var`
//! is paired with `cargo:rerun-if-env-changed`.
use std::path::{Path, PathBuf};
use std::process::Command;
const KERNELS: &[&str] = &[
"mamba2_alpha_kernel", // Mamba2 SSM scan kernel (used by PerceptionTrainer's encoder prefix)
"snap_feature_assemble",
"cfc_step",
"multi_horizon_heads",
"projection",
"bce_loss_multi_horizon", // Kendall σ-weighted multi-horizon BCE (axis A)
"adamw_step",
"grad_norm",
"horizon_lambda", // ISV-driven per-horizon gradient scaler (EMA + lambda)
"layer_norm", // Phase 1: trunk pre-CfC normalisation
"variable_selection", // Phase 2D: TFT-style per-feature gating
"attention_pool", // Phase 3: single-Q learned content summary at CfC k=0
"reduce_axis0", // Phase B: cross-batch param-grad reducer
"output_smoothness", // CRT.train: per-horizon adjacent-position prob-jitter penalty
"smoothness_lambda_controller", // CRT.train: ISV-driven λ controller anchored on h30 jitter
"gpu_log_ring", // GPU diagnostic log ring — tick kernel + log_record helper
];
// Cache bust v14 (2026-05-21): gpu_log_ring.cu added — unified in-kernel diagnostic logging.
fn main() {
println!("cargo:rerun-if-changed=build.rs");
println!("cargo:rerun-if-env-changed=CARGO_FEATURE_CUDA");
if std::env::var("CARGO_FEATURE_CUDA").is_err() {
eprintln!(" ml-alpha: cuda feature disabled, skipping kernel build");
return;
}
println!("cargo:rerun-if-env-changed=CUDA_COMPUTE_CAP");
println!("cargo:rerun-if-env-changed=CUDA_HOME");
let cap = std::env::var("CUDA_COMPUTE_CAP").unwrap_or_else(|_| "80".to_string());
let arch = format!("sm_{cap}");
let nvcc = match find_nvcc() {
Some(p) => p,
None => {
eprintln!(" ml-alpha: nvcc not found, skipping kernel build (set CUDA_HOME or install CUDA toolkit)");
return;
}
};
let out = PathBuf::from(std::env::var("OUT_DIR").expect("OUT_DIR not set by cargo"));
for k in KERNELS {
let src = PathBuf::from(format!("cuda/{k}.cu"));
if !src.exists() {
eprintln!(" ml-alpha: skipping {k} — source not yet present");
continue;
}
println!("cargo:rerun-if-changed={}", src.display());
let cubin = out.join(format!("{k}.cubin"));
compile(&nvcc, &src, &cubin, &arch);
}
}
fn compile(nvcc: &Path, src: &Path, cubin: &Path, arch: &str) {
let status = Command::new(nvcc)
.args([
"-cubin",
&format!("-arch={arch}"),
"-O3",
"--use_fast_math",
"--ftz=true",
"--fmad=true",
"-o",
cubin.to_str().unwrap(),
src.to_str().unwrap(),
])
.status()
.unwrap_or_else(|e| panic!("nvcc spawn failed for {}: {e}", src.display()));
if !status.success() {
panic!(
"nvcc failed for {} (exit {})",
src.display(),
status.code().unwrap_or(-1)
);
}
eprintln!(
" ml-alpha: compiled {} -> {} ({arch})",
src.display(),
cubin.display()
);
}
fn find_nvcc() -> Option<PathBuf> {
if let Ok(home) = std::env::var("CUDA_HOME") {
let p = PathBuf::from(home).join("bin/nvcc");
if p.exists() {
return Some(p);
}
}
for cand in ["/usr/local/cuda/bin/nvcc", "/usr/bin/nvcc"] {
let p = PathBuf::from(cand);
if p.exists() {
return Some(p);
}
}
Command::new("nvcc")
.arg("--version")
.output()
.ok()
.filter(|o| o.status.success())
.map(|_| PathBuf::from("nvcc"))
}