use rayon::prelude::*;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use super::super::{solver::PoissonSolver, types::*};
use super::utils::finite_difference_acceleration;
use std::f64::consts::PI;
pub struct SphericalHarmonicsPoisson {
pub l_max: usize,
pub n_radial: usize,
pub r_max: f64,
pub shape: [usize; 3],
pub dx: [f64; 3],
progress: Option<Arc<super::super::progress::StepProgress>>,
}
impl SphericalHarmonicsPoisson {
pub fn new(l_max: usize, n_radial: usize, shape: [usize; 3], dx: [f64; 3]) -> Self {
let half_extents = [
shape[0] as f64 * dx[0] / 2.0,
shape[1] as f64 * dx[1] / 2.0,
shape[2] as f64 * dx[2] / 2.0,
];
let max_half = half_extents.iter().cloned().fold(0.0_f64, f64::max);
let r_max = 3.0_f64.sqrt() * max_half;
Self {
l_max,
n_radial,
r_max,
shape,
dx,
progress: None,
}
}
#[inline]
fn n_harmonics(&self) -> usize {
(self.l_max + 1) * (self.l_max + 1)
}
#[inline]
fn lm_index(l: usize, m: i32) -> usize {
((l * l) as i32 + l as i32 + m) as usize
}
}
#[inline]
fn associated_legendre(l: usize, m: usize, x: f64) -> f64 {
debug_assert!(m <= l, "m must be <= l");
let mut pmm = 1.0_f64;
if m > 0 {
let somx2 = (1.0 - x * x).max(0.0).sqrt();
let mut fact = 1.0;
for i in 1..=m {
pmm *= -fact * somx2;
fact += 2.0;
let _ = i; }
}
if l == m {
return pmm;
}
let mut pmmp1 = x * (2 * m + 1) as f64 * pmm;
if l == m + 1 {
return pmmp1;
}
let mut pll = 0.0;
for ll in (m + 2)..=l {
pll = ((2 * ll - 1) as f64 * x * pmmp1 - (ll + m - 1) as f64 * pmm) / (ll - m) as f64;
pmm = pmmp1;
pmmp1 = pll;
}
pll
}
#[inline]
fn normalization(l: usize, m_abs: usize) -> f64 {
let mut ratio = 1.0_f64;
for k in (l - m_abs + 1)..=(l + m_abs) {
ratio *= k as f64;
}
((2 * l + 1) as f64 / (4.0 * PI * ratio)).sqrt()
}
#[inline]
fn real_spherical_harmonic(l: usize, m: i32, theta: f64, phi: f64) -> f64 {
let m_abs = m.unsigned_abs() as usize;
let n_lm = normalization(l, m_abs);
let plm = associated_legendre(l, m_abs, theta.cos());
if m > 0 {
2.0_f64.sqrt() * n_lm * plm * (m_abs as f64 * phi).cos()
} else if m < 0 {
2.0_f64.sqrt() * n_lm * plm * (m_abs as f64 * phi).sin()
} else {
n_lm * plm
}
}
#[inline]
fn cell_to_xyz(
ix: usize,
iy: usize,
iz: usize,
shape: &[usize; 3],
dx: &[f64; 3],
) -> (f64, f64, f64) {
let x = (ix as f64 - shape[0] as f64 / 2.0 + 0.5) * dx[0];
let y = (iy as f64 - shape[1] as f64 / 2.0 + 0.5) * dx[1];
let z = (iz as f64 - shape[2] as f64 / 2.0 + 0.5) * dx[2];
(x, y, z)
}
#[inline]
fn xyz_to_spherical(x: f64, y: f64, z: f64) -> (f64, f64, f64) {
let r = (x * x + y * y + z * z).sqrt();
if r < 1e-30 {
return (0.0, 0.0, 0.0);
}
let theta = (z / r).clamp(-1.0, 1.0).acos();
let phi = y.atan2(x); (r, theta, phi)
}
fn decompose_density(
density: &DensityField,
shape: &[usize; 3],
dx: &[f64; 3],
l_max: usize,
n_radial: usize,
r_max: f64,
) -> Vec<Vec<f64>> {
let n_harm = (l_max + 1) * (l_max + 1);
let dr = r_max / n_radial as f64;
let cell_vol = dx[0] * dx[1] * dx[2];
let [nx, ny, nz] = *shape;
let (mut rho_lm, _shell_vol) = (0..nx)
.into_par_iter()
.fold(
|| (vec![vec![0.0; n_radial]; n_harm], vec![0.0; n_radial]),
|(mut rho_lm_local, mut shell_vol_local), ix| {
for iy in 0..ny {
for iz in 0..nz {
let (x, y, z) = cell_to_xyz(ix, iy, iz, shape, dx);
let (r, theta, phi) = xyz_to_spherical(x, y, z);
if r >= r_max {
continue;
}
let r_idx = ((r / dr) as usize).min(n_radial - 1);
let cell_idx = ix * ny * nz + iy * nz + iz;
let rho_val = density.data[cell_idx];
shell_vol_local[r_idx] += cell_vol;
for l in 0..=l_max {
for m in -(l as i32)..=(l as i32) {
let ylm = real_spherical_harmonic(l, m, theta, phi);
let h_idx = SphericalHarmonicsPoisson::lm_index(l, m);
rho_lm_local[h_idx][r_idx] += rho_val * ylm * cell_vol;
}
}
}
}
(rho_lm_local, shell_vol_local)
},
)
.reduce(
|| (vec![vec![0.0; n_radial]; n_harm], vec![0.0; n_radial]),
|(mut a_rho, mut a_sv), (b_rho, b_sv)| {
for h in 0..n_harm {
for r in 0..n_radial {
a_rho[h][r] += b_rho[h][r];
}
}
for r in 0..n_radial {
a_sv[r] += b_sv[r];
}
(a_rho, a_sv)
},
);
for r_idx in 0..n_radial {
let r_center = (r_idx as f64 + 0.5) * dr;
let r2_dr = r_center * r_center * dr;
if r2_dr > 1e-30 {
for harm in rho_lm.iter_mut() {
harm[r_idx] /= r2_dr;
}
}
}
rho_lm
}
fn radial_poisson_solve(rho_lm: &[f64], l: usize, n_radial: usize, r_max: f64, g: f64) -> Vec<f64> {
let dr = r_max / n_radial as f64;
let mut phi_lm = vec![0.0f64; n_radial];
let r_centers: Vec<f64> = (0..n_radial).map(|i| (i as f64 + 0.5) * dr).collect();
let inner_integrand: Vec<f64> = (0..n_radial)
.map(|i| rho_lm[i] * r_centers[i].powi(l as i32 + 2) * dr)
.collect();
let outer_integrand: Vec<f64> = (0..n_radial)
.map(|i| {
let s = r_centers[i];
if s < 1e-30 && l > 1 {
0.0
} else {
rho_lm[i] * s.powi(1 - l as i32) * dr
}
})
.collect();
let mut i_inner = vec![0.0f64; n_radial];
i_inner[0] = inner_integrand[0];
for i in 1..n_radial {
i_inner[i] = i_inner[i - 1] + inner_integrand[i];
}
let mut i_outer = vec![0.0f64; n_radial];
for i in (0..n_radial - 1).rev() {
i_outer[i] = i_outer[i + 1] + outer_integrand[i + 1];
}
let prefactor = -4.0 * PI * g / (2 * l + 1) as f64;
for i in 0..n_radial {
let r = r_centers[i];
if r < 1e-30 {
if l == 0 {
phi_lm[i] = prefactor * i_outer[i];
}
continue;
}
let r_neg_lp1 = r.powi(-(l as i32 + 1));
let r_pos_l = r.powi(l as i32);
phi_lm[i] = prefactor * (r_neg_lp1 * i_inner[i] + r_pos_l * i_outer[i]);
}
phi_lm
}
fn reconstruct_potential(
phi_lm: &[Vec<f64>],
l_max: usize,
shape: &[usize; 3],
dx: &[f64; 3],
n_radial: usize,
r_max: f64,
progress: &Option<Arc<super::super::progress::StepProgress>>,
) -> PotentialField {
let [nx, ny, nz] = *shape;
let n_total = nx * ny * nz;
let dr = r_max / n_radial as f64;
let mut pot_data = vec![0.0f64; n_total];
let total = n_total as u64;
let counter = AtomicU64::new(0);
let report_interval = (total / 100).max(1);
pot_data.par_iter_mut().enumerate().for_each(|(flat, val)| {
let ix = flat / (ny * nz);
let iy = (flat / nz) % ny;
let iz = flat % nz;
let (x, y, z) = cell_to_xyz(ix, iy, iz, shape, dx);
let (r, theta, phi) = xyz_to_spherical(x, y, z);
let r_frac = r / dr - 0.5;
let r_idx_lo = r_frac.floor().max(0.0) as usize;
let r_idx_hi = (r_idx_lo + 1).min(n_radial - 1);
let t = (r_frac - r_idx_lo as f64).clamp(0.0, 1.0);
let mut sum = 0.0;
for l in 0..=l_max {
for m in -(l as i32)..=(l as i32) {
let h_idx = SphericalHarmonicsPoisson::lm_index(l, m);
let ylm = real_spherical_harmonic(l, m, theta, phi);
let phi_r = (1.0 - t) * phi_lm[h_idx][r_idx_lo] + t * phi_lm[h_idx][r_idx_hi];
sum += phi_r * ylm;
}
}
*val = sum;
if let Some(p) = progress {
let c = counter.fetch_add(1, Ordering::Relaxed);
if c.is_multiple_of(report_interval) {
p.set_intra_progress(c, total);
}
}
});
PotentialField {
data: pot_data,
shape: *shape,
}
}
impl PoissonSolver for SphericalHarmonicsPoisson {
fn set_progress(&mut self, p: std::sync::Arc<super::super::progress::StepProgress>) {
self.progress = Some(p);
}
fn solve(&self, density: &DensityField, g: f64) -> PotentialField {
let _span = tracing::info_span!("spherical_harmonics_solve").entered();
let rho_lm = decompose_density(
density,
&self.shape,
&self.dx,
self.l_max,
self.n_radial,
self.r_max,
);
let n_harm = self.n_harmonics();
let mut phi_lm: Vec<Vec<f64>> = Vec::with_capacity(n_harm);
for l in 0..=self.l_max {
for m in -(l as i32)..=(l as i32) {
let h_idx = Self::lm_index(l, m);
let radial_pot =
radial_poisson_solve(&rho_lm[h_idx], l, self.n_radial, self.r_max, g);
debug_assert_eq!(phi_lm.len(), h_idx);
phi_lm.push(radial_pot);
}
}
reconstruct_potential(
&phi_lm,
self.l_max,
&self.shape,
&self.dx,
self.n_radial,
self.r_max,
&self.progress,
)
}
fn compute_acceleration(&self, potential: &PotentialField) -> AccelerationField {
finite_difference_acceleration(potential, &self.dx)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn spherical_point_mass() {
let n = 16;
let dx = [0.5; 3]; let shape = [n; 3];
let mut rho = vec![0.0; n * n * n];
let mid = n / 2;
let cell_vol = dx[0] * dx[1] * dx[2];
let m_total = 1.0;
rho[mid * n * n + mid * n + mid] = m_total / cell_vol;
let solver = SphericalHarmonicsPoisson::new(4, 32, shape, dx);
let density = DensityField { data: rho, shape };
let pot = solver.solve(&density, 1.0);
let test_ix = mid + 3;
let r = (test_ix as f64 - mid as f64 + 0.5) * dx[0];
let expected = -m_total / r; let idx = test_ix * n * n + mid * n + mid;
let err = (pot.data[idx] - expected).abs() / expected.abs();
assert!(
err < 0.5,
"Point mass error {err} at r={r} (expected {expected}, got {})",
pot.data[idx]
);
}
#[test]
fn spherical_plummer() {
let n = 16;
let dx = [0.5; 3];
let shape = [n; 3];
let g_val = 1.0;
let m_total = 1.0;
let a: f64 = 1.0;
let mut rho = vec![0.0; n * n * n];
for ix in 0..n {
for iy in 0..n {
for iz in 0..n {
let x = (ix as f64 - n as f64 / 2.0 + 0.5) * dx[0];
let y = (iy as f64 - n as f64 / 2.0 + 0.5) * dx[1];
let z = (iz as f64 - n as f64 / 2.0 + 0.5) * dx[2];
let r2 = x * x + y * y + z * z;
rho[ix * n * n + iy * n + iz] =
3.0 * m_total / (4.0 * PI * a.powi(3)) * (1.0 + r2 / (a * a)).powf(-2.5);
}
}
}
let solver = SphericalHarmonicsPoisson::new(0, 32, shape, dx);
let density = DensityField { data: rho, shape };
let pot = solver.solve(&density, g_val);
let mid = n / 2;
let test_ix = mid + 2;
let r = (test_ix as f64 - mid as f64 + 0.5) * dx[0];
let expected = -g_val * m_total / (r * r + a * a).sqrt();
let idx = test_ix * n * n + mid * n + mid;
let err = (pot.data[idx] - expected).abs() / expected.abs();
assert!(err < 1.0, "Plummer potential error {err} at r={r}");
}
#[test]
fn spherical_vs_fft_isolated() {
let n = 8;
let dx = [1.0; 3];
let shape = [n; 3];
let mut rho = vec![0.0; n * n * n];
for ix in 0..n {
for iy in 0..n {
for iz in 0..n {
let x = (ix as f64 - n as f64 / 2.0 + 0.5) * dx[0];
let y = (iy as f64 - n as f64 / 2.0 + 0.5) * dx[1];
let z = (iz as f64 - n as f64 / 2.0 + 0.5) * dx[2];
let r2 = x * x + y * y + z * z;
rho[ix * n * n + iy * n + iz] = (-r2 / 4.0).exp();
}
}
}
let solver = SphericalHarmonicsPoisson::new(2, 16, shape, dx);
let density = DensityField { data: rho, shape };
let pot = solver.solve(&density, 1.0);
assert!(
pot.data.iter().all(|x| x.is_finite()),
"Potential must be finite"
);
assert!(
pot.data.iter().any(|x| *x != 0.0),
"Potential must be non-zero"
);
}
#[test]
fn spherical_harmonics_y00() {
let y00 = real_spherical_harmonic(0, 0, 0.5, 1.0);
let expected = 1.0 / (4.0 * PI).sqrt();
assert!(
(y00 - expected).abs() < 1e-12,
"Y_0^0 = {y00}, expected {expected}"
);
}
#[test]
fn spherical_harmonics_orthogonality() {
let n_theta = 50;
let n_phi = 100;
let d_theta = PI / n_theta as f64;
let d_phi = 2.0 * PI / n_phi as f64;
let l_max = 2;
for l1 in 0..=l_max {
for m1 in -(l1 as i32)..=(l1 as i32) {
for l2 in 0..=l_max {
for m2 in -(l2 as i32)..=(l2 as i32) {
let mut integral = 0.0;
for it in 0..n_theta {
let theta = (it as f64 + 0.5) * d_theta;
let sin_theta = theta.sin();
for ip in 0..n_phi {
let phi = (ip as f64 + 0.5) * d_phi;
let y1 = real_spherical_harmonic(l1, m1, theta, phi);
let y2 = real_spherical_harmonic(l2, m2, theta, phi);
integral += y1 * y2 * sin_theta * d_theta * d_phi;
}
}
let expected = if l1 == l2 && m1 == m2 { 1.0 } else { 0.0 };
assert!(
(integral - expected).abs() < 0.02,
"Orthogonality failed for ({l1},{m1}),({l2},{m2}): got {integral}, expected {expected}"
);
}
}
}
}
}
#[test]
fn spherical_acceleration_finite() {
let n = 8;
let dx = [1.0; 3];
let shape = [n; 3];
let mut rho = vec![0.0; n * n * n];
let mid = n / 2;
let cell_vol = dx[0] * dx[1] * dx[2];
rho[mid * n * n + mid * n + mid] = 1.0 / cell_vol;
let solver = SphericalHarmonicsPoisson::new(0, 16, shape, dx);
let density = DensityField { data: rho, shape };
let pot = solver.solve(&density, 1.0);
let acc = solver.compute_acceleration(&pot);
assert!(acc.gx.iter().all(|x| x.is_finite()));
assert!(acc.gy.iter().all(|x| x.is_finite()));
assert!(acc.gz.iter().all(|x| x.is_finite()));
}
#[test]
fn spherical_monopole_symmetry() {
let n = 12;
let dx = [0.5; 3];
let shape = [n; 3];
let mut rho = vec![0.0; n * n * n];
let r_sphere = 2.0;
for ix in 0..n {
for iy in 0..n {
for iz in 0..n {
let x = (ix as f64 - n as f64 / 2.0 + 0.5) * dx[0];
let y = (iy as f64 - n as f64 / 2.0 + 0.5) * dx[1];
let z = (iz as f64 - n as f64 / 2.0 + 0.5) * dx[2];
let r = (x * x + y * y + z * z).sqrt();
if r < r_sphere {
rho[ix * n * n + iy * n + iz] = 1.0;
}
}
}
}
let solver = SphericalHarmonicsPoisson::new(0, 24, shape, dx);
let density = DensityField { data: rho, shape };
let pot = solver.solve(&density, 1.0);
let mid = n / 2;
let d = 2;
let pot_x = pot.data[(mid + d) * n * n + mid * n + mid];
let pot_y = pot.data[mid * n * n + (mid + d) * n + mid];
let pot_z = pot.data[mid * n * n + mid * n + (mid + d)];
let mean = (pot_x + pot_y + pot_z) / 3.0;
let max_dev = [
(pot_x - mean).abs(),
(pot_y - mean).abs(),
(pot_z - mean).abs(),
]
.iter()
.cloned()
.fold(0.0_f64, f64::max);
let rel_dev = max_dev / mean.abs().max(1e-30);
assert!(
rel_dev < 0.15,
"Symmetry broken: pot_x={pot_x}, pot_y={pot_y}, pot_z={pot_z}, rel_dev={rel_dev}"
);
}
}