1use serde::{Deserialize, Serialize};
8
9const GEOMETRY_DIM: usize = 64;
10const NUM_COORDS: usize = 3;
11
12#[derive(Debug, Clone)]
18struct Linear {
19 weights: Vec<f32>,
20 bias: Vec<f32>,
21 in_f: usize,
22 #[allow(dead_code)]
23 out_f: usize,
24}
25
26impl Linear {
27 fn new(in_f: usize, out_f: usize, seed: u64) -> Self {
29 let k = (1.0 / in_f as f32).sqrt();
30 Linear {
31 weights: det_uniform(in_f * out_f, -k, k, seed),
32 bias: vec![0.0; out_f],
33 in_f,
34 out_f,
35 }
36 }
37
38 fn forward(&self, x: &[f32]) -> Vec<f32> {
39 debug_assert_eq!(x.len(), self.in_f);
40 let mut y = self.bias.clone();
41 for (j, yj) in y.iter_mut().enumerate() {
42 let off = j * self.in_f;
43 let s: f32 = x
44 .iter()
45 .zip(self.weights[off..off + self.in_f].iter())
46 .map(|(&xi, &wi)| xi * wi)
47 .sum();
48 *yj += s;
49 }
50 y
51 }
52}
53
54fn det_uniform(n: usize, lo: f32, hi: f32, seed: u64) -> Vec<f32> {
57 let r = hi - lo;
58 let mut s = seed.wrapping_add(0x9E37_79B9_7F4A_7C15);
59 (0..n)
60 .map(|_| {
61 s ^= s << 13;
62 s ^= s >> 7;
63 s ^= s << 17;
64 lo + (s >> 40) as f32 / (1u64 << 24) as f32 * r
65 })
66 .collect()
67}
68
69fn relu(v: &mut [f32]) {
70 for x in v.iter_mut() {
71 if *x < 0.0 {
72 *x = 0.0;
73 }
74 }
75}
76
77#[derive(Debug, Clone, Serialize, Deserialize)]
83pub struct MeridianGeometryConfig {
84 pub n_frequencies: usize,
86 pub scale: f32,
88 pub geometry_dim: usize,
90 pub seed: u64,
92}
93
94impl Default for MeridianGeometryConfig {
95 fn default() -> Self {
96 MeridianGeometryConfig {
97 n_frequencies: 10,
98 scale: 1.0,
99 geometry_dim: GEOMETRY_DIM,
100 seed: 42,
101 }
102 }
103}
104
105pub struct FourierPositionalEncoding {
114 n_frequencies: usize,
115 scale: f32,
116 output_dim: usize,
117}
118
119impl FourierPositionalEncoding {
120 pub fn new(cfg: &MeridianGeometryConfig) -> Self {
122 FourierPositionalEncoding {
123 n_frequencies: cfg.n_frequencies,
124 scale: cfg.scale,
125 output_dim: cfg.geometry_dim,
126 }
127 }
128
129 pub fn encode(&self, coords: &[f32; 3]) -> Vec<f32> {
131 let raw = NUM_COORDS * 2 * self.n_frequencies;
132 let mut enc = Vec::with_capacity(raw.max(self.output_dim));
133 for &c in coords {
134 let sc = c * self.scale;
135 for l in 0..self.n_frequencies {
136 let f = (2.0f32).powi(l as i32) * std::f32::consts::PI * sc;
137 enc.push(f.sin());
138 enc.push(f.cos());
139 }
140 }
141 enc.resize(self.output_dim, 0.0);
142 enc
143 }
144}
145
146pub struct DeepSets {
152 phi: Linear,
153 rho: Linear,
154 dim: usize,
155}
156
157impl DeepSets {
158 pub fn new(cfg: &MeridianGeometryConfig) -> Self {
160 let d = cfg.geometry_dim;
161 DeepSets {
162 phi: Linear::new(d, d, cfg.seed.wrapping_add(1)),
163 rho: Linear::new(d, d, cfg.seed.wrapping_add(2)),
164 dim: d,
165 }
166 }
167
168 pub fn encode(&self, ap_embeddings: &[Vec<f32>]) -> Vec<f32> {
177 assert!(
178 !ap_embeddings.is_empty(),
179 "DeepSets: input set must be non-empty"
180 );
181 let n = ap_embeddings.len() as f32;
182 let mut pooled = vec![0.0f32; self.dim];
183 for emb in ap_embeddings {
184 debug_assert_eq!(emb.len(), self.dim);
185 let mut t = self.phi.forward(emb);
186 relu(&mut t);
187 for (p, v) in pooled.iter_mut().zip(t.iter()) {
188 *p += *v;
189 }
190 }
191 for p in pooled.iter_mut() {
192 *p /= n;
193 }
194 let mut out = self.rho.forward(&pooled);
195 relu(&mut out);
196 out
197 }
198}
199
200pub struct GeometryEncoder {
206 pos_embed: FourierPositionalEncoding,
207 set_encoder: DeepSets,
208}
209
210impl GeometryEncoder {
211 pub fn new(cfg: &MeridianGeometryConfig) -> Self {
213 GeometryEncoder {
214 pos_embed: FourierPositionalEncoding::new(cfg),
215 set_encoder: DeepSets::new(cfg),
216 }
217 }
218
219 pub fn encode(&self, ap_positions: &[[f32; 3]]) -> Vec<f32> {
221 let embs: Vec<Vec<f32>> = ap_positions
222 .iter()
223 .map(|p| self.pos_embed.encode(p))
224 .collect();
225 self.set_encoder.encode(&embs)
226 }
227}
228
229pub struct FilmLayer {
235 gamma_proj: Linear,
236 beta_proj: Linear,
237}
238
239impl FilmLayer {
240 pub fn new(cfg: &MeridianGeometryConfig) -> Self {
242 let d = cfg.geometry_dim;
243 let mut gamma_proj = Linear::new(d, d, cfg.seed.wrapping_add(3));
244 for b in gamma_proj.bias.iter_mut() {
245 *b = 1.0;
246 }
247 FilmLayer {
248 gamma_proj,
249 beta_proj: Linear::new(d, d, cfg.seed.wrapping_add(4)),
250 }
251 }
252
253 pub fn modulate(&self, features: &[f32], geometry: &[f32]) -> Vec<f32> {
255 let gamma = self.gamma_proj.forward(geometry);
256 let beta = self.beta_proj.forward(geometry);
257 features
258 .iter()
259 .zip(gamma.iter())
260 .zip(beta.iter())
261 .map(|((&f, &g), &b)| g * f + b)
262 .collect()
263 }
264}
265
266#[cfg(test)]
271mod tests {
272 use super::*;
273
274 fn cfg() -> MeridianGeometryConfig {
275 MeridianGeometryConfig::default()
276 }
277
278 #[test]
279 fn fourier_output_dimension_is_64() {
280 let c = cfg();
281 let out = FourierPositionalEncoding::new(&c).encode(&[1.0, 2.0, 3.0]);
282 assert_eq!(out.len(), c.geometry_dim);
283 }
284
285 #[test]
286 fn fourier_different_coords_different_outputs() {
287 let enc = FourierPositionalEncoding::new(&cfg());
288 let a = enc.encode(&[0.0, 0.0, 0.0]);
289 let b = enc.encode(&[1.0, 0.0, 0.0]);
290 let c = enc.encode(&[0.0, 1.0, 0.0]);
291 let d = enc.encode(&[0.0, 0.0, 1.0]);
292 assert_ne!(a, b);
293 assert_ne!(a, c);
294 assert_ne!(a, d);
295 assert_ne!(b, c);
296 }
297
298 #[test]
299 fn fourier_values_bounded() {
300 let out = FourierPositionalEncoding::new(&cfg()).encode(&[5.5, -3.2, 0.1]);
301 for &v in &out {
302 assert!(v.abs() <= 1.0 + 1e-6, "got {v}");
303 }
304 }
305
306 #[test]
307 fn deepsets_permutation_invariant() {
308 let c = cfg();
309 let enc = FourierPositionalEncoding::new(&c);
310 let ds = DeepSets::new(&c);
311 let (a, b, d) = (
312 enc.encode(&[1.0, 0.0, 0.0]),
313 enc.encode(&[0.0, 2.0, 0.0]),
314 enc.encode(&[0.0, 0.0, 3.0]),
315 );
316 let abc = ds.encode(&[a.clone(), b.clone(), d.clone()]);
317 let cba = ds.encode(&[d.clone(), b.clone(), a.clone()]);
318 let bac = ds.encode(&[b.clone(), a.clone(), d.clone()]);
319 for i in 0..c.geometry_dim {
320 assert!(
321 (abc[i] - cba[i]).abs() < 1e-5,
322 "dim {i}: abc={} cba={}",
323 abc[i],
324 cba[i]
325 );
326 assert!(
327 (abc[i] - bac[i]).abs() < 1e-5,
328 "dim {i}: abc={} bac={}",
329 abc[i],
330 bac[i]
331 );
332 }
333 }
334
335 #[test]
336 fn deepsets_variable_ap_count() {
337 let c = cfg();
338 let enc = FourierPositionalEncoding::new(&c);
339 let ds = DeepSets::new(&c);
340 let one = ds.encode(&[enc.encode(&[1.0, 0.0, 0.0])]);
341 assert_eq!(one.len(), c.geometry_dim);
342 let three = ds.encode(&[
343 enc.encode(&[1.0, 0.0, 0.0]),
344 enc.encode(&[0.0, 2.0, 0.0]),
345 enc.encode(&[0.0, 0.0, 3.0]),
346 ]);
347 assert_eq!(three.len(), c.geometry_dim);
348 let six = ds.encode(&[
349 enc.encode(&[1.0, 0.0, 0.0]),
350 enc.encode(&[0.0, 2.0, 0.0]),
351 enc.encode(&[0.0, 0.0, 3.0]),
352 enc.encode(&[-1.0, 0.0, 0.0]),
353 enc.encode(&[0.0, -2.0, 0.0]),
354 enc.encode(&[0.0, 0.0, -3.0]),
355 ]);
356 assert_eq!(six.len(), c.geometry_dim);
357 assert_ne!(one, three);
358 assert_ne!(three, six);
359 }
360
361 #[test]
362 fn geometry_encoder_end_to_end() {
363 let c = cfg();
364 let g =
365 GeometryEncoder::new(&c).encode(&[[1.0, 0.0, 2.5], [0.0, 3.0, 2.5], [-2.0, 1.0, 2.5]]);
366 assert_eq!(g.len(), c.geometry_dim);
367 for &v in &g {
368 assert!(v.is_finite());
369 }
370 }
371
372 #[test]
373 fn geometry_encoder_single_ap() {
374 let c = cfg();
375 assert_eq!(
376 GeometryEncoder::new(&c).encode(&[[0.0, 0.0, 0.0]]).len(),
377 c.geometry_dim
378 );
379 }
380
381 #[test]
382 fn film_identity_when_geometry_zero() {
383 let c = cfg();
384 let film = FilmLayer::new(&c);
385 let feat = vec![1.0f32; c.geometry_dim];
386 let out = film.modulate(&feat, &vec![0.0f32; c.geometry_dim]);
387 assert_eq!(out.len(), c.geometry_dim);
388 for i in 0..c.geometry_dim {
390 assert!(
391 (out[i] - feat[i]).abs() < 1e-5,
392 "dim {i}: expected {}, got {}",
393 feat[i],
394 out[i]
395 );
396 }
397 }
398
399 #[test]
400 fn film_nontrivial_modulation() {
401 let c = cfg();
402 let film = FilmLayer::new(&c);
403 let feat: Vec<f32> = (0..c.geometry_dim).map(|i| i as f32 * 0.1).collect();
404 let geom: Vec<f32> = (0..c.geometry_dim)
405 .map(|i| (i as f32 - 32.0) * 0.01)
406 .collect();
407 let out = film.modulate(&feat, &geom);
408 assert_eq!(out.len(), c.geometry_dim);
409 assert!(out
410 .iter()
411 .zip(feat.iter())
412 .any(|(o, f)| (o - f).abs() > 1e-6));
413 for &v in &out {
414 assert!(v.is_finite());
415 }
416 }
417
418 #[test]
419 fn film_explicit_gamma_beta() {
420 let c = MeridianGeometryConfig {
421 geometry_dim: 4,
422 ..cfg()
423 };
424 let mut film = FilmLayer::new(&c);
425 film.gamma_proj.weights = vec![0.0; 16];
426 film.gamma_proj.bias = vec![2.0, 3.0, 0.5, 1.0];
427 film.beta_proj.weights = vec![0.0; 16];
428 film.beta_proj.bias = vec![10.0, 20.0, 30.0, 40.0];
429 let out = film.modulate(&[1.0, 2.0, 3.0, 4.0], &[999.0; 4]);
430 let exp = [12.0, 26.0, 31.5, 44.0];
431 for i in 0..4 {
432 assert!((out[i] - exp[i]).abs() < 1e-5, "dim {i}");
433 }
434 }
435
436 #[test]
437 fn config_defaults() {
438 let c = MeridianGeometryConfig::default();
439 assert_eq!(c.n_frequencies, 10);
440 assert!((c.scale - 1.0).abs() < 1e-6);
441 assert_eq!(c.geometry_dim, 64);
442 assert_eq!(c.seed, 42);
443 }
444
445 #[test]
446 fn config_serde_round_trip() {
447 let c = MeridianGeometryConfig {
448 n_frequencies: 8,
449 scale: 0.5,
450 geometry_dim: 32,
451 seed: 123,
452 };
453 let j = serde_json::to_string(&c).unwrap();
454 let d: MeridianGeometryConfig = serde_json::from_str(&j).unwrap();
455 assert_eq!(d.n_frequencies, 8);
456 assert!((d.scale - 0.5).abs() < 1e-6);
457 assert_eq!(d.geometry_dim, 32);
458 assert_eq!(d.seed, 123);
459 }
460
461 #[test]
462 fn linear_forward_dim() {
463 assert_eq!(Linear::new(8, 4, 0).forward(&[1.0; 8]).len(), 4);
464 }
465
466 #[test]
467 fn linear_zero_input_gives_bias() {
468 let lin = Linear::new(4, 3, 0);
469 let out = lin.forward(&[0.0; 4]);
470 for (oi, bi) in out.iter().zip(lin.bias.iter()) {
471 assert!((oi - bi).abs() < 1e-6);
472 }
473 }
474}