use crate::config::SimpleConvConfig;
#[derive(Clone, Debug, serde::Serialize, serde::Deserialize)]
pub struct V1TransformerConfig {
pub name: String,
pub alibi_pos_bias: bool,
pub depth: usize,
pub heads: usize,
#[serde(default = "default_ff_mult")]
pub ff_mult: usize,
#[serde(default = "default_true")]
pub use_scalenorm: bool,
#[serde(default = "default_false")]
pub rotary_pos_emb: bool,
#[serde(default = "default_true")]
pub scale_residual: bool,
}
fn default_ff_mult() -> usize {
4
}
fn default_true() -> bool {
true
}
fn default_false() -> bool {
false
}
impl Default for V1TransformerConfig {
fn default() -> Self {
Self {
name: "TransformerEncoder".into(),
alibi_pos_bias: true,
depth: 4,
heads: 2,
ff_mult: 4,
use_scalenorm: true,
rotary_pos_emb: true,
scale_residual: true,
}
}
}
#[derive(Clone, Debug, serde::Serialize, serde::Deserialize)]
pub struct V1EncoderConfig {
pub name: String,
pub time_agg_out: String,
pub encoder_config: SimpleConvConfig,
}
#[derive(Clone, Debug, serde::Serialize, serde::Deserialize)]
pub struct V1Config {
pub brain_model_config: V1EncoderConfig,
pub transformer_config: V1TransformerConfig,
pub num_classes: usize,
pub n_in_channels: usize,
}
impl V1Config {
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 production() -> Self {
Self::from_yaml(&format!(
"{}/../../data/v1_config.yaml",
env!("CARGO_MANIFEST_DIR")
))
.expect("v1_config.yaml")
}
pub fn tiny() -> Self {
let mut cfg = Self::production();
cfg.brain_model_config.encoder_config.hidden = 32;
cfg.brain_model_config.encoder_config.depth = 2;
cfg.brain_model_config.encoder_config.initial_linear = 16;
cfg.brain_model_config
.encoder_config
.merger_config
.n_virtual_channels = 16;
cfg.brain_model_config
.encoder_config
.merger_config
.fourier_emb_config
.total_dim = Some(32);
cfg.transformer_config.depth = 1;
cfg.transformer_config.heads = 1;
cfg.transformer_config.rotary_pos_emb = false;
cfg
}
}