diff --git a/crates/ml/build.rs b/crates/ml/build.rs index 24a9666fc..fcc0e680c 100644 --- a/crates/ml/build.rs +++ b/crates/ml/build.rs @@ -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}"); } } }