Skip to main content

proof_engine/svogi/
sh.rs

1use glam::{Vec3, Mat4};
2
3/// 2nd order spherical harmonics (L=0,1,2 => 9 coefficients).
4#[derive(Debug, Clone, Copy)]
5pub struct SH2 {
6    pub coeffs: [f32; 9],
7}
8
9impl Default for SH2 {
10    fn default() -> Self {
11        Self { coeffs: [0.0; 9] }
12    }
13}
14
15/// 3rd order spherical harmonics (L=0,1,2,3 => 16 coefficients).
16#[derive(Debug, Clone, Copy)]
17pub struct SH3 {
18    pub coeffs: [f32; 16],
19}
20
21impl Default for SH3 {
22    fn default() -> Self {
23        Self { coeffs: [0.0; 16] }
24    }
25}
26
27/// Factorial helper.
28fn factorial(n: u32) -> f64 {
29    (1..=n as u64).fold(1.0f64, |acc, x| acc * x as f64)
30}
31
32/// Double factorial: n!! = n * (n-2) * (n-4) * ...
33fn double_factorial(n: i32) -> f64 {
34    if n <= 0 { return 1.0; }
35    let mut result = 1.0f64;
36    let mut k = n;
37    while k > 0 {
38        result *= k as f64;
39        k -= 2;
40    }
41    result
42}
43
44/// Associated Legendre polynomial P_l^m(x).
45pub fn legendre_p(l: i32, m: i32, x: f64) -> f64 {
46    let m_abs = m.abs();
47
48    if m_abs > l {
49        return 0.0;
50    }
51
52    // Compute P_m^m using the formula P_m^m(x) = (-1)^m * (2m-1)!! * (1-x^2)^(m/2)
53    let mut pmm = 1.0f64;
54    if m_abs > 0 {
55        let somx2 = ((1.0 - x) * (1.0 + x)).sqrt();
56        let mut fact = 1.0f64;
57        for i in 1..=m_abs {
58            pmm *= -fact * somx2;
59            fact += 2.0;
60        }
61    }
62
63    if l == m_abs {
64        if m < 0 {
65            let sign = if m_abs % 2 == 0 { 1.0 } else { -1.0 };
66            return sign * factorial((l - m_abs) as u32) as f64 / factorial((l + m_abs) as u32) as f64 * pmm;
67        }
68        return pmm;
69    }
70
71    // P_{m+1}^m(x) = x * (2m+1) * P_m^m(x)
72    let mut pmmp1 = x * (2 * m_abs + 1) as f64 * pmm;
73
74    if l == m_abs + 1 {
75        if m < 0 {
76            let sign = if m_abs % 2 == 0 { 1.0 } else { -1.0 };
77            return sign * factorial((l - m_abs) as u32) as f64 / factorial((l + m_abs) as u32) as f64 * pmmp1;
78        }
79        return pmmp1;
80    }
81
82    // Use recurrence: (l-m)*P_l^m = x*(2l-1)*P_{l-1}^m - (l+m-1)*P_{l-2}^m
83    let mut pll = 0.0f64;
84    for ll in (m_abs + 2)..=l {
85        pll = (x * (2 * ll - 1) as f64 * pmmp1 - (ll + m_abs - 1) as f64 * pmm) / (ll - m_abs) as f64;
86        pmm = pmmp1;
87        pmmp1 = pll;
88    }
89
90    if m < 0 {
91        let sign = if m_abs % 2 == 0 { 1.0 } else { -1.0 };
92        return sign * factorial((l - m_abs) as u32) as f64 / factorial((l + m_abs) as u32) as f64 * pll;
93    }
94    pll
95}
96
97/// SH normalization constant K_l^m.
98fn sh_k(l: i32, m: i32) -> f64 {
99    let m_abs = m.abs();
100    let num = (2 * l + 1) as f64 * factorial((l - m_abs) as u32) as f64;
101    let den = 4.0 * std::f64::consts::PI * factorial((l + m_abs) as u32) as f64;
102    (num / den).sqrt()
103}
104
105/// Evaluate the real SH basis function Y_l^m at direction (theta, phi).
106fn sh_basis_real(l: i32, m: i32, theta: f64, phi: f64) -> f64 {
107    let k = sh_k(l, m);
108    let p = legendre_p(l, m.abs(), theta.cos());
109    if m > 0 {
110        std::f64::consts::SQRT_2 * k * (m as f64 * phi).cos() * p
111    } else if m < 0 {
112        std::f64::consts::SQRT_2 * k * ((-m) as f64 * phi).sin() * p
113    } else {
114        k * p
115    }
116}
117
118/// Direction to spherical coordinates (theta, phi).
119fn dir_to_spherical(dir: Vec3) -> (f64, f64) {
120    let d = dir.normalize_or_zero();
121    let theta = (d.z as f64).acos();
122    let phi = (d.y as f64).atan2(d.x as f64);
123    (theta, phi)
124}
125
126/// Evaluate 2nd order SH basis at a direction. Returns 9 values.
127pub fn sh_basis_2(dir: Vec3) -> [f32; 9] {
128    let d = dir.normalize_or_zero();
129    let x = d.x;
130    let y = d.y;
131    let z = d.z;
132
133    [
134        0.282095,                            // Y_0^0
135        0.488603 * y,                        // Y_1^{-1}
136        0.488603 * z,                        // Y_1^0
137        0.488603 * x,                        // Y_1^1
138        1.092548 * x * y,                    // Y_2^{-2}
139        1.092548 * y * z,                    // Y_2^{-1}
140        0.315392 * (3.0 * z * z - 1.0),     // Y_2^0
141        1.092548 * x * z,                    // Y_2^1
142        0.546274 * (x * x - y * y),          // Y_2^2
143    ]
144}
145
146/// Evaluate 3rd order SH basis at a direction. Returns 16 values.
147pub fn sh_basis_3(dir: Vec3) -> [f32; 16] {
148    let d = dir.normalize_or_zero();
149    let x = d.x;
150    let y = d.y;
151    let z = d.z;
152
153    let mut result = [0.0f32; 16];
154
155    // L=0
156    result[0] = 0.282095;
157
158    // L=1
159    result[1] = 0.488603 * y;
160    result[2] = 0.488603 * z;
161    result[3] = 0.488603 * x;
162
163    // L=2
164    result[4] = 1.092548 * x * y;
165    result[5] = 1.092548 * y * z;
166    result[6] = 0.315392 * (3.0 * z * z - 1.0);
167    result[7] = 1.092548 * x * z;
168    result[8] = 0.546274 * (x * x - y * y);
169
170    // L=3
171    result[9]  = 0.590044 * y * (3.0 * x * x - y * y);
172    result[10] = 2.890611 * x * y * z;
173    result[11] = 0.457046 * y * (5.0 * z * z - 1.0);
174    result[12] = 0.373176 * z * (5.0 * z * z - 3.0);
175    result[13] = 0.457046 * x * (5.0 * z * z - 1.0);
176    result[14] = 1.445306 * z * (x * x - y * y);
177    result[15] = 0.590044 * x * (x * x - 3.0 * y * y);
178
179    result
180}
181
182/// Evaluate SH at a direction given coefficients.
183pub fn sh_evaluate(coeffs: &[f32], dir: Vec3) -> f32 {
184    if coeffs.len() >= 16 {
185        let basis = sh_basis_3(dir);
186        coeffs.iter().zip(basis.iter()).take(16).map(|(c, b)| c * b).sum()
187    } else if coeffs.len() >= 9 {
188        let basis = sh_basis_2(dir);
189        coeffs.iter().zip(basis.iter()).take(9).map(|(c, b)| c * b).sum()
190    } else {
191        let basis = sh_basis_2(dir);
192        coeffs.iter().zip(basis.iter()).map(|(c, b)| c * b).sum()
193    }
194}
195
196/// Project a function into 2nd order SH via Monte Carlo sampling.
197pub fn sh_project_function(
198    sample_fn: impl Fn(Vec3) -> f32,
199    num_samples: usize,
200) -> SH2 {
201    let mut result = SH2::default();
202    let weight = 4.0 * std::f32::consts::PI / num_samples as f32;
203
204    // Use stratified sampling on the sphere
205    let n_sqrt = (num_samples as f32).sqrt().ceil() as usize;
206    let mut count = 0;
207    for i in 0..n_sqrt {
208        for j in 0..n_sqrt {
209            if count >= num_samples {
210                break;
211            }
212            // Stratified spherical coordinates
213            let u = (i as f32 + 0.5) / n_sqrt as f32;
214            let v = (j as f32 + 0.5) / n_sqrt as f32;
215
216            let theta = (1.0 - 2.0 * u).acos();
217            let phi = 2.0 * std::f32::consts::PI * v;
218
219            let dir = Vec3::new(
220                theta.sin() * phi.cos(),
221                theta.sin() * phi.sin(),
222                theta.cos(),
223            );
224
225            let value = sample_fn(dir);
226            let basis = sh_basis_2(dir);
227            for k in 0..9 {
228                result.coeffs[k] += value * basis[k] * weight;
229            }
230            count += 1;
231        }
232    }
233
234    result
235}
236
237/// Convolve SH with a zonal harmonic kernel.
238pub fn sh_convolve(a: &SH2, kernel: &[f32]) -> SH2 {
239    let mut result = SH2::default();
240    // Band 0: 1 coefficient
241    if kernel.len() > 0 {
242        result.coeffs[0] = a.coeffs[0] * kernel[0];
243    }
244    // Band 1: 3 coefficients
245    if kernel.len() > 1 {
246        for i in 1..4 {
247            result.coeffs[i] = a.coeffs[i] * kernel[1];
248        }
249    }
250    // Band 2: 5 coefficients
251    if kernel.len() > 2 {
252        for i in 4..9 {
253            result.coeffs[i] = a.coeffs[i] * kernel[2];
254        }
255    }
256    result
257}
258
259/// Rotate SH by a rotation matrix. Uses the ZYZ Euler angle decomposition.
260pub fn sh_rotate(coeffs: &SH2, rotation: Mat4) -> SH2 {
261    let mut result = SH2::default();
262
263    // Band 0 is rotation-invariant
264    result.coeffs[0] = coeffs.coeffs[0];
265
266    // Band 1: rotate using the upper-left 3x3 of the matrix
267    // The band-1 SH transform directly as a 3-vector
268    let r = rotation;
269    let sh1 = [coeffs.coeffs[3], coeffs.coeffs[1], coeffs.coeffs[2]]; // (x, y, z) order
270
271    // Apply rotation
272    let rx = r.x_axis;
273    let ry = r.y_axis;
274    let rz = r.z_axis;
275
276    let rotated_x = rx.x * sh1[0] + ry.x * sh1[1] + rz.x * sh1[2];
277    let rotated_y = rx.y * sh1[0] + ry.y * sh1[1] + rz.y * sh1[2];
278    let rotated_z = rx.z * sh1[0] + ry.z * sh1[1] + rz.z * sh1[2];
279
280    result.coeffs[3] = rotated_x; // Y_1^1  -> x
281    result.coeffs[1] = rotated_y; // Y_1^-1 -> y
282    result.coeffs[2] = rotated_z; // Y_1^0  -> z
283
284    // Band 2: approximate rotation via reprojection
285    // For each of the 5 band-2 coefficients, we reproject
286    // This is a simplified approach using the real SH rotation property
287    let dirs = [
288        Vec3::X, Vec3::Y, Vec3::Z,
289        Vec3::new(1.0, 1.0, 0.0).normalize(),
290        Vec3::new(1.0, 0.0, 1.0).normalize(),
291        Vec3::new(0.0, 1.0, 1.0).normalize(),
292        Vec3::new(1.0, -1.0, 0.0).normalize(),
293        Vec3::new(-1.0, 0.0, 1.0).normalize(),
294        Vec3::new(0.0, -1.0, 1.0).normalize(),
295    ];
296
297    // Compute band-2 rotation matrix using sampling
298    let mut band2_coeffs = [0.0f32; 5];
299    let weight = 4.0 * std::f32::consts::PI / dirs.len() as f32;
300    for &dir in &dirs {
301        let original_basis = sh_basis_2(dir);
302        let original_val: f32 = (4..9).map(|i| coeffs.coeffs[i] * original_basis[i]).sum();
303
304        let rot3 = glam::Mat3::from_mat4(rotation);
305        let rotated_dir = rot3 * dir;
306        let rotated_basis = sh_basis_2(rotated_dir);
307
308        for i in 0..5 {
309            band2_coeffs[i] += original_val * rotated_basis[i + 4] * weight;
310        }
311    }
312
313    for i in 0..5 {
314        result.coeffs[4 + i] = band2_coeffs[i];
315    }
316
317    result
318}
319
320/// Add two SH2.
321pub fn sh_add(a: &SH2, b: &SH2) -> SH2 {
322    let mut result = SH2::default();
323    for i in 0..9 {
324        result.coeffs[i] = a.coeffs[i] + b.coeffs[i];
325    }
326    result
327}
328
329/// Scale SH2.
330pub fn sh_scale(a: &SH2, s: f32) -> SH2 {
331    let mut result = SH2::default();
332    for i in 0..9 {
333        result.coeffs[i] = a.coeffs[i] * s;
334    }
335    result
336}
337
338/// Inner product of two SH2.
339pub fn sh_dot(a: &SH2, b: &SH2) -> f32 {
340    (0..9).map(|i| a.coeffs[i] * b.coeffs[i]).sum()
341}
342
343/// SH representation of a clamped cosine lobe (for diffuse transfer).
344pub fn cosine_lobe_sh() -> SH2 {
345    // Clamped cosine along +Z in SH
346    let mut sh = SH2::default();
347    sh.coeffs[0] = 0.886227;   // sqrt(pi) / 2
348    sh.coeffs[2] = 1.023326;   // sqrt(pi/3)
349    sh.coeffs[6] = 0.495415;   // sqrt(5*pi) / 8
350    sh
351}
352
353/// SH irradiance probe.
354#[derive(Debug, Clone)]
355pub struct SHProbe {
356    pub position: Vec3,
357    pub sh_r: SH2,
358    pub sh_g: SH2,
359    pub sh_b: SH2,
360}
361
362impl SHProbe {
363    pub fn new(position: Vec3) -> Self {
364        Self {
365            position,
366            sh_r: SH2::default(),
367            sh_g: SH2::default(),
368            sh_b: SH2::default(),
369        }
370    }
371
372    /// Evaluate irradiance in a given direction.
373    pub fn evaluate(&self, direction: Vec3) -> Vec3 {
374        let basis = sh_basis_2(direction);
375        let r: f32 = (0..9).map(|i| self.sh_r.coeffs[i] * basis[i]).sum();
376        let g: f32 = (0..9).map(|i| self.sh_g.coeffs[i] * basis[i]).sum();
377        let b: f32 = (0..9).map(|i| self.sh_b.coeffs[i] * basis[i]).sum();
378        Vec3::new(r.max(0.0), g.max(0.0), b.max(0.0))
379    }
380
381    /// Add a directional light sample.
382    pub fn add_sample(&mut self, direction: Vec3, color: Vec3) {
383        let basis = sh_basis_2(direction);
384        for i in 0..9 {
385            self.sh_r.coeffs[i] += color.x * basis[i];
386            self.sh_g.coeffs[i] += color.y * basis[i];
387            self.sh_b.coeffs[i] += color.z * basis[i];
388        }
389    }
390}
391
392/// Evaluate 3-channel irradiance from 3 sets of SH2 coefficients.
393pub fn sh_to_color_9(sh_r: &SH2, sh_g: &SH2, sh_b: &SH2, dir: Vec3) -> Vec3 {
394    let basis = sh_basis_2(dir);
395    let r: f32 = (0..9).map(|i| sh_r.coeffs[i] * basis[i]).sum();
396    let g: f32 = (0..9).map(|i| sh_g.coeffs[i] * basis[i]).sum();
397    let b: f32 = (0..9).map(|i| sh_b.coeffs[i] * basis[i]).sum();
398    Vec3::new(r.max(0.0), g.max(0.0), b.max(0.0))
399}
400
401impl SH2 {
402    /// Total energy (L2 norm).
403    pub fn energy(&self) -> f32 {
404        self.coeffs.iter().map(|c| c * c).sum()
405    }
406
407    /// Evaluate at a direction.
408    pub fn evaluate(&self, dir: Vec3) -> f32 {
409        sh_evaluate(&self.coeffs, dir)
410    }
411
412    /// Project a direction with a value.
413    pub fn project(&mut self, dir: Vec3, value: f32) {
414        let basis = sh_basis_2(dir);
415        for i in 0..9 {
416            self.coeffs[i] += value * basis[i];
417        }
418    }
419}
420
421impl SH3 {
422    pub fn energy(&self) -> f32 {
423        self.coeffs.iter().map(|c| c * c).sum()
424    }
425
426    pub fn evaluate(&self, dir: Vec3) -> f32 {
427        sh_evaluate(&self.coeffs, dir)
428    }
429
430    pub fn project(&mut self, dir: Vec3, value: f32) {
431        let basis = sh_basis_3(dir);
432        for i in 0..16 {
433            self.coeffs[i] += value * basis[i];
434        }
435    }
436}
437
438#[cfg(test)]
439mod tests {
440    use super::*;
441
442    #[test]
443    fn test_sh_basis_2_normalization() {
444        // The L=0 coefficient should be constant 0.282095
445        let b = sh_basis_2(Vec3::X);
446        assert!((b[0] - 0.282095).abs() < 1e-4);
447
448        let b2 = sh_basis_2(Vec3::Y);
449        assert!((b2[0] - 0.282095).abs() < 1e-4);
450    }
451
452    #[test]
453    fn test_sh_basis_3_length() {
454        let b = sh_basis_3(Vec3::Z);
455        assert_eq!(b.len(), 16);
456    }
457
458    #[test]
459    fn test_sh_orthogonality() {
460        // SH basis functions should be approximately orthogonal
461        // Integrate Y_i * Y_j over the sphere using Monte Carlo
462        let n = 10000;
463        let n_sqrt = (n as f32).sqrt().ceil() as usize;
464        let weight = 4.0 * std::f32::consts::PI / (n_sqrt * n_sqrt) as f32;
465
466        // Test orthogonality between band 0 and band 1
467        let mut dot_01 = 0.0f32;
468        let mut dot_00 = 0.0f32;
469        let mut dot_11 = 0.0f32;
470
471        for i in 0..n_sqrt {
472            for j in 0..n_sqrt {
473                let u = (i as f32 + 0.5) / n_sqrt as f32;
474                let v = (j as f32 + 0.5) / n_sqrt as f32;
475                let theta = (1.0 - 2.0 * u).acos();
476                let phi = 2.0 * std::f32::consts::PI * v;
477                let dir = Vec3::new(
478                    theta.sin() * phi.cos(),
479                    theta.sin() * phi.sin(),
480                    theta.cos(),
481                );
482                let basis = sh_basis_2(dir);
483                dot_00 += basis[0] * basis[0] * weight;
484                dot_01 += basis[0] * basis[1] * weight;
485                dot_11 += basis[1] * basis[1] * weight;
486            }
487        }
488
489        assert!((dot_00 - 1.0).abs() < 0.15, "Y00 self-dot should be ~1, got {dot_00}");
490        assert!(dot_01.abs() < 0.15, "Y00.Y1-1 should be ~0, got {dot_01}");
491        assert!((dot_11 - 1.0).abs() < 0.15, "Y1-1 self-dot should be ~1, got {dot_11}");
492    }
493
494    #[test]
495    fn test_sh_project_then_evaluate() {
496        // Project a function that's strong in the +Y direction
497        let sh = sh_project_function(
498            |dir| dir.y.max(0.0),
499            2500,
500        );
501
502        let val_up = sh.evaluate(Vec3::Y);
503        let val_down = sh.evaluate(-Vec3::Y);
504
505        assert!(val_up > val_down, "Projected function should be stronger in +Y: up={val_up}, down={val_down}");
506        assert!(val_up > 0.0);
507    }
508
509    #[test]
510    fn test_sh_rotation_preserves_energy() {
511        let mut sh = SH2::default();
512        sh.project(Vec3::new(1.0, 1.0, 0.0).normalize(), 1.0);
513        let original_energy = sh.energy();
514
515        let rotation = Mat4::from_rotation_z(std::f32::consts::FRAC_PI_2);
516        let rotated = sh_rotate(&sh, rotation);
517
518        let rotated_energy = rotated.energy();
519        // Energy should be approximately preserved
520        let ratio = rotated_energy / original_energy;
521        assert!(
522            ratio > 0.5 && ratio < 2.0,
523            "Energy should be roughly preserved: original={original_energy}, rotated={rotated_energy}"
524        );
525    }
526
527    #[test]
528    fn test_sh_add_scale() {
529        let mut a = SH2::default();
530        a.coeffs[0] = 1.0;
531        a.coeffs[1] = 2.0;
532
533        let mut b = SH2::default();
534        b.coeffs[0] = 3.0;
535        b.coeffs[1] = 4.0;
536
537        let sum = sh_add(&a, &b);
538        assert!((sum.coeffs[0] - 4.0).abs() < 1e-6);
539        assert!((sum.coeffs[1] - 6.0).abs() < 1e-6);
540
541        let scaled = sh_scale(&a, 2.0);
542        assert!((scaled.coeffs[0] - 2.0).abs() < 1e-6);
543        assert!((scaled.coeffs[1] - 4.0).abs() < 1e-6);
544    }
545
546    #[test]
547    fn test_sh_dot() {
548        let mut a = SH2::default();
549        a.coeffs[0] = 1.0;
550        let mut b = SH2::default();
551        b.coeffs[0] = 2.0;
552        assert!((sh_dot(&a, &b) - 2.0).abs() < 1e-6);
553    }
554
555    #[test]
556    fn test_cosine_lobe() {
557        let cosine = cosine_lobe_sh();
558        let val_up = cosine.evaluate(Vec3::Z);
559        let val_side = cosine.evaluate(Vec3::X);
560        let val_down = cosine.evaluate(-Vec3::Z);
561
562        assert!(val_up > val_side, "Cosine lobe should be strongest at Z");
563        assert!(val_side >= val_down, "Cosine lobe should be weaker below horizon");
564    }
565
566    #[test]
567    fn test_sh_probe() {
568        let mut probe = SHProbe::new(Vec3::ZERO);
569        probe.add_sample(Vec3::Y, Vec3::new(1.0, 0.0, 0.0));
570
571        let color_up = probe.evaluate(Vec3::Y);
572        let color_down = probe.evaluate(-Vec3::Y);
573        assert!(color_up.x > color_down.x, "Probe should be brighter in sample direction");
574    }
575
576    #[test]
577    fn test_sh_convolve() {
578        let mut sh = SH2::default();
579        sh.coeffs[0] = 1.0;
580        sh.coeffs[2] = 0.5;
581
582        let kernel = [
583            std::f32::consts::PI,
584            2.0 * std::f32::consts::PI / 3.0,
585            std::f32::consts::PI / 4.0,
586        ];
587        let convolved = sh_convolve(&sh, &kernel);
588        assert!((convolved.coeffs[0] - std::f32::consts::PI).abs() < 1e-4);
589    }
590
591    #[test]
592    fn test_legendre_p() {
593        // P_0^0(x) = 1
594        assert!((legendre_p(0, 0, 0.5) - 1.0).abs() < 1e-6);
595        // P_1^0(x) = x
596        assert!((legendre_p(1, 0, 0.5) - 0.5).abs() < 1e-6);
597        // P_2^0(x) = (3x^2 - 1)/2
598        let x = 0.5;
599        let expected = (3.0 * x * x - 1.0) / 2.0;
600        assert!((legendre_p(2, 0, x) - expected).abs() < 1e-6);
601    }
602
603    #[test]
604    fn test_sh3_project_evaluate() {
605        let mut sh = SH3::default();
606        sh.project(Vec3::Z, 1.0);
607        let val = sh.evaluate(Vec3::Z);
608        assert!(val > 0.0);
609    }
610
611    #[test]
612    fn test_sh_to_color_9() {
613        let mut r = SH2::default();
614        r.coeffs[0] = 1.0;
615        let g = SH2::default();
616        let b = SH2::default();
617        let color = sh_to_color_9(&r, &g, &b, Vec3::X);
618        assert!(color.x > 0.0);
619        assert!((color.y).abs() < 1e-6);
620    }
621
622    #[test]
623    fn test_factorial() {
624        assert!((factorial(0) - 1.0).abs() < 1e-10);
625        assert!((factorial(5) - 120.0).abs() < 1e-10);
626    }
627}