1use anyhow::{Result, anyhow};
54use pest::Parser;
55use pest_derive::Parser;
56use rust_embed::Embed;
57use serde::{Deserialize, Serialize};
58use std::collections::HashMap;
59
60#[derive(Embed)]
62#[folder = "well-known/"]
63#[include = "*.mal"]
64struct WellKnown;
65
66#[derive(Parser)]
67#[grammar = "mal.pest"]
68pub struct MalParser;
69
70#[derive(Debug, Clone, Serialize, Deserialize)]
76pub enum PositionEncoding {
77 Rope { theta: f64, scaling: Option<f64> },
78 Alibi { learned_slopes: bool },
79 Learned { max_positions: usize },
80 None,
81}
82
83impl Default for PositionEncoding {
84 fn default() -> Self {
85 Self::Rope {
86 theta: 10000.0,
87 scaling: None,
88 }
89 }
90}
91
92#[derive(Debug, Clone, Serialize, Deserialize)]
94pub struct AttentionDef {
95 pub name: String,
96 pub num_heads: Option<usize>,
97 pub num_kv_heads: Option<usize>,
98 pub head_dim: Option<usize>,
99 pub dropout: f64,
100 pub bias: bool,
101 pub position_encoding: PositionEncoding,
102 pub window_size: Option<usize>,
103 pub causal: bool,
104 #[serde(default)]
106 pub qk_norm: bool,
107}
108
109impl Default for AttentionDef {
110 fn default() -> Self {
111 Self {
112 name: "default".to_string(),
113 num_heads: None,
114 num_kv_heads: None,
115 head_dim: None,
116 dropout: 0.0,
117 bias: false,
118 position_encoding: PositionEncoding::default(),
119 window_size: None,
120 causal: true,
121 qk_norm: false,
122 }
123 }
124}
125
126#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
128pub enum NormType {
129 #[default]
130 RmsNorm,
131 LayerNorm,
132 None,
133}
134
135#[derive(Debug, Clone, Serialize, Deserialize, Default)]
137pub struct NormConfig {
138 pub norm_type: NormType,
139 pub eps: f64,
140}
141
142#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
144pub enum Activation {
145 #[default]
146 SwiGLU,
147 GELU,
148 SiLU,
149 ReLU,
150 GELUNew,
151 GELUTanh,
152}
153
154#[derive(Debug, Clone, Serialize, Deserialize)]
156pub struct SsmDef {
157 pub name: String,
158 pub state_dim: usize,
160 pub conv_kernel: usize,
162 pub expand: usize,
164 pub dt_rank: Option<usize>,
166}
167
168impl Default for SsmDef {
169 fn default() -> Self {
170 Self {
171 name: "default".to_string(),
172 state_dim: 16,
173 conv_kernel: 4,
174 expand: 2,
175 dt_rank: None,
176 }
177 }
178}
179
180#[derive(Debug, Clone, Serialize, Deserialize)]
182pub struct FfnDef {
183 pub name: String,
184 pub hidden_dim: Option<usize>,
185 pub activation: Activation,
186 pub bias: bool,
187 pub dropout: f64,
188 pub gate: bool,
189}
190
191impl Default for FfnDef {
192 fn default() -> Self {
193 Self {
194 name: "default".to_string(),
195 hidden_dim: None,
196 activation: Activation::default(),
197 bias: false,
198 dropout: 0.0,
199 gate: true,
200 }
201 }
202}
203
204#[derive(Debug, Clone, Serialize, Deserialize)]
209pub struct BlockDef {
210 pub name: String,
211 pub attention: AttentionDef,
212 #[serde(default)]
213 pub ssm: Option<SsmDef>,
214 pub ffn: FfnDef,
215 pub norm: NormConfig,
216 pub norm_position: NormPosition,
217 pub residual: bool,
218 pub dropout: f64,
219}
220
221impl BlockDef {
222 pub fn is_ssm(&self) -> bool {
223 self.ssm.is_some()
224 }
225
226 pub fn num_heads(&self) -> usize {
229 self.attention.num_heads.unwrap_or(12)
230 }
231
232 pub fn num_kv_heads(&self) -> usize {
233 self.attention.num_kv_heads.unwrap_or(self.num_heads())
234 }
235
236 pub fn head_dim(&self, hidden_size: usize) -> usize {
237 self.attention
238 .head_dim
239 .unwrap_or(hidden_size / self.num_heads())
240 }
241
242 pub fn intermediate_size(&self, hidden_size: usize) -> usize {
243 self.ffn.hidden_dim.unwrap_or(hidden_size * 4)
244 }
245
246 pub fn norm_eps(&self) -> f64 {
247 if self.norm.eps > 0.0 {
248 self.norm.eps
249 } else {
250 1e-5
251 }
252 }
253
254 pub fn rope_theta(&self) -> f64 {
255 match &self.attention.position_encoding {
256 PositionEncoding::Rope { theta, .. } => *theta,
257 _ => 10000.0,
258 }
259 }
260
261 pub fn rope_scaling(&self) -> Option<f64> {
262 match &self.attention.position_encoding {
263 PositionEncoding::Rope { scaling, .. } => *scaling,
264 _ => None,
265 }
266 }
267}
268
269#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
271pub enum NormPosition {
272 #[default]
273 Pre,
274 Post,
275}
276
277impl Default for BlockDef {
278 fn default() -> Self {
279 Self {
280 name: "default".to_string(),
281 attention: AttentionDef::default(),
282 ssm: None,
283 ffn: FfnDef::default(),
284 norm: NormConfig {
285 norm_type: NormType::RmsNorm,
286 eps: 1e-5,
287 },
288 norm_position: NormPosition::Pre,
289 residual: true,
290 dropout: 0.0,
291 }
292 }
293}
294
295#[derive(Debug, Clone, Serialize, Deserialize, Default)]
297pub struct EmbeddingsConfig {
298 pub tie_weights: bool,
299 pub dropout: f64,
300 pub scale: Option<f64>,
301}
302
303#[derive(Debug, Clone, Serialize, Deserialize, Default)]
305pub struct OutputConfig {
306 pub bias: bool,
307 pub norm: Option<NormConfig>,
308}
309
310#[derive(Debug, Clone, Serialize, Deserialize)]
312pub struct ModelDef {
313 pub name: String,
314 pub description: Option<String>,
315 pub vocab_size: usize,
316 pub max_seq_len: usize,
317 pub hidden_size: usize,
318 pub num_layers: usize,
319 pub block: BlockDef,
320 #[serde(default)]
323 pub pattern: Option<Vec<BlockDef>>,
324 pub embeddings: EmbeddingsConfig,
325 pub output: OutputConfig,
326}
327
328impl Default for ModelDef {
329 fn default() -> Self {
330 Self {
331 name: "default".to_string(),
332 description: None,
333 vocab_size: 32000,
334 max_seq_len: 2048,
335 hidden_size: 768,
336 num_layers: 12,
337 block: BlockDef::default(),
338 pattern: None,
339 embeddings: EmbeddingsConfig::default(),
340 output: OutputConfig::default(),
341 }
342 }
343}
344
345impl std::fmt::Display for ModelDef {
346 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
347 writeln!(f, "model {} {{", self.name)?;
348 if let Some(desc) = &self.description {
349 writeln!(f, " description: \"{}\"", desc)?;
350 }
351 writeln!(f, " vocab_size: {}", self.vocab_size)?;
352 writeln!(f, " max_seq_len: {}", self.max_seq_len)?;
353 writeln!(f, " hidden_size: {}", self.hidden_size)?;
354 writeln!(f, " num_layers: {}", self.num_layers)?;
355 writeln!(f, "}}")?;
356 writeln!(f)?;
357
358 writeln!(f, "attention {{")?;
360 if let Some(h) = self.block.attention.num_heads {
361 writeln!(f, " num_heads: {}", h)?;
362 }
363 if let Some(kv) = self.block.attention.num_kv_heads {
364 writeln!(f, " num_kv_heads: {}", kv)?;
365 }
366 if let Some(hd) = self.block.attention.head_dim {
367 writeln!(f, " head_dim: {}", hd)?;
368 }
369 writeln!(f, " bias: {}", self.block.attention.bias)?;
370 writeln!(f, "}}")?;
371 writeln!(f)?;
372
373 writeln!(f, "ffn {{")?;
375 if let Some(dim) = self.block.ffn.hidden_dim {
376 writeln!(f, " hidden_dim: {}", dim)?;
377 }
378 writeln!(f, " activation: {:?}", self.block.ffn.activation)?;
379 writeln!(f, " bias: {}", self.block.ffn.bias)?;
380 writeln!(f, "}}")?;
381 writeln!(f)?;
382
383 writeln!(f, "block {{")?;
385 writeln!(f, " norm: {:?}", self.block.norm.norm_type)?;
386 writeln!(f, " norm_position: {:?}", self.block.norm_position)?;
387 writeln!(f, " residual: {}", self.block.residual)?;
388 writeln!(f, "}}")?;
389 writeln!(f)?;
390
391 let params = self.estimated_params();
393 writeln!(
394 f,
395 "Estimated parameters: {:.2}B",
396 params as f64 / 1_000_000_000.0
397 )
398 }
399}
400
401impl ModelDef {
402 pub fn block_for_layer(&self, i: usize) -> &BlockDef {
409 match &self.pattern {
410 Some(p) if !p.is_empty() => &p[i % p.len()],
411 _ => &self.block,
412 }
413 }
414
415 pub fn dt_rank(&self, ssm: &SsmDef) -> usize {
417 ssm.dt_rank.unwrap_or(self.hidden_size.div_ceil(16))
418 }
419
420 pub fn num_heads(&self) -> usize {
421 self.block.attention.num_heads.unwrap_or(12)
422 }
423
424 pub fn num_kv_heads(&self) -> usize {
425 self.block
426 .attention
427 .num_kv_heads
428 .unwrap_or(self.num_heads())
429 }
430
431 pub fn head_dim(&self) -> usize {
432 self.block
433 .attention
434 .head_dim
435 .unwrap_or(self.hidden_size / self.num_heads())
436 }
437
438 pub fn intermediate_size(&self) -> usize {
439 self.block.ffn.hidden_dim.unwrap_or(self.hidden_size * 4)
440 }
441
442 pub fn norm_eps(&self) -> f64 {
443 if self.block.norm.eps > 0.0 {
444 self.block.norm.eps
445 } else {
446 1e-5
447 }
448 }
449
450 pub fn rope_theta(&self) -> f64 {
451 match &self.block.attention.position_encoding {
452 PositionEncoding::Rope { theta, .. } => *theta,
453 _ => 10000.0,
454 }
455 }
456
457 pub fn estimated_params(&self) -> usize {
459 let h = self.hidden_size;
460 let embed_params = self.vocab_size * h;
461 let norm_params = |norm: &NormConfig| match norm.norm_type {
462 NormType::RmsNorm => h,
463 NormType::LayerNorm => 2 * h,
464 NormType::None => 0,
465 };
466
467 let mut layer_params = 0usize;
468 for i in 0..self.num_layers {
469 let block = self.block_for_layer(i);
470 let mixer = match &block.ssm {
471 Some(ssm) => {
473 let d_inner = ssm.expand * h;
474 let dt_rank = self.dt_rank(ssm);
475 2 * h * d_inner + d_inner * h + d_inner * ssm.conv_kernel
478 + d_inner * (dt_rank + 2 * ssm.state_dim) + dt_rank * d_inner + d_inner + d_inner * ssm.state_dim + d_inner + d_inner }
484 None => {
486 let q = block.num_heads() * block.head_dim(h);
487 let kv = block.num_kv_heads() * block.head_dim(h);
488 let weights = h * q + 2 * h * kv + q * h;
489 let bias = if block.attention.bias {
490 2 * q + 2 * kv
491 } else {
492 0
493 };
494 let qk_norm = if block.attention.qk_norm {
495 2 * block.head_dim(h)
496 } else {
497 0
498 };
499 weights + bias + qk_norm
500 }
501 };
502 let intermediate = block.intermediate_size(h);
503 let projections = if block.ffn.gate { 3 } else { 2 };
504 let ff_weights = projections * h * intermediate;
505 let ff_bias = if block.ffn.bias {
506 (if block.ffn.gate { 2 } else { 1 }) * intermediate + h
507 } else {
508 0
509 };
510 layer_params += mixer + ff_weights + ff_bias + 2 * norm_params(&block.norm);
511 }
512
513 let final_norm = self
514 .output
515 .norm
516 .as_ref()
517 .unwrap_or(&self.block_for_layer(0).norm);
518 let head_weights = (!self.embeddings.tie_weights) as usize * h * self.vocab_size;
519 let head_bias = self.output.bias as usize * self.vocab_size;
520 embed_params + layer_params + norm_params(final_norm) + head_weights + head_bias
521 }
522
523 pub fn from_json<P: AsRef<std::path::Path>>(path: P) -> Result<Self> {
525 let content = std::fs::read_to_string(path)?;
526 Ok(serde_json::from_str(&content)?)
527 }
528
529 pub fn save_json(&self, path: &str) -> Result<()> {
531 let content = serde_json::to_string_pretty(self)?;
532 std::fs::write(path, content)?;
533 Ok(())
534 }
535}
536
537#[derive(Debug, Clone, Default)]
539pub struct MalFile {
540 pub attentions: HashMap<String, AttentionDef>,
541 pub ssms: HashMap<String, SsmDef>,
542 pub ffns: HashMap<String, FfnDef>,
543 pub blocks: HashMap<String, BlockDef>,
544 pub models: HashMap<String, ModelDef>,
545}
546
547fn parse_activation(s: &str) -> Activation {
553 match s {
554 "swiglu" => Activation::SwiGLU,
555 "gelu" => Activation::GELU,
556 "silu" => Activation::SiLU,
557 "relu" => Activation::ReLU,
558 "gelu_new" => Activation::GELUNew,
559 "gelu_tanh" => Activation::GELUTanh,
560 _ => Activation::SwiGLU,
561 }
562}
563
564fn parse_model_prop(
566 pair: pest::iterators::Pair<Rule>,
567 def: &mut ModelDef,
568 file: &MalFile,
569) -> Result<()> {
570 for inner in pair.into_inner() {
571 match inner.as_rule() {
572 Rule::vocab_size_prop => {
573 if let Some(val) = inner.into_inner().next() {
574 def.vocab_size = val.as_str().parse()?;
575 }
576 }
577 Rule::max_seq_len_prop => {
578 if let Some(val) = inner.into_inner().next() {
579 def.max_seq_len = val.as_str().parse()?;
580 }
581 }
582 Rule::hidden_size_prop => {
583 if let Some(val) = inner.into_inner().next() {
584 def.hidden_size = val.as_str().parse()?;
585 }
586 }
587 Rule::num_layers_prop => {
588 if let Some(val) = inner.into_inner().next() {
589 def.num_layers = val.as_str().parse()?;
590 }
591 }
592 Rule::block_ref_prop => {
593 for child in inner.into_inner() {
594 match child.as_rule() {
595 Rule::identifier => {
596 let name = child.as_str();
597 def.block = file
598 .blocks
599 .get(name)
600 .ok_or_else(|| anyhow!("undefined block '{name}'"))?
601 .clone();
602 }
603 Rule::inline_block => {
604 let mut block = BlockDef::default();
605 for prop in child.into_inner() {
606 if prop.as_rule() == Rule::block_prop {
607 parse_block_prop(prop, &mut block, file)?;
608 }
609 }
610 def.block = block;
611 }
612 _ => {}
613 }
614 }
615 }
616 Rule::pattern_prop => {
617 let mut blocks = Vec::new();
618 for child in inner.into_inner() {
619 if child.as_rule() == Rule::identifier {
620 let name = child.as_str();
621 let block = file.blocks.get(name).ok_or_else(|| {
622 anyhow!("pattern references undefined block '{}'", name)
623 })?;
624 blocks.push(block.clone());
625 }
626 }
627 if !blocks.is_empty() {
628 def.pattern = Some(blocks);
629 }
630 }
631 Rule::embeddings_prop => {
632 for param in inner.into_inner() {
633 for child in param.into_inner() {
634 match child.as_rule() {
635 Rule::tie_weights_prop => {
636 if let Some(val) = child.into_inner().next() {
637 def.embeddings.tie_weights = val.as_str() == "true";
638 }
639 }
640 Rule::dropout_prop => {
641 if let Some(val) = child.into_inner().next() {
642 def.embeddings.dropout = val.as_str().parse()?;
643 }
644 }
645 Rule::scale_prop => {
646 if let Some(val) = child.into_inner().next() {
647 def.embeddings.scale = Some(val.as_str().parse()?);
648 }
649 }
650 _ => {}
651 }
652 }
653 }
654 }
655 Rule::output_prop => {
656 for param in inner.into_inner() {
657 for child in param.into_inner() {
658 match child.as_rule() {
659 Rule::bias_prop => {
660 if let Some(val) = child.into_inner().next() {
661 def.output.bias = val.as_str() == "true";
662 }
663 }
664 Rule::norm_prop => {
665 if let Some(cfg) = child.into_inner().next() {
666 def.output.norm = Some(parse_norm_config(cfg)?);
667 }
668 }
669 _ => {}
670 }
671 }
672 }
673 }
674 Rule::description_prop => {
675 if let Some(val) = inner.into_inner().next() {
676 let s = val.as_str();
677 def.description = Some(s[1..s.len() - 1].to_string());
678 }
679 }
680 _ => {}
681 }
682 }
683 Ok(())
684}
685
686fn parse_model_def(pair: pest::iterators::Pair<Rule>, file: &MalFile) -> Result<ModelDef> {
688 let mut def = ModelDef::default();
689 let mut inner = pair.into_inner();
690
691 if let Some(name) = inner.next() {
693 def.name = name.as_str().to_string();
694 }
695
696 for prop in inner {
698 if prop.as_rule() == Rule::model_prop {
699 parse_model_prop(prop, &mut def, file)?;
700 }
701 }
702
703 Ok(def)
704}
705
706pub fn parse_mal(input: &str) -> Result<ModelDef> {
708 let file = parse_mal_full(input)?;
709 file.models
710 .into_values()
711 .next()
712 .ok_or_else(|| anyhow!("No model definition found"))
713}
714
715pub fn parse_mal_full(input: &str) -> Result<MalFile> {
717 let pairs = MalParser::parse(Rule::file, input).map_err(|e| anyhow!("Parse error: {}", e))?;
718
719 let mut file = MalFile::default();
720
721 for pair in pairs {
722 if pair.as_rule() == Rule::file {
723 for inner in pair.into_inner() {
724 if inner.as_rule() == Rule::definition {
725 for def in inner.into_inner() {
726 match def.as_rule() {
727 Rule::model_def => {
728 let model = parse_model_def(def, &file)?;
729 file.models.insert(model.name.clone(), model);
730 }
731 Rule::attention_def => {
732 let attn = parse_attention_def(def)?;
733 file.attentions.insert(attn.name.clone(), attn);
734 }
735 Rule::ssm_def => {
736 let ssm = parse_ssm_def(def)?;
737 file.ssms.insert(ssm.name.clone(), ssm);
738 }
739 Rule::ffn_def => {
740 let ffn = parse_ffn_def(def)?;
741 file.ffns.insert(ffn.name.clone(), ffn);
742 }
743 Rule::block_def => {
744 let block = parse_block_def(def, &file)?;
745 file.blocks.insert(block.name.clone(), block);
746 }
747 _ => {}
748 }
749 }
750 }
751 }
752 }
753 }
754
755 Ok(file)
756}
757
758fn parse_attention_def(pair: pest::iterators::Pair<Rule>) -> Result<AttentionDef> {
760 let mut def = AttentionDef::default();
761 let mut inner = pair.into_inner();
762
763 if let Some(name) = inner.next() {
764 def.name = name.as_str().to_string();
765 }
766
767 for prop in inner {
768 if prop.as_rule() == Rule::attention_prop {
769 parse_attention_prop(prop, &mut def)?;
770 }
771 }
772
773 Ok(def)
774}
775
776fn parse_position_encoding(pair: pest::iterators::Pair<Rule>) -> Result<PositionEncoding> {
778 let Some(config) = pair.into_inner().next() else {
781 return Ok(PositionEncoding::None);
782 };
783 match config.as_rule() {
784 Rule::rope_config => {
785 let mut theta = 10000.0;
786 let mut scaling = None;
787 for param in config.into_inner() {
788 for inner in param.into_inner() {
789 match inner.as_rule() {
790 Rule::rope_theta_prop | Rule::rope_base_prop => {
791 if let Some(val) = inner.into_inner().next() {
792 theta = val.as_str().parse()?;
793 }
794 }
795 Rule::rope_scaling_prop => {
796 if let Some(val) = inner.into_inner().next() {
797 scaling = Some(val.as_str().parse()?);
798 }
799 }
800 _ => {}
801 }
802 }
803 }
804 Ok(PositionEncoding::Rope { theta, scaling })
805 }
806 Rule::alibi_config => {
807 let learned_slopes = config.as_str().contains("learned");
808 Ok(PositionEncoding::Alibi { learned_slopes })
809 }
810 Rule::learned_config => {
811 let mut max_positions = 0;
812 for param in config.into_inner() {
813 for inner in param.into_inner() {
814 if inner.as_rule() == Rule::max_positions_prop
815 && let Some(val) = inner.into_inner().next()
816 {
817 max_positions = val.as_str().parse()?;
818 }
819 }
820 }
821 Ok(PositionEncoding::Learned { max_positions })
822 }
823 _ => Ok(PositionEncoding::None),
824 }
825}
826
827fn parse_attention_prop(pair: pest::iterators::Pair<Rule>, def: &mut AttentionDef) -> Result<()> {
829 for inner in pair.into_inner() {
830 match inner.as_rule() {
831 Rule::num_heads_prop => {
832 if let Some(val) = inner.into_inner().next() {
833 def.num_heads = Some(val.as_str().parse()?);
834 }
835 }
836 Rule::num_kv_heads_prop => {
837 if let Some(val) = inner.into_inner().next() {
838 def.num_kv_heads = Some(val.as_str().parse()?);
839 }
840 }
841 Rule::head_dim_prop => {
842 if let Some(val) = inner.into_inner().next() {
843 def.head_dim = Some(val.as_str().parse()?);
844 }
845 }
846 Rule::dropout_prop => {
847 if let Some(val) = inner.into_inner().next() {
848 def.dropout = val.as_str().parse()?;
849 }
850 }
851 Rule::bias_prop => {
852 if let Some(val) = inner.into_inner().next() {
853 def.bias = val.as_str() == "true";
854 }
855 }
856 Rule::causal_prop => {
857 if let Some(val) = inner.into_inner().next() {
858 def.causal = val.as_str() == "true";
859 }
860 }
861 Rule::window_size_prop => {
862 if let Some(val) = inner.into_inner().next() {
863 def.window_size = Some(val.as_str().parse()?);
864 }
865 }
866 Rule::position_encoding_prop => {
867 if let Some(config) = inner.into_inner().next() {
868 def.position_encoding = parse_position_encoding(config)?;
869 }
870 }
871 Rule::qk_norm_prop => {
872 if let Some(val) = inner.into_inner().next() {
873 def.qk_norm = val.as_str() == "true";
874 }
875 }
876 _ => {}
877 }
878 }
879 Ok(())
880}
881
882fn parse_ssm_def(pair: pest::iterators::Pair<Rule>) -> Result<SsmDef> {
884 let mut def = SsmDef::default();
885 let mut inner = pair.into_inner();
886
887 if let Some(name) = inner.next() {
888 def.name = name.as_str().to_string();
889 }
890
891 for prop in inner {
892 if prop.as_rule() == Rule::ssm_prop {
893 parse_ssm_prop(prop, &mut def)?;
894 }
895 }
896
897 Ok(def)
898}
899
900fn parse_ssm_prop(pair: pest::iterators::Pair<Rule>, def: &mut SsmDef) -> Result<()> {
902 for inner in pair.into_inner() {
903 match inner.as_rule() {
904 Rule::state_dim_prop => {
905 if let Some(val) = inner.into_inner().next() {
906 def.state_dim = val.as_str().parse()?;
907 }
908 }
909 Rule::conv_kernel_prop => {
910 if let Some(val) = inner.into_inner().next() {
911 def.conv_kernel = val.as_str().parse()?;
912 }
913 }
914 Rule::expand_prop => {
915 if let Some(val) = inner.into_inner().next() {
916 def.expand = val.as_str().parse()?;
917 }
918 }
919 Rule::dt_rank_prop => {
920 if let Some(val) = inner.into_inner().next() {
921 def.dt_rank = Some(val.as_str().parse()?);
922 }
923 }
924 _ => {}
925 }
926 }
927 Ok(())
928}
929
930fn parse_norm_config(pair: pest::iterators::Pair<Rule>) -> Result<NormConfig> {
935 let mut norm = NormConfig::default();
936 match pair.into_inner().next() {
937 None => norm.norm_type = NormType::None,
939 Some(cfg) => {
940 norm.norm_type = match cfg.as_rule() {
941 Rule::rmsnorm_config => NormType::RmsNorm,
942 Rule::layernorm_config => NormType::LayerNorm,
943 other => anyhow::bail!("unexpected norm config rule: {other:?}"),
944 };
945 for param in cfg.into_inner() {
946 if let Some(eps) = param
948 .into_inner()
949 .next()
950 .and_then(|p| p.into_inner().next().and_then(|n| n.as_str().parse().ok()))
951 {
952 norm.eps = eps;
953 }
954 }
955 }
956 }
957 Ok(norm)
958}
959
960fn parse_ffn_def(pair: pest::iterators::Pair<Rule>) -> Result<FfnDef> {
961 let mut def = FfnDef::default();
962 let mut inner = pair.into_inner();
963
964 if let Some(name) = inner.next() {
965 def.name = name.as_str().to_string();
966 }
967
968 for prop in inner {
969 if prop.as_rule() == Rule::ffn_prop {
970 parse_ffn_prop(prop, &mut def)?;
971 }
972 }
973
974 Ok(def)
975}
976
977fn parse_ffn_prop(pair: pest::iterators::Pair<Rule>, def: &mut FfnDef) -> Result<()> {
979 for inner in pair.into_inner() {
980 match inner.as_rule() {
981 Rule::hidden_dim_prop => {
982 if let Some(val) = inner.into_inner().next() {
983 def.hidden_dim = Some(val.as_str().parse()?);
984 }
985 }
986 Rule::activation_prop => {
987 if let Some(val) = inner.into_inner().next() {
988 def.activation = parse_activation(val.as_str());
989 }
990 }
991 Rule::bias_prop => {
992 if let Some(val) = inner.into_inner().next() {
993 def.bias = val.as_str() == "true";
994 }
995 }
996 Rule::dropout_prop => {
997 if let Some(val) = inner.into_inner().next() {
998 def.dropout = val.as_str().parse()?;
999 }
1000 }
1001 Rule::gate_prop => {
1002 if let Some(val) = inner.into_inner().next() {
1003 def.gate = val.as_str() == "true";
1004 }
1005 }
1006 _ => {}
1007 }
1008 }
1009 Ok(())
1010}
1011
1012fn parse_block_def(pair: pest::iterators::Pair<Rule>, file: &MalFile) -> Result<BlockDef> {
1014 let mut def = BlockDef::default();
1015 let mut inner = pair.into_inner();
1016
1017 if let Some(name) = inner.next() {
1018 def.name = name.as_str().to_string();
1019 }
1020
1021 for prop in inner {
1022 if prop.as_rule() == Rule::block_prop {
1023 parse_block_prop(prop, &mut def, file)?;
1024 }
1025 }
1026
1027 Ok(def)
1028}
1029
1030fn parse_block_prop(
1032 pair: pest::iterators::Pair<Rule>,
1033 def: &mut BlockDef,
1034 file: &MalFile,
1035) -> Result<()> {
1036 for inner in pair.into_inner() {
1037 match inner.as_rule() {
1038 Rule::attention_ref_prop => {
1039 for child in inner.into_inner() {
1041 match child.as_rule() {
1042 Rule::identifier => {
1043 let name = child.as_str();
1044 def.attention = file
1045 .attentions
1046 .get(name)
1047 .ok_or_else(|| anyhow!("undefined attention '{name}'"))?
1048 .clone();
1049 }
1050 Rule::inline_attention => {
1051 let mut attn = AttentionDef::default();
1052 for prop in child.into_inner() {
1053 if prop.as_rule() == Rule::attention_prop {
1054 parse_attention_prop(prop, &mut attn)?;
1055 }
1056 }
1057 def.attention = attn;
1058 }
1059 _ => {}
1060 }
1061 }
1062 }
1063 Rule::ssm_ref_prop => {
1064 for child in inner.into_inner() {
1065 match child.as_rule() {
1066 Rule::identifier => {
1067 let name = child.as_str();
1068 def.ssm = Some(
1069 file.ssms
1070 .get(name)
1071 .ok_or_else(|| anyhow!("undefined ssm '{name}'"))?
1072 .clone(),
1073 );
1074 }
1075 Rule::inline_ssm => {
1076 let mut ssm = SsmDef::default();
1077 for prop in child.into_inner() {
1078 if prop.as_rule() == Rule::ssm_prop {
1079 parse_ssm_prop(prop, &mut ssm)?;
1080 }
1081 }
1082 def.ssm = Some(ssm);
1083 }
1084 _ => {}
1085 }
1086 }
1087 }
1088 Rule::ffn_ref_prop => {
1089 for child in inner.into_inner() {
1090 match child.as_rule() {
1091 Rule::identifier => {
1092 let name = child.as_str();
1093 def.ffn = file
1094 .ffns
1095 .get(name)
1096 .ok_or_else(|| anyhow!("undefined ffn '{name}'"))?
1097 .clone();
1098 }
1099 Rule::inline_ffn => {
1100 let mut ffn = FfnDef::default();
1101 for prop in child.into_inner() {
1102 if prop.as_rule() == Rule::ffn_prop {
1103 parse_ffn_prop(prop, &mut ffn)?;
1104 }
1105 }
1106 def.ffn = ffn;
1107 }
1108 _ => {}
1109 }
1110 }
1111 }
1112 Rule::norm_prop => {
1113 if let Some(cfg) = inner.into_inner().next() {
1115 def.norm = parse_norm_config(cfg)?;
1116 }
1117 }
1118 Rule::norm_position_prop => {
1119 if let Some(val) = inner.into_inner().next() {
1120 def.norm_position = match val.as_str() {
1121 "pre" => NormPosition::Pre,
1122 "post" => NormPosition::Post,
1123 _ => NormPosition::Pre,
1124 };
1125 }
1126 }
1127 Rule::residual_prop => {
1128 if let Some(val) = inner.into_inner().next() {
1129 def.residual = val.as_str() == "true";
1130 }
1131 }
1132 Rule::dropout_prop => {
1133 if let Some(val) = inner.into_inner().next() {
1134 def.dropout = val.as_str().parse()?;
1135 }
1136 }
1137 _ => {}
1138 }
1139 }
1140 Ok(())
1141}
1142
1143pub fn parse_mal_file<P: AsRef<std::path::Path>>(path: P) -> Result<ModelDef> {
1145 let content = std::fs::read_to_string(path)?;
1146 parse_mal(&content)
1147}
1148
1149pub fn get_builtin_model(name: &str) -> Option<ModelDef> {
1160 let mal = get_wellknown_mal(name)?;
1161 parse_mal(&mal).ok()
1162}
1163
1164pub fn get_wellknown_mal(name: &str) -> Option<String> {
1168 let name = name.strip_prefix("well-known/").unwrap_or(name);
1170 let filename = if name.ends_with(".mal") {
1171 name.to_string()
1172 } else {
1173 format!("{}.mal", name.replace('-', "_"))
1175 };
1176
1177 WellKnown::get(&filename).map(|f| String::from_utf8_lossy(&f.data).into_owned())
1178}
1179
1180pub fn list_wellknown_models() -> Vec<String> {
1182 WellKnown::iter()
1183 .filter_map(|path| {
1184 let path: &str = path.as_ref();
1185 if path.ends_with(".mal") {
1186 Some(path.strip_suffix(".mal").unwrap().replace('_', "-"))
1187 } else {
1188 None
1189 }
1190 })
1191 .collect()
1192}
1193
1194#[cfg(test)]
1195mod tests {
1196 use super::*;
1197
1198 #[test]
1199 fn test_parse_simple_model() {
1200 let mal = r#"
1201 attention test_attn {
1202 num_heads: 8
1203 bias: false
1204 }
1205
1206 ffn test_ffn {
1207 hidden_dim: 2048
1208 activation: gelu
1209 }
1210
1211 block test_block {
1212 attention: test_attn
1213 ffn: test_ffn
1214 norm_position: pre
1215 }
1216
1217 model test {
1218 vocab_size: 32000
1219 hidden_size: 512
1220 num_layers: 8
1221 block: test_block
1222 }
1223 "#;
1224
1225 let def = parse_mal(mal).unwrap();
1226 assert_eq!(def.name, "test");
1227 assert_eq!(def.vocab_size, 32000);
1228 assert_eq!(def.hidden_size, 512);
1229 assert_eq!(def.num_layers, 8);
1230 }
1231
1232 #[test]
1233 fn test_parse_with_block_props() {
1234 let mal = r#"
1235 attention full_attn {
1236 num_heads: 16
1237 num_kv_heads: 4
1238 bias: true
1239 dropout: 0.1
1240 }
1241
1242 ffn full_ffn {
1243 hidden_dim: 4096
1244 activation: gelu
1245 bias: true
1246 dropout: 0.1
1247 }
1248
1249 block full_block {
1250 attention: full_attn
1251 ffn: full_ffn
1252 norm: layernorm { eps: 1e-6 }
1253 norm_position: pre
1254 residual: true
1255 }
1256
1257 model full_test {
1258 description: "A test model"
1259 vocab_size: 50000
1260 max_seq_len: 4096
1261 hidden_size: 1024
1262 num_layers: 12
1263 block: full_block
1264 }
1265 "#;
1266
1267 let def = parse_mal(mal).unwrap();
1268 assert_eq!(def.description, Some("A test model".to_string()));
1269 assert_eq!(def.vocab_size, 50000);
1270 assert_eq!(def.max_seq_len, 4096);
1271 assert_eq!(def.block.attention.num_heads, Some(16));
1272 assert_eq!(def.block.attention.num_kv_heads, Some(4));
1273 assert_eq!(def.block.ffn.hidden_dim, Some(4096));
1274 assert!(matches!(def.block.ffn.activation, Activation::GELU));
1275 assert!(matches!(def.block.norm.norm_type, NormType::LayerNorm));
1278 assert_eq!(def.block.norm.eps, 1e-6);
1279 }
1280
1281 #[test]
1282 fn test_norm_none_and_rmsnorm() {
1283 let base = |norm: &str| {
1284 format!(
1285 r#"
1286 block b {{ attention: {{ num_heads: 4 }} ffn: {{ hidden_dim: 64 }} norm: {norm} }}
1287 model m {{ vocab_size: 100 max_seq_len: 64 hidden_size: 32 num_layers: 2 block: b }}
1288 "#
1289 )
1290 };
1291 let d = parse_mal(&base("rmsnorm { eps: 1e-5 }")).unwrap();
1292 assert!(matches!(d.block.norm.norm_type, NormType::RmsNorm));
1293 let d = parse_mal(&base("none")).unwrap();
1294 assert!(matches!(d.block.norm.norm_type, NormType::None));
1295 }
1296
1297 #[test]
1298 fn test_embedding_and_output_configuration() {
1299 let def = parse_mal(
1300 r#"
1301 block b { attention: { num_heads: 4 } ffn: { hidden_dim: 64 } }
1302 model m {
1303 vocab_size: 100
1304 max_seq_len: 64
1305 hidden_size: 32
1306 num_layers: 2
1307 block: b
1308 embeddings { tie_weights: true dropout: 0.2 scale: 5.5 }
1309 output { bias: true norm: none }
1310 }
1311 "#,
1312 )
1313 .unwrap();
1314
1315 assert!(def.embeddings.tie_weights);
1316 assert_eq!(def.embeddings.dropout, 0.2);
1317 assert_eq!(def.embeddings.scale, Some(5.5));
1318 assert!(def.output.bias);
1319 assert!(matches!(def.output.norm.unwrap().norm_type, NormType::None));
1320 }
1321
1322 #[test]
1323 fn test_undefined_refs_error() {
1324 let cases = [
1327 "model m { vocab_size: 100 max_seq_len: 64 hidden_size: 32 num_layers: 2 block: nope }",
1328 "block b { attention: nope ffn: { hidden_dim: 64 } }\n\
1329 model m { vocab_size: 100 max_seq_len: 64 hidden_size: 32 num_layers: 2 block: b }",
1330 "block b { attention: { num_heads: 4 } ffn: nope }\n\
1331 model m { vocab_size: 100 max_seq_len: 64 hidden_size: 32 num_layers: 2 block: b }",
1332 "block b { ssm: nope ffn: { hidden_dim: 64 } }\n\
1333 model m { vocab_size: 100 max_seq_len: 64 hidden_size: 32 num_layers: 2 block: b }",
1334 ];
1335 for mal in cases {
1336 let err = parse_mal(mal).unwrap_err().to_string();
1337 assert!(
1338 err.contains("undefined"),
1339 "expected undefined-ref error, got: {err}"
1340 );
1341 }
1342 }
1343
1344 #[test]
1345 fn test_wellknown_models() {
1346 for name in list_wellknown_models() {
1347 let def = get_builtin_model(&name).unwrap_or_else(|| panic!("Failed to get {}", name));
1348 assert!(def.num_heads() > 0);
1350 assert!(def.intermediate_size() > 0);
1351 }
1352 }
1353
1354 #[test]
1355 fn test_model_properties() {
1356 let def = get_builtin_model("tiny").unwrap();
1357
1358 assert_eq!(def.vocab_size, 32000);
1359 assert_eq!(def.hidden_size, 128);
1360 assert_eq!(def.num_layers, 4);
1361 assert_eq!(def.num_heads(), 4);
1362 }
1363
1364 #[test]
1365 fn test_comments() {
1366 let mal = r#"
1367 # This is a comment
1368 attention test_attn {
1369 # Comment in attention
1370 num_heads: 2
1371 }
1372
1373 ffn test_ffn {
1374 hidden_dim: 256
1375 }
1376
1377 block test_block {
1378 attention: test_attn
1379 ffn: test_ffn
1380 }
1381
1382 # Comment before model
1383 model test {
1384 vocab_size: 1000
1385 hidden_size: 64
1386 num_layers: 2
1387 block: test_block
1388 }
1389 "#;
1390
1391 let def = parse_mal(mal).unwrap();
1392 assert_eq!(def.vocab_size, 1000);
1393 }
1394
1395 #[test]
1396 fn test_parse_position_encoding_and_tie_weights() {
1397 let mal = r#"
1398 attention pe_attn {
1399 num_heads: 8
1400 position_encoding: rope { theta: 100000.0 }
1401 }
1402
1403 ffn pe_ffn {
1404 hidden_dim: 1024
1405 }
1406
1407 block pe_block {
1408 attention: pe_attn
1409 ffn: pe_ffn
1410 }
1411
1412 model pe_test {
1413 vocab_size: 1000
1414 hidden_size: 256
1415 num_layers: 2
1416 block: pe_block
1417 embeddings {
1418 tie_weights: true
1419 }
1420 }
1421 "#;
1422
1423 let def = parse_mal(mal).unwrap();
1424 assert_eq!(def.rope_theta(), 100000.0, "theta must not be dropped");
1425
1426 let with_qk = parse_mal(
1428 r#"
1429 attention qk { num_heads: 4
1430 qk_norm: true }
1431 ffn f { hidden_dim: 64 }
1432 block b { attention: qk
1433 ffn: f }
1434 model m { vocab_size: 100
1435 hidden_size: 64
1436 num_layers: 1
1437 block: b }
1438 "#,
1439 )
1440 .unwrap();
1441 assert!(with_qk.block.attention.qk_norm);
1442
1443 assert!(
1444 def.embeddings.tie_weights,
1445 "tie_weights must not be dropped"
1446 );
1447
1448 let mal2 = r#"
1450 attention a { rope_theta: 500000.0 }
1451 "#;
1452 assert!(
1455 parse_mal_full(mal2).is_err() || {
1456 let f = parse_mal_full(mal2).unwrap();
1457 f.attentions.is_empty()
1458 }
1459 );
1460
1461 let mal3 = r#"
1462 attention nopos { position_encoding: none }
1463 ffn f { hidden_dim: 64 }
1464 block b { attention: nopos
1465 ffn: f }
1466 model m { vocab_size: 100
1467 hidden_size: 64
1468 num_layers: 1
1469 block: b }
1470 "#;
1471 let def3 = parse_mal(mal3).unwrap();
1472 assert!(matches!(
1473 def3.block.attention.position_encoding,
1474 PositionEncoding::None
1475 ));
1476 }
1477
1478 #[test]
1479 fn test_parse_hybrid_ssm_pattern() {
1480 let mal = r#"
1481 attention h_attn {
1482 num_heads: 4
1483 bias: false
1484 }
1485
1486 ssm h_ssm {
1487 state_dim: 16
1488 conv_kernel: 4
1489 expand: 2
1490 }
1491
1492 ffn h_ffn {
1493 hidden_dim: 512
1494 activation: swiglu
1495 bias: false
1496 }
1497
1498 block attn_block {
1499 attention: h_attn
1500 ffn: h_ffn
1501 norm: rmsnorm { eps: 1e-5 }
1502 norm_position: pre
1503 }
1504
1505 block mamba_block {
1506 ssm: h_ssm
1507 ffn: h_ffn
1508 norm: rmsnorm { eps: 1e-5 }
1509 norm_position: pre
1510 }
1511
1512 model hybrid {
1513 vocab_size: 1000
1514 max_seq_len: 128
1515 hidden_size: 64
1516 num_layers: 6
1517 block: attn_block
1518 pattern: [mamba_block, mamba_block, attn_block]
1519 }
1520 "#;
1521
1522 let def = parse_mal(mal).unwrap();
1523 assert_eq!(def.num_layers, 6);
1524
1525 let pattern = def.pattern.as_ref().unwrap();
1526 assert_eq!(pattern.len(), 3);
1527 assert!(pattern[0].is_ssm());
1528 assert!(pattern[1].is_ssm());
1529 assert!(!pattern[2].is_ssm());
1530
1531 assert!(def.block_for_layer(0).is_ssm());
1533 assert!(!def.block_for_layer(2).is_ssm());
1534 assert!(def.block_for_layer(3).is_ssm());
1535 assert!(!def.block_for_layer(5).is_ssm());
1536
1537 let ssm = pattern[0].ssm.as_ref().unwrap();
1538 assert_eq!(ssm.state_dim, 16);
1539 assert_eq!(ssm.conv_kernel, 4);
1540 assert_eq!(ssm.expand, 2);
1541 assert_eq!(def.dt_rank(ssm), 4); let json = serde_json::to_string(&def).unwrap();
1545 let back: ModelDef = serde_json::from_str(&json).unwrap();
1546 assert!(back.pattern.as_ref().unwrap()[0].is_ssm());
1547
1548 let legacy: ModelDef = serde_json::from_str(
1550 &serde_json::to_string(&get_builtin_model("tiny").unwrap()).unwrap(),
1551 )
1552 .unwrap();
1553 assert!(legacy.pattern.is_none());
1554 assert!(!legacy.block.is_ssm());
1555 }
1556
1557 #[test]
1558 fn test_composable_architecture() {
1559 let mal = r#"
1560 attention my_attn {
1561 num_heads: 16
1562 num_kv_heads: 4
1563 head_dim: 128
1564 bias: false
1565 }
1566
1567 ffn my_ffn {
1568 hidden_dim: 11008
1569 activation: swiglu
1570 bias: false
1571 }
1572
1573 block my_block {
1574 attention: my_attn
1575 ffn: my_ffn
1576 norm: rmsnorm { eps: 1e-5 }
1577 norm_position: pre
1578 residual: true
1579 }
1580
1581 model my_model {
1582 description: "LLaMA 7B architecture"
1583 vocab_size: 32000
1584 max_seq_len: 4096
1585 hidden_size: 4096
1586 num_layers: 32
1587 block: my_block
1588 }
1589 "#;
1590
1591 let file = parse_mal_full(mal).unwrap();
1592
1593 assert!(file.attentions.contains_key("my_attn"));
1594 assert!(file.ffns.contains_key("my_ffn"));
1595 assert!(file.blocks.contains_key("my_block"));
1596 assert!(file.models.contains_key("my_model"));
1597
1598 let attn = file.attentions.get("my_attn").unwrap();
1599 assert_eq!(attn.num_heads, Some(16));
1600 assert_eq!(attn.num_kv_heads, Some(4));
1601
1602 let ffn = file.ffns.get("my_ffn").unwrap();
1603 assert_eq!(ffn.hidden_dim, Some(11008));
1604 assert!(matches!(ffn.activation, Activation::SwiGLU));
1605
1606 let block = file.blocks.get("my_block").unwrap();
1607 assert!(matches!(block.norm_position, NormPosition::Pre));
1608 assert!(block.residual);
1609 }
1610}