brain2qwerty 0.0.1

Brain2Qwerty V1/V2 MEG neural decoding inference in Rust (parity-tested vs Python)
Documentation
//! Brain2Qwerty V1 configuration (mirrors `brain2qwerty_v1/config/model_config.py`).

use crate::config::SimpleConvConfig;

/// x-transformers `TransformerEncoder` settings for the sentence-level stack.
#[derive(Clone, Debug, serde::Serialize, serde::Deserialize)]
pub struct V1TransformerConfig {
    pub name: String,
    pub alibi_pos_bias: bool,
    pub depth: usize,
    pub heads: usize,
    #[serde(default = "default_ff_mult")]
    pub ff_mult: usize,
    #[serde(default = "default_true")]
    pub use_scalenorm: bool,
    #[serde(default = "default_false")]
    pub rotary_pos_emb: bool,
    #[serde(default = "default_true")]
    pub scale_residual: bool,
}

fn default_ff_mult() -> usize {
    4
}
fn default_true() -> bool {
    true
}
fn default_false() -> bool {
    false
}

impl Default for V1TransformerConfig {
    fn default() -> Self {
        Self {
            name: "TransformerEncoder".into(),
            alibi_pos_bias: true,
            depth: 4,
            heads: 2,
            ff_mult: 4,
            use_scalenorm: true,
            rotary_pos_emb: true,
            scale_residual: true,
        }
    }
}

#[derive(Clone, Debug, serde::Serialize, serde::Deserialize)]
pub struct V1EncoderConfig {
    pub name: String,
    /// Temporal pooling mode; production uses `"att"` (Bahdanau attention).
    pub time_agg_out: String,
    pub encoder_config: SimpleConvConfig,
}

#[derive(Clone, Debug, serde::Serialize, serde::Deserialize)]
pub struct V1Config {
    pub brain_model_config: V1EncoderConfig,
    pub transformer_config: V1TransformerConfig,
    pub num_classes: usize,
    pub n_in_channels: usize,
}

impl V1Config {
    /// Parse `data/v1_config.yaml` (or any compatible YAML path).
    pub fn from_yaml(path: &str) -> anyhow::Result<Self> {
        let s = std::fs::read_to_string(path)?;
        Ok(serde_yaml::from_str(&s)?)
    }

    pub fn production() -> Self {
        Self::from_yaml(&format!(
            "{}/../../data/v1_config.yaml",
            env!("CARGO_MANIFEST_DIR")
        ))
        .expect("v1_config.yaml")
    }

    /// Tiny CI config (mirrors `tests/test_v1_model.py`).
    pub fn tiny() -> Self {
        let mut cfg = Self::production();
        cfg.brain_model_config.encoder_config.hidden = 32;
        cfg.brain_model_config.encoder_config.depth = 2;
        cfg.brain_model_config.encoder_config.initial_linear = 16;
        cfg.brain_model_config
            .encoder_config
            .merger_config
            .n_virtual_channels = 16;
        cfg.brain_model_config
            .encoder_config
            .merger_config
            .fourier_emb_config
            .total_dim = Some(32);
        cfg.transformer_config.depth = 1;
        cfg.transformer_config.heads = 1;
        cfg.transformer_config.rotary_pos_emb = false;
        cfg
    }
}