use crate::config::{Brain2QwertyConfig, ConformerConfig, EncoderConfig, SimpleConvConfig};
use crate::model::channel_merger::ChannelMerger;
use crate::model::conformer::{ConformerLayer, ConformerStack, ConvModule, FeedForward};
use crate::model::fourier::FourierEmb;
use crate::model::simple_conv::{ConvBlock, SimpleConvEncoder};
use crate::model::subject_layers::SubjectLayers;
use crate::tensor::Tensor;
use crate::weights::{get, load_safetensors, param_to_tensor, take, WeightStore};
#[derive(Debug, Clone)]
pub struct ConvConformerOutput {
pub z: Tensor,
pub z_enc: Tensor,
pub z_transformer_in: Tensor,
pub z_final: Tensor,
pub c_out: Tensor,
pub z_aux: Option<Tensor>,
}
#[derive(Debug, Clone)]
pub struct EncoderPrefixOutput {
pub z: Tensor,
pub z_enc: Tensor,
pub z_transformer_in: Tensor,
pub z_aux: Option<Tensor>,
}
pub struct ConvConformer {
pub dim: usize,
pub aux_prediction: bool,
encoder: SimpleConvEncoder,
td_kernel: Tensor,
td_bias: Option<Tensor>,
td_ln_w: Option<Tensor>,
td_ln_b: Option<Tensor>,
shared_ln_w: Option<Tensor>,
shared_ln_b: Option<Tensor>,
intermediate_w: Option<Tensor>,
intermediate_b: Option<Tensor>,
pub conformer: ConformerStack,
output_w: Tensor,
output_b: Tensor,
td_kernel_size: usize,
td_stride: usize,
}
impl ConvConformer {
pub fn from_config_and_weights(
cfg: &EncoderConfig,
store: &mut WeightStore,
prefix: &str,
) -> anyhow::Result<Self> {
let enc_prefix = format!("{prefix}encoder.");
let encoder = load_simple_conv(&cfg.encoder_config, store, &enc_prefix)?;
let td_prefix = format!("{prefix}temporal_downsampling.");
let td_w = param_to_tensor(get(store, &format!("{td_prefix}agg.weight"))?);
let td_b = get(store, &format!("{td_prefix}agg.bias"))
.ok()
.map(param_to_tensor);
let td_ln_w = get(store, &format!("{td_prefix}layer_norm.weight"))
.ok()
.map(param_to_tensor);
let td_ln_b = get(store, &format!("{td_prefix}layer_norm.bias"))
.ok()
.map(param_to_tensor);
let conformer = load_conformer(
&cfg.transformer_config,
store,
&format!("{prefix}transformer."),
)?;
let output_w = param_to_tensor(get(store, &format!("{prefix}output_layer.weight"))?);
let output_b = param_to_tensor(get(store, &format!("{prefix}output_layer.bias"))?);
let (shared_ln_w, shared_ln_b, intermediate_w, intermediate_b) = if cfg.aux_prediction {
(
Some(param_to_tensor(get(
store,
&format!("{prefix}shared_layer_norm.weight"),
)?)),
Some(param_to_tensor(get(
store,
&format!("{prefix}shared_layer_norm.bias"),
)?)),
Some(param_to_tensor(get(
store,
&format!("{prefix}intermediate_linear.weight"),
)?)),
Some(param_to_tensor(get(
store,
&format!("{prefix}intermediate_linear.bias"),
)?)),
)
} else {
(None, None, None, None)
};
Ok(Self {
dim: cfg.dim,
aux_prediction: cfg.aux_prediction,
encoder,
td_kernel: td_w,
td_bias: td_b,
td_ln_w,
td_ln_b,
shared_ln_w,
shared_ln_b,
intermediate_w,
intermediate_b,
conformer,
output_w,
output_b,
td_kernel_size: cfg.temporal_downsampling_config.kernel_size,
td_stride: cfg.temporal_downsampling_config.stride,
})
}
pub fn from_pretrained(config_path: &str, weights_path: &str) -> anyhow::Result<Self> {
let cfg = Brain2QwertyConfig::from_yaml(config_path)?;
let mut store = load_safetensors(weights_path)?;
Self::from_config_and_weights(&cfg.brain_model_config, &mut store, "")
}
pub fn from_tiny_weights(weights_path: &str) -> anyhow::Result<Self> {
let cfg = Brain2QwertyConfig::tiny();
let mut store = load_safetensors(weights_path)?;
Self::from_config_and_weights(&cfg.brain_model_config, &mut store, "")
}
pub fn forward(
&self,
neuros: &Tensor,
subject_ids: &[usize],
chan_pos: Option<&Tensor>,
) -> ConvConformerOutput {
let prefix = self.forward_prefix(neuros, subject_ids, chan_pos);
let z_final = self.conformer.forward(&prefix.z_transformer_in);
let c_out = self.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 forward_prefix(
&self,
neuros: &Tensor,
subject_ids: &[usize],
chan_pos: Option<&Tensor>,
) -> EncoderPrefixOutput {
let x = neuros.transpose(&[0, 2, 1]);
let z_enc = self.encoder_and_downsample(&x, subject_ids, chan_pos);
let mut z = z_enc.clone();
let mut z_aux = None;
if self.aux_prediction {
let ln_w = self.shared_ln_w.as_ref().unwrap();
let ln_b = self.shared_ln_b.as_ref().unwrap();
z = z.layer_norm(ln_w, ln_b, 1e-5);
z_aux = Some(z.linear(&self.output_w, Some(&self.output_b)));
let sm = z_aux.as_ref().unwrap().softmax(2);
let blended = sm.linear(
self.intermediate_w.as_ref().unwrap(),
Some(self.intermediate_b.as_ref().unwrap()),
);
z = z.add(&blended);
}
let z_transformer_in = z.clone();
let z_out = if self.aux_prediction {
z_aux.clone().unwrap()
} else {
z_enc.clone()
};
EncoderPrefixOutput {
z: z_out,
z_enc,
z_transformer_in,
z_aux,
}
}
fn encoder_and_downsample(
&self,
x: &Tensor,
subject_ids: &[usize],
chan_pos: Option<&Tensor>,
) -> Tensor {
let mut z = self.encoder.forward(x, subject_ids, chan_pos);
z = z.transpose(&[0, 2, 1]);
z = self.temporal_downsample(&z);
z
}
fn temporal_downsample(&self, z: &Tensor) -> Tensor {
let (b, t, f) = (z.shape[0], z.shape[1], z.shape[2]);
let z4 = z.reshape(&[b, 1, t, f]);
let mut y = z4.conv2d(
&self.td_kernel,
self.td_bias.as_ref(),
[self.td_stride, 1],
[0, 0],
);
y = y.reshape(&[b, y.shape[2], f]);
if let (Some(w), Some(b)) = (&self.td_ln_w, &self.td_ln_b) {
y = y.layer_norm(w, b, 1e-8);
}
y.gelu()
}
pub fn forward_transformer(&self, z: &Tensor) -> Tensor {
self.conformer.forward(z)
}
pub fn forward_head(&self, z_final: &Tensor) -> Tensor {
let mut c_out = z_final.clone();
if self.aux_prediction {
let ln_w = self.shared_ln_w.as_ref().unwrap();
let ln_b = self.shared_ln_b.as_ref().unwrap();
c_out = c_out.layer_norm(ln_w, ln_b, 1e-5);
}
c_out.linear(&self.output_w, Some(&self.output_b))
}
pub fn compute_output_lens(&self, neuro_sizes: &[usize]) -> Vec<usize> {
neuro_sizes
.iter()
.map(|&n| crate::weights::compute_output_lens(n, self.td_kernel_size, self.td_stride))
.collect()
}
}
pub fn load_simple_conv(
cfg: &SimpleConvConfig,
store: &mut WeightStore,
prefix: &str,
) -> anyhow::Result<SimpleConvEncoder> {
let merger = if cfg.merger_config.n_virtual_channels > 0 {
let n_freqs = cfg.merger_config.fourier_emb_config.resolved_n_freqs()?;
let emb = FourierEmb::new(
n_freqs,
cfg.merger_config.fourier_emb_config.n_dims,
cfg.merger_config.fourier_emb_config.margin,
);
let heads = param_to_tensor(get(store, &format!("{prefix}merger.heads"))?);
Some(ChannelMerger {
embedding: emb,
heads,
per_subject: cfg.merger_config.per_subject,
n_virtual: cfg.merger_config.n_virtual_channels,
invalid_value: -0.1,
})
} else {
None
};
let initial_linear = if cfg.initial_linear > 0 {
let w = param_to_tensor(get(store, &format!("{prefix}initial_linear.0.weight"))?);
let b = param_to_tensor(get(store, &format!("{prefix}initial_linear.0.bias"))?);
Some((w, b))
} else {
None
};
let subject_layers = if cfg.subject_layers_config.is_some() {
let w = param_to_tensor(get(store, &format!("{prefix}subject_layers.weights"))?);
let b = param_to_tensor(get(store, &format!("{prefix}subject_layers.bias"))?);
Some(SubjectLayers {
weights: w,
bias: Some(b),
average_subjects: cfg
.subject_layers_config
.as_ref()
.map(|s| s.average_subjects)
.unwrap_or(false),
n_subjects: cfg
.subject_layers_config
.as_ref()
.map(|s| s.n_subjects)
.unwrap_or(200),
})
} else {
None
};
let mut blocks = Vec::new();
let mut dilation = 1usize;
for k in 0..cfg.depth {
if cfg.dilation_period.is_some_and(|p| k % p == 0) {
dilation = 1;
}
let is_last = k + 1 == cfg.depth;
let conv_idx = if k == 0 && cfg.dropout_input > 0.0 {
1
} else {
0
};
let bn_idx = conv_idx + 1;
let seq_prefix = format!("{prefix}encoder.sequence.{k}.");
let conv_w = param_to_tensor(&take(store, &format!("{seq_prefix}{conv_idx}.weight"))?);
let conv_b = param_to_tensor(&take(store, &format!("{seq_prefix}{conv_idx}.bias"))?);
let chin = conv_w.shape[1];
let chout = conv_w.shape[0];
let pad = cfg.kernel_size / 2 * dilation;
let use_activation = !is_last;
let block_skip = cfg.skip && chin == chout;
let scale_idx = if use_activation {
conv_idx + 4
} else {
conv_idx + 1
};
let bn_w = if use_activation && cfg.batch_norm {
get(store, &format!("{seq_prefix}{bn_idx}.weight"))
.ok()
.map(param_to_tensor)
} else {
None
};
let bn_b = if use_activation && cfg.batch_norm {
get(store, &format!("{seq_prefix}{bn_idx}.bias"))
.ok()
.map(param_to_tensor)
} else {
None
};
let bn_running_mean = if use_activation && cfg.batch_norm {
get(store, &format!("{seq_prefix}{bn_idx}.running_mean"))
.ok()
.map(param_to_tensor)
} else {
None
};
let bn_running_var = if use_activation && cfg.batch_norm {
get(store, &format!("{seq_prefix}{bn_idx}.running_var"))
.ok()
.map(param_to_tensor)
} else {
None
};
let layer_scale = if block_skip && cfg.scale.is_some() {
get(store, &format!("{seq_prefix}{scale_idx}.scale"))
.ok()
.map(param_to_tensor)
} else {
None
};
blocks.push(ConvBlock {
conv_w,
conv_b,
padding: pad,
dilation,
bn_w,
bn_b,
bn_running_mean,
bn_running_var,
layer_scale,
use_gelu: cfg.gelu,
leakiness: cfg.relu_leakiness,
use_activation,
skip: block_skip,
});
dilation *= cfg.dilation_growth;
}
Ok(SimpleConvEncoder {
merger,
initial_linear,
subject_layers,
blocks,
out_channels: cfg.hidden,
})
}
fn load_conformer(
cfg: &ConformerConfig,
store: &WeightStore,
prefix: &str,
) -> anyhow::Result<ConformerStack> {
let mut layers = Vec::new();
for i in 0..cfg.num_layers {
let p = format!("{prefix}conformer_layers.{i}.");
let ffn = |n: &str| -> anyhow::Result<FeedForward> {
Ok(FeedForward {
ln_w: param_to_tensor(get(store, &format!("{p}{n}sequential.0.weight"))?),
ln_b: param_to_tensor(get(store, &format!("{p}{n}sequential.0.bias"))?),
w1: param_to_tensor(get(store, &format!("{p}{n}sequential.1.weight"))?),
b1: param_to_tensor(get(store, &format!("{p}{n}sequential.1.bias"))?),
w2: param_to_tensor(get(store, &format!("{p}{n}sequential.4.weight"))?),
b2: param_to_tensor(get(store, &format!("{p}{n}sequential.4.bias"))?),
})
};
let conv = ConvModule {
ln_w: param_to_tensor(get(store, &format!("{p}conv_module.layer_norm.weight"))?),
ln_b: param_to_tensor(get(store, &format!("{p}conv_module.layer_norm.bias"))?),
p1_w: param_to_tensor(get(store, &format!("{p}conv_module.sequential.0.weight"))?),
p1_b: param_to_tensor(get(store, &format!("{p}conv_module.sequential.0.bias"))?),
dw_w: param_to_tensor(get(store, &format!("{p}conv_module.sequential.2.weight"))?),
dw_b: param_to_tensor(get(store, &format!("{p}conv_module.sequential.2.bias"))?),
gn_w: param_to_tensor(get(store, &format!("{p}conv_module.sequential.3.weight"))?),
gn_b: param_to_tensor(get(store, &format!("{p}conv_module.sequential.3.bias"))?),
p2_w: param_to_tensor(get(store, &format!("{p}conv_module.sequential.5.weight"))?),
p2_b: param_to_tensor(get(store, &format!("{p}conv_module.sequential.5.bias"))?),
kernel: cfg.depthwise_conv_kernel_size,
use_group_norm: cfg.use_group_norm,
};
layers.push(ConformerLayer {
ffn1: ffn("ffn1.")?,
sa_ln_w: param_to_tensor(get(store, &format!("{p}self_attn_layer_norm.weight"))?),
sa_ln_b: param_to_tensor(get(store, &format!("{p}self_attn_layer_norm.bias"))?),
in_proj_w: param_to_tensor(get(store, &format!("{p}self_attn.in_proj_weight"))?),
in_proj_b: param_to_tensor(get(store, &format!("{p}self_attn.in_proj_bias"))?),
out_proj_w: param_to_tensor(get(store, &format!("{p}self_attn.out_proj.weight"))?),
out_proj_b: param_to_tensor(get(store, &format!("{p}self_attn.out_proj.bias"))?),
conv,
ffn2: ffn("ffn2.")?,
final_ln_w: param_to_tensor(get(store, &format!("{p}final_layer_norm.weight"))?),
final_ln_b: param_to_tensor(get(store, &format!("{p}final_layer_norm.bias"))?),
num_heads: cfg.num_heads,
convolution_first: cfg.convolution_first,
});
}
Ok(ConformerStack { layers })
}