use std::path::PathBuf;
use brain2qwerty::pipeline::{InferenceInput, Pipeline, PipelineOptions};
use brain2qwerty::tensor::{load_tensor_bin, save_tensor_bin};
use clap::Parser;
#[derive(Parser)]
#[command(
name = "brain2qwerty-infer",
about = "Brain2Qwerty V2 MEG → text inference",
long_about = "Run the V2 pipeline (ConvConformer + CTC + optional TinyLlama beam search).\n\
Input tensors use the custom .bin format (see tensor::load_tensor_bin)."
)]
struct Args {
#[arg(long, default_value = "data/config.yaml")]
config: PathBuf,
#[arg(long, default_value = "data/encoder.safetensors")]
encoder_weights: PathBuf,
#[arg(long)]
llm_weights: Option<PathBuf>,
#[arg(long)]
input: PathBuf,
#[arg(long)]
chan_pos: Option<PathBuf>,
#[arg(long, default_value = "0")]
subject: usize,
#[arg(long)]
output: PathBuf,
#[arg(long, default_value = "rlx")]
backend: String,
#[arg(long, default_value = "auto")]
device: String,
}
fn main() -> anyhow::Result<()> {
let args = Args::parse();
let llm_dir = args
.llm_weights
.as_ref()
.map(|p| p.to_string_lossy().into_owned());
let opts = PipelineOptions {
rlx_device: args.device.clone(),
use_rust_encoder: args.backend.eq_ignore_ascii_case("rust"),
};
let mut pipeline = Pipeline::from_paths_with_options(
&args.config.to_string_lossy(),
&args.encoder_weights.to_string_lossy(),
llm_dir.as_deref(),
opts,
)?;
let neuros = load_tensor_bin(&args.input)?;
let chan_pos = args
.chan_pos
.as_ref()
.map(|p| load_tensor_bin(p))
.transpose()?;
let input = InferenceInput {
neuros,
subject_ids: vec![args.subject],
chan_pos,
};
let out = pipeline.run(&input)?;
let json = serde_json::json!({
"ctc_text": out.ctc_text,
"pred_text": out.pred_text,
"backend": pipeline.encoder_backend,
"device": args.device,
});
std::fs::write(&args.output, serde_json::to_string_pretty(&json)?)?;
save_tensor_bin(&args.output.with_extension("z_final.bin"), &out.z_final)?;
println!("{}", out.pred_text);
Ok(())
}