use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use super::net::NeuralPrior;
use crate::game::moves::Move;
use crate::game::rules::Variant;
use crate::game::state::GameState;
use crate::search::nrpa;
use crate::search::SearchState;
#[derive(Debug, Clone)]
pub struct TabulaRasaConfig {
pub variant: Variant,
pub rounds: usize,
pub secs_per_round: f64,
pub islands: usize,
pub level: usize,
pub epochs: usize,
pub lr: f64,
pub elite: usize,
pub scale: f64,
pub scale_min: f64,
}
impl Default for TabulaRasaConfig {
fn default() -> Self {
Self {
variant: Variant::T5,
rounds: 12,
secs_per_round: 60.0,
islands: 4,
level: 3,
epochs: 30,
lr: 1e-3,
elite: 40,
scale: 4.0,
scale_min: 4.0,
}
}
}
impl TabulaRasaConfig {
pub fn smoke(variant: Variant) -> Self {
Self {
variant,
rounds: 2,
secs_per_round: 2.0,
islands: 1,
level: 1,
epochs: 1,
lr: 1e-3,
elite: 4,
scale: 4.0,
scale_min: 4.0,
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct Rung {
pub round: usize,
pub found: usize,
pub best_ever: usize,
pub scale: f64,
pub elite_max: usize,
pub elite_min: usize,
pub elite_size: usize,
}
pub fn train(
cfg: &TabulaRasaConfig,
cancel: &AtomicBool,
mut progress: impl FnMut(Rung),
) -> candle_core::Result<(Arc<NeuralPrior>, Vec<Vec<Move>>)> {
let scale_at = |round: usize| -> f64 {
if cfg.rounds <= 2 {
return cfg.scale;
}
let t = ((round.saturating_sub(1)) as f64 / (cfg.rounds - 2) as f64).clamp(0.0, 1.0);
cfg.scale + (cfg.scale_min - cfg.scale) * t
};
let islands = cfg.islands.max(1);
let per = (cfg.secs_per_round / islands as f64).max(0.5);
let mut elite: Vec<Vec<Move>> = Vec::new();
let mut current: Option<Arc<NeuralPrior>> = None; let mut last_prior: Option<Arc<NeuralPrior>> = None;
let mut best_ever = 0usize;
for round in 0..cfg.rounds {
if cancel.load(Ordering::Relaxed) {
break;
}
let scale = scale_at(round);
super::set_scale(scale);
super::prior::arm(current.clone());
let mut round_max = 0usize;
for i in 0..islands {
if cancel.load(Ordering::Relaxed) {
break;
}
let g = if round == 0 || elite.is_empty() {
generate_one(cfg.variant, cfg.level, per, cancel)
} else {
let seed = elite[i % elite.len()].clone();
refine_one(cfg.variant, cfg.level, per, seed, cancel)
};
round_max = round_max.max(g.len());
if !g.is_empty() {
elite.push(g);
}
}
super::prior::arm(None);
best_ever = best_ever.max(round_max);
let mut deduped: Vec<Vec<Move>> = Vec::with_capacity(elite.len());
for g in elite.drain(..) {
if !deduped.contains(&g) {
deduped.push(g);
}
}
deduped.sort_by_key(|g| std::cmp::Reverse(g.len()));
deduped.truncate(cfg.elite);
elite = deduped;
if elite.is_empty() {
continue; }
let prior = Arc::new(super::prior::train_on_games(
cfg.variant,
&elite,
cfg.epochs,
cfg.lr,
)?);
current = Some(prior.clone());
last_prior = Some(prior);
progress(Rung {
round,
found: round_max,
best_ever,
scale,
elite_max: elite.first().map(|g| g.len()).unwrap_or(0),
elite_min: elite.last().map(|g| g.len()).unwrap_or(0),
elite_size: elite.len(),
});
}
super::prior::arm(None);
super::reset_scale(); let prior = last_prior.ok_or_else(|| {
candle_core::Error::Msg("tabula-rasa produced no prior (cancelled before round 0)".into())
})?;
Ok((prior, elite))
}
fn generate_one(variant: Variant, level: usize, secs: f64, cancel: &AtomicBool) -> Vec<Move> {
let search = SearchState::new();
let s2 = search.clone();
let st = GameState::new(variant);
search.running.store(true, Ordering::Relaxed);
let handle = std::thread::spawn(move || nrpa::run(&st, s2, level));
let start = Instant::now();
while start.elapsed().as_secs_f64() < secs && !cancel.load(Ordering::Relaxed) {
std::thread::sleep(Duration::from_millis(50));
}
search.running.store(false, Ordering::Relaxed);
let _ = handle.join();
let g = search.best_sequence.read().unwrap().clone();
g
}
fn refine_one(
variant: Variant,
level: usize,
secs: f64,
seed: Vec<Move>,
cancel: &AtomicBool,
) -> Vec<Move> {
let search = SearchState::new();
let s2 = search.clone();
search.running.store(true, Ordering::Relaxed);
let handle = std::thread::spawn(move || nrpa::run_perturbation(s2, level, seed, variant));
let start = Instant::now();
while start.elapsed().as_secs_f64() < secs && !cancel.load(Ordering::Relaxed) {
std::thread::sleep(Duration::from_millis(50));
}
search.running.store(false, Ordering::Relaxed);
let _ = handle.join();
let g = search.best_sequence.read().unwrap().clone();
g
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
#[ignore = "tabula-rasa smoke (a few seconds), run with --features neural -- --ignored --nocapture"]
fn tabula_rasa_smoke() {
let cfg = TabulaRasaConfig::smoke(Variant::T5);
let cancel = AtomicBool::new(false);
let mut rungs = 0usize;
let (prior, corpus) = train(&cfg, &cancel, |r| {
rungs += 1;
println!(
"round {} found={} best_ever={} elite=[{}..{}]x{}",
r.round, r.found, r.best_ever, r.elite_min, r.elite_max, r.elite_size
);
})
.expect("tabula-rasa training should produce a prior");
assert!(rungs >= 1, "at least one round should complete");
assert!(!corpus.is_empty(), "the corpus (elite) should be non-empty");
let st = GameState::new(Variant::T5);
let moves = crate::game::moves::legal_moves(&st);
let feats: Vec<Vec<f32>> = moves
.iter()
.map(|m| crate::search::neural::features::encode(&st, m))
.collect();
let logits = prior.logits(&feats).expect("prior forward pass");
assert_eq!(logits.len(), moves.len());
}
}