feat: build.rs precompiles all 25 CUDA kernels at cargo build time
Extend build.rs to compile all 25 .cu kernel files to cubins via nvcc. Common header prepended to 24 kernels; experience_kernels.cu is standalone. Total cubin size ~1.1 MB embedded in binary. - cargo:rerun-if-changed tracks all .cu and .cuh files - CUDA_COMPUTE_CAP env var selects target arch (sm_86, sm_90) - Graceful skip when nvcc unavailable (non-CUDA builds) Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -28,65 +28,127 @@ fn main() {
|
||||
}
|
||||
};
|
||||
|
||||
// Proof-of-concept: compile epsilon_greedy_kernel.cu (no #define dependencies)
|
||||
let simple_kernels = [
|
||||
"epsilon_greedy_kernel.cu",
|
||||
];
|
||||
|
||||
// Read common header once
|
||||
let common_src = std::fs::read_to_string(&common_header)
|
||||
.unwrap_or_else(|e| panic!("Failed to read {}: {e}", common_header.display()));
|
||||
|
||||
println!("cargo:rerun-if-changed={}", common_header.display());
|
||||
|
||||
for kernel_name in &simple_kernels {
|
||||
let kernel_path = kernel_dir.join(kernel_name);
|
||||
let cubin_name = kernel_name.replace(".cu", ".cubin");
|
||||
let cubin_path = out_dir.join(&cubin_name);
|
||||
// All 25 kernels to precompile.
|
||||
// Kernels marked "standalone" have their own helpers and don't need common header.
|
||||
// All others get common_device_functions.cuh prepended.
|
||||
let kernels_with_common = [
|
||||
"epsilon_greedy_kernel.cu",
|
||||
"backtest_env_kernel.cu",
|
||||
"backtest_forward_ppo_kernel.cu",
|
||||
"backtest_forward_supervised_kernel.cu",
|
||||
"backtest_gather_kernel.cu",
|
||||
"backtest_metrics_kernel.cu",
|
||||
"dt_kernels.cu",
|
||||
"ensemble_kernels.cu",
|
||||
"her_episode_kernel.cu",
|
||||
"her_relabel_kernel.cu",
|
||||
"signal_adapter_kernel.cu",
|
||||
"statistics_kernel.cu",
|
||||
"training_guard_kernel.cu",
|
||||
"c51_loss_kernel.cu",
|
||||
"mse_loss_kernel.cu",
|
||||
"curiosity_training_kernel.cu",
|
||||
"dqn_utility_kernels.cu",
|
||||
"attention_kernel.cu",
|
||||
"attention_backward_kernel.cu",
|
||||
"iql_value_kernel.cu",
|
||||
"iqn_dual_head_kernel.cu",
|
||||
"monitoring_kernel.cu",
|
||||
"nstep_kernel.cu",
|
||||
"ppo_experience_kernel.cu",
|
||||
];
|
||||
|
||||
println!("cargo:rerun-if-changed={}", kernel_path.display());
|
||||
// experience_kernels.cu is standalone -- it has its own inline helpers
|
||||
// and explicitly does NOT use common_device_functions.cuh.
|
||||
let standalone_kernels = [
|
||||
"experience_kernels.cu",
|
||||
];
|
||||
|
||||
let kernel_src = std::fs::read_to_string(&kernel_path)
|
||||
.unwrap_or_else(|e| panic!("Failed to read {}: {e}", kernel_path.display()));
|
||||
// Compile kernels that need common header prepended
|
||||
for kernel_name in &kernels_with_common {
|
||||
compile_kernel(
|
||||
&nvcc,
|
||||
kernel_dir,
|
||||
kernel_name,
|
||||
&arch,
|
||||
&out_dir,
|
||||
Some(&common_src),
|
||||
);
|
||||
}
|
||||
|
||||
// Compose source: common header + kernel
|
||||
// The common header requires STATE_DIM, MARKET_DIM, PORTFOLIO_DIM to be defined.
|
||||
// For kernels that DON'T use these symbols, we inject dummy defines so the
|
||||
// header compiles without errors. These kernels don't reference the values.
|
||||
let defines = "\
|
||||
#define STATE_DIM 72\n\
|
||||
#define MARKET_DIM 42\n\
|
||||
#define PORTFOLIO_DIM 8\n";
|
||||
let full_source = format!("{defines}{common_src}\n{kernel_src}");
|
||||
// Compile standalone kernels (no common header)
|
||||
for kernel_name in &standalone_kernels {
|
||||
compile_kernel(
|
||||
&nvcc,
|
||||
kernel_dir,
|
||||
kernel_name,
|
||||
&arch,
|
||||
&out_dir,
|
||||
None,
|
||||
);
|
||||
}
|
||||
|
||||
// Write composed source to temp file
|
||||
let tmp_src = out_dir.join(format!("_{kernel_name}"));
|
||||
std::fs::write(&tmp_src, &full_source).unwrap();
|
||||
eprintln!(" Precompiled {} CUDA kernels ({arch})", kernels_with_common.len() + standalone_kernels.len());
|
||||
}
|
||||
|
||||
// Compile with nvcc
|
||||
let status = Command::new(&nvcc)
|
||||
.args([
|
||||
"-cubin",
|
||||
&format!("-arch={arch}"),
|
||||
"-O3",
|
||||
"--use_fast_math",
|
||||
"--ftz=true",
|
||||
"--fmad=true",
|
||||
"-o", cubin_path.to_str().unwrap(),
|
||||
tmp_src.to_str().unwrap(),
|
||||
])
|
||||
.status();
|
||||
/// Compile a single .cu kernel file to a .cubin via nvcc.
|
||||
///
|
||||
/// If `common_header` is Some, it is prepended to the kernel source.
|
||||
fn compile_kernel(
|
||||
nvcc: &Path,
|
||||
kernel_dir: &Path,
|
||||
kernel_name: &str,
|
||||
arch: &str,
|
||||
out_dir: &Path,
|
||||
common_header: Option<&str>,
|
||||
) {
|
||||
let kernel_path = kernel_dir.join(kernel_name);
|
||||
let cubin_name = kernel_name.replace(".cu", ".cubin");
|
||||
let cubin_path = out_dir.join(&cubin_name);
|
||||
|
||||
match status {
|
||||
Ok(s) if s.success() => {
|
||||
eprintln!(" Compiled {kernel_name} -> {cubin_name} ({arch})");
|
||||
}
|
||||
Ok(s) => {
|
||||
panic!("nvcc failed to compile {kernel_name} (exit={})", s.code().unwrap_or(-1));
|
||||
}
|
||||
Err(e) => {
|
||||
panic!("Failed to execute nvcc: {e}");
|
||||
}
|
||||
println!("cargo:rerun-if-changed={}", kernel_path.display());
|
||||
|
||||
let kernel_src = std::fs::read_to_string(&kernel_path)
|
||||
.unwrap_or_else(|e| panic!("Failed to read {}: {e}", kernel_path.display()));
|
||||
|
||||
// Compose source: optional common header + kernel
|
||||
let full_source = match common_header {
|
||||
Some(header) => format!("{header}\n{kernel_src}"),
|
||||
None => kernel_src,
|
||||
};
|
||||
|
||||
// Write composed source to temp file
|
||||
let tmp_src = out_dir.join(format!("_{kernel_name}"));
|
||||
std::fs::write(&tmp_src, &full_source).unwrap();
|
||||
|
||||
// Compile with nvcc
|
||||
let status = Command::new(nvcc)
|
||||
.args([
|
||||
"-cubin",
|
||||
&format!("-arch={arch}"),
|
||||
"-O3",
|
||||
"--use_fast_math",
|
||||
"--ftz=true",
|
||||
"--fmad=true",
|
||||
"-o", cubin_path.to_str().unwrap(),
|
||||
tmp_src.to_str().unwrap(),
|
||||
])
|
||||
.status();
|
||||
|
||||
match status {
|
||||
Ok(s) if s.success() => {
|
||||
eprintln!(" Compiled {kernel_name} -> {cubin_name} ({arch})");
|
||||
}
|
||||
Ok(s) => {
|
||||
panic!("nvcc failed to compile {kernel_name} (exit={})", s.code().unwrap_or(-1));
|
||||
}
|
||||
Err(e) => {
|
||||
panic!("Failed to execute nvcc: {e}");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user