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> {
170 assert!(
171 !ap_embeddings.is_empty(),
172 "DeepSets: input set must be non-empty"
173 );
174 let n = ap_embeddings.len() as f32;
175 let mut pooled = vec![0.0f32; self.dim];
176 for emb in ap_embeddings {
177 debug_assert_eq!(emb.len(), self.dim);
178 let mut t = self.phi.forward(emb);
179 relu(&mut t);
180 for (p, v) in pooled.iter_mut().zip(t.iter()) {
181 *p += *v;
182 }
183 }
184 for p in pooled.iter_mut() {
185 *p /= n;
186 }
187 let mut out = self.rho.forward(&pooled);
188 relu(&mut out);
189 out
190 }
191}
192
193pub struct GeometryEncoder {
199 pos_embed: FourierPositionalEncoding,
200 set_encoder: DeepSets,
201}
202
203impl GeometryEncoder {
204 pub fn new(cfg: &MeridianGeometryConfig) -> Self {
206 GeometryEncoder {
207 pos_embed: FourierPositionalEncoding::new(cfg),
208 set_encoder: DeepSets::new(cfg),
209 }
210 }
211
212 pub fn encode(&self, ap_positions: &[[f32; 3]]) -> Vec<f32> {
214 let embs: Vec<Vec<f32>> = ap_positions
215 .iter()
216 .map(|p| self.pos_embed.encode(p))
217 .collect();
218 self.set_encoder.encode(&embs)
219 }
220}
221
222pub struct FilmLayer {
228 gamma_proj: Linear,
229 beta_proj: Linear,
230}
231
232impl FilmLayer {
233 pub fn new(cfg: &MeridianGeometryConfig) -> Self {
235 let d = cfg.geometry_dim;
236 let mut gamma_proj = Linear::new(d, d, cfg.seed.wrapping_add(3));
237 for b in gamma_proj.bias.iter_mut() {
238 *b = 1.0;
239 }
240 FilmLayer {
241 gamma_proj,
242 beta_proj: Linear::new(d, d, cfg.seed.wrapping_add(4)),
243 }
244 }
245
246 pub fn modulate(&self, features: &[f32], geometry: &[f32]) -> Vec<f32> {
248 let gamma = self.gamma_proj.forward(geometry);
249 let beta = self.beta_proj.forward(geometry);
250 features
251 .iter()
252 .zip(gamma.iter())
253 .zip(beta.iter())
254 .map(|((&f, &g), &b)| g * f + b)
255 .collect()
256 }
257}
258
259#[cfg(test)]
264mod tests {
265 use super::*;
266
267 fn cfg() -> MeridianGeometryConfig {
268 MeridianGeometryConfig::default()
269 }
270
271 #[test]
272 fn fourier_output_dimension_is_64() {
273 let c = cfg();
274 let out = FourierPositionalEncoding::new(&c).encode(&[1.0, 2.0, 3.0]);
275 assert_eq!(out.len(), c.geometry_dim);
276 }
277
278 #[test]
279 fn fourier_different_coords_different_outputs() {
280 let enc = FourierPositionalEncoding::new(&cfg());
281 let a = enc.encode(&[0.0, 0.0, 0.0]);
282 let b = enc.encode(&[1.0, 0.0, 0.0]);
283 let c = enc.encode(&[0.0, 1.0, 0.0]);
284 let d = enc.encode(&[0.0, 0.0, 1.0]);
285 assert_ne!(a, b);
286 assert_ne!(a, c);
287 assert_ne!(a, d);
288 assert_ne!(b, c);
289 }
290
291 #[test]
292 fn fourier_values_bounded() {
293 let out = FourierPositionalEncoding::new(&cfg()).encode(&[5.5, -3.2, 0.1]);
294 for &v in &out {
295 assert!(v.abs() <= 1.0 + 1e-6, "got {v}");
296 }
297 }
298
299 #[test]
300 fn deepsets_permutation_invariant() {
301 let c = cfg();
302 let enc = FourierPositionalEncoding::new(&c);
303 let ds = DeepSets::new(&c);
304 let (a, b, d) = (
305 enc.encode(&[1.0, 0.0, 0.0]),
306 enc.encode(&[0.0, 2.0, 0.0]),
307 enc.encode(&[0.0, 0.0, 3.0]),
308 );
309 let abc = ds.encode(&[a.clone(), b.clone(), d.clone()]);
310 let cba = ds.encode(&[d.clone(), b.clone(), a.clone()]);
311 let bac = ds.encode(&[b.clone(), a.clone(), d.clone()]);
312 for i in 0..c.geometry_dim {
313 assert!(
314 (abc[i] - cba[i]).abs() < 1e-5,
315 "dim {i}: abc={} cba={}",
316 abc[i],
317 cba[i]
318 );
319 assert!(
320 (abc[i] - bac[i]).abs() < 1e-5,
321 "dim {i}: abc={} bac={}",
322 abc[i],
323 bac[i]
324 );
325 }
326 }
327
328 #[test]
329 fn deepsets_variable_ap_count() {
330 let c = cfg();
331 let enc = FourierPositionalEncoding::new(&c);
332 let ds = DeepSets::new(&c);
333 let one = ds.encode(&[enc.encode(&[1.0, 0.0, 0.0])]);
334 assert_eq!(one.len(), c.geometry_dim);
335 let three = ds.encode(&[
336 enc.encode(&[1.0, 0.0, 0.0]),
337 enc.encode(&[0.0, 2.0, 0.0]),
338 enc.encode(&[0.0, 0.0, 3.0]),
339 ]);
340 assert_eq!(three.len(), c.geometry_dim);
341 let six = ds.encode(&[
342 enc.encode(&[1.0, 0.0, 0.0]),
343 enc.encode(&[0.0, 2.0, 0.0]),
344 enc.encode(&[0.0, 0.0, 3.0]),
345 enc.encode(&[-1.0, 0.0, 0.0]),
346 enc.encode(&[0.0, -2.0, 0.0]),
347 enc.encode(&[0.0, 0.0, -3.0]),
348 ]);
349 assert_eq!(six.len(), c.geometry_dim);
350 assert_ne!(one, three);
351 assert_ne!(three, six);
352 }
353
354 #[test]
355 fn geometry_encoder_end_to_end() {
356 let c = cfg();
357 let g =
358 GeometryEncoder::new(&c).encode(&[[1.0, 0.0, 2.5], [0.0, 3.0, 2.5], [-2.0, 1.0, 2.5]]);
359 assert_eq!(g.len(), c.geometry_dim);
360 for &v in &g {
361 assert!(v.is_finite());
362 }
363 }
364
365 #[test]
366 fn geometry_encoder_single_ap() {
367 let c = cfg();
368 assert_eq!(
369 GeometryEncoder::new(&c).encode(&[[0.0, 0.0, 0.0]]).len(),
370 c.geometry_dim
371 );
372 }
373
374 #[test]
375 fn film_identity_when_geometry_zero() {
376 let c = cfg();
377 let film = FilmLayer::new(&c);
378 let feat = vec![1.0f32; c.geometry_dim];
379 let out = film.modulate(&feat, &vec![0.0f32; c.geometry_dim]);
380 assert_eq!(out.len(), c.geometry_dim);
381 for i in 0..c.geometry_dim {
383 assert!(
384 (out[i] - feat[i]).abs() < 1e-5,
385 "dim {i}: expected {}, got {}",
386 feat[i],
387 out[i]
388 );
389 }
390 }
391
392 #[test]
393 fn film_nontrivial_modulation() {
394 let c = cfg();
395 let film = FilmLayer::new(&c);
396 let feat: Vec<f32> = (0..c.geometry_dim).map(|i| i as f32 * 0.1).collect();
397 let geom: Vec<f32> = (0..c.geometry_dim)
398 .map(|i| (i as f32 - 32.0) * 0.01)
399 .collect();
400 let out = film.modulate(&feat, &geom);
401 assert_eq!(out.len(), c.geometry_dim);
402 assert!(out
403 .iter()
404 .zip(feat.iter())
405 .any(|(o, f)| (o - f).abs() > 1e-6));
406 for &v in &out {
407 assert!(v.is_finite());
408 }
409 }
410
411 #[test]
412 fn film_explicit_gamma_beta() {
413 let c = MeridianGeometryConfig {
414 geometry_dim: 4,
415 ..cfg()
416 };
417 let mut film = FilmLayer::new(&c);
418 film.gamma_proj.weights = vec![0.0; 16];
419 film.gamma_proj.bias = vec![2.0, 3.0, 0.5, 1.0];
420 film.beta_proj.weights = vec![0.0; 16];
421 film.beta_proj.bias = vec![10.0, 20.0, 30.0, 40.0];
422 let out = film.modulate(&[1.0, 2.0, 3.0, 4.0], &[999.0; 4]);
423 let exp = [12.0, 26.0, 31.5, 44.0];
424 for i in 0..4 {
425 assert!((out[i] - exp[i]).abs() < 1e-5, "dim {i}");
426 }
427 }
428
429 #[test]
430 fn config_defaults() {
431 let c = MeridianGeometryConfig::default();
432 assert_eq!(c.n_frequencies, 10);
433 assert!((c.scale - 1.0).abs() < 1e-6);
434 assert_eq!(c.geometry_dim, 64);
435 assert_eq!(c.seed, 42);
436 }
437
438 #[test]
439 fn config_serde_round_trip() {
440 let c = MeridianGeometryConfig {
441 n_frequencies: 8,
442 scale: 0.5,
443 geometry_dim: 32,
444 seed: 123,
445 };
446 let j = serde_json::to_string(&c).unwrap();
447 let d: MeridianGeometryConfig = serde_json::from_str(&j).unwrap();
448 assert_eq!(d.n_frequencies, 8);
449 assert!((d.scale - 0.5).abs() < 1e-6);
450 assert_eq!(d.geometry_dim, 32);
451 assert_eq!(d.seed, 123);
452 }
453
454 #[test]
455 fn linear_forward_dim() {
456 assert_eq!(Linear::new(8, 4, 0).forward(&[1.0; 8]).len(), 4);
457 }
458
459 #[test]
460 fn linear_zero_input_gives_bias() {
461 let lin = Linear::new(4, 3, 0);
462 let out = lin.forward(&[0.0; 4]);
463 for (oi, bi) in out.iter().zip(lin.bias.iter()) {
464 assert!((oi - bi).abs() < 1e-6);
465 }
466 }
467}