Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,12 @@ rand_distr = "0.5"
clap = { version = "4", features = ["derive"] }
bincode = "1"

[features]
# Train on the GPU (cubecl-cuda, Nvidia). Off by default: CPU (ndarray) is
# faster for small models and needs no GPU/toolkit; enable for bigger models
# where the matmuls become compute-bound. Requires the CUDA deps in shell.nix.
cuda = ["burn/cuda"]

# Enable optimizations in dev for physics sim performance
[profile.dev]
opt-level = 1
Expand Down
16 changes: 13 additions & 3 deletions shell.nix
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@ let
pkgs = import (fetchTarball {
url = "https://github.com/NixOS/nixpkgs/archive/d6c71932130818840fc8fe9509cf50be8c64634f.tar.gz";
sha256 = "1klgyhj98j3gfsql5sn9rapyx62qk5g8adk5zh9mnc4d0fj61gdr";
}) {};
}) { config.allowUnfree = true; }; # allowUnfree: CUDA toolkit is unfree
in
pkgs.mkShell {
buildInputs = with pkgs; [
Expand Down Expand Up @@ -33,14 +33,24 @@ pkgs.mkShell {
# Build tools
clang
mold

# CUDA — burn's `cuda` backend (cubecl-cuda). cudatoolkit provides
# nvcc/nvrtc/cudart; libcuda (the driver) is NOT here, it ships with the
# host NVIDIA driver at /run/opengl-driver/lib (on LD_LIBRARY_PATH below).
cudaPackages.cudatoolkit
];

# Point Vulkan ICD loader at the right drivers
# cubecl resolves the CUDA toolkit through CUDA_PATH.
CUDA_PATH = pkgs.cudaPackages.cudatoolkit;

# Vulkan ICD + CUDA runtime libs. libcuda.so (driver) is host-only, so append
# the raw /run/opengl-driver/lib path after the nix-store libs.
LD_LIBRARY_PATH = pkgs.lib.makeLibraryPath (with pkgs; [
vulkan-loader
udev
alsa-lib
libxkbcommon
wayland
]);
cudaPackages.cudatoolkit
]) + ":/run/opengl-driver/lib";
}
37 changes: 26 additions & 11 deletions src/training/session.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,6 @@ use std::path::{Path, PathBuf};
use bevy::app::AppExit;
use bevy::prelude::*;
use burn::backend::Autodiff;
use burn::backend::ndarray::{NdArray, NdArrayDevice};
use burn::grad_clipping::GradientClippingConfig;
use burn::module::AutodiffModule;
use burn::optim::adaptor::OptimizerAdaptor;
Expand Down Expand Up @@ -170,8 +169,17 @@ impl ObsNormalizer {
}
}

/// Backend type aliases.
pub type TrainBackend = Autodiff<NdArray>;
/// Training backend. CPU (ndarray) by default — fast for small models and needs
/// no GPU; build with `--features cuda` to train on the GPU (cubecl-cuda), which
/// pays off as models grow. `InferBackend` and `Device` derive from the choice,
/// so the rest of the code is backend-agnostic.
#[cfg(not(feature = "cuda"))]
pub type TrainBackend = Autodiff<burn::backend::ndarray::NdArray>;
#[cfg(feature = "cuda")]
pub type TrainBackend = Autodiff<burn::backend::cuda::Cuda>;

pub type InferBackend = <TrainBackend as burn::tensor::backend::AutodiffBackend>::InnerBackend;
pub type Device = <TrainBackend as burn::tensor::backend::Backend>::Device;

/// Concrete optimizer type.
type CrabOptimizer = OptimizerAdaptor<Adam, CrabBrain<TrainBackend>, TrainBackend>;
Expand Down Expand Up @@ -247,7 +255,7 @@ pub struct TrainingState {
pub brain: CrabBrain<TrainBackend>,
pub config: PpoConfig,
pub rollout: RolloutBuffer,
pub device: NdArrayDevice,
pub device: Device,
optimizer: CrabOptimizer,

pub episode_reward: f32,
Expand Down Expand Up @@ -275,7 +283,7 @@ pub struct TrainingState {

impl TrainingState {
pub fn new(args: &Args) -> Self {
let device = NdArrayDevice::Cpu;
let device = Device::default();
let mut brain: CrabBrain<TrainBackend> = CrabBrain::new(&device);
let optimizer: CrabOptimizer = AdamConfig::new()
.with_grad_clipping(Some(GradientClippingConfig::Norm(0.5)))
Expand Down Expand Up @@ -568,17 +576,20 @@ pub fn brain_step(
carapace_q: Query<(&Transform, &bevy_rapier3d::prelude::Velocity), With<CrabCarapace>>,
) {
let raw_obs = obs.values;
let device = training.device;
// Clone needed under --features cuda (CudaDevice isn't Copy); a no-op copy
// for the default NdArrayDevice, hence the allow.
#[allow(clippy::clone_on_copy)]
let device = training.device.clone();

let obs_array = training.obs_normalizer.normalize(&raw_obs);

let inference_brain = training.brain.valid();

let obs_tensor = Tensor::<NdArray, 1>::from_floats(obs_array.as_slice(), &device);
let obs_tensor = Tensor::<InferBackend, 1>::from_floats(obs_array.as_slice(), &device);
let obs_batch = obs_tensor.clone().unsqueeze::<2>();

let (means_batch, log_std) = inference_brain.policy(obs_batch);
let means: Tensor<NdArray, 1> = means_batch.flatten(0, 1);
let means: Tensor<InferBackend, 1> = means_batch.flatten(0, 1);

let obs_batch2 = obs_tensor.clone().unsqueeze::<2>();
let value = inference_brain
Expand Down Expand Up @@ -779,6 +790,10 @@ pub fn save_on_exit(
#[cfg(test)]
mod tests {
use super::*;
use burn::backend::ndarray::{NdArray, NdArrayDevice};

// Backend-agnostic tests run on CPU so `cargo test` needs no GPU.
type TestBackend = Autodiff<NdArray>;

#[test]
fn reward_velocity_penalty_is_significant() {
Expand Down Expand Up @@ -846,7 +861,7 @@ mod tests {
std::fs::create_dir_all(&dir).unwrap();

let device = NdArrayDevice::Cpu;
let brain: CrabBrain<TrainBackend> = CrabBrain::new(&device);
let brain: CrabBrain<TestBackend> = CrabBrain::new(&device);

let stem = dir.join(BRAIN_STEM);
let recorder = BinFileRecorder::<FullPrecisionSettings>::default();
Expand All @@ -860,9 +875,9 @@ mod tests {
);

let loaded_record = recorder.load(stem, &device).expect("load brain");
let loaded = CrabBrain::<TrainBackend>::new(&device).load_record(loaded_record);
let loaded = CrabBrain::<TestBackend>::new(&device).load_record(loaded_record);

let test_obs = Tensor::<TrainBackend, 2>::zeros([1, OBS_SIZE], &device);
let test_obs = Tensor::<TestBackend, 2>::zeros([1, OBS_SIZE], &device);
let (orig_means, orig_log_std) = brain.policy(test_obs.clone());
let (loaded_means, loaded_log_std) = loaded.policy(test_obs);

Expand Down