Files
foxhunt/vendor/cudarc/examples/03-launch-kernel.rs
jgrusewski 07d0e60fe4 feat(bf16): remove nvrtc from entire workspace + wire ml-core precompiled cubins
- Fork cudarc locally (vendor/cudarc): add CudaContext::load_cubin()
  that calls cuModuleLoadData directly — zero nvrtc dependency
- Remove "nvrtc" feature from ml-core, ml-dqn, ml-ppo Cargo.toml
- Replace all 89 Ptx::from_binary + load_module calls with load_cubin
- ml-core cuda_autograd: wire 9 stub constructors to precompiled cubins
  (activation, elementwise, linear, loss, reduction, dropout, layer_norm, optimizer)
- ml-core build.rs: compile 8 BF16-native CUDA kernels via nvcc
- cubin_loader.rs: thin wrapper around CudaContext::load_cubin()
- Fix size_of::<f32> in gpu_tensor.rs, stream_ops.rs, layer_norm.rs
- Fix test data: Vec<f32> → Vec<half::bf16> for memcpy_htod
- Stub ml-ppo/ml-dqn runtime compile_ptx calls (dead code)
- backtest_metrics_kernel.cu: full native BF16 rewrite (no float)
- backtest_env_kernel.cu: shared memory → __nv_bfloat16

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-03-28 10:11:46 +01:00

39 lines
1.1 KiB
Rust

use cudarc::{
driver::{CudaContext, DriverError, LaunchConfig, PushKernelArg},
nvrtc::Ptx,
};
fn main() -> Result<(), DriverError> {
let ctx = CudaContext::new(0)?;
let stream = ctx.default_stream();
// You can load a function from a pre-compiled PTX like so:
let module = ctx.load_module(Ptx::from_file("./examples/sin.ptx"))?;
// and then load a function from it:
let f = module.load_function("sin_kernel").unwrap();
let a_host = [1.0, 2.0, 3.0];
let a_dev = stream.clone_htod(&a_host)?;
let mut b_dev = a_dev.clone();
// we use a buidler pattern to launch kernels.
let n = 3i32;
let cfg = LaunchConfig::for_num_elems(n as u32);
let mut launch_args = stream.launch_builder(&f);
launch_args.arg(&mut b_dev);
launch_args.arg(&a_dev);
launch_args.arg(&n);
unsafe { launch_args.launch(cfg) }?;
let a_host_2 = stream.clone_dtoh(&a_dev)?;
let b_host = stream.clone_dtoh(&b_dev)?;
println!("Found {b_host:?}");
println!("Expected {:?}", a_host.map(f32::sin));
assert_eq!(&a_host, a_host_2.as_slice());
Ok(())
}