Skip to main content

hermes_mal/
lib.rs

1//! Model Architecture Language (MAL) for Hermes LLM
2//!
3//! A composable DSL for defining LLM model architectures using pest parser.
4//!
5//! # Example MAL - Minimal inline style
6//!
7//! ```text
8//! model tiny {
9//!     vocab_size: 32000
10//!     max_seq_len: 2048
11//!     hidden_size: 128
12//!     num_layers: 4
13//!     block: {
14//!         attention: { num_heads: 4 }
15//!         ffn: { hidden_dim: 512 }
16//!     }
17//! }
18//! ```
19//!
20//! # Example MAL - Composable style
21//!
22//! ```text
23//! # Define attention mechanism
24//! attention gqa {
25//!     num_heads: 32
26//!     num_kv_heads: 8
27//!     head_dim: 128
28//!     position_encoding: rope { theta: 10000.0 }
29//! }
30//!
31//! # Define FFN
32//! ffn swiglu_mlp {
33//!     hidden_dim: 14336
34//!     activation: swiglu
35//!     bias: false
36//! }
37//!
38//! # Define transformer block
39//! block llama_block {
40//!     attention: gqa
41//!     ffn: swiglu_mlp
42//!     norm: rmsnorm { eps: 1e-5 }
43//!     norm_position: pre
44//! }
45//!
46//! # Define model using the block
47//! model llama_7b {
48//!     vocab_size: 32000
49//!     max_seq_len: 4096
50//!     hidden_size: 4096
51//!     block: llama_block
52//!     num_layers: 32
53//! }
54//! ```
55
56use anyhow::{Result, anyhow};
57use pest::Parser;
58use pest_derive::Parser;
59use rust_embed::Embed;
60use serde::{Deserialize, Serialize};
61use std::collections::HashMap;
62
63/// Embedded well-known model definitions
64#[derive(Embed)]
65#[folder = "well-known/"]
66#[include = "*.mal"]
67struct WellKnown;
68
69#[derive(Parser)]
70#[grammar = "mal.pest"]
71pub struct MalParser;
72
73// ============================================================================
74// AST Types
75// ============================================================================
76
77/// Position encoding type
78#[derive(Debug, Clone, Serialize, Deserialize)]
79pub enum PositionEncoding {
80    Rope { theta: f64, scaling: Option<f64> },
81    Alibi { learned_slopes: bool },
82    Learned { max_positions: usize },
83    None,
84}
85
86impl Default for PositionEncoding {
87    fn default() -> Self {
88        Self::Rope {
89            theta: 10000.0,
90            scaling: None,
91        }
92    }
93}
94
95/// Attention mechanism definition
96#[derive(Debug, Clone, Serialize, Deserialize)]
97pub struct AttentionDef {
98    pub name: String,
99    pub num_heads: Option<usize>,
100    pub num_kv_heads: Option<usize>,
101    pub head_dim: Option<usize>,
102    pub dropout: f64,
103    pub bias: bool,
104    pub position_encoding: PositionEncoding,
105    pub window_size: Option<usize>,
106    pub causal: bool,
107    /// Per-head RMSNorm on Q and K before RoPE (OLMo2/Gemma-style stabilizer)
108    #[serde(default)]
109    pub qk_norm: bool,
110}
111
112impl Default for AttentionDef {
113    fn default() -> Self {
114        Self {
115            name: "default".to_string(),
116            num_heads: None,
117            num_kv_heads: None,
118            head_dim: None,
119            dropout: 0.0,
120            bias: false,
121            position_encoding: PositionEncoding::default(),
122            window_size: None,
123            causal: true,
124            qk_norm: false,
125        }
126    }
127}
128
129/// Normalization type
130#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
131pub enum NormType {
132    #[default]
133    RmsNorm,
134    LayerNorm,
135    None,
136}
137
138/// Normalization configuration
139#[derive(Debug, Clone, Serialize, Deserialize, Default)]
140pub struct NormConfig {
141    pub norm_type: NormType,
142    pub eps: f64,
143}
144
145/// Activation function
146#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
147pub enum Activation {
148    #[default]
149    SwiGLU,
150    GELU,
151    SiLU,
152    ReLU,
153    GELUNew,
154    GELUTanh,
155}
156
157/// Selective state-space (Mamba) mixer definition
158#[derive(Debug, Clone, Serialize, Deserialize)]
159pub struct SsmDef {
160    pub name: String,
161    /// SSM state dimension N (default 16)
162    pub state_dim: usize,
163    /// Depthwise causal conv kernel width (default 4)
164    pub conv_kernel: usize,
165    /// Inner expansion factor: d_inner = expand * hidden_size (default 2)
166    pub expand: usize,
167    /// Δ projection rank (default ceil(hidden_size / 16))
168    pub dt_rank: Option<usize>,
169}
170
171impl Default for SsmDef {
172    fn default() -> Self {
173        Self {
174            name: "default".to_string(),
175            state_dim: 16,
176            conv_kernel: 4,
177            expand: 2,
178            dt_rank: None,
179        }
180    }
181}
182
183/// Feed-forward network definition
184#[derive(Debug, Clone, Serialize, Deserialize)]
185pub struct FfnDef {
186    pub name: String,
187    pub hidden_dim: Option<usize>,
188    pub activation: Activation,
189    pub bias: bool,
190    pub dropout: f64,
191    pub gate: bool,
192    /// Sparse token-choice routing. Omitted for an ordinary dense FFN.
193    #[serde(default, skip_serializing_if = "Option::is_none")]
194    pub moe: Option<MoeDef>,
195}
196
197/// Configurable dropless token-choice mixture of experts.
198///
199/// `experts` are routed experts; `shared_experts` are always active. Router
200/// regularization belongs to the architecture config so every training entry
201/// point applies the same stable objective.
202#[derive(Debug, Clone, Serialize, Deserialize)]
203pub struct MoeDef {
204    pub experts: usize,
205    pub top_k: usize,
206    #[serde(default)]
207    pub shared_experts: usize,
208    #[serde(default)]
209    pub load_balance_loss_weight: f64,
210    #[serde(default)]
211    pub router_z_loss_weight: f64,
212}
213
214impl Default for FfnDef {
215    fn default() -> Self {
216        Self {
217            name: "default".to_string(),
218            hidden_dim: None,
219            activation: Activation::default(),
220            bias: false,
221            dropout: 0.0,
222            gate: true,
223            moe: None,
224        }
225    }
226}
227
228/// Transformer block definition
229///
230/// The mixer is attention by default; when `ssm` is set the block is a
231/// Mamba block (attention settings are ignored).
232#[derive(Debug, Clone, Serialize, Deserialize)]
233pub struct BlockDef {
234    pub name: String,
235    pub attention: AttentionDef,
236    #[serde(default)]
237    pub ssm: Option<SsmDef>,
238    pub ffn: FfnDef,
239    pub norm: NormConfig,
240    pub norm_position: NormPosition,
241    pub residual: bool,
242    pub dropout: f64,
243}
244
245impl BlockDef {
246    pub fn is_ssm(&self) -> bool {
247        self.ssm.is_some()
248    }
249
250    // Per-block computed properties (pattern-aware model construction)
251
252    pub fn num_heads(&self) -> usize {
253        self.attention.num_heads.unwrap_or(12)
254    }
255
256    pub fn num_kv_heads(&self) -> usize {
257        self.attention.num_kv_heads.unwrap_or(self.num_heads())
258    }
259
260    pub fn head_dim(&self, hidden_size: usize) -> usize {
261        self.attention
262            .head_dim
263            .unwrap_or(hidden_size / self.num_heads())
264    }
265
266    pub fn intermediate_size(&self, hidden_size: usize) -> usize {
267        self.ffn.hidden_dim.unwrap_or(hidden_size * 4)
268    }
269
270    pub fn norm_eps(&self) -> f64 {
271        if self.norm.eps > 0.0 {
272            self.norm.eps
273        } else {
274            1e-5
275        }
276    }
277
278    pub fn rope_theta(&self) -> f64 {
279        match &self.attention.position_encoding {
280            PositionEncoding::Rope { theta, .. } => *theta,
281            _ => 10000.0,
282        }
283    }
284
285    pub fn rope_scaling(&self) -> Option<f64> {
286        match &self.attention.position_encoding {
287            PositionEncoding::Rope { scaling, .. } => *scaling,
288            _ => None,
289        }
290    }
291}
292
293/// Normalization position in block
294#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
295pub enum NormPosition {
296    #[default]
297    Pre,
298    Post,
299}
300
301impl Default for BlockDef {
302    fn default() -> Self {
303        Self {
304            name: "default".to_string(),
305            attention: AttentionDef::default(),
306            ssm: None,
307            ffn: FfnDef::default(),
308            norm: NormConfig {
309                norm_type: NormType::RmsNorm,
310                eps: 1e-5,
311            },
312            norm_position: NormPosition::Pre,
313            residual: true,
314            dropout: 0.0,
315        }
316    }
317}
318
319/// Embeddings configuration
320#[derive(Debug, Clone, Serialize, Deserialize, Default)]
321pub struct EmbeddingsConfig {
322    pub tie_weights: bool,
323    pub dropout: f64,
324    pub scale: Option<f64>,
325}
326
327/// Output head configuration
328#[derive(Debug, Clone, Serialize, Deserialize, Default)]
329pub struct OutputConfig {
330    pub bias: bool,
331    pub norm: Option<NormConfig>,
332}
333
334/// Parsed model definition from MAL
335#[derive(Debug, Clone, Serialize, Deserialize)]
336pub struct ModelDef {
337    pub name: String,
338    pub description: Option<String>,
339    /// Number of token IDs exposed by the tokenizer and model API. Parameter
340    /// storage is padded internally for efficient accelerator kernels.
341    pub vocab_size: usize,
342    pub max_seq_len: usize,
343    pub hidden_size: usize,
344    pub num_layers: usize,
345    pub block: BlockDef,
346    /// Optional heterogeneous layer pattern, repeated cyclically across
347    /// num_layers (e.g. [mamba, mamba, attn]). Overrides `block` when set.
348    #[serde(default)]
349    pub pattern: Option<Vec<BlockDef>>,
350    pub embeddings: EmbeddingsConfig,
351    pub output: OutputConfig,
352}
353
354impl Default for ModelDef {
355    fn default() -> Self {
356        Self {
357            name: "default".to_string(),
358            description: None,
359            vocab_size: 32000,
360            max_seq_len: 2048,
361            hidden_size: 768,
362            num_layers: 12,
363            block: BlockDef::default(),
364            pattern: None,
365            embeddings: EmbeddingsConfig::default(),
366            output: OutputConfig::default(),
367        }
368    }
369}
370
371impl std::fmt::Display for ModelDef {
372    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
373        writeln!(f, "model {} {{", self.name)?;
374        if let Some(desc) = &self.description {
375            writeln!(f, "    description: \"{}\"", desc)?;
376        }
377        writeln!(f, "    vocab_size: {}", self.vocab_size)?;
378        writeln!(f, "    max_seq_len: {}", self.max_seq_len)?;
379        writeln!(f, "    hidden_size: {}", self.hidden_size)?;
380        writeln!(f, "    num_layers: {}", self.num_layers)?;
381        writeln!(f, "}}")?;
382        writeln!(f)?;
383
384        // Attention
385        writeln!(f, "attention {{")?;
386        if let Some(h) = self.block.attention.num_heads {
387            writeln!(f, "    num_heads: {}", h)?;
388        }
389        if let Some(kv) = self.block.attention.num_kv_heads {
390            writeln!(f, "    num_kv_heads: {}", kv)?;
391        }
392        if let Some(hd) = self.block.attention.head_dim {
393            writeln!(f, "    head_dim: {}", hd)?;
394        }
395        writeln!(f, "    bias: {}", self.block.attention.bias)?;
396        writeln!(f, "}}")?;
397        writeln!(f)?;
398
399        // FFN
400        writeln!(f, "ffn {{")?;
401        if let Some(dim) = self.block.ffn.hidden_dim {
402            writeln!(f, "    hidden_dim: {}", dim)?;
403        }
404        writeln!(f, "    activation: {:?}", self.block.ffn.activation)?;
405        writeln!(f, "    bias: {}", self.block.ffn.bias)?;
406        writeln!(f, "}}")?;
407        writeln!(f)?;
408
409        // Block
410        writeln!(f, "block {{")?;
411        writeln!(f, "    norm: {:?}", self.block.norm.norm_type)?;
412        writeln!(f, "    norm_position: {:?}", self.block.norm_position)?;
413        writeln!(f, "    residual: {}", self.block.residual)?;
414        writeln!(f, "}}")?;
415        writeln!(f)?;
416
417        // Estimated parameters
418        let params = self.estimated_params();
419        writeln!(
420            f,
421            "Estimated parameters: {:.2}B",
422            params as f64 / 1_000_000_000.0
423        )
424    }
425}
426
427impl ModelDef {
428    // ========================================================================
429    // Computed properties for model construction
430    // ========================================================================
431
432    /// Block definition for layer `i`: cycles through `pattern` when set,
433    /// otherwise the homogeneous `block`.
434    pub fn block_for_layer(&self, i: usize) -> &BlockDef {
435        match &self.pattern {
436            Some(p) if !p.is_empty() => &p[i % p.len()],
437            _ => &self.block,
438        }
439    }
440
441    /// Effective Δ rank for an SSM mixer (paper default: ceil(hidden/16))
442    pub fn dt_rank(&self, ssm: &SsmDef) -> usize {
443        ssm.dt_rank.unwrap_or(self.hidden_size.div_ceil(16))
444    }
445
446    pub fn num_heads(&self) -> usize {
447        self.block.attention.num_heads.unwrap_or(12)
448    }
449
450    pub fn num_kv_heads(&self) -> usize {
451        self.block
452            .attention
453            .num_kv_heads
454            .unwrap_or(self.num_heads())
455    }
456
457    pub fn head_dim(&self) -> usize {
458        self.block
459            .attention
460            .head_dim
461            .unwrap_or(self.hidden_size / self.num_heads())
462    }
463
464    pub fn intermediate_size(&self) -> usize {
465        self.block.ffn.hidden_dim.unwrap_or(self.hidden_size * 4)
466    }
467
468    pub fn norm_eps(&self) -> f64 {
469        if self.block.norm.eps > 0.0 {
470            self.block.norm.eps
471        } else {
472            1e-5
473        }
474    }
475
476    pub fn rope_theta(&self) -> f64 {
477        match &self.block.attention.position_encoding {
478            PositionEncoding::Rope { theta, .. } => *theta,
479            _ => 10000.0,
480        }
481    }
482
483    /// Vocabulary rows stored by embeddings and the output projection.
484    ///
485    /// Keeping this derived preserves a single logical vocabulary in MAL and
486    /// checkpoint configs while giving GPU kernels an aligned output dimension.
487    pub fn padded_vocab_size(&self) -> usize {
488        self.vocab_size.next_multiple_of(64)
489    }
490
491    /// Count trainable parameters implied by the model definition.
492    pub fn estimated_params(&self) -> usize {
493        let h = self.hidden_size;
494        let stored_vocab_size = self.padded_vocab_size();
495        let embed_params = stored_vocab_size * h;
496        let norm_params = |norm: &NormConfig| match norm.norm_type {
497            NormType::RmsNorm => h,
498            NormType::LayerNorm => 2 * h,
499            NormType::None => 0,
500        };
501
502        let mut layer_params = 0usize;
503        for i in 0..self.num_layers {
504            let block = self.block_for_layer(i);
505            let mixer = match &block.ssm {
506                // Mamba: in_proj + out_proj + conv + x_proj + dt_proj + A + D.
507                Some(ssm) => {
508                    let d_inner = ssm.expand * h;
509                    let dt_rank = self.dt_rank(ssm);
510                    2 * h * d_inner            // in_proj (x, z)
511                        + d_inner * h          // out_proj
512                        + d_inner * ssm.conv_kernel
513                        + d_inner * (dt_rank + 2 * ssm.state_dim) // x_proj
514                        + dt_rank * d_inner + d_inner // dt_proj weight + bias
515                        + d_inner * ssm.state_dim // A_log
516                        + d_inner // D
517                        + d_inner // depthwise conv bias
518                }
519                // Attention: q/k/v/o. Uses GQA kv width when configured.
520                None => {
521                    let q = block.num_heads() * block.head_dim(h);
522                    let kv = block.num_kv_heads() * block.head_dim(h);
523                    let weights = h * q + 2 * h * kv + q * h;
524                    let bias = if block.attention.bias {
525                        2 * q + 2 * kv
526                    } else {
527                        0
528                    };
529                    let qk_norm = if block.attention.qk_norm {
530                        2 * block.head_dim(h)
531                    } else {
532                        0
533                    };
534                    weights + bias + qk_norm
535                }
536            };
537            let intermediate = block.intermediate_size(h);
538            let projections = if block.ffn.gate { 3 } else { 2 };
539            let expert_count = block
540                .ffn
541                .moe
542                .as_ref()
543                .map_or(1, |moe| moe.experts + moe.shared_experts);
544            let ff_weights = expert_count * projections * h * intermediate;
545            let ff_bias = if block.ffn.bias {
546                expert_count * ((if block.ffn.gate { 2 } else { 1 }) * intermediate + h)
547            } else {
548                0
549            };
550            let router = block.ffn.moe.as_ref().map_or(0, |moe| h * moe.experts);
551            layer_params += mixer + ff_weights + ff_bias + router + 2 * norm_params(&block.norm);
552        }
553
554        let final_norm = self
555            .output
556            .norm
557            .as_ref()
558            .unwrap_or(&self.block_for_layer(0).norm);
559        let head_weights = (!self.embeddings.tie_weights) as usize * h * stored_vocab_size;
560        let head_bias = self.output.bias as usize * stored_vocab_size;
561        embed_params + layer_params + norm_params(final_norm) + head_weights + head_bias
562    }
563
564    /// Load from JSON file
565    pub fn from_json<P: AsRef<std::path::Path>>(path: P) -> Result<Self> {
566        let content = std::fs::read_to_string(path)?;
567        Ok(serde_json::from_str(&content)?)
568    }
569
570    /// Save to JSON file
571    pub fn save_json(&self, path: &str) -> Result<()> {
572        let content = serde_json::to_string_pretty(self)?;
573        std::fs::write(path, content)?;
574        Ok(())
575    }
576}
577
578/// Complete parsed MAL file with all definitions
579#[derive(Debug, Clone, Default)]
580pub struct MalFile {
581    pub attentions: HashMap<String, AttentionDef>,
582    pub ssms: HashMap<String, SsmDef>,
583    pub ffns: HashMap<String, FfnDef>,
584    pub blocks: HashMap<String, BlockDef>,
585    pub models: HashMap<String, ModelDef>,
586}
587
588// ============================================================================
589// Parsing Functions
590// ============================================================================
591
592/// Parse activation type from string
593fn parse_activation(s: &str) -> Activation {
594    match s {
595        "swiglu" => Activation::SwiGLU,
596        "gelu" => Activation::GELU,
597        "silu" => Activation::SiLU,
598        "relu" => Activation::ReLU,
599        "gelu_new" => Activation::GELUNew,
600        "gelu_tanh" => Activation::GELUTanh,
601        _ => Activation::SwiGLU,
602    }
603}
604
605/// Parse a model property (block-based only)
606fn parse_model_prop(
607    pair: pest::iterators::Pair<Rule>,
608    def: &mut ModelDef,
609    file: &MalFile,
610) -> Result<()> {
611    for inner in pair.into_inner() {
612        match inner.as_rule() {
613            Rule::vocab_size_prop => {
614                if let Some(val) = inner.into_inner().next() {
615                    def.vocab_size = val.as_str().parse()?;
616                }
617            }
618            Rule::max_seq_len_prop => {
619                if let Some(val) = inner.into_inner().next() {
620                    def.max_seq_len = val.as_str().parse()?;
621                }
622            }
623            Rule::hidden_size_prop => {
624                if let Some(val) = inner.into_inner().next() {
625                    def.hidden_size = val.as_str().parse()?;
626                }
627            }
628            Rule::num_layers_prop => {
629                if let Some(val) = inner.into_inner().next() {
630                    def.num_layers = val.as_str().parse()?;
631                }
632            }
633            Rule::block_ref_prop => {
634                for child in inner.into_inner() {
635                    match child.as_rule() {
636                        Rule::identifier => {
637                            let name = child.as_str();
638                            def.block = file
639                                .blocks
640                                .get(name)
641                                .ok_or_else(|| anyhow!("undefined block '{name}'"))?
642                                .clone();
643                        }
644                        Rule::inline_block => {
645                            let mut block = BlockDef::default();
646                            for prop in child.into_inner() {
647                                if prop.as_rule() == Rule::block_prop {
648                                    parse_block_prop(prop, &mut block, file)?;
649                                }
650                            }
651                            def.block = block;
652                        }
653                        _ => {}
654                    }
655                }
656            }
657            Rule::pattern_prop => {
658                let mut blocks = Vec::new();
659                for child in inner.into_inner() {
660                    if child.as_rule() == Rule::identifier {
661                        let name = child.as_str();
662                        let block = file.blocks.get(name).ok_or_else(|| {
663                            anyhow!("pattern references undefined block '{}'", name)
664                        })?;
665                        blocks.push(block.clone());
666                    }
667                }
668                if !blocks.is_empty() {
669                    def.pattern = Some(blocks);
670                }
671            }
672            Rule::embeddings_prop => {
673                for param in inner.into_inner() {
674                    for child in param.into_inner() {
675                        match child.as_rule() {
676                            Rule::tie_weights_prop => {
677                                if let Some(val) = child.into_inner().next() {
678                                    def.embeddings.tie_weights = val.as_str() == "true";
679                                }
680                            }
681                            Rule::dropout_prop => {
682                                if let Some(val) = child.into_inner().next() {
683                                    def.embeddings.dropout = val.as_str().parse()?;
684                                }
685                            }
686                            Rule::scale_prop => {
687                                if let Some(val) = child.into_inner().next() {
688                                    def.embeddings.scale = Some(val.as_str().parse()?);
689                                }
690                            }
691                            _ => {}
692                        }
693                    }
694                }
695            }
696            Rule::output_prop => {
697                for param in inner.into_inner() {
698                    for child in param.into_inner() {
699                        match child.as_rule() {
700                            Rule::bias_prop => {
701                                if let Some(val) = child.into_inner().next() {
702                                    def.output.bias = val.as_str() == "true";
703                                }
704                            }
705                            Rule::norm_prop => {
706                                if let Some(cfg) = child.into_inner().next() {
707                                    def.output.norm = Some(parse_norm_config(cfg)?);
708                                }
709                            }
710                            _ => {}
711                        }
712                    }
713                }
714            }
715            Rule::description_prop => {
716                if let Some(val) = inner.into_inner().next() {
717                    let s = val.as_str();
718                    def.description = Some(s[1..s.len() - 1].to_string());
719                }
720            }
721            _ => {}
722        }
723    }
724    Ok(())
725}
726
727/// Parse a model definition from pest pair
728fn parse_model_def(pair: pest::iterators::Pair<Rule>, file: &MalFile) -> Result<ModelDef> {
729    let mut def = ModelDef::default();
730    let mut inner = pair.into_inner();
731
732    // Get model name
733    if let Some(name) = inner.next() {
734        def.name = name.as_str().to_string();
735    }
736
737    // Parse properties
738    for prop in inner {
739        if prop.as_rule() == Rule::model_prop {
740            parse_model_prop(prop, &mut def, file)?;
741        }
742    }
743
744    Ok(def)
745}
746
747/// Parse a MAL string containing exactly one model definition.
748///
749/// Use [`parse_mal_full`] when a source intentionally defines multiple models.
750pub fn parse_mal(input: &str) -> Result<ModelDef> {
751    let file = parse_mal_full(input)?;
752    match file.models.len() {
753        0 => Err(anyhow!("no model definition found")),
754        1 => Ok(file.models.into_values().next().expect("length checked")),
755        _ => {
756            let mut names = file.models.keys().cloned().collect::<Vec<_>>();
757            names.sort();
758            Err(anyhow!(
759                "multiple model definitions found ({}); use parse_mal_full to select one",
760                names.join(", ")
761            ))
762        }
763    }
764}
765
766fn insert_unique<T>(
767    definitions: &mut HashMap<String, T>,
768    kind: &str,
769    name: String,
770    value: T,
771) -> Result<()> {
772    match definitions.entry(name) {
773        std::collections::hash_map::Entry::Vacant(entry) => {
774            entry.insert(value);
775            Ok(())
776        }
777        std::collections::hash_map::Entry::Occupied(entry) => {
778            Err(anyhow!("duplicate {kind} '{}'", entry.key()))
779        }
780    }
781}
782
783/// Parse complete MAL file with all definitions
784pub fn parse_mal_full(input: &str) -> Result<MalFile> {
785    let pairs = MalParser::parse(Rule::file, input).map_err(|e| anyhow!("Parse error: {}", e))?;
786
787    let mut file = MalFile::default();
788
789    for pair in pairs {
790        if pair.as_rule() == Rule::file {
791            for inner in pair.into_inner() {
792                if inner.as_rule() == Rule::definition {
793                    for def in inner.into_inner() {
794                        match def.as_rule() {
795                            Rule::model_def => {
796                                let model = parse_model_def(def, &file)?;
797                                insert_unique(
798                                    &mut file.models,
799                                    "model",
800                                    model.name.clone(),
801                                    model,
802                                )?;
803                            }
804                            Rule::attention_def => {
805                                let attn = parse_attention_def(def)?;
806                                insert_unique(
807                                    &mut file.attentions,
808                                    "attention",
809                                    attn.name.clone(),
810                                    attn,
811                                )?;
812                            }
813                            Rule::ssm_def => {
814                                let ssm = parse_ssm_def(def)?;
815                                insert_unique(&mut file.ssms, "ssm", ssm.name.clone(), ssm)?;
816                            }
817                            Rule::ffn_def => {
818                                let ffn = parse_ffn_def(def)?;
819                                insert_unique(&mut file.ffns, "ffn", ffn.name.clone(), ffn)?;
820                            }
821                            Rule::block_def => {
822                                let block = parse_block_def(def, &file)?;
823                                insert_unique(
824                                    &mut file.blocks,
825                                    "block",
826                                    block.name.clone(),
827                                    block,
828                                )?;
829                            }
830                            _ => {}
831                        }
832                    }
833                }
834            }
835        }
836    }
837
838    Ok(file)
839}
840
841/// Parse an attention definition
842fn parse_attention_def(pair: pest::iterators::Pair<Rule>) -> Result<AttentionDef> {
843    let mut def = AttentionDef::default();
844    let mut inner = pair.into_inner();
845
846    if let Some(name) = inner.next() {
847        def.name = name.as_str().to_string();
848    }
849
850    for prop in inner {
851        if prop.as_rule() == Rule::attention_prop {
852            parse_attention_prop(prop, &mut def)?;
853        }
854    }
855
856    Ok(def)
857}
858
859/// Parse a position-encoding config (rope { theta, scaling } | alibi | learned | none)
860fn parse_position_encoding(pair: pest::iterators::Pair<Rule>) -> Result<PositionEncoding> {
861    // position_encoding_config contains one of the variant configs; the bare
862    // "none" literal produces no inner pair.
863    let Some(config) = pair.into_inner().next() else {
864        return Ok(PositionEncoding::None);
865    };
866    match config.as_rule() {
867        Rule::rope_config => {
868            let mut theta = 10000.0;
869            let mut scaling = None;
870            for param in config.into_inner() {
871                for inner in param.into_inner() {
872                    match inner.as_rule() {
873                        Rule::rope_theta_prop | Rule::rope_base_prop => {
874                            if let Some(val) = inner.into_inner().next() {
875                                theta = val.as_str().parse()?;
876                            }
877                        }
878                        Rule::rope_scaling_prop => {
879                            if let Some(val) = inner.into_inner().next() {
880                                scaling = Some(val.as_str().parse()?);
881                            }
882                        }
883                        _ => {}
884                    }
885                }
886            }
887            Ok(PositionEncoding::Rope { theta, scaling })
888        }
889        Rule::alibi_config => {
890            let learned_slopes = config.as_str().contains("learned");
891            Ok(PositionEncoding::Alibi { learned_slopes })
892        }
893        Rule::learned_config => {
894            let mut max_positions = 0;
895            for param in config.into_inner() {
896                for inner in param.into_inner() {
897                    if inner.as_rule() == Rule::max_positions_prop
898                        && let Some(val) = inner.into_inner().next()
899                    {
900                        max_positions = val.as_str().parse()?;
901                    }
902                }
903            }
904            Ok(PositionEncoding::Learned { max_positions })
905        }
906        _ => Ok(PositionEncoding::None),
907    }
908}
909
910/// Parse attention properties
911fn parse_attention_prop(pair: pest::iterators::Pair<Rule>, def: &mut AttentionDef) -> Result<()> {
912    for inner in pair.into_inner() {
913        match inner.as_rule() {
914            Rule::num_heads_prop => {
915                if let Some(val) = inner.into_inner().next() {
916                    def.num_heads = Some(val.as_str().parse()?);
917                }
918            }
919            Rule::num_kv_heads_prop => {
920                if let Some(val) = inner.into_inner().next() {
921                    def.num_kv_heads = Some(val.as_str().parse()?);
922                }
923            }
924            Rule::head_dim_prop => {
925                if let Some(val) = inner.into_inner().next() {
926                    def.head_dim = Some(val.as_str().parse()?);
927                }
928            }
929            Rule::dropout_prop => {
930                if let Some(val) = inner.into_inner().next() {
931                    def.dropout = val.as_str().parse()?;
932                }
933            }
934            Rule::bias_prop => {
935                if let Some(val) = inner.into_inner().next() {
936                    def.bias = val.as_str() == "true";
937                }
938            }
939            Rule::causal_prop => {
940                if let Some(val) = inner.into_inner().next() {
941                    def.causal = val.as_str() == "true";
942                }
943            }
944            Rule::window_size_prop => {
945                if let Some(val) = inner.into_inner().next() {
946                    def.window_size = Some(val.as_str().parse()?);
947                }
948            }
949            Rule::position_encoding_prop => {
950                if let Some(config) = inner.into_inner().next() {
951                    def.position_encoding = parse_position_encoding(config)?;
952                }
953            }
954            Rule::qk_norm_prop => {
955                if let Some(val) = inner.into_inner().next() {
956                    def.qk_norm = val.as_str() == "true";
957                }
958            }
959            _ => {}
960        }
961    }
962    Ok(())
963}
964
965/// Parse an SSM (Mamba) definition
966fn parse_ssm_def(pair: pest::iterators::Pair<Rule>) -> Result<SsmDef> {
967    let mut def = SsmDef::default();
968    let mut inner = pair.into_inner();
969
970    if let Some(name) = inner.next() {
971        def.name = name.as_str().to_string();
972    }
973
974    for prop in inner {
975        if prop.as_rule() == Rule::ssm_prop {
976            parse_ssm_prop(prop, &mut def)?;
977        }
978    }
979
980    Ok(def)
981}
982
983/// Parse SSM properties
984fn parse_ssm_prop(pair: pest::iterators::Pair<Rule>, def: &mut SsmDef) -> Result<()> {
985    for inner in pair.into_inner() {
986        match inner.as_rule() {
987            Rule::state_dim_prop => {
988                if let Some(val) = inner.into_inner().next() {
989                    def.state_dim = val.as_str().parse()?;
990                }
991            }
992            Rule::conv_kernel_prop => {
993                if let Some(val) = inner.into_inner().next() {
994                    def.conv_kernel = val.as_str().parse()?;
995                }
996            }
997            Rule::expand_prop => {
998                if let Some(val) = inner.into_inner().next() {
999                    def.expand = val.as_str().parse()?;
1000                }
1001            }
1002            Rule::dt_rank_prop => {
1003                if let Some(val) = inner.into_inner().next() {
1004                    def.dt_rank = Some(val.as_str().parse()?);
1005                }
1006            }
1007            _ => {}
1008        }
1009    }
1010    Ok(())
1011}
1012
1013/// Parse an FFN definition
1014/// Parse a `norm_config` (`rmsnorm { eps: … }` | `layernorm { … }` | `none`)
1015/// into a `NormConfig`. An omitted epsilon is represented as zero and resolved
1016/// by model construction.
1017fn parse_norm_config(pair: pest::iterators::Pair<Rule>) -> Result<NormConfig> {
1018    let mut norm = NormConfig::default();
1019    match pair.into_inner().next() {
1020        // Bare `none` literal produces no inner rule.
1021        None => norm.norm_type = NormType::None,
1022        Some(cfg) => {
1023            norm.norm_type = match cfg.as_rule() {
1024                Rule::rmsnorm_config => NormType::RmsNorm,
1025                Rule::layernorm_config => NormType::LayerNorm,
1026                other => anyhow::bail!("unexpected norm config rule: {other:?}"),
1027            };
1028            for param in cfg.into_inner() {
1029                // norm_param -> norm_eps_prop -> number
1030                if let Some(prop) = param.into_inner().next()
1031                    && let Some(number) = prop.into_inner().next()
1032                {
1033                    norm.eps = number.as_str().parse()?;
1034                }
1035            }
1036        }
1037    }
1038    Ok(norm)
1039}
1040
1041fn parse_ffn_def(pair: pest::iterators::Pair<Rule>) -> Result<FfnDef> {
1042    let mut def = FfnDef::default();
1043    let mut inner = pair.into_inner();
1044
1045    if let Some(name) = inner.next() {
1046        def.name = name.as_str().to_string();
1047    }
1048
1049    for prop in inner {
1050        if prop.as_rule() == Rule::ffn_prop {
1051            parse_ffn_prop(prop, &mut def)?;
1052        }
1053    }
1054
1055    Ok(def)
1056}
1057
1058/// Parse FFN properties
1059fn parse_ffn_prop(pair: pest::iterators::Pair<Rule>, def: &mut FfnDef) -> Result<()> {
1060    for inner in pair.into_inner() {
1061        match inner.as_rule() {
1062            Rule::hidden_dim_prop => {
1063                if let Some(val) = inner.into_inner().next() {
1064                    def.hidden_dim = Some(val.as_str().parse()?);
1065                }
1066            }
1067            Rule::activation_prop => {
1068                if let Some(val) = inner.into_inner().next() {
1069                    def.activation = parse_activation(val.as_str());
1070                }
1071            }
1072            Rule::bias_prop => {
1073                if let Some(val) = inner.into_inner().next() {
1074                    def.bias = val.as_str() == "true";
1075                }
1076            }
1077            Rule::dropout_prop => {
1078                if let Some(val) = inner.into_inner().next() {
1079                    def.dropout = val.as_str().parse()?;
1080                }
1081            }
1082            Rule::gate_prop => {
1083                if let Some(val) = inner.into_inner().next() {
1084                    def.gate = val.as_str() == "true";
1085                }
1086            }
1087            Rule::moe_prop => {
1088                let mut moe = MoeDef {
1089                    experts: 0,
1090                    top_k: 0,
1091                    shared_experts: 0,
1092                    load_balance_loss_weight: 0.0,
1093                    router_z_loss_weight: 0.0,
1094                };
1095                for param in inner.into_inner() {
1096                    let Some(prop) = param.into_inner().next() else {
1097                        continue;
1098                    };
1099                    let value = prop
1100                        .clone()
1101                        .into_inner()
1102                        .next()
1103                        .map(|value| value.as_str())
1104                        .unwrap_or_default();
1105                    match prop.as_rule() {
1106                        Rule::experts_prop => moe.experts = value.parse()?,
1107                        Rule::top_k_prop => moe.top_k = value.parse()?,
1108                        Rule::shared_experts_prop => moe.shared_experts = value.parse()?,
1109                        Rule::load_balance_loss_weight_prop => {
1110                            moe.load_balance_loss_weight = value.parse()?
1111                        }
1112                        Rule::router_z_loss_weight_prop => {
1113                            moe.router_z_loss_weight = value.parse()?
1114                        }
1115                        _ => {}
1116                    }
1117                }
1118                def.moe = Some(moe);
1119            }
1120            _ => {}
1121        }
1122    }
1123    Ok(())
1124}
1125
1126/// Parse a block definition
1127fn parse_block_def(pair: pest::iterators::Pair<Rule>, file: &MalFile) -> Result<BlockDef> {
1128    let mut def = BlockDef::default();
1129    let mut inner = pair.into_inner();
1130
1131    if let Some(name) = inner.next() {
1132        def.name = name.as_str().to_string();
1133    }
1134
1135    for prop in inner {
1136        if prop.as_rule() == Rule::block_prop {
1137            parse_block_prop(prop, &mut def, file)?;
1138        }
1139    }
1140
1141    Ok(def)
1142}
1143
1144/// Parse block properties
1145fn parse_block_prop(
1146    pair: pest::iterators::Pair<Rule>,
1147    def: &mut BlockDef,
1148    file: &MalFile,
1149) -> Result<()> {
1150    for inner in pair.into_inner() {
1151        match inner.as_rule() {
1152            Rule::attention_ref_prop => {
1153                // Can be identifier or inline definition
1154                for child in inner.into_inner() {
1155                    match child.as_rule() {
1156                        Rule::identifier => {
1157                            let name = child.as_str();
1158                            def.attention = file
1159                                .attentions
1160                                .get(name)
1161                                .ok_or_else(|| anyhow!("undefined attention '{name}'"))?
1162                                .clone();
1163                        }
1164                        Rule::inline_attention => {
1165                            let mut attn = AttentionDef::default();
1166                            for prop in child.into_inner() {
1167                                if prop.as_rule() == Rule::attention_prop {
1168                                    parse_attention_prop(prop, &mut attn)?;
1169                                }
1170                            }
1171                            def.attention = attn;
1172                        }
1173                        _ => {}
1174                    }
1175                }
1176            }
1177            Rule::ssm_ref_prop => {
1178                for child in inner.into_inner() {
1179                    match child.as_rule() {
1180                        Rule::identifier => {
1181                            let name = child.as_str();
1182                            def.ssm = Some(
1183                                file.ssms
1184                                    .get(name)
1185                                    .ok_or_else(|| anyhow!("undefined ssm '{name}'"))?
1186                                    .clone(),
1187                            );
1188                        }
1189                        Rule::inline_ssm => {
1190                            let mut ssm = SsmDef::default();
1191                            for prop in child.into_inner() {
1192                                if prop.as_rule() == Rule::ssm_prop {
1193                                    parse_ssm_prop(prop, &mut ssm)?;
1194                                }
1195                            }
1196                            def.ssm = Some(ssm);
1197                        }
1198                        _ => {}
1199                    }
1200                }
1201            }
1202            Rule::ffn_ref_prop => {
1203                for child in inner.into_inner() {
1204                    match child.as_rule() {
1205                        Rule::identifier => {
1206                            let name = child.as_str();
1207                            def.ffn = file
1208                                .ffns
1209                                .get(name)
1210                                .ok_or_else(|| anyhow!("undefined ffn '{name}'"))?
1211                                .clone();
1212                        }
1213                        Rule::inline_ffn => {
1214                            let mut ffn = FfnDef::default();
1215                            for prop in child.into_inner() {
1216                                if prop.as_rule() == Rule::ffn_prop {
1217                                    parse_ffn_prop(prop, &mut ffn)?;
1218                                }
1219                            }
1220                            def.ffn = ffn;
1221                        }
1222                        _ => {}
1223                    }
1224                }
1225            }
1226            Rule::norm_prop => {
1227                // norm_prop -> norm_config -> (rmsnorm_config | layernorm_config | "none")
1228                if let Some(cfg) = inner.into_inner().next() {
1229                    def.norm = parse_norm_config(cfg)?;
1230                }
1231            }
1232            Rule::norm_position_prop => {
1233                if let Some(val) = inner.into_inner().next() {
1234                    def.norm_position = match val.as_str() {
1235                        "pre" => NormPosition::Pre,
1236                        "post" => NormPosition::Post,
1237                        _ => NormPosition::Pre,
1238                    };
1239                }
1240            }
1241            Rule::residual_prop => {
1242                if let Some(val) = inner.into_inner().next() {
1243                    def.residual = val.as_str() == "true";
1244                }
1245            }
1246            Rule::dropout_prop => {
1247                if let Some(val) = inner.into_inner().next() {
1248                    def.dropout = val.as_str().parse()?;
1249                }
1250            }
1251            _ => {}
1252        }
1253    }
1254    Ok(())
1255}
1256
1257/// Parse MAL from a file
1258pub fn parse_mal_file<P: AsRef<std::path::Path>>(path: P) -> Result<ModelDef> {
1259    let content = std::fs::read_to_string(path)?;
1260    parse_mal(&content)
1261}
1262
1263// ============================================================================
1264// Built-in model definitions
1265// ============================================================================
1266
1267/// Get a well-known model definition by name
1268///
1269/// Accepts:
1270/// - Short names: "nano", "tiny", "gpt2-small", etc.
1271/// - Well-known paths: "well-known/nano.mal", "well-known/gpt2_small.mal"
1272/// - Filenames: "nano.mal", "gpt2_small.mal"
1273pub fn get_builtin_model(name: &str) -> Option<ModelDef> {
1274    let mal = get_wellknown_mal(name)?;
1275    parse_mal(&mal).ok()
1276}
1277
1278/// Get the raw MAL content for a well-known model
1279///
1280/// Dynamically loads from embedded well-known/ directory.
1281pub fn get_wellknown_mal(name: &str) -> Option<String> {
1282    // Normalize: strip well-known/ prefix, ensure .mal suffix
1283    let name = name.strip_prefix("well-known/").unwrap_or(name);
1284    let filename = if name.ends_with(".mal") {
1285        name.to_string()
1286    } else {
1287        // Convert kebab-case to snake_case for filename
1288        format!("{}.mal", name.replace('-', "_"))
1289    };
1290
1291    WellKnown::get(&filename).map(|f| String::from_utf8_lossy(&f.data).into_owned())
1292}
1293
1294/// List all well-known model names (auto-discovered from embedded files)
1295pub fn list_wellknown_models() -> Vec<String> {
1296    WellKnown::iter()
1297        .filter_map(|path| {
1298            let path: &str = path.as_ref();
1299            if path.ends_with(".mal") {
1300                Some(path.strip_suffix(".mal").unwrap().replace('_', "-"))
1301            } else {
1302                None
1303            }
1304        })
1305        .collect()
1306}
1307
1308#[cfg(test)]
1309mod tests {
1310    use super::*;
1311
1312    #[test]
1313    fn test_parse_simple_model() {
1314        let mal = r#"
1315            attention test_attn {
1316                num_heads: 8
1317                bias: false
1318            }
1319
1320            ffn test_ffn {
1321                hidden_dim: 2048
1322                activation: gelu
1323            }
1324
1325            block test_block {
1326                attention: test_attn
1327                ffn: test_ffn
1328                norm_position: pre
1329            }
1330
1331            model test {
1332                vocab_size: 32000
1333                hidden_size: 512
1334                num_layers: 8
1335                block: test_block
1336            }
1337        "#;
1338
1339        let def = parse_mal(mal).unwrap();
1340        assert_eq!(def.name, "test");
1341        assert_eq!(def.vocab_size, 32000);
1342        assert_eq!(def.hidden_size, 512);
1343        assert_eq!(def.num_layers, 8);
1344    }
1345
1346    #[test]
1347    fn test_parse_with_block_props() {
1348        let mal = r#"
1349            attention full_attn {
1350                num_heads: 16
1351                num_kv_heads: 4
1352                bias: true
1353                dropout: 0.1
1354            }
1355
1356            ffn full_ffn {
1357                hidden_dim: 4096
1358                activation: gelu
1359                bias: true
1360                dropout: 0.1
1361            }
1362
1363            block full_block {
1364                attention: full_attn
1365                ffn: full_ffn
1366                norm: layernorm { eps: 1e-6 }
1367                norm_position: pre
1368                residual: true
1369            }
1370
1371            model full_test {
1372                description: "A test model"
1373                vocab_size: 50000
1374                max_seq_len: 4096
1375                hidden_size: 1024
1376                num_layers: 12
1377                block: full_block
1378            }
1379        "#;
1380
1381        let def = parse_mal(mal).unwrap();
1382        assert_eq!(def.description, Some("A test model".to_string()));
1383        assert_eq!(def.vocab_size, 50000);
1384        assert_eq!(def.max_seq_len, 4096);
1385        assert_eq!(def.block.attention.num_heads, Some(16));
1386        assert_eq!(def.block.attention.num_kv_heads, Some(4));
1387        assert_eq!(def.block.ffn.hidden_dim, Some(4096));
1388        assert!(matches!(def.block.ffn.activation, Activation::GELU));
1389        // Regression: `norm:` was silently dropped (no Rule::norm_prop arm),
1390        // building every block as the default RMSNorm regardless of config.
1391        assert!(matches!(def.block.norm.norm_type, NormType::LayerNorm));
1392        assert_eq!(def.block.norm.eps, 1e-6);
1393    }
1394
1395    #[test]
1396    fn moe_is_optional_and_fully_configurable() {
1397        let dense = parse_mal(
1398            "model d { vocab_size: 64 max_seq_len: 16 hidden_size: 8 num_layers: 1 \
1399             block: { attention: { num_heads: 1 } ffn: { hidden_dim: 12 } } }",
1400        )
1401        .unwrap();
1402        assert!(dense.block.ffn.moe.is_none());
1403
1404        let moe = parse_mal(
1405            r#"
1406            ffn experts {
1407                hidden_dim: 12
1408                activation: swiglu
1409                moe {
1410                    experts: 8
1411                    top_k: 2
1412                    shared_experts: 1
1413                    load_balance_loss_weight: 0.01
1414                    router_z_loss_weight: 0.001
1415                }
1416            }
1417            model m {
1418                vocab_size: 64 max_seq_len: 16 hidden_size: 8 num_layers: 1
1419                block: { attention: { num_heads: 1 } ffn: experts }
1420            }
1421            "#,
1422        )
1423        .unwrap();
1424        let config = moe.block.ffn.moe.as_ref().unwrap();
1425        assert_eq!(config.experts, 8);
1426        assert_eq!(config.top_k, 2);
1427        assert_eq!(config.shared_experts, 1);
1428        assert_eq!(config.load_balance_loss_weight, 0.01);
1429        assert_eq!(config.router_z_loss_weight, 0.001);
1430        assert!(moe.estimated_params() > dense.estimated_params());
1431    }
1432
1433    #[test]
1434    fn retriever_200m_moe_has_the_intended_sparse_budget() {
1435        let model = get_builtin_model("retriever-200m-moe").unwrap();
1436        assert_eq!(model.estimated_params(), 200_795_648);
1437        assert_eq!(model.num_layers, 24);
1438        assert_eq!(model.pattern.as_ref().unwrap().len(), 6);
1439
1440        let moe_layers: Vec<_> = (0..model.num_layers)
1441            .filter_map(|layer| model.block_for_layer(layer).ffn.moe.as_ref())
1442            .collect();
1443        assert_eq!(moe_layers.len(), 12);
1444        assert!(moe_layers.iter().all(|moe| moe.experts == 8));
1445        assert!(moe_layers.iter().all(|moe| moe.top_k == 2));
1446        assert!(moe_layers.iter().all(|moe| moe.shared_experts == 0));
1447        assert_eq!(model.block_for_layer(23).name, "attn_moe_block");
1448    }
1449
1450    #[test]
1451    fn test_norm_none_and_rmsnorm() {
1452        let base = |norm: &str| {
1453            format!(
1454                r#"
1455                block b {{ attention: {{ num_heads: 4 }} ffn: {{ hidden_dim: 64 }} norm: {norm} }}
1456                model m {{ vocab_size: 100 max_seq_len: 64 hidden_size: 32 num_layers: 2 block: b }}
1457                "#
1458            )
1459        };
1460        let d = parse_mal(&base("rmsnorm { eps: 1e-5 }")).unwrap();
1461        assert!(matches!(d.block.norm.norm_type, NormType::RmsNorm));
1462        let d = parse_mal(&base("none")).unwrap();
1463        assert!(matches!(d.block.norm.norm_type, NormType::None));
1464    }
1465
1466    #[test]
1467    fn test_embedding_and_output_configuration() {
1468        let def = parse_mal(
1469            r#"
1470            block b { attention: { num_heads: 4 } ffn: { hidden_dim: 64 } }
1471            model m {
1472                vocab_size: 100
1473                max_seq_len: 64
1474                hidden_size: 32
1475                num_layers: 2
1476                block: b
1477                embeddings { tie_weights: true dropout: 0.2 scale: 5.5 }
1478                output { bias: true norm: none }
1479            }
1480            "#,
1481        )
1482        .unwrap();
1483
1484        assert!(def.embeddings.tie_weights);
1485        assert_eq!(def.embeddings.dropout, 0.2);
1486        assert_eq!(def.embeddings.scale, Some(5.5));
1487        assert!(def.output.bias);
1488        assert!(matches!(def.output.norm.unwrap().norm_type, NormType::None));
1489    }
1490
1491    #[test]
1492    fn test_undefined_refs_error() {
1493        // Undefined block/attention/ffn/ssm references must fail loud, not
1494        // silently fall back to defaults.
1495        let cases = [
1496            "model m { vocab_size: 100 max_seq_len: 64 hidden_size: 32 num_layers: 2 block: nope }",
1497            "block b { attention: nope ffn: { hidden_dim: 64 } }\n\
1498             model m { vocab_size: 100 max_seq_len: 64 hidden_size: 32 num_layers: 2 block: b }",
1499            "block b { attention: { num_heads: 4 } ffn: nope }\n\
1500             model m { vocab_size: 100 max_seq_len: 64 hidden_size: 32 num_layers: 2 block: b }",
1501            "block b { ssm: nope ffn: { hidden_dim: 64 } }\n\
1502             model m { vocab_size: 100 max_seq_len: 64 hidden_size: 32 num_layers: 2 block: b }",
1503        ];
1504        for mal in cases {
1505            let err = parse_mal(mal).unwrap_err().to_string();
1506            assert!(
1507                err.contains("undefined"),
1508                "expected undefined-ref error, got: {err}"
1509            );
1510        }
1511    }
1512
1513    #[test]
1514    fn test_parse_mal_rejects_multiple_models() {
1515        let err = parse_mal(
1516            r#"
1517            model alpha { vocab_size: 10 hidden_size: 8 num_layers: 1 block: { attention: { num_heads: 1 } ffn: { hidden_dim: 16 } } }
1518            model beta { vocab_size: 10 hidden_size: 8 num_layers: 1 block: { attention: { num_heads: 1 } ffn: { hidden_dim: 16 } } }
1519            "#,
1520        )
1521        .unwrap_err()
1522        .to_string();
1523        assert!(err.contains("multiple model definitions"), "{err}");
1524        assert!(err.contains("alpha") && err.contains("beta"), "{err}");
1525    }
1526
1527    #[test]
1528    fn test_parse_mal_full_rejects_duplicate_definitions() {
1529        let err = parse_mal_full("attention repeated {} attention repeated {}")
1530            .unwrap_err()
1531            .to_string();
1532        assert!(err.contains("duplicate attention 'repeated'"), "{err}");
1533    }
1534
1535    #[test]
1536    fn test_wellknown_models() {
1537        for name in list_wellknown_models() {
1538            let def = get_builtin_model(&name).unwrap_or_else(|| panic!("Failed to get {}", name));
1539            // Verify computed properties work
1540            assert!(def.num_heads() > 0);
1541            assert!(def.intermediate_size() > 0);
1542        }
1543    }
1544
1545    #[test]
1546    fn test_model_properties() {
1547        let def = get_builtin_model("tiny").unwrap();
1548
1549        assert_eq!(def.vocab_size, 32000);
1550        assert_eq!(def.hidden_size, 128);
1551        assert_eq!(def.num_layers, 4);
1552        assert_eq!(def.num_heads(), 4);
1553    }
1554
1555    #[test]
1556    fn test_comments() {
1557        let mal = r#"
1558            # This is a comment
1559            attention test_attn {
1560                # Comment in attention
1561                num_heads: 2
1562            }
1563
1564            ffn test_ffn {
1565                hidden_dim: 256
1566            }
1567
1568            block test_block {
1569                attention: test_attn
1570                ffn: test_ffn
1571            }
1572
1573            # Comment before model
1574            model test {
1575                vocab_size: 1000
1576                hidden_size: 64
1577                num_layers: 2
1578                block: test_block
1579            }
1580        "#;
1581
1582        let def = parse_mal(mal).unwrap();
1583        assert_eq!(def.vocab_size, 1000);
1584    }
1585
1586    #[test]
1587    fn test_parse_position_encoding_and_tie_weights() {
1588        let mal = r#"
1589            attention pe_attn {
1590                num_heads: 8
1591                position_encoding: rope { theta: 100000.0 }
1592            }
1593
1594            ffn pe_ffn {
1595                hidden_dim: 1024
1596            }
1597
1598            block pe_block {
1599                attention: pe_attn
1600                ffn: pe_ffn
1601            }
1602
1603            model pe_test {
1604                vocab_size: 1000
1605                hidden_size: 256
1606                num_layers: 2
1607                block: pe_block
1608                embeddings {
1609                    tie_weights: true
1610                }
1611            }
1612        "#;
1613
1614        let def = parse_mal(mal).unwrap();
1615        assert_eq!(def.rope_theta(), 100000.0, "theta must not be dropped");
1616
1617        // qk_norm parses and lands
1618        let with_qk = parse_mal(
1619            r#"
1620            attention qk { num_heads: 4
1621                           qk_norm: true }
1622            ffn f { hidden_dim: 64 }
1623            block b { attention: qk
1624                      ffn: f }
1625            model m { vocab_size: 100
1626                      hidden_size: 64
1627                      num_layers: 1
1628                      block: b }
1629        "#,
1630        )
1631        .unwrap();
1632        assert!(with_qk.block.attention.qk_norm);
1633
1634        assert!(
1635            def.embeddings.tie_weights,
1636            "tie_weights must not be dropped"
1637        );
1638
1639        // Alternate spellings and variants
1640        let mal2 = r#"
1641            attention a { rope_theta: 500000.0 }
1642        "#;
1643        // rope_theta at attention level requires position_encoding wrapper;
1644        // bare form is not part of the grammar — this should fail to parse
1645        assert!(
1646            parse_mal_full(mal2).is_err() || {
1647                let f = parse_mal_full(mal2).unwrap();
1648                f.attentions.is_empty()
1649            }
1650        );
1651
1652        let mal3 = r#"
1653            attention nopos { position_encoding: none }
1654            ffn f { hidden_dim: 64 }
1655            block b { attention: nopos
1656                      ffn: f }
1657            model m { vocab_size: 100
1658                      hidden_size: 64
1659                      num_layers: 1
1660                      block: b }
1661        "#;
1662        let def3 = parse_mal(mal3).unwrap();
1663        assert!(matches!(
1664            def3.block.attention.position_encoding,
1665            PositionEncoding::None
1666        ));
1667    }
1668
1669    #[test]
1670    fn test_parse_hybrid_ssm_pattern() {
1671        let mal = r#"
1672            attention h_attn {
1673                num_heads: 4
1674                bias: false
1675            }
1676
1677            ssm h_ssm {
1678                state_dim: 16
1679                conv_kernel: 4
1680                expand: 2
1681            }
1682
1683            ffn h_ffn {
1684                hidden_dim: 512
1685                activation: swiglu
1686                bias: false
1687            }
1688
1689            block attn_block {
1690                attention: h_attn
1691                ffn: h_ffn
1692                norm: rmsnorm { eps: 1e-5 }
1693                norm_position: pre
1694            }
1695
1696            block mamba_block {
1697                ssm: h_ssm
1698                ffn: h_ffn
1699                norm: rmsnorm { eps: 1e-5 }
1700                norm_position: pre
1701            }
1702
1703            model hybrid {
1704                vocab_size: 1000
1705                max_seq_len: 128
1706                hidden_size: 64
1707                num_layers: 6
1708                block: attn_block
1709                pattern: [mamba_block, mamba_block, attn_block]
1710            }
1711        "#;
1712
1713        let def = parse_mal(mal).unwrap();
1714        assert_eq!(def.num_layers, 6);
1715
1716        let pattern = def.pattern.as_ref().unwrap();
1717        assert_eq!(pattern.len(), 3);
1718        assert!(pattern[0].is_ssm());
1719        assert!(pattern[1].is_ssm());
1720        assert!(!pattern[2].is_ssm());
1721
1722        // Cyclic layer assignment
1723        assert!(def.block_for_layer(0).is_ssm());
1724        assert!(!def.block_for_layer(2).is_ssm());
1725        assert!(def.block_for_layer(3).is_ssm());
1726        assert!(!def.block_for_layer(5).is_ssm());
1727
1728        let ssm = pattern[0].ssm.as_ref().unwrap();
1729        assert_eq!(ssm.state_dim, 16);
1730        assert_eq!(ssm.conv_kernel, 4);
1731        assert_eq!(ssm.expand, 2);
1732        assert_eq!(def.dt_rank(ssm), 4); // ceil(64/16)
1733
1734        // JSON roundtrip keeps the hybrid structure
1735        let json = serde_json::to_string(&def).unwrap();
1736        let back: ModelDef = serde_json::from_str(&json).unwrap();
1737        assert!(back.pattern.as_ref().unwrap()[0].is_ssm());
1738
1739        // Attention-only JSON leaves the optional hybrid fields empty.
1740        let attention_only: ModelDef = serde_json::from_str(
1741            &serde_json::to_string(&get_builtin_model("tiny").unwrap()).unwrap(),
1742        )
1743        .unwrap();
1744        assert!(attention_only.pattern.is_none());
1745        assert!(!attention_only.block.is_ssm());
1746    }
1747
1748    #[test]
1749    fn test_composable_architecture() {
1750        let mal = r#"
1751            attention my_attn {
1752                num_heads: 16
1753                num_kv_heads: 4
1754                head_dim: 128
1755                bias: false
1756            }
1757
1758            ffn my_ffn {
1759                hidden_dim: 11008
1760                activation: swiglu
1761                bias: false
1762            }
1763
1764            block my_block {
1765                attention: my_attn
1766                ffn: my_ffn
1767                norm: rmsnorm { eps: 1e-5 }
1768                norm_position: pre
1769                residual: true
1770            }
1771
1772            model my_model {
1773                description: "LLaMA 7B architecture"
1774                vocab_size: 32000
1775                max_seq_len: 4096
1776                hidden_size: 4096
1777                num_layers: 32
1778                block: my_block
1779            }
1780        "#;
1781
1782        let file = parse_mal_full(mal).unwrap();
1783
1784        assert!(file.attentions.contains_key("my_attn"));
1785        assert!(file.ffns.contains_key("my_ffn"));
1786        assert!(file.blocks.contains_key("my_block"));
1787        assert!(file.models.contains_key("my_model"));
1788
1789        let attn = file.attentions.get("my_attn").unwrap();
1790        assert_eq!(attn.num_heads, Some(16));
1791        assert_eq!(attn.num_kv_heads, Some(4));
1792
1793        let ffn = file.ffns.get("my_ffn").unwrap();
1794        assert_eq!(ffn.hidden_dim, Some(11008));
1795        assert!(matches!(ffn.activation, Activation::SwiGLU));
1796
1797        let block = file.blocks.get("my_block").unwrap();
1798        assert!(matches!(block.norm_position, NormPosition::Pre));
1799        assert!(block.residual);
1800    }
1801
1802    #[test]
1803    fn vocabulary_storage_alignment_is_derived() {
1804        let mut model = ModelDef {
1805            vocab_size: 50_277,
1806            ..ModelDef::default()
1807        };
1808        assert_eq!(model.padded_vocab_size(), 50_304);
1809
1810        model.vocab_size = 32_000;
1811        assert_eq!(model.padded_vocab_size(), 32_000);
1812
1813        let serialized = serde_json::to_value(&model).unwrap();
1814        assert_eq!(serialized["vocab_size"], 32_000);
1815        assert!(serialized.get("padded_vocab_size").is_none());
1816    }
1817}