use crate::CfdScalar;
use alloc::format;
use alloc::vec;
use alloc::vec::Vec;
use deep_causality_algebra::ConjugateScalar;
use deep_causality_physics::PhysicsError;
use deep_causality_tensor::{
CausalTensor, CausalTensorTrain, CausalTensorTrainOperator, TensorTrain, TensorTrainOperator,
Truncation,
};
fn build_core<R, F>(rl: usize, rr: usize, fill: F) -> CausalTensor<R>
where
R: ConjugateScalar,
F: Fn(usize, usize, usize, usize) -> bool,
{
let mut data = vec![R::zero(); rl * 2 * 2 * rr];
for cl in 0..rl {
for o in 0..2 {
for i in 0..2 {
for cr in 0..rr {
if fill(cl, o, i, cr) {
data[((cl * 2 + o) * 2 + i) * rr + cr] = R::one();
}
}
}
}
}
CausalTensor::new(data, vec![rl, 2, 2, rr]).unwrap()
}
pub fn shift_plus<R>(l: usize) -> Result<CausalTensorTrainOperator<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
if l == 0 {
return Err(PhysicsError::DimensionMismatch(format!(
"shift operator requires l >= 1, got {l}"
)));
}
if l == 1 {
let c = build_core::<R, _>(1, 1, |_cl, o, i, _cr| o == (i ^ 1));
return Ok(CausalTensorTrainOperator::from_cores(vec![c])?);
}
let mut cores = Vec::with_capacity(l);
for k in 0..l {
let c = if k == 0 {
build_core::<R, _>(1, 2, |_cl, o, i, cr| o == (i ^ cr))
} else if k == l - 1 {
build_core::<R, _>(2, 1, |cl, o, i, _cr| o == (i ^ 1) && cl == i)
} else {
build_core::<R, _>(2, 2, |cl, o, i, cr| o == (i ^ cr) && cl == (i & cr))
};
cores.push(c);
}
Ok(CausalTensorTrainOperator::from_cores(cores)?)
}
pub fn shift_minus<R>(l: usize) -> Result<CausalTensorTrainOperator<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
Ok(shift_plus::<R>(l)?.transpose())
}
pub fn gradient<R>(
l: usize,
dx: R,
trunc: &Truncation<R>,
) -> Result<CausalTensorTrainOperator<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
let two = R::one() + R::one();
let half_inv_dx = R::one() / (two * dx);
let g = shift_minus::<R>(l)?
.sub(&shift_plus::<R>(l)?)?
.scale(half_inv_dx);
Ok(g.round(trunc)?)
}
pub fn laplacian<R>(
l: usize,
dx: R,
trunc: &Truncation<R>,
) -> Result<CausalTensorTrainOperator<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
let two = R::one() + R::one();
let inv_dx2 = R::one() / (dx * dx);
let id = CausalTensorTrainOperator::<R>::identity(&vec![2usize; l]);
let lap = shift_plus::<R>(l)?
.add(&shift_minus::<R>(l)?)?
.sub(&id.scale(two))?
.scale(inv_dx2);
Ok(lap.round(trunc)?)
}
fn identity_core<R: ConjugateScalar>() -> CausalTensor<R> {
build_core::<R, _>(1, 1, |_cl, o, i, _cr| o == i)
}
pub(crate) fn lift_leading<R>(
op: &CausalTensorTrainOperator<R>,
m: usize,
) -> Result<CausalTensorTrainOperator<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
let mut cores = op.cores().to_vec();
cores.extend((0..m).map(|_| identity_core::<R>()));
Ok(CausalTensorTrainOperator::from_cores(cores)?)
}
pub(crate) fn lift_trailing<R>(
op: &CausalTensorTrainOperator<R>,
m: usize,
) -> Result<CausalTensorTrainOperator<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
let mut cores: Vec<CausalTensor<R>> = (0..m).map(|_| identity_core::<R>()).collect();
cores.extend(op.cores().to_vec());
Ok(CausalTensorTrainOperator::from_cores(cores)?)
}
pub fn gradient_x<R>(
lx: usize,
ly: usize,
dx: R,
trunc: &Truncation<R>,
) -> Result<CausalTensorTrainOperator<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
lift_leading(&gradient::<R>(lx, dx, trunc)?, ly)
}
pub fn gradient_y<R>(
lx: usize,
ly: usize,
dy: R,
trunc: &Truncation<R>,
) -> Result<CausalTensorTrainOperator<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
lift_trailing(&gradient::<R>(ly, dy, trunc)?, lx)
}
pub fn laplacian_2d<R>(
lx: usize,
ly: usize,
dx: R,
dy: R,
trunc: &Truncation<R>,
) -> Result<CausalTensorTrainOperator<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
let lap_x = lift_leading(&laplacian::<R>(lx, dx, trunc)?, ly)?;
let lap_y = lift_trailing(&laplacian::<R>(ly, dy, trunc)?, lx)?;
Ok(lap_x.add(&lap_y)?.round(trunc)?)
}
pub(crate) fn lift_block<R>(
op: &CausalTensorTrainOperator<R>,
lead: usize,
trail: usize,
) -> Result<CausalTensorTrainOperator<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
let mut cores: Vec<CausalTensor<R>> = (0..lead).map(|_| identity_core::<R>()).collect();
cores.extend(op.cores().to_vec());
cores.extend((0..trail).map(|_| identity_core::<R>()));
Ok(CausalTensorTrainOperator::from_cores(cores)?)
}
pub fn gradient_x_3d<R>(
lx: usize,
ly: usize,
lz: usize,
dx: R,
trunc: &Truncation<R>,
) -> Result<CausalTensorTrainOperator<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
lift_block(&gradient::<R>(lx, dx, trunc)?, 0, ly + lz)
}
pub fn gradient_y_3d<R>(
lx: usize,
ly: usize,
lz: usize,
dy: R,
trunc: &Truncation<R>,
) -> Result<CausalTensorTrainOperator<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
lift_block(&gradient::<R>(ly, dy, trunc)?, lx, lz)
}
pub fn gradient_z_3d<R>(
lx: usize,
ly: usize,
lz: usize,
dz: R,
trunc: &Truncation<R>,
) -> Result<CausalTensorTrainOperator<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
lift_block(&gradient::<R>(lz, dz, trunc)?, lx + ly, 0)
}
pub fn laplacian_3d<R>(
lx: usize,
ly: usize,
lz: usize,
dx: R,
dy: R,
dz: R,
trunc: &Truncation<R>,
) -> Result<CausalTensorTrainOperator<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
let lap_x = lift_block(&laplacian::<R>(lx, dx, trunc)?, 0, ly + lz)?;
let lap_y = lift_block(&laplacian::<R>(ly, dy, trunc)?, lx, lz)?;
let lap_z = lift_block(&laplacian::<R>(lz, dz, trunc)?, lx + ly, 0)?;
Ok(lap_x.add(&lap_y)?.add(&lap_z)?.round(trunc)?)
}
pub fn divergence_3d<R>(
fx: &CausalTensorTrain<R>,
fy: &CausalTensorTrain<R>,
fz: &CausalTensorTrain<R>,
grad_x: &CausalTensorTrainOperator<R>,
grad_y: &CausalTensorTrainOperator<R>,
grad_z: &CausalTensorTrainOperator<R>,
trunc: &Truncation<R>,
) -> Result<CausalTensorTrain<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
let dfx = grad_x.apply(fx, trunc)?;
let dfy = grad_y.apply(fy, trunc)?;
let dfz = grad_z.apply(fz, trunc)?;
Ok(dfx.add(&dfy)?.add(&dfz)?.round(trunc)?)
}