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