brain2qwerty 0.0.1

Brain2Qwerty V1/V2 MEG neural decoding inference in Rust (parity-tested vs Python)
Documentation
//! Brain2Qwerty V2 inference CLI.

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,
    /// Encoder backend: `rlx` (default) or `rust` (pure reference).
    #[arg(long, default_value = "rlx")]
    backend: String,
    /// RLX device when backend=rlx: `auto`, `cpu`, `metal`, `mlx`, `cuda`, …
    #[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(())
}