brain2qwerty 0.0.1

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

use std::collections::HashMap;

use crate::config::Brain2QwertyConfig;
use crate::model::conv_conformer::{ConvConformer, ConvConformerOutput};
use crate::model_rlx::graph::{
    build_conformer_tail_graph, build_conv_module_graph, build_post_conv_graph,
    build_pre_conv_graph, EncoderSpec,
};
use crate::model_rlx::weights::{apply_params, build_rlx_params, load_encoder_weights};
use crate::tensor::Tensor;

type SubgraphKey = (String, usize, usize, usize);

const SUBGRAPH_KINDS_RUST_CONV: [&str; 2] = ["pre_conv", "post_conv"];
const SUBGRAPH_KINDS_RLX_CONV: [&str; 3] = ["pre_conv", "conv", "post_conv"];

fn use_rlx_conv(device: rlx::Device) -> bool {
    // Full-GPU conv only where validated; others use Rust conv between RLX blocks.
    matches!(device, rlx::Device::Cuda | rlx::Device::Ane)
}

fn use_monolithic_tail(device: rlx::Device) -> bool {
    // Fused tail graph: best on ANE/CUDA; Metal/wgpu prefer orchestrated + Rust conv.
    matches!(device, rlx::Device::Ane | rlx::Device::Cuda)
}

pub struct ConvConformerRlx {
    pub rust: ConvConformer,
    device: rlx::Device,
    session: rlx::Session,
    params: crate::model_rlx::weights::ParamMap,
    subgraph_cache: HashMap<SubgraphKey, rlx::CompiledGraph>,
    tail_graph: Option<rlx::CompiledGraph>,
    spec_template: EncoderSpec,
    precompiled: Option<(usize, usize)>,
    act_a: Vec<f32>,
    act_b: Vec<f32>,
}

impl ConvConformerRlx {
    pub fn from_tiny_weights(weights_path: &str) -> anyhow::Result<Self> {
        Self::from_config_and_weights(&Brain2QwertyConfig::tiny(), weights_path)
    }

    pub fn from_pretrained(config_path: &str, weights_path: &str) -> anyhow::Result<Self> {
        let cfg = Brain2QwertyConfig::from_yaml(config_path)?;
        Self::from_config_and_weights(&cfg, weights_path)
    }

    pub fn from_config_and_weights(
        cfg: &Brain2QwertyConfig,
        weights_path: &str,
    ) -> anyhow::Result<Self> {
        let mut store = load_encoder_weights(weights_path)?;
        let rust = ConvConformer::from_config_and_weights(&cfg.brain_model_config, &mut store, "")?;
        let tc = &cfg.brain_model_config.transformer_config;
        let spec_template = EncoderSpec {
            b: 1,
            t: 27,
            d: cfg.brain_model_config.dim,
            n_classes: cfg.inference.num_classes,
            num_layers: tc.num_layers,
            num_heads: tc.num_heads,
            ffn_dim: tc.ffn_dim,
            dw_kernel: tc.depthwise_conv_kernel_size,
            aux: cfg.brain_model_config.aux_prediction,
        };
        let mut raw = load_encoder_weights(weights_path)?;
        let params = build_rlx_params(&mut raw, &spec_template)?;
        Ok(Self {
            rust,
            device: rlx::Device::Cpu,
            session: rlx::Session::new(rlx::Device::Cpu),
            params,
            subgraph_cache: HashMap::new(),
            tail_graph: None,
            spec_template,
            precompiled: None,
            act_a: Vec::new(),
            act_b: Vec::new(),
        })
    }

    pub fn with_device(mut self, device: rlx::Device) -> Self {
        self.device = device;
        self.session = rlx::Session::new(device);
        self.subgraph_cache.clear();
        self.tail_graph = None;
        self.precompiled = None;
        self
    }

    /// Pre-compile all orchestrated RLX subgraphs for the given `(B, T)` spec.
    pub fn precompile_tail(&mut self, spec: &EncoderSpec) -> anyhow::Result<()> {
        if self.precompiled == Some((spec.b, spec.t)) {
            return Ok(());
        }
        if use_monolithic_tail(self.device) {
            let graph = build_conformer_tail_graph(spec);
            let mut compiled = self.session.compile(graph);
            apply_params(&mut compiled, &self.params);
            compiled.finalize_params();
            self.tail_graph = Some(compiled);
        } else {
            let kinds = if use_rlx_conv(self.device) {
                SUBGRAPH_KINDS_RLX_CONV.as_slice()
            } else {
                SUBGRAPH_KINDS_RUST_CONV.as_slice()
            };
            for layer in 0..spec.num_layers {
                for kind in kinds {
                    self.ensure_subgraph(kind, layer, spec)?;
                }
            }
        }
        self.precompiled = Some((spec.b, spec.t));
        Ok(())
    }

    pub fn forward(
        &mut self,
        neuros: &Tensor,
        subject_ids: &[usize],
        chan_pos: Option<&Tensor>,
    ) -> ConvConformerOutput {
        let prefix = self.rust.forward_prefix(neuros, subject_ids, chan_pos);
        let b = prefix.z_transformer_in.shape[0];
        let t = prefix.z_transformer_in.shape[1];
        let spec = EncoderSpec {
            b,
            t,
            ..self.spec_template
        };
        match self.run_rlx_tail(&prefix.z_transformer_in, &spec) {
            Ok((z_final, c_out)) => ConvConformerOutput {
                z: prefix.z,
                z_enc: prefix.z_enc,
                z_transformer_in: prefix.z_transformer_in,
                z_final,
                c_out,
                z_aux: prefix.z_aux,
            },
            Err(e) => {
                tracing::warn!("RLX tail failed ({e}); using rust reference");
                let z_final = self.rust.conformer.forward(&prefix.z_transformer_in);
                let c_out = self.rust.forward_head(&z_final);
                ConvConformerOutput {
                    z: prefix.z,
                    z_enc: prefix.z_enc,
                    z_transformer_in: prefix.z_transformer_in,
                    z_final,
                    c_out,
                    z_aux: prefix.z_aux,
                }
            }
        }
    }

    pub fn run_tail(&mut self, z: &Tensor) -> anyhow::Result<(Tensor, Tensor)> {
        let spec = EncoderSpec {
            b: z.shape[0],
            t: z.shape[1],
            ..self.spec_template
        };
        self.run_rlx_tail(z, &spec)
    }

    fn run_rlx_tail(&mut self, z: &Tensor, spec: &EncoderSpec) -> anyhow::Result<(Tensor, Tensor)> {
        self.precompile_tail(spec)?;
        if use_monolithic_tail(self.device) {
            let compiled = self.tail_graph.as_mut().unwrap();
            let outs = compiled.run(&[("z", z.data.as_slice())]);
            let z_final = Tensor::from_vec(outs[0].clone(), vec![spec.b, spec.t, spec.d]);
            let c_shape = vec![spec.b, spec.t, spec.n_classes];
            let c_out = Tensor::from_vec(outs[1].clone(), c_shape);
            return Ok((z_final, c_out));
        }

        let n = z.numel();
        if self.act_a.len() != n {
            self.act_a.resize(n, 0.0);
            self.act_b.resize(n, 0.0);
        }
        self.act_a.copy_from_slice(&z.data);

        for layer in 0..spec.num_layers {
            self.run_layer(layer, spec)?;
        }

        let z_final = Tensor::from_vec(
            std::mem::take(&mut self.act_a),
            vec![spec.b, spec.t, spec.d],
        );
        let c_out = self.rust.forward_head(&z_final);
        Ok((z_final, c_out))
    }

    fn run_layer(&mut self, layer: usize, spec: &EncoderSpec) -> anyhow::Result<()> {
        self.run_subgraph_into("pre_conv", layer, spec)?;
        std::mem::swap(&mut self.act_a, &mut self.act_b);

        if use_rlx_conv(self.device) {
            let residual = self.act_a.clone();
            self.run_subgraph_into("conv", layer, spec)?;
            std::mem::swap(&mut self.act_a, &mut self.act_b);
            for (out, r) in self.act_a.iter_mut().zip(&residual) {
                *out += *r;
            }
        } else {
            let conv = &self.rust.conformer.layers[layer].conv;
            let h = Tensor::from_vec(
                std::mem::take(&mut self.act_a),
                vec![spec.b, spec.t, spec.d],
            );
            let residual = h.clone();
            self.act_a = conv.forward(&h).add(&residual).data;
        }

        self.run_subgraph_into("post_conv", layer, spec)?;
        std::mem::swap(&mut self.act_a, &mut self.act_b);
        Ok(())
    }

    fn ensure_subgraph(
        &mut self,
        kind: &str,
        layer: usize,
        spec: &EncoderSpec,
    ) -> anyhow::Result<()> {
        let key = (kind.to_string(), layer, spec.b, spec.t);
        if self.subgraph_cache.contains_key(&key) {
            return Ok(());
        }
        let graph = match kind {
            "pre_conv" => build_pre_conv_graph(spec, layer),
            "conv" => build_conv_module_graph(spec, layer),
            "post_conv" => build_post_conv_graph(spec, layer),
            other => anyhow::bail!("unknown RLX subgraph {other}"),
        };
        let mut compiled = self.session.compile(graph);
        apply_params(&mut compiled, &self.params);
        compiled.finalize_params();
        self.subgraph_cache.insert(key, compiled);
        Ok(())
    }

    fn run_subgraph_into(
        &mut self,
        kind: &str,
        layer: usize,
        spec: &EncoderSpec,
    ) -> anyhow::Result<()> {
        self.ensure_subgraph(kind, layer, spec)?;
        let key = (kind.to_string(), layer, spec.b, spec.t);
        let input = self.act_a.clone();
        let compiled = self.subgraph_cache.get_mut(&key).unwrap();
        let outs = compiled.run(&[("z", input.as_slice())]);
        let out = &outs[0];
        if self.act_b.len() != out.len() {
            self.act_b.resize(out.len(), 0.0);
        }
        self.act_b.copy_from_slice(out);
        Ok(())
    }
}