polydat_nodes/sampling/
lut.rs1use polydat::ast::{CompiledU64Op, NodeMeta, PolydatNode, Port, PortType, Slot, Value};
16
17pub struct LutF64 {
21 lut: Vec<f64>,
24}
25
26impl LutF64 {
27 pub fn from_fn(f: impl Fn(f64) -> f64, resolution: usize) -> Self {
33 assert!(resolution > 0, "resolution must be positive");
34 let mut lut = Vec::with_capacity(resolution + 1);
35 for i in 0..=resolution {
36 let p = i as f64 / resolution as f64;
37 lut.push(f(p));
38 }
39 Self::sanitize(&mut lut);
41 Self { lut }
42 }
43
44 pub fn from_values(values: &[f64]) -> Self {
46 assert!(values.len() >= 2, "LUT must have at least 2 entries");
47 let mut lut = values.to_vec();
48 Self::sanitize(&mut lut);
49 Self { lut }
50 }
51
52 fn sanitize(lut: &mut [f64]) {
54 let mut last_finite = 0.0;
56 let mut found_first = false;
57 for v in lut.iter_mut() {
58 if v.is_finite() {
59 last_finite = *v;
60 found_first = true;
61 } else if found_first {
62 *v = last_finite;
63 }
64 }
65 let mut last_finite = 0.0;
67 for v in lut.iter_mut().rev() {
68 if v.is_finite() {
69 last_finite = *v;
70 } else {
71 *v = last_finite;
72 }
73 }
74 }
75
76 #[inline]
80 pub fn sample(&self, u: f64) -> f64 {
81 let u = u.clamp(0.0, 1.0);
82 let n = (self.lut.len() - 1) as f64;
83 let pos = u * n;
84 let idx = (pos as usize).min(self.lut.len() - 2);
85 let frac = pos - idx as f64;
86 self.lut[idx] * (1.0 - frac) + self.lut[idx + 1] * frac
87 }
88
89 pub fn len(&self) -> usize {
91 self.lut.len()
92 }
93
94 pub fn is_empty(&self) -> bool {
96 self.lut.is_empty()
97 }
98
99 pub fn as_ptr(&self) -> *const f64 {
101 self.lut.as_ptr()
102 }
103
104 pub fn resolution(&self) -> usize {
106 self.lut.len() - 1
107 }
108}
109
110pub struct LutSample {
132 meta: NodeMeta,
133 table: LutF64,
134}
135
136impl LutSample {
137 pub fn new(table: LutF64) -> Self {
139 Self {
140 meta: NodeMeta {
141 name: "lut_sample".into(),
142 outs: vec![Port::new("output", PortType::F64)],
143 ins: vec![Slot::Wire(Port::new("input", PortType::F64))],
144 },
145 table,
146 }
147 }
148}
149
150impl PolydatNode for LutSample {
151 fn meta(&self) -> &NodeMeta {
152 &self.meta
153 }
154
155 fn eval(&self, inputs: &[Value], outputs: &mut [Value]) {
156 outputs[0] = Value::F64(self.table.sample(inputs[0].as_f64()));
157 }
158
159 fn compiled_u64(&self) -> Option<CompiledU64Op> {
160 let lut_addr = self.table.lut.as_ptr() as usize;
164 let lut_len = self.table.lut.len();
165 Some(Box::new(move |inputs, outputs| {
166 let u = f64::from_bits(inputs[0]).clamp(0.0, 1.0);
167 let n = (lut_len - 1) as f64;
168 let pos = u * n;
169 let idx = (pos as usize).min(lut_len - 2);
170 let frac = pos - idx as f64;
171 let result = unsafe {
172 let ptr = lut_addr as *const f64;
173 let a = *ptr.add(idx);
174 let b = *ptr.add(idx + 1);
175 a * (1.0 - frac) + b * frac
176 };
177 outputs[0] = result.to_bits();
178 }))
179 }
180
181 fn jit_constants(&self) -> Vec<u64> {
182 vec![self.table.lut.as_ptr() as u64, self.table.lut.len() as u64]
183 }
184}
185
186impl polydat::derive_support::PolydatSetup for LutF64 {}
191
192fn parse_empirical_lut(spec: &str) -> LutF64 {
195 let mut values: Vec<f64> = spec
196 .split([' ', ',', ';'])
197 .filter(|s| !s.trim().is_empty())
198 .map(|s| {
199 s.trim()
200 .parse::<f64>()
201 .expect("invalid empirical data point")
202 })
203 .collect();
204 assert!(
205 values.len() >= 2,
206 "empirical distribution needs at least 2 data points"
207 );
208 values.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
209 LutF64::from_values(&values)
210}
211
212fn dist_empirical_jit_constants(node: &DistEmpirical) -> Vec<u64> {
215 vec![node.table.as_ptr() as u64, node.table.len() as u64]
216}
217
218#[polydat::polydat_node(category = Probability, jit_constants = dist_empirical_jit_constants)]
223fn dist_empirical(
224 input: f64,
225 spec: polydat::derive_support::Const<&str>,
226 #[poly_const(parse_empirical_lut, from = spec)] table: &LutF64,
227) -> f64 {
228 table.sample(input)
229}
230
231#[cfg(test)]
232mod tests {
233 use super::*;
234
235 #[test]
236 fn lut_identity() {
237 let table = LutF64::from_fn(|p| p, 100);
238 assert!((table.sample(0.0) - 0.0).abs() < 1e-10);
239 assert!((table.sample(0.5) - 0.5).abs() < 0.01);
240 assert!((table.sample(1.0) - 1.0).abs() < 1e-10);
241 }
242
243 #[test]
244 fn lut_quadratic() {
245 let table = LutF64::from_fn(|p| p * p, 1000);
246 assert!((table.sample(0.5) - 0.25).abs() < 0.001);
247 assert!((table.sample(0.0) - 0.0).abs() < 1e-10);
248 assert!((table.sample(1.0) - 1.0).abs() < 0.001);
249 }
250
251 #[test]
252 fn lut_clamps_input() {
253 let table = LutF64::from_fn(|p| p * 10.0, 100);
254 assert!((table.sample(-0.5) - 0.0).abs() < 1e-10);
256 assert!((table.sample(1.5) - 10.0).abs() < 1e-10);
258 }
259
260 #[test]
261 fn lut_sanitizes_infinities() {
262 let table = LutF64::from_fn(
263 |p| {
264 if !(0.01..=0.99).contains(&p) {
265 f64::INFINITY
266 } else {
267 p
268 }
269 },
270 100,
271 );
272 assert!(table.sample(0.0).is_finite());
274 assert!(table.sample(1.0).is_finite());
275 }
276
277 #[test]
278 fn lut_from_values() {
279 let table = LutF64::from_values(&[0.0, 5.0, 10.0]);
280 assert!((table.sample(0.0) - 0.0).abs() < 1e-10);
281 assert!((table.sample(0.5) - 5.0).abs() < 1e-10);
282 assert!((table.sample(1.0) - 10.0).abs() < 1e-10);
283 assert!((table.sample(0.25) - 2.5).abs() < 1e-10);
285 }
286
287 #[test]
288 fn lut_node_eval() {
289 let table = LutF64::from_fn(|p| p * 100.0, 1000);
290 let node = LutSample::new(table);
291 let mut out = [Value::None];
292 node.eval(&[Value::F64(0.5)], &mut out);
293 assert!((out[0].as_f64() - 50.0).abs() < 0.1);
294 }
295}