Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 10 additions & 11 deletions bin/trainer.rs
Original file line number Diff line number Diff line change
Expand Up @@ -59,7 +59,7 @@ impl SparseInputType for Features {
type RequiredDataType = ChessBoard;

fn num_inputs(&self) -> usize {
Self::LEN + 768
Self::LEN + PSQFeature::LEN
}

fn max_active(&self) -> usize {
Expand Down Expand Up @@ -93,7 +93,7 @@ impl SparseInputType for Features {
let mut remaining = Bitboard::from(pawns);
for s in remaining.iter() {
remaining &= !s.bitboard();
for t in PPFeature::MASK[s.file()].bitand(remaining).iter() {
for t in PPFeature::WINDOW[s.file()].bitand(remaining).iter() {
let pft1 = pfts.map(|ft| Num::new(ft[s]));
let pft2 = pfts.map(|ft| Num::new(ft[t]));

Expand Down Expand Up @@ -158,7 +158,7 @@ impl SparseInputType for Features {
}

fn shorthand(&self) -> String {
format!("768x{}ti", Self::BUCKETS)
format!("768x{}ti-pp", Self::BUCKETS)
}

fn description(&self) -> String {
Expand Down Expand Up @@ -487,7 +487,7 @@ impl Orchestrator {
weights
})
.round()
.quantise::<i16>(FTQ),
.quantise::<i8>(FTQ),
SavedFormat::id("ftw")
.transform(|_, weights| {
let mut merged = Features.merge_factoriser(weights);
Expand Down Expand Up @@ -550,19 +550,18 @@ impl Orchestrator {
(out, err * err + 0.15 * err_relu * err_relu + 0.005 * l1_reg)
});

let params = AdamWParams {
trainer.optimiser.set_params(AdamWParams {
min_weight: f32::MIN,
max_weight: f32::MAX,
..AdamWParams::default()
};
});

trainer.optimiser.set_params(params);
trainer.optimiser.set_params_for_weight(
"ftw",
AdamWParams {
min_weight: -1.0,
max_weight: 1.0,
..params
min_weight: i8::MIN as f32 / FTQ as f32,
max_weight: i8::MAX as f32 / FTQ as f32,
..AdamWParams::default()
},
);

Expand All @@ -571,7 +570,7 @@ impl Orchestrator {
AdamWParams {
min_weight: i8::MIN as f32 / HLQ as f32,
max_weight: i8::MAX as f32 / HLQ as f32,
..params
..AdamWParams::default()
},
);

Expand Down
14 changes: 10 additions & 4 deletions lib/nnue.rs
Original file line number Diff line number Diff line change
Expand Up @@ -25,10 +25,16 @@ pub use synapse::*;
pub use transformer::*;

/// Quantization scale for the feature transformer.
pub const FTQ: i16 = 255;
pub const FTQ: i16 = 127;

/// Quantization scale for the hidden layers.
pub const HLQ: i16 = 64;
pub const HLQ: i16 = 75;

/// Conversion factor from quantized to floating point.
pub const I2F: f32 = (1 << 7) as f32 / (FTQ as f32 * FTQ as f32 * HLQ as f32);

/// Eval scale.
pub const F2V: f32 = 75.0;

const unsafe fn copy_bytes<T>(dst: &mut T, src: &[u8]) -> usize {
let len = size_of_val(dst);
Expand Down Expand Up @@ -187,8 +193,8 @@ mod tests {
let (mut lower, mut upper) = (bias, bias);

let mut ka = Vec::from_iter(transformer.ka.iter().map(|a| a[i]));
let mut ti = Vec::from_iter(transformer.ti.iter().map(|a| a[i]));
let mut pp = Vec::from_iter(transformer.ti.iter().map(|a| a[i]));
let mut ti = Vec::from_iter(transformer.ti.iter().map(|a| a[i] as i16));
let mut pp = Vec::from_iter(transformer.ti.iter().map(|a| a[i] as i16));

for (n, ws) in [(32, &mut ka), (64, &mut ti), (24, &mut pp)] {
let (small, _, _) = ws.select_nth_unstable(n);
Expand Down
131 changes: 74 additions & 57 deletions lib/nnue/evaluator.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@ use crate::{chess::*, nnue::*, params::Params, search::Ply, simd::*};
use bytemuck::Zeroable;
use derive_more::with_trait::{Debug, Deref};
use std::hash::{Hash, Hasher};
use std::ops::{Index, Range};
use std::ops::{BitAnd, Index, Range};
use std::{array, mem::replace, str::FromStr};

#[cfg(test)]
Expand Down Expand Up @@ -341,17 +341,13 @@ impl Evaluator {
let bucket = KingBucket::new(side, ksq);
let cache = &mut self.cache[side][bucket];

let diff = accumulate_ka(side, ksq, &cache.attacks, new, |sub, add| {
Nnue::transformer().accumulate_ka_in_place(&mut cache.accumulator, sub, add);
});

if !diff.is_empty() {
if accumulate_ka_in_place(side, ksq, &cache.attacks, new, &mut cache.accumulator).any() {
let old = replace(&mut cache.attacks, Attacks::new(new));
accumulate_ti(side, ksq, &old, &cache.attacks, |sub, add| {
accumulate_ti_in_place(side, ksq, &old, &cache.attacks, |sub, add| {
Nnue::transformer().accumulate_ti_in_place(&mut cache.accumulator, sub, add);
});

accumulate_pp(side, ksq, &old, &cache.attacks, |sub, add| {
accumulate_pp_in_place(side, ksq, &old, &cache.attacks, |sub, add| {
Nnue::transformer().accumulate_pp_in_place(&mut cache.accumulator, sub, add);
});
}
Expand All @@ -367,22 +363,17 @@ impl Evaluator {
(ply > 0).assume();
let new = &self.positions[ply];
let old = &self.positions[ply - 1];
let ksq = new.king(side);
let (left, right) = self.accumulator[side].split_at_mut(ply.cast());
let (src, dst) = (&left[left.len() - 1], &mut right[0]);
let diff = accumulate_ka(side, ksq, old, new, |sub, add| {
Nnue::transformer().accumulate_ka(src, dst, sub, add);
});

if diff.is_empty() {
*dst = *src;
} else {
let ksq = new.king(side);
if accumulate_ka(side, ksq, old, new, src, dst).any() {
let (old, new) = (Attacks::new(old), Attacks::new(new));
accumulate_ti(side, ksq, &old, &new, |sub, add| {
accumulate_ti_in_place(side, ksq, &old, &new, |sub, add| {
Nnue::transformer().accumulate_ti_in_place(dst, sub, add);
});

accumulate_pp(side, ksq, &old, &new, |sub, add| {
accumulate_pp_in_place(side, ksq, &old, &new, |sub, add| {
Nnue::transformer().accumulate_pp_in_place(dst, sub, add);
});
}
Expand All @@ -391,44 +382,75 @@ impl Evaluator {

#[inline(always)]
#[cfg_attr(feature = "no_panic", no_panic::no_panic)]
fn accumulate_ka<F>(
fn accumulate_ka(
side: Color,
ksq: Square,
old: &Placement,
new: &Placement,
mut acc: F,
) -> Bitboard
where
F: FnMut([Option<KAFeature>; 2], [Option<KAFeature>; 2]),
{
src: &Accumulator,
dst: &mut Accumulator,
) -> M8x64 {
let diff: M8x64 = old.pieces().simd_ne(new.pieces()).into();

let kafts = KAFeature::lut(side, ksq, old).to_array();
let to_sub = Bitboard::from(diff & old.occupied());
let mut to_sub = to_sub.iter().map(|sq| Num::new(kafts[sq]));
if !diff.any() {
*dst = *src;
} else {
let kafts_to_sub = KAFeature::lut(side, ksq, old).to_array();
let mut to_sub = Bitboard::from(diff & old.occupied()).iter();
(1..=2).contains(&to_sub.len()).assume();

let kafts = KAFeature::lut(side, ksq, new).to_array();
let to_add = Bitboard::from(diff & new.occupied());
let mut to_add = to_add.iter().map(|sq| Num::new(kafts[sq]));
let kafts_to_add = KAFeature::lut(side, ksq, new).to_array();
let mut to_add = Bitboard::from(diff & new.occupied()).iter();
(1..=2).contains(&to_add.len()).assume();

loop {
let sub = array::from_fn(|_| to_sub.next());
let add = array::from_fn(|_| to_add.next());
if sub != [None; 2] || add != [None; 2] {
acc(sub, add);
} else {
break;
let sub = array::from_fn(|_| Some(Num::new(kafts_to_sub[to_sub.next()?])));
let add = array::from_fn(|_| Some(Num::new(kafts_to_add[to_add.next()?])));

Nnue::transformer().accumulate_ka(src, dst, sub, add);
}

diff
}

#[inline(always)]
#[cfg_attr(feature = "no_panic", no_panic::no_panic)]
fn accumulate_ka_in_place(
side: Color,
ksq: Square,
old: &Placement,
new: &Placement,
dst: &mut Accumulator,
) -> M8x64 {
let diff: M8x64 = old.pieces().simd_ne(new.pieces()).into();

if diff.any() {
let kafts_to_sub = KAFeature::lut(side, ksq, old).to_array();
let to_sub = Bitboard::from(diff & old.occupied());

let kafts_to_add = KAFeature::lut(side, ksq, new).to_array();
let to_add = Bitboard::from(diff & new.occupied());

let mut to_sub = to_sub.iter().map(|sq| Num::new(kafts_to_sub[sq]));
let mut to_add = to_add.iter().map(|sq| Num::new(kafts_to_add[sq]));

loop {
let (sub, add) = (to_sub.next(), to_add.next());
if sub.is_some() || add.is_some() {
Nnue::transformer().accumulate_ka_in_place(dst, sub, add);
} else {
break;
}
}
}

diff.into()
diff
}

#[inline(always)]
#[cfg_attr(feature = "no_panic", no_panic::no_panic)]
fn accumulate_ti<F>(side: Color, ksq: Square, old: &Attacks, new: &Attacks, mut acc: F)
fn accumulate_ti_in_place<F>(side: Color, ksq: Square, old: &Attacks, new: &Attacks, mut acc: F)
where
F: FnMut([Option<TIFeature>; 2], [Option<TIFeature>; 2]),
F: FnMut(Option<TIFeature>, Option<TIFeature>),
{
let captured = old.occupied() & new.occupied() & old.pieces().simd_ne(new.pieces());

Expand Down Expand Up @@ -465,9 +487,8 @@ where
});

loop {
let sub = array::from_fn(|_| to_sub.next());
let add = array::from_fn(|_| to_add.next());
if sub != [None; 2] || add != [None; 2] {
let (sub, add) = (to_sub.next(), to_add.next());
if sub.is_some() || add.is_some() {
acc(sub, add);
} else {
break;
Expand All @@ -478,47 +499,43 @@ where

#[inline(always)]
#[cfg_attr(feature = "no_panic", no_panic::no_panic)]
fn accumulate_pp<F>(side: Color, ksq: Square, old: &Placement, new: &Placement, mut acc: F)
fn accumulate_pp_in_place<F>(side: Color, ksq: Square, old: &Placement, new: &Placement, mut acc: F)
where
F: FnMut([Option<PPFeature>; 2], [Option<PPFeature>; 2]),
F: FnMut(Option<PPFeature>, Option<PPFeature>),
{
let old_white_pawns = old.by_piece(Piece::WhitePawn);
let old_black_pawns = old.by_piece(Piece::BlackPawn);
let new_white_pawns = new.by_piece(Piece::WhitePawn);
let new_black_pawns: M<i8, 64> = new.by_piece(Piece::BlackPawn);

let old_pawns = old_white_pawns | old_black_pawns;
let new_pawns = new_white_pawns | new_black_pawns;
let new_black_pawns = new.by_piece(Piece::BlackPawn);
let diff = (old_white_pawns ^ new_white_pawns) | (old_black_pawns ^ new_black_pawns);

let pfts = PFeature::lut(side, ksq, old).to_array();
let mut remaining = Bitboard::from(old_pawns);
let mut to_sub = Bitboard::from(diff & old_pawns).iter().flat_map(|s| {
let mut remaining = Bitboard::from(old_white_pawns | old_black_pawns);
let mut to_sub = remaining.bitand(diff).iter().flat_map(|s| {
remaining &= !s.bitboard();
let pft1 = Num::new(pfts[s]);
let visible = PPFeature::MASK[s.file()] & remaining;
let visible = PPFeature::WINDOW[s.file()] & remaining;
visible.iter().map(move |t| {
let pft2 = Num::new(pfts[t]);
PPFeature::new(pft1, pft2)
})
});

let pfts = PFeature::lut(side, ksq, new).to_array();
let mut remaining = Bitboard::from(new_pawns);
let mut to_add = Bitboard::from(diff & new_pawns).iter().flat_map(|s| {
let mut remaining = Bitboard::from(new_white_pawns | new_black_pawns);
let mut to_add = remaining.bitand(diff).iter().flat_map(|s| {
remaining &= !s.bitboard();
let pft1 = Num::new(pfts[s]);
let visible = PPFeature::MASK[s.file()] & remaining;
let visible = PPFeature::WINDOW[s.file()] & remaining;
visible.iter().map(move |t| {
let pft2 = Num::new(pfts[t]);
PPFeature::new(pft1, pft2)
})
});

loop {
let sub = array::from_fn(|_| to_sub.next());
let add = array::from_fn(|_| to_add.next());
if sub != [None; 2] || add != [None; 2] {
let (sub, add) = (to_sub.next(), to_add.next());
if sub.is_some() || add.is_some() {
acc(sub, add);
} else {
break;
Expand Down
2 changes: 1 addition & 1 deletion lib/nnue/feature.rs
Original file line number Diff line number Diff line change
Expand Up @@ -379,7 +379,7 @@ impl PPFeature {
pub const LEN: usize = PFeature::LEN * (PFeature::LEN - 1) / 2;

/// A mask for pawns visible from a [`File`].
pub const MASK: [Bitboard; 8] = [
pub const WINDOW: [Bitboard; 8] = [
File::A.bitboard() | File::B.bitboard(),
File::A.bitboard() | File::B.bitboard() | File::C.bitboard(),
File::B.bitboard() | File::C.bitboard() | File::D.bitboard(),
Expand Down
17 changes: 7 additions & 10 deletions lib/nnue/lin.rs
Original file line number Diff line number Diff line change
@@ -1,13 +1,11 @@
use crate::nnue::{FTQ, HLQ, Layer, Li, Ln, Synapse};
use crate::nnue::{F2V, FTQ, I2F, Layer, Li, Ln, Synapse};
use crate::{simd::*, util::Assume};
use bytemuck::Zeroable;
use std::{array, ops::Mul};
use std::array;

const I: usize = Li::LEN;
const O: usize = Ln::LEN / 2;

const I2F: f32 = (1 << 9) as f32 / (FTQ as f32 * FTQ as f32 * HLQ as f32);

/// The input connection.
#[derive(Debug, Zeroable)]
pub struct Lin<S> {
Expand Down Expand Up @@ -35,16 +33,16 @@ impl<S: for<'a> Synapse<Input<'a> = Ln<'a>, Output = V2<f32>>> Synapse for Lin<S
let xh00 = us[1][2 * i].simd_clamp(Simd::splat(0), Simd::splat(FTQ));
let xh01 = us[1][2 * i + 1].simd_clamp(Simd::splat(0), Simd::splat(FTQ));

let x00 = xl00.mul_high::<9>(xh00);
let x01 = xl01.mul_high::<9>(xh01);
let x00 = xl00.mul_high::<7>(xh00);
let x01 = xl01.mul_high::<7>(xh01);

let xl10 = them[0][2 * i].simd_min(Simd::splat(FTQ));
let xl11 = them[0][2 * i + 1].simd_min(Simd::splat(FTQ));
let xh10 = them[1][2 * i].simd_clamp(Simd::splat(0), Simd::splat(FTQ));
let xh11 = them[1][2 * i + 1].simd_clamp(Simd::splat(0), Simd::splat(FTQ));

let x10 = xl10.mul_high::<9>(xh10);
let x11 = xl11.mul_high::<9>(xh11);
let x10 = xl10.mul_high::<7>(xh10);
let x11 = xl11.mul_high::<7>(xh11);

[Pack::pack(x00, x10), Pack::pack(x01, x11)]
}));
Expand Down Expand Up @@ -121,7 +119,6 @@ impl<S: for<'a> Synapse<Input<'a> = Ln<'a>, Output = V2<f32>>> Synapse for Lin<S
output.map(|i| i.simd_min(Simd::splat(0.)).powi::<2>()),
]);

let result = self.next.forward(active.cast_ref()).reduce_sum();
result.mul(HLQ as f32)
F2V * self.next.forward(active.cast_ref()).reduce_sum()
}
}
Binary file modified lib/nnue/nnue.bin.zst
Binary file not shown.
Loading
Loading