use std::collections::{BTreeMap, HashMap, HashSet};
use std::fs::{self, File};
use std::io::{BufRead, BufReader, BufWriter, Write};
use std::path::PathBuf;
use anyhow::{Context, Result};
use clap::{Parser, Subcommand};
use shogiesa_core::{Observation, PositionRecord, sfen::Sfen};
use shogiesa_usi::UsiEngine;
use tracing::info;
#[derive(Parser)]
#[command(
name = "shogiesa",
about = "Shogi training data feed for NNUE engines."
)]
struct Cli {
#[command(subcommand)]
command: Commands,
}
#[derive(Subcommand)]
enum Commands {
Extract(ExtractArgs),
Label(LabelArgs),
Report(ReportArgs),
Validate(ValidateArgs),
}
#[derive(clap::Args)]
struct ExtractArgs {
#[arg(short, long)]
input: PathBuf,
#[arg(short, long)]
out: PathBuf,
#[arg(long, default_value = "1")]
min_ply: u32,
#[arg(long)]
max_ply: Option<u32>,
#[arg(long, default_value = "1", name = "every-n-plies")]
every_n_plies: u32,
#[arg(long)]
dedup: bool,
}
#[derive(clap::Args)]
struct ReportArgs {
#[arg(short, long)]
input: PathBuf,
}
#[derive(clap::Args)]
struct ValidateArgs {
#[arg(short, long)]
input: PathBuf,
#[arg(long)]
strict: bool,
}
#[derive(clap::Args)]
struct LabelArgs {
#[arg(short, long)]
input: PathBuf,
#[arg(long)]
engine: PathBuf,
#[arg(long)]
engine_name: Option<String>,
#[arg(long)]
depths: String,
#[arg(long, default_value = "10000")]
timeout_ms: u64,
#[arg(short, long)]
out: PathBuf,
}
fn main() -> Result<()> {
tracing_subscriber::fmt()
.with_env_filter(tracing_subscriber::EnvFilter::from_default_env())
.with_writer(std::io::stderr)
.init();
let cli = Cli::parse();
match cli.command {
Commands::Extract(args) => cmd_extract(args),
Commands::Label(args) => cmd_label(args),
Commands::Report(args) => cmd_report(args),
Commands::Validate(args) => cmd_validate(args),
}
}
fn cmd_extract(args: ExtractArgs) -> Result<()> {
let config = shogiesa_csa::ExtractConfig {
min_ply: args.min_ply,
max_ply: args.max_ply,
every_n: args.every_n_plies,
dedup: args.dedup,
};
let paths = collect_csa_paths(&args.input)?;
if paths.is_empty() {
anyhow::bail!("no .csa files found in {:?}", args.input);
}
let out_file =
File::create(&args.out).with_context(|| format!("cannot create {:?}", args.out))?;
let mut writer = BufWriter::new(out_file);
let mut seen = HashSet::new();
let mut total_games = 0usize;
let mut total_positions = 0usize;
let mut skipped = 0usize;
for path in &paths {
total_games += 1;
match shogiesa_csa::extract_from_path(path, &config, &mut seen) {
Ok(records) => {
for rec in &records {
serde_json::to_writer(&mut writer, rec)?;
writer.write_all(b"\n")?;
}
total_positions += records.len();
}
Err(e) => {
tracing::warn!(path = %path.display(), "skipped: {e}");
skipped += 1;
}
}
info!(
games = total_games,
positions = total_positions,
"processed {}",
path.display()
);
}
writer.flush()?;
eprintln!(
"done: {} games, {} positions extracted, {} skipped → {:?}",
total_games, total_positions, skipped, args.out
);
Ok(())
}
fn cmd_label(args: LabelArgs) -> Result<()> {
let depths: Vec<u32> = args
.depths
.split(',')
.filter_map(|s| s.trim().parse().ok())
.collect();
if depths.is_empty() {
anyhow::bail!("--depths must contain at least one valid integer, e.g. '4,6,8'");
}
let name = args.engine_name.unwrap_or_default();
let mut engine = UsiEngine::launch(&args.engine, name, args.timeout_ms)
.with_context(|| format!("failed to launch engine {:?}", args.engine))?;
info!(
engine = engine.engine_name,
depths = ?depths,
"labeling started"
);
let reader = BufReader::new(
File::open(&args.input).with_context(|| format!("cannot open {:?}", args.input))?,
);
let out_file =
File::create(&args.out).with_context(|| format!("cannot create {:?}", args.out))?;
let mut writer = BufWriter::new(out_file);
let mut total = 0usize;
let mut labeled = 0usize;
let mut skipped = 0usize;
for (i, line) in reader.lines().enumerate() {
let line = line?;
if line.trim().is_empty() {
continue;
}
total += 1;
let mut rec: PositionRecord = match serde_json::from_str(&line) {
Ok(r) => r,
Err(e) => {
tracing::warn!(line = i + 1, "JSON parse error: {e}");
skipped += 1;
continue;
}
};
if Sfen::parse(&rec.sfen).is_err() {
tracing::warn!(line = i + 1, sfen = rec.sfen, "invalid SFEN, skipping");
skipped += 1;
continue;
}
for &depth in &depths {
match engine.analyse(&rec.sfen, depth, args.timeout_ms) {
Ok(result) => {
rec.observations.push(Observation {
engine: engine.engine_name.clone(),
engine_version: engine.engine_version.clone(),
depth: result.depth,
score: result.score,
bestmove: result.bestmove,
nodes: result.nodes,
time_ms: result.time_ms,
pv: result.pv,
});
}
Err(e) => {
tracing::warn!(line = i + 1, depth, "analysis error: {e}");
}
}
}
serde_json::to_writer(&mut writer, &rec)?;
writer.write_all(b"\n")?;
labeled += 1;
}
writer.flush()?;
engine.quit();
eprintln!(
"done: {total} positions, {labeled} labeled, {skipped} skipped → {:?}",
args.out
);
Ok(())
}
fn collect_csa_paths(input: &PathBuf) -> Result<Vec<PathBuf>> {
if input.is_file() {
return Ok(vec![input.clone()]);
}
let mut paths = Vec::new();
for entry in
fs::read_dir(input).with_context(|| format!("cannot read directory {:?}", input))?
{
let entry = entry?;
let p = entry.path();
if p.extension().and_then(|e| e.to_str()) == Some("csa") {
paths.push(p);
}
}
paths.sort();
Ok(paths)
}
fn load_records(path: &PathBuf) -> Result<(Vec<PositionRecord>, usize)> {
let content = fs::read_to_string(path).with_context(|| format!("cannot read {:?}", path))?;
let non_empty: Vec<&str> = content.lines().filter(|l| !l.trim().is_empty()).collect();
let broken = non_empty
.iter()
.filter(|l| serde_json::from_str::<PositionRecord>(l).is_err())
.count();
let records: Vec<PositionRecord> = non_empty
.iter()
.enumerate()
.filter_map(|(i, line)| {
serde_json::from_str::<PositionRecord>(line)
.map_err(|e| tracing::warn!(line = i + 1, "parse error: {e}"))
.ok()
})
.collect();
Ok((records, broken))
}
fn cmd_report(args: ReportArgs) -> Result<()> {
let (records, broken) = load_records(&args.input)?;
if records.is_empty() {
println!("no valid records in {:?}", args.input);
return Ok(());
}
let n = records.len();
let mut phases = BTreeMap::<String, usize>::new();
let mut sides = BTreeMap::<String, usize>::new();
let mut schema_versions = BTreeMap::<u32, usize>::new();
let mut ply_sum = 0u64;
let mut ply_min = u32::MAX;
let mut ply_max = 0u32;
let mut sfen_counts: HashMap<&str, usize> = HashMap::new();
let mut tag_mismatches = 0usize;
let mut invalid_sfens = 0usize;
for rec in &records {
*phases.entry(format!("{}", rec.tags.phase)).or_default() += 1;
*sides
.entry(format!("{}", rec.tags.side_to_move))
.or_default() += 1;
*schema_versions.entry(rec.schema_version).or_default() += 1;
let ply = rec.source.ply;
ply_sum += ply as u64;
ply_min = ply_min.min(ply);
ply_max = ply_max.max(ply);
*sfen_counts.entry(rec.sfen.as_str()).or_default() += 1;
match Sfen::parse(&rec.sfen) {
Ok(sfen) => {
if sfen.side_to_move() != rec.tags.side_to_move {
tag_mismatches += 1;
}
}
Err(_) => invalid_sfens += 1,
}
}
let duplicate_sfens: usize = sfen_counts
.values()
.filter(|&&c| c > 1)
.map(|&c| c - 1)
.sum();
let duplicate_rate = duplicate_sfens as f64 / n as f64 * 100.0;
let mut sources = BTreeMap::<&str, usize>::new();
for rec in &records {
*sources.entry(rec.source.path.as_str()).or_default() += 1;
}
let top_source_pct = sources.values().max().copied().unwrap_or(0) as f64 / n as f64 * 100.0;
let opening_pct = phases.get("opening").copied().unwrap_or(0) as f64 / n as f64 * 100.0;
let black_count = sides.get("black").copied().unwrap_or(0);
let white_count = sides.get("white").copied().unwrap_or(0);
println!("=== shogiesa report ===");
println!("positions : {n}");
println!("broken lines : {broken}");
println!(
"ply range : {ply_min}–{ply_max} (avg {:.1})",
ply_sum as f64 / n as f64
);
println!("invalid SFENs : {invalid_sfens}");
println!("duplicate SFENs: {duplicate_sfens}");
println!("tag mismatches : {tag_mismatches} (side_to_move vs SFEN)");
println!();
println!("schema versions: {schema_versions:?}");
println!();
println!("phase distribution:");
for (phase, count) in &phases {
println!(
" {phase:<12} {count:>6} ({:.1}%)",
*count as f64 / n as f64 * 100.0
);
}
println!();
println!("side to move:");
for (side, count) in &sides {
println!(
" {side:<12} {count:>6} ({:.1}%)",
*count as f64 / n as f64 * 100.0
);
}
println!();
println!("source files: {}", sources.len());
for (path, count) in sources.iter().take(10) {
println!(" {path}: {count}");
}
if sources.len() > 10 {
println!(" … and {} more", sources.len() - 10);
}
println!();
println!("source dominance:");
let top_warn = if top_source_pct > 50.0 {
"WARN: too concentrated"
} else {
"OK"
};
println!(" top source : {top_source_pct:.1}% {top_warn}");
println!();
println!("balance warnings:");
let opening_warn = if opening_pct > 50.0 {
"WARN: too high"
} else {
"OK"
};
println!(" opening ratio : {opening_pct:.1}% {opening_warn}");
let (b_pct, w_pct) = (
black_count as f64 / n as f64 * 100.0,
white_count as f64 / n as f64 * 100.0,
);
let side_warn = if b_pct > 65.0 || w_pct > 65.0 {
"WARN"
} else {
"OK"
};
println!(" side imbalance : {b_pct:.1}% / {w_pct:.1}% {side_warn}");
let dup_warn = if duplicate_rate > 5.0 {
"WARN: too high"
} else {
"OK"
};
println!(" duplicate rate : {duplicate_rate:.1}% {dup_warn}");
Ok(())
}
fn cmd_validate(args: ValidateArgs) -> Result<()> {
let content =
fs::read_to_string(&args.input).with_context(|| format!("cannot read {:?}", args.input))?;
let total_lines = content.lines().filter(|l| !l.trim().is_empty()).count();
let mut valid_json = 0usize;
let mut valid_records = 0usize;
let mut tag_mismatches = 0usize;
let mut invalid_sfens = 0usize;
let mut schema_versions = BTreeMap::<u32, usize>::new();
let mut seen_sfens: HashSet<String> = HashSet::new();
let mut duplicate_sfens = 0usize;
for line in content.lines().filter(|l| !l.trim().is_empty()) {
let Ok(val) = serde_json::from_str::<serde_json::Value>(line) else {
continue;
};
valid_json += 1;
let Ok(rec) = serde_json::from_value::<PositionRecord>(val) else {
continue;
};
valid_records += 1;
*schema_versions.entry(rec.schema_version).or_default() += 1;
if !seen_sfens.insert(rec.sfen.clone()) {
duplicate_sfens += 1;
}
match Sfen::parse(&rec.sfen) {
Ok(sfen) => {
if sfen.side_to_move() != rec.tags.side_to_move {
tag_mismatches += 1;
}
}
Err(_) => invalid_sfens += 1,
}
}
let broken = total_lines - valid_json;
let has_problems = tag_mismatches > 0 || broken > 0 || invalid_sfens > 0;
println!("=== shogiesa validate ===");
println!("total lines : {total_lines}");
println!("valid JSON : {valid_json}");
println!("valid records : {valid_records}");
println!("broken lines : {broken}");
println!("invalid SFENs : {invalid_sfens}");
println!("duplicate SFENs: {duplicate_sfens}");
println!("tag mismatches : {tag_mismatches} (side_to_move vs SFEN)");
println!("schema versions: {schema_versions:?}");
if has_problems {
println!();
if broken > 0 {
println!("WARN: {broken} broken lines");
}
if invalid_sfens > 0 {
println!("WARN: {invalid_sfens} invalid SFENs");
}
if tag_mismatches > 0 {
println!("WARN: {tag_mismatches} side_to_move tag mismatches");
}
if args.strict {
std::process::exit(1);
}
} else {
println!();
println!("OK");
}
Ok(())
}