1use glam::{Vec3, Mat4};
2
3#[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#[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
27fn factorial(n: u32) -> f64 {
29 (1..=n as u64).fold(1.0f64, |acc, x| acc * x as f64)
30}
31
32fn 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
44pub 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 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 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 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
97fn 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
105fn 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
118fn 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
126pub 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, 0.488603 * y, 0.488603 * z, 0.488603 * x, 1.092548 * x * y, 1.092548 * y * z, 0.315392 * (3.0 * z * z - 1.0), 1.092548 * x * z, 0.546274 * (x * x - y * y), ]
144}
145
146pub 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 result[0] = 0.282095;
157
158 result[1] = 0.488603 * y;
160 result[2] = 0.488603 * z;
161 result[3] = 0.488603 * x;
162
163 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 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
182pub 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
196pub 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 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 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
237pub fn sh_convolve(a: &SH2, kernel: &[f32]) -> SH2 {
239 let mut result = SH2::default();
240 if kernel.len() > 0 {
242 result.coeffs[0] = a.coeffs[0] * kernel[0];
243 }
244 if kernel.len() > 1 {
246 for i in 1..4 {
247 result.coeffs[i] = a.coeffs[i] * kernel[1];
248 }
249 }
250 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
259pub fn sh_rotate(coeffs: &SH2, rotation: Mat4) -> SH2 {
261 let mut result = SH2::default();
262
263 result.coeffs[0] = coeffs.coeffs[0];
265
266 let r = rotation;
269 let sh1 = [coeffs.coeffs[3], coeffs.coeffs[1], coeffs.coeffs[2]]; 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; result.coeffs[1] = rotated_y; result.coeffs[2] = rotated_z; 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 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
320pub 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
329pub 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
338pub fn sh_dot(a: &SH2, b: &SH2) -> f32 {
340 (0..9).map(|i| a.coeffs[i] * b.coeffs[i]).sum()
341}
342
343pub fn cosine_lobe_sh() -> SH2 {
345 let mut sh = SH2::default();
347 sh.coeffs[0] = 0.886227; sh.coeffs[2] = 1.023326; sh.coeffs[6] = 0.495415; sh
351}
352
353#[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 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 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
392pub 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 pub fn energy(&self) -> f32 {
404 self.coeffs.iter().map(|c| c * c).sum()
405 }
406
407 pub fn evaluate(&self, dir: Vec3) -> f32 {
409 sh_evaluate(&self.coeffs, dir)
410 }
411
412 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 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 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 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 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 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 assert!((legendre_p(0, 0, 0.5) - 1.0).abs() < 1e-6);
595 assert!((legendre_p(1, 0, 0.5) - 0.5).abs() < 1e-6);
597 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}