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