diff --git a/crates/ml-ppo/src/cuda_nn/linear.rs b/crates/ml-ppo/src/cuda_nn/linear.rs index f404bc348..926bd53f6 100644 --- a/crates/ml-ppo/src/cuda_nn/linear.rs +++ b/crates/ml-ppo/src/cuda_nn/linear.rs @@ -73,7 +73,7 @@ extern "C" __global__ void bias_add( fn compile_bias_add(ctx: &GpuContext) -> Result { let context = ctx.stream.context(); let ptx_result = LINEAR_BIAS_ADD_PTX.get_or_init(|| { - cudarc::nvrtc::compile_ptx(BIAS_ADD_KERNEL) + ml_core::cuda_compile::compile_ptx_for_device(BIAS_ADD_KERNEL, &context) .map_err(|e| format!("Failed to compile bias_add kernel: {e}")) }); let ptx = ptx_result.as_ref().map_err(|e| { diff --git a/crates/ml-ppo/src/cuda_nn/lstm.rs b/crates/ml-ppo/src/cuda_nn/lstm.rs index 33c9484d1..fb4910e44 100644 --- a/crates/ml-ppo/src/cuda_nn/lstm.rs +++ b/crates/ml-ppo/src/cuda_nn/lstm.rs @@ -156,7 +156,7 @@ impl CudaLSTM { let context = stream.context(); let gate_ptx_result = LSTM_GATE_PTX.get_or_init(|| { - cudarc::nvrtc::compile_ptx(LSTM_GATE_KERNEL) + ml_core::cuda_compile::compile_ptx_for_device(LSTM_GATE_KERNEL, &context) .map_err(|e| format!("compile gate kernel: {e}")) }); let gate_ptx = gate_ptx_result.as_ref().map_err(|e| { @@ -170,7 +170,7 @@ impl CudaLSTM { })?; let bias_ptx_result = LSTM_BIAS_ADD_PTX.get_or_init(|| { - cudarc::nvrtc::compile_ptx(BIAS_ADD_KERNEL) + ml_core::cuda_compile::compile_ptx_for_device(BIAS_ADD_KERNEL, &context) .map_err(|e| format!("compile bias_add: {e}")) }); let bias_ptx = bias_ptx_result.as_ref().map_err(|e| { diff --git a/crates/ml-ppo/src/cuda_nn/softmax.rs b/crates/ml-ppo/src/cuda_nn/softmax.rs index 8f18dc10a..907122028 100644 --- a/crates/ml-ppo/src/cuda_nn/softmax.rs +++ b/crates/ml-ppo/src/cuda_nn/softmax.rs @@ -119,7 +119,7 @@ struct SoftmaxKernels { fn compile_softmax_kernels(stream: &Arc) -> Result { let context = stream.context(); let ptx_result = SOFTMAX_PTX.get_or_init(|| { - cudarc::nvrtc::compile_ptx(SOFTMAX_KERNEL) + ml_core::cuda_compile::compile_ptx_for_device(SOFTMAX_KERNEL, &context) .map_err(|e| format!("Failed to compile softmax kernels: {e}")) }); let ptx = ptx_result.as_ref().map_err(|e| {