use super::codec::{dequantize_2d, quantize_2d};
use super::operators::{gradient_x, gradient_y};
use crate::CfdScalar;
use alloc::format;
use alloc::vec;
use alloc::vec::Vec;
use deep_causality_algebra::ConjugateScalar;
use deep_causality_fft::RfftPlanNd;
use deep_causality_num::FromPrimitive;
use deep_causality_num_complex::Complex;
use deep_causality_physics::PhysicsError;
use deep_causality_tensor::{
CausalTensor, CausalTensorTrain, CausalTensorTrainOperator, TensorTrain, TensorTrainOperator,
Truncation,
};
pub struct QttProjector2d<R>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
lx: usize,
ly: usize,
dx: R,
dy: R,
gx: CausalTensorTrainOperator<R>,
gy: CausalTensorTrainOperator<R>,
trunc: Truncation<R>,
}
impl<R> QttProjector2d<R>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
pub fn new(
lx: usize,
ly: usize,
dx: R,
dy: R,
trunc: Truncation<R>,
) -> Result<Self, PhysicsError> {
let gx = gradient_x::<R>(lx, ly, dx, &trunc)?;
let gy = gradient_y::<R>(lx, ly, dy, &trunc)?;
Ok(Self {
lx,
ly,
dx,
dy,
gx,
gy,
trunc,
})
}
pub fn divergence(
&self,
u: &CausalTensorTrain<R>,
v: &CausalTensorTrain<R>,
) -> Result<CausalTensorTrain<R>, PhysicsError> {
let du = self.gx.apply(u, &self.trunc)?;
let dv = self.gy.apply(v, &self.trunc)?;
Ok(du.add(&dv)?.round(&self.trunc)?)
}
pub fn solve_poisson(
&self,
rhs: &CausalTensorTrain<R>,
) -> Result<CausalTensorTrain<R>, PhysicsError> {
let dense = dequantize_2d(rhs, self.lx, self.ly)?;
let nx = 1usize << self.lx;
let ny = 1usize << self.ly;
let p = spectral_poisson::<R>(dense.as_slice(), nx, ny, self.dx, self.dy)?;
let field = CausalTensor::new(p, vec![nx, ny])?;
quantize_2d(&field, &self.trunc)
}
pub fn project(
&self,
u: &CausalTensorTrain<R>,
v: &CausalTensorTrain<R>,
) -> Result<(CausalTensorTrain<R>, CausalTensorTrain<R>), PhysicsError> {
let div = self.divergence(u, v)?;
let p = self.solve_poisson(&div)?;
let neg = R::zero() - R::one();
let un = u
.add(&self.gx.apply(&p, &self.trunc)?.scale(neg))?
.round(&self.trunc)?;
let vn = v
.add(&self.gy.apply(&p, &self.trunc)?.scale(neg))?
.round(&self.trunc)?;
Ok((un, vn))
}
}
fn spectral_poisson<R>(
rhs: &[R],
nx: usize,
ny: usize,
dx: R,
dy: R,
) -> Result<Vec<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
let plan = RfftPlanNd::<R>::new(&[nx, ny])
.map_err(|e| PhysicsError::CalculationError(format!("rfft plan: {e:?}")))?;
let zero = Complex::from_real(R::zero());
let mut spec = vec![zero; plan.spectrum_len()];
let mut scratch = vec![zero; plan.scratch_len()];
plan.execute(rhs, &mut spec, &mut scratch)
.map_err(|e| PhysicsError::CalculationError(format!("rfft forward: {e:?}")))?;
let hy = plan.spectrum_shape()[1]; let two = R::one() + R::one();
let tau = R::pi() * two;
let nxf = from_usize::<R>(nx);
let nyf = from_usize::<R>(ny);
let dx2 = dx * dx;
let dy2 = dy * dy;
let (half_x, half_y) = (nx / 2, ny / 2);
for kx in 0..nx {
let sx = (tau * from_usize::<R>(kx) / nxf).sin();
let lamx = sx * sx / dx2;
for ky in 0..hy {
let sy = (tau * from_usize::<R>(ky) / nyf).sin();
let lamy = sy * sy / dy2;
let idx = kx * hy + ky;
let is_null = (kx == 0 || kx == half_x) && (ky == 0 || ky == half_y);
if is_null {
spec[idx] = zero;
} else {
let inv = R::zero() - R::one() / (lamx + lamy);
spec[idx] = Complex::new(spec[idx].re * inv, spec[idx].im * inv);
}
}
}
let mut out = vec![R::zero(); nx * ny];
plan.execute_inverse(&mut spec, &mut out, &mut scratch)
.map_err(|e| PhysicsError::CalculationError(format!("rfft inverse: {e:?}")))?;
Ok(out)
}
fn from_usize<R: FromPrimitive>(n: usize) -> R {
<R as FromPrimitive>::from_f64(n as f64).unwrap()
}