use crate::CfdScalar;
use alloc::format;
use alloc::vec;
use deep_causality_algebra::ConjugateScalar;
use deep_causality_physics::PhysicsError;
use deep_causality_tensor::{CausalTensor, CausalTensorTrain, Tensor, TensorTrain, Truncation};
pub fn quantize<R>(
field: &CausalTensor<R>,
trunc: &Truncation<R>,
) -> Result<CausalTensorTrain<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
let n = field.as_slice().len();
if n == 0 || !n.is_power_of_two() {
return Err(PhysicsError::DimensionMismatch(format!(
"quantize requires a power-of-two field length, got {n}"
)));
}
let l = n.trailing_zeros() as usize;
let modes = vec![2usize; l];
let reshaped = field.reshape(&modes)?;
Ok(CausalTensorTrain::from_dense(&reshaped, trunc)?)
}
pub fn dequantize<R>(train: &CausalTensorTrain<R>) -> Result<CausalTensor<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
let dense = train.to_dense()?;
let n: usize = dense.shape().iter().product();
Ok(dense.reshape(&[n])?)
}
pub fn quantize_2d<R>(
field: &CausalTensor<R>,
trunc: &Truncation<R>,
) -> Result<CausalTensorTrain<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
let shape = field.shape();
if shape.len() != 2 {
return Err(PhysicsError::DimensionMismatch(format!(
"quantize_2d requires a 2-D field, got {} dims",
shape.len()
)));
}
let (nx, ny) = (shape[0], shape[1]);
if nx == 0 || ny == 0 || !nx.is_power_of_two() || !ny.is_power_of_two() {
return Err(PhysicsError::DimensionMismatch(format!(
"quantize_2d requires power-of-two extents, got {nx} x {ny}"
)));
}
let modes = vec![2usize; nx.trailing_zeros() as usize + ny.trailing_zeros() as usize];
let reshaped = field.reshape(&modes)?;
Ok(CausalTensorTrain::from_dense(&reshaped, trunc)?)
}
pub fn dequantize_2d<R>(
train: &CausalTensorTrain<R>,
lx: usize,
ly: usize,
) -> Result<CausalTensor<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
let dense = train.to_dense()?;
Ok(dense.reshape(&[1usize << lx, 1usize << ly])?)
}
pub fn quantize_3d<R>(
field: &CausalTensor<R>,
trunc: &Truncation<R>,
) -> Result<CausalTensorTrain<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
let shape = field.shape();
if shape.len() != 3 {
return Err(PhysicsError::DimensionMismatch(format!(
"quantize_3d requires a 3-D field, got {} dims",
shape.len()
)));
}
let (nx, ny, nz) = (shape[0], shape[1], shape[2]);
if nx == 0
|| ny == 0
|| nz == 0
|| !nx.is_power_of_two()
|| !ny.is_power_of_two()
|| !nz.is_power_of_two()
{
return Err(PhysicsError::DimensionMismatch(format!(
"quantize_3d requires power-of-two extents, got {nx} x {ny} x {nz}"
)));
}
let l =
nx.trailing_zeros() as usize + ny.trailing_zeros() as usize + nz.trailing_zeros() as usize;
let modes = vec![2usize; l];
let reshaped = field.reshape(&modes)?;
Ok(CausalTensorTrain::from_dense(&reshaped, trunc)?)
}
pub fn dequantize_3d<R>(
train: &CausalTensorTrain<R>,
lx: usize,
ly: usize,
lz: usize,
) -> Result<CausalTensor<R>, PhysicsError>
where
R: CfdScalar + ConjugateScalar<Real = R>,
{
let dense = train.to_dense()?;
Ok(dense.reshape(&[1usize << lx, 1usize << ly, 1usize << lz])?)
}