brain2qwerty 0.0.1

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

use serde::{Deserialize, Serialize};

#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Brain2QwertyConfig {
    pub brain_model_config: EncoderConfig,
    #[serde(default)]
    pub inference: InferenceConfig,
}

#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BuildArgs {
    pub n_in_channels: usize,
    pub n_classes: usize,
    pub dim: usize,
    #[serde(default = "default_merger_virtual")]
    pub merger_n_virtual: usize,
    #[serde(default = "default_fourier_dim")]
    pub fourier_total_dim: usize,
    #[serde(default = "default_td_kernel")]
    pub temporal_downsample_kernel: usize,
    #[serde(default = "default_td_stride")]
    pub temporal_downsample_stride: usize,
}

fn default_merger_virtual() -> usize {
    270
}
fn default_fourier_dim() -> usize {
    2048
}
fn default_td_kernel() -> usize {
    16
}
fn default_td_stride() -> usize {
    4
}

impl BuildArgs {
    pub fn from_json(path: &str) -> anyhow::Result<Self> {
        let s = std::fs::read_to_string(path)?;
        Ok(serde_json::from_str(&s)?)
    }
}

#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EncoderConfig {
    pub name: String,
    pub dim: usize,
    #[serde(default)]
    pub aux_prediction: bool,
    pub encoder_config: SimpleConvConfig,
    pub temporal_downsampling_config: TemporalDownsamplingConfig,
    pub transformer_config: ConformerConfig,
}

#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SimpleConvConfig {
    pub name: String,
    #[serde(default)]
    pub dropout_input: f32,
    #[serde(default)]
    pub conv_dropout: f32,
    pub hidden: usize,
    #[serde(default)]
    pub batch_norm: bool,
    pub depth: usize,
    #[serde(default)]
    pub dilation_period: Option<usize>,
    #[serde(default = "default_dilation_growth")]
    pub dilation_growth: usize,
    pub kernel_size: usize,
    #[serde(default)]
    pub relu_leakiness: f32,
    #[serde(default)]
    pub initial_linear: usize,
    #[serde(default)]
    pub gelu: bool,
    #[serde(default)]
    pub skip: bool,
    #[serde(default)]
    pub scale: Option<f32>,
    #[serde(default)]
    pub subject_layers_config: Option<SubjectLayersConfig>,
    pub merger_config: MergerConfig,
}

#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct SubjectLayersConfig {
    #[serde(default = "default_n_subjects")]
    pub n_subjects: usize,
    #[serde(default = "default_true")]
    pub bias: bool,
    #[serde(default)]
    pub average_subjects: bool,
}

fn default_n_subjects() -> usize {
    200
}
fn default_true() -> bool {
    true
}

#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MergerConfig {
    pub n_virtual_channels: usize,
    pub fourier_emb_config: FourierEmbConfig,
    #[serde(default)]
    pub dropout: f32,
    #[serde(default)]
    pub usage_penalty: f32,
    #[serde(default)]
    pub per_subject: bool,
    #[serde(default)]
    pub embed_ref: bool,
    #[serde(default = "default_n_subjects")]
    pub n_subjects: usize,
}

#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FourierEmbConfig {
    pub n_freqs: Option<usize>,
    pub total_dim: Option<usize>,
    pub n_dims: usize,
    #[serde(default = "default_margin")]
    pub margin: f32,
}

fn default_dilation_growth() -> usize {
    2
}
fn default_margin() -> f32 {
    0.2
}

impl FourierEmbConfig {
    pub fn resolved_n_freqs(&self) -> anyhow::Result<usize> {
        if let Some(n) = self.n_freqs {
            return Ok(n);
        }
        if let Some(total) = self.total_dim {
            let n = (total as f64 / 2.0).powf(1.0 / self.n_dims as f64);
            let rounded = n.round();
            if (n - rounded).abs() > 1e-6 {
                anyhow::bail!("(total_dim / 2) ** (1 / n_dims) must be integer");
            }
            return Ok(rounded as usize);
        }
        anyhow::bail!("Exactly one of n_freqs and total_dim must be set")
    }

    pub fn total_dim(&self) -> anyhow::Result<usize> {
        let n = self.resolved_n_freqs()?;
        Ok((n.pow(self.n_dims as u32)) * 2)
    }
}

#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TemporalDownsamplingConfig {
    pub kernel_size: usize,
    pub stride: usize,
    #[serde(default = "default_true")]
    pub layer_norm: bool,
    #[serde(default = "default_true")]
    pub layer_norm_affine: bool,
    #[serde(default = "default_true")]
    pub gelu: bool,
}

#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ConformerConfig {
    pub name: String,
    pub ffn_dim: usize,
    pub num_heads: usize,
    pub num_layers: usize,
    pub depthwise_conv_kernel_size: usize,
    #[serde(default)]
    pub dropout: f32,
    #[serde(default)]
    pub use_group_norm: bool,
    #[serde(default)]
    pub convolution_first: bool,
}

#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct InferenceConfig {
    #[serde(default = "default_classes")]
    pub num_classes: usize,
    #[serde(default = "default_channels")]
    pub n_in_channels: usize,
    #[serde(default = "default_llm")]
    pub llm_name: String,
    #[serde(default = "default_lora_rank")]
    pub lora_rank: usize,
    #[serde(default = "default_lora_alpha")]
    pub lora_alpha: usize,
    #[serde(default)]
    pub lora_dropout: f32,
    #[serde(default = "default_lora_targets")]
    pub lora_target_modules: Vec<String>,
    #[serde(default = "default_max_tokens")]
    pub max_new_tokens: usize,
    #[serde(default = "default_beams")]
    pub num_beams: usize,
    #[serde(default = "default_val_beams")]
    pub val_num_beams: usize,
    #[serde(default = "default_length_penalty")]
    pub length_penalty: f32,
    #[serde(default = "default_sys")]
    pub sys_prompt: String,
    #[serde(default = "default_mid")]
    pub mid_prompt: String,
    #[serde(default = "default_resp")]
    pub resp_prompt: String,
    #[serde(default = "default_pool_layers")]
    pub word_pool_n_layers: usize,
    #[serde(default = "default_true")]
    pub seg_include_blanks: bool,
}

impl Default for InferenceConfig {
    fn default() -> Self {
        Self {
            num_classes: default_classes(),
            n_in_channels: default_channels(),
            llm_name: default_llm(),
            lora_rank: default_lora_rank(),
            lora_alpha: default_lora_alpha(),
            lora_dropout: 0.0,
            lora_target_modules: default_lora_targets(),
            max_new_tokens: default_max_tokens(),
            num_beams: default_beams(),
            val_num_beams: default_val_beams(),
            length_penalty: default_length_penalty(),
            sys_prompt: default_sys(),
            mid_prompt: default_mid(),
            resp_prompt: default_resp(),
            word_pool_n_layers: default_pool_layers(),
            seg_include_blanks: true,
        }
    }
}

fn default_classes() -> usize {
    29
}
fn default_channels() -> usize {
    306
}
fn default_llm() -> String {
    "TinyLlama/TinyLlama-1.1B-Chat-v1.0".into()
}
fn default_lora_rank() -> usize {
    2
}
fn default_lora_alpha() -> usize {
    4
}
fn default_lora_targets() -> Vec<String> {
    vec![
        "q_proj".into(),
        "v_proj".into(),
        "k_proj".into(),
        "o_proj".into(),
    ]
}
fn default_max_tokens() -> usize {
    60
}
fn default_beams() -> usize {
    16
}
fn default_val_beams() -> usize {
    1
}
fn default_length_penalty() -> f32 {
    0.2
}
fn default_sys() -> String {
    "CTC: ".into()
}
fn default_mid() -> String {
    "\nMEG: ".into()
}
fn default_resp() -> String {
    "\nOutput: ".into()
}
fn default_pool_layers() -> usize {
    2
}

impl Brain2QwertyConfig {
    pub fn from_yaml(path: &str) -> anyhow::Result<Self> {
        let s = std::fs::read_to_string(path)?;
        Ok(serde_yaml::from_str(&s)?)
    }

    /// Tiny config for CI (mirrors tests/test_v2_model.py).
    pub fn tiny() -> Self {
        let mut cfg = Self::from_yaml(&format!(
            "{}/../../data/config.yaml",
            env!("CARGO_MANIFEST_DIR")
        ))
        .unwrap_or_else(|_| Self {
            brain_model_config: EncoderConfig {
                name: "ConvConformer".into(),
                dim: 1024,
                aux_prediction: true,
                encoder_config: SimpleConvConfig {
                    name: "SimpleConv".into(),
                    dropout_input: 0.2,
                    conv_dropout: 0.5,
                    hidden: 1500,
                    batch_norm: true,
                    depth: 4,
                    dilation_period: Some(3),
                    dilation_growth: 2,
                    kernel_size: 5,
                    relu_leakiness: 0.01,
                    initial_linear: 512,
                    gelu: true,
                    skip: true,
                    scale: Some(0.1),
                    subject_layers_config: Some(SubjectLayersConfig::default()),
                    merger_config: MergerConfig {
                        n_virtual_channels: 270,
                        fourier_emb_config: FourierEmbConfig {
                            n_freqs: None,
                            total_dim: Some(2048),
                            n_dims: 2,
                            margin: 0.2,
                        },
                        dropout: 0.2,
                        usage_penalty: 1.0,
                        per_subject: true,
                        embed_ref: false,
                        n_subjects: 200,
                    },
                },
                temporal_downsampling_config: TemporalDownsamplingConfig {
                    kernel_size: 16,
                    stride: 4,
                    layer_norm: true,
                    layer_norm_affine: true,
                    gelu: true,
                },
                transformer_config: ConformerConfig {
                    name: "Conformer".into(),
                    ffn_dim: 1024,
                    num_heads: 4,
                    num_layers: 4,
                    depthwise_conv_kernel_size: 17,
                    dropout: 0.3,
                    use_group_norm: true,
                    convolution_first: false,
                },
            },
            inference: InferenceConfig::default(),
        });
        let ec = &mut cfg.brain_model_config;
        ec.dim = 32;
        ec.encoder_config.hidden = 64;
        ec.encoder_config.depth = 2;
        ec.encoder_config.initial_linear = 16;
        ec.encoder_config.merger_config.n_virtual_channels = 16;
        ec.encoder_config.merger_config.fourier_emb_config.total_dim = Some(32);
        ec.transformer_config.ffn_dim = 32;
        ec.transformer_config.num_heads = 2;
        ec.transformer_config.num_layers = 1;
        cfg.inference.max_new_tokens = 12;
        cfg.inference.num_beams = 4;
        cfg.inference.length_penalty = 0.2;
        cfg
    }
}