#[macro_export]
macro_rules! tensor {
() => {{
$crate::core::Tensor::<f32>::from_data(vec![], &[0usize])
.expect("failed to create empty tensor from macro")
}};
(dtype = $ty:ty ; $([$($col:expr),+ $(,)?]),+ $(,)?) => {{
let rows: Vec<Vec<$ty>> = vec![$( vec![$($col as $ty),+] ),+];
let n_rows = rows.len();
let n_cols = rows[0].len();
for row in rows.iter() {
assert_eq!(
row.len(),
n_cols,
"tensor! macro: all rows must have equal length (expected {}, got {})",
n_cols,
row.len(),
);
}
let flat: Vec<$ty> = rows.into_iter().flatten().collect();
$crate::core::Tensor::<$ty>::from_data(flat, &[n_rows, n_cols])
.expect("failed to create 2-D tensor from macro: internal shape mismatch")
}};
(dtype = $ty:ty ; $($val:expr),+ $(,)?) => {{
let data: Vec<$ty> = vec![$($val as $ty),+];
let len = data.len();
$crate::core::Tensor::<$ty>::from_data(data, &[len])
.expect("failed to create tensor from macro: internal shape mismatch")
}};
($([$($col:expr),+ $(,)?]),+ $(,)?) => {{
let rows: Vec<Vec<_>> = vec![$( vec![$($col),+] ),+];
let n_rows = rows.len();
let n_cols = rows[0].len();
for row in rows.iter() {
assert_eq!(
row.len(),
n_cols,
"tensor! macro: all rows must have equal length (expected {}, got {})",
n_cols,
row.len(),
);
}
let flat: Vec<_> = rows.into_iter().flatten().collect();
$crate::core::Tensor::from_data(flat, &[n_rows, n_cols])
.expect("failed to create 2-D tensor from macro: internal shape mismatch")
}};
($($val:expr),+ $(,)?) => {{
let data: Vec<_> = vec![$($val),+];
let len = data.len();
$crate::core::Tensor::from_data(data, &[len])
.expect("failed to create tensor from macro: internal shape mismatch")
}};
}