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