Skip to main content

sim_lib_numbers_stats/
function.rs

1//! Runtime library surface for statistics functions.
2
3use std::{any::Any, sync::Arc};
4
5use sim_kernel::{
6    AbiVersion, Args, Callable, ClassRef, Cx, DefaultFactory, Dependency, Export, Expr, Factory,
7    Lib, LibManifest, LibTarget, Linker, Object, RawArgs, Result, Symbol, Value, Version,
8};
9
10use super::{agent_fixtures, runtime};
11
12/// Returns the symbol bound to the `stats/disparate-impact-claim` operation.
13pub fn stats_disparate_impact_claim_symbol() -> Symbol {
14    Symbol::qualified("stats", "disparate-impact-claim")
15}
16
17/// Returns the symbol bound to the `stats/mean-claim` operation.
18pub fn stats_mean_claim_symbol() -> Symbol {
19    Symbol::qualified("stats", "mean-claim")
20}
21
22/// Returns the symbol bound to the `stats/variance-claim` operation.
23pub fn stats_variance_claim_symbol() -> Symbol {
24    Symbol::qualified("stats", "variance-claim")
25}
26
27/// Returns the symbol bound to the `stats/entropy-claim` operation.
28pub fn stats_entropy_claim_symbol() -> Symbol {
29    Symbol::qualified("stats", "entropy-claim")
30}
31
32/// Returns the symbol bound to the `stats/claims` batch operation.
33pub fn stats_claims_symbol() -> Symbol {
34    Symbol::qualified("stats", "claims")
35}
36
37/// Returns the symbol bound to deterministic bounded k-means.
38pub fn stats_kmeans_symbol() -> Symbol {
39    Symbol::qualified("stats", "kmeans")
40}
41
42/// Returns the symbol bound to regularized Gaussian-mixture EM.
43pub fn stats_gmm_symbol() -> Symbol {
44    Symbol::qualified("stats", "gmm")
45}
46
47/// Returns the symbol bound to exact finite binary intervals.
48pub fn stats_exact_binary_interval_symbol() -> Symbol {
49    Symbol::qualified("stats", "exact-binary-interval")
50}
51/// Returns the symbol bound to paired bootstrap intervals.
52pub fn stats_paired_bootstrap_symbol() -> Symbol {
53    Symbol::qualified("stats", "paired-bootstrap")
54}
55/// Returns the symbol bound to cluster-preserving bootstrap intervals.
56pub fn stats_clustered_bootstrap_symbol() -> Symbol {
57    Symbol::qualified("stats", "clustered-bootstrap")
58}
59/// Returns the symbol bound to registered-look intervals.
60pub fn stats_registered_look_symbol() -> Symbol {
61    Symbol::qualified("stats", "registered-look-interval")
62}
63/// Returns the symbol bound to weighted isotonic fitting.
64pub fn stats_isotonic_symbol() -> Symbol {
65    Symbol::qualified("stats", "isotonic")
66}
67
68fn function_symbols() -> [Symbol; 16] {
69    [
70        stats_mean_claim_symbol(),
71        stats_variance_claim_symbol(),
72        stats_entropy_claim_symbol(),
73        stats_disparate_impact_claim_symbol(),
74        stats_claims_symbol(),
75        stats_kmeans_symbol(),
76        stats_gmm_symbol(),
77        stats_exact_binary_interval_symbol(),
78        stats_paired_bootstrap_symbol(),
79        stats_clustered_bootstrap_symbol(),
80        stats_registered_look_symbol(),
81        stats_isotonic_symbol(),
82        agent_fixtures::fixture_symbols()[0].clone(),
83        agent_fixtures::fixture_symbols()[1].clone(),
84        agent_fixtures::fixture_symbols()[2].clone(),
85        agent_fixtures::fixture_symbols()[3].clone(),
86    ]
87}
88
89#[derive(Clone)]
90struct StatsFunction {
91    symbol: Symbol,
92}
93
94impl Object for StatsFunction {
95    fn display(&self, _cx: &mut Cx) -> Result<String> {
96        Ok(format!("#<function {}>", self.symbol))
97    }
98
99    fn as_any(&self) -> &dyn Any {
100        self
101    }
102}
103
104impl sim_kernel::ObjectCompat for StatsFunction {
105    fn class(&self, cx: &mut Cx) -> Result<ClassRef> {
106        if let Some(value) = cx
107            .registry()
108            .class_by_symbol(&Symbol::qualified("core", "Function"))
109        {
110            return Ok(value.clone());
111        }
112        DefaultFactory.class_stub(
113            sim_kernel::CORE_FUNCTION_CLASS_ID,
114            Symbol::qualified("core", "Function"),
115        )
116    }
117
118    fn as_expr(&self, _cx: &mut Cx) -> Result<Expr> {
119        Ok(Expr::Symbol(self.symbol.clone()))
120    }
121
122    fn as_callable(&self) -> Option<&dyn Callable> {
123        Some(self)
124    }
125}
126
127impl Callable for StatsFunction {
128    fn call(&self, cx: &mut Cx, args: Args) -> Result<Value> {
129        if self.symbol == stats_mean_claim_symbol() {
130            return runtime::call_stats_mean_claim(cx, args);
131        }
132        if self.symbol == stats_variance_claim_symbol() {
133            return runtime::call_stats_variance_claim(cx, args);
134        }
135        if self.symbol == stats_entropy_claim_symbol() {
136            return runtime::call_stats_entropy_claim(cx, args);
137        }
138        if self.symbol == stats_disparate_impact_claim_symbol() {
139            return runtime::call_stats_disparate_impact_claim(cx, args);
140        }
141        if self.symbol == stats_claims_symbol() {
142            return runtime::call_stats_claims(cx, args);
143        }
144        if self.symbol == stats_kmeans_symbol() {
145            return super::runtime_clustering::call_kmeans_values(cx, args.into_vec());
146        }
147        if self.symbol == stats_gmm_symbol() {
148            return super::runtime_clustering::call_gmm_values(cx, args.into_vec());
149        }
150        if super::runtime_decision::is_symbol(&self.symbol) {
151            return super::runtime_decision::call(cx, &self.symbol, args.into_vec());
152        }
153        if let Some(value) = agent_fixtures::call_fixture(cx, &self.symbol, args)? {
154            return Ok(value);
155        }
156        unreachable!("unregistered stats function {}", self.symbol)
157    }
158
159    fn call_exprs(&self, cx: &mut Cx, args: RawArgs) -> Result<Value> {
160        if self.symbol == stats_kmeans_symbol() {
161            return super::runtime_clustering::call_kmeans_exprs(cx, args.into_exprs());
162        }
163        if self.symbol == stats_gmm_symbol() {
164            return super::runtime_clustering::call_gmm_exprs(cx, args.into_exprs());
165        }
166        let values = args
167            .into_exprs()
168            .into_iter()
169            .map(|expr| cx.eval_expr(expr))
170            .collect::<Result<Vec<_>>>()?;
171        self.call(cx, Args::new(values))
172    }
173}
174
175/// Library that installs the runtime statistics functions.
176pub struct StatsNumbersLib;
177
178impl StatsNumbersLib {
179    /// Creates a new statistics runtime library.
180    pub fn new() -> Self {
181        Self
182    }
183}
184
185impl Default for StatsNumbersLib {
186    fn default() -> Self {
187        Self::new()
188    }
189}
190
191impl Lib for StatsNumbersLib {
192    fn manifest(&self) -> LibManifest {
193        LibManifest {
194            id: Symbol::qualified("numbers", "stats"),
195            version: Version(env!("CARGO_PKG_VERSION").to_owned()),
196            abi: AbiVersion { major: 0, minor: 1 },
197            target: LibTarget::HostRegistered,
198            requires: Vec::<Dependency>::new(),
199            capabilities: Vec::new(),
200            exports: function_symbols()
201                .into_iter()
202                .map(|symbol| Export::Function {
203                    symbol,
204                    function_id: None,
205                })
206                .collect(),
207        }
208    }
209
210    fn load(&self, _cx: &mut sim_kernel::LoadCx, linker: &mut Linker<'_>) -> Result<()> {
211        for symbol in function_symbols() {
212            linker.function_value(
213                symbol.clone(),
214                DefaultFactory.opaque(Arc::new(StatsFunction { symbol }))?,
215            )?;
216        }
217        Ok(())
218    }
219}