diff --git a/.cargo/config.toml b/.cargo/config.toml new file mode 100644 index 0000000..80a81b5 --- /dev/null +++ b/.cargo/config.toml @@ -0,0 +1,8 @@ +# for Linux +[target.x86_64-unknown-linux-gnu] +linker = "clang" +rustflags = ["-C", "link-arg=-fuse-ld=lld"] + +# for Windows +[target.x86_64-pc-windows-msvc] +linker = "rust-lld.exe" diff --git a/Cargo.lock b/Cargo.lock index e61107f..53c7626 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2167,6 +2167,8 @@ dependencies = [ "ctrlc", "rand 0.10.1", "rand_distr 0.6.0", + "serde", + "serde_json", ] [[package]] diff --git a/engine/Cargo.toml b/engine/Cargo.toml index c858ba0..a2752c9 100644 --- a/engine/Cargo.toml +++ b/engine/Cargo.toml @@ -10,3 +10,5 @@ burn-ndarray = "0.21.0" rand_distr = "0.6.0" rand = "0.10.1" ctrlc = "3.5.2" +serde_json = "1.0.150" +serde = { version = "1.0.228", features = ["derive"] } diff --git a/engine/src/main.rs b/engine/src/main.rs index 6cd3487..999aa86 100644 --- a/engine/src/main.rs +++ b/engine/src/main.rs @@ -1,8 +1,7 @@ #![recursion_limit = "256"] -use burn::backend::Autodiff; +use burn::backend::{Autodiff, Cuda}; use burn::optim::AdamConfig; -use burn_ndarray::NdArray; use engine::mcts::MctsConfig; use engine::training::train::{train, TrainingConfig}; // fn main() { @@ -16,14 +15,14 @@ use engine::training::train::{train, TrainingConfig}; fn main() { // type MyBackend = Wgpu; - // type MyBackend = Cuda; - type MyBackend = NdArray; + type MyBackend = Cuda; + // type MyBackend = NdArray; type MyAutodiffBackend = Autodiff; // let device = burn::backend::wgpu::WgpuDevice::default(); - let device = burn::backend::ndarray::NdArrayDevice::default(); - // let device = burn::backend::cuda::CudaDevice::default(); + // let device = burn::backend::ndarray::NdArrayDevice::default(); + let device = burn::backend::cuda::CudaDevice::default(); - let mcts_config = MctsConfig::new(100, 1.0, 0.05, 0.25); + let mcts_config = MctsConfig::new(400, 1.0, 0.05, 0.25, 128); let adam_config = AdamConfig::new(); @@ -41,6 +40,7 @@ fn main() { mcts_config, optimizer: adam_config, lr: 2e-4, + seed: None, }; train::(training_config, device); diff --git a/engine/src/mcts.rs b/engine/src/mcts.rs index 7833ed6..288818e 100644 --- a/engine/src/mcts.rs +++ b/engine/src/mcts.rs @@ -80,6 +80,7 @@ pub struct MctsConfig { pub c_puct: f32, pub dirichlet_alpha: f32, pub dirichlet_epsilon: f32, + pub batch_max: usize, } impl MctsConfig { @@ -88,19 +89,21 @@ impl MctsConfig { c_puct: f32, dirichlet_alpha: f32, dirichlet_epsilon: f32, + batch_max: usize, ) -> MctsConfig { MctsConfig { num_simulations, c_puct, dirichlet_alpha, dirichlet_epsilon, + batch_max, } } } impl Default for MctsConfig { fn default() -> MctsConfig { - MctsConfig::new(400, 1.0, 0.05, 0.25) + MctsConfig::new(400, 1.0, 0.05, 0.25, 64) } } @@ -128,10 +131,8 @@ impl Mcts { self.add_dirichlet_noise(root, &mut nodes); // We'll batch leaf evaluations to reduce per-leaf model calls and device-host syncs. - let mut sims_done = 0usize; + let mut sims_done: usize = 0; let num_sims = self.config.num_simulations; - // Tunable batch size for NN evaluation. Small value is safe; larger values increase throughput on GPU. - let batch_max = 32usize; while sims_done < num_sims { // Collect a batch of leaf nodes (and their selection paths) @@ -139,7 +140,7 @@ impl Mcts { let mut leaf_paths: Vec> = Vec::new(); let mut leaf_states: Vec> = Vec::new(); - while leaf_nodes.len() < std::cmp::min(batch_max, num_sims - sims_done) { + while leaf_nodes.len() < std::cmp::min(self.config.batch_max, num_sims - sims_done) { let mut path = vec![root]; let mut current = root; @@ -153,8 +154,9 @@ impl Mcts { leaf_paths.push(path.clone()); // Prepare state tensor for this leaf - let state: Tensor = encode_board_state_perspective(&nodes[current].board_state, device) - .reshape([1, 18, 8, 8]); + let state: Tensor = + encode_board_state_perspective(&nodes[current].board_state, device) + .reshape([1, 18, 8, 8]); leaf_states.push(state); sims_done += 1; @@ -184,32 +186,51 @@ impl Mcts { let logits = &policy_data[start..end]; // Convert logits to probabilities with a numerically-stable softmax on host - let mut max_logit = std::f32::NEG_INFINITY; + let mut max_logit = f32::NEG_INFINITY; for &v in logits.iter() { if v > max_logit { max_logit = v; } } - let mut exps_sum = 0.0f32; - // We'll build a Vec of probabilities lazily when needed - let mut probs: Vec = Vec::new(); - probs.resize(num_moves, 0.0); - for (j, &v) in logits.iter().enumerate() { - let e = (v - max_logit).exp(); - probs[j] = e; - exps_sum += e; + if logits.is_empty() { + println!("logits empty") } - if exps_sum > 0.0 { - for p in probs.iter_mut() { + if !max_logit.is_finite() { + println!("max logits not finite") + } + + let mut exps_sum = 0.0f32; + let mut probs = vec![0.0f32; num_moves]; + + for (j, &v) in logits.iter().enumerate() { + let x = (v - max_logit).exp(); + + if !x.is_finite() { + probs[j] = 0.0; + } else { + probs[j] = x; + exps_sum += x; + } + } + + if exps_sum > 0.0 && exps_sum.is_finite() { + for p in &mut probs { *p /= exps_sum; } + } else { + // fallback: uniform or zero + let uniform = 1.0 / num_moves as f32; + for p in &mut probs { + *p = uniform; + } } // Expand: add legal moves as children with prior from probs - let legal_moves: Vec = MoveGen::new_legal(&nodes[node_idx].board_state.board).collect(); + let legal_moves: Vec = + MoveGen::new_legal(&nodes[node_idx].board_state.board).collect(); for mv in legal_moves { let stm = nodes[node_idx].board_state.board.side_to_move(); - let idx = encode_move(mv, stm); + let idx = encode_move(mv, stm).expect("Invalid move"); let prior = probs[idx]; let mut new_board = nodes[node_idx].board_state.clone(); @@ -233,7 +254,7 @@ impl Mcts { let denom = self.config.num_simulations as f32; for child_idx in nodes[root].children.iter() { let mv = nodes[*child_idx].last_move.expect("move didnt exist"); - let enc = encode_move(mv, stm); + let enc = encode_move(mv, stm).expect("Invalid move"); let prob = nodes[*child_idx].visit_count as f32 / denom; move_dist.push((enc, prob)); } @@ -280,7 +301,7 @@ impl Mcts { for mv in legal_moves { let stm = arena[node_idx].board_state.board.side_to_move(); - let idx = encode_move(mv, stm); + let idx = encode_move(mv, stm).expect("Invalid move"); let prior = policy[idx]; let mut new_board = arena[node_idx].board_state.clone(); diff --git a/engine/src/net/encoding.rs b/engine/src/net/encoding.rs index f4cdb13..b1526ba 100644 --- a/engine/src/net/encoding.rs +++ b/engine/src/net/encoding.rs @@ -20,6 +20,38 @@ Output: 65-73: underpromotions */ +const PLANES_PER_SQUARE: usize = 73; + +const SLIDING_DIRS: [(i8, i8); 8] = [ + (0, 1), // N + (1, 1), // NE + (1, 0), // E + (1, -1), // SE + (0, -1), // S + (-1, -1), // SW + (-1, 0), // W + (-1, 1), // NW +]; + +const KNIGHT_DIRS: [(i8, i8); 8] = [ + (1, 2), + (2, 1), + (2, -1), + (1, -2), + (-1, -2), + (-2, -1), + (-2, 1), + (-1, 2), +]; + +const UNDERPROMO_DIRS: [(i8, i8); 3] = [ + (-1, 1), // capture left + (0, 1), // forward + (1, 1), // capture right +]; + +const UNDERPROMO_PIECES: [Piece; 3] = [Piece::Knight, Piece::Bishop, Piece::Rook]; + pub fn encode_board_state_perspective( state: &BoardState, device: &B::Device, @@ -120,180 +152,163 @@ fn fill_plane(buffer: &mut [f32], plane: usize) { } } -pub fn encode_move(mv: ChessMove, side_to_move: Color) -> usize { - let from = mv.get_source().to_index(); - let to = mv.get_dest().to_index(); +/// Rotate a square into the side-to-move perspective. +fn orient_square(sq: Square, stm: Color) -> (i8, i8) { + let file = sq.get_file().to_index() as i8; + let rank = sq.get_rank().to_index() as i8; - let mut from_rank = from / 8; - let mut from_file = from % 8; - let mut to_rank = to / 8; - let mut to_file = to % 8; - - if side_to_move == Color::Black { - from_rank = 7 - from_rank; - from_file = 7 - from_file; - - to_rank = 7 - to_rank; - to_file = 7 - to_file; + match stm { + Color::White => (file, rank), + Color::Black => (7 - file, 7 - rank), } - - let delta_rank = to_rank as i32 - from_rank as i32; - let delta_file = to_file as i32 - from_file as i32; - - let plane = encode_move_type(delta_rank, delta_file, mv.get_promotion()); - - plane * 64 + (from_rank * 8 + from_file) } -fn encode_move_type(dr: i32, df: i32, promotion: Option) -> usize { - // Knight moves - const KNIGHT_DELTAS: [(i32, i32); 8] = [ - (2, 1), - (1, 2), - (-1, 2), - (-2, 1), - (-2, -1), - (-1, -2), - (1, -2), - (2, -1), - ]; +/// Convert a perspective-space coordinate back into a real square. +fn deorient_square(file: i8, rank: i8, stm: Color) -> Square { + let (file, rank) = match stm { + Color::White => (file, rank), + Color::Black => (7 - file, 7 - rank), + }; - for (i, (r, f)) in KNIGHT_DELTAS.iter().enumerate() { - if dr == *r && df == *f { - return 56 + i; - } - } + Square::make_square( + Rank::from_index(rank as usize), + File::from_index(file as usize), + ) +} - // UNDERPromotions - if let Some(promo) = promotion { +/// Encode a move into [0, 4672). +pub fn encode_move(mv: ChessMove, stm: Color) -> Option { + let (ff, fr) = orient_square(mv.get_source(), stm); + let (tf, tr) = orient_square(mv.get_dest(), stm); + + let dx = tf - ff; + let dy = tr - fr; + + let plane = if let Some(promo) = mv.get_promotion() { if promo != Piece::Queen { - let dir = if df == 0 { - 0 - } else if df < 0 { - 1 - } else { - 2 - }; + let dir = UNDERPROMO_DIRS + .iter() + .position(|&(x, y)| x == dx && y == dy)?; - let piece_index = match promo { - // Piece::Queen => 0, - Piece::Rook => 0, - Piece::Bishop => 1, - Piece::Knight => 2, - _ => unreachable!(), - }; + let piece = UNDERPROMO_PIECES.iter().position(|&p| p == promo)?; - return 64 + dir * 3 + piece_index; + 64 + piece * 3 + dir + } else { + encode_non_underpromo(dx, dy)? + } + } else { + encode_non_underpromo(dx, dy)? + }; + + let from_sq = (fr as usize) * 8 + (ff as usize); + Some(from_sq * PLANES_PER_SQUARE + plane) +} + +fn encode_non_underpromo(dx: i8, dy: i8) -> Option { + // Knight planes: 56..63 + if let Some(idx) = KNIGHT_DIRS.iter().position(|&(x, y)| x == dx && y == dy) { + return Some(56 + idx); + } + + // Sliding planes: 0..55 + for (dir_idx, &(sx, sy)) in SLIDING_DIRS.iter().enumerate() { + for dist in 1..=7 { + if dx == sx * dist && dy == sy * dist { + return Some(dir_idx * 7 + (dist as usize - 1)); + } } } - // Sliding - let direction_index = match (dr.signum(), df.signum()) { - (1, 0) => 0, // N - (1, 1) => 1, - (0, 1) => 2, - (-1, 1) => 3, - (-1, 0) => 4, - (-1, -1) => 5, - (0, -1) => 6, - (1, -1) => 7, - _ => panic!("Invalid move delta"), + None +} + +/// Decode an index in [0, 4672) back into a ChessMove. +pub fn decode_move(idx: usize, stm: Color) -> Option { + if idx >= 4672 { + return None; + } + + let from_idx = idx / PLANES_PER_SQUARE; + let plane = idx % PLANES_PER_SQUARE; + + let ff = (from_idx % 8) as i8; + let fr = (from_idx / 8) as i8; + + let (dx, dy, promo) = if plane < 56 { + let dir = plane / 7; + let dist = (plane % 7 + 1) as i8; + + let (sx, sy) = SLIDING_DIRS[dir]; + (sx * dist, sy * dist, None) + } else if plane < 64 { + let k = plane - 56; + let (dx, dy) = KNIGHT_DIRS[k]; + (dx, dy, None) + } else { + let p = plane - 64; + + let piece = UNDERPROMO_PIECES[p / 3]; + let (dx, dy) = UNDERPROMO_DIRS[p % 3]; + + (dx, dy, Some(piece)) }; - let distance = dr.abs().max(df.abs()) as usize - 1; + let tf = ff + dx; + let tr = fr + dy; - direction_index * 7 + distance -} - -pub fn decode_move(index: usize, side_to_move: Color) -> ChessMove { - let from_index = index % 64; - let plane = index / 64; - - // Perspective-space coordinates - let mut from_rank = from_index / 8; - let mut from_file = from_index % 8; - - let (mut dr, mut df, promotion) = decode_move_type(plane); - - // Convert from perspective coordinates back to absolute board coordinates - if side_to_move == Color::Black { - from_rank = 7 - from_rank; - from_file = 7 - from_file; - - dr = -dr; - df = -df; + if !(0..8).contains(&tf) || !(0..8).contains(&tr) { + return None; } - let to_rank = (from_rank as i32 + dr) as usize; - let to_file = (from_file as i32 + df) as usize; + let from = deorient_square(ff, fr, stm); + let to = deorient_square(tf, tr, stm); - let from = Square::make_square(Rank::from_index(from_rank), File::from_index(from_file)); - - let to = Square::make_square(Rank::from_index(to_rank), File::from_index(to_file)); - - ChessMove::new(from, to, promotion) + Some(ChessMove::new(from, to, promo)) } -fn decode_move_type(plane: usize) -> (i32, i32, Option) { - // Knight moves - const KNIGHT_DELTAS: [(i32, i32); 8] = [ - (2, 1), - (1, 2), - (-1, 2), - (-2, 1), - (-2, -1), - (-1, -2), - (1, -2), - (2, -1), - ]; +#[cfg(test)] +mod tests { + use super::*; + use chess::ALL_COLORS; - // 0–55: sliding moves - if plane < 56 { - let direction = plane / 7; - let distance = (plane % 7) + 1; - - let (dr, df) = match direction { - 0 => (1, 0), - 1 => (1, 1), - 2 => (0, 1), - 3 => (-1, 1), - 4 => (-1, 0), - 5 => (-1, -1), - 6 => (0, -1), - 7 => (1, -1), - _ => unreachable!(), - }; - - return (dr * distance as i32, df * distance as i32, None); + #[test] + fn encoding_roundtrips() { + for color in ALL_COLORS { + for action in 0..4672 { + let decoded = decode_move(action, color); + if decoded.is_none() { + continue; + } + let decoded = decoded.unwrap(); + let encoded = encode_move(decoded, color).unwrap(); + // eprintln!( + // "orig idx = {}, plane={}, from_idx={}", + // action, + // action / 64, + // action % 64 + // ); + // eprintln!( + // "move = {}, from={}, to={}, promo={:?}", + // decoded, + // decoded.get_source(), + // decoded.get_dest(), + // decoded.get_promotion() + // ); + // eprintln!( + // "decoded = {}, from={}, to={}, promo={:?}", + // decoded, + // decoded.get_source(), + // decoded.get_dest(), + // decoded.get_promotion() + // ); + // eprintln!( + // "re-encoded idx = {}, plane={}, from_idx={}", + // encoded, + // encoded / 64, + // encoded % 64 + // ); + assert_eq!(action, encoded); + } + } } - - // 56–63: knight moves - if plane < 64 { - let (dr, df) = KNIGHT_DELTAS[plane - 56]; - return (dr, df, None); - } - - // 64–72: underpromotions - let promo_plane = plane - 64; - - let dir = promo_plane / 3; - let piece_index = promo_plane % 3; - - let df = match dir { - 0 => 0, - 1 => -1, - 2 => 1, - _ => unreachable!(), - }; - - let dr = 1; // always forward (important: assumes white perspective) - - let promotion = Some(match piece_index { - 0 => Piece::Rook, - 1 => Piece::Bishop, - 2 => Piece::Knight, - _ => unreachable!(), - }); - - (dr, df, promotion) } diff --git a/engine/src/net/model.rs b/engine/src/net/model.rs index 0db3344..12cd024 100644 --- a/engine/src/net/model.rs +++ b/engine/src/net/model.rs @@ -32,6 +32,12 @@ Output: 65-73: underpromotions */ +#[derive(serde::Serialize, serde::Deserialize)] +pub struct ModelMetadata { + pub(crate) name: String, + pub(crate) iterations: usize, +} + #[derive(Module, Debug)] pub struct ResidualBlock { conv1: Conv2d, diff --git a/engine/src/training/train.rs b/engine/src/training/train.rs index a59deec..faadde2 100644 --- a/engine/src/training/train.rs +++ b/engine/src/training/train.rs @@ -1,6 +1,8 @@ use crate::mcts::{BoardState, BoardStateStatus, Mcts, MctsConfig, MctsResults}; use crate::net::encoding::decode_move; -use crate::net::model::{ChessBatcher, ChessModel, ChessModelConfig, TrainingSample}; +use crate::net::model::{ + ChessBatcher, ChessModel, ChessModelConfig, ModelMetadata, TrainingSample, +}; use burn::data::dataloader::batcher::Batcher; use burn::module::{AutodiffModule, Module}; use burn::optim::{AdamConfig, GradientsParams, Optimizer}; @@ -11,6 +13,8 @@ use rand::rngs::SmallRng; use rand::seq::SliceRandom; use rand::{RngExt, SeedableRng}; use std::collections::VecDeque; +use std::fs::File; +use std::io::BufReader; use std::marker::PhantomData; use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::Arc; @@ -19,7 +23,7 @@ use std::time::{SystemTime, UNIX_EPOCH}; pub struct TrainingConfig { pub max_time_s: Option, - pub num_iters: Option, + pub num_iters: Option, pub max_depth: u16, // unused pub model_name: String, pub load_model: bool, @@ -31,10 +35,12 @@ pub struct TrainingConfig { pub mcts_config: MctsConfig, pub optimizer: AdamConfig, pub lr: f64, + pub seed: Option, } pub fn train(training_config: TrainingConfig, device: B::Device) { let model_path = format!("artifacts/{}", training_config.model_name.as_str()); + let metadata_path = format!("{}.json", model_path); println!("Creating model..."); let mut model: ChessModel = ChessModelConfig::init( training_config.num_blocks, @@ -53,7 +59,20 @@ pub fn train(training_config: TrainingConfig, device: B::Dev let train = Arc::new(AtomicBool::new(true)); let train_signal = Arc::clone(&train); - let mut iter: u32 = 0; + let mut iter: usize = 0; + + if training_config.load_model { + let file = File::open(&metadata_path).unwrap(); + + // 2. Use a buffered reader for efficiency + let reader = BufReader::new(file); + + // 3. Deserialize JSON directly from the file reader + let metadata: ModelMetadata = serde_json::from_reader(reader).unwrap(); + iter = metadata.iterations; + println!("Loaded model had {} iters", iter); + } + let start_time = Instant::now(); ctrlc::set_handler(move || { @@ -72,7 +91,8 @@ pub fn train(training_config: TrainingConfig, device: B::Dev // Create RNG once and reuse it for sampling and shuffling // Seed from system time (platform default entropy may be unavailable in some contexts) let now = SystemTime::now().duration_since(UNIX_EPOCH).unwrap(); - let seed = now.as_nanos() as u64; + let seed = training_config.seed.unwrap_or(now.as_nanos() as u64); + println!("using seed: {}", seed); let mut rng = SmallRng::seed_from_u64(seed); // Initialize optimizer once so state (moments) persist across steps @@ -80,6 +100,7 @@ pub fn train(training_config: TrainingConfig, device: B::Dev println!("Starting training..."); while train.load(Ordering::Relaxed) { + println!("Iteration: {}", iter); let infer_model = model.valid(); // Gen samples println!("Generating {} games...", training_config.num_episodes); @@ -91,7 +112,9 @@ pub fn train(training_config: TrainingConfig, device: B::Dev let mut game_hist = String::new(); while board_state.status == BoardStateStatus::Ongoing { + // let before = Instant::now(); let results = mcts.search(&board_state, &infer_model, &device); + // println!("{}", before.elapsed().as_millis()); episode_buffer.push(results); let temp = if board_state.halfmove_clock < 30 { @@ -112,7 +135,9 @@ pub fn train(training_config: TrainingConfig, device: B::Dev board_state.apply_move(mv) } - println!("Game history of first game of iteration: {}", game_hist); + if episode == 0 { + println!("Game history of first game of iteration: {}", game_hist); + } for result in episode_buffer.iter().enumerate() { if board_state.status == BoardStateStatus::Stalemate @@ -188,6 +213,19 @@ pub fn train(training_config: TrainingConfig, device: B::Dev } println!("Saving model..."); + + let metadata = ModelMetadata { + name: training_config.model_name, + iterations: iter, + }; + + std::fs::write( + metadata_path, + serde_json::to_string_pretty(&metadata) + .expect("Should be able to convert metadata to JSON string"), + ) + .expect("Should be able to write metadata"); + // Save model in MessagePack format with full precision let recorder = NamedMpkFileRecorder::::new(); model @@ -243,10 +281,12 @@ fn sample_move( for (idx, p) in dist { r -= *p; if r <= 0.0 { - return Some(decode_move(*idx, side_to_move)); + return Some(decode_move(*idx, side_to_move).expect("Invalid move")); } } // fallback due to floating point drift - dist.get(0).map(|(idx, _)| decode_move(*idx, side_to_move)) + dist.get(0) + .map(|(idx, _)| decode_move(*idx, side_to_move).expect("Invalid move")) } +