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 {
matches!(device, rlx::Device::Cuda | rlx::Device::Ane)
}
fn use_monolithic_tail(device: rlx::Device) -> bool {
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
}
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(())
}
}