use candle_core::{Device, Result, Tensor};
use candle_nn::{linear, Linear, Module, VarBuilder, VarMap};
use super::features::{PatchKey, FEATURE_LEN};
use crate::game::{moves::Move, state::GameState};
pub const HIDDEN: usize = 64;
pub trait MovePrior {
fn biases(&self, features: &[Vec<f32>]) -> Vec<f64>;
}
pub trait FeatureSource: Send + Sync {
fn feat_dim(&self) -> usize;
fn warm_theta(&self, scale: f64) -> Vec<f64>;
fn key_and_input(&self, scratch: &GameState, mv: &Move) -> (PatchKey, Vec<f32>);
fn compute_features(&self, inputs: &[Vec<f32>]) -> Vec<Vec<f32>>;
}
pub struct PolicyNet {
l1: Linear,
l2: Linear,
l3: Linear,
}
impl PolicyNet {
pub fn new(vb: VarBuilder) -> Result<Self> {
Ok(Self {
l1: linear(FEATURE_LEN, HIDDEN, vb.pp("l1"))?,
l2: linear(HIDDEN, HIDDEN, vb.pp("l2"))?,
l3: linear(HIDDEN, 1, vb.pp("l3"))?,
})
}
pub fn forward_features(&self, x: &Tensor) -> Result<Tensor> {
let x = self.l1.forward(x)?.relu()?;
self.l2.forward(&x)?.relu() }
pub fn forward(&self, x: &Tensor) -> Result<Tensor> {
let x = self.forward_features(x)?; let x = self.l3.forward(&x)?; x.squeeze(1) }
pub fn head(&self) -> &Linear {
&self.l3
}
}
pub struct NeuralPrior {
net: PolicyNet,
#[allow(dead_code)]
varmap: VarMap,
device: Device,
}
impl NeuralPrior {
fn stack(&self, features: &[Vec<f32>]) -> Result<Tensor> {
let n = features.len();
let flat: Vec<f32> = features.iter().flat_map(|f| f.iter().copied()).collect();
Tensor::from_vec(flat, (n, FEATURE_LEN), &self.device)
}
#[allow(dead_code)]
pub fn new(net: PolicyNet, varmap: VarMap, device: Device) -> Self {
Self {
net,
varmap,
device,
}
}
pub fn save(&self, path: &str) -> Result<()> {
self.varmap.save(path)
}
pub fn load(path: &str, device: Device) -> Result<Self> {
let mut varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::F32, &device);
let net = PolicyNet::new(vb)?;
varmap.load(path)?;
Ok(Self {
net,
varmap,
device,
})
}
pub fn logits(&self, features: &[Vec<f32>]) -> Result<Vec<f32>> {
if features.is_empty() {
return Ok(Vec::new());
}
let x = self.stack(features)?;
self.net.forward(&x)?.to_vec1::<f32>()
}
fn penult(&self, features: &[Vec<f32>]) -> Result<Vec<Vec<f32>>> {
if features.is_empty() {
return Ok(Vec::new());
}
let x = self.stack(features)?;
self.net.forward_features(&x)?.to_vec2::<f32>() }
fn head_weight(&self) -> Vec<f64> {
match self
.net
.head()
.weight()
.flatten_all()
.and_then(|t| t.to_vec1::<f32>())
{
Ok(w) => w.into_iter().map(|x| x as f64).collect(),
Err(e) => {
log::error!("head_weight read failed: {e}");
Vec::new()
}
}
}
fn width(&self) -> usize {
self.net.head().weight().dims().get(1).copied().unwrap_or(0)
}
}
impl FeatureSource for NeuralPrior {
fn feat_dim(&self) -> usize {
self.width()
}
fn warm_theta(&self, scale: f64) -> Vec<f64> {
self.head_weight().into_iter().map(|w| w * scale).collect()
}
fn key_and_input(&self, scratch: &GameState, mv: &Move) -> (PatchKey, Vec<f32>) {
super::features::encode_keyed(scratch, mv)
}
fn compute_features(&self, inputs: &[Vec<f32>]) -> Vec<Vec<f32>> {
self.penult(inputs).unwrap_or_else(|e| {
log::error!("feat penult inference failed: {e}");
Vec::new()
})
}
}
impl MovePrior for NeuralPrior {
fn biases(&self, features: &[Vec<f32>]) -> Vec<f64> {
match self.logits(features) {
Ok(v) => v.into_iter().map(|x| x as f64).collect(),
Err(e) => {
log::error!("neural prior inference failed: {e}");
vec![0.0; features.len()]
}
}
}
}
pub struct ValueNet {
l1: Linear,
l2: Linear,
l3: Linear,
}
impl ValueNet {
pub fn new(vb: VarBuilder) -> Result<Self> {
Ok(Self {
l1: linear(super::position::VALUE_LEN, HIDDEN, vb.pp("v1"))?,
l2: linear(HIDDEN, HIDDEN, vb.pp("v2"))?,
l3: linear(HIDDEN, 1, vb.pp("v3"))?,
})
}
pub fn forward(&self, x: &Tensor) -> Result<Tensor> {
let x = self.l1.forward(x)?.relu()?;
let x = self.l2.forward(&x)?.relu()?;
let x = candle_nn::ops::sigmoid(&self.l3.forward(&x)?)?; x.squeeze(1)
}
}
pub struct ValuePredictor {
net: ValueNet,
#[allow(dead_code)]
varmap: VarMap,
device: Device,
}
impl ValuePredictor {
pub fn new(net: ValueNet, varmap: VarMap, device: Device) -> Self {
Self {
net,
varmap,
device,
}
}
pub fn save(&self, path: &str) -> Result<()> {
self.varmap.save(path)
}
pub fn load(path: &str, device: Device) -> Result<Self> {
let mut varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::F32, &device);
let net = ValueNet::new(vb)?;
varmap.load(path)?;
Ok(Self {
net,
varmap,
device,
})
}
pub fn value(&self, state: &crate::game::state::GameState) -> f32 {
let f = super::position::encode_value_natural(state);
let run = || -> Result<f32> {
let x = Tensor::from_vec(f, (1, super::position::VALUE_LEN), &self.device)?;
Ok(self.net.forward(&x)?.to_vec1::<f32>()?[0])
};
run().unwrap_or(0.5)
}
}
#[cfg(test)]
mod tests {
use super::*;
use candle_core::{Device, Tensor};
#[test]
fn forward_shapes() {
let dev = Device::Cpu;
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::F32, &dev);
let net = PolicyNet::new(vb).unwrap();
let n = 5usize;
let x = Tensor::arange(0f32, (n * FEATURE_LEN) as f32, &dev)
.unwrap()
.reshape((n, FEATURE_LEN))
.unwrap();
let out = net.forward(&x).unwrap().to_vec1::<f32>().unwrap();
assert_eq!(out.len(), n);
}
}