#![allow(clippy::needless_range_loop)]
use std::collections::HashMap;
use sekirei_core::{
board::Board,
color::Color,
movegen::{generate_legal_moves, is_in_check},
nnue::{INPUT, L1, L2, NnueWeights, feature_index, hand_feature_index},
piece::PieceKind,
search::{SearchConfig, Searcher},
sfen::board_to_sfen,
tt::Tt,
};
use crate::csa::{CsaGame, GameResult};
use crate::diagnostics;
const SLOW_SEARCH_LOG_THRESHOLD: std::time::Duration = std::time::Duration::from_secs(5);
fn wdl_target_cp(result: GameResult, stm: Color, scale: f32) -> Option<f32> {
let wdl = match result {
GameResult::BlackWin => {
if stm == Color::Black {
1.0
} else {
0.0
}
}
GameResult::WhiteWin => {
if stm == Color::White {
1.0
} else {
0.0
}
}
GameResult::Draw => 0.5,
GameResult::Unknown => return None,
};
Some((wdl - 0.5) * scale)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum LrSchedule {
Constant,
StepHalf,
Cosine,
}
impl LrSchedule {
pub fn parse(s: &str) -> Option<Self> {
match s {
"constant" => Some(LrSchedule::Constant),
"step-half" => Some(LrSchedule::StepHalf),
"cosine" => Some(LrSchedule::Cosine),
_ => None,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FreezeLayer {
Ft,
L2,
Out,
}
impl FreezeLayer {
pub fn parse(s: &str) -> Option<Self> {
match s {
"ft" => Some(FreezeLayer::Ft),
"l2" => Some(FreezeLayer::L2),
"out" => Some(FreezeLayer::Out),
_ => None,
}
}
pub fn as_str(&self) -> &'static str {
match self {
FreezeLayer::Ft => "ft",
FreezeLayer::L2 => "l2",
FreezeLayer::Out => "out",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ReplayComponent {
Cp,
Wdl,
}
impl ReplayComponent {
pub fn parse(s: &str) -> Option<Self> {
match s {
"cp" => Some(ReplayComponent::Cp),
"wdl" => Some(ReplayComponent::Wdl),
_ => None,
}
}
pub fn as_str(&self) -> &'static str {
match self {
ReplayComponent::Cp => "cp",
ReplayComponent::Wdl => "wdl",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ConflictMaskLayer {
Ft,
FtAndL2,
}
impl ConflictMaskLayer {
pub fn parse(s: &str) -> Option<Self> {
match s {
"ft" => Some(ConflictMaskLayer::Ft),
"ft-l2" => Some(ConflictMaskLayer::FtAndL2),
_ => None,
}
}
pub fn as_str(&self) -> &'static str {
match self {
ConflictMaskLayer::Ft => "ft",
ConflictMaskLayer::FtAndL2 => "ft-l2",
}
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct ConflictGroupStats {
pub count: u64,
pub cp_residual_abs_sum: f64,
pub cp_residual_abs_sq_sum: f64,
pub wdl_residual_abs_sum: f64,
pub wdl_residual_abs_sq_sum: f64,
pub ft_grad_norm_sum: f64,
pub ft_grad_norm_sq_sum: f64,
pub l2_grad_norm_sum: f64,
pub l2_grad_norm_sq_sum: f64,
pub new_dead_ft_sum: u64,
pub new_dead_l2_sum: u64,
}
pub fn compute_lr(
schedule: LrSchedule,
base_lr: f32,
min_lr: f32,
epoch: u32,
total_epochs: u32,
warmup_epochs: u32,
) -> f32 {
if warmup_epochs > 0 && epoch <= warmup_epochs {
return (base_lr * epoch as f32 / warmup_epochs as f32).max(min_lr);
}
let e = epoch.saturating_sub(warmup_epochs).max(1);
let post_total = total_epochs.saturating_sub(warmup_epochs).max(1);
let lr = match schedule {
LrSchedule::Constant => base_lr,
LrSchedule::StepHalf => base_lr * 0.5_f32.powi((e - 1) as i32),
LrSchedule::Cosine => {
let denom = post_total.saturating_sub(1).max(1) as f32;
let progress = ((e - 1) as f32 / denom).min(1.0);
min_lr + 0.5 * (base_lr - min_lr) * (1.0 + (std::f32::consts::PI * progress).cos())
}
};
lr.max(min_lr)
}
pub fn resolve_schedule_epochs(
epochs: u32,
requested: Option<u32>,
warmup_epochs: u32,
) -> Result<u32, String> {
if epochs == 0 {
return Ok(requested.unwrap_or(0));
}
let schedule_epochs = requested.unwrap_or(epochs);
if schedule_epochs == 0 {
return Err("--lr-schedule-epochs must be greater than 0".to_string());
}
if warmup_epochs > schedule_epochs {
return Err(format!(
"--warmup-epochs ({warmup_epochs}) cannot exceed --lr-schedule-epochs ({schedule_epochs})"
));
}
if schedule_epochs < epochs {
return Err(format!(
"--lr-schedule-epochs ({schedule_epochs}) cannot be less than --epochs ({epochs}) -- \
use a schedule horizon at least as long as the run, or omit the flag to default it to --epochs"
));
}
Ok(schedule_epochs)
}
struct Lcg(u64);
impl Lcg {
fn next_u64(&mut self) -> u64 {
self.0 = self
.0
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
self.0
}
fn uniform(&mut self, bound: f32) -> f32 {
let u = self.next_u64() as f64 / u64::MAX as f64; ((u * 2.0 - 1.0) as f32) * bound
}
}
fn he_bound(fan_in: usize) -> f32 {
(6.0 / fan_in as f32).sqrt()
}
pub fn shuffled_order(n: usize, seed: u64) -> Vec<usize> {
let mut order: Vec<usize> = (0..n).collect();
let mut rng = Lcg(seed ^ 0xD1B5_4A32_D192_ED03);
for i in (1..n).rev() {
let j = (rng.next_u64() % (i as u64 + 1)) as usize;
order.swap(i, j);
}
order
}
#[derive(Clone)]
pub struct TrainWeights {
ft: Vec<f32>, ft_bias: Vec<f32>, l2: Vec<f32>, l2_bias: Vec<f32>, out: Vec<f32>, out_bias: f32,
ft_m: Vec<f32>,
ft_v: Vec<f32>,
bias_m: Vec<f32>,
bias_v: Vec<f32>,
l2_m: Vec<f32>,
l2_v: Vec<f32>,
l2bias_m: Vec<f32>,
l2bias_v: Vec<f32>,
out_m: Vec<f32>,
out_v: Vec<f32>,
obias_m: f32,
obias_v: f32,
step: u64,
}
impl TrainWeights {
pub fn new_seeded(seed: u64, l2_bias_init: f32) -> Self {
let ft_len = INPUT * L1;
let l2_len = 2 * L1 * L2;
let out_len = L2;
let mut rng = Lcg(seed ^ 0x9E37_79B9_7F4A_7C15);
let ft_bound = he_bound(INPUT);
let l2_bound = he_bound(2 * L1);
let out_bound = he_bound(L2);
TrainWeights {
ft: (0..ft_len).map(|_| rng.uniform(ft_bound)).collect(),
ft_bias: vec![0.5; L1],
l2: (0..l2_len).map(|_| rng.uniform(l2_bound)).collect(),
l2_bias: vec![l2_bias_init; L2],
out: (0..out_len).map(|_| rng.uniform(out_bound)).collect(),
out_bias: 0.0,
ft_m: vec![0.0; ft_len],
ft_v: vec![0.0; ft_len],
bias_m: vec![0.0; L1],
bias_v: vec![0.0; L1],
l2_m: vec![0.0; l2_len],
l2_v: vec![0.0; l2_len],
l2bias_m: vec![0.0; L2],
l2bias_v: vec![0.0; L2],
out_m: vec![0.0; out_len],
out_v: vec![0.0; out_len],
obias_m: 0.0,
obias_v: 0.0,
step: 0,
}
}
pub fn from_nnue_weights(w: &NnueWeights) -> Self {
const FT_SCALE: f32 = 64.0;
let ft_len = INPUT * L1;
let l2_len = 2 * L1 * L2;
let out_len = L2;
let mut ft = vec![0.0f32; ft_len];
for i in 0..INPUT {
for j in 0..L1 {
ft[i * L1 + j] = w.ft[i][j] as f32 / FT_SCALE;
}
}
let ft_bias: Vec<f32> = w.ft_bias.iter().map(|&v| v as f32 / FT_SCALE).collect();
let mut l2 = vec![0.0f32; l2_len];
for i in 0..2 * L1 {
for o in 0..L2 {
l2[i * L2 + o] = w.l2[i][o];
}
}
TrainWeights {
ft,
ft_bias,
l2,
l2_bias: w.l2_bias.to_vec(),
out: w.out.to_vec(),
out_bias: w.out_bias,
ft_m: vec![0.0; ft_len],
ft_v: vec![0.0; ft_len],
bias_m: vec![0.0; L1],
bias_v: vec![0.0; L1],
l2_m: vec![0.0; l2_len],
l2_v: vec![0.0; l2_len],
l2bias_m: vec![0.0; L2],
l2bias_v: vec![0.0; L2],
out_m: vec![0.0; out_len],
out_v: vec![0.0; out_len],
obias_m: 0.0,
obias_v: 0.0,
step: 0,
}
}
pub fn to_nnue_weights(&self) -> NnueWeights {
const FT_SCALE: f32 = 64.0;
let mut ft = vec![[0i16; L1]; INPUT];
for i in 0..INPUT {
for j in 0..L1 {
ft[i][j] = (self.ft[i * L1 + j] * FT_SCALE).clamp(-32767.0, 32767.0) as i16;
}
}
let mut ft_bias = [0i16; L1];
for (i, &v) in self.ft_bias.iter().enumerate() {
ft_bias[i] = (v * FT_SCALE).clamp(-32767.0, 32767.0) as i16;
}
let mut l2 = vec![[0.0f32; L2]; 2 * L1];
for i in 0..2 * L1 {
for o in 0..L2 {
l2[i][o] = self.l2[i * L2 + o];
}
}
let mut l2_bias = [0.0f32; L2];
l2_bias.copy_from_slice(&self.l2_bias);
let mut out = [0.0f32; L2];
out.copy_from_slice(&self.out);
NnueWeights {
ft,
ft_bias,
l2,
l2_bias,
out,
out_bias: self.out_bias,
}
}
pub fn snapshot_params(&self) -> Vec<f32> {
let mut v = Vec::with_capacity(
self.ft.len()
+ self.ft_bias.len()
+ self.l2.len()
+ self.l2_bias.len()
+ self.out.len()
+ 1,
);
v.extend_from_slice(&self.ft);
v.extend_from_slice(&self.ft_bias);
v.extend_from_slice(&self.l2);
v.extend_from_slice(&self.l2_bias);
v.extend_from_slice(&self.out);
v.push(self.out_bias);
v
}
pub fn l2(&self) -> &[f32] {
&self.l2
}
pub fn l2_bias(&self) -> &[f32] {
&self.l2_bias
}
pub fn out(&self) -> &[f32] {
&self.out
}
pub fn out_bias(&self) -> f32 {
self.out_bias
}
}
#[derive(Debug, Clone, Copy)]
pub struct ValidStats {
pub loss_sum: f64,
pub count: u64,
pub cp_mse_sum: f64,
pub wdl_loss_sum: f64,
pub wdl_count: u64,
pub output_sum: f64,
pub output_sum_sq: f64,
pub output_min: f32,
pub output_max: f32,
}
impl Default for ValidStats {
fn default() -> Self {
ValidStats {
loss_sum: 0.0,
count: 0,
cp_mse_sum: 0.0,
wdl_loss_sum: 0.0,
wdl_count: 0,
output_sum: 0.0,
output_sum_sq: 0.0,
output_min: f32::INFINITY,
output_max: f32::NEG_INFINITY,
}
}
}
impl std::ops::Add for ValidStats {
type Output = ValidStats;
fn add(self, other: ValidStats) -> ValidStats {
ValidStats {
loss_sum: self.loss_sum + other.loss_sum,
count: self.count + other.count,
cp_mse_sum: self.cp_mse_sum + other.cp_mse_sum,
wdl_loss_sum: self.wdl_loss_sum + other.wdl_loss_sum,
wdl_count: self.wdl_count + other.wdl_count,
output_sum: self.output_sum + other.output_sum,
output_sum_sq: self.output_sum_sq + other.output_sum_sq,
output_min: self.output_min.min(other.output_min),
output_max: self.output_max.max(other.output_max),
}
}
}
pub struct Trainer {
pub weights: TrainWeights,
pub total_loss: f64,
pub total_count: u64,
pub total_weight: f64, pub dropped_missing: u64, pub lr: f32,
pub grad_clip_norm: Option<f32>,
pub grad_clip_count: u64,
pub ft_clip_norm: Option<f32>,
pub l2_clip_norm: Option<f32>,
pub out_clip_norm: Option<f32>,
pub ft_clip_count: u64,
pub l2_clip_count: u64,
pub out_clip_count: u64,
pub out_grad_norm_values: Vec<f32>,
pub out_grad_norm_after_sum: f64,
pub out_grad_norm_after_sum_sq: f64,
pub ft_ever_active: Vec<bool>,
pub ft_ever_saturated: Vec<bool>,
pub l2_ever_active: Vec<bool>,
pub l2_ever_saturated: Vec<bool>,
pub output_sum: f64,
pub output_sum_sq: f64,
pub l2_zero_count: Vec<u64>,
pub l2_sat_count: Vec<u64>,
pub l2_sample_count: u64,
pub l2_values: Vec<Vec<f32>>,
pub ft_grad_norm_sum: f64,
pub ft_grad_norm_sum_sq: f64,
pub l2_grad_norm_sum: f64,
pub l2_grad_norm_sum_sq: f64,
pub out_grad_norm_sum: f64,
pub out_grad_norm_sum_sq: f64,
pub global_grad_norm_values: Vec<f32>,
pub ft_update_norm_sum: f64,
pub ft_update_norm_sum_sq: f64,
pub l2_update_norm_sum: f64,
pub l2_update_norm_sum_sq: f64,
pub out_update_norm_sum: f64,
pub out_update_norm_sum_sq: f64,
pub target_sum: f64,
pub target_sum_sq: f64,
pub eval_teacher_sum: f64,
pub eval_teacher_sum_sq: f64,
pub pred_eval_prod_sum: f64,
pub cp_component_sum: f64,
pub wdl_component_sum: f64,
pub wdl_component_count: u64,
pub cache_hits: u64,
pub cache_misses: u64,
pub search_time_ns: u64,
pub trace_positions: std::collections::HashSet<u64>,
pub trace_snapshots: Vec<diagnostics::TraceSnapshot>,
pub weight_snapshot_trace: bool,
pub weight_snapshots: Vec<(u64, TrainWeights)>,
pub l2_weighted_input_values: Vec<Vec<f32>>,
pub l2_dacc_sum: Vec<f64>,
pub l2_dacc_sq_sum: Vec<f64>,
pub l2_dacc_pos_count: Vec<u64>,
pub l2_dacc_neg_count: Vec<u64>,
pub ft_dacc_sum: Vec<f64>,
pub ft_dacc_sq_sum: Vec<f64>,
pub ft_dacc_pos_count: Vec<u64>,
pub ft_dacc_neg_count: Vec<u64>,
pub l2_bias_update_sq_sum: Vec<f64>,
pub ft_bias_update_sq_sum: Vec<f64>,
pub ft_zero_count: Vec<u64>,
pub ft_sat_count: Vec<u64>,
pub l2_input_norm_sum: f64,
pub l2_input_norm_sq_sum: f64,
pub ft_output_sum: f64,
pub ft_output_sum_sq: f64,
pub ft_output_count: u64,
pub cp_wdl_grad_trace: bool,
pub l2_cp_dacc_sum: Vec<f64>,
pub l2_cp_dacc_sq_sum: Vec<f64>,
pub l2_cp_dacc_pos_count: Vec<u64>,
pub l2_cp_dacc_neg_count: Vec<u64>,
pub l2_wdl_dacc_sum: Vec<f64>,
pub l2_wdl_dacc_sq_sum: Vec<f64>,
pub l2_wdl_dacc_pos_count: Vec<u64>,
pub l2_wdl_dacc_neg_count: Vec<u64>,
pub l2_cp_wdl_dot_sum: Vec<f64>,
pub ft_cp_dacc_sum: Vec<f64>,
pub ft_cp_dacc_sq_sum: Vec<f64>,
pub ft_cp_dacc_pos_count: Vec<u64>,
pub ft_cp_dacc_neg_count: Vec<u64>,
pub ft_wdl_dacc_sum: Vec<f64>,
pub ft_wdl_dacc_sq_sum: Vec<f64>,
pub ft_wdl_dacc_pos_count: Vec<u64>,
pub ft_wdl_dacc_neg_count: Vec<u64>,
pub ft_cp_wdl_dot_sum: Vec<f64>,
pub cp_ft_grad_norm_sum: f64,
pub cp_ft_grad_norm_sum_sq: f64,
pub wdl_ft_grad_norm_sum: f64,
pub wdl_ft_grad_norm_sum_sq: f64,
pub cp_l2_grad_norm_sum: f64,
pub cp_l2_grad_norm_sum_sq: f64,
pub wdl_l2_grad_norm_sum: f64,
pub wdl_l2_grad_norm_sum_sq: f64,
pub cp_out_grad_norm_sum: f64,
pub cp_out_grad_norm_sum_sq: f64,
pub wdl_out_grad_norm_sum: f64,
pub wdl_out_grad_norm_sum_sq: f64,
pub cp_target_sum: f64,
pub cp_target_sum_sq: f64,
pub wdl_target_sum: f64,
pub wdl_target_sum_sq: f64,
pub prediction_sum: f64,
pub prediction_sum_sq: f64,
pub cp_residual_sum: f64,
pub cp_residual_sum_sq: f64,
pub wdl_residual_sum: f64,
pub wdl_residual_sum_sq: f64,
pub cp_d_output_sum: f64,
pub cp_d_output_sum_sq: f64,
pub wdl_d_output_sum: f64,
pub wdl_d_output_sum_sq: f64,
pub sample_grad_trace_limit: u64,
pub sample_grad_records: Vec<diagnostics::SampleGradRecord>,
sample_grad_prev_d_l2_acc: Option<[f32; L2]>,
sample_grad_running_mean_d_l2_acc: [f32; L2],
sample_grad_running_count: u64,
pub diagnostic_freeze_layer: Option<FreezeLayer>,
pub diagnostic_freeze_from_position: u64,
pub diagnostic_freeze_until_position: u64,
pub diagnostic_ft_active_block: u64,
pub diagnostic_ft_frozen_block: u64,
pub diagnostic_ft_frozen_first: bool,
pub diagnostic_ft_reactivate_from_position: u64,
pub diagnostic_ft_reactivate_until_position: u64,
pub diagnostic_ft_reactivate2_from_position: u64,
pub diagnostic_ft_reactivate2_until_position: u64,
pub diagnostic_replay_component: Option<ReplayComponent>,
pub diagnostic_replay_from_position: u64,
pub diagnostic_replay_until_position: u64,
pub diagnostic_shadow_trace_from_position: u64,
pub diagnostic_shadow_trace_until_position: u64,
pub diagnostic_shadow_trace_wdl_lambda: f32,
pub diagnostic_shadow_trace_probe_boards: Vec<Board>,
pub shadow_trace_records: Vec<diagnostics::ShadowTraceRecord>,
pub diagnostic_conflict_mask: Option<ConflictMaskLayer>,
pub diagnostic_rate_matched_mask_count: u64,
pub diagnostic_rate_matched_mask_total: u64,
pub diagnostic_rate_matched_mask_seed: u64,
rate_matched_remaining_needed: u64,
rate_matched_remaining_pool: u64,
rate_matched_rng: Lcg,
pub masked_position_count: u64,
pub conflict_group: ConflictGroupStats,
pub nonconflict_group: ConflictGroupStats,
pending_conflict_dead_before: Option<(bool, u64, u64)>,
searcher: Searcher,
}
impl Trainer {
pub fn new(seed: u64, l2_bias_init: f32) -> Self {
let tt = Tt::new(4); Trainer {
weights: TrainWeights::new_seeded(seed, l2_bias_init),
total_loss: 0.0,
total_count: 0,
total_weight: 0.0,
dropped_missing: 0,
lr: 0.001,
grad_clip_norm: None,
grad_clip_count: 0,
ft_clip_norm: None,
l2_clip_norm: None,
out_clip_norm: None,
ft_clip_count: 0,
l2_clip_count: 0,
out_clip_count: 0,
out_grad_norm_values: Vec::new(),
out_grad_norm_after_sum: 0.0,
out_grad_norm_after_sum_sq: 0.0,
ft_ever_active: vec![false; L1],
ft_ever_saturated: vec![false; L1],
l2_ever_active: vec![false; L2],
l2_ever_saturated: vec![false; L2],
output_sum: 0.0,
output_sum_sq: 0.0,
l2_zero_count: vec![0; L2],
l2_sat_count: vec![0; L2],
l2_sample_count: 0,
l2_values: vec![Vec::new(); L2],
ft_grad_norm_sum: 0.0,
ft_grad_norm_sum_sq: 0.0,
l2_grad_norm_sum: 0.0,
l2_grad_norm_sum_sq: 0.0,
out_grad_norm_sum: 0.0,
out_grad_norm_sum_sq: 0.0,
global_grad_norm_values: Vec::new(),
ft_update_norm_sum: 0.0,
ft_update_norm_sum_sq: 0.0,
l2_update_norm_sum: 0.0,
l2_update_norm_sum_sq: 0.0,
out_update_norm_sum: 0.0,
out_update_norm_sum_sq: 0.0,
target_sum: 0.0,
target_sum_sq: 0.0,
eval_teacher_sum: 0.0,
eval_teacher_sum_sq: 0.0,
pred_eval_prod_sum: 0.0,
cp_component_sum: 0.0,
wdl_component_sum: 0.0,
wdl_component_count: 0,
cache_hits: 0,
cache_misses: 0,
search_time_ns: 0,
trace_positions: std::collections::HashSet::new(),
trace_snapshots: Vec::new(),
weight_snapshot_trace: false,
weight_snapshots: Vec::new(),
l2_weighted_input_values: vec![Vec::new(); L2],
l2_dacc_sum: vec![0.0; L2],
l2_dacc_sq_sum: vec![0.0; L2],
l2_dacc_pos_count: vec![0; L2],
l2_dacc_neg_count: vec![0; L2],
ft_dacc_sum: vec![0.0; L1],
ft_dacc_sq_sum: vec![0.0; L1],
ft_dacc_pos_count: vec![0; L1],
ft_dacc_neg_count: vec![0; L1],
l2_bias_update_sq_sum: vec![0.0; L2],
ft_bias_update_sq_sum: vec![0.0; L1],
ft_zero_count: vec![0; L1],
ft_sat_count: vec![0; L1],
l2_input_norm_sum: 0.0,
l2_input_norm_sq_sum: 0.0,
ft_output_sum: 0.0,
ft_output_sum_sq: 0.0,
ft_output_count: 0,
cp_wdl_grad_trace: false,
l2_cp_dacc_sum: vec![0.0; L2],
l2_cp_dacc_sq_sum: vec![0.0; L2],
l2_cp_dacc_pos_count: vec![0; L2],
l2_cp_dacc_neg_count: vec![0; L2],
l2_wdl_dacc_sum: vec![0.0; L2],
l2_wdl_dacc_sq_sum: vec![0.0; L2],
l2_wdl_dacc_pos_count: vec![0; L2],
l2_wdl_dacc_neg_count: vec![0; L2],
l2_cp_wdl_dot_sum: vec![0.0; L2],
ft_cp_dacc_sum: vec![0.0; L1],
ft_cp_dacc_sq_sum: vec![0.0; L1],
ft_cp_dacc_pos_count: vec![0; L1],
ft_cp_dacc_neg_count: vec![0; L1],
ft_wdl_dacc_sum: vec![0.0; L1],
ft_wdl_dacc_sq_sum: vec![0.0; L1],
ft_wdl_dacc_pos_count: vec![0; L1],
ft_wdl_dacc_neg_count: vec![0; L1],
ft_cp_wdl_dot_sum: vec![0.0; L1],
cp_ft_grad_norm_sum: 0.0,
cp_ft_grad_norm_sum_sq: 0.0,
wdl_ft_grad_norm_sum: 0.0,
wdl_ft_grad_norm_sum_sq: 0.0,
cp_l2_grad_norm_sum: 0.0,
cp_l2_grad_norm_sum_sq: 0.0,
wdl_l2_grad_norm_sum: 0.0,
wdl_l2_grad_norm_sum_sq: 0.0,
cp_out_grad_norm_sum: 0.0,
cp_out_grad_norm_sum_sq: 0.0,
wdl_out_grad_norm_sum: 0.0,
wdl_out_grad_norm_sum_sq: 0.0,
cp_target_sum: 0.0,
cp_target_sum_sq: 0.0,
wdl_target_sum: 0.0,
wdl_target_sum_sq: 0.0,
prediction_sum: 0.0,
prediction_sum_sq: 0.0,
cp_residual_sum: 0.0,
cp_residual_sum_sq: 0.0,
wdl_residual_sum: 0.0,
wdl_residual_sum_sq: 0.0,
cp_d_output_sum: 0.0,
cp_d_output_sum_sq: 0.0,
wdl_d_output_sum: 0.0,
wdl_d_output_sum_sq: 0.0,
sample_grad_trace_limit: 0,
sample_grad_records: Vec::new(),
sample_grad_prev_d_l2_acc: None,
sample_grad_running_mean_d_l2_acc: [0.0; L2],
sample_grad_running_count: 0,
diagnostic_freeze_layer: None,
diagnostic_freeze_from_position: 0,
diagnostic_freeze_until_position: 0,
diagnostic_ft_active_block: 0,
diagnostic_ft_frozen_block: 0,
diagnostic_ft_frozen_first: false,
diagnostic_ft_reactivate_from_position: 0,
diagnostic_ft_reactivate_until_position: 0,
diagnostic_ft_reactivate2_from_position: 0,
diagnostic_ft_reactivate2_until_position: 0,
diagnostic_replay_component: None,
diagnostic_replay_from_position: 0,
diagnostic_replay_until_position: 0,
diagnostic_shadow_trace_from_position: 0,
diagnostic_shadow_trace_until_position: 0,
diagnostic_shadow_trace_wdl_lambda: 0.0,
diagnostic_shadow_trace_probe_boards: Vec::new(),
shadow_trace_records: Vec::new(),
diagnostic_conflict_mask: None,
diagnostic_rate_matched_mask_count: 0,
diagnostic_rate_matched_mask_total: 0,
diagnostic_rate_matched_mask_seed: 0,
rate_matched_remaining_needed: 0,
rate_matched_remaining_pool: 0,
rate_matched_rng: Lcg(0),
masked_position_count: 0,
conflict_group: ConflictGroupStats::default(),
nonconflict_group: ConflictGroupStats::default(),
pending_conflict_dead_before: None,
searcher: Searcher::new(tt),
}
}
#[allow(clippy::too_many_arguments)]
pub fn train_positions(
&mut self,
samples: &[crate::positions::PositionSample],
label_depth: u32,
scored: &HashMap<String, f32>,
stability_weighted: bool,
phase_weights: &HashMap<String, f32>,
side_weights: &HashMap<String, f32>,
teacher_cache: &HashMap<String, i32>,
new_entries: &mut Vec<(String, i32)>,
) {
for sample in samples {
let sfen = sekirei_core::sfen::board_to_sfen(&sample.board);
let stability = if scored.is_empty() {
1.0f32
} else {
match scored.get(&sfen) {
Some(&s) => {
if stability_weighted {
s
} else {
1.0
}
}
None => {
self.dropped_missing += 1;
continue;
}
}
};
let phase_w = phase_weights.get(&sample.phase).copied().unwrap_or(1.0);
let side_w = side_weights
.get(&sample.side_to_move)
.copied()
.unwrap_or(1.0);
let weight = stability * phase_w * side_w;
let score_cp = if let Some(&cp) = teacher_cache.get(&sfen) {
cp
} else {
let config = SearchConfig {
max_depth: label_depth,
time_limit: None,
soft_limit: None,
multi_pv: 1,
};
let mut b = sample.board.clone();
let cp = self.searcher.search(&mut b, config).score;
new_entries.push((sfen, cp));
cp
};
let teacher = (score_cp as f32).clamp(-600.0, 600.0);
self.train_position(
&sample.board,
teacher,
weight,
teacher,
None,
0,
GameResult::Unknown,
);
}
}
pub fn eval_positions(
&mut self,
samples: &[crate::positions::PositionSample],
label_depth: u32,
phase_weights: &HashMap<String, f32>,
side_weights: &HashMap<String, f32>,
teacher_cache: &HashMap<String, i32>,
new_entries: &mut Vec<(String, i32)>,
) -> (f64, f64, u64) {
let mut loss_raw = 0.0f64;
let mut loss_weighted = 0.0f64;
let mut total_w = 0.0f64;
let mut count = 0u64;
for sample in samples {
let sfen = sekirei_core::sfen::board_to_sfen(&sample.board);
let teacher_cp = if let Some(&cp) = teacher_cache.get(&sfen) {
cp
} else {
let config = SearchConfig {
max_depth: label_depth,
time_limit: None,
soft_limit: None,
multi_pv: 1,
};
let mut b = sample.board.clone();
let cp = self.searcher.search(&mut b, config).score;
new_entries.push((sfen, cp));
cp
};
let teacher = (teacher_cp as f32).clamp(-600.0, 600.0);
let score = self.forward(&sample.board);
let err2 = ((score - teacher) * (score - teacher)) as f64;
loss_raw += err2;
let w = phase_weights.get(&sample.phase).copied().unwrap_or(1.0)
* side_weights
.get(&sample.side_to_move)
.copied()
.unwrap_or(1.0);
loss_weighted += w as f64 * err2;
total_w += w as f64;
count += 1;
}
let raw = if count > 0 {
loss_raw / count as f64
} else {
0.0
};
let weighted = if total_w > 0.0 {
loss_weighted / total_w
} else {
0.0
};
(raw, weighted, count)
}
fn position_teacher_components(
&mut self,
board: &mut Board,
result: GameResult,
label_depth: u32,
cache: &mut HashMap<String, i32>,
wdl_target_scale: f32,
) -> (f32, Option<f32>) {
let sfen = board_to_sfen(board);
let score_cp = if let Some(&cp) = cache.get(&sfen) {
self.cache_hits += 1;
cp
} else {
self.cache_misses += 1;
let config = SearchConfig {
max_depth: label_depth,
time_limit: None,
soft_limit: None,
multi_pv: 1,
};
let search_start = std::time::Instant::now();
let info = self.searcher.search(board, config);
let search_elapsed = search_start.elapsed();
self.search_time_ns += search_elapsed.as_nanos() as u64;
if search_elapsed >= SLOW_SEARCH_LOG_THRESHOLD {
let legal_move_count = generate_legal_moves(board).len();
eprintln!(
" slow search: {:.1}s depth={} nodes={} stm={:?} legal_moves={legal_move_count} sfen={sfen}",
search_elapsed.as_secs_f64(),
info.depth,
info.nodes,
board.side_to_move,
);
}
let cp = info.score;
cache.insert(sfen, cp);
cp
};
let eval_teacher = (score_cp as f32).clamp(-600.0, 600.0);
(
eval_teacher,
wdl_target_cp(result, board.side_to_move, wdl_target_scale),
)
}
fn replay_override(
&self,
position: u64,
wdl_lambda: Option<f32>,
eval_teacher: f32,
wdl_target: Option<f32>,
weight: f32,
default_teacher: f32,
) -> (f32, f32) {
let (Some(component), Some(lambda), Some(wdl_target)) =
(self.diagnostic_replay_component, wdl_lambda, wdl_target)
else {
return (default_teacher, weight);
};
if position < self.diagnostic_replay_from_position
|| position > self.diagnostic_replay_until_position
{
return (default_teacher, weight);
}
match component {
ReplayComponent::Cp => (eval_teacher, weight * lambda),
ReplayComponent::Wdl => (wdl_target, weight * (1.0 - lambda)),
}
}
#[allow(clippy::too_many_arguments)]
pub fn train_game(
&mut self,
game_id: u64,
game: &CsaGame,
sample_every: usize,
quiet: bool,
min_ply: usize,
label_depth: u32,
scored: &HashMap<String, f32>,
stability_weighted: bool,
wdl_lambda: Option<f32>,
wdl_target_scale: f32,
cache: &mut HashMap<String, i32>,
) {
let mut board = Board::startpos();
for (ply, &mv) in game.moves.iter().enumerate() {
if ply < min_ply || ply % sample_every != 0 {
board.do_move(mv);
continue;
}
if quiet {
if is_in_check(&board, board.side_to_move) {
board.do_move(mv);
continue;
}
if board.piece_at(mv.to).is_some() {
board.do_move(mv);
continue;
}
}
let weight = if scored.is_empty() {
1.0f32
} else {
let sfen = board_to_sfen(&board);
match scored.get(&sfen) {
Some(&s) => {
if stability_weighted {
s
} else {
1.0
}
}
None => {
self.dropped_missing += 1;
board.do_move(mv);
continue; }
}
};
let (eval_teacher, wdl_target) = self.position_teacher_components(
&mut board,
game.result,
label_depth,
cache,
wdl_target_scale,
);
let teacher = match (wdl_lambda, wdl_target) {
(Some(lambda), Some(wdl_target)) => {
lambda * eval_teacher + (1.0 - lambda) * wdl_target
}
_ => eval_teacher,
};
let (teacher, weight) = self.replay_override(
self.l2_sample_count + 1,
wdl_lambda,
eval_teacher,
wdl_target,
weight,
teacher,
);
self.train_position(
&board,
teacher,
weight,
eval_teacher,
wdl_target,
game_id,
game.result,
);
board.do_move(mv);
}
}
#[allow(clippy::too_many_arguments)]
#[allow(clippy::too_many_arguments)]
pub fn eval_game(
&mut self,
game: &CsaGame,
sample_every: usize,
quiet: bool,
min_ply: usize,
label_depth: u32,
wdl_lambda: Option<f32>,
wdl_target_scale: f32,
cache: &mut HashMap<String, i32>,
) -> ValidStats {
let mut board = Board::startpos();
let mut stats = ValidStats::default();
for (ply, &mv) in game.moves.iter().enumerate() {
if ply < min_ply || ply % sample_every != 0 {
board.do_move(mv);
continue;
}
if quiet {
if is_in_check(&board, board.side_to_move) {
board.do_move(mv);
continue;
}
if board.piece_at(mv.to).is_some() {
board.do_move(mv);
continue;
}
}
let (eval_teacher, wdl_target) = self.position_teacher_components(
&mut board,
game.result,
label_depth,
cache,
wdl_target_scale,
);
let teacher = match (wdl_lambda, wdl_target) {
(Some(lambda), Some(wdl_target)) => {
lambda * eval_teacher + (1.0 - lambda) * wdl_target
}
_ => eval_teacher,
};
let score = self.forward(&board);
let err = (score - teacher) as f64;
stats.loss_sum += err * err;
stats.count += 1;
let cp_err = (score - eval_teacher) as f64;
stats.cp_mse_sum += cp_err * cp_err;
if let Some(wdl_target) = wdl_target {
let wdl_err = (score - wdl_target) as f64;
stats.wdl_loss_sum += wdl_err * wdl_err;
stats.wdl_count += 1;
}
stats.output_sum += score as f64;
stats.output_sum_sq += (score as f64) * (score as f64);
stats.output_min = stats.output_min.min(score);
stats.output_max = stats.output_max.max(score);
board.do_move(mv);
}
stats
}
fn forward(&self, board: &Board) -> f32 {
let stm = board.side_to_move;
let w = &self.weights;
let mut acc_us = w.ft_bias.clone();
let mut acc_them = acc_us.clone();
for feat in &active_features(board, stm) {
let base = feat * L1;
for j in 0..L1 {
acc_us[j] += w.ft[base + j];
}
}
for feat in &active_features(board, stm.flip()) {
let base = feat * L1;
for j in 0..L1 {
acc_them[j] += w.ft[base + j];
}
}
let relu_us: Vec<f32> = acc_us.iter().map(|&x| x.clamp(0.0, 127.0)).collect();
let relu_them: Vec<f32> = acc_them.iter().map(|&x| x.clamp(0.0, 127.0)).collect();
let mut l2_acc = w.l2_bias.clone();
for j in 0..L1 {
let base_us = j * L2;
let base_them = (L1 + j) * L2;
for o in 0..L2 {
l2_acc[o] += relu_us[j] * w.l2[base_us + o];
l2_acc[o] += relu_them[j] * w.l2[base_them + o];
}
}
let relu_l2: Vec<f32> = l2_acc.iter().map(|&x| x.clamp(0.0, 127.0)).collect();
let mut output = w.out_bias;
for o in 0..L2 {
output += relu_l2[o] * w.out[o];
}
output / 64.0
}
fn ft_periodic_active_phase(&self) -> bool {
if self.diagnostic_ft_active_block == 0 || self.diagnostic_ft_frozen_block == 0 {
return false;
}
let cycle = self.diagnostic_ft_active_block + self.diagnostic_ft_frozen_block;
let offset = self.l2_sample_count - self.diagnostic_freeze_from_position;
let phase = offset % cycle;
if self.diagnostic_ft_frozen_first {
phase >= self.diagnostic_ft_frozen_block
} else {
phase < self.diagnostic_ft_active_block
}
}
fn ft_reactivated(&self) -> bool {
let window1 = self.l2_sample_count >= self.diagnostic_ft_reactivate_from_position
&& self.l2_sample_count <= self.diagnostic_ft_reactivate_until_position;
let window2 = self.l2_sample_count >= self.diagnostic_ft_reactivate2_from_position
&& self.l2_sample_count <= self.diagnostic_ft_reactivate2_until_position;
window1 || window2
}
#[allow(clippy::too_many_arguments)]
fn train_position(
&mut self,
board: &Board,
teacher: f32,
weight: f32,
eval_teacher: f32,
wdl_target: Option<f32>,
game_id: u64,
game_result: GameResult,
) {
let stm = board.side_to_move;
let w = &self.weights;
let mut acc_us = w.ft_bias.clone();
let mut acc_them = acc_us.clone();
let active_us = active_features(board, stm);
let active_them = active_features(board, stm.flip());
for feat in &active_us {
let base = feat * L1;
for j in 0..L1 {
acc_us[j] += w.ft[base + j];
}
}
for feat in &active_them {
let base = feat * L1;
for j in 0..L1 {
acc_them[j] += w.ft[base + j];
}
}
let relu_us: Vec<f32> = acc_us.iter().map(|&x| x.clamp(0.0, 127.0)).collect();
let relu_them: Vec<f32> = acc_them.iter().map(|&x| x.clamp(0.0, 127.0)).collect();
for &x in relu_us.iter().chain(relu_them.iter()) {
self.ft_output_sum += x as f64;
self.ft_output_sum_sq += (x as f64) * (x as f64);
}
self.ft_output_count += 2 * L1 as u64;
for j in 0..L1 {
if relu_us[j] > 0.0 || relu_them[j] > 0.0 {
self.ft_ever_active[j] = true;
}
if relu_us[j] >= 127.0 || relu_them[j] >= 127.0 {
self.ft_ever_saturated[j] = true;
}
if acc_us[j] <= 0.0 && acc_them[j] <= 0.0 {
self.ft_zero_count[j] += 1;
}
if acc_us[j] >= 127.0 || acc_them[j] >= 127.0 {
self.ft_sat_count[j] += 1;
}
}
let mut l2_acc = w.l2_bias.clone(); for j in 0..L1 {
let a = relu_us[j];
let b = relu_them[j];
let base_us = j * L2;
let base_them = (L1 + j) * L2;
for o in 0..L2 {
l2_acc[o] += a * w.l2[base_us + o];
l2_acc[o] += b * w.l2[base_them + o];
}
}
let relu_l2: Vec<f32> = l2_acc.iter().map(|&x| x.clamp(0.0, 127.0)).collect();
for o in 0..L2 {
if relu_l2[o] > 0.0 {
self.l2_ever_active[o] = true;
}
if relu_l2[o] >= 127.0 {
self.l2_ever_saturated[o] = true;
}
let pre = l2_acc[o];
if pre <= 0.0 {
self.l2_zero_count[o] += 1;
}
if pre >= 127.0 {
self.l2_sat_count[o] += 1;
}
self.l2_values[o].push(pre);
self.l2_weighted_input_values[o].push(pre - w.l2_bias[o]);
}
self.l2_sample_count += 1;
let l2_input_norm_sq: f64 = relu_us
.iter()
.chain(relu_them.iter())
.map(|&x| (x as f64).powi(2))
.sum();
self.l2_input_norm_sum += l2_input_norm_sq.sqrt();
self.l2_input_norm_sq_sum += l2_input_norm_sq;
let mut output = w.out_bias;
for o in 0..L2 {
output += relu_l2[o] * w.out[o];
}
let score = output / 64.0;
self.output_sum += score as f64;
self.output_sum_sq += (score as f64) * (score as f64);
let err = score - teacher;
self.total_loss += (weight as f64) * (err * err) as f64;
self.total_count += 1;
self.total_weight += weight as f64;
self.target_sum += teacher as f64;
self.target_sum_sq += (teacher as f64) * (teacher as f64);
self.eval_teacher_sum += eval_teacher as f64;
self.eval_teacher_sum_sq += (eval_teacher as f64) * (eval_teacher as f64);
self.pred_eval_prod_sum += (score as f64) * (eval_teacher as f64);
let cp_err = (score - eval_teacher) as f64;
self.cp_component_sum += cp_err * cp_err;
if let Some(wdl_target) = wdl_target {
let wdl_err = (score - wdl_target) as f64;
self.wdl_component_sum += wdl_err * wdl_err;
self.wdl_component_count += 1;
if self.cp_wdl_grad_trace {
let cp = diagnostic_backward(
&self.weights,
&l2_acc,
&relu_l2,
&acc_us,
&acc_them,
&relu_us,
&relu_them,
&active_us,
&active_them,
score - eval_teacher,
weight,
);
let wdl = diagnostic_backward(
&self.weights,
&l2_acc,
&relu_l2,
&acc_us,
&acc_them,
&relu_us,
&relu_them,
&active_us,
&active_them,
score - wdl_target,
weight,
);
for o in 0..L2 {
let gc = cp.d_l2_acc[o] as f64;
let gw = wdl.d_l2_acc[o] as f64;
self.l2_cp_dacc_sum[o] += gc;
self.l2_cp_dacc_sq_sum[o] += gc * gc;
self.l2_wdl_dacc_sum[o] += gw;
self.l2_wdl_dacc_sq_sum[o] += gw * gw;
self.l2_cp_wdl_dot_sum[o] += gc * gw;
if gc > 0.0 {
self.l2_cp_dacc_pos_count[o] += 1;
} else if gc < 0.0 {
self.l2_cp_dacc_neg_count[o] += 1;
}
if gw > 0.0 {
self.l2_wdl_dacc_pos_count[o] += 1;
} else if gw < 0.0 {
self.l2_wdl_dacc_neg_count[o] += 1;
}
}
for j in 0..L1 {
let gc = cp.d_ft_acc[j] as f64;
let gw = wdl.d_ft_acc[j] as f64;
self.ft_cp_dacc_sum[j] += gc;
self.ft_cp_dacc_sq_sum[j] += gc * gc;
self.ft_wdl_dacc_sum[j] += gw;
self.ft_wdl_dacc_sq_sum[j] += gw * gw;
self.ft_cp_wdl_dot_sum[j] += gc * gw;
if gc > 0.0 {
self.ft_cp_dacc_pos_count[j] += 1;
} else if gc < 0.0 {
self.ft_cp_dacc_neg_count[j] += 1;
}
if gw > 0.0 {
self.ft_wdl_dacc_pos_count[j] += 1;
} else if gw < 0.0 {
self.ft_wdl_dacc_neg_count[j] += 1;
}
}
self.cp_ft_grad_norm_sum += cp.ft_grad_norm;
self.cp_ft_grad_norm_sum_sq += cp.ft_grad_norm * cp.ft_grad_norm;
self.wdl_ft_grad_norm_sum += wdl.ft_grad_norm;
self.wdl_ft_grad_norm_sum_sq += wdl.ft_grad_norm * wdl.ft_grad_norm;
self.cp_l2_grad_norm_sum += cp.l2_grad_norm;
self.cp_l2_grad_norm_sum_sq += cp.l2_grad_norm * cp.l2_grad_norm;
self.wdl_l2_grad_norm_sum += wdl.l2_grad_norm;
self.wdl_l2_grad_norm_sum_sq += wdl.l2_grad_norm * wdl.l2_grad_norm;
self.cp_out_grad_norm_sum += cp.out_grad_norm;
self.cp_out_grad_norm_sum_sq += cp.out_grad_norm * cp.out_grad_norm;
self.wdl_out_grad_norm_sum += wdl.out_grad_norm;
self.wdl_out_grad_norm_sum_sq += wdl.out_grad_norm * wdl.out_grad_norm;
let eval_teacher_f64 = eval_teacher as f64;
let wdl_target_f64 = wdl_target as f64;
let score_f64 = score as f64;
self.cp_target_sum += eval_teacher_f64;
self.cp_target_sum_sq += eval_teacher_f64 * eval_teacher_f64;
self.wdl_target_sum += wdl_target_f64;
self.wdl_target_sum_sq += wdl_target_f64 * wdl_target_f64;
self.prediction_sum += score_f64;
self.prediction_sum_sq += score_f64 * score_f64;
self.cp_residual_sum += cp_err;
self.cp_residual_sum_sq += cp_err * cp_err;
self.wdl_residual_sum += wdl_err;
self.wdl_residual_sum_sq += wdl_err * wdl_err;
let cp_d_output = cp.d_output as f64;
let wdl_d_output = wdl.d_output as f64;
self.cp_d_output_sum += cp_d_output;
self.cp_d_output_sum_sq += cp_d_output * cp_d_output;
self.wdl_d_output_sum += wdl_d_output;
self.wdl_d_output_sum_sq += wdl_d_output * wdl_d_output;
}
}
let d_score = weight * 2.0 * err;
let d_output = d_score / 64.0;
let mut d_out = vec![0.0f32; L2];
for o in 0..L2 {
d_out[o] = d_output * relu_l2[o];
}
let mut d_out_bias = d_output;
let mut d_l2_acc = [0.0f32; L2];
for o in 0..L2 {
if l2_acc[o] > 0.0 && l2_acc[o] < 127.0 {
d_l2_acc[o] = d_output * self.weights.out[o];
}
}
for o in 0..L2 {
let g = d_l2_acc[o] as f64;
self.l2_dacc_sum[o] += g;
self.l2_dacc_sq_sum[o] += g * g;
if g > 0.0 {
self.l2_dacc_pos_count[o] += 1;
} else if g < 0.0 {
self.l2_dacc_neg_count[o] += 1;
}
}
if self.sample_grad_trace_limit > 0 && self.l2_sample_count <= self.sample_grad_trace_limit
{
let cp_d_output = weight * 2.0 * (score - eval_teacher) / 64.0;
let wdl_d_output = wdl_target.map(|t| weight * 2.0 * (score - t) / 64.0);
let l2_grad_norm = (d_l2_acc.iter().map(|&x| (x as f64).powi(2)).sum::<f64>()).sqrt();
let cosine_prev = self
.sample_grad_prev_d_l2_acc
.as_ref()
.map(|prev| diagnostics::vector_cosine_similarity(prev, &d_l2_acc));
let cosine_running_mean = if self.sample_grad_running_count > 0 {
Some(diagnostics::vector_cosine_similarity(
&self.sample_grad_running_mean_d_l2_acc,
&d_l2_acc,
))
} else {
None
};
let l2_gate: Vec<i8> = l2_acc
.iter()
.map(|&x| {
if x <= 0.0 {
-1
} else if x >= 127.0 {
1
} else {
0
}
})
.collect();
self.sample_grad_records
.push(diagnostics::SampleGradRecord {
game_id,
game_result: format!("{game_result:?}"),
position_index: self.l2_sample_count,
prediction: score,
cp_target: eval_teacher,
wdl_target,
cp_d_output,
wdl_d_output,
l2_grad_vector: d_l2_acc.to_vec(),
l2_grad_norm,
cosine_prev,
cosine_running_mean,
l2_gate,
});
self.sample_grad_prev_d_l2_acc = Some(d_l2_acc);
self.sample_grad_running_count += 1;
let n = self.sample_grad_running_count as f32;
for o in 0..L2 {
self.sample_grad_running_mean_d_l2_acc[o] +=
(d_l2_acc[o] - self.sample_grad_running_mean_d_l2_acc[o]) / n;
}
}
let mut d_l2 = vec![0.0f32; 2 * L1 * L2];
let mut d_l2_bias = vec![0.0f32; L2];
let mut d_relu_us = vec![0.0f32; L1];
let mut d_relu_them = vec![0.0f32; L1];
for j in 0..L1 {
let base_us = j * L2;
let base_them = (L1 + j) * L2;
for o in 0..L2 {
let g = d_l2_acc[o];
d_l2[base_us + o] += g * relu_us[j];
d_l2[base_them + o] += g * relu_them[j];
d_relu_us[j] += g * self.weights.l2[base_us + o];
d_relu_them[j] += g * self.weights.l2[base_them + o];
}
}
d_l2_bias[..L2].copy_from_slice(&d_l2_acc[..L2]);
let mut d_acc_us = vec![0.0f32; L1];
let mut d_acc_them = vec![0.0f32; L1];
for j in 0..L1 {
if acc_us[j] > 0.0 && acc_us[j] < 127.0 {
d_acc_us[j] = d_relu_us[j];
}
if acc_them[j] > 0.0 && acc_them[j] < 127.0 {
d_acc_them[j] = d_relu_them[j];
}
}
let mut d_ft = vec![0.0f32; INPUT * L1];
let mut d_bias = vec![0.0f32; L1];
for feat in &active_us {
let base = feat * L1;
for j in 0..L1 {
d_ft[base + j] += d_acc_us[j];
}
}
for feat in &active_them {
let base = feat * L1;
for j in 0..L1 {
d_ft[base + j] += d_acc_them[j];
}
}
for j in 0..L1 {
d_bias[j] = d_acc_us[j] + d_acc_them[j];
}
for j in 0..L1 {
let g = d_bias[j] as f64;
self.ft_dacc_sum[j] += g;
self.ft_dacc_sq_sum[j] += g * g;
if g > 0.0 {
self.ft_dacc_pos_count[j] += 1;
} else if g < 0.0 {
self.ft_dacc_neg_count[j] += 1;
}
}
let d_acc_us_sq: f64 = d_acc_us.iter().map(|&x| (x as f64).powi(2)).sum();
let d_acc_them_sq: f64 = d_acc_them.iter().map(|&x| (x as f64).powi(2)).sum();
let d_bias_sq: f64 = d_bias.iter().map(|&x| (x as f64).powi(2)).sum();
let ft_grad_sq = d_acc_us_sq * active_us.len() as f64
+ d_acc_them_sq * active_them.len() as f64
+ d_bias_sq;
let l2_grad_sq: f64 = d_l2.iter().map(|&x| (x as f64).powi(2)).sum::<f64>()
+ d_l2_bias.iter().map(|&x| (x as f64).powi(2)).sum::<f64>();
let out_grad_sq: f64 =
d_out.iter().map(|&x| (x as f64).powi(2)).sum::<f64>() + (d_out_bias as f64).powi(2);
let ft_grad_norm = ft_grad_sq.sqrt();
let l2_grad_norm = l2_grad_sq.sqrt();
let out_grad_norm = out_grad_sq.sqrt();
self.ft_grad_norm_sum += ft_grad_norm;
self.ft_grad_norm_sum_sq += ft_grad_norm * ft_grad_norm;
self.l2_grad_norm_sum += l2_grad_norm;
self.l2_grad_norm_sum_sq += l2_grad_norm * l2_grad_norm;
self.out_grad_norm_sum += out_grad_norm;
self.out_grad_norm_sum_sq += out_grad_norm * out_grad_norm;
self.out_grad_norm_values.push(out_grad_norm as f32);
let global_grad_norm = (ft_grad_sq + l2_grad_sq + out_grad_sq).sqrt();
self.global_grad_norm_values.push(global_grad_norm as f32);
if let Some(clip_norm) = self.ft_clip_norm {
let clip_norm = clip_norm as f64;
if ft_grad_norm > clip_norm {
self.ft_clip_count += 1;
let scale = (clip_norm / ft_grad_norm) as f32;
d_ft.iter_mut().for_each(|x| *x *= scale);
d_bias.iter_mut().for_each(|x| *x *= scale);
}
}
if let Some(clip_norm) = self.l2_clip_norm {
let clip_norm = clip_norm as f64;
if l2_grad_norm > clip_norm {
self.l2_clip_count += 1;
let scale = (clip_norm / l2_grad_norm) as f32;
d_l2.iter_mut().for_each(|x| *x *= scale);
d_l2_bias.iter_mut().for_each(|x| *x *= scale);
}
}
let mut out_grad_norm_after = out_grad_norm;
if let Some(clip_norm) = self.out_clip_norm {
let clip_norm = clip_norm as f64;
if out_grad_norm > clip_norm {
self.out_clip_count += 1;
let scale = (clip_norm / out_grad_norm) as f32;
d_out.iter_mut().for_each(|x| *x *= scale);
d_out_bias *= scale;
out_grad_norm_after = clip_norm;
}
}
self.out_grad_norm_after_sum += out_grad_norm_after;
self.out_grad_norm_after_sum_sq += out_grad_norm_after * out_grad_norm_after;
if let Some(clip_norm) = self.grad_clip_norm {
let clip_norm = clip_norm as f64;
if global_grad_norm > clip_norm {
self.grad_clip_count += 1;
let scale = (clip_norm / global_grad_norm) as f32;
d_ft.iter_mut().for_each(|x| *x *= scale);
d_bias.iter_mut().for_each(|x| *x *= scale);
d_l2.iter_mut().for_each(|x| *x *= scale);
d_l2_bias.iter_mut().for_each(|x| *x *= scale);
d_out.iter_mut().for_each(|x| *x *= scale);
d_out_bias *= scale;
}
}
let eligible = wdl_target.is_some();
if eligible {
let wdl_target = wdl_target.expect("eligible checked wdl_target.is_some()");
let cp_residual = (score - eval_teacher) as f64;
let wdl_residual = (score - wdl_target) as f64;
let is_conflict = cp_residual * wdl_residual < 0.0;
let (mask_ft, mask_l2) = match self.diagnostic_conflict_mask {
Some(ConflictMaskLayer::Ft) => (is_conflict, false),
Some(ConflictMaskLayer::FtAndL2) => (is_conflict, is_conflict),
None if self.diagnostic_rate_matched_mask_count > 0 => {
(self.rate_matched_should_mask(), false)
}
None => (false, false),
};
let dead_before_ft = acc_us
.iter()
.chain(acc_them.iter())
.filter(|&&x| x.clamp(0.0, 127.0) == 0.0)
.count() as u64;
let dead_before_l2 = l2_acc.iter().filter(|&&x| x <= 0.0).count() as u64;
let group = if is_conflict {
&mut self.conflict_group
} else {
&mut self.nonconflict_group
};
group.count += 1;
group.cp_residual_abs_sum += cp_residual.abs();
group.cp_residual_abs_sq_sum += cp_residual * cp_residual;
group.wdl_residual_abs_sum += wdl_residual.abs();
group.wdl_residual_abs_sq_sum += wdl_residual * wdl_residual;
group.ft_grad_norm_sum += ft_grad_norm;
group.ft_grad_norm_sq_sum += ft_grad_norm * ft_grad_norm;
group.l2_grad_norm_sum += l2_grad_norm;
group.l2_grad_norm_sq_sum += l2_grad_norm * l2_grad_norm;
if mask_ft {
d_ft.iter_mut().for_each(|x| *x = 0.0);
d_bias.iter_mut().for_each(|x| *x = 0.0);
}
if mask_l2 {
d_l2.iter_mut().for_each(|x| *x = 0.0);
d_l2_bias.iter_mut().for_each(|x| *x = 0.0);
}
if mask_ft || mask_l2 {
self.masked_position_count += 1;
}
self.pending_conflict_dead_before = Some((is_conflict, dead_before_ft, dead_before_l2));
} else {
self.pending_conflict_dead_before = None;
}
self.weights.step += 1;
let t = self.weights.step;
let lr = self.lr;
let shadow_trace_active = !self.diagnostic_shadow_trace_probe_boards.is_empty()
&& self.l2_sample_count >= self.diagnostic_shadow_trace_from_position
&& self.l2_sample_count <= self.diagnostic_shadow_trace_until_position
&& wdl_target.is_some();
let shadow_pending = shadow_trace_active.then(|| {
compute_shadow_trace(
&self.weights,
&l2_acc,
&relu_us,
&relu_them,
&acc_us,
&acc_them,
&active_us,
&active_them,
score,
eval_teacher,
wdl_target.expect("shadow_trace_active checked wdl_target.is_some()"),
weight,
self.diagnostic_shadow_trace_wdl_lambda,
lr,
t,
&d_ft,
&d_bias,
&d_l2,
&d_l2_bias,
&self.diagnostic_shadow_trace_probe_boards,
self.l2_sample_count,
)
});
let freeze_active = self.diagnostic_freeze_layer.is_some()
&& self.l2_sample_count >= self.diagnostic_freeze_from_position
&& self.l2_sample_count <= self.diagnostic_freeze_until_position;
let ft_targeted = freeze_active && self.diagnostic_freeze_layer == Some(FreezeLayer::Ft);
let ft_frozen = ft_targeted && !self.ft_periodic_active_phase() && !self.ft_reactivated();
let l2_frozen = freeze_active && self.diagnostic_freeze_layer == Some(FreezeLayer::L2);
let out_frozen = freeze_active && self.diagnostic_freeze_layer == Some(FreezeLayer::Out);
let (ft_update_sq, ft_bias_update_sq) = if ft_frozen {
d_bias.iter_mut().for_each(|x| *x = 0.0);
(0.0, 0.0)
} else {
let ft_update_sq = adam_update_slice(
&mut self.weights.ft,
&mut self.weights.ft_m,
&mut self.weights.ft_v,
&mut d_ft,
lr,
t,
);
let ft_bias_update_sq = adam_update_slice(
&mut self.weights.ft_bias,
&mut self.weights.bias_m,
&mut self.weights.bias_v,
&mut d_bias,
lr,
t,
);
(ft_update_sq, ft_bias_update_sq)
};
for j in 0..L1 {
self.ft_bias_update_sq_sum[j] += (d_bias[j] as f64).powi(2);
}
let (l2_update_sq, l2_bias_update_sq) = if l2_frozen {
d_l2_bias.iter_mut().for_each(|x| *x = 0.0);
(0.0, 0.0)
} else {
let l2_update_sq = adam_update_slice(
&mut self.weights.l2,
&mut self.weights.l2_m,
&mut self.weights.l2_v,
&mut d_l2,
lr,
t,
);
let l2_bias_update_sq = adam_update_slice(
&mut self.weights.l2_bias,
&mut self.weights.l2bias_m,
&mut self.weights.l2bias_v,
&mut d_l2_bias,
lr,
t,
);
(l2_update_sq, l2_bias_update_sq)
};
for o in 0..L2 {
self.l2_bias_update_sq_sum[o] += (d_l2_bias[o] as f64).powi(2);
}
if let Some(mut pending) = shadow_pending {
pending.record.blend_matches_real_ft = self.weights.ft == pending.shadow_blend_ft
&& self.weights.ft_bias == pending.shadow_blend_ft_bias;
pending.record.blend_matches_real_l2 = self.weights.l2 == pending.shadow_blend_l2
&& self.weights.l2_bias == pending.shadow_blend_l2_bias;
assert!(
pending.record.blend_matches_real_ft && pending.record.blend_matches_real_l2,
"shadow trace: Blend branch diverged from the real applied update at position {}",
pending.record.position_index
);
self.shadow_trace_records.push(pending.record);
}
if let Some((is_conflict, dead_before_ft, dead_before_l2)) =
self.pending_conflict_dead_before.take()
{
let dead_after_ft = ft_dead_count(&self.weights.ft, &self.weights.ft_bias, board);
let dead_after_l2 = l2_state_for_board(
&self.weights.ft,
&self.weights.ft_bias,
&self.weights.l2,
&self.weights.l2_bias,
board,
)
.0 as u64;
let group = if is_conflict {
&mut self.conflict_group
} else {
&mut self.nonconflict_group
};
group.new_dead_ft_sum += dead_after_ft.saturating_sub(dead_before_ft);
group.new_dead_l2_sum += dead_after_l2.saturating_sub(dead_before_l2);
}
let (out_update_sq, out_bias_delta) = if out_frozen {
(0.0, 0.0f32)
} else {
let out_update_sq = adam_update_slice(
&mut self.weights.out,
&mut self.weights.out_m,
&mut self.weights.out_v,
&mut d_out,
lr,
t,
);
let out_bias_delta = adam_update_scalar(
&mut self.weights.out_bias,
&mut self.weights.obias_m,
&mut self.weights.obias_v,
d_out_bias,
lr,
t,
);
(out_update_sq, out_bias_delta)
};
let ft_update_norm = (ft_update_sq + ft_bias_update_sq).sqrt();
let l2_update_norm = (l2_update_sq + l2_bias_update_sq).sqrt();
let out_update_norm = (out_update_sq + (out_bias_delta as f64).powi(2)).sqrt();
self.ft_update_norm_sum += ft_update_norm;
self.ft_update_norm_sum_sq += ft_update_norm * ft_update_norm;
self.l2_update_norm_sum += l2_update_norm;
self.l2_update_norm_sum_sq += l2_update_norm * l2_update_norm;
self.out_update_norm_sum += out_update_norm;
self.out_update_norm_sum_sq += out_update_norm * out_update_norm;
self.maybe_trace_snapshot();
}
fn maybe_trace_snapshot(&mut self) {
if !self.trace_positions.contains(&self.l2_sample_count) {
return;
}
if self.weight_snapshot_trace {
self.weight_snapshots
.push((self.l2_sample_count, self.weights.clone()));
}
let l2_weight_row_norm: Vec<f32> = (0..L2)
.map(|o| {
(0..2 * L1)
.map(|j| self.weights.l2[j * L2 + o].powi(2))
.sum::<f32>()
.sqrt()
})
.collect();
let ft_weight_row_norm: Vec<f32> = (0..L1)
.map(|j| {
(0..INPUT)
.map(|feat| self.weights.ft[feat * L1 + j].powi(2))
.sum::<f32>()
.sqrt()
})
.collect();
let l2 = diagnostics::build_trace_layer_snapshot(
&self.l2_values,
&self.l2_weighted_input_values,
&self.l2_zero_count,
&self.l2_sat_count,
self.l2_sample_count,
l2_weight_row_norm,
self.weights.l2_bias.clone(),
&self.l2_dacc_sum,
&self.l2_dacc_sq_sum,
&self.l2_dacc_pos_count,
&self.l2_dacc_neg_count,
&self.l2_bias_update_sq_sum,
);
let ft = diagnostics::build_trace_layer_snapshot(
&[], &[], &self.ft_zero_count,
&self.ft_sat_count,
self.l2_sample_count,
ft_weight_row_norm,
self.weights.ft_bias.clone(),
&self.ft_dacc_sum,
&self.ft_dacc_sq_sum,
&self.ft_dacc_pos_count,
&self.ft_dacc_neg_count,
&self.ft_bias_update_sq_sum,
);
let (l2_input_norm_mean, l2_input_norm_std) = diagnostics::mean_std(
self.l2_input_norm_sum,
self.l2_input_norm_sq_sum,
self.l2_sample_count,
);
let (ft_output_mean, ft_output_std) = diagnostics::mean_std(
self.ft_output_sum,
self.ft_output_sum_sq,
self.ft_output_count,
);
let cp_wdl = if self.cp_wdl_grad_trace && self.wdl_component_count > 0 {
let n = self.wdl_component_count;
let (cp_ft_grad_rms, _) =
diagnostics::mean_std(self.cp_ft_grad_norm_sum, self.cp_ft_grad_norm_sum_sq, n);
let (wdl_ft_grad_rms, _) =
diagnostics::mean_std(self.wdl_ft_grad_norm_sum, self.wdl_ft_grad_norm_sum_sq, n);
let (cp_l2_grad_rms, _) =
diagnostics::mean_std(self.cp_l2_grad_norm_sum, self.cp_l2_grad_norm_sum_sq, n);
let (wdl_l2_grad_rms, _) =
diagnostics::mean_std(self.wdl_l2_grad_norm_sum, self.wdl_l2_grad_norm_sum_sq, n);
let (cp_out_grad_rms, _) =
diagnostics::mean_std(self.cp_out_grad_norm_sum, self.cp_out_grad_norm_sum_sq, n);
let (wdl_out_grad_rms, _) =
diagnostics::mean_std(self.wdl_out_grad_norm_sum, self.wdl_out_grad_norm_sum_sq, n);
let (cp_target_mean, cp_target_std) =
diagnostics::mean_std(self.cp_target_sum, self.cp_target_sum_sq, n);
let (wdl_target_mean, wdl_target_std) =
diagnostics::mean_std(self.wdl_target_sum, self.wdl_target_sum_sq, n);
let (prediction_mean, prediction_std) =
diagnostics::mean_std(self.prediction_sum, self.prediction_sum_sq, n);
let (cp_residual_mean, cp_residual_std) =
diagnostics::mean_std(self.cp_residual_sum, self.cp_residual_sum_sq, n);
let (wdl_residual_mean, wdl_residual_std) =
diagnostics::mean_std(self.wdl_residual_sum, self.wdl_residual_sum_sq, n);
let (cp_d_output_mean, cp_d_output_std) =
diagnostics::mean_std(self.cp_d_output_sum, self.cp_d_output_sum_sq, n);
let (wdl_d_output_mean, wdl_d_output_std) =
diagnostics::mean_std(self.wdl_d_output_sum, self.wdl_d_output_sum_sq, n);
Some(diagnostics::CpWdlTrace {
l2: diagnostics::build_cp_wdl_layer_trace(
&self.l2_cp_dacc_sum,
&self.l2_cp_dacc_sq_sum,
&self.l2_cp_dacc_pos_count,
&self.l2_cp_dacc_neg_count,
&self.l2_wdl_dacc_sum,
&self.l2_wdl_dacc_sq_sum,
&self.l2_wdl_dacc_pos_count,
&self.l2_wdl_dacc_neg_count,
&self.l2_cp_wdl_dot_sum,
n,
),
ft: diagnostics::build_cp_wdl_layer_trace(
&self.ft_cp_dacc_sum,
&self.ft_cp_dacc_sq_sum,
&self.ft_cp_dacc_pos_count,
&self.ft_cp_dacc_neg_count,
&self.ft_wdl_dacc_sum,
&self.ft_wdl_dacc_sq_sum,
&self.ft_wdl_dacc_pos_count,
&self.ft_wdl_dacc_neg_count,
&self.ft_cp_wdl_dot_sum,
n,
),
cp_ft_grad_rms,
wdl_ft_grad_rms,
cp_l2_grad_rms,
wdl_l2_grad_rms,
cp_out_grad_rms,
wdl_out_grad_rms,
cp_target_mean,
cp_target_std,
wdl_target_mean,
wdl_target_std,
prediction_mean,
prediction_std,
cp_residual_mean,
cp_residual_std,
wdl_residual_mean,
wdl_residual_std,
cp_d_output_mean,
cp_d_output_std,
wdl_d_output_mean,
wdl_d_output_std,
})
} else {
None
};
self.trace_snapshots.push(diagnostics::TraceSnapshot {
position_index: self.l2_sample_count,
l2,
ft,
l2_input_norm_mean,
l2_input_norm_std,
ft_output_mean,
ft_output_std,
cp_wdl,
});
}
pub fn avg_loss(&self) -> f64 {
if self.total_weight > 0.0 {
self.total_loss / self.total_weight
} else {
0.0
}
}
pub fn reset_epoch_stats(&mut self) {
self.total_loss = 0.0;
self.total_count = 0;
self.total_weight = 0.0;
self.dropped_missing = 0;
self.ft_ever_active.iter_mut().for_each(|b| *b = false);
self.ft_ever_saturated.iter_mut().for_each(|b| *b = false);
self.l2_ever_active.iter_mut().for_each(|b| *b = false);
self.l2_ever_saturated.iter_mut().for_each(|b| *b = false);
self.output_sum = 0.0;
self.output_sum_sq = 0.0;
self.l2_zero_count.iter_mut().for_each(|c| *c = 0);
self.l2_sat_count.iter_mut().for_each(|c| *c = 0);
self.l2_sample_count = 0;
self.l2_values.iter_mut().for_each(|v| v.clear());
self.ft_grad_norm_sum = 0.0;
self.ft_grad_norm_sum_sq = 0.0;
self.l2_grad_norm_sum = 0.0;
self.l2_grad_norm_sum_sq = 0.0;
self.out_grad_norm_sum = 0.0;
self.out_grad_norm_sum_sq = 0.0;
self.global_grad_norm_values.clear();
self.ft_update_norm_sum = 0.0;
self.ft_update_norm_sum_sq = 0.0;
self.l2_update_norm_sum = 0.0;
self.l2_update_norm_sum_sq = 0.0;
self.out_update_norm_sum = 0.0;
self.out_update_norm_sum_sq = 0.0;
self.target_sum = 0.0;
self.target_sum_sq = 0.0;
self.eval_teacher_sum = 0.0;
self.eval_teacher_sum_sq = 0.0;
self.pred_eval_prod_sum = 0.0;
self.cp_component_sum = 0.0;
self.wdl_component_sum = 0.0;
self.wdl_component_count = 0;
self.grad_clip_count = 0;
self.ft_clip_count = 0;
self.l2_clip_count = 0;
self.out_clip_count = 0;
self.out_grad_norm_values.clear();
self.out_grad_norm_after_sum = 0.0;
self.out_grad_norm_after_sum_sq = 0.0;
self.cache_hits = 0;
self.cache_misses = 0;
self.search_time_ns = 0;
self.trace_snapshots.clear();
self.weight_snapshots.clear();
self.l2_weighted_input_values
.iter_mut()
.for_each(|v| v.clear());
self.l2_dacc_sum.iter_mut().for_each(|x| *x = 0.0);
self.l2_dacc_sq_sum.iter_mut().for_each(|x| *x = 0.0);
self.l2_dacc_pos_count.iter_mut().for_each(|x| *x = 0);
self.l2_dacc_neg_count.iter_mut().for_each(|x| *x = 0);
self.ft_dacc_sum.iter_mut().for_each(|x| *x = 0.0);
self.ft_dacc_sq_sum.iter_mut().for_each(|x| *x = 0.0);
self.ft_dacc_pos_count.iter_mut().for_each(|x| *x = 0);
self.ft_dacc_neg_count.iter_mut().for_each(|x| *x = 0);
self.l2_bias_update_sq_sum.iter_mut().for_each(|x| *x = 0.0);
self.ft_bias_update_sq_sum.iter_mut().for_each(|x| *x = 0.0);
self.ft_zero_count.iter_mut().for_each(|x| *x = 0);
self.ft_sat_count.iter_mut().for_each(|x| *x = 0);
self.l2_input_norm_sum = 0.0;
self.l2_input_norm_sq_sum = 0.0;
self.ft_output_sum = 0.0;
self.ft_output_sum_sq = 0.0;
self.ft_output_count = 0;
self.l2_cp_dacc_sum.iter_mut().for_each(|x| *x = 0.0);
self.l2_cp_dacc_sq_sum.iter_mut().for_each(|x| *x = 0.0);
self.l2_cp_dacc_pos_count.iter_mut().for_each(|x| *x = 0);
self.l2_cp_dacc_neg_count.iter_mut().for_each(|x| *x = 0);
self.l2_wdl_dacc_sum.iter_mut().for_each(|x| *x = 0.0);
self.l2_wdl_dacc_sq_sum.iter_mut().for_each(|x| *x = 0.0);
self.l2_wdl_dacc_pos_count.iter_mut().for_each(|x| *x = 0);
self.l2_wdl_dacc_neg_count.iter_mut().for_each(|x| *x = 0);
self.l2_cp_wdl_dot_sum.iter_mut().for_each(|x| *x = 0.0);
self.ft_cp_dacc_sum.iter_mut().for_each(|x| *x = 0.0);
self.ft_cp_dacc_sq_sum.iter_mut().for_each(|x| *x = 0.0);
self.ft_cp_dacc_pos_count.iter_mut().for_each(|x| *x = 0);
self.ft_cp_dacc_neg_count.iter_mut().for_each(|x| *x = 0);
self.ft_wdl_dacc_sum.iter_mut().for_each(|x| *x = 0.0);
self.ft_wdl_dacc_sq_sum.iter_mut().for_each(|x| *x = 0.0);
self.ft_wdl_dacc_pos_count.iter_mut().for_each(|x| *x = 0);
self.ft_wdl_dacc_neg_count.iter_mut().for_each(|x| *x = 0);
self.ft_cp_wdl_dot_sum.iter_mut().for_each(|x| *x = 0.0);
self.cp_ft_grad_norm_sum = 0.0;
self.cp_ft_grad_norm_sum_sq = 0.0;
self.wdl_ft_grad_norm_sum = 0.0;
self.wdl_ft_grad_norm_sum_sq = 0.0;
self.cp_l2_grad_norm_sum = 0.0;
self.cp_l2_grad_norm_sum_sq = 0.0;
self.wdl_l2_grad_norm_sum = 0.0;
self.wdl_l2_grad_norm_sum_sq = 0.0;
self.cp_out_grad_norm_sum = 0.0;
self.cp_out_grad_norm_sum_sq = 0.0;
self.wdl_out_grad_norm_sum = 0.0;
self.wdl_out_grad_norm_sum_sq = 0.0;
self.cp_target_sum = 0.0;
self.cp_target_sum_sq = 0.0;
self.wdl_target_sum = 0.0;
self.wdl_target_sum_sq = 0.0;
self.prediction_sum = 0.0;
self.prediction_sum_sq = 0.0;
self.cp_residual_sum = 0.0;
self.cp_residual_sum_sq = 0.0;
self.wdl_residual_sum = 0.0;
self.wdl_residual_sum_sq = 0.0;
self.cp_d_output_sum = 0.0;
self.cp_d_output_sum_sq = 0.0;
self.wdl_d_output_sum = 0.0;
self.wdl_d_output_sum_sq = 0.0;
self.sample_grad_records.clear();
self.sample_grad_prev_d_l2_acc = None;
self.sample_grad_running_mean_d_l2_acc = [0.0; L2];
self.sample_grad_running_count = 0;
self.shadow_trace_records.clear();
self.masked_position_count = 0;
self.conflict_group = ConflictGroupStats::default();
self.nonconflict_group = ConflictGroupStats::default();
self.rate_matched_remaining_needed = self.diagnostic_rate_matched_mask_count;
self.rate_matched_remaining_pool = self.diagnostic_rate_matched_mask_total;
self.rate_matched_rng = Lcg(self.diagnostic_rate_matched_mask_seed ^ 0xA5A5_5A5A_1234_5678);
}
fn rate_matched_should_mask(&mut self) -> bool {
if self.rate_matched_remaining_pool == 0 {
return false;
}
let r = self.rate_matched_rng.next_u64() as f64 / u64::MAX as f64;
let threshold =
self.rate_matched_remaining_needed as f64 / self.rate_matched_remaining_pool as f64;
let select = r < threshold;
if select {
self.rate_matched_remaining_needed =
self.rate_matched_remaining_needed.saturating_sub(1);
}
self.rate_matched_remaining_pool -= 1;
select
}
}
fn active_features(board: &Board, perspective: Color) -> Vec<usize> {
const ALL_KINDS: [PieceKind; 14] = [
PieceKind::Fu,
PieceKind::Kyou,
PieceKind::Kei,
PieceKind::Gin,
PieceKind::Kin,
PieceKind::Kaku,
PieceKind::Hisha,
PieceKind::Ou,
PieceKind::Tokin,
PieceKind::Narikyo,
PieceKind::Narikei,
PieceKind::Narigin,
PieceKind::Uma,
PieceKind::Ryu,
];
const HAND_KINDS: [PieceKind; 7] = [
PieceKind::Fu,
PieceKind::Kyou,
PieceKind::Kei,
PieceKind::Gin,
PieceKind::Kin,
PieceKind::Kaku,
PieceKind::Hisha,
];
let mut features = Vec::with_capacity(60);
for &kind in &ALL_KINDS {
for color in [Color::Black, Color::White] {
let mut bb = board.pieces(color, kind);
while let Some(sq) = bb.pop_lsb() {
features.push(feature_index(sq, kind, color, perspective));
}
}
}
for &kind in &HAND_KINDS {
for color in [Color::Black, Color::White] {
let count = board.hand(color).get(kind);
for n in 1..=count {
features.push(hand_feature_index(kind, n, color, perspective));
}
}
}
features
}
struct DiagnosticGrad {
d_l2_acc: [f32; L2],
d_ft_acc: Vec<f32>,
l2_grad_norm: f64,
ft_grad_norm: f64,
out_grad_norm: f64,
d_output: f32,
}
#[allow(clippy::too_many_arguments)]
fn diagnostic_backward(
w: &TrainWeights,
l2_acc: &[f32],
relu_l2: &[f32],
acc_us: &[f32],
acc_them: &[f32],
relu_us: &[f32],
relu_them: &[f32],
active_us: &[usize],
active_them: &[usize],
err: f32,
weight: f32,
) -> DiagnosticGrad {
let d_score = weight * 2.0 * err;
let d_output = d_score / 64.0;
let d_out: Vec<f32> = relu_l2.iter().map(|&r| d_output * r).collect();
let d_out_bias = d_output;
let mut d_l2_acc = [0.0f32; L2];
for o in 0..L2 {
if l2_acc[o] > 0.0 && l2_acc[o] < 127.0 {
d_l2_acc[o] = d_output * w.out[o];
}
}
let mut d_l2 = vec![0.0f32; 2 * L1 * L2];
let mut d_l2_bias = [0.0f32; L2];
let mut d_relu_us = vec![0.0f32; L1];
let mut d_relu_them = vec![0.0f32; L1];
for j in 0..L1 {
let base_us = j * L2;
let base_them = (L1 + j) * L2;
for o in 0..L2 {
let g = d_l2_acc[o];
d_l2[base_us + o] += g * relu_us[j];
d_l2[base_them + o] += g * relu_them[j];
d_relu_us[j] += g * w.l2[base_us + o];
d_relu_them[j] += g * w.l2[base_them + o];
}
}
d_l2_bias[..L2].copy_from_slice(&d_l2_acc[..L2]);
let mut d_acc_us = vec![0.0f32; L1];
let mut d_acc_them = vec![0.0f32; L1];
for j in 0..L1 {
if acc_us[j] > 0.0 && acc_us[j] < 127.0 {
d_acc_us[j] = d_relu_us[j];
}
if acc_them[j] > 0.0 && acc_them[j] < 127.0 {
d_acc_them[j] = d_relu_them[j];
}
}
let mut d_ft_acc = vec![0.0f32; L1];
for j in 0..L1 {
d_ft_acc[j] = d_acc_us[j] + d_acc_them[j];
}
let d_acc_us_sq: f64 = d_acc_us.iter().map(|&x| (x as f64).powi(2)).sum();
let d_acc_them_sq: f64 = d_acc_them.iter().map(|&x| (x as f64).powi(2)).sum();
let d_bias_sq: f64 = d_ft_acc.iter().map(|&x| (x as f64).powi(2)).sum();
let ft_grad_sq =
d_acc_us_sq * active_us.len() as f64 + d_acc_them_sq * active_them.len() as f64 + d_bias_sq;
let l2_grad_sq: f64 = d_l2.iter().map(|&x| (x as f64).powi(2)).sum::<f64>()
+ d_l2_bias.iter().map(|&x| (x as f64).powi(2)).sum::<f64>();
let out_grad_sq: f64 =
d_out.iter().map(|&x| (x as f64).powi(2)).sum::<f64>() + (d_out_bias as f64).powi(2);
DiagnosticGrad {
d_l2_acc,
d_ft_acc,
l2_grad_norm: l2_grad_sq.sqrt(),
ft_grad_norm: ft_grad_sq.sqrt(),
out_grad_norm: out_grad_sq.sqrt(),
d_output,
}
}
#[allow(clippy::too_many_arguments)]
fn shadow_component_grad(
w: &TrainWeights,
l2_acc: &[f32],
relu_us: &[f32],
relu_them: &[f32],
acc_us: &[f32],
acc_them: &[f32],
active_us: &[usize],
active_them: &[usize],
score: f32,
target: f32,
weight: f32,
) -> (Vec<f32>, Vec<f32>, Vec<f32>, Vec<f32>) {
let d_score = weight * 2.0 * (score - target);
let d_output = d_score / 64.0;
let mut d_l2_acc = [0.0f32; L2];
for o in 0..L2 {
if l2_acc[o] > 0.0 && l2_acc[o] < 127.0 {
d_l2_acc[o] = d_output * w.out[o];
}
}
let mut d_l2 = vec![0.0f32; 2 * L1 * L2];
let mut d_l2_bias = vec![0.0f32; L2];
let mut d_relu_us = vec![0.0f32; L1];
let mut d_relu_them = vec![0.0f32; L1];
for j in 0..L1 {
let base_us = j * L2;
let base_them = (L1 + j) * L2;
for o in 0..L2 {
let g = d_l2_acc[o];
d_l2[base_us + o] += g * relu_us[j];
d_l2[base_them + o] += g * relu_them[j];
d_relu_us[j] += g * w.l2[base_us + o];
d_relu_them[j] += g * w.l2[base_them + o];
}
}
d_l2_bias.copy_from_slice(&d_l2_acc);
let mut d_acc_us = vec![0.0f32; L1];
let mut d_acc_them = vec![0.0f32; L1];
for j in 0..L1 {
if acc_us[j] > 0.0 && acc_us[j] < 127.0 {
d_acc_us[j] = d_relu_us[j];
}
if acc_them[j] > 0.0 && acc_them[j] < 127.0 {
d_acc_them[j] = d_relu_them[j];
}
}
let mut d_ft = vec![0.0f32; INPUT * L1];
let mut d_bias = vec![0.0f32; L1];
for &feat in active_us {
let base = feat * L1;
for j in 0..L1 {
d_ft[base + j] += d_acc_us[j];
}
}
for &feat in active_them {
let base = feat * L1;
for j in 0..L1 {
d_ft[base + j] += d_acc_them[j];
}
}
for j in 0..L1 {
d_bias[j] = d_acc_us[j] + d_acc_them[j];
}
(d_ft, d_bias, d_l2, d_l2_bias)
}
fn apply_shadow_adam(
w: &TrainWeights,
d_ft: &mut [f32],
d_bias: &mut [f32],
d_l2: &mut [f32],
d_l2_bias: &mut [f32],
lr: f32,
t: u64,
) -> (Vec<f32>, Vec<f32>, Vec<f32>, Vec<f32>) {
let mut ft = w.ft.clone();
let mut ft_m = w.ft_m.clone();
let mut ft_v = w.ft_v.clone();
adam_update_slice(&mut ft, &mut ft_m, &mut ft_v, d_ft, lr, t);
let mut ft_bias = w.ft_bias.clone();
let mut bias_m = w.bias_m.clone();
let mut bias_v = w.bias_v.clone();
adam_update_slice(&mut ft_bias, &mut bias_m, &mut bias_v, d_bias, lr, t);
let mut l2 = w.l2.clone();
let mut l2_m = w.l2_m.clone();
let mut l2_v = w.l2_v.clone();
adam_update_slice(&mut l2, &mut l2_m, &mut l2_v, d_l2, lr, t);
let mut l2_bias = w.l2_bias.clone();
let mut l2bias_m = w.l2bias_m.clone();
let mut l2bias_v = w.l2bias_v.clone();
adam_update_slice(&mut l2_bias, &mut l2bias_m, &mut l2bias_v, d_l2_bias, lr, t);
(ft, ft_bias, l2, l2_bias)
}
fn ft_dead_mask(ft: &[f32], ft_bias: &[f32], board: &Board) -> Vec<bool> {
let stm = board.side_to_move;
let active_us = active_features(board, stm);
let active_them = active_features(board, stm.flip());
let mut acc_us = ft_bias.to_vec();
let mut acc_them = acc_us.clone();
for &feat in &active_us {
let base = feat * L1;
for j in 0..L1 {
acc_us[j] += ft[base + j];
}
}
for &feat in &active_them {
let base = feat * L1;
for j in 0..L1 {
acc_them[j] += ft[base + j];
}
}
acc_us
.iter()
.chain(acc_them.iter())
.map(|&x| x.clamp(0.0, 127.0) == 0.0)
.collect()
}
fn ft_dead_count(ft: &[f32], ft_bias: &[f32], board: &Board) -> u64 {
ft_dead_mask(ft, ft_bias, board)
.iter()
.filter(|&&d| d)
.count() as u64
}
fn l2_state_for_board(
ft: &[f32],
ft_bias: &[f32],
l2: &[f32],
l2_bias: &[f32],
board: &Board,
) -> (u32, f64) {
let stm = board.side_to_move;
let active_us = active_features(board, stm);
let active_them = active_features(board, stm.flip());
let mut acc_us = ft_bias.to_vec();
let mut acc_them = acc_us.clone();
for &feat in &active_us {
let base = feat * L1;
for j in 0..L1 {
acc_us[j] += ft[base + j];
}
}
for &feat in &active_them {
let base = feat * L1;
for j in 0..L1 {
acc_them[j] += ft[base + j];
}
}
let relu_us: Vec<f32> = acc_us.iter().map(|&x| x.clamp(0.0, 127.0)).collect();
let relu_them: Vec<f32> = acc_them.iter().map(|&x| x.clamp(0.0, 127.0)).collect();
let mut l2_acc = l2_bias.to_vec();
for j in 0..L1 {
let a = relu_us[j];
let b = relu_them[j];
let base_us = j * L2;
let base_them = (L1 + j) * L2;
for o in 0..L2 {
l2_acc[o] += a * l2[base_us + o];
l2_acc[o] += b * l2[base_them + o];
}
}
let mut dead = 0u32;
let mut weighted_input_sum = 0.0f64;
for o in 0..L2 {
if l2_acc[o] <= 0.0 {
dead += 1;
}
weighted_input_sum += (l2_acc[o] - l2_bias[o]) as f64;
}
(dead, weighted_input_sum)
}
fn vec_norm_f64(parts: &[&[f32]]) -> f64 {
parts
.iter()
.flat_map(|p| p.iter())
.map(|&x| (x as f64).powi(2))
.sum::<f64>()
.sqrt()
}
fn vec_dot_f64(a_parts: &[&[f32]], b_parts: &[&[f32]]) -> f64 {
a_parts
.iter()
.zip(b_parts.iter())
.map(|(a, b)| {
a.iter()
.zip(b.iter())
.map(|(&x, &y)| (x as f64) * (y as f64))
.sum::<f64>()
})
.sum()
}
struct ShadowTracePending {
record: diagnostics::ShadowTraceRecord,
shadow_blend_ft: Vec<f32>,
shadow_blend_ft_bias: Vec<f32>,
shadow_blend_l2: Vec<f32>,
shadow_blend_l2_bias: Vec<f32>,
}
#[allow(clippy::too_many_arguments)]
fn compute_shadow_trace(
w: &TrainWeights,
l2_acc: &[f32],
relu_us: &[f32],
relu_them: &[f32],
acc_us: &[f32],
acc_them: &[f32],
active_us: &[usize],
active_them: &[usize],
score: f32,
eval_teacher: f32,
wdl_target: f32,
weight: f32,
lambda: f32,
lr: f32,
t: u64,
real_d_ft: &[f32],
real_d_bias: &[f32],
real_d_l2: &[f32],
real_d_l2_bias: &[f32],
probe_boards: &[Board],
position_index: u64,
) -> ShadowTracePending {
let (mut d_ft_cp, mut d_bias_cp, mut d_l2_cp, mut d_l2_bias_cp) = shadow_component_grad(
w,
l2_acc,
relu_us,
relu_them,
acc_us,
acc_them,
active_us,
active_them,
score,
eval_teacher,
weight * lambda,
);
let (mut d_ft_wdl, mut d_bias_wdl, mut d_l2_wdl, mut d_l2_bias_wdl) = shadow_component_grad(
w,
l2_acc,
relu_us,
relu_them,
acc_us,
acc_them,
active_us,
active_them,
score,
wdl_target,
weight * (1.0 - lambda),
);
let g_cp_norm = vec_norm_f64(&[&d_ft_cp, &d_bias_cp, &d_l2_cp, &d_l2_bias_cp]);
let g_wdl_norm = vec_norm_f64(&[&d_ft_wdl, &d_bias_wdl, &d_l2_wdl, &d_l2_bias_wdl]);
let g_dot = vec_dot_f64(
&[&d_ft_cp, &d_bias_cp, &d_l2_cp, &d_l2_bias_cp],
&[&d_ft_wdl, &d_bias_wdl, &d_l2_wdl, &d_l2_bias_wdl],
);
let cos_g_cp_wdl = if g_cp_norm > 0.0 && g_wdl_norm > 0.0 {
g_dot / (g_cp_norm * g_wdl_norm)
} else {
0.0
};
let (cp_ft, cp_ft_bias, cp_l2, cp_l2_bias) = apply_shadow_adam(
w,
&mut d_ft_cp,
&mut d_bias_cp,
&mut d_l2_cp,
&mut d_l2_bias_cp,
lr,
t,
);
let (wdl_ft, wdl_ft_bias, wdl_l2, wdl_l2_bias) = apply_shadow_adam(
w,
&mut d_ft_wdl,
&mut d_bias_wdl,
&mut d_l2_wdl,
&mut d_l2_bias_wdl,
lr,
t,
);
let delta_cp_norm = vec_norm_f64(&[&d_ft_cp, &d_bias_cp]);
let delta_wdl_norm = vec_norm_f64(&[&d_ft_wdl, &d_bias_wdl]);
let delta_dot = vec_dot_f64(&[&d_ft_cp, &d_bias_cp], &[&d_ft_wdl, &d_bias_wdl]);
let cos_delta_cp_wdl = if delta_cp_norm > 0.0 && delta_wdl_norm > 0.0 {
delta_dot / (delta_cp_norm * delta_wdl_norm)
} else {
0.0
};
let mut d_ft_blend = real_d_ft.to_vec();
let mut d_bias_blend = real_d_bias.to_vec();
let mut d_l2_blend = real_d_l2.to_vec();
let mut d_l2_bias_blend = real_d_l2_bias.to_vec();
let (blend_ft, blend_ft_bias, blend_l2, blend_l2_bias) = apply_shadow_adam(
w,
&mut d_ft_blend,
&mut d_bias_blend,
&mut d_l2_blend,
&mut d_l2_bias_blend,
lr,
t,
);
let delta_blend_norm = vec_norm_f64(&[&d_ft_blend, &d_bias_blend]);
let linpred_ft: Vec<f32> =
w.ft.iter()
.zip(d_ft_cp.iter())
.zip(d_ft_wdl.iter())
.map(|((&a, &dc), &dw)| a + dc + dw)
.collect();
let linpred_ft_bias: Vec<f32> = w
.ft_bias
.iter()
.zip(d_bias_cp.iter())
.zip(d_bias_wdl.iter())
.map(|((&a, &dc), &dw)| a + dc + dw)
.collect();
let mut contingency = [0u64; 8];
let mut blend_dead_linpred_alive = 0u64;
let mut blend_dead_linpred_dead = 0u64;
let mut blend_alive_linpred_dead = 0u64;
let mut blend_alive_linpred_alive = 0u64;
let mut n_alive_at_anchor = 0u64;
let mut l2_dead_cp = 0u32;
let mut l2_dead_wdl = 0u32;
let mut l2_dead_blend = 0u32;
let mut l2_wsum_cp = 0.0f64;
let mut l2_wsum_wdl = 0.0f64;
let mut l2_wsum_blend = 0.0f64;
for board in probe_boards {
let anchor_mask = ft_dead_mask(&w.ft, &w.ft_bias, board);
let cp_mask = ft_dead_mask(&cp_ft, &cp_ft_bias, board);
let wdl_mask = ft_dead_mask(&wdl_ft, &wdl_ft_bias, board);
let blend_mask = ft_dead_mask(&blend_ft, &blend_ft_bias, board);
let linpred_mask = ft_dead_mask(&linpred_ft, &linpred_ft_bias, board);
for u in 0..2 * L1 {
if anchor_mask[u] {
continue;
}
n_alive_at_anchor += 1;
let idx =
(cp_mask[u] as usize) * 4 + (wdl_mask[u] as usize) * 2 + (blend_mask[u] as usize);
contingency[idx] += 1;
match (blend_mask[u], linpred_mask[u]) {
(true, false) => blend_dead_linpred_alive += 1,
(true, true) => blend_dead_linpred_dead += 1,
(false, true) => blend_alive_linpred_dead += 1,
(false, false) => blend_alive_linpred_alive += 1,
}
}
let (dead, wsum) = l2_state_for_board(&cp_ft, &cp_ft_bias, &cp_l2, &cp_l2_bias, board);
l2_dead_cp += dead;
l2_wsum_cp += wsum;
let (dead, wsum) = l2_state_for_board(&wdl_ft, &wdl_ft_bias, &wdl_l2, &wdl_l2_bias, board);
l2_dead_wdl += dead;
l2_wsum_wdl += wsum;
let (dead, wsum) =
l2_state_for_board(&blend_ft, &blend_ft_bias, &blend_l2, &blend_l2_bias, board);
l2_dead_blend += dead;
l2_wsum_blend += wsum;
}
let l2_total = (probe_boards.len() * L2) as f64;
let record = diagnostics::ShadowTraceRecord {
position_index,
g_cp_norm,
g_wdl_norm,
cos_g_cp_wdl,
delta_cp_norm,
delta_wdl_norm,
delta_blend_norm,
cos_delta_cp_wdl,
contingency_cp_wdl_blend: contingency,
blend_dead_linpred_alive,
blend_dead_linpred_dead,
blend_alive_linpred_dead,
blend_alive_linpred_alive,
n_alive_at_anchor,
l2_dead_frac_cp: l2_dead_cp as f64 / l2_total,
l2_dead_frac_wdl: l2_dead_wdl as f64 / l2_total,
l2_dead_frac_blend: l2_dead_blend as f64 / l2_total,
l2_weighted_input_mean_cp: l2_wsum_cp / l2_total,
l2_weighted_input_mean_wdl: l2_wsum_wdl / l2_total,
l2_weighted_input_mean_blend: l2_wsum_blend / l2_total,
blend_matches_real_ft: false,
blend_matches_real_l2: false,
};
ShadowTracePending {
record,
shadow_blend_ft: blend_ft,
shadow_blend_ft_bias: blend_ft_bias,
shadow_blend_l2: blend_l2,
shadow_blend_l2_bias: blend_l2_bias,
}
}
fn adam_update_slice(
params: &mut [f32],
m: &mut [f32],
v: &mut [f32],
grads: &mut [f32],
lr: f32,
t: u64,
) -> f64 {
let mut delta_sq_sum = 0.0f64;
for i in 0..params.len() {
let delta = adam_update_scalar(&mut params[i], &mut m[i], &mut v[i], grads[i], lr, t);
delta_sq_sum += (delta as f64) * (delta as f64);
grads[i] = delta;
}
delta_sq_sum
}
#[inline]
fn adam_update_scalar(
param: &mut f32,
m: &mut f32,
v: &mut f32,
grad: f32,
lr: f32,
t: u64,
) -> f32 {
const B1: f32 = 0.9;
const B2: f32 = 0.999;
const EPS: f32 = 1e-8;
*m = B1 * *m + (1.0 - B1) * grad;
*v = B2 * *v + (1.0 - B2) * grad * grad;
let m_hat = *m / (1.0 - B1.powi(t as i32));
let v_hat = *v / (1.0 - B2.powi(t as i32));
let delta = -lr * m_hat / (v_hat.sqrt() + EPS);
*param += delta;
delta
}
#[cfg(test)]
mod tests {
use super::*;
fn variance(xs: &[f32]) -> f32 {
let mean = xs.iter().sum::<f32>() / xs.len() as f32;
xs.iter().map(|x| (x - mean).powi(2)).sum::<f32>() / xs.len() as f32
}
#[test]
fn seeded_init_breaks_symmetry_within_each_layer() {
let w = TrainWeights::new_seeded(42, 0.5);
assert!(variance(&w.ft[0..L1]) > 0.0);
assert!(variance(&w.l2[0..L2]) > 0.0);
assert!(variance(&w.out) > 0.0);
}
#[test]
fn seeded_init_is_deterministic() {
let a = TrainWeights::new_seeded(42, 0.5);
let b = TrainWeights::new_seeded(42, 0.5);
assert_eq!(a.ft, b.ft);
assert_eq!(a.l2, b.l2);
assert_eq!(a.out, b.out);
}
#[test]
fn l2_bias_init_only_touches_l2_bias() {
let default_bias = TrainWeights::new_seeded(42, 0.5);
let custom_bias = TrainWeights::new_seeded(42, 3.0);
assert_eq!(custom_bias.l2_bias, vec![3.0; L2]);
assert_eq!(default_bias.l2_bias, vec![0.5; L2]);
assert_eq!(default_bias.ft, custom_bias.ft);
assert_eq!(default_bias.l2, custom_bias.l2);
assert_eq!(default_bias.out, custom_bias.out);
assert_eq!(default_bias.ft_bias, custom_bias.ft_bias);
}
#[test]
fn from_nnue_weights_round_trips_forward_output() {
let mut t = Trainer::new(42, 0.5);
let board = Board::startpos();
let before = t.forward(&board);
let nn = t.weights.to_nnue_weights();
t.weights = TrainWeights::from_nnue_weights(&nn);
let after = t.forward(&board);
assert!(
(before - after).abs() < 1.0,
"before={before} after={after}"
);
}
#[test]
fn seeded_init_differs_across_seeds() {
let a = TrainWeights::new_seeded(1, 0.5);
let b = TrainWeights::new_seeded(2, 0.5);
assert_ne!(a.ft, b.ft);
}
#[test]
fn wdl_target_black_win_from_black_perspective_is_max() {
assert_eq!(
wdl_target_cp(GameResult::BlackWin, Color::Black, 1200.0),
Some(600.0)
);
}
#[test]
fn wdl_target_black_win_from_white_perspective_is_min() {
assert_eq!(
wdl_target_cp(GameResult::BlackWin, Color::White, 1200.0),
Some(-600.0)
);
}
#[test]
fn wdl_target_white_win_from_white_perspective_is_max() {
assert_eq!(
wdl_target_cp(GameResult::WhiteWin, Color::White, 1200.0),
Some(600.0)
);
}
#[test]
fn wdl_target_white_win_from_black_perspective_is_min() {
assert_eq!(
wdl_target_cp(GameResult::WhiteWin, Color::Black, 1200.0),
Some(-600.0)
);
}
#[test]
fn wdl_target_draw_is_zero_regardless_of_perspective() {
assert_eq!(
wdl_target_cp(GameResult::Draw, Color::Black, 1200.0),
Some(0.0)
);
assert_eq!(
wdl_target_cp(GameResult::Draw, Color::White, 1200.0),
Some(0.0)
);
}
#[test]
fn wdl_target_unknown_result_has_no_signal() {
assert_eq!(
wdl_target_cp(GameResult::Unknown, Color::Black, 1200.0),
None
);
assert_eq!(
wdl_target_cp(GameResult::Unknown, Color::White, 1200.0),
None
);
}
#[test]
fn compute_lr_step_half_matches_original_hardcoded_formula() {
assert_eq!(
compute_lr(LrSchedule::StepHalf, 0.001, 0.0, 1, 20, 0),
0.001
);
assert_eq!(
compute_lr(LrSchedule::StepHalf, 0.001, 0.0, 2, 20, 0),
0.0005
);
assert!((compute_lr(LrSchedule::StepHalf, 0.001, 0.0, 3, 20, 0) - 0.00025).abs() < 1e-9);
}
#[test]
fn compute_lr_constant_ignores_epoch() {
assert_eq!(
compute_lr(LrSchedule::Constant, 0.001, 0.0, 1, 20, 0),
0.001
);
assert_eq!(
compute_lr(LrSchedule::Constant, 0.001, 0.0, 20, 20, 0),
0.001
);
}
#[test]
fn compute_lr_min_lr_floors_step_half_too() {
let lr = compute_lr(LrSchedule::StepHalf, 0.001, 0.0001, 20, 20, 0);
assert_eq!(lr, 0.0001);
}
#[test]
fn compute_lr_cosine_starts_at_base_and_ends_at_min_lr_exactly() {
let first = compute_lr(LrSchedule::Cosine, 0.001, 0.00001, 1, 20, 0);
let last = compute_lr(LrSchedule::Cosine, 0.001, 0.00001, 20, 20, 0);
assert!((first - 0.001).abs() < 1e-9);
assert_eq!(last, 0.00001);
}
#[test]
fn compute_lr_warmup_ramps_linearly_and_lands_on_base_lr() {
let half = compute_lr(LrSchedule::StepHalf, 0.001, 0.0, 2, 20, 4);
assert!((half - 0.0005).abs() < 1e-9); let at_boundary = compute_lr(LrSchedule::StepHalf, 0.001, 0.0, 4, 20, 4);
assert!((at_boundary - 0.001).abs() < 1e-9); let first_post_warmup = compute_lr(LrSchedule::StepHalf, 0.001, 0.0, 5, 20, 4);
assert!((first_post_warmup - 0.001).abs() < 1e-9); }
#[test]
fn compute_lr_single_epoch_run_uses_base_lr_for_every_schedule() {
assert_eq!(compute_lr(LrSchedule::StepHalf, 0.001, 0.0, 1, 1, 0), 0.001);
assert_eq!(compute_lr(LrSchedule::Constant, 0.001, 0.0, 1, 1, 0), 0.001);
assert!((compute_lr(LrSchedule::Cosine, 0.001, 0.0, 1, 1, 0) - 0.001).abs() < 1e-9);
}
#[test]
fn compute_lr_warmup_equals_total_epochs_never_panics() {
for epoch in 1..=5u32 {
let lr = compute_lr(LrSchedule::Cosine, 0.001, 0.0, epoch, 5, 5);
assert!(lr.is_finite() && lr >= 0.0);
}
assert_eq!(compute_lr(LrSchedule::Cosine, 0.001, 0.0, 5, 5, 5), 0.001);
}
#[test]
fn compute_lr_short_run_reproduces_epoch3_of_the_real_20_epoch_schedule() {
let lr = compute_lr(LrSchedule::Cosine, 0.001, 0.00001, 3, 20, 1);
assert!(
(lr - 0.000992).abs() < 1e-6,
"epoch3 lr={lr}, expected ~0.000992 (not the min_lr floor 0.00001)"
);
}
#[test]
fn compute_lr_first_3_epochs_of_20_match_hand_computed_prefix() {
let expected = [0.001, 0.001, 0.000992];
for (i, want) in expected.iter().enumerate() {
let epoch = (i + 1) as u32;
let got = compute_lr(LrSchedule::Cosine, 0.001, 0.00001, epoch, 20, 1);
assert!(
(got - want).abs() < 1e-6,
"epoch {epoch}: got {got}, want {want}"
);
}
}
#[test]
fn resolve_schedule_epochs_defaults_to_epochs_when_omitted() {
assert_eq!(resolve_schedule_epochs(3, None, 1).unwrap(), 3);
assert_eq!(resolve_schedule_epochs(20, None, 0).unwrap(), 20);
}
#[test]
fn resolve_schedule_epochs_accepts_a_longer_explicit_horizon() {
assert_eq!(resolve_schedule_epochs(3, Some(20), 1).unwrap(), 20);
}
#[test]
fn resolve_schedule_epochs_rejects_zero() {
assert!(resolve_schedule_epochs(3, Some(0), 0).is_err());
}
#[test]
fn resolve_schedule_epochs_rejects_warmup_exceeding_schedule_epochs() {
assert!(resolve_schedule_epochs(3, Some(5), 6).is_err());
}
#[test]
fn resolve_schedule_epochs_rejects_schedule_epochs_less_than_epochs() {
assert!(resolve_schedule_epochs(20, Some(3), 0).is_err());
}
#[test]
fn resolve_schedule_epochs_epochs_zero_never_errors() {
assert_eq!(resolve_schedule_epochs(0, None, 0).unwrap(), 0);
assert_eq!(resolve_schedule_epochs(0, None, 5).unwrap(), 0);
assert_eq!(resolve_schedule_epochs(0, Some(20), 0).unwrap(), 20);
}
#[test]
fn lr_schedule_parse_roundtrips_known_names_and_rejects_unknown() {
assert_eq!(LrSchedule::parse("constant"), Some(LrSchedule::Constant));
assert_eq!(LrSchedule::parse("step-half"), Some(LrSchedule::StepHalf));
assert_eq!(LrSchedule::parse("cosine"), Some(LrSchedule::Cosine));
assert_eq!(LrSchedule::parse("bogus"), None);
}
#[test]
fn position_teacher_reuses_cached_search_on_repeated_position() {
let mut trainer = Trainer::new(1, 0.5);
let mut cache: HashMap<String, i32> = HashMap::new();
let mut board = Board::startpos();
let (first, _) = trainer.position_teacher_components(
&mut board,
GameResult::Unknown,
2,
&mut cache,
1200.0,
);
assert_eq!(trainer.cache_misses, 1);
assert_eq!(trainer.cache_hits, 0);
assert_eq!(cache.len(), 1);
let mut board_again = Board::startpos();
let (second, _) = trainer.position_teacher_components(
&mut board_again,
GameResult::Unknown,
2,
&mut cache,
1200.0,
);
assert_eq!(trainer.cache_misses, 1, "second call must not re-search");
assert_eq!(trainer.cache_hits, 1);
assert_eq!(cache.len(), 1);
assert_eq!(first, second);
}
#[test]
fn train_position_grad_clip_norm_shrinks_the_applied_update() {
let board = Board::startpos();
let mut unclipped = Trainer::new(1, 0.5);
unclipped.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
let mut clipped = Trainer::new(1, 0.5);
clipped.grad_clip_norm = Some(1.0); clipped.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
assert_eq!(
clipped.grad_clip_count, 1,
"the tiny threshold must trigger"
);
assert_eq!(unclipped.grad_clip_count, 0);
assert!(
clipped.ft_update_norm_sum < unclipped.ft_update_norm_sum,
"clipped={} unclipped={}",
clipped.ft_update_norm_sum,
unclipped.ft_update_norm_sum
);
assert_eq!(
clipped.global_grad_norm_values[0],
unclipped.global_grad_norm_values[0]
);
}
#[test]
fn train_position_out_clip_norm_leaves_ft_and_l2_untouched() {
let board = Board::startpos();
let mut unclipped = Trainer::new(1, 0.5);
unclipped.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
let mut clipped = Trainer::new(1, 0.5);
clipped.out_clip_norm = Some(1.0); clipped.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
assert_eq!(clipped.out_clip_count, 1, "the tiny threshold must trigger");
assert_eq!(clipped.ft_clip_count, 0);
assert_eq!(clipped.l2_clip_count, 0);
assert!(clipped.out_update_norm_sum < unclipped.out_update_norm_sum);
assert_eq!(clipped.ft_update_norm_sum, unclipped.ft_update_norm_sum);
assert_eq!(clipped.l2_update_norm_sum, unclipped.l2_update_norm_sum);
assert_eq!(
clipped.out_grad_norm_values[0],
unclipped.out_grad_norm_values[0]
);
assert!(clipped.out_grad_norm_after_sum < unclipped.out_grad_norm_after_sum);
}
#[test]
fn diagnostic_freeze_layer_unset_is_byte_identical_to_no_freeze() {
let board = Board::startpos();
let mut baseline = Trainer::new(1, 0.5);
baseline.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
let mut untouched = Trainer::new(1, 0.5);
untouched.diagnostic_freeze_until_position = 999; untouched.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
assert_eq!(baseline.weights.ft, untouched.weights.ft);
assert_eq!(baseline.weights.l2, untouched.weights.l2);
assert_eq!(baseline.weights.out, untouched.weights.out);
assert_eq!(baseline.total_loss, untouched.total_loss);
}
#[test]
fn train_position_freeze_layer_l2_leaves_l2_unchanged_but_ft_and_out_still_update() {
let board = Board::startpos();
let mut trainer = Trainer::new(1, 0.5);
trainer.diagnostic_freeze_layer = Some(FreezeLayer::L2);
trainer.diagnostic_freeze_until_position = 10;
let l2_before = trainer.weights.l2.clone();
let l2_bias_before = trainer.weights.l2_bias.clone();
let ft_before = trainer.weights.ft.clone();
let out_before = trainer.weights.out.clone();
trainer.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
assert_eq!(
trainer.weights.l2, l2_before,
"frozen L2 weights must not move"
);
assert_eq!(
trainer.weights.l2_bias, l2_bias_before,
"frozen L2 bias must not move"
);
assert_ne!(
trainer.weights.ft, ft_before,
"FT must still update -- gradient must flow through the frozen L2 weights, not be cut"
);
assert_ne!(
trainer.weights.out, out_before,
"Output must still update normally, unaffected by an L2 freeze"
);
}
#[test]
fn train_position_freeze_layer_ft_leaves_ft_unchanged_but_l2_and_out_still_update() {
let board = Board::startpos();
let mut trainer = Trainer::new(1, 0.5);
trainer.diagnostic_freeze_layer = Some(FreezeLayer::Ft);
trainer.diagnostic_freeze_until_position = 10;
let ft_before = trainer.weights.ft.clone();
let ft_bias_before = trainer.weights.ft_bias.clone();
let l2_before = trainer.weights.l2.clone();
let out_before = trainer.weights.out.clone();
trainer.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
assert_eq!(
trainer.weights.ft, ft_before,
"frozen FT weights must not move"
);
assert_eq!(
trainer.weights.ft_bias, ft_bias_before,
"frozen FT bias must not move"
);
assert_ne!(
trainer.weights.l2, l2_before,
"L2 must still update normally, unaffected by an FT freeze"
);
assert_ne!(
trainer.weights.out, out_before,
"Output must still update normally, unaffected by an FT freeze"
);
}
#[test]
fn train_position_freeze_layer_out_leaves_out_unchanged_but_ft_and_l2_still_update() {
let board = Board::startpos();
let mut trainer = Trainer::new(1, 0.5);
trainer.diagnostic_freeze_layer = Some(FreezeLayer::Out);
trainer.diagnostic_freeze_until_position = 10;
let out_before = trainer.weights.out.clone();
let out_bias_before = trainer.weights.out_bias;
let l2_before = trainer.weights.l2.clone();
let ft_before = trainer.weights.ft.clone();
trainer.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
assert_eq!(
trainer.weights.out, out_before,
"frozen Output weights must not move"
);
assert_eq!(
trainer.weights.out_bias, out_bias_before,
"frozen Output bias must not move"
);
assert_ne!(
trainer.weights.l2, l2_before,
"L2 must still update -- gradient must flow through the frozen Out weights, not be cut"
);
assert_ne!(
trainer.weights.ft, ft_before,
"FT must still update -- gradient must flow all the way through, not be cut"
);
}
#[test]
fn diagnostic_freeze_layer_resumes_updating_once_the_position_bound_is_passed() {
let board = Board::startpos();
let mut trainer = Trainer::new(1, 0.5);
trainer.diagnostic_freeze_layer = Some(FreezeLayer::L2);
trainer.diagnostic_freeze_until_position = 1;
trainer.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
assert_eq!(trainer.l2_sample_count, 1);
let l2_after_frozen_position = trainer.weights.l2.clone();
trainer.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
assert_eq!(trainer.l2_sample_count, 2);
assert_ne!(
trainer.weights.l2, l2_after_frozen_position,
"position 2 is past diagnostic_freeze_until_position=1, L2 must resume updating"
);
}
#[test]
fn diagnostic_freeze_from_position_unset_is_byte_identical_to_no_freeze() {
let board = Board::startpos();
let mut baseline = Trainer::new(1, 0.5);
baseline.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
let mut untouched = Trainer::new(1, 0.5);
untouched.diagnostic_freeze_from_position = 1;
untouched.diagnostic_freeze_until_position = 999; untouched.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
assert_eq!(baseline.weights.ft, untouched.weights.ft);
assert_eq!(baseline.weights.l2, untouched.weights.l2);
assert_eq!(baseline.weights.out, untouched.weights.out);
assert_eq!(baseline.total_loss, untouched.total_loss);
}
#[test]
fn diagnostic_freeze_window_only_freezes_between_from_and_until_positions() {
let board = Board::startpos();
let mut trainer = Trainer::new(1, 0.5);
trainer.diagnostic_freeze_layer = Some(FreezeLayer::Ft);
trainer.diagnostic_freeze_from_position = 2;
trainer.diagnostic_freeze_until_position = 3;
trainer.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
assert_eq!(trainer.l2_sample_count, 1);
let ft_after_position_1 = trainer.weights.ft.clone();
trainer.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
assert_eq!(trainer.l2_sample_count, 2);
assert_eq!(
trainer.weights.ft, ft_after_position_1,
"position 2 is inside [from=2, until=3], FT must stay frozen"
);
let ft_after_position_2 = trainer.weights.ft.clone();
trainer.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
assert_eq!(trainer.l2_sample_count, 3);
assert_eq!(
trainer.weights.ft, ft_after_position_2,
"position 3 is inside [from=2, until=3], FT must stay frozen"
);
trainer.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
assert_eq!(trainer.l2_sample_count, 4);
assert_ne!(
trainer.weights.ft, ft_after_position_2,
"position 4 is past diagnostic_freeze_until_position=3, FT must resume updating"
);
}
#[test]
fn diagnostic_ft_periodic_freeze_cycles_active_and_frozen_sub_blocks() {
let board = Board::startpos();
let mut trainer = Trainer::new(1, 0.5);
trainer.diagnostic_freeze_layer = Some(FreezeLayer::Ft);
trainer.diagnostic_freeze_from_position = 2;
trainer.diagnostic_freeze_until_position = 9;
trainer.diagnostic_ft_active_block = 2;
trainer.diagnostic_ft_frozen_block = 2;
let mut ft_after = Vec::new();
for _ in 1..=9 {
trainer.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
ft_after.push(trainer.weights.ft.clone());
}
assert_ne!(ft_after[1], ft_after[0], "position 2 (active) must update");
assert_ne!(ft_after[2], ft_after[1], "position 3 (active) must update");
assert_eq!(
ft_after[3], ft_after[2],
"position 4 (frozen) must not move"
);
assert_eq!(
ft_after[4], ft_after[2],
"position 5 (frozen) must not move"
);
assert_ne!(ft_after[5], ft_after[4], "position 6 (active) must update");
assert_ne!(ft_after[6], ft_after[5], "position 7 (active) must update");
assert_eq!(
ft_after[7], ft_after[6],
"position 8 (frozen) must not move"
);
assert_eq!(
ft_after[8], ft_after[6],
"position 9 (frozen) must not move"
);
}
#[test]
fn diagnostic_ft_periodic_freeze_unset_blocks_is_byte_identical_to_plain_window_freeze() {
let board = Board::startpos();
let mut plain = Trainer::new(1, 0.5);
plain.diagnostic_freeze_layer = Some(FreezeLayer::Ft);
plain.diagnostic_freeze_from_position = 2;
plain.diagnostic_freeze_until_position = 5;
let mut with_unset_blocks = Trainer::new(1, 0.5);
with_unset_blocks.diagnostic_freeze_layer = Some(FreezeLayer::Ft);
with_unset_blocks.diagnostic_freeze_from_position = 2;
with_unset_blocks.diagnostic_freeze_until_position = 5;
with_unset_blocks.diagnostic_ft_active_block = 0;
with_unset_blocks.diagnostic_ft_frozen_block = 0;
for _ in 1..=6 {
plain.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
with_unset_blocks.train_position(
&board,
-600.0,
1.0,
-600.0,
None,
0,
GameResult::Unknown,
);
assert_eq!(plain.weights.ft, with_unset_blocks.weights.ft);
}
}
#[test]
fn diagnostic_ft_frozen_first_produces_the_exact_complement_pattern() {
let board = Board::startpos();
let mut trainer = Trainer::new(1, 0.5);
trainer.diagnostic_freeze_layer = Some(FreezeLayer::Ft);
trainer.diagnostic_freeze_from_position = 2;
trainer.diagnostic_freeze_until_position = 9;
trainer.diagnostic_ft_active_block = 2;
trainer.diagnostic_ft_frozen_block = 2;
trainer.diagnostic_ft_frozen_first = true;
let mut ft_after = Vec::new();
for _ in 1..=9 {
trainer.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
ft_after.push(trainer.weights.ft.clone());
}
assert_eq!(
ft_after[1], ft_after[0],
"position 2 (frozen, frozen-first) must not move"
);
assert_eq!(
ft_after[2], ft_after[0],
"position 3 (frozen, frozen-first) must not move"
);
assert_ne!(
ft_after[3], ft_after[2],
"position 4 (active, frozen-first) must update"
);
assert_ne!(
ft_after[4], ft_after[3],
"position 5 (active, frozen-first) must update"
);
assert_eq!(
ft_after[5], ft_after[4],
"position 6 (frozen, frozen-first) must not move"
);
assert_eq!(
ft_after[6], ft_after[4],
"position 7 (frozen, frozen-first) must not move"
);
assert_ne!(
ft_after[7], ft_after[6],
"position 8 (active, frozen-first) must update"
);
assert_ne!(
ft_after[8], ft_after[7],
"position 9 (active, frozen-first) must update"
);
}
#[test]
fn diagnostic_ft_frozen_first_unset_is_byte_identical_to_active_first_default() {
let board = Board::startpos();
let mut default_run = Trainer::new(1, 0.5);
default_run.diagnostic_freeze_layer = Some(FreezeLayer::Ft);
default_run.diagnostic_freeze_from_position = 2;
default_run.diagnostic_freeze_until_position = 9;
default_run.diagnostic_ft_active_block = 2;
default_run.diagnostic_ft_frozen_block = 2;
let mut explicit_false = Trainer::new(1, 0.5);
explicit_false.diagnostic_freeze_layer = Some(FreezeLayer::Ft);
explicit_false.diagnostic_freeze_from_position = 2;
explicit_false.diagnostic_freeze_until_position = 9;
explicit_false.diagnostic_ft_active_block = 2;
explicit_false.diagnostic_ft_frozen_block = 2;
explicit_false.diagnostic_ft_frozen_first = false;
for _ in 1..=9 {
default_run.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
explicit_false.train_position(
&board,
-600.0,
1.0,
-600.0,
None,
0,
GameResult::Unknown,
);
assert_eq!(default_run.weights.ft, explicit_false.weights.ft);
}
}
#[test]
fn diagnostic_ft_reactivate_window_reopens_a_single_hole_in_an_otherwise_frozen_span() {
let board = Board::startpos();
let mut trainer = Trainer::new(1, 0.5);
trainer.diagnostic_freeze_layer = Some(FreezeLayer::Ft);
trainer.diagnostic_freeze_from_position = 2;
trainer.diagnostic_freeze_until_position = 9;
trainer.diagnostic_ft_reactivate_from_position = 5;
trainer.diagnostic_ft_reactivate_until_position = 6;
let mut ft_after = Vec::new();
for _ in 1..=9 {
trainer.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
ft_after.push(trainer.weights.ft.clone());
}
assert_eq!(
ft_after[1], ft_after[0],
"position 2 (frozen) must not move"
);
assert_eq!(
ft_after[2], ft_after[0],
"position 3 (frozen) must not move"
);
assert_eq!(
ft_after[3], ft_after[0],
"position 4 (frozen) must not move"
);
assert_ne!(
ft_after[4], ft_after[3],
"position 5 (reactivated) must update"
);
assert_ne!(
ft_after[5], ft_after[4],
"position 6 (reactivated) must update"
);
assert_eq!(
ft_after[6], ft_after[5],
"position 7 (frozen) must not move"
);
assert_eq!(
ft_after[7], ft_after[5],
"position 8 (frozen) must not move"
);
assert_eq!(
ft_after[8], ft_after[5],
"position 9 (frozen) must not move"
);
}
#[test]
fn diagnostic_ft_reactivate_window_unset_is_byte_identical_to_plain_window_freeze() {
let board = Board::startpos();
let mut plain = Trainer::new(1, 0.5);
plain.diagnostic_freeze_layer = Some(FreezeLayer::Ft);
plain.diagnostic_freeze_from_position = 2;
plain.diagnostic_freeze_until_position = 9;
let mut with_unset_reactivate = Trainer::new(1, 0.5);
with_unset_reactivate.diagnostic_freeze_layer = Some(FreezeLayer::Ft);
with_unset_reactivate.diagnostic_freeze_from_position = 2;
with_unset_reactivate.diagnostic_freeze_until_position = 9;
with_unset_reactivate.diagnostic_ft_reactivate_from_position = 0;
with_unset_reactivate.diagnostic_ft_reactivate_until_position = 0;
for _ in 1..=9 {
plain.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
with_unset_reactivate.train_position(
&board,
-600.0,
1.0,
-600.0,
None,
0,
GameResult::Unknown,
);
assert_eq!(plain.weights.ft, with_unset_reactivate.weights.ft);
}
}
#[test]
fn diagnostic_ft_reactivate2_window_reopens_a_second_disjoint_hole() {
let board = Board::startpos();
let mut trainer = Trainer::new(1, 0.5);
trainer.diagnostic_freeze_layer = Some(FreezeLayer::Ft);
trainer.diagnostic_freeze_from_position = 2;
trainer.diagnostic_freeze_until_position = 13;
trainer.diagnostic_ft_reactivate_from_position = 4;
trainer.diagnostic_ft_reactivate_until_position = 5;
trainer.diagnostic_ft_reactivate2_from_position = 10;
trainer.diagnostic_ft_reactivate2_until_position = 11;
let mut ft_after = Vec::new();
for _ in 1..=13 {
trainer.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
ft_after.push(trainer.weights.ft.clone());
}
let active_positions = [4, 5, 10, 11];
for pos in 2..=13u64 {
let idx = (pos - 1) as usize;
let prev_idx = idx - 1;
if active_positions.contains(&pos) {
assert_ne!(
ft_after[idx], ft_after[prev_idx],
"position {pos} (reactivated) must update"
);
} else {
assert_eq!(
ft_after[idx], ft_after[prev_idx],
"position {pos} (frozen) must not move"
);
}
}
}
#[test]
fn diagnostic_ft_reactivate2_window_unset_is_byte_identical_to_single_reactivate_window() {
let board = Board::startpos();
let mut single_window = Trainer::new(1, 0.5);
single_window.diagnostic_freeze_layer = Some(FreezeLayer::Ft);
single_window.diagnostic_freeze_from_position = 2;
single_window.diagnostic_freeze_until_position = 9;
single_window.diagnostic_ft_reactivate_from_position = 5;
single_window.diagnostic_ft_reactivate_until_position = 6;
let mut with_unset_window2 = Trainer::new(1, 0.5);
with_unset_window2.diagnostic_freeze_layer = Some(FreezeLayer::Ft);
with_unset_window2.diagnostic_freeze_from_position = 2;
with_unset_window2.diagnostic_freeze_until_position = 9;
with_unset_window2.diagnostic_ft_reactivate_from_position = 5;
with_unset_window2.diagnostic_ft_reactivate_until_position = 6;
with_unset_window2.diagnostic_ft_reactivate2_from_position = 0;
with_unset_window2.diagnostic_ft_reactivate2_until_position = 0;
for _ in 1..=9 {
single_window.train_position(&board, -600.0, 1.0, -600.0, None, 0, GameResult::Unknown);
with_unset_window2.train_position(
&board,
-600.0,
1.0,
-600.0,
None,
0,
GameResult::Unknown,
);
assert_eq!(single_window.weights.ft, with_unset_window2.weights.ft);
}
}
#[test]
fn replay_override_unset_leaves_teacher_and_weight_unchanged() {
let trainer = Trainer::new(1, 0.5);
let (teacher, weight) = trainer.replay_override(5, Some(0.7), 100.0, Some(50.0), 1.0, 85.0);
assert_eq!(teacher, 85.0);
assert_eq!(weight, 1.0);
}
#[test]
fn replay_override_cp_component_uses_eval_teacher_scaled_by_lambda() {
let mut trainer = Trainer::new(1, 0.5);
trainer.diagnostic_replay_component = Some(ReplayComponent::Cp);
trainer.diagnostic_replay_from_position = 10;
trainer.diagnostic_replay_until_position = 20;
let (teacher, weight) =
trainer.replay_override(15, Some(0.7), 100.0, Some(50.0), 2.0, 85.0);
assert_eq!(teacher, 100.0);
assert!((weight - 2.0 * 0.7).abs() < 1e-6);
}
#[test]
fn replay_override_wdl_component_uses_wdl_target_scaled_by_one_minus_lambda() {
let mut trainer = Trainer::new(1, 0.5);
trainer.diagnostic_replay_component = Some(ReplayComponent::Wdl);
trainer.diagnostic_replay_from_position = 10;
trainer.diagnostic_replay_until_position = 20;
let (teacher, weight) =
trainer.replay_override(15, Some(0.7), 100.0, Some(50.0), 2.0, 85.0);
assert_eq!(teacher, 50.0);
assert!((weight - 2.0 * 0.3).abs() < 1e-6);
}
#[test]
fn replay_override_out_of_window_leaves_teacher_and_weight_unchanged() {
let mut trainer = Trainer::new(1, 0.5);
trainer.diagnostic_replay_component = Some(ReplayComponent::Cp);
trainer.diagnostic_replay_from_position = 10;
trainer.diagnostic_replay_until_position = 20;
let (before, _) = trainer.replay_override(9, Some(0.7), 100.0, Some(50.0), 1.0, 85.0);
let (after, _) = trainer.replay_override(21, Some(0.7), 100.0, Some(50.0), 1.0, 85.0);
assert_eq!(before, 85.0);
assert_eq!(after, 85.0);
}
#[test]
fn train_position_wdl_component_only_accumulates_when_target_present() {
let mut trainer = Trainer::new(1, 0.5);
let board = Board::startpos();
trainer.train_position(&board, 10.0, 1.0, 10.0, None, 0, GameResult::Unknown);
assert_eq!(trainer.wdl_component_count, 0);
assert_eq!(trainer.wdl_component_sum, 0.0);
trainer.train_position(&board, 5.0, 1.0, 20.0, Some(-30.0), 0, GameResult::Unknown);
assert_eq!(trainer.wdl_component_count, 1);
assert!(trainer.wdl_component_sum > 0.0);
assert!(trainer.cp_component_sum > 0.0);
}
#[test]
fn train_position_records_exactly_one_grad_norm_sample_per_call() {
let mut trainer = Trainer::new(1, 0.5);
let board = Board::startpos();
trainer.train_position(&board, 10.0, 1.0, 10.0, None, 0, GameResult::Unknown);
trainer.train_position(&board, 10.0, 1.0, 10.0, None, 0, GameResult::Unknown);
assert_eq!(trainer.global_grad_norm_values.len(), 2);
assert!(
trainer
.global_grad_norm_values
.iter()
.all(|&g| g >= 0.0 && g.is_finite())
);
assert!(trainer.ft_grad_norm_sum_sq >= 0.0);
}
#[test]
fn trace_positions_snapshots_exactly_the_requested_points() {
let mut trainer = Trainer::new(1, 0.5);
trainer.trace_positions = [2u64, 5].into_iter().collect();
let board = Board::startpos();
for _ in 0..6 {
trainer.train_position(&board, 10.0, 1.0, 10.0, None, 0, GameResult::Unknown);
}
let indices: Vec<u64> = trainer
.trace_snapshots
.iter()
.map(|s| s.position_index)
.collect();
assert_eq!(indices, vec![2, 5]);
for snapshot in &trainer.trace_snapshots {
assert_eq!(snapshot.l2.bias.len(), L2);
assert_eq!(snapshot.ft.bias.len(), L1);
}
let mut trainer2 = Trainer::new(1, 0.5);
trainer2.trace_positions = [0u64].into_iter().collect();
trainer2.train_position(&board, 10.0, 1.0, 10.0, None, 0, GameResult::Unknown);
assert!(trainer2.trace_snapshots.is_empty());
}
#[test]
fn trace_positions_omitted_writes_no_snapshots() {
let mut trainer = Trainer::new(1, 0.5);
assert!(trainer.trace_positions.is_empty());
let board = Board::startpos();
for _ in 0..10 {
trainer.train_position(&board, 10.0, 1.0, 10.0, None, 0, GameResult::Unknown);
}
assert!(trainer.trace_snapshots.is_empty());
}
#[test]
fn shuffled_order_is_a_permutation() {
let order = shuffled_order(500, 42);
let mut sorted = order.clone();
sorted.sort_unstable();
assert_eq!(sorted, (0..500).collect::<Vec<_>>());
}
#[test]
fn shuffled_order_is_deterministic_for_the_same_seed() {
assert_eq!(shuffled_order(200, 7), shuffled_order(200, 7));
}
#[test]
fn shuffled_order_differs_across_seeds() {
assert_ne!(shuffled_order(200, 1), shuffled_order(200, 2));
}
#[test]
fn shuffled_order_handles_zero_and_one() {
assert_eq!(shuffled_order(0, 42), Vec::<usize>::new());
assert_eq!(shuffled_order(1, 42), vec![0]);
}
#[test]
fn l2_preactivation_gradient_matches_doutput_times_out_weight_times_clippedrelu_derivative() {
let mut trainer = Trainer::new(3, 0.5);
let board = Board::startpos();
let teacher = 42.0;
let weight = 1.0;
let score_before = trainer.forward(&board);
let out_before = trainer.weights.out.clone();
let d_output_expected = weight * 2.0 * (score_before - teacher) / 64.0;
trainer.train_position(
&board,
teacher,
weight,
teacher,
None,
0,
GameResult::Unknown,
);
for o in 0..L2 {
let actual = trainer.l2_dacc_sum[o];
let unclamped = (d_output_expected * out_before[o]) as f64;
let matches_unclamped = (actual - unclamped).abs() < 1e-4;
let is_zero = actual == 0.0;
assert!(
matches_unclamped || is_zero,
"neuron {o}: l2_dacc_sum={actual} does not match d_output*out_weight={unclamped} and isn't 0 (dead/saturated)"
);
}
}
#[test]
fn cp_wdl_grad_trace_blended_gradient_is_the_expected_weighted_sum() {
let mut trainer = Trainer::new(1, 0.5);
trainer.cp_wdl_grad_trace = true;
trainer.trace_positions = [3u64].into_iter().collect();
let board = Board::startpos();
let lambda = 0.7f32;
let eval_teacher = 40.0f32;
let wdl_target = -120.0f32;
let teacher = lambda * eval_teacher + (1.0 - lambda) * wdl_target;
for _ in 0..3 {
trainer.train_position(
&board,
teacher,
1.0,
eval_teacher,
Some(wdl_target),
0,
GameResult::Unknown,
);
}
let cp_wdl = trainer.trace_snapshots[0]
.cp_wdl
.as_ref()
.expect("cp_wdl populated");
assert!((cp_wdl.cp_target_mean - eval_teacher as f64).abs() < 1e-6);
assert!((cp_wdl.wdl_target_mean - wdl_target as f64).abs() < 1e-6);
assert!(cp_wdl.prediction_mean.is_finite());
assert!(cp_wdl.cp_residual_std.is_finite() && cp_wdl.cp_residual_std >= 0.0);
assert!(cp_wdl.wdl_residual_std.is_finite() && cp_wdl.wdl_residual_std >= 0.0);
assert!(cp_wdl.cp_d_output_mean.is_finite());
assert!(cp_wdl.wdl_d_output_mean.is_finite());
for o in 0..L2 {
let expected = lambda as f64 * trainer.l2_cp_dacc_sum[o]
+ (1.0 - lambda) as f64 * trainer.l2_wdl_dacc_sum[o];
assert!(
(trainer.l2_dacc_sum[o] - expected).abs() < 1e-3,
"l2 neuron {o}: blended={} expected={}",
trainer.l2_dacc_sum[o],
expected
);
}
for j in 0..L1 {
let expected = lambda as f64 * trainer.ft_cp_dacc_sum[j]
+ (1.0 - lambda) as f64 * trainer.ft_wdl_dacc_sum[j];
assert!(
(trainer.ft_dacc_sum[j] - expected).abs() < 1e-3,
"ft neuron {j}: blended={} expected={}",
trainer.ft_dacc_sum[j],
expected
);
}
}
#[test]
fn cp_wdl_grad_trace_does_not_alter_training_state() {
let board = Board::startpos();
let lambda = 0.7f32;
let eval_teacher = 40.0f32;
let wdl_target = -120.0f32;
let teacher = lambda * eval_teacher + (1.0 - lambda) * wdl_target;
let mut plain = Trainer::new(1, 0.5);
let mut traced = Trainer::new(1, 0.5);
traced.cp_wdl_grad_trace = true;
for _ in 0..5 {
plain.train_position(
&board,
teacher,
1.0,
eval_teacher,
Some(wdl_target),
0,
GameResult::Unknown,
);
traced.train_position(
&board,
teacher,
1.0,
eval_teacher,
Some(wdl_target),
0,
GameResult::Unknown,
);
}
assert_eq!(
plain.weights.snapshot_params(),
traced.weights.snapshot_params()
);
assert_eq!(plain.total_loss, traced.total_loss);
assert_eq!(plain.l2_dacc_sum, traced.l2_dacc_sum);
assert_eq!(plain.ft_grad_norm_sum, traced.ft_grad_norm_sum);
}
#[test]
fn shadow_trace_active_run_is_byte_identical_to_inactive() {
let board = Board::startpos();
let lambda = 0.7f32;
let eval_teacher = 40.0f32;
let wdl_target = -120.0f32;
let teacher = lambda * eval_teacher + (1.0 - lambda) * wdl_target;
let mut plain = Trainer::new(1, 0.5);
let mut traced = Trainer::new(1, 0.5);
traced.diagnostic_shadow_trace_from_position = 1;
traced.diagnostic_shadow_trace_until_position = 5;
traced.diagnostic_shadow_trace_wdl_lambda = lambda;
traced.diagnostic_shadow_trace_probe_boards = vec![Board::startpos()];
for _ in 0..5 {
plain.train_position(
&board,
teacher,
1.0,
eval_teacher,
Some(wdl_target),
0,
GameResult::Unknown,
);
traced.train_position(
&board,
teacher,
1.0,
eval_teacher,
Some(wdl_target),
0,
GameResult::Unknown,
);
}
assert_eq!(
plain.weights.snapshot_params(),
traced.weights.snapshot_params()
);
assert_eq!(traced.shadow_trace_records.len(), 5);
}
#[test]
fn shadow_trace_unset_probe_boards_records_nothing() {
let mut trainer = Trainer::new(1, 0.5);
trainer.diagnostic_shadow_trace_from_position = 1;
trainer.diagnostic_shadow_trace_until_position = 5;
trainer.diagnostic_shadow_trace_wdl_lambda = 0.7;
let board = Board::startpos();
trainer.train_position(&board, 4.0, 1.0, 40.0, Some(-120.0), 0, GameResult::Unknown);
assert!(trainer.shadow_trace_records.is_empty());
}
#[test]
fn shadow_trace_records_pass_the_blend_correctness_guard_and_a_full_contingency() {
let mut trainer = Trainer::new(1, 0.5);
trainer.diagnostic_shadow_trace_from_position = 1;
trainer.diagnostic_shadow_trace_until_position = 1;
trainer.diagnostic_shadow_trace_wdl_lambda = 0.7;
trainer.diagnostic_shadow_trace_probe_boards = vec![Board::startpos()];
let board = Board::startpos();
trainer.train_position(&board, 4.0, 1.0, 40.0, Some(-120.0), 0, GameResult::Unknown);
assert_eq!(trainer.shadow_trace_records.len(), 1);
let record = &trainer.shadow_trace_records[0];
assert!(record.blend_matches_real_ft);
assert!(record.blend_matches_real_l2);
assert_eq!(
record.contingency_cp_wdl_blend.iter().sum::<u64>(),
record.n_alive_at_anchor
);
assert_eq!(
record.blend_dead_linpred_alive
+ record.blend_dead_linpred_dead
+ record.blend_alive_linpred_dead
+ record.blend_alive_linpred_alive,
record.n_alive_at_anchor
);
}
#[test]
fn conflict_mask_unset_is_byte_identical_to_no_mask() {
let board = Board::startpos();
let teacher = 0.7 * 40.0 + 0.3 * (-120.0);
let mut plain = Trainer::new(1, 0.5);
let mut masked = Trainer::new(1, 0.5); for _ in 0..5 {
plain.train_position(
&board,
teacher,
1.0,
40.0,
Some(-120.0),
0,
GameResult::Unknown,
);
masked.train_position(
&board,
teacher,
1.0,
40.0,
Some(-120.0),
0,
GameResult::Unknown,
);
}
assert_eq!(
plain.weights.snapshot_params(),
masked.weights.snapshot_params()
);
assert_eq!(masked.masked_position_count, 0);
}
#[test]
fn conflict_mask_ft_zeroes_ft_update_only_at_a_guaranteed_conflicting_position() {
let mut control = Trainer::new(1, 0.5);
let mut masked = Trainer::new(1, 0.5);
masked.diagnostic_conflict_mask = Some(ConflictMaskLayer::Ft);
let board = Board::startpos();
let eval_teacher = -1.0e9;
let wdl_target = 1.0e9;
let teacher = 0.5 * eval_teacher + 0.5 * wdl_target;
control.train_position(
&board,
teacher,
1.0,
eval_teacher,
Some(wdl_target),
0,
GameResult::Unknown,
);
masked.train_position(
&board,
teacher,
1.0,
eval_teacher,
Some(wdl_target),
0,
GameResult::Unknown,
);
assert_eq!(masked.weights.ft, Trainer::new(1, 0.5).weights.ft);
assert_eq!(masked.weights.ft_bias, Trainer::new(1, 0.5).weights.ft_bias);
assert_eq!(masked.weights.l2, control.weights.l2);
assert_eq!(masked.weights.out, control.weights.out);
assert_eq!(masked.masked_position_count, 1);
assert_eq!(masked.conflict_group.count, 1);
assert_eq!(masked.nonconflict_group.count, 0);
}
#[test]
fn conflict_mask_ft_and_l2_zeroes_both_at_a_guaranteed_conflicting_position() {
let mut control = Trainer::new(1, 0.5);
let mut masked = Trainer::new(1, 0.5);
masked.diagnostic_conflict_mask = Some(ConflictMaskLayer::FtAndL2);
let board = Board::startpos();
let eval_teacher = -1.0e9;
let wdl_target = 1.0e9;
let teacher = 0.5 * eval_teacher + 0.5 * wdl_target;
control.train_position(
&board,
teacher,
1.0,
eval_teacher,
Some(wdl_target),
0,
GameResult::Unknown,
);
masked.train_position(
&board,
teacher,
1.0,
eval_teacher,
Some(wdl_target),
0,
GameResult::Unknown,
);
let fresh = Trainer::new(1, 0.5);
assert_eq!(masked.weights.ft, fresh.weights.ft);
assert_eq!(masked.weights.l2, fresh.weights.l2);
assert_eq!(masked.weights.out, control.weights.out);
}
#[test]
fn conflict_mask_ft_passes_through_unchanged_at_a_guaranteed_nonconflicting_position() {
let mut control = Trainer::new(1, 0.5);
let mut masked = Trainer::new(1, 0.5);
masked.diagnostic_conflict_mask = Some(ConflictMaskLayer::Ft);
let board = Board::startpos();
let eval_teacher = -1.0e9;
let wdl_target = -2.0e9;
let teacher = 0.5 * eval_teacher + 0.5 * wdl_target;
control.train_position(
&board,
teacher,
1.0,
eval_teacher,
Some(wdl_target),
0,
GameResult::Unknown,
);
masked.train_position(
&board,
teacher,
1.0,
eval_teacher,
Some(wdl_target),
0,
GameResult::Unknown,
);
assert_eq!(
control.weights.snapshot_params(),
masked.weights.snapshot_params()
);
assert_eq!(masked.masked_position_count, 0);
assert_eq!(masked.conflict_group.count, 0);
assert_eq!(masked.nonconflict_group.count, 1);
}
#[test]
fn rate_matched_mask_selects_exactly_k_of_n_and_is_deterministic() {
let mut trainer = Trainer::new(1, 0.5);
trainer.diagnostic_rate_matched_mask_count = 17;
trainer.diagnostic_rate_matched_mask_total = 200;
trainer.diagnostic_rate_matched_mask_seed = 42;
trainer.rate_matched_remaining_needed = trainer.diagnostic_rate_matched_mask_count;
trainer.rate_matched_remaining_pool = trainer.diagnostic_rate_matched_mask_total;
trainer.rate_matched_rng =
Lcg(trainer.diagnostic_rate_matched_mask_seed ^ 0xA5A5_5A5A_1234_5678);
let selected: Vec<bool> = (0..200)
.map(|_| trainer.rate_matched_should_mask())
.collect();
assert_eq!(selected.iter().filter(|&&s| s).count(), 17);
assert!(!trainer.rate_matched_should_mask());
let mut again = Trainer::new(1, 0.5);
again.diagnostic_rate_matched_mask_count = 17;
again.diagnostic_rate_matched_mask_total = 200;
again.diagnostic_rate_matched_mask_seed = 42;
again.rate_matched_remaining_needed = again.diagnostic_rate_matched_mask_count;
again.rate_matched_remaining_pool = again.diagnostic_rate_matched_mask_total;
again.rate_matched_rng =
Lcg(again.diagnostic_rate_matched_mask_seed ^ 0xA5A5_5A5A_1234_5678);
let selected_again: Vec<bool> =
(0..200).map(|_| again.rate_matched_should_mask()).collect();
assert_eq!(selected, selected_again);
}
#[test]
fn sample_grad_trace_records_expected_fields_and_cosine_semantics() {
let mut trainer = Trainer::new(1, 0.5);
trainer.sample_grad_trace_limit = 3;
let board = Board::startpos();
trainer.train_position(&board, 10.0, 1.0, 10.0, None, 42, GameResult::WhiteWin);
trainer.train_position(&board, 10.0, 1.0, 10.0, None, 42, GameResult::WhiteWin);
assert_eq!(trainer.sample_grad_records.len(), 2);
let first = &trainer.sample_grad_records[0];
let second = &trainer.sample_grad_records[1];
assert_eq!(first.game_id, 42);
assert_eq!(first.game_result, "WhiteWin");
assert_eq!(first.position_index, 1);
assert_eq!(second.position_index, 2);
assert_eq!(first.cosine_prev, None);
assert_eq!(first.cosine_running_mean, None);
assert!(second.cosine_prev.is_some());
assert!(second.cosine_running_mean.is_some());
assert_eq!(first.l2_gate.len(), L2);
}
#[test]
fn sample_grad_trace_does_not_alter_training_state() {
let board = Board::startpos();
let lambda = 0.7f32;
let eval_teacher = 40.0f32;
let wdl_target = -120.0f32;
let teacher = lambda * eval_teacher + (1.0 - lambda) * wdl_target;
let mut plain = Trainer::new(1, 0.5);
let mut traced = Trainer::new(1, 0.5);
traced.sample_grad_trace_limit = 5;
for i in 0..8 {
plain.train_position(
&board,
teacher,
1.0,
eval_teacher,
Some(wdl_target),
7,
GameResult::BlackWin,
);
traced.train_position(
&board,
teacher,
1.0,
eval_teacher,
Some(wdl_target),
7,
GameResult::BlackWin,
);
let expected_records = (i + 1).min(5);
assert_eq!(traced.sample_grad_records.len(), expected_records);
}
assert_eq!(
plain.weights.snapshot_params(),
traced.weights.snapshot_params()
);
assert_eq!(plain.total_loss, traced.total_loss);
assert_eq!(plain.l2_dacc_sum, traced.l2_dacc_sum);
assert_eq!(plain.ft_grad_norm_sum, traced.ft_grad_norm_sum);
}
}