1use anyhow::{Result, anyhow};
82use pest::Parser;
83use pest_derive::Parser;
84use rust_embed::Embed;
85use serde::{Deserialize, Serialize};
86use std::collections::HashMap;
87
88#[derive(Embed)]
90#[folder = "well-known/"]
91#[include = "*.mal"]
92struct WellKnown;
93
94#[derive(Parser)]
95#[grammar = "mal.pest"]
96pub struct MalParser;
97
98#[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#[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 #[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#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
156pub enum NormType {
157 #[default]
158 RmsNorm,
159 LayerNorm,
160 None,
161}
162
163#[derive(Debug, Clone, Serialize, Deserialize, Default)]
165pub struct NormConfig {
166 pub norm_type: NormType,
167 pub eps: f64,
168}
169
170#[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#[derive(Debug, Clone, Serialize, Deserialize)]
184pub struct SsmDef {
185 pub name: String,
186 pub state_dim: usize,
188 pub conv_kernel: usize,
190 pub expand: usize,
192 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#[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 #[serde(default, skip_serializing_if = "Option::is_none")]
219 pub moe: Option<MoeDef>,
220}
221
222#[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#[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#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
255pub enum MemoryTierInit {
256 #[default]
257 Default,
258 ResidualZero,
261}
262
263#[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#[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#[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 #[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 pub fn is_ssm(&self) -> bool {
318 self.ssm.is_some()
319 }
320
321 pub fn num_heads(&self) -> usize {
325 self.attention.num_heads.unwrap_or(12)
326 }
327
328 pub fn num_kv_heads(&self) -> usize {
330 self.attention.num_kv_heads.unwrap_or(self.num_heads())
331 }
332
333 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 pub fn intermediate_size(&self, hidden_size: usize) -> usize {
342 self.ffn.hidden_dim.unwrap_or(hidden_size * 4)
343 }
344
345 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 pub fn rope_theta(&self) -> f64 {
356 match &self.attention.position_encoding {
357 PositionEncoding::Rope { theta, .. } => *theta,
358 _ => 10000.0,
359 }
360 }
361
362 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#[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#[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#[derive(Debug, Clone, Serialize, Deserialize, Default)]
408pub struct OutputConfig {
409 pub bias: bool,
410 pub norm: Option<NormConfig>,
411}
412
413#[derive(Debug, Clone, Serialize, Deserialize)]
415pub struct ModelDef {
416 pub name: String,
417 pub description: Option<String>,
418 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 #[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 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 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 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 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 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 pub fn dt_rank(&self, ssm: &SsmDef) -> usize {
522 ssm.dt_rank.unwrap_or(self.hidden_size.div_ceil(16))
523 }
524
525 pub fn num_heads(&self) -> usize {
530 self.block.num_heads()
531 }
532
533 pub fn num_kv_heads(&self) -> usize {
538 self.block.num_kv_heads()
539 }
540
541 pub fn head_dim(&self) -> usize {
546 self.block.head_dim(self.hidden_size)
547 }
548
549 pub fn intermediate_size(&self) -> usize {
554 self.block.intermediate_size(self.hidden_size)
555 }
556
557 pub fn norm_eps(&self) -> f64 {
562 self.block.norm_eps()
563 }
564
565 pub fn rope_theta(&self) -> f64 {
570 self.block.rope_theta()
571 }
572
573 pub fn padded_vocab_size(&self) -> usize {
578 self.vocab_size.next_multiple_of(64)
579 }
580
581 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 Some(ssm) => {
598 let d_inner = ssm.expand * h;
599 let dt_rank = self.dt_rank(ssm);
600 2 * h * d_inner + d_inner * h + d_inner * ssm.conv_kernel
603 + d_inner * (dt_rank + 2 * ssm.state_dim) + dt_rank * d_inner + d_inner + d_inner * ssm.state_dim + d_inner + d_inner }
609 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 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 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#[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
693fn 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
710fn 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
832fn 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 if let Some(name) = inner.next() {
839 def.name = name.as_str().to_string();
840 }
841
842 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
852pub 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
888pub 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
955fn 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
973fn parse_position_encoding(pair: pest::iterators::Pair<Rule>) -> Result<PositionEncoding> {
975 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
1024fn 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
1079fn 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
1097fn 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
1127fn parse_norm_config(pair: pest::iterators::Pair<Rule>) -> Result<NormConfig> {
1132 let mut norm = NormConfig::default();
1133 match pair.into_inner().next() {
1134 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 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
1172fn 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
1345fn 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
1363fn 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 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 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
1503pub 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
1509pub fn get_builtin_model(name: &str) -> Option<ModelDef> {
1520 let mal = get_wellknown_mal(name)?;
1521 parse_mal(&mal).ok()
1522}
1523
1524pub fn get_wellknown_mal(name: &str) -> Option<String> {
1528 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 format!("{}.mal", name.replace('-', "_"))
1535 };
1536
1537 WellKnown::get(&filename).map(|f| String::from_utf8_lossy(&f.data).into_owned())
1538}
1539
1540pub 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 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 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 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 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 let mal2 = r#"
2031 attention a { rope_theta: 500000.0 }
2032 "#;
2033 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 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); 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 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}