use std::path::Path;
use clap::Parser;
use lumamba::eval::{load_manifest, load_signals, run};
use lumamba::init_threads;
use lumamba::LuMambaEncoder;
#[derive(Parser, Debug)]
#[command(about = "Evaluate a fine-tuned LuMamba classifier on a labeled EEG eval set")]
struct Args {
#[arg(long, default_value = "cpu", value_parser = lumamba::parse_device)]
device: rlx::Device,
#[arg(long, env = "LUMAMBA_WEIGHTS")]
weights: String,
#[arg(long, env = "LUMAMBA_CONFIG")]
config: String,
#[arg(long)]
manifest: String,
#[arg(long, env = "RAYON_NUM_THREADS")]
threads: Option<usize>,
}
fn main() -> anyhow::Result<()> {
let args = Args::parse();
let n_threads = init_threads(args.threads);
let device = args.device;
eprintln!("Device : {device:?} ({n_threads} threads)");
let config_path = lumamba::hf::resolve(&args.config)?;
let weights_path = lumamba::hf::resolve(&args.weights)?;
let (mut model, ms) = LuMambaEncoder::load(&config_path, &weights_path, device)?;
eprintln!("Model : {} ({ms:.0} ms)", model.describe());
eprintln!("Classifier head: {:?}", model.classifier_kind);
let manifest = load_manifest(Path::new(&args.manifest))?;
let base = Path::new(&args.manifest).parent().unwrap_or_else(|| Path::new("."));
let signals_path = base.join(&manifest.signals_path);
let signals = load_signals(&signals_path)?;
eprintln!(
"Eval set : {} task, {} epochs × {} ch × {} samples, metric={}",
manifest.task, signals.n, signals.c, signals.t, manifest.metric
);
let report = run(&mut model, &manifest, &signals)?;
println!("───────────────────────────────────────────");
report.print();
Ok(())
}