use anyhow::{Result, anyhow};
use pest::Parser;
use pest_derive::Parser;
use rust_embed::Embed;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
#[derive(Embed)]
#[folder = "well-known/"]
#[include = "*.mal"]
struct WellKnown;
#[derive(Parser)]
#[grammar = "mal.pest"]
pub struct MalParser;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum PositionEncoding {
Rope { theta: f64, scaling: Option<f64> },
Alibi { learned_slopes: bool },
Learned { max_positions: usize },
None,
}
impl Default for PositionEncoding {
fn default() -> Self {
Self::Rope {
theta: 10000.0,
scaling: None,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AttentionDef {
pub name: String,
pub num_heads: Option<usize>,
pub num_kv_heads: Option<usize>,
pub head_dim: Option<usize>,
pub dropout: f64,
pub bias: bool,
pub position_encoding: PositionEncoding,
pub window_size: Option<usize>,
pub causal: bool,
#[serde(default)]
pub qk_norm: bool,
}
impl Default for AttentionDef {
fn default() -> Self {
Self {
name: "default".to_string(),
num_heads: None,
num_kv_heads: None,
head_dim: None,
dropout: 0.0,
bias: false,
position_encoding: PositionEncoding::default(),
window_size: None,
causal: true,
qk_norm: false,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
pub enum NormType {
#[default]
RmsNorm,
LayerNorm,
None,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct NormConfig {
pub norm_type: NormType,
pub eps: f64,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
pub enum Activation {
#[default]
SwiGLU,
GELU,
SiLU,
ReLU,
GELUNew,
GELUTanh,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SsmDef {
pub name: String,
pub state_dim: usize,
pub conv_kernel: usize,
pub expand: usize,
pub dt_rank: Option<usize>,
}
impl Default for SsmDef {
fn default() -> Self {
Self {
name: "default".to_string(),
state_dim: 16,
conv_kernel: 4,
expand: 2,
dt_rank: None,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FfnDef {
pub name: String,
pub hidden_dim: Option<usize>,
pub activation: Activation,
pub bias: bool,
pub dropout: f64,
pub gate: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub moe: Option<MoeDef>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MoeDef {
pub experts: usize,
pub top_k: usize,
#[serde(default)]
pub shared_experts: usize,
#[serde(default)]
pub load_balance_loss_weight: f64,
#[serde(default)]
pub router_z_loss_weight: f64,
}
impl Default for FfnDef {
fn default() -> Self {
Self {
name: "default".to_string(),
hidden_dim: None,
activation: Activation::default(),
bias: false,
dropout: 0.0,
gate: true,
moe: None,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BlockDef {
pub name: String,
pub attention: AttentionDef,
#[serde(default)]
pub ssm: Option<SsmDef>,
pub ffn: FfnDef,
pub norm: NormConfig,
pub norm_position: NormPosition,
pub residual: bool,
pub dropout: f64,
}
impl BlockDef {
pub fn is_ssm(&self) -> bool {
self.ssm.is_some()
}
pub fn num_heads(&self) -> usize {
self.attention.num_heads.unwrap_or(12)
}
pub fn num_kv_heads(&self) -> usize {
self.attention.num_kv_heads.unwrap_or(self.num_heads())
}
pub fn head_dim(&self, hidden_size: usize) -> usize {
self.attention
.head_dim
.unwrap_or(hidden_size / self.num_heads())
}
pub fn intermediate_size(&self, hidden_size: usize) -> usize {
self.ffn.hidden_dim.unwrap_or(hidden_size * 4)
}
pub fn norm_eps(&self) -> f64 {
if self.norm.eps > 0.0 {
self.norm.eps
} else {
1e-5
}
}
pub fn rope_theta(&self) -> f64 {
match &self.attention.position_encoding {
PositionEncoding::Rope { theta, .. } => *theta,
_ => 10000.0,
}
}
pub fn rope_scaling(&self) -> Option<f64> {
match &self.attention.position_encoding {
PositionEncoding::Rope { scaling, .. } => *scaling,
_ => None,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
pub enum NormPosition {
#[default]
Pre,
Post,
}
impl Default for BlockDef {
fn default() -> Self {
Self {
name: "default".to_string(),
attention: AttentionDef::default(),
ssm: None,
ffn: FfnDef::default(),
norm: NormConfig {
norm_type: NormType::RmsNorm,
eps: 1e-5,
},
norm_position: NormPosition::Pre,
residual: true,
dropout: 0.0,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct EmbeddingsConfig {
pub tie_weights: bool,
pub dropout: f64,
pub scale: Option<f64>,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct OutputConfig {
pub bias: bool,
pub norm: Option<NormConfig>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelDef {
pub name: String,
pub description: Option<String>,
pub vocab_size: usize,
pub max_seq_len: usize,
pub hidden_size: usize,
pub num_layers: usize,
pub block: BlockDef,
#[serde(default)]
pub pattern: Option<Vec<BlockDef>>,
pub embeddings: EmbeddingsConfig,
pub output: OutputConfig,
}
impl Default for ModelDef {
fn default() -> Self {
Self {
name: "default".to_string(),
description: None,
vocab_size: 32000,
max_seq_len: 2048,
hidden_size: 768,
num_layers: 12,
block: BlockDef::default(),
pattern: None,
embeddings: EmbeddingsConfig::default(),
output: OutputConfig::default(),
}
}
}
impl std::fmt::Display for ModelDef {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
writeln!(f, "model {} {{", self.name)?;
if let Some(desc) = &self.description {
writeln!(f, " description: \"{}\"", desc)?;
}
writeln!(f, " vocab_size: {}", self.vocab_size)?;
writeln!(f, " max_seq_len: {}", self.max_seq_len)?;
writeln!(f, " hidden_size: {}", self.hidden_size)?;
writeln!(f, " num_layers: {}", self.num_layers)?;
writeln!(f, "}}")?;
writeln!(f)?;
writeln!(f, "attention {{")?;
if let Some(h) = self.block.attention.num_heads {
writeln!(f, " num_heads: {}", h)?;
}
if let Some(kv) = self.block.attention.num_kv_heads {
writeln!(f, " num_kv_heads: {}", kv)?;
}
if let Some(hd) = self.block.attention.head_dim {
writeln!(f, " head_dim: {}", hd)?;
}
writeln!(f, " bias: {}", self.block.attention.bias)?;
writeln!(f, "}}")?;
writeln!(f)?;
writeln!(f, "ffn {{")?;
if let Some(dim) = self.block.ffn.hidden_dim {
writeln!(f, " hidden_dim: {}", dim)?;
}
writeln!(f, " activation: {:?}", self.block.ffn.activation)?;
writeln!(f, " bias: {}", self.block.ffn.bias)?;
writeln!(f, "}}")?;
writeln!(f)?;
writeln!(f, "block {{")?;
writeln!(f, " norm: {:?}", self.block.norm.norm_type)?;
writeln!(f, " norm_position: {:?}", self.block.norm_position)?;
writeln!(f, " residual: {}", self.block.residual)?;
writeln!(f, "}}")?;
writeln!(f)?;
let params = self.estimated_params();
writeln!(
f,
"Estimated parameters: {:.2}B",
params as f64 / 1_000_000_000.0
)
}
}
impl ModelDef {
pub fn block_for_layer(&self, i: usize) -> &BlockDef {
match &self.pattern {
Some(p) if !p.is_empty() => &p[i % p.len()],
_ => &self.block,
}
}
pub fn dt_rank(&self, ssm: &SsmDef) -> usize {
ssm.dt_rank.unwrap_or(self.hidden_size.div_ceil(16))
}
pub fn num_heads(&self) -> usize {
self.block.attention.num_heads.unwrap_or(12)
}
pub fn num_kv_heads(&self) -> usize {
self.block
.attention
.num_kv_heads
.unwrap_or(self.num_heads())
}
pub fn head_dim(&self) -> usize {
self.block
.attention
.head_dim
.unwrap_or(self.hidden_size / self.num_heads())
}
pub fn intermediate_size(&self) -> usize {
self.block.ffn.hidden_dim.unwrap_or(self.hidden_size * 4)
}
pub fn norm_eps(&self) -> f64 {
if self.block.norm.eps > 0.0 {
self.block.norm.eps
} else {
1e-5
}
}
pub fn rope_theta(&self) -> f64 {
match &self.block.attention.position_encoding {
PositionEncoding::Rope { theta, .. } => *theta,
_ => 10000.0,
}
}
pub fn padded_vocab_size(&self) -> usize {
self.vocab_size.next_multiple_of(64)
}
pub fn estimated_params(&self) -> usize {
let h = self.hidden_size;
let stored_vocab_size = self.padded_vocab_size();
let embed_params = stored_vocab_size * h;
let norm_params = |norm: &NormConfig| match norm.norm_type {
NormType::RmsNorm => h,
NormType::LayerNorm => 2 * h,
NormType::None => 0,
};
let mut layer_params = 0usize;
for i in 0..self.num_layers {
let block = self.block_for_layer(i);
let mixer = match &block.ssm {
Some(ssm) => {
let d_inner = ssm.expand * h;
let dt_rank = self.dt_rank(ssm);
2 * h * d_inner + d_inner * h + d_inner * ssm.conv_kernel
+ d_inner * (dt_rank + 2 * ssm.state_dim) + dt_rank * d_inner + d_inner + d_inner * ssm.state_dim + d_inner + d_inner }
None => {
let q = block.num_heads() * block.head_dim(h);
let kv = block.num_kv_heads() * block.head_dim(h);
let weights = h * q + 2 * h * kv + q * h;
let bias = if block.attention.bias {
2 * q + 2 * kv
} else {
0
};
let qk_norm = if block.attention.qk_norm {
2 * block.head_dim(h)
} else {
0
};
weights + bias + qk_norm
}
};
let intermediate = block.intermediate_size(h);
let projections = if block.ffn.gate { 3 } else { 2 };
let expert_count = block
.ffn
.moe
.as_ref()
.map_or(1, |moe| moe.experts + moe.shared_experts);
let ff_weights = expert_count * projections * h * intermediate;
let ff_bias = if block.ffn.bias {
expert_count * ((if block.ffn.gate { 2 } else { 1 }) * intermediate + h)
} else {
0
};
let router = block.ffn.moe.as_ref().map_or(0, |moe| h * moe.experts);
layer_params += mixer + ff_weights + ff_bias + router + 2 * norm_params(&block.norm);
}
let final_norm = self
.output
.norm
.as_ref()
.unwrap_or(&self.block_for_layer(0).norm);
let head_weights = (!self.embeddings.tie_weights) as usize * h * stored_vocab_size;
let head_bias = self.output.bias as usize * stored_vocab_size;
embed_params + layer_params + norm_params(final_norm) + head_weights + head_bias
}
pub fn from_json<P: AsRef<std::path::Path>>(path: P) -> Result<Self> {
let content = std::fs::read_to_string(path)?;
Ok(serde_json::from_str(&content)?)
}
pub fn save_json(&self, path: &str) -> Result<()> {
let content = serde_json::to_string_pretty(self)?;
std::fs::write(path, content)?;
Ok(())
}
}
#[derive(Debug, Clone, Default)]
pub struct MalFile {
pub attentions: HashMap<String, AttentionDef>,
pub ssms: HashMap<String, SsmDef>,
pub ffns: HashMap<String, FfnDef>,
pub blocks: HashMap<String, BlockDef>,
pub models: HashMap<String, ModelDef>,
}
fn parse_activation(s: &str) -> Activation {
match s {
"swiglu" => Activation::SwiGLU,
"gelu" => Activation::GELU,
"silu" => Activation::SiLU,
"relu" => Activation::ReLU,
"gelu_new" => Activation::GELUNew,
"gelu_tanh" => Activation::GELUTanh,
_ => Activation::SwiGLU,
}
}
fn parse_model_prop(
pair: pest::iterators::Pair<Rule>,
def: &mut ModelDef,
file: &MalFile,
) -> Result<()> {
for inner in pair.into_inner() {
match inner.as_rule() {
Rule::vocab_size_prop => {
if let Some(val) = inner.into_inner().next() {
def.vocab_size = val.as_str().parse()?;
}
}
Rule::max_seq_len_prop => {
if let Some(val) = inner.into_inner().next() {
def.max_seq_len = val.as_str().parse()?;
}
}
Rule::hidden_size_prop => {
if let Some(val) = inner.into_inner().next() {
def.hidden_size = val.as_str().parse()?;
}
}
Rule::num_layers_prop => {
if let Some(val) = inner.into_inner().next() {
def.num_layers = val.as_str().parse()?;
}
}
Rule::block_ref_prop => {
for child in inner.into_inner() {
match child.as_rule() {
Rule::identifier => {
let name = child.as_str();
def.block = file
.blocks
.get(name)
.ok_or_else(|| anyhow!("undefined block '{name}'"))?
.clone();
}
Rule::inline_block => {
let mut block = BlockDef::default();
for prop in child.into_inner() {
if prop.as_rule() == Rule::block_prop {
parse_block_prop(prop, &mut block, file)?;
}
}
def.block = block;
}
_ => {}
}
}
}
Rule::pattern_prop => {
let mut blocks = Vec::new();
for child in inner.into_inner() {
if child.as_rule() == Rule::identifier {
let name = child.as_str();
let block = file.blocks.get(name).ok_or_else(|| {
anyhow!("pattern references undefined block '{}'", name)
})?;
blocks.push(block.clone());
}
}
if !blocks.is_empty() {
def.pattern = Some(blocks);
}
}
Rule::embeddings_prop => {
for param in inner.into_inner() {
for child in param.into_inner() {
match child.as_rule() {
Rule::tie_weights_prop => {
if let Some(val) = child.into_inner().next() {
def.embeddings.tie_weights = val.as_str() == "true";
}
}
Rule::dropout_prop => {
if let Some(val) = child.into_inner().next() {
def.embeddings.dropout = val.as_str().parse()?;
}
}
Rule::scale_prop => {
if let Some(val) = child.into_inner().next() {
def.embeddings.scale = Some(val.as_str().parse()?);
}
}
_ => {}
}
}
}
}
Rule::output_prop => {
for param in inner.into_inner() {
for child in param.into_inner() {
match child.as_rule() {
Rule::bias_prop => {
if let Some(val) = child.into_inner().next() {
def.output.bias = val.as_str() == "true";
}
}
Rule::norm_prop => {
if let Some(cfg) = child.into_inner().next() {
def.output.norm = Some(parse_norm_config(cfg)?);
}
}
_ => {}
}
}
}
}
Rule::description_prop => {
if let Some(val) = inner.into_inner().next() {
let s = val.as_str();
def.description = Some(s[1..s.len() - 1].to_string());
}
}
_ => {}
}
}
Ok(())
}
fn parse_model_def(pair: pest::iterators::Pair<Rule>, file: &MalFile) -> Result<ModelDef> {
let mut def = ModelDef::default();
let mut inner = pair.into_inner();
if let Some(name) = inner.next() {
def.name = name.as_str().to_string();
}
for prop in inner {
if prop.as_rule() == Rule::model_prop {
parse_model_prop(prop, &mut def, file)?;
}
}
Ok(def)
}
pub fn parse_mal(input: &str) -> Result<ModelDef> {
let file = parse_mal_full(input)?;
match file.models.len() {
0 => Err(anyhow!("no model definition found")),
1 => Ok(file.models.into_values().next().expect("length checked")),
_ => {
let mut names = file.models.keys().cloned().collect::<Vec<_>>();
names.sort();
Err(anyhow!(
"multiple model definitions found ({}); use parse_mal_full to select one",
names.join(", ")
))
}
}
}
fn insert_unique<T>(
definitions: &mut HashMap<String, T>,
kind: &str,
name: String,
value: T,
) -> Result<()> {
match definitions.entry(name) {
std::collections::hash_map::Entry::Vacant(entry) => {
entry.insert(value);
Ok(())
}
std::collections::hash_map::Entry::Occupied(entry) => {
Err(anyhow!("duplicate {kind} '{}'", entry.key()))
}
}
}
pub fn parse_mal_full(input: &str) -> Result<MalFile> {
let pairs = MalParser::parse(Rule::file, input).map_err(|e| anyhow!("Parse error: {}", e))?;
let mut file = MalFile::default();
for pair in pairs {
if pair.as_rule() == Rule::file {
for inner in pair.into_inner() {
if inner.as_rule() == Rule::definition {
for def in inner.into_inner() {
match def.as_rule() {
Rule::model_def => {
let model = parse_model_def(def, &file)?;
insert_unique(
&mut file.models,
"model",
model.name.clone(),
model,
)?;
}
Rule::attention_def => {
let attn = parse_attention_def(def)?;
insert_unique(
&mut file.attentions,
"attention",
attn.name.clone(),
attn,
)?;
}
Rule::ssm_def => {
let ssm = parse_ssm_def(def)?;
insert_unique(&mut file.ssms, "ssm", ssm.name.clone(), ssm)?;
}
Rule::ffn_def => {
let ffn = parse_ffn_def(def)?;
insert_unique(&mut file.ffns, "ffn", ffn.name.clone(), ffn)?;
}
Rule::block_def => {
let block = parse_block_def(def, &file)?;
insert_unique(
&mut file.blocks,
"block",
block.name.clone(),
block,
)?;
}
_ => {}
}
}
}
}
}
}
Ok(file)
}
fn parse_attention_def(pair: pest::iterators::Pair<Rule>) -> Result<AttentionDef> {
let mut def = AttentionDef::default();
let mut inner = pair.into_inner();
if let Some(name) = inner.next() {
def.name = name.as_str().to_string();
}
for prop in inner {
if prop.as_rule() == Rule::attention_prop {
parse_attention_prop(prop, &mut def)?;
}
}
Ok(def)
}
fn parse_position_encoding(pair: pest::iterators::Pair<Rule>) -> Result<PositionEncoding> {
let Some(config) = pair.into_inner().next() else {
return Ok(PositionEncoding::None);
};
match config.as_rule() {
Rule::rope_config => {
let mut theta = 10000.0;
let mut scaling = None;
for param in config.into_inner() {
for inner in param.into_inner() {
match inner.as_rule() {
Rule::rope_theta_prop | Rule::rope_base_prop => {
if let Some(val) = inner.into_inner().next() {
theta = val.as_str().parse()?;
}
}
Rule::rope_scaling_prop => {
if let Some(val) = inner.into_inner().next() {
scaling = Some(val.as_str().parse()?);
}
}
_ => {}
}
}
}
Ok(PositionEncoding::Rope { theta, scaling })
}
Rule::alibi_config => {
let learned_slopes = config.as_str().contains("learned");
Ok(PositionEncoding::Alibi { learned_slopes })
}
Rule::learned_config => {
let mut max_positions = 0;
for param in config.into_inner() {
for inner in param.into_inner() {
if inner.as_rule() == Rule::max_positions_prop
&& let Some(val) = inner.into_inner().next()
{
max_positions = val.as_str().parse()?;
}
}
}
Ok(PositionEncoding::Learned { max_positions })
}
_ => Ok(PositionEncoding::None),
}
}
fn parse_attention_prop(pair: pest::iterators::Pair<Rule>, def: &mut AttentionDef) -> Result<()> {
for inner in pair.into_inner() {
match inner.as_rule() {
Rule::num_heads_prop => {
if let Some(val) = inner.into_inner().next() {
def.num_heads = Some(val.as_str().parse()?);
}
}
Rule::num_kv_heads_prop => {
if let Some(val) = inner.into_inner().next() {
def.num_kv_heads = Some(val.as_str().parse()?);
}
}
Rule::head_dim_prop => {
if let Some(val) = inner.into_inner().next() {
def.head_dim = Some(val.as_str().parse()?);
}
}
Rule::dropout_prop => {
if let Some(val) = inner.into_inner().next() {
def.dropout = val.as_str().parse()?;
}
}
Rule::bias_prop => {
if let Some(val) = inner.into_inner().next() {
def.bias = val.as_str() == "true";
}
}
Rule::causal_prop => {
if let Some(val) = inner.into_inner().next() {
def.causal = val.as_str() == "true";
}
}
Rule::window_size_prop => {
if let Some(val) = inner.into_inner().next() {
def.window_size = Some(val.as_str().parse()?);
}
}
Rule::position_encoding_prop => {
if let Some(config) = inner.into_inner().next() {
def.position_encoding = parse_position_encoding(config)?;
}
}
Rule::qk_norm_prop => {
if let Some(val) = inner.into_inner().next() {
def.qk_norm = val.as_str() == "true";
}
}
_ => {}
}
}
Ok(())
}
fn parse_ssm_def(pair: pest::iterators::Pair<Rule>) -> Result<SsmDef> {
let mut def = SsmDef::default();
let mut inner = pair.into_inner();
if let Some(name) = inner.next() {
def.name = name.as_str().to_string();
}
for prop in inner {
if prop.as_rule() == Rule::ssm_prop {
parse_ssm_prop(prop, &mut def)?;
}
}
Ok(def)
}
fn parse_ssm_prop(pair: pest::iterators::Pair<Rule>, def: &mut SsmDef) -> Result<()> {
for inner in pair.into_inner() {
match inner.as_rule() {
Rule::state_dim_prop => {
if let Some(val) = inner.into_inner().next() {
def.state_dim = val.as_str().parse()?;
}
}
Rule::conv_kernel_prop => {
if let Some(val) = inner.into_inner().next() {
def.conv_kernel = val.as_str().parse()?;
}
}
Rule::expand_prop => {
if let Some(val) = inner.into_inner().next() {
def.expand = val.as_str().parse()?;
}
}
Rule::dt_rank_prop => {
if let Some(val) = inner.into_inner().next() {
def.dt_rank = Some(val.as_str().parse()?);
}
}
_ => {}
}
}
Ok(())
}
fn parse_norm_config(pair: pest::iterators::Pair<Rule>) -> Result<NormConfig> {
let mut norm = NormConfig::default();
match pair.into_inner().next() {
None => norm.norm_type = NormType::None,
Some(cfg) => {
norm.norm_type = match cfg.as_rule() {
Rule::rmsnorm_config => NormType::RmsNorm,
Rule::layernorm_config => NormType::LayerNorm,
other => anyhow::bail!("unexpected norm config rule: {other:?}"),
};
for param in cfg.into_inner() {
if let Some(prop) = param.into_inner().next()
&& let Some(number) = prop.into_inner().next()
{
norm.eps = number.as_str().parse()?;
}
}
}
}
Ok(norm)
}
fn parse_ffn_def(pair: pest::iterators::Pair<Rule>) -> Result<FfnDef> {
let mut def = FfnDef::default();
let mut inner = pair.into_inner();
if let Some(name) = inner.next() {
def.name = name.as_str().to_string();
}
for prop in inner {
if prop.as_rule() == Rule::ffn_prop {
parse_ffn_prop(prop, &mut def)?;
}
}
Ok(def)
}
fn parse_ffn_prop(pair: pest::iterators::Pair<Rule>, def: &mut FfnDef) -> Result<()> {
for inner in pair.into_inner() {
match inner.as_rule() {
Rule::hidden_dim_prop => {
if let Some(val) = inner.into_inner().next() {
def.hidden_dim = Some(val.as_str().parse()?);
}
}
Rule::activation_prop => {
if let Some(val) = inner.into_inner().next() {
def.activation = parse_activation(val.as_str());
}
}
Rule::bias_prop => {
if let Some(val) = inner.into_inner().next() {
def.bias = val.as_str() == "true";
}
}
Rule::dropout_prop => {
if let Some(val) = inner.into_inner().next() {
def.dropout = val.as_str().parse()?;
}
}
Rule::gate_prop => {
if let Some(val) = inner.into_inner().next() {
def.gate = val.as_str() == "true";
}
}
Rule::moe_prop => {
let mut moe = MoeDef {
experts: 0,
top_k: 0,
shared_experts: 0,
load_balance_loss_weight: 0.0,
router_z_loss_weight: 0.0,
};
for param in inner.into_inner() {
let Some(prop) = param.into_inner().next() else {
continue;
};
let value = prop
.clone()
.into_inner()
.next()
.map(|value| value.as_str())
.unwrap_or_default();
match prop.as_rule() {
Rule::experts_prop => moe.experts = value.parse()?,
Rule::top_k_prop => moe.top_k = value.parse()?,
Rule::shared_experts_prop => moe.shared_experts = value.parse()?,
Rule::load_balance_loss_weight_prop => {
moe.load_balance_loss_weight = value.parse()?
}
Rule::router_z_loss_weight_prop => {
moe.router_z_loss_weight = value.parse()?
}
_ => {}
}
}
def.moe = Some(moe);
}
_ => {}
}
}
Ok(())
}
fn parse_block_def(pair: pest::iterators::Pair<Rule>, file: &MalFile) -> Result<BlockDef> {
let mut def = BlockDef::default();
let mut inner = pair.into_inner();
if let Some(name) = inner.next() {
def.name = name.as_str().to_string();
}
for prop in inner {
if prop.as_rule() == Rule::block_prop {
parse_block_prop(prop, &mut def, file)?;
}
}
Ok(def)
}
fn parse_block_prop(
pair: pest::iterators::Pair<Rule>,
def: &mut BlockDef,
file: &MalFile,
) -> Result<()> {
for inner in pair.into_inner() {
match inner.as_rule() {
Rule::attention_ref_prop => {
for child in inner.into_inner() {
match child.as_rule() {
Rule::identifier => {
let name = child.as_str();
def.attention = file
.attentions
.get(name)
.ok_or_else(|| anyhow!("undefined attention '{name}'"))?
.clone();
}
Rule::inline_attention => {
let mut attn = AttentionDef::default();
for prop in child.into_inner() {
if prop.as_rule() == Rule::attention_prop {
parse_attention_prop(prop, &mut attn)?;
}
}
def.attention = attn;
}
_ => {}
}
}
}
Rule::ssm_ref_prop => {
for child in inner.into_inner() {
match child.as_rule() {
Rule::identifier => {
let name = child.as_str();
def.ssm = Some(
file.ssms
.get(name)
.ok_or_else(|| anyhow!("undefined ssm '{name}'"))?
.clone(),
);
}
Rule::inline_ssm => {
let mut ssm = SsmDef::default();
for prop in child.into_inner() {
if prop.as_rule() == Rule::ssm_prop {
parse_ssm_prop(prop, &mut ssm)?;
}
}
def.ssm = Some(ssm);
}
_ => {}
}
}
}
Rule::ffn_ref_prop => {
for child in inner.into_inner() {
match child.as_rule() {
Rule::identifier => {
let name = child.as_str();
def.ffn = file
.ffns
.get(name)
.ok_or_else(|| anyhow!("undefined ffn '{name}'"))?
.clone();
}
Rule::inline_ffn => {
let mut ffn = FfnDef::default();
for prop in child.into_inner() {
if prop.as_rule() == Rule::ffn_prop {
parse_ffn_prop(prop, &mut ffn)?;
}
}
def.ffn = ffn;
}
_ => {}
}
}
}
Rule::norm_prop => {
if let Some(cfg) = inner.into_inner().next() {
def.norm = parse_norm_config(cfg)?;
}
}
Rule::norm_position_prop => {
if let Some(val) = inner.into_inner().next() {
def.norm_position = match val.as_str() {
"pre" => NormPosition::Pre,
"post" => NormPosition::Post,
_ => NormPosition::Pre,
};
}
}
Rule::residual_prop => {
if let Some(val) = inner.into_inner().next() {
def.residual = val.as_str() == "true";
}
}
Rule::dropout_prop => {
if let Some(val) = inner.into_inner().next() {
def.dropout = val.as_str().parse()?;
}
}
_ => {}
}
}
Ok(())
}
pub fn parse_mal_file<P: AsRef<std::path::Path>>(path: P) -> Result<ModelDef> {
let content = std::fs::read_to_string(path)?;
parse_mal(&content)
}
pub fn get_builtin_model(name: &str) -> Option<ModelDef> {
let mal = get_wellknown_mal(name)?;
parse_mal(&mal).ok()
}
pub fn get_wellknown_mal(name: &str) -> Option<String> {
let name = name.strip_prefix("well-known/").unwrap_or(name);
let filename = if name.ends_with(".mal") {
name.to_string()
} else {
format!("{}.mal", name.replace('-', "_"))
};
WellKnown::get(&filename).map(|f| String::from_utf8_lossy(&f.data).into_owned())
}
pub fn list_wellknown_models() -> Vec<String> {
WellKnown::iter()
.filter_map(|path| {
let path: &str = path.as_ref();
if path.ends_with(".mal") {
Some(path.strip_suffix(".mal").unwrap().replace('_', "-"))
} else {
None
}
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_simple_model() {
let mal = r#"
attention test_attn {
num_heads: 8
bias: false
}
ffn test_ffn {
hidden_dim: 2048
activation: gelu
}
block test_block {
attention: test_attn
ffn: test_ffn
norm_position: pre
}
model test {
vocab_size: 32000
hidden_size: 512
num_layers: 8
block: test_block
}
"#;
let def = parse_mal(mal).unwrap();
assert_eq!(def.name, "test");
assert_eq!(def.vocab_size, 32000);
assert_eq!(def.hidden_size, 512);
assert_eq!(def.num_layers, 8);
}
#[test]
fn test_parse_with_block_props() {
let mal = r#"
attention full_attn {
num_heads: 16
num_kv_heads: 4
bias: true
dropout: 0.1
}
ffn full_ffn {
hidden_dim: 4096
activation: gelu
bias: true
dropout: 0.1
}
block full_block {
attention: full_attn
ffn: full_ffn
norm: layernorm { eps: 1e-6 }
norm_position: pre
residual: true
}
model full_test {
description: "A test model"
vocab_size: 50000
max_seq_len: 4096
hidden_size: 1024
num_layers: 12
block: full_block
}
"#;
let def = parse_mal(mal).unwrap();
assert_eq!(def.description, Some("A test model".to_string()));
assert_eq!(def.vocab_size, 50000);
assert_eq!(def.max_seq_len, 4096);
assert_eq!(def.block.attention.num_heads, Some(16));
assert_eq!(def.block.attention.num_kv_heads, Some(4));
assert_eq!(def.block.ffn.hidden_dim, Some(4096));
assert!(matches!(def.block.ffn.activation, Activation::GELU));
assert!(matches!(def.block.norm.norm_type, NormType::LayerNorm));
assert_eq!(def.block.norm.eps, 1e-6);
}
#[test]
fn moe_is_optional_and_fully_configurable() {
let dense = parse_mal(
"model d { vocab_size: 64 max_seq_len: 16 hidden_size: 8 num_layers: 1 \
block: { attention: { num_heads: 1 } ffn: { hidden_dim: 12 } } }",
)
.unwrap();
assert!(dense.block.ffn.moe.is_none());
let moe = parse_mal(
r#"
ffn experts {
hidden_dim: 12
activation: swiglu
moe {
experts: 8
top_k: 2
shared_experts: 1
load_balance_loss_weight: 0.01
router_z_loss_weight: 0.001
}
}
model m {
vocab_size: 64 max_seq_len: 16 hidden_size: 8 num_layers: 1
block: { attention: { num_heads: 1 } ffn: experts }
}
"#,
)
.unwrap();
let config = moe.block.ffn.moe.as_ref().unwrap();
assert_eq!(config.experts, 8);
assert_eq!(config.top_k, 2);
assert_eq!(config.shared_experts, 1);
assert_eq!(config.load_balance_loss_weight, 0.01);
assert_eq!(config.router_z_loss_weight, 0.001);
assert!(moe.estimated_params() > dense.estimated_params());
}
#[test]
fn retriever_200m_moe_has_the_intended_sparse_budget() {
let model = get_builtin_model("retriever-200m-moe").unwrap();
assert_eq!(model.estimated_params(), 200_795_648);
assert_eq!(model.num_layers, 24);
assert_eq!(model.pattern.as_ref().unwrap().len(), 6);
let moe_layers: Vec<_> = (0..model.num_layers)
.filter_map(|layer| model.block_for_layer(layer).ffn.moe.as_ref())
.collect();
assert_eq!(moe_layers.len(), 12);
assert!(moe_layers.iter().all(|moe| moe.experts == 8));
assert!(moe_layers.iter().all(|moe| moe.top_k == 2));
assert!(moe_layers.iter().all(|moe| moe.shared_experts == 0));
assert_eq!(model.block_for_layer(23).name, "attn_moe_block");
}
#[test]
fn test_norm_none_and_rmsnorm() {
let base = |norm: &str| {
format!(
r#"
block b {{ attention: {{ num_heads: 4 }} ffn: {{ hidden_dim: 64 }} norm: {norm} }}
model m {{ vocab_size: 100 max_seq_len: 64 hidden_size: 32 num_layers: 2 block: b }}
"#
)
};
let d = parse_mal(&base("rmsnorm { eps: 1e-5 }")).unwrap();
assert!(matches!(d.block.norm.norm_type, NormType::RmsNorm));
let d = parse_mal(&base("none")).unwrap();
assert!(matches!(d.block.norm.norm_type, NormType::None));
}
#[test]
fn test_embedding_and_output_configuration() {
let def = parse_mal(
r#"
block b { attention: { num_heads: 4 } ffn: { hidden_dim: 64 } }
model m {
vocab_size: 100
max_seq_len: 64
hidden_size: 32
num_layers: 2
block: b
embeddings { tie_weights: true dropout: 0.2 scale: 5.5 }
output { bias: true norm: none }
}
"#,
)
.unwrap();
assert!(def.embeddings.tie_weights);
assert_eq!(def.embeddings.dropout, 0.2);
assert_eq!(def.embeddings.scale, Some(5.5));
assert!(def.output.bias);
assert!(matches!(def.output.norm.unwrap().norm_type, NormType::None));
}
#[test]
fn test_undefined_refs_error() {
let cases = [
"model m { vocab_size: 100 max_seq_len: 64 hidden_size: 32 num_layers: 2 block: nope }",
"block b { attention: nope ffn: { hidden_dim: 64 } }\n\
model m { vocab_size: 100 max_seq_len: 64 hidden_size: 32 num_layers: 2 block: b }",
"block b { attention: { num_heads: 4 } ffn: nope }\n\
model m { vocab_size: 100 max_seq_len: 64 hidden_size: 32 num_layers: 2 block: b }",
"block b { ssm: nope ffn: { hidden_dim: 64 } }\n\
model m { vocab_size: 100 max_seq_len: 64 hidden_size: 32 num_layers: 2 block: b }",
];
for mal in cases {
let err = parse_mal(mal).unwrap_err().to_string();
assert!(
err.contains("undefined"),
"expected undefined-ref error, got: {err}"
);
}
}
#[test]
fn test_parse_mal_rejects_multiple_models() {
let err = parse_mal(
r#"
model alpha { vocab_size: 10 hidden_size: 8 num_layers: 1 block: { attention: { num_heads: 1 } ffn: { hidden_dim: 16 } } }
model beta { vocab_size: 10 hidden_size: 8 num_layers: 1 block: { attention: { num_heads: 1 } ffn: { hidden_dim: 16 } } }
"#,
)
.unwrap_err()
.to_string();
assert!(err.contains("multiple model definitions"), "{err}");
assert!(err.contains("alpha") && err.contains("beta"), "{err}");
}
#[test]
fn test_parse_mal_full_rejects_duplicate_definitions() {
let err = parse_mal_full("attention repeated {} attention repeated {}")
.unwrap_err()
.to_string();
assert!(err.contains("duplicate attention 'repeated'"), "{err}");
}
#[test]
fn test_wellknown_models() {
for name in list_wellknown_models() {
let def = get_builtin_model(&name).unwrap_or_else(|| panic!("Failed to get {}", name));
assert!(def.num_heads() > 0);
assert!(def.intermediate_size() > 0);
}
}
#[test]
fn test_model_properties() {
let def = get_builtin_model("tiny").unwrap();
assert_eq!(def.vocab_size, 32000);
assert_eq!(def.hidden_size, 128);
assert_eq!(def.num_layers, 4);
assert_eq!(def.num_heads(), 4);
}
#[test]
fn test_comments() {
let mal = r#"
# This is a comment
attention test_attn {
# Comment in attention
num_heads: 2
}
ffn test_ffn {
hidden_dim: 256
}
block test_block {
attention: test_attn
ffn: test_ffn
}
# Comment before model
model test {
vocab_size: 1000
hidden_size: 64
num_layers: 2
block: test_block
}
"#;
let def = parse_mal(mal).unwrap();
assert_eq!(def.vocab_size, 1000);
}
#[test]
fn test_parse_position_encoding_and_tie_weights() {
let mal = r#"
attention pe_attn {
num_heads: 8
position_encoding: rope { theta: 100000.0 }
}
ffn pe_ffn {
hidden_dim: 1024
}
block pe_block {
attention: pe_attn
ffn: pe_ffn
}
model pe_test {
vocab_size: 1000
hidden_size: 256
num_layers: 2
block: pe_block
embeddings {
tie_weights: true
}
}
"#;
let def = parse_mal(mal).unwrap();
assert_eq!(def.rope_theta(), 100000.0, "theta must not be dropped");
let with_qk = parse_mal(
r#"
attention qk { num_heads: 4
qk_norm: true }
ffn f { hidden_dim: 64 }
block b { attention: qk
ffn: f }
model m { vocab_size: 100
hidden_size: 64
num_layers: 1
block: b }
"#,
)
.unwrap();
assert!(with_qk.block.attention.qk_norm);
assert!(
def.embeddings.tie_weights,
"tie_weights must not be dropped"
);
let mal2 = r#"
attention a { rope_theta: 500000.0 }
"#;
assert!(
parse_mal_full(mal2).is_err() || {
let f = parse_mal_full(mal2).unwrap();
f.attentions.is_empty()
}
);
let mal3 = r#"
attention nopos { position_encoding: none }
ffn f { hidden_dim: 64 }
block b { attention: nopos
ffn: f }
model m { vocab_size: 100
hidden_size: 64
num_layers: 1
block: b }
"#;
let def3 = parse_mal(mal3).unwrap();
assert!(matches!(
def3.block.attention.position_encoding,
PositionEncoding::None
));
}
#[test]
fn test_parse_hybrid_ssm_pattern() {
let mal = r#"
attention h_attn {
num_heads: 4
bias: false
}
ssm h_ssm {
state_dim: 16
conv_kernel: 4
expand: 2
}
ffn h_ffn {
hidden_dim: 512
activation: swiglu
bias: false
}
block attn_block {
attention: h_attn
ffn: h_ffn
norm: rmsnorm { eps: 1e-5 }
norm_position: pre
}
block mamba_block {
ssm: h_ssm
ffn: h_ffn
norm: rmsnorm { eps: 1e-5 }
norm_position: pre
}
model hybrid {
vocab_size: 1000
max_seq_len: 128
hidden_size: 64
num_layers: 6
block: attn_block
pattern: [mamba_block, mamba_block, attn_block]
}
"#;
let def = parse_mal(mal).unwrap();
assert_eq!(def.num_layers, 6);
let pattern = def.pattern.as_ref().unwrap();
assert_eq!(pattern.len(), 3);
assert!(pattern[0].is_ssm());
assert!(pattern[1].is_ssm());
assert!(!pattern[2].is_ssm());
assert!(def.block_for_layer(0).is_ssm());
assert!(!def.block_for_layer(2).is_ssm());
assert!(def.block_for_layer(3).is_ssm());
assert!(!def.block_for_layer(5).is_ssm());
let ssm = pattern[0].ssm.as_ref().unwrap();
assert_eq!(ssm.state_dim, 16);
assert_eq!(ssm.conv_kernel, 4);
assert_eq!(ssm.expand, 2);
assert_eq!(def.dt_rank(ssm), 4);
let json = serde_json::to_string(&def).unwrap();
let back: ModelDef = serde_json::from_str(&json).unwrap();
assert!(back.pattern.as_ref().unwrap()[0].is_ssm());
let attention_only: ModelDef = serde_json::from_str(
&serde_json::to_string(&get_builtin_model("tiny").unwrap()).unwrap(),
)
.unwrap();
assert!(attention_only.pattern.is_none());
assert!(!attention_only.block.is_ssm());
}
#[test]
fn test_composable_architecture() {
let mal = r#"
attention my_attn {
num_heads: 16
num_kv_heads: 4
head_dim: 128
bias: false
}
ffn my_ffn {
hidden_dim: 11008
activation: swiglu
bias: false
}
block my_block {
attention: my_attn
ffn: my_ffn
norm: rmsnorm { eps: 1e-5 }
norm_position: pre
residual: true
}
model my_model {
description: "LLaMA 7B architecture"
vocab_size: 32000
max_seq_len: 4096
hidden_size: 4096
num_layers: 32
block: my_block
}
"#;
let file = parse_mal_full(mal).unwrap();
assert!(file.attentions.contains_key("my_attn"));
assert!(file.ffns.contains_key("my_ffn"));
assert!(file.blocks.contains_key("my_block"));
assert!(file.models.contains_key("my_model"));
let attn = file.attentions.get("my_attn").unwrap();
assert_eq!(attn.num_heads, Some(16));
assert_eq!(attn.num_kv_heads, Some(4));
let ffn = file.ffns.get("my_ffn").unwrap();
assert_eq!(ffn.hidden_dim, Some(11008));
assert!(matches!(ffn.activation, Activation::SwiGLU));
let block = file.blocks.get("my_block").unwrap();
assert!(matches!(block.norm_position, NormPosition::Pre));
assert!(block.residual);
}
#[test]
fn vocabulary_storage_alignment_is_derived() {
let mut model = ModelDef {
vocab_size: 50_277,
..ModelDef::default()
};
assert_eq!(model.padded_vocab_size(), 50_304);
model.vocab_size = 32_000;
assert_eq!(model.padded_vocab_size(), 32_000);
let serialized = serde_json::to_value(&model).unwrap();
assert_eq!(serialized["vocab_size"], 32_000);
assert!(serialized.get("padded_vocab_size").is_none());
}
}