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