omena_transform_egg/
mdl_cost.rs1use egg::{CostFunction, Id, Language};
7use serde::Serialize;
8
9use crate::CssRewriteLanguage;
10
11#[derive(Debug, Default, Clone, Copy, PartialEq, Eq, Serialize)]
12#[serde(rename_all = "camelCase")]
13pub enum MdlExtractionModeV0 {
14 #[default]
15 AstSize,
16 #[cfg(feature = "mdl")]
17 TwoPartUniform,
18}
19
20#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
21#[serde(rename_all = "camelCase")]
22pub struct MdlExtractionModeSummaryV0 {
23 pub schema_version: &'static str,
24 pub product: &'static str,
25 pub layer_marker: &'static str,
26 pub default_mode: MdlExtractionModeV0,
27 pub alternative_modes: Vec<MdlExtractionModeV0>,
28 pub unit: &'static str,
29 pub feature_gate: &'static str,
30 pub default_preserves_ast_size: bool,
31}
32
33#[derive(Debug, Clone, Copy)]
34pub struct MdlExtractionCostV0 {
35 mode: MdlExtractionModeV0,
36}
37
38impl MdlExtractionCostV0 {
39 pub const fn new(mode: MdlExtractionModeV0) -> Self {
40 Self { mode }
41 }
42
43 pub const fn default_ast_size() -> Self {
44 Self::new(MdlExtractionModeV0::AstSize)
45 }
46}
47
48impl Default for MdlExtractionCostV0 {
49 fn default() -> Self {
50 Self::default_ast_size()
51 }
52}
53
54impl CostFunction<CssRewriteLanguage> for MdlExtractionCostV0 {
55 type Cost = usize;
56
57 fn cost<C>(&mut self, enode: &CssRewriteLanguage, mut costs: C) -> Self::Cost
58 where
59 C: FnMut(Id) -> Self::Cost,
60 {
61 let ast_size = 1 + enode
62 .children()
63 .iter()
64 .copied()
65 .map(&mut costs)
66 .sum::<usize>();
67
68 match self.mode {
69 MdlExtractionModeV0::AstSize => ast_size,
70 #[cfg(feature = "mdl")]
71 MdlExtractionModeV0::TwoPartUniform => ast_size + two_part_uniform_model_penalty(enode),
72 }
73 }
74}
75
76#[cfg(feature = "mdl")]
77fn two_part_uniform_model_penalty(enode: &CssRewriteLanguage) -> usize {
78 match enode {
79 CssRewriteLanguage::Num(_) | CssRewriteLanguage::Symbol(_) => 1,
80 CssRewriteLanguage::Add(_)
81 | CssRewriteLanguage::Sub(_)
82 | CssRewriteLanguage::Mul(_)
83 | CssRewriteLanguage::Div(_)
84 | CssRewriteLanguage::Unit(_)
85 | CssRewriteLanguage::List(_) => 2,
86 CssRewriteLanguage::Calc(_) | CssRewriteLanguage::Is(_) | CssRewriteLanguage::Where(_) => 3,
87 }
88}
89
90pub fn summarize_mdl_extraction_mode() -> MdlExtractionModeSummaryV0 {
91 #[cfg(feature = "mdl")]
92 let alternative_modes = vec![MdlExtractionModeV0::TwoPartUniform];
93 #[cfg(not(feature = "mdl"))]
94 let alternative_modes = Vec::new();
95
96 MdlExtractionModeSummaryV0 {
97 schema_version: "0",
98 product: "omena-transform-egg.mdl-extraction",
99 layer_marker: "mdl-bits",
100 default_mode: MdlExtractionModeV0::AstSize,
101 alternative_modes,
102 unit: "bit",
103 feature_gate: "mdl",
104 default_preserves_ast_size: true,
105 }
106}