Compare commits
4
Commits
be589de6f4
..
main
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c40a9667bd | ||
|
|
b9961d5645 | ||
|
|
ba3b962e86 | ||
|
|
23135b4386 |
@@ -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"
|
||||||
Generated
+2
@@ -2167,6 +2167,8 @@ dependencies = [
|
|||||||
"ctrlc",
|
"ctrlc",
|
||||||
"rand 0.10.1",
|
"rand 0.10.1",
|
||||||
"rand_distr 0.6.0",
|
"rand_distr 0.6.0",
|
||||||
|
"serde",
|
||||||
|
"serde_json",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
|
|||||||
@@ -10,3 +10,5 @@ burn-ndarray = "0.21.0"
|
|||||||
rand_distr = "0.6.0"
|
rand_distr = "0.6.0"
|
||||||
rand = "0.10.1"
|
rand = "0.10.1"
|
||||||
ctrlc = "3.5.2"
|
ctrlc = "3.5.2"
|
||||||
|
serde_json = "1.0.150"
|
||||||
|
serde = { version = "1.0.228", features = ["derive"] }
|
||||||
|
|||||||
+7
-6
@@ -1,6 +1,6 @@
|
|||||||
#![recursion_limit = "256"]
|
#![recursion_limit = "256"]
|
||||||
|
|
||||||
use burn::backend::{Autodiff, Wgpu};
|
use burn::backend::{Autodiff, Cuda};
|
||||||
use burn::optim::AdamConfig;
|
use burn::optim::AdamConfig;
|
||||||
use engine::mcts::MctsConfig;
|
use engine::mcts::MctsConfig;
|
||||||
use engine::training::train::{train, TrainingConfig};
|
use engine::training::train::{train, TrainingConfig};
|
||||||
@@ -14,15 +14,15 @@ use engine::training::train::{train, TrainingConfig};
|
|||||||
// }
|
// }
|
||||||
|
|
||||||
fn main() {
|
fn main() {
|
||||||
type MyBackend = Wgpu<f32, i32>;
|
// type MyBackend = Wgpu<f32, i32>;
|
||||||
// type MyBackend = Cuda<f32, i32>;
|
type MyBackend = Cuda<f32, i32>;
|
||||||
// type MyBackend = NdArray<f32, i32>;
|
// type MyBackend = NdArray<f32, i32>;
|
||||||
type MyAutodiffBackend = Autodiff<MyBackend>;
|
type MyAutodiffBackend = Autodiff<MyBackend>;
|
||||||
let device = burn::backend::wgpu::WgpuDevice::default();
|
// let device = burn::backend::wgpu::WgpuDevice::default();
|
||||||
// let device = burn::backend::ndarray::NdArrayDevice::default();
|
// let device = burn::backend::ndarray::NdArrayDevice::default();
|
||||||
// let device = burn::backend::cuda::CudaDevice::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();
|
let adam_config = AdamConfig::new();
|
||||||
|
|
||||||
@@ -40,6 +40,7 @@ fn main() {
|
|||||||
mcts_config,
|
mcts_config,
|
||||||
optimizer: adam_config,
|
optimizer: adam_config,
|
||||||
lr: 2e-4,
|
lr: 2e-4,
|
||||||
|
seed: None,
|
||||||
};
|
};
|
||||||
|
|
||||||
train::<MyAutodiffBackend>(training_config, device);
|
train::<MyAutodiffBackend>(training_config, device);
|
||||||
|
|||||||
+179
-49
@@ -5,9 +5,10 @@ use crate::net::model::ChessModel;
|
|||||||
use burn::prelude::Backend;
|
use burn::prelude::Backend;
|
||||||
use burn::Tensor;
|
use burn::Tensor;
|
||||||
use chess::BoardStatus::{Checkmate, Stalemate};
|
use chess::BoardStatus::{Checkmate, Stalemate};
|
||||||
use chess::Color::White;
|
use chess::Color::{Black, White};
|
||||||
use chess::Piece::{Bishop, Knight, Pawn, Queen, Rook};
|
use chess::Piece::{Bishop, Knight, Pawn, Queen, Rook};
|
||||||
use chess::{Board, ChessMove, Color, MoveGen, Piece, ALL_COLORS, ALL_PIECES};
|
use chess::{Board, ChessMove, Color, MoveGen, Piece, ALL_COLORS, ALL_PIECES};
|
||||||
|
use rand::SeedableRng;
|
||||||
use std::collections::HashMap;
|
use std::collections::HashMap;
|
||||||
use std::marker::PhantomData;
|
use std::marker::PhantomData;
|
||||||
|
|
||||||
@@ -58,16 +59,13 @@ impl Clone for Node {
|
|||||||
#[derive(Clone, Debug)]
|
#[derive(Clone, Debug)]
|
||||||
pub struct MctsResults {
|
pub struct MctsResults {
|
||||||
pub board_state: BoardState,
|
pub board_state: BoardState,
|
||||||
pub move_dist: HashMap<ChessMove, f32>,
|
// Compact encoded move distribution: (encoded_move_index, probability)
|
||||||
|
pub move_dist: Vec<(usize, f32)>,
|
||||||
pub value: f32,
|
pub value: f32,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl MctsResults {
|
impl MctsResults {
|
||||||
pub fn new(
|
pub fn new(board_state: BoardState, move_dist: Vec<(usize, f32)>, value: f32) -> MctsResults {
|
||||||
board_state: BoardState,
|
|
||||||
move_dist: HashMap<ChessMove, f32>,
|
|
||||||
value: f32,
|
|
||||||
) -> MctsResults {
|
|
||||||
MctsResults {
|
MctsResults {
|
||||||
board_state,
|
board_state,
|
||||||
move_dist,
|
move_dist,
|
||||||
@@ -82,6 +80,7 @@ pub struct MctsConfig {
|
|||||||
pub c_puct: f32,
|
pub c_puct: f32,
|
||||||
pub dirichlet_alpha: f32,
|
pub dirichlet_alpha: f32,
|
||||||
pub dirichlet_epsilon: f32,
|
pub dirichlet_epsilon: f32,
|
||||||
|
pub batch_max: usize,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl MctsConfig {
|
impl MctsConfig {
|
||||||
@@ -90,19 +89,21 @@ impl MctsConfig {
|
|||||||
c_puct: f32,
|
c_puct: f32,
|
||||||
dirichlet_alpha: f32,
|
dirichlet_alpha: f32,
|
||||||
dirichlet_epsilon: f32,
|
dirichlet_epsilon: f32,
|
||||||
|
batch_max: usize,
|
||||||
) -> MctsConfig {
|
) -> MctsConfig {
|
||||||
MctsConfig {
|
MctsConfig {
|
||||||
num_simulations,
|
num_simulations,
|
||||||
c_puct,
|
c_puct,
|
||||||
dirichlet_alpha,
|
dirichlet_alpha,
|
||||||
dirichlet_epsilon,
|
dirichlet_epsilon,
|
||||||
|
batch_max,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl Default for MctsConfig {
|
impl Default for MctsConfig {
|
||||||
fn default() -> MctsConfig {
|
fn default() -> MctsConfig {
|
||||||
MctsConfig::new(400, 1.0, 0.05, 0.25)
|
MctsConfig::new(400, 1.0, 0.05, 0.25, 64)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -123,32 +124,139 @@ impl<B: Backend> Mcts<B> {
|
|||||||
let root = 0;
|
let root = 0;
|
||||||
nodes.push(Node::new(0.0, board_state.clone(), None));
|
nodes.push(Node::new(0.0, board_state.clone(), None));
|
||||||
|
|
||||||
|
// Expand root to create initial children and priors
|
||||||
self.expand(root, &mut nodes, model, device);
|
self.expand(root, &mut nodes, model, device);
|
||||||
|
|
||||||
// 👇 APPLY DIRICHLET NOISE HERE
|
// Apply Dirichlet noise to root children
|
||||||
self.add_dirichlet_noise(root, &mut nodes);
|
self.add_dirichlet_noise(root, &mut nodes);
|
||||||
|
|
||||||
for _ in 0..self.config.num_simulations {
|
// We'll batch leaf evaluations to reduce per-leaf model calls and device-host syncs.
|
||||||
let mut path = vec![root];
|
let mut sims_done: usize = 0;
|
||||||
let mut current = root;
|
let num_sims = self.config.num_simulations;
|
||||||
|
|
||||||
while !nodes[current].children.is_empty() {
|
while sims_done < num_sims {
|
||||||
current = nodes[current].select_child(&nodes, &self.config.c_puct);
|
// Collect a batch of leaf nodes (and their selection paths)
|
||||||
path.push(current);
|
let mut leaf_nodes: Vec<usize> = Vec::new();
|
||||||
|
let mut leaf_paths: Vec<Vec<usize>> = Vec::new();
|
||||||
|
let mut leaf_states: Vec<Tensor<B, 4>> = Vec::new();
|
||||||
|
|
||||||
|
while leaf_nodes.len() < std::cmp::min(self.config.batch_max, num_sims - sims_done) {
|
||||||
|
let mut path = vec![root];
|
||||||
|
let mut current = root;
|
||||||
|
|
||||||
|
while !nodes[current].children.is_empty() {
|
||||||
|
current = nodes[current].select_child(&nodes, &self.config.c_puct);
|
||||||
|
path.push(current);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Record leaf node and its path
|
||||||
|
leaf_nodes.push(current);
|
||||||
|
leaf_paths.push(path.clone());
|
||||||
|
|
||||||
|
// Prepare state tensor for this leaf
|
||||||
|
let state: Tensor<B, 4> =
|
||||||
|
encode_board_state_perspective(&nodes[current].board_state, device)
|
||||||
|
.reshape([1, 18, 8, 8]);
|
||||||
|
leaf_states.push(state);
|
||||||
|
|
||||||
|
sims_done += 1;
|
||||||
}
|
}
|
||||||
|
|
||||||
let value: f32 = self.expand(current, &mut nodes, model, device);
|
if leaf_nodes.is_empty() {
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
|
||||||
let color = nodes[current].board_state.board.side_to_move();
|
// Batch evaluate the collected leaf states
|
||||||
self.backpropagate(&mut nodes, &path, value, color);
|
let batch = Tensor::cat(leaf_states, 0);
|
||||||
|
let (policy_batch, value_batch) = model.forward(batch);
|
||||||
|
|
||||||
|
// Move tensors to host once per batch
|
||||||
|
let policy_data = policy_batch.into_data().to_vec::<f32>().unwrap();
|
||||||
|
let value_data = value_batch.into_data().to_vec::<f32>().unwrap();
|
||||||
|
|
||||||
|
let num_moves = policy_data.len() / leaf_nodes.len();
|
||||||
|
|
||||||
|
// Process each evaluated leaf: expand and backpropagate
|
||||||
|
for (i, &node_idx) in leaf_nodes.iter().enumerate() {
|
||||||
|
let path = &leaf_paths[i];
|
||||||
|
|
||||||
|
// slice for this sample's logits
|
||||||
|
let start = i * num_moves;
|
||||||
|
let end = start + num_moves;
|
||||||
|
let logits = &policy_data[start..end];
|
||||||
|
|
||||||
|
// Convert logits to probabilities with a numerically-stable softmax on host
|
||||||
|
let mut max_logit = f32::NEG_INFINITY;
|
||||||
|
for &v in logits.iter() {
|
||||||
|
if v > max_logit {
|
||||||
|
max_logit = v;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if logits.is_empty() {
|
||||||
|
println!("logits empty")
|
||||||
|
}
|
||||||
|
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<ChessMove> =
|
||||||
|
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).expect("Invalid move");
|
||||||
|
let prior = probs[idx];
|
||||||
|
|
||||||
|
let mut new_board = nodes[node_idx].board_state.clone();
|
||||||
|
new_board.apply_move(mv);
|
||||||
|
|
||||||
|
let child_idx = nodes.len();
|
||||||
|
nodes.push(Node::new(prior, new_board, Some(mv)));
|
||||||
|
nodes[node_idx].children.push(child_idx);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Backpropagate the value for this leaf
|
||||||
|
let value = value_data[i];
|
||||||
|
let color = nodes[node_idx].board_state.board.side_to_move();
|
||||||
|
self.backpropagate(&mut nodes, path, value, color);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
let mut move_dist: HashMap<ChessMove, f32> = HashMap::new(); // TODO: make vec<(Chessmove, f32)>
|
// Build compact move distribution: encoded move index -> probability (visits / num_simulations)
|
||||||
for idx in nodes[root].children.iter() {
|
let mut move_dist: Vec<(usize, f32)> = Vec::with_capacity(nodes[root].children.len());
|
||||||
move_dist.insert(
|
let stm = board_state.board.side_to_move();
|
||||||
nodes[*idx].last_move.expect("move didnt exist"),
|
let denom = self.config.num_simulations as f32;
|
||||||
nodes[*idx].visit_count as f32 / 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).expect("Invalid move");
|
||||||
|
let prob = nodes[*child_idx].visit_count as f32 / denom;
|
||||||
|
move_dist.push((enc, prob));
|
||||||
}
|
}
|
||||||
|
|
||||||
MctsResults::new(board_state.clone(), move_dist, nodes[root].value())
|
MctsResults::new(board_state.clone(), move_dist, nodes[root].value())
|
||||||
@@ -161,33 +269,52 @@ impl<B: Backend> Mcts<B> {
|
|||||||
model: &ChessModel<B>,
|
model: &ChessModel<B>,
|
||||||
device: &B::Device,
|
device: &B::Device,
|
||||||
) -> f32 {
|
) -> f32 {
|
||||||
let state: Tensor<B, 4> =
|
if arena[node_idx].board_state.status == BoardStateStatus::Stalemate
|
||||||
encode_board_state_perspective(&arena[node_idx].board_state, device)
|
|| arena[node_idx].board_state.status == BoardStateStatus::Threefold
|
||||||
.reshape([1, 18, 8, 8]);
|
|| arena[node_idx].board_state.status == BoardStateStatus::FiftyMove
|
||||||
// let start = Instant::now();
|
{
|
||||||
let (policy_head, value_head) = model.forward(state);
|
0.0
|
||||||
// println!("time: {:?}", start.elapsed());
|
} else if arena[node_idx].board_state.status == BoardStateStatus::WhiteWinner {
|
||||||
|
if arena[node_idx].board_state.board.side_to_move() == Black {
|
||||||
|
1.0
|
||||||
|
} else {
|
||||||
|
-1.0
|
||||||
|
}
|
||||||
|
} else if arena[node_idx].board_state.status == BoardStateStatus::BlackWinner {
|
||||||
|
if arena[node_idx].board_state.board.side_to_move() == White {
|
||||||
|
1.0
|
||||||
|
} else {
|
||||||
|
-1.0
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
let state: Tensor<B, 4> =
|
||||||
|
encode_board_state_perspective(&arena[node_idx].board_state, device)
|
||||||
|
.reshape([1, 18, 8, 8]);
|
||||||
|
// let start = Instant::now();
|
||||||
|
let (policy_head, value_head) = model.forward(state);
|
||||||
|
// println!("time: {:?}", start.elapsed());
|
||||||
|
|
||||||
let legal_moves: Vec<ChessMove> =
|
let legal_moves: Vec<ChessMove> =
|
||||||
MoveGen::new_legal(&arena[node_idx].board_state.board).collect();
|
MoveGen::new_legal(&arena[node_idx].board_state.board).collect();
|
||||||
|
|
||||||
let policy = policy_head.into_data().to_vec::<f32>().unwrap();
|
let policy = policy_head.into_data().to_vec::<f32>().unwrap();
|
||||||
|
|
||||||
for mv in legal_moves {
|
for mv in legal_moves {
|
||||||
let stm = arena[node_idx].board_state.board.side_to_move();
|
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 prior = policy[idx];
|
||||||
|
|
||||||
let mut new_board = arena[node_idx].board_state.clone();
|
let mut new_board = arena[node_idx].board_state.clone();
|
||||||
new_board.apply_move(mv);
|
new_board.apply_move(mv);
|
||||||
|
|
||||||
let child_idx = arena.len();
|
let child_idx = arena.len();
|
||||||
|
|
||||||
arena.push(Node::new(prior, new_board, Some(mv)));
|
arena.push(Node::new(prior, new_board, Some(mv)));
|
||||||
arena[node_idx].children.push(child_idx);
|
arena[node_idx].children.push(child_idx);
|
||||||
|
}
|
||||||
|
|
||||||
|
value_head.into_data().to_vec().unwrap()[0]
|
||||||
}
|
}
|
||||||
|
|
||||||
value_head.into_data().to_vec().unwrap()[0]
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn backpropagate(&mut self, nodes: &mut [Node], path: &[usize], value: f32, color: Color) {
|
fn backpropagate(&mut self, nodes: &mut [Node], path: &[usize], value: f32, color: Color) {
|
||||||
@@ -229,9 +356,14 @@ fn dirichlet_sample(size: usize, alpha: f32) -> Vec<f32> {
|
|||||||
|
|
||||||
let gamma = Gamma::new(alpha as f64, 1.0).unwrap();
|
let gamma = Gamma::new(alpha as f64, 1.0).unwrap();
|
||||||
|
|
||||||
let mut samples: Vec<f32> = (0..size)
|
// Use a single SmallRng seeded from system time (avoid depending on thread_rng helper)
|
||||||
.map(|_| gamma.sample(&mut rand::rng()) as f32)
|
let now = std::time::SystemTime::now()
|
||||||
.collect();
|
.duration_since(std::time::UNIX_EPOCH)
|
||||||
|
.unwrap();
|
||||||
|
let seed = now.as_nanos() as u64;
|
||||||
|
let mut rng = rand::rngs::SmallRng::seed_from_u64(seed);
|
||||||
|
|
||||||
|
let mut samples: Vec<f32> = (0..size).map(|_| gamma.sample(&mut rng) as f32).collect();
|
||||||
|
|
||||||
let sum: f32 = samples.iter().sum();
|
let sum: f32 = samples.iter().sum();
|
||||||
|
|
||||||
@@ -334,8 +466,6 @@ pub fn heuristic_eval(board: &Board, perspective: Color) -> f32 {
|
|||||||
}
|
}
|
||||||
|
|
||||||
value
|
value
|
||||||
|
|
||||||
// board.checkers()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||||
|
|||||||
+169
-154
@@ -20,6 +20,38 @@ Output:
|
|||||||
65-73: underpromotions
|
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<B: Backend>(
|
pub fn encode_board_state_perspective<B: Backend>(
|
||||||
state: &BoardState,
|
state: &BoardState,
|
||||||
device: &B::Device,
|
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 {
|
/// Rotate a square into the side-to-move perspective.
|
||||||
let from = mv.get_source().to_index();
|
fn orient_square(sq: Square, stm: Color) -> (i8, i8) {
|
||||||
let to = mv.get_dest().to_index();
|
let file = sq.get_file().to_index() as i8;
|
||||||
|
let rank = sq.get_rank().to_index() as i8;
|
||||||
|
|
||||||
let mut from_rank = from / 8;
|
match stm {
|
||||||
let mut from_file = from % 8;
|
Color::White => (file, rank),
|
||||||
let mut to_rank = to / 8;
|
Color::Black => (7 - file, 7 - rank),
|
||||||
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;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
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<Piece>) -> usize {
|
/// Convert a perspective-space coordinate back into a real square.
|
||||||
// Knight moves
|
fn deorient_square(file: i8, rank: i8, stm: Color) -> Square {
|
||||||
const KNIGHT_DELTAS: [(i32, i32); 8] = [
|
let (file, rank) = match stm {
|
||||||
(2, 1),
|
Color::White => (file, rank),
|
||||||
(1, 2),
|
Color::Black => (7 - file, 7 - rank),
|
||||||
(-1, 2),
|
};
|
||||||
(-2, 1),
|
|
||||||
(-2, -1),
|
|
||||||
(-1, -2),
|
|
||||||
(1, -2),
|
|
||||||
(2, -1),
|
|
||||||
];
|
|
||||||
|
|
||||||
for (i, (r, f)) in KNIGHT_DELTAS.iter().enumerate() {
|
Square::make_square(
|
||||||
if dr == *r && df == *f {
|
Rank::from_index(rank as usize),
|
||||||
return 56 + i;
|
File::from_index(file as usize),
|
||||||
}
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
// UNDERPromotions
|
/// Encode a move into [0, 4672).
|
||||||
if let Some(promo) = promotion {
|
pub fn encode_move(mv: ChessMove, stm: Color) -> Option<usize> {
|
||||||
|
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 {
|
if promo != Piece::Queen {
|
||||||
let dir = if df == 0 {
|
let dir = UNDERPROMO_DIRS
|
||||||
0
|
.iter()
|
||||||
} else if df < 0 {
|
.position(|&(x, y)| x == dx && y == dy)?;
|
||||||
1
|
|
||||||
} else {
|
|
||||||
2
|
|
||||||
};
|
|
||||||
|
|
||||||
let piece_index = match promo {
|
let piece = UNDERPROMO_PIECES.iter().position(|&p| p == promo)?;
|
||||||
// Piece::Queen => 0,
|
|
||||||
Piece::Rook => 0,
|
|
||||||
Piece::Bishop => 1,
|
|
||||||
Piece::Knight => 2,
|
|
||||||
_ => unreachable!(),
|
|
||||||
};
|
|
||||||
|
|
||||||
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<usize> {
|
||||||
|
// 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
|
None
|
||||||
let direction_index = match (dr.signum(), df.signum()) {
|
}
|
||||||
(1, 0) => 0, // N
|
|
||||||
(1, 1) => 1,
|
/// Decode an index in [0, 4672) back into a ChessMove.
|
||||||
(0, 1) => 2,
|
pub fn decode_move(idx: usize, stm: Color) -> Option<ChessMove> {
|
||||||
(-1, 1) => 3,
|
if idx >= 4672 {
|
||||||
(-1, 0) => 4,
|
return None;
|
||||||
(-1, -1) => 5,
|
}
|
||||||
(0, -1) => 6,
|
|
||||||
(1, -1) => 7,
|
let from_idx = idx / PLANES_PER_SQUARE;
|
||||||
_ => panic!("Invalid move delta"),
|
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
|
if !(0..8).contains(&tf) || !(0..8).contains(&tr) {
|
||||||
}
|
return None;
|
||||||
|
|
||||||
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;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
let to_rank = (from_rank as i32 + dr) as usize;
|
let from = deorient_square(ff, fr, stm);
|
||||||
let to_file = (from_file as i32 + df) as usize;
|
let to = deorient_square(tf, tr, stm);
|
||||||
|
|
||||||
let from = Square::make_square(Rank::from_index(from_rank), File::from_index(from_file));
|
Some(ChessMove::new(from, to, promo))
|
||||||
|
|
||||||
let to = Square::make_square(Rank::from_index(to_rank), File::from_index(to_file));
|
|
||||||
|
|
||||||
ChessMove::new(from, to, promotion)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn decode_move_type(plane: usize) -> (i32, i32, Option<Piece>) {
|
#[cfg(test)]
|
||||||
// Knight moves
|
mod tests {
|
||||||
const KNIGHT_DELTAS: [(i32, i32); 8] = [
|
use super::*;
|
||||||
(2, 1),
|
use chess::ALL_COLORS;
|
||||||
(1, 2),
|
|
||||||
(-1, 2),
|
|
||||||
(-2, 1),
|
|
||||||
(-2, -1),
|
|
||||||
(-1, -2),
|
|
||||||
(1, -2),
|
|
||||||
(2, -1),
|
|
||||||
];
|
|
||||||
|
|
||||||
// 0–55: sliding moves
|
#[test]
|
||||||
if plane < 56 {
|
fn encoding_roundtrips() {
|
||||||
let direction = plane / 7;
|
for color in ALL_COLORS {
|
||||||
let distance = (plane % 7) + 1;
|
for action in 0..4672 {
|
||||||
|
let decoded = decode_move(action, color);
|
||||||
let (dr, df) = match direction {
|
if decoded.is_none() {
|
||||||
0 => (1, 0),
|
continue;
|
||||||
1 => (1, 1),
|
}
|
||||||
2 => (0, 1),
|
let decoded = decoded.unwrap();
|
||||||
3 => (-1, 1),
|
let encoded = encode_move(decoded, color).unwrap();
|
||||||
4 => (-1, 0),
|
// eprintln!(
|
||||||
5 => (-1, -1),
|
// "orig idx = {}, plane={}, from_idx={}",
|
||||||
6 => (0, -1),
|
// action,
|
||||||
7 => (1, -1),
|
// action / 64,
|
||||||
_ => unreachable!(),
|
// action % 64
|
||||||
};
|
// );
|
||||||
|
// eprintln!(
|
||||||
return (dr * distance as i32, df * distance as i32, None);
|
// "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)
|
|
||||||
}
|
}
|
||||||
|
|||||||
+13
-8
@@ -1,5 +1,5 @@
|
|||||||
use crate::mcts::{BoardState, MctsResults};
|
use crate::mcts::{BoardState, MctsResults};
|
||||||
use crate::net::encoding::{encode_board_state_perspective, encode_move};
|
use crate::net::encoding::encode_board_state_perspective;
|
||||||
use burn::data::dataloader::batcher::Batcher;
|
use burn::data::dataloader::batcher::Batcher;
|
||||||
use burn::nn::conv::Conv2dConfig;
|
use burn::nn::conv::Conv2dConfig;
|
||||||
use burn::nn::loss::{MseLoss, Reduction};
|
use burn::nn::loss::{MseLoss, Reduction};
|
||||||
@@ -13,8 +13,6 @@ use burn::{
|
|||||||
prelude::*,
|
prelude::*,
|
||||||
};
|
};
|
||||||
use burn_ndarray::NdArray;
|
use burn_ndarray::NdArray;
|
||||||
use chess::ChessMove;
|
|
||||||
use std::collections::HashMap;
|
|
||||||
/*
|
/*
|
||||||
Input planes:
|
Input planes:
|
||||||
1-6: your pieces (Pawn, Knight, Bishop, Rook, Queen, King)
|
1-6: your pieces (Pawn, Knight, Bishop, Rook, Queen, King)
|
||||||
@@ -34,6 +32,12 @@ Output:
|
|||||||
65-73: underpromotions
|
65-73: underpromotions
|
||||||
*/
|
*/
|
||||||
|
|
||||||
|
#[derive(serde::Serialize, serde::Deserialize)]
|
||||||
|
pub struct ModelMetadata {
|
||||||
|
pub(crate) name: String,
|
||||||
|
pub(crate) iterations: usize,
|
||||||
|
}
|
||||||
|
|
||||||
#[derive(Module, Debug)]
|
#[derive(Module, Debug)]
|
||||||
pub struct ResidualBlock<B: Backend> {
|
pub struct ResidualBlock<B: Backend> {
|
||||||
conv1: Conv2d<B>,
|
conv1: Conv2d<B>,
|
||||||
@@ -255,14 +259,15 @@ impl<B: Backend> InferenceStep for ChessModel<B> {
|
|||||||
#[derive(Clone)]
|
#[derive(Clone)]
|
||||||
pub struct TrainingSample {
|
pub struct TrainingSample {
|
||||||
pub board_state: BoardState,
|
pub board_state: BoardState,
|
||||||
pub policy_target: HashMap<ChessMove, f32>,
|
// Compact representation: list of (encoded_move_index, probability)
|
||||||
|
pub policy_target: Vec<(usize, f32)>,
|
||||||
pub value_target: f32,
|
pub value_target: f32,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl TrainingSample {
|
impl TrainingSample {
|
||||||
pub fn new(
|
pub fn new(
|
||||||
board_state: BoardState,
|
board_state: BoardState,
|
||||||
policy_target: HashMap<ChessMove, f32>,
|
policy_target: Vec<(usize, f32)>,
|
||||||
value_target: f32,
|
value_target: f32,
|
||||||
) -> Self {
|
) -> Self {
|
||||||
TrainingSample {
|
TrainingSample {
|
||||||
@@ -273,6 +278,7 @@ impl TrainingSample {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub fn from_mcts_with_outcome(mcts_results: MctsResults, outcome: f32) -> Self {
|
pub fn from_mcts_with_outcome(mcts_results: MctsResults, outcome: f32) -> Self {
|
||||||
|
// move_dist is already a compact Vec<(encoded_move_index, prob)>
|
||||||
TrainingSample::new(mcts_results.board_state, mcts_results.move_dist, outcome)
|
TrainingSample::new(mcts_results.board_state, mcts_results.move_dist, outcome)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -301,9 +307,8 @@ impl<B: Backend> Batcher<B, TrainingSample, ChessBatch<B>> for ChessBatcher {
|
|||||||
.cloned()
|
.cloned()
|
||||||
.map(|item| {
|
.map(|item| {
|
||||||
let mut policy = vec![0.0f32; 4672];
|
let mut policy = vec![0.0f32; 4672];
|
||||||
let stm = item.board_state.board.side_to_move();
|
for (idx, prob) in item.policy_target.iter() {
|
||||||
for (mv, prob) in item.policy_target {
|
policy[*idx] = *prob;
|
||||||
policy[encode_move(mv, stm)] = prob;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Normalize
|
// Normalize
|
||||||
|
|||||||
@@ -1,23 +1,29 @@
|
|||||||
use crate::mcts::{BoardState, BoardStateStatus, Mcts, MctsConfig, MctsResults};
|
use crate::mcts::{BoardState, BoardStateStatus, Mcts, MctsConfig, MctsResults};
|
||||||
use crate::net::model::{ChessBatcher, ChessModel, ChessModelConfig, TrainingSample};
|
use crate::net::encoding::decode_move;
|
||||||
|
use crate::net::model::{
|
||||||
|
ChessBatcher, ChessModel, ChessModelConfig, ModelMetadata, TrainingSample,
|
||||||
|
};
|
||||||
use burn::data::dataloader::batcher::Batcher;
|
use burn::data::dataloader::batcher::Batcher;
|
||||||
use burn::module::{AutodiffModule, Module};
|
use burn::module::{AutodiffModule, Module};
|
||||||
use burn::optim::{AdamConfig, GradientsParams, Optimizer};
|
use burn::optim::{AdamConfig, GradientsParams, Optimizer};
|
||||||
use burn::record::{FullPrecisionSettings, NamedMpkFileRecorder};
|
use burn::record::{FullPrecisionSettings, NamedMpkFileRecorder};
|
||||||
use burn::tensor::backend::AutodiffBackend;
|
use burn::tensor::backend::AutodiffBackend;
|
||||||
use chess::ChessMove;
|
use chess::{ChessMove, Color};
|
||||||
use rand::rngs::ThreadRng;
|
use rand::rngs::SmallRng;
|
||||||
use rand::seq::SliceRandom;
|
use rand::seq::SliceRandom;
|
||||||
use rand::RngExt;
|
use rand::{RngExt, SeedableRng};
|
||||||
use std::collections::{HashMap, VecDeque};
|
use std::collections::VecDeque;
|
||||||
|
use std::fs::File;
|
||||||
|
use std::io::BufReader;
|
||||||
use std::marker::PhantomData;
|
use std::marker::PhantomData;
|
||||||
use std::sync::atomic::{AtomicBool, Ordering};
|
use std::sync::atomic::{AtomicBool, Ordering};
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
use std::time::Instant;
|
use std::time::Instant;
|
||||||
|
use std::time::{SystemTime, UNIX_EPOCH};
|
||||||
|
|
||||||
pub struct TrainingConfig {
|
pub struct TrainingConfig {
|
||||||
pub max_time_s: Option<u64>,
|
pub max_time_s: Option<u64>,
|
||||||
pub num_iters: Option<u32>,
|
pub num_iters: Option<usize>,
|
||||||
pub max_depth: u16, // unused
|
pub max_depth: u16, // unused
|
||||||
pub model_name: String,
|
pub model_name: String,
|
||||||
pub load_model: bool,
|
pub load_model: bool,
|
||||||
@@ -29,14 +35,16 @@ pub struct TrainingConfig {
|
|||||||
pub mcts_config: MctsConfig,
|
pub mcts_config: MctsConfig,
|
||||||
pub optimizer: AdamConfig,
|
pub optimizer: AdamConfig,
|
||||||
pub lr: f64,
|
pub lr: f64,
|
||||||
|
pub seed: Option<u64>,
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn train<B: AutodiffBackend>(training_config: TrainingConfig, device: B::Device) {
|
pub fn train<B: AutodiffBackend>(training_config: TrainingConfig, device: B::Device) {
|
||||||
let model_path = format!("artifacts/{}", training_config.model_name.as_str());
|
let model_path = format!("artifacts/{}", training_config.model_name.as_str());
|
||||||
|
let metadata_path = format!("{}.json", model_path);
|
||||||
println!("Creating model...");
|
println!("Creating model...");
|
||||||
let mut model: ChessModel<B> = ChessModelConfig::init(
|
let mut model: ChessModel<B> = ChessModelConfig::init(
|
||||||
training_config.hidden_channels,
|
|
||||||
training_config.num_blocks,
|
training_config.num_blocks,
|
||||||
|
training_config.hidden_channels,
|
||||||
&device,
|
&device,
|
||||||
);
|
);
|
||||||
if training_config.load_model {
|
if training_config.load_model {
|
||||||
@@ -51,7 +59,20 @@ pub fn train<B: AutodiffBackend>(training_config: TrainingConfig, device: B::Dev
|
|||||||
let train = Arc::new(AtomicBool::new(true));
|
let train = Arc::new(AtomicBool::new(true));
|
||||||
let train_signal = Arc::clone(&train);
|
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();
|
let start_time = Instant::now();
|
||||||
|
|
||||||
ctrlc::set_handler(move || {
|
ctrlc::set_handler(move || {
|
||||||
@@ -67,10 +88,19 @@ pub fn train<B: AutodiffBackend>(training_config: TrainingConfig, device: B::Dev
|
|||||||
_marker: PhantomData,
|
_marker: PhantomData,
|
||||||
};
|
};
|
||||||
|
|
||||||
let mut rng = rand::rng();
|
// 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 = 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
|
||||||
|
let mut optim = training_config.optimizer.init();
|
||||||
|
|
||||||
println!("Starting training...");
|
println!("Starting training...");
|
||||||
while train.load(Ordering::Relaxed) {
|
while train.load(Ordering::Relaxed) {
|
||||||
|
println!("Iteration: {}", iter);
|
||||||
let infer_model = model.valid();
|
let infer_model = model.valid();
|
||||||
// Gen samples
|
// Gen samples
|
||||||
println!("Generating {} games...", training_config.num_episodes);
|
println!("Generating {} games...", training_config.num_episodes);
|
||||||
@@ -79,8 +109,12 @@ pub fn train<B: AutodiffBackend>(training_config: TrainingConfig, device: B::Dev
|
|||||||
let mut board_state = BoardState::default();
|
let mut board_state = BoardState::default();
|
||||||
let mut episode_buffer: Vec<MctsResults> = vec![];
|
let mut episode_buffer: Vec<MctsResults> = vec![];
|
||||||
|
|
||||||
|
let mut game_hist = String::new();
|
||||||
|
|
||||||
while board_state.status == BoardStateStatus::Ongoing {
|
while board_state.status == BoardStateStatus::Ongoing {
|
||||||
|
// let before = Instant::now();
|
||||||
let results = mcts.search(&board_state, &infer_model, &device);
|
let results = mcts.search(&board_state, &infer_model, &device);
|
||||||
|
// println!("{}", before.elapsed().as_millis());
|
||||||
episode_buffer.push(results);
|
episode_buffer.push(results);
|
||||||
|
|
||||||
let temp = if board_state.halfmove_clock < 30 {
|
let temp = if board_state.halfmove_clock < 30 {
|
||||||
@@ -91,10 +125,19 @@ pub fn train<B: AutodiffBackend>(training_config: TrainingConfig, device: B::Dev
|
|||||||
|
|
||||||
let adjusted = apply_temperature(&episode_buffer.last().unwrap().move_dist, temp);
|
let adjusted = apply_temperature(&episode_buffer.last().unwrap().move_dist, temp);
|
||||||
|
|
||||||
let mv = sample_move(&adjusted, &mut rng).unwrap();
|
let stm = board_state.board.side_to_move();
|
||||||
println!("playing move: {}", mv);
|
let mv = sample_move(&adjusted, &mut rng, stm).unwrap();
|
||||||
|
// println!("playing move: {}", mv);
|
||||||
|
if episode == 0 {
|
||||||
|
game_hist.push_str(mv.to_string().as_str());
|
||||||
|
game_hist.push(' ');
|
||||||
|
}
|
||||||
|
|
||||||
board_state.apply_move(mv)
|
board_state.apply_move(mv)
|
||||||
}
|
}
|
||||||
|
if episode == 0 {
|
||||||
|
println!("Game history of first game of iteration: {}", game_hist);
|
||||||
|
}
|
||||||
|
|
||||||
for result in episode_buffer.iter().enumerate() {
|
for result in episode_buffer.iter().enumerate() {
|
||||||
if board_state.status == BoardStateStatus::Stalemate
|
if board_state.status == BoardStateStatus::Stalemate
|
||||||
@@ -140,8 +183,6 @@ pub fn train<B: AutodiffBackend>(training_config: TrainingConfig, device: B::Dev
|
|||||||
|
|
||||||
let batch = batcher.batch(samples, &device);
|
let batch = batcher.batch(samples, &device);
|
||||||
|
|
||||||
let mut optim = training_config.optimizer.init();
|
|
||||||
|
|
||||||
let output = model.forward_chess(batch.states, batch.policy_targets, batch.value_targets);
|
let output = model.forward_chess(batch.states, batch.policy_targets, batch.value_targets);
|
||||||
|
|
||||||
let grads = output.loss.backward();
|
let grads = output.loss.backward();
|
||||||
@@ -172,6 +213,19 @@ pub fn train<B: AutodiffBackend>(training_config: TrainingConfig, device: B::Dev
|
|||||||
}
|
}
|
||||||
|
|
||||||
println!("Saving model...");
|
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
|
// Save model in MessagePack format with full precision
|
||||||
let recorder = NamedMpkFileRecorder::<FullPrecisionSettings>::new();
|
let recorder = NamedMpkFileRecorder::<FullPrecisionSettings>::new();
|
||||||
model
|
model
|
||||||
@@ -181,56 +235,58 @@ pub fn train<B: AutodiffBackend>(training_config: TrainingConfig, device: B::Dev
|
|||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
|
|
||||||
fn apply_temperature(
|
fn apply_temperature(visits: &[(usize, f32)], temperature: f32) -> Vec<(usize, f32)> {
|
||||||
visits: &HashMap<ChessMove, f32>,
|
|
||||||
temperature: f32,
|
|
||||||
) -> HashMap<ChessMove, f32> {
|
|
||||||
if visits.is_empty() {
|
if visits.is_empty() {
|
||||||
return HashMap::new();
|
return Vec::new();
|
||||||
}
|
}
|
||||||
|
|
||||||
// Special case: deterministic selection
|
// Special case: deterministic selection
|
||||||
if temperature == 0.0 {
|
if temperature == 0.0 {
|
||||||
let (&best_move, _) = visits
|
let (&best_idx, _) = visits
|
||||||
.iter()
|
.iter()
|
||||||
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
|
.max_by(|a, b| a.1.partial_cmp(&b.1).unwrap())
|
||||||
|
.map(|(i, p)| (i, p))
|
||||||
.unwrap();
|
.unwrap();
|
||||||
|
|
||||||
let mut out = HashMap::new();
|
return vec![(best_idx, 1.0)];
|
||||||
out.insert(best_move.clone(), 1.0);
|
|
||||||
return out;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
let inv_temp = 1.0 / temperature;
|
let inv_temp = 1.0 / temperature;
|
||||||
|
|
||||||
// Step 1: apply exponent
|
// Step 1: apply exponent
|
||||||
let mut adjusted: HashMap<ChessMove, f32> =
|
let mut adjusted: Vec<(usize, f32)> =
|
||||||
visits.iter().map(|(m, v)| (*m, v.powf(inv_temp))).collect();
|
visits.iter().map(|(i, v)| (*i, v.powf(inv_temp))).collect();
|
||||||
|
|
||||||
// Step 2: normalize
|
// Step 2: normalize
|
||||||
let sum: f32 = adjusted.values().sum();
|
let sum: f32 = adjusted.iter().map(|(_, v)| *v).sum();
|
||||||
|
|
||||||
if sum <= 0.0 {
|
if sum <= 0.0 {
|
||||||
return adjusted; // fallback (shouldn't happen in normal MCTS)
|
return adjusted; // fallback (shouldn't happen in normal MCTS)
|
||||||
}
|
}
|
||||||
|
|
||||||
for v in adjusted.values_mut() {
|
for (_, v) in adjusted.iter_mut() {
|
||||||
*v /= sum;
|
*v /= sum;
|
||||||
}
|
}
|
||||||
|
|
||||||
adjusted
|
adjusted
|
||||||
}
|
}
|
||||||
|
|
||||||
fn sample_move(dist: &HashMap<ChessMove, f32>, rng: &mut ThreadRng) -> Option<ChessMove> {
|
fn sample_move(
|
||||||
|
dist: &[(usize, f32)],
|
||||||
|
rng: &mut SmallRng,
|
||||||
|
side_to_move: Color,
|
||||||
|
) -> Option<ChessMove> {
|
||||||
let mut r: f32 = rng.random_range(0.0..1.0);
|
let mut r: f32 = rng.random_range(0.0..1.0);
|
||||||
|
|
||||||
for (m, p) in dist {
|
for (idx, p) in dist {
|
||||||
r -= p;
|
r -= *p;
|
||||||
if r <= 0.0 {
|
if r <= 0.0 {
|
||||||
return Some(m.clone());
|
return Some(decode_move(*idx, side_to_move).expect("Invalid move"));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// fallback due to floating point drift
|
// fallback due to floating point drift
|
||||||
dist.keys().next().cloned()
|
dist.get(0)
|
||||||
|
.map(|(idx, _)| decode_move(*idx, side_to_move).expect("Invalid move"))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user