use super::features::{encode, encode_orientation};
use crate::game::moves::{legal_moves, Move};
use crate::game::rules::Variant;
use crate::game::state::GameState;
#[derive(Debug, Clone)]
pub struct StateSample {
pub moves: Vec<Vec<f32>>,
pub chosen: usize,
#[allow(dead_code)] pub final_score: u32,
}
pub fn samples_from_game(variant: Variant, history: &[Move]) -> Vec<StateSample> {
let final_score = history.len() as u32;
let mut st = GameState::new(variant);
let mut out = Vec::with_capacity(history.len());
for &mv in history {
let legal = legal_moves(&st);
let Some(chosen) = legal.iter().position(|m| *m == mv) else {
break; };
let moves = legal.iter().map(|m| encode(&st, m)).collect();
out.push(StateSample {
moves,
chosen,
final_score,
});
if !st.apply(mv) {
break; }
}
out
}
pub fn samples_from_corpus(variant: Variant) -> Vec<StateSample> {
let mut out = Vec::new();
for rec in morpion_solitaire_records::RECORDS.iter() {
let Ok(g) = crate::game::io::import_save(rec.2) else {
continue;
};
if g.variant != variant {
continue;
}
out.extend(samples_from_game(g.variant, &g.history));
}
out
}
fn samples_from_game_oriented(variant: Variant, history: &[Move], t: usize) -> Vec<StateSample> {
let final_score = history.len() as u32;
let mut st = GameState::new(variant);
let mut out = Vec::with_capacity(history.len());
for &mv in history {
let legal = legal_moves(&st);
let Some(chosen) = legal.iter().position(|m| *m == mv) else {
break;
};
let moves = legal
.iter()
.map(|m| encode_orientation(&st, m, t))
.collect();
out.push(StateSample {
moves,
chosen,
final_score,
});
if !st.apply(mv) {
break;
}
}
out
}
pub fn augmented_samples_from_corpus(variant: Variant) -> Vec<StateSample> {
let mut out = Vec::new();
for rec in morpion_solitaire_records::RECORDS.iter() {
let Ok(g) = crate::game::io::import_save(rec.2) else {
continue;
};
if g.variant != variant {
continue;
}
for t in 0..8 {
out.extend(samples_from_game_oriented(g.variant, &g.history, t));
}
}
out
}
pub fn augmented_samples_from_games(variant: Variant, games: &[Vec<Move>]) -> Vec<StateSample> {
let mut out = Vec::new();
for g in games {
for t in 0..8 {
out.extend(samples_from_game_oriented(variant, g, t));
}
}
out
}
#[derive(Debug, Clone)]
pub struct ValueSample {
pub features: Vec<f32>,
pub target: f32,
}
pub fn value_samples_from_corpus(variant: Variant, augment: bool) -> Vec<ValueSample> {
use super::position::{encode_value, value_target};
let orients = if augment { 0..8 } else { 0..1 };
let mut out = Vec::new();
for rec in morpion_solitaire_records::RECORDS.iter() {
let Ok(g) = crate::game::io::import_save(rec.2) else {
continue;
};
if g.variant != variant {
continue;
}
let target = value_target(g.history.len() as u32);
for t in orients.clone() {
let mut st = GameState::new(g.variant);
for &mv in &g.history {
out.push(ValueSample {
features: encode_value(&st, t),
target,
});
if !st.apply(mv) {
break;
}
}
out.push(ValueSample {
features: encode_value(&st, t),
target,
});
}
}
out
}
pub fn value_samples_from_games(
variant: Variant,
games: &[Vec<Move>],
augment: bool,
) -> Vec<ValueSample> {
use super::position::{encode_value, value_target};
let orients = if augment { 0..8 } else { 0..1 };
let mut out = Vec::new();
for g in games {
let target = value_target(g.len() as u32);
for t in orients.clone() {
let mut st = GameState::new(variant);
for &mv in g {
out.push(ValueSample {
features: encode_value(&st, t),
target,
});
if !st.apply(mv) {
break;
}
}
out.push(ValueSample {
features: encode_value(&st, t),
target,
});
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
use crate::game::moves::legal_moves;
use crate::search::neural::features::FEATURE_LEN;
#[test]
fn samples_match_a_replayed_game() {
let mut st = GameState::new(Variant::T5);
let mut history = Vec::new();
for _ in 0..15 {
let ms = legal_moves(&st);
if ms.is_empty() {
break;
}
history.push(ms[0]);
st.apply(ms[0]);
}
let samples = samples_from_game(Variant::T5, &history);
assert_eq!(samples.len(), history.len());
for s in &samples {
assert!(s.chosen < s.moves.len());
assert_eq!(s.final_score, history.len() as u32);
for f in &s.moves {
assert_eq!(f.len(), FEATURE_LEN);
}
}
}
#[test]
fn corpus_yields_5t_samples() {
let samples = samples_from_corpus(Variant::T5);
assert!(
samples.len() > 178,
"expected hundreds of 5T decision points, got {}",
samples.len()
);
assert!(samples.iter().all(|s| s.chosen < s.moves.len()));
}
}