use super::kfvs::MacroState;
use rayon::prelude::*;
pub const N_CONSERVED: usize = 5;
pub fn moment_preserving_projection(
f: &[f64],
spatial_shape: [usize; 3],
velocity_shape: [usize; 3],
target_moments: &[MacroState],
dv: [f64; 3],
v_min: [f64; 3],
) -> Vec<f64> {
let [nx, ny, nz] = spatial_shape;
let [nv1, nv2, nv3] = velocity_shape;
let n_spatial = nx * ny * nz;
let n_vel = nv1 * nv2 * nv3;
let dv3 = dv[0] * dv[1] * dv[2];
assert_eq!(f.len(), n_spatial * n_vel);
assert_eq!(target_moments.len(), n_spatial);
let mut result = f.to_vec();
let mut psi = vec![[0.0f64; N_CONSERVED]; n_vel]; let mut v_coords = vec![[0.0f64; 3]; n_vel];
for iv1 in 0..nv1 {
for iv2 in 0..nv2 {
for iv3 in 0..nv3 {
let iv = iv1 * nv2 * nv3 + iv2 * nv3 + iv3;
let v1 = v_min[0] + (iv1 as f64 + 0.5) * dv[0];
let v2 = v_min[1] + (iv2 as f64 + 0.5) * dv[1];
let v3 = v_min[2] + (iv3 as f64 + 0.5) * dv[2];
v_coords[iv] = [v1, v2, v3];
psi[iv] = [1.0, v1, v2, v3, 0.5 * (v1 * v1 + v2 * v2 + v3 * v3)];
}
}
}
let mut gram = [[0.0f64; N_CONSERVED]; N_CONSERVED];
for psi_v in psi.iter() {
for i in 0..N_CONSERVED {
for j in 0..N_CONSERVED {
gram[i][j] += psi_v[i] * psi_v[j] * dv3;
}
}
}
let gram_inv = invert_5x5(&gram);
result
.par_chunks_mut(n_vel)
.enumerate()
.for_each(|(ix, cell)| {
let target = &target_moments[ix];
let mut current = [0.0f64; N_CONSERVED];
for iv in 0..n_vel {
let fval = cell[iv];
for m in 0..N_CONSERVED {
current[m] += fval * psi[iv][m] * dv3;
}
}
let target_vec = [
target.density,
target.momentum[0],
target.momentum[1],
target.momentum[2],
target.energy,
];
let mut delta = [0.0f64; N_CONSERVED];
for m in 0..N_CONSERVED {
delta[m] = target_vec[m] - current[m];
}
let mut coeffs = [0.0f64; N_CONSERVED];
for i in 0..N_CONSERVED {
for j in 0..N_CONSERVED {
coeffs[i] += gram_inv[i][j] * delta[j];
}
}
for iv in 0..n_vel {
let mut correction = 0.0;
for m in 0..N_CONSERVED {
correction += coeffs[m] * psi[iv][m];
}
cell[iv] += correction;
}
});
result
}
pub fn conservative_truncation(
f_truncated: &[f64],
spatial_shape: [usize; 3],
velocity_shape: [usize; 3],
target_moments: &[MacroState],
dv: [f64; 3],
v_min: [f64; 3],
) -> Vec<f64> {
moment_preserving_projection(
f_truncated,
spatial_shape,
velocity_shape,
target_moments,
dv,
v_min,
)
}
pub fn extract_moments(
f: &[f64],
spatial_shape: [usize; 3],
velocity_shape: [usize; 3],
dv: [f64; 3],
v_min: [f64; 3],
) -> Vec<MacroState> {
let [nv1, nv2, nv3] = velocity_shape;
let n_vel = nv1 * nv2 * nv3;
let dv3 = dv[0] * dv[1] * dv[2];
let v1_coords: Vec<f64> = (0..nv1)
.map(|i| v_min[0] + (i as f64 + 0.5) * dv[0])
.collect();
let v2_coords: Vec<f64> = (0..nv2)
.map(|i| v_min[1] + (i as f64 + 0.5) * dv[1])
.collect();
let v3_coords: Vec<f64> = (0..nv3)
.map(|i| v_min[2] + (i as f64 + 0.5) * dv[2])
.collect();
f.par_chunks(n_vel)
.map(|cell| {
let mut rho = 0.0;
let mut mom = [0.0f64; 3];
let mut energy = 0.0;
for (iv1, &v1) in v1_coords.iter().enumerate() {
for (iv2, &v2) in v2_coords.iter().enumerate() {
for (iv3, &v3) in v3_coords.iter().enumerate() {
let iv = iv1 * nv2 * nv3 + iv2 * nv3 + iv3;
let fval = cell[iv];
rho += fval * dv3;
mom[0] += fval * v1 * dv3;
mom[1] += fval * v2 * dv3;
mom[2] += fval * v3 * dv3;
energy += 0.5 * fval * (v1 * v1 + v2 * v2 + v3 * v3) * dv3;
}
}
}
MacroState {
density: rho,
momentum: mom,
energy,
}
})
.collect()
}
fn invert_5x5(a: &[[f64; N_CONSERVED]; N_CONSERVED]) -> [[f64; N_CONSERVED]; N_CONSERVED] {
let n = N_CONSERVED;
let mut aug = [[0.0f64; 2 * N_CONSERVED]; N_CONSERVED];
for i in 0..n {
for j in 0..n {
aug[i][j] = a[i][j];
}
aug[i][n + i] = 1.0;
}
for col in 0..n {
let (max_row, _) = aug[col..n]
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| {
a[col]
.abs()
.partial_cmp(&b[col].abs())
.unwrap_or(std::cmp::Ordering::Equal)
})
.map(|(i, row)| (i + col, row[col].abs()))
.unwrap_or((col, 0.0));
aug.swap(col, max_row);
let pivot = aug[col][col];
if pivot.abs() < 1e-30 {
let mut result = [[0.0f64; N_CONSERVED]; N_CONSERVED];
for (i, row) in result.iter_mut().enumerate() {
row[i] = 1.0;
}
return result;
}
for val in aug[col].iter_mut().take(2 * n) {
*val /= pivot;
}
for row in 0..n {
if row == col {
continue;
}
let factor = aug[row][col];
let pivot_row: [f64; 2 * N_CONSERVED] = aug[col];
for j in 0..2 * n {
aug[row][j] -= factor * pivot_row[j];
}
}
}
let mut inv = [[0.0f64; N_CONSERVED]; N_CONSERVED];
for (inv_row, aug_row) in inv.iter_mut().zip(aug.iter()) {
inv_row[..n].copy_from_slice(&aug_row[n..2 * n]);
}
inv
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn conservative_svd_preserves_moments() {
let spatial_shape = [2, 2, 2];
let velocity_shape = [4, 4, 4];
let n_spatial = 8;
let n_vel = 64;
let dv = [1.0; 3];
let v_min = [-2.0; 3];
let mut f = vec![0.0; n_spatial * n_vel];
for ix in 0..n_spatial {
for iv1 in 0..4 {
for iv2 in 0..4 {
for iv3 in 0..4 {
let iv = iv1 * 16 + iv2 * 4 + iv3;
let v1 = v_min[0] + (iv1 as f64 + 0.5) * dv[0];
let v2 = v_min[1] + (iv2 as f64 + 0.5) * dv[1];
let v3 = v_min[2] + (iv3 as f64 + 0.5) * dv[2];
let v2_total = v1 * v1 + v2 * v2 + v3 * v3;
f[ix * n_vel + iv] = (-v2_total / 2.0).exp();
}
}
}
}
let original_moments = extract_moments(&f, spatial_shape, velocity_shape, dv, v_min);
let mut f_perturbed = f.clone();
for i in 0..f_perturbed.len() {
if i % 3 == 0 {
f_perturbed[i] *= 0.5;
}
}
let perturbed_moments =
extract_moments(&f_perturbed, spatial_shape, velocity_shape, dv, v_min);
let mass_diff = (perturbed_moments[0].density - original_moments[0].density).abs();
assert!(
mass_diff > 1e-10,
"Perturbation should change moments: diff={mass_diff}"
);
let f_corrected = conservative_truncation(
&f_perturbed,
spatial_shape,
velocity_shape,
&original_moments,
dv,
v_min,
);
let corrected_moments =
extract_moments(&f_corrected, spatial_shape, velocity_shape, dv, v_min);
for ix in 0..n_spatial {
let orig = &original_moments[ix];
let corr = &corrected_moments[ix];
assert!(
(corr.density - orig.density).abs() < 1e-12,
"Cell {ix}: density {:.6e} vs {:.6e}",
corr.density,
orig.density
);
for d in 0..3 {
assert!(
(corr.momentum[d] - orig.momentum[d]).abs() < 1e-12,
"Cell {ix}: momentum[{d}] {:.6e} vs {:.6e}",
corr.momentum[d],
orig.momentum[d]
);
}
assert!(
(corr.energy - orig.energy).abs() < 1e-11,
"Cell {ix}: energy {:.6e} vs {:.6e}",
corr.energy,
orig.energy
);
}
}
#[test]
fn extract_moments_maxwellian() {
let spatial_shape = [1, 1, 1];
let velocity_shape = [8, 8, 8];
let dv = [0.5; 3];
let v_min = [-2.0; 3];
let n_vel = 512;
let mut f = vec![0.0; n_vel];
for iv1 in 0..8 {
for iv2 in 0..8 {
for iv3 in 0..8 {
let iv = iv1 * 64 + iv2 * 8 + iv3;
let v1 = v_min[0] + (iv1 as f64 + 0.5) * dv[0];
let v2 = v_min[1] + (iv2 as f64 + 0.5) * dv[1];
let v3 = v_min[2] + (iv3 as f64 + 0.5) * dv[2];
f[iv] = (-0.5 * (v1 * v1 + v2 * v2 + v3 * v3)).exp();
}
}
}
let moments = extract_moments(&f, spatial_shape, velocity_shape, dv, v_min);
assert_eq!(moments.len(), 1);
assert!(moments[0].density > 0.0);
for d in 0..3 {
assert!(
moments[0].momentum[d].abs() < 1e-14,
"Momentum[{d}] = {}, expected ~0",
moments[0].momentum[d]
);
}
}
#[test]
fn gram_inverse_identity() {
let id = [
[1.0, 0.0, 0.0, 0.0, 0.0],
[0.0, 1.0, 0.0, 0.0, 0.0],
[0.0, 0.0, 1.0, 0.0, 0.0],
[0.0, 0.0, 0.0, 1.0, 0.0],
[0.0, 0.0, 0.0, 0.0, 1.0],
];
let inv = invert_5x5(&id);
for i in 0..5 {
for j in 0..5 {
let expected = if i == j { 1.0 } else { 0.0 };
assert!(
(inv[i][j] - expected).abs() < 1e-14,
"inv[{i}][{j}] = {}, expected {expected}",
inv[i][j]
);
}
}
}
}