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