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)?)
}
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
}
}