pub mod dataset;
pub mod embedded;
pub mod feat;
pub mod features;
pub mod global_features;
pub mod net;
pub mod plugin;
pub mod position;
pub mod puct;
pub mod selfplay;
pub mod tabula_rasa;
pub mod train;
use std::cell::RefCell;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, OnceLock, RwLock};
use rustc_hash::FxHashMap;
use crate::game::{moves::Move, state::GameState};
use crate::search::plugin::{registry, BiasModifier, OptionValue};
use crate::search::SearchState;
use features::PatchKey;
use net::{MovePrior, NeuralPrior, ValuePredictor};
pub const DEFAULT_SCALE: f64 = 4.0;
fn armed() -> &'static RwLock<Option<Arc<NeuralPrior>>> {
static A: OnceLock<RwLock<Option<Arc<NeuralPrior>>>> = OnceLock::new();
A.get_or_init(|| RwLock::new(None))
}
pub fn is_armed() -> bool {
armed().read().unwrap().is_some()
}
static GENERATION: AtomicU64 = AtomicU64::new(1);
fn bump_generation() {
GENERATION.fetch_add(1, Ordering::Relaxed);
}
struct BiasCache {
generation: u64,
map: FxHashMap<PatchKey, f64>,
}
thread_local! {
static BIAS_CACHE: RefCell<BiasCache> = RefCell::new(BiasCache {
generation: 0,
map: FxHashMap::default(),
});
}
pub fn set_scale(scale: f64) {
registry().set_value("neural-scale", OptionValue::Float(scale));
}
pub fn reset_scale() {
set_scale(DEFAULT_SCALE);
}
pub fn armed_arc() -> Option<Arc<NeuralPrior>> {
armed().read().unwrap().clone()
}
fn value_slot() -> &'static RwLock<Option<Arc<ValuePredictor>>> {
static V: OnceLock<RwLock<Option<Arc<ValuePredictor>>>> = OnceLock::new();
V.get_or_init(|| RwLock::new(None))
}
pub fn armed_value() -> Option<Arc<ValuePredictor>> {
value_slot().read().unwrap().clone()
}
pub fn install_value(v: Option<ValuePredictor>) {
*value_slot().write().unwrap() = v.map(Arc::new);
}
pub fn load_value(path: &str) -> candle_core::Result<ValuePredictor> {
ValuePredictor::load(path, candle_core::Device::Cpu)
}
pub fn train_value_net(
variant: crate::game::rules::Variant,
prior: Option<&NeuralPrior>,
n: usize,
epochs: usize,
) -> candle_core::Result<ValuePredictor> {
use dataset::value_samples_from_games;
use selfplay::varied_games;
use train::{train_value, TrainConfig};
let temps = [0.5f64, 1.0, 2.0];
let prior_dyn: Option<&dyn MovePrior> = prior.map(|p| p as &dyn MovePrior);
let games = varied_games(variant, n, prior_dyn, &temps);
let samples = value_samples_from_games(variant, &games, true);
let (varmap, net) = train_value(
&samples,
&TrainConfig { epochs, lr: 1e-3 },
256,
candle_core::Device::Cpu,
)?;
Ok(ValuePredictor::new(net, varmap, candle_core::Device::Cpu))
}
struct UniformPrior;
impl MovePrior for UniformPrior {
fn biases(&self, features: &[Vec<f32>]) -> Vec<f64> {
vec![0.0; features.len()]
}
}
static UNIFORM_PRIOR: UniformPrior = UniformPrior;
pub fn run_puct_armed(search: Arc<SearchState>, variant: crate::game::rules::Variant) {
let c_puct = registry().value_f64("c-puct", 1.5);
let value = armed_value();
let cfg = puct::PuctConfig {
c_puct,
leaf: if value.is_some() {
puct::LeafEval::Value
} else {
puct::LeafEval::Rollout
},
rollout_inv_temp: 1.0,
};
let vp = value.as_deref();
match armed_arc() {
Some(p) => puct::run_puct(search, variant, p.as_ref(), vp, &cfg),
None => puct::run_puct(search, variant, &UNIFORM_PRIOR, vp, &cfg),
}
}
pub struct NeuralBias;
impl BiasModifier for NeuralBias {
fn active(&self) -> bool {
is_armed()
}
fn biases(&self, state: &GameState, moves: &[Move], out: &mut Vec<f64>) {
out.clear();
let guard = armed().read().unwrap();
let Some(prior) = guard.as_ref() else {
return; };
let scale = registry().value_f64("neural-scale", DEFAULT_SCALE);
let generation = GENERATION.load(Ordering::Relaxed);
BIAS_CACHE.with(|cell| {
let mut cache = cell.borrow_mut();
if cache.generation != generation {
cache.generation = generation;
cache.map.clear();
}
let mut keys: Vec<PatchKey> = Vec::with_capacity(moves.len());
let mut miss_keys: Vec<PatchKey> = Vec::new();
let mut miss_feats: Vec<Vec<f32>> = Vec::new();
for mv in moves {
let (key, feat) = features::encode_keyed(state, mv);
if !cache.map.contains_key(&key) {
miss_keys.push(key);
miss_feats.push(feat);
}
keys.push(key);
}
if !miss_feats.is_empty() {
let logits = prior
.logits(&miss_feats)
.unwrap_or_else(|_| vec![0.0; miss_feats.len()]);
for (k, l) in miss_keys.iter().zip(logits) {
cache.map.insert(*k, l as f64);
}
}
out.extend(
keys.iter()
.map(|k| cache.map.get(k).copied().unwrap_or(0.0) * scale),
);
});
}
}
pub static NEURAL_BIAS: NeuralBias = NeuralBias;
pub mod prior {
use super::dataset::{
augmented_samples_from_corpus, augmented_samples_from_games, StateSample,
};
use super::net::NeuralPrior;
use super::train::{train, TrainConfig};
use super::{armed, Arc};
use crate::game::moves::Move;
use crate::game::rules::Variant;
use candle_core::Device;
fn train_and_freeze(
samples: &[StateSample],
epochs: usize,
lr: f64,
) -> candle_core::Result<NeuralPrior> {
let pr = train(samples, &TrainConfig { epochs, lr }, Device::Cpu)?;
let tmp =
std::env::temp_dir().join(format!("morpion_prior_{}.safetensors", std::process::id()));
let tmps = tmp.to_string_lossy().into_owned();
pr.save(&tmps)?;
let loaded = NeuralPrior::load(&tmps, Device::Cpu);
let _ = std::fs::remove_file(&tmp);
loaded
}
pub fn train_on_corpus(
variant: Variant,
epochs: usize,
lr: f64,
) -> candle_core::Result<NeuralPrior> {
train_and_freeze(&augmented_samples_from_corpus(variant), epochs, lr)
}
pub fn train_on_games(
variant: Variant,
games: &[Vec<Move>],
epochs: usize,
lr: f64,
) -> candle_core::Result<NeuralPrior> {
train_and_freeze(&augmented_samples_from_games(variant, games), epochs, lr)
}
pub fn train_on_bundled_corpus(
variant: Variant,
epochs: usize,
lr: f64,
) -> candle_core::Result<NeuralPrior> {
let games = super::embedded::corpus(variant);
if games.is_empty() {
return Err(candle_core::Error::Msg(format!(
"no bundled from-scratch corpus for {} yet",
variant.name()
)));
}
train_on_games(variant, &games, epochs, lr)
}
pub fn bundled(variant: Variant) -> Option<NeuralPrior> {
super::embedded::prior(variant)
}
pub fn load(path: &str) -> candle_core::Result<NeuralPrior> {
NeuralPrior::load(path, Device::Cpu)
}
pub fn save(prior: &NeuralPrior, path: &str) -> candle_core::Result<()> {
prior.save(path)
}
pub fn arm(prior: Option<Arc<NeuralPrior>>) {
*armed().write().unwrap() = prior;
super::bump_generation();
}
pub fn install(prior: Option<NeuralPrior>) {
arm(prior.map(Arc::new));
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::game::moves::legal_moves;
use crate::game::rules::Variant;
use std::sync::Mutex;
static ARM_TEST_LOCK: Mutex<()> = Mutex::new(());
#[test]
fn bundled_prior_biases_legal_moves() {
let _g = ARM_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let Some(p) = prior::bundled(Variant::T5) else {
panic!("bundled 5T prior should be committed");
};
assert!(!is_armed());
prior::install(Some(p));
assert!(is_armed());
let st = GameState::new(Variant::T5);
let moves = legal_moves(&st);
let mut out = Vec::new();
NEURAL_BIAS.biases(&st, &moves, &mut out);
assert_eq!(out.len(), moves.len(), "one bias per legal move");
assert!(out.iter().all(|b| b.is_finite()), "biases must be finite");
prior::arm(None);
assert!(!is_armed());
}
#[test]
fn trained_prior_round_trips() {
let mut games = Vec::new();
for _ in 0..2 {
let mut st = GameState::new(Variant::T5);
let mut h = Vec::new();
for _ in 0..12 {
let ms = legal_moves(&st);
if ms.is_empty() {
break;
}
h.push(ms[0]);
st.apply(ms[0]);
}
games.push(h);
}
let p = prior::train_on_games(Variant::T5, &games, 1, 1e-3).expect("train");
let st = GameState::new(Variant::T5);
let moves = legal_moves(&st);
let feats: Vec<Vec<f32>> = moves.iter().map(|m| features::encode(&st, m)).collect();
let b = net::MovePrior::biases(&p, &feats);
assert_eq!(b.len(), moves.len());
assert!(b.iter().all(|x| x.is_finite()));
}
#[test]
fn bias_cache_is_transparent_and_generation_aware() {
let _g = ARM_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let p = prior::bundled(Variant::T5).expect("bundled prior");
prior::install(Some(p));
let st = GameState::new(Variant::T5);
let moves = legal_moves(&st);
let mut first = Vec::new();
NEURAL_BIAS.biases(&st, &moves, &mut first); let mut second = Vec::new();
NEURAL_BIAS.biases(&st, &moves, &mut second); assert_eq!(first, second, "cache hit must match the computed result");
assert!(first.iter().all(|x| x.is_finite()));
prior::arm(None);
prior::install(prior::bundled(Variant::T5));
let mut third = Vec::new();
NEURAL_BIAS.biases(&st, &moves, &mut third);
assert_eq!(
first, third,
"a fresh generation reproduces the same logits"
);
prior::arm(None);
}
#[test]
fn feat_warm_reproduces_frozen_prior() {
let _g = ARM_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let reg = registry();
prior::install(prior::bundled(Variant::T5));
reg.set_value("feat-adapt", OptionValue::Toggle(true));
let scale = reg.value_f64("neural-scale", DEFAULT_SCALE);
feat::restart(); let st = GameState::new(Variant::T5);
let moves = legal_moves(&st);
let mut adaptive = Vec::new();
feat::logits(&st, &moves, &mut adaptive);
let p = prior::bundled(Variant::T5).unwrap();
let feats: Vec<Vec<f32>> = moves.iter().map(|m| features::encode(&st, m)).collect();
let frozen = net::MovePrior::biases(&p, &feats);
for i in 1..moves.len() {
let lhs = adaptive[i] - adaptive[0];
let rhs = scale * (frozen[i] - frozen[0]);
assert!(
(lhs - rhs).abs() < 1e-2,
"warm θ·φ should track scale·frozen: move {i} {lhs} vs {rhs}"
);
}
reg.set_value("feat-adapt", OptionValue::Toggle(false));
prior::arm(None);
}
#[test]
fn feat_adapt_updates_head_finitely() {
let _g = ARM_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
let reg = registry();
prior::install(prior::bundled(Variant::T5));
reg.set_value("feat-adapt", OptionValue::Toggle(true));
feat::restart();
let st = GameState::new(Variant::T5);
let moves = legal_moves(&st);
let mut before = Vec::new();
feat::logits(&st, &moves, &mut before);
let probs = vec![1.0 / moves.len() as f64; moves.len()];
feat::adapt(&st, &moves, &moves[0], &probs);
let mut after = Vec::new();
feat::logits(&st, &moves, &mut after);
assert!(
after.iter().all(|x| x.is_finite()),
"θ·φ must stay finite after adapt"
);
assert!(
after.iter().zip(&before).any(|(a, b)| (a - b).abs() > 1e-9),
"adapt should move at least one logit"
);
let mean_delta =
after.iter().zip(&before).map(|(a, b)| a - b).sum::<f64>() / moves.len() as f64;
assert!(
(after[0] - before[0]) >= mean_delta - 1e-9,
"the chosen move's logit should rise at least as much as the mean"
);
reg.set_value("feat-adapt", OptionValue::Toggle(false));
prior::arm(None);
}
}