use crate::error::{Result, TbError};
use ndarray::{Array1, Array2};
#[inline(always)]
pub fn gen_kplane(
origin: &Array1<f64>,
vec1: &Array1<f64>,
vec2: &Array1<f64>,
n1: usize,
n2: usize,
) -> Result<Array2<f64>> {
let dim = origin.len();
if vec1.len() != dim || vec2.len() != dim {
return Err(TbError::KVectorLengthMismatch {
expected: dim,
actual: vec1.len().max(vec2.len()),
});
}
let nk = n1 * n2;
let mut kvec = Array2::<f64>::zeros((nk, dim));
let n1_f = n1 as f64;
let n2_f = n2 as f64;
for j in 0..n2 {
for i in 0..n1 {
let idx = i + j * n1;
let ti = i as f64 / n1_f;
let tj = j as f64 / n2_f;
for d in 0..dim {
kvec[[idx, d]] = origin[d] + ti * vec1[d] + tj * vec2[d];
}
}
}
Ok(kvec)
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::arr1;
#[test]
fn test_gen_kplane_2d() {
let origin = arr1(&[0.0, 0.0]);
let vec1 = arr1(&[1.0, 0.0]);
let vec2 = arr1(&[0.0, 1.0]);
let k = gen_kplane(&origin, &vec1, &vec2, 2, 2).unwrap();
assert_eq!(k.shape(), &[4, 2]);
assert!((k[[0, 0]] - 0.0).abs() < 1e-10);
assert!((k[[0, 1]] - 0.0).abs() < 1e-10);
assert!((k[[1, 0]] - 0.5).abs() < 1e-10);
assert!((k[[1, 1]] - 0.0).abs() < 1e-10);
assert!((k[[2, 0]] - 0.0).abs() < 1e-10);
assert!((k[[2, 1]] - 0.5).abs() < 1e-10);
assert!((k[[3, 0]] - 0.5).abs() < 1e-10);
assert!((k[[3, 1]] - 0.5).abs() < 1e-10);
}
#[test]
fn test_gen_kplane_3d() {
let origin = arr1(&[0.0, 0.0, 0.5]);
let vec1 = arr1(&[1.0, 0.0, 0.0]);
let vec2 = arr1(&[0.0, 1.0, 0.0]);
let k = gen_kplane(&origin, &vec1, &vec2, 2, 2).unwrap();
assert_eq!(k.shape(), &[4, 3]);
assert!((k[[0, 2]] - 0.5).abs() < 1e-10);
assert!((k[[3, 2]] - 0.5).abs() < 1e-10);
}
#[test]
fn test_gen_kplane_length_mismatch() {
let origin = arr1(&[0.0, 0.0]);
let vec1 = arr1(&[1.0, 0.0]);
let vec2 = arr1(&[0.0, 1.0, 0.0]); let result = gen_kplane(&origin, &vec1, &vec2, 2, 2);
assert!(result.is_err());
}
}