Skip to main content

omena_transform_egg/
mdl_cost.rs

1//! MDL-oriented extraction cost surface for optional e-graph rewrites.
2//!
3//! The default cost preserves plain AST-size extraction, while feature-gated
4//! modes can layer additional model penalties without affecting the core path.
5
6use 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}