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
47fn function_symbols() -> [Symbol; 11] {
48    [
49        stats_mean_claim_symbol(),
50        stats_variance_claim_symbol(),
51        stats_entropy_claim_symbol(),
52        stats_disparate_impact_claim_symbol(),
53        stats_claims_symbol(),
54        stats_kmeans_symbol(),
55        stats_gmm_symbol(),
56        agent_fixtures::fixture_symbols()[0].clone(),
57        agent_fixtures::fixture_symbols()[1].clone(),
58        agent_fixtures::fixture_symbols()[2].clone(),
59        agent_fixtures::fixture_symbols()[3].clone(),
60    ]
61}
62
63#[derive(Clone)]
64struct StatsFunction {
65    symbol: Symbol,
66}
67
68impl Object for StatsFunction {
69    fn display(&self, _cx: &mut Cx) -> Result<String> {
70        Ok(format!("#<function {}>", self.symbol))
71    }
72
73    fn as_any(&self) -> &dyn Any {
74        self
75    }
76}
77
78impl sim_kernel::ObjectCompat for StatsFunction {
79    fn class(&self, cx: &mut Cx) -> Result<ClassRef> {
80        if let Some(value) = cx
81            .registry()
82            .class_by_symbol(&Symbol::qualified("core", "Function"))
83        {
84            return Ok(value.clone());
85        }
86        DefaultFactory.class_stub(
87            sim_kernel::CORE_FUNCTION_CLASS_ID,
88            Symbol::qualified("core", "Function"),
89        )
90    }
91
92    fn as_expr(&self, _cx: &mut Cx) -> Result<Expr> {
93        Ok(Expr::Symbol(self.symbol.clone()))
94    }
95
96    fn as_callable(&self) -> Option<&dyn Callable> {
97        Some(self)
98    }
99}
100
101impl Callable for StatsFunction {
102    fn call(&self, cx: &mut Cx, args: Args) -> Result<Value> {
103        if self.symbol == stats_mean_claim_symbol() {
104            return runtime::call_stats_mean_claim(cx, args);
105        }
106        if self.symbol == stats_variance_claim_symbol() {
107            return runtime::call_stats_variance_claim(cx, args);
108        }
109        if self.symbol == stats_entropy_claim_symbol() {
110            return runtime::call_stats_entropy_claim(cx, args);
111        }
112        if self.symbol == stats_disparate_impact_claim_symbol() {
113            return runtime::call_stats_disparate_impact_claim(cx, args);
114        }
115        if self.symbol == stats_claims_symbol() {
116            return runtime::call_stats_claims(cx, args);
117        }
118        if self.symbol == stats_kmeans_symbol() {
119            return super::runtime_clustering::call_kmeans_values(cx, args.into_vec());
120        }
121        if self.symbol == stats_gmm_symbol() {
122            return super::runtime_clustering::call_gmm_values(cx, args.into_vec());
123        }
124        if let Some(value) = agent_fixtures::call_fixture(cx, &self.symbol, args)? {
125            return Ok(value);
126        }
127        unreachable!("unregistered stats function {}", self.symbol)
128    }
129
130    fn call_exprs(&self, cx: &mut Cx, args: RawArgs) -> Result<Value> {
131        if self.symbol == stats_kmeans_symbol() {
132            return super::runtime_clustering::call_kmeans_exprs(cx, args.into_exprs());
133        }
134        if self.symbol == stats_gmm_symbol() {
135            return super::runtime_clustering::call_gmm_exprs(cx, args.into_exprs());
136        }
137        let values = args
138            .into_exprs()
139            .into_iter()
140            .map(|expr| cx.eval_expr(expr))
141            .collect::<Result<Vec<_>>>()?;
142        self.call(cx, Args::new(values))
143    }
144}
145
146/// Library that installs the runtime statistics functions.
147pub struct StatsNumbersLib;
148
149impl StatsNumbersLib {
150    /// Creates a new statistics runtime library.
151    pub fn new() -> Self {
152        Self
153    }
154}
155
156impl Default for StatsNumbersLib {
157    fn default() -> Self {
158        Self::new()
159    }
160}
161
162impl Lib for StatsNumbersLib {
163    fn manifest(&self) -> LibManifest {
164        LibManifest {
165            id: Symbol::qualified("numbers", "stats"),
166            version: Version(env!("CARGO_PKG_VERSION").to_owned()),
167            abi: AbiVersion { major: 0, minor: 1 },
168            target: LibTarget::HostRegistered,
169            requires: Vec::<Dependency>::new(),
170            capabilities: Vec::new(),
171            exports: function_symbols()
172                .into_iter()
173                .map(|symbol| Export::Function {
174                    symbol,
175                    function_id: None,
176                })
177                .collect(),
178        }
179    }
180
181    fn load(&self, _cx: &mut sim_kernel::LoadCx, linker: &mut Linker<'_>) -> Result<()> {
182        for symbol in function_symbols() {
183            linker.function_value(
184                symbol.clone(),
185                DefaultFactory.opaque(Arc::new(StatsFunction { symbol }))?,
186            )?;
187        }
188        Ok(())
189    }
190}