use std::sync::Arc;
use torsh_core::{
device::DeviceType,
dtype::{FloatElement, TensorElement},
error::{Result, TorshError},
};
use crate::{
core_ops::{Tensor, UnaryKind},
storage::TensorStorage,
};
const PARALLEL_SUM_THRESHOLD: usize = 65_536;
#[cfg(feature = "simd")]
const SIMD_SUM_THRESHOLD: usize = 1_024;
fn sum_slice<T>(data: &[T]) -> T
where
T: TensorElement + Copy + std::ops::Add<Output = T> + num_traits::Zero,
{
let zero = <T as num_traits::Zero>::zero();
#[cfg(feature = "simd")]
{
if data.len() >= SIMD_SUM_THRESHOLD
&& std::any::TypeId::of::<T>() == std::any::TypeId::of::<f32>()
{
use scirs2_core::simd_ops::SimdUnifiedOps;
let as_f32: &[f32] =
unsafe { std::slice::from_raw_parts(data.as_ptr() as *const f32, data.len()) };
let total =
<f32 as SimdUnifiedOps>::simd_sum(&scirs2_core::ndarray::ArrayView1::from(as_f32));
if let Some(value) = <T as TensorElement>::from_f64(f64::from(total)) {
return value;
}
}
}
#[cfg(feature = "parallel")]
{
if data.len() >= PARALLEL_SUM_THRESHOLD {
use scirs2_core::parallel_ops::*;
return data
.par_chunks(16_384)
.map(|chunk| chunk.iter().fold(zero, |acc, &x| acc + x))
.reduce(|| zero, |a, b| a + b);
}
}
data.iter().fold(zero, |acc, &x| acc + x)
}
impl<T: FloatElement + Copy> Tensor<T> {
pub fn scalar(value: T) -> Result<Self> {
Self::from_data(vec![value], vec![], DeviceType::Cpu)
}
pub fn as_ndarray(&self) -> Result<scirs2_core::ndarray::ArrayD<T>> {
use scirs2_core::ndarray::ArrayD;
let data = self.data()?;
let shape_obj = self.shape().clone();
let shape = shape_obj.dims();
ArrayD::from_shape_vec(shape, data.to_vec())
.map_err(|e| TorshError::InvalidShape(format!("ndarray conversion failed: {}", e)))
}
pub fn from_ndarray(
array: scirs2_core::ndarray::ArrayD<T>,
device: DeviceType,
) -> Result<Self> {
let shape = array.shape().to_vec();
let (data, _offset) = array.into_raw_vec_and_offset();
Self::from_data(data, shape, device)
}
}
impl<T: TensorElement + Copy> Tensor<T>
where
T: PartialEq + num_traits::Zero,
{
pub fn all(&self) -> Result<Tensor<bool>> {
let data = self.to_vec()?;
let zero = <T as num_traits::Zero>::zero();
let all_true = data.iter().all(|&x| x != zero);
Tensor::from_data(vec![all_true], vec![], self.device())
}
pub fn any(&self) -> Result<Tensor<bool>> {
let data = self.to_vec()?;
let zero = <T as num_traits::Zero>::zero();
let any_true = data.iter().any(|&x| x != zero);
Tensor::from_data(vec![any_true], vec![], self.device())
}
pub fn all_dim(&self, dim: i32, keepdim: bool) -> Result<Tensor<bool>> {
let shape_binding = self.shape();
let input_shape = shape_binding.dims();
let normalized_dim = if dim < 0 {
(input_shape.len() as i32 + dim) as usize
} else {
dim as usize
};
if normalized_dim >= input_shape.len() {
return Err(torsh_core::error::TorshError::InvalidDimension {
dim: normalized_dim,
ndim: input_shape.len(),
});
}
let data = self.data()?;
let zero = <T as num_traits::Zero>::zero();
let outer_size: usize = input_shape[..normalized_dim].iter().product();
let dim_size = input_shape[normalized_dim];
let inner_size: usize = input_shape[normalized_dim + 1..].iter().product();
let output_size = outer_size * inner_size;
let mut result_data = vec![true; output_size];
for outer in 0..outer_size {
for inner in 0..inner_size {
let all_nonzero = (0..dim_size).all(|d| {
let idx = outer * dim_size * inner_size + d * inner_size + inner;
data[idx] != zero
});
let out_idx = outer * inner_size + inner;
result_data[out_idx] = all_nonzero;
}
}
let mut output_shape = input_shape.to_vec();
if keepdim {
output_shape[normalized_dim] = 1;
} else {
output_shape.remove(normalized_dim);
}
Tensor::<bool>::from_data(result_data, output_shape, self.device())
}
pub fn any_dim(&self, dim: i32, keepdim: bool) -> Result<Tensor<bool>> {
let shape_binding = self.shape();
let input_shape = shape_binding.dims();
let normalized_dim = if dim < 0 {
(input_shape.len() as i32 + dim) as usize
} else {
dim as usize
};
if normalized_dim >= input_shape.len() {
return Err(torsh_core::error::TorshError::InvalidDimension {
dim: normalized_dim,
ndim: input_shape.len(),
});
}
let data = self.data()?;
let zero = <T as num_traits::Zero>::zero();
let outer_size: usize = input_shape[..normalized_dim].iter().product();
let dim_size = input_shape[normalized_dim];
let inner_size: usize = input_shape[normalized_dim + 1..].iter().product();
let output_size = outer_size * inner_size;
let mut result_data = vec![false; output_size];
for outer in 0..outer_size {
for inner in 0..inner_size {
let any_nonzero = (0..dim_size).any(|d| {
let idx = outer * dim_size * inner_size + d * inner_size + inner;
data[idx] != zero
});
let out_idx = outer * inner_size + inner;
result_data[out_idx] = any_nonzero;
}
}
let mut output_shape = input_shape.to_vec();
if keepdim {
output_shape[normalized_dim] = 1;
} else {
output_shape.remove(normalized_dim);
}
Tensor::<bool>::from_data(result_data, output_shape, self.device())
}
}
impl<T: TensorElement + Copy> Tensor<T> {
pub fn sum(&self) -> Result<Self>
where
T: std::ops::Add<Output = T> + num_traits::Zero,
{
let sum_value = self.with_contiguous_data(|data| Ok(sum_slice(data)))?;
let mut result = Tensor::from_data(vec![sum_value], vec![], self.device())?;
if crate::should_record_grad(self.requires_grad) {
result.requires_grad = true;
result.operation = crate::core_ops::Operation::Sum {
input: Arc::new(self.clone()),
};
}
Ok(result)
}
pub fn min(&self) -> Result<Self>
where
T: std::cmp::PartialOrd + Copy,
{
let min_val = self.with_contiguous_data(|data| {
let mut iter = data.iter().copied();
let first = iter.next().ok_or_else(|| {
TorshError::InvalidOperation("Cannot compute min of empty tensor".to_string())
})?;
Ok(iter.fold(first, |acc, x| if x < acc { x } else { acc }))
})?;
Self::from_data(vec![min_val], vec![], self.device)
}
pub fn t(&self) -> Result<Self>
where
T: Copy + num_traits::Zero,
{
let shape = self.shape();
let dims = shape.dims();
if dims.len() != 2 {
return Err(TorshError::InvalidOperation(
"Transpose operation only supported for 2D tensors".to_string(),
));
}
self.transpose(0, 1)
}
pub fn shares_storage(&self, other: &Self) -> bool {
match (&self.storage, &other.storage) {
(TensorStorage::InMemory(a), TensorStorage::InMemory(b)) => Arc::ptr_eq(a, b),
(TensorStorage::MemoryMapped(a), TensorStorage::MemoryMapped(b)) => Arc::ptr_eq(a, b),
#[cfg(feature = "gpu")]
(TensorStorage::Device { buffer: a, .. }, TensorStorage::Device { buffer: b, .. }) => {
Arc::ptr_eq(a, b)
}
_ => false,
}
}
pub fn data(&self) -> Result<Vec<T>>
where
T: Copy,
{
self.to_vec()
}
pub(crate) fn with_contiguous_data<R, F>(&self, f: F) -> Result<R>
where
F: FnOnce(&[T]) -> Result<R>,
T: Copy,
{
let numel = self.numel();
if self.is_view() || self.strides.is_some() || self.storage_offset != 0 {
let data = self.to_vec()?;
return f(&data);
}
self.storage.with_slice(|slice| {
if slice.len() < numel {
return Err(TorshError::InvalidOperation(format!(
"storage holds {} elements but the tensor shape needs {}",
slice.len(),
numel
)));
}
f(&slice[..numel])
})
}
pub(crate) fn materialize_contiguous(&mut self) -> Result<()>
where
T: Copy,
{
let data_vec = self.to_vec()?;
self.storage = Self::mutable_storage(data_vec)?;
self.strides = None;
self.storage_offset = 0;
self.base_tensor = None;
Ok(())
}
fn mutable_storage(data: Vec<T>) -> Result<TensorStorage<T>> {
#[cfg(feature = "simd")]
{
const MEMORY_MAPPING_BYTES: usize = 1024 * 1024 * 1024;
if data.len() * std::mem::size_of::<T>() < MEMORY_MAPPING_BYTES {
return TensorStorage::aligned(data);
}
}
TensorStorage::create_optimal(data)
}
pub fn data_mut_apply<F>(&mut self, mut func: F) -> Result<()>
where
F: FnMut(&mut T),
T: Copy,
{
self.prepare_for_inplace()?;
match &self.storage {
TensorStorage::MemoryMapped(_) => {
let mut new_data = self.to_vec()?;
for item in new_data.iter_mut() {
func(item);
}
self.storage = TensorStorage::create_optimal(new_data)?;
Ok(())
}
_ => {
let numel = self.numel();
self.storage.with_slice_mut(|slice| {
for item in slice.iter_mut().take(numel) {
func(item);
}
Ok(())
})
}
}
}
pub fn clone_data(&self) -> Self
where
T: Copy,
{
let data = self
.to_vec()
.expect("tensor to vec conversion should succeed");
Self::from_data(data, self.shape().dims().to_vec(), self.device)
.expect("tensor creation should succeed")
}
pub fn make_unique(&mut self) -> Result<()> {
if self.is_view() || self.strides.is_some() || self.storage_offset != 0 {
return self.materialize_contiguous();
}
match &self.storage {
TensorStorage::InMemory(data) => {
if Arc::strong_count(data) > 1 {
let data_vec = self.to_vec()?;
self.storage = Self::mutable_storage(data_vec)?;
}
}
TensorStorage::MemoryMapped(storage) => {
if Arc::strong_count(storage) > 1 {
let data_vec = self.to_vec()?;
self.storage = TensorStorage::create_optimal(data_vec)?;
}
}
#[cfg(feature = "simd")]
TensorStorage::Aligned(data) => {
if Arc::strong_count(data) > 1 {
let data_vec = self.to_vec()?;
self.storage = Self::mutable_storage(data_vec)?;
}
}
#[cfg(feature = "simd")]
TensorStorage::SimdOptimized(storage) => {
let promoted = storage.with_slice(TensorStorage::aligned_from_slice)?;
self.storage = promoted;
}
#[cfg(feature = "gpu")]
TensorStorage::Device { .. } => {
let data_vec = self.to_vec()?;
self.storage = Self::mutable_storage(data_vec)?;
}
}
Ok(())
}
pub(crate) fn prepare_for_inplace(&mut self) -> Result<()>
where
T: Copy,
{
self.make_unique()
}
pub fn apply_<F>(&mut self, func: F) -> Result<()>
where
F: Fn(T) -> T,
T: Copy,
{
self.prepare_for_inplace()?;
if matches!(self.storage, TensorStorage::MemoryMapped(_)) {
let data = self.to_vec()?;
let new_data: Vec<T> = data.into_iter().map(func).collect();
self.storage = TensorStorage::create_optimal(new_data)?;
return Ok(());
}
let numel = self.numel();
self.storage.with_slice_mut(|slice| {
for value in slice.iter_mut().take(numel) {
*value = func(*value);
}
Ok(())
})
}
pub fn map<F>(&self, func: F) -> Result<Self>
where
F: Fn(T) -> T,
T: Copy,
{
let new_data = self.with_contiguous_data(|data| {
let mut out = Vec::with_capacity(data.len());
out.extend(data.iter().map(|&x| func(x)));
Ok(out)
})?;
Self::from_data(new_data, self.shape().dims().to_vec(), self.device)
}
pub fn item(&self) -> Result<T>
where
T: Copy,
{
let data = self.data()?;
if data.len() != 1 {
return Err(TorshError::InvalidArgument(format!(
"item() can only be called on single-element tensors, got {} elements",
data.len()
)));
}
Ok(data[0])
}
pub fn cat(tensors: &[&Self], dim: i32) -> Result<Self>
where
T: Copy,
{
if tensors.is_empty() {
return Err(TorshError::InvalidArgument(
"Cannot concatenate empty tensor list".to_string(),
));
}
let first_shape_binding = tensors[0].shape();
let first_shape = first_shape_binding.dims();
let ndim = first_shape.len();
let actual_dim = if dim < 0 {
(ndim as i32 + dim) as usize
} else {
dim as usize
};
if actual_dim >= ndim {
return Err(TorshError::InvalidArgument(format!(
"Dimension {} out of range for {}-dimensional tensor",
dim, ndim
)));
}
for (i, tensor) in tensors.iter().enumerate().skip(1) {
let shape_binding = tensor.shape();
let shape = shape_binding.dims();
if shape.len() != ndim {
return Err(TorshError::InvalidArgument(format!(
"Tensor {} has {} dimensions but first tensor has {}",
i,
shape.len(),
ndim
)));
}
for (d, (&s1, &s2)) in first_shape.iter().zip(shape.iter()).enumerate() {
if d != actual_dim && s1 != s2 {
return Err(TorshError::ShapeMismatch {
expected: first_shape.to_vec(),
got: shape.to_vec(),
});
}
}
}
let cat_dim_total: usize = tensors.iter().map(|t| t.shape().dims()[actual_dim]).sum();
let mut result_shape = first_shape.to_vec();
result_shape[actual_dim] = cat_dim_total;
let outer_size: usize = first_shape[..actual_dim].iter().product();
let inner_size: usize = first_shape[actual_dim + 1..].iter().product();
let total_numel: usize = result_shape.iter().product();
let mut result_data = Vec::with_capacity(total_numel);
let sources: Vec<Vec<T>> = tensors
.iter()
.map(|tensor| tensor.to_vec())
.collect::<Result<Vec<_>>>()?;
let cat_sizes: Vec<usize> = tensors
.iter()
.map(|tensor| tensor.shape().dims()[actual_dim])
.collect();
for outer in 0..outer_size {
for (source, &cat_size) in sources.iter().zip(cat_sizes.iter()) {
let run = cat_size * inner_size;
let start = outer * run;
result_data.extend_from_slice(&source[start..start + run]);
}
}
let mut result = Self::from_data(result_data, result_shape, tensors[0].device)?;
let any_requires_grad = tensors.iter().any(|t| t.requires_grad);
if crate::should_record_grad(any_requires_grad) {
result.requires_grad = true;
result.operation = crate::core_ops::Operation::Concat {
inputs: tensors.iter().map(|t| Arc::new((*t).clone())).collect(),
dim: actual_dim,
};
}
Ok(result)
}
}
impl<T: TensorElement + Copy> Tensor<T>
where
T: num_traits::Float,
{
pub fn norm(&self) -> Result<Self> {
let data = self.data()?;
let sum_squares: T = data
.iter()
.map(|&x| x * x)
.fold(num_traits::Zero::zero(), |acc, x| acc + x);
let norm_value = sum_squares.sqrt();
Tensor::from_data(vec![norm_value], vec![], self.device())
}
pub fn norm_lp(&self, p: f64, dims: Option<&[usize]>, keepdim: bool) -> Result<Self>
where
T: num_traits::FromPrimitive,
{
let shape_binding = self.shape();
let input_shape = shape_binding.dims().to_vec();
let ndim = input_shape.len();
let reduce_dims: Vec<usize> = match dims {
Some(requested) => {
for &dim in requested {
if dim >= ndim {
return Err(TorshError::InvalidOperation(format!(
"Dimension {} out of range for {}-dimensional tensor",
dim, ndim
)));
}
}
let mut normalized = requested.to_vec();
normalized.sort_unstable();
normalized.dedup();
normalized
}
None => (0..ndim).collect(),
};
let convert = |value: f64| -> Result<T> {
<T as num_traits::FromPrimitive>::from_f64(value).ok_or_else(|| {
TorshError::InvalidOperation(format!(
"norm: p={} cannot be represented in this tensor's element type",
p
))
})
};
let zero = <T as num_traits::Zero>::zero();
let one = <T as num_traits::One>::one();
#[derive(Clone, Copy, PartialEq)]
enum NormKind {
L0,
L1,
L2,
MaxAbs,
MinAbs,
General,
}
let kind = if p == 1.0 {
NormKind::L1
} else if p == 2.0 {
NormKind::L2
} else if p == 0.0 {
NormKind::L0
} else if p == f64::INFINITY {
NormKind::MaxAbs
} else if p == f64::NEG_INFINITY {
NormKind::MinAbs
} else {
NormKind::General
};
let (p_t, inv_p_t) = if kind == NormKind::General {
(convert(p)?, convert(1.0 / p)?)
} else {
(zero, zero)
};
let init: T = if kind == NormKind::MinAbs {
<T as num_traits::Float>::infinity()
} else {
zero
};
let elem = |x: T| -> T {
match kind {
NormKind::L1 | NormKind::MaxAbs | NormKind::MinAbs => x.abs(),
NormKind::L2 => x * x,
NormKind::L0 => {
if x == zero {
zero
} else {
one
}
}
NormKind::General => x.abs().powf(p_t),
}
};
let combine = |a: T, b: T| -> T {
match kind {
NormKind::MaxAbs => a.max(b),
NormKind::MinAbs => a.min(b),
_ => a + b,
}
};
let finalize = |s: T| -> T {
match kind {
NormKind::L2 => s.sqrt(),
NormKind::General => s.powf(inv_p_t),
_ => s,
}
};
let data = self.data()?;
if ndim == 0 || reduce_dims.len() == ndim {
let acc = data.iter().fold(init, |acc, &x| combine(acc, elem(x)));
let value = finalize(acc);
let out_shape = if keepdim { vec![1; ndim] } else { vec![] };
return Self::from_data(vec![value], out_shape, self.device());
}
let mut is_reduced = vec![false; ndim];
for &d in &reduce_dims {
is_reduced[d] = true;
}
let mut input_strides = vec![1usize; ndim];
for i in (0..ndim - 1).rev() {
input_strides[i] = input_strides[i + 1] * input_shape[i + 1];
}
let mut output_shape_keepdim = input_shape.clone();
for &d in &reduce_dims {
output_shape_keepdim[d] = 1;
}
let mut output_strides = vec![1usize; ndim];
for i in (0..ndim - 1).rev() {
output_strides[i] = output_strides[i + 1] * output_shape_keepdim[i + 1];
}
let output_size: usize = output_shape_keepdim.iter().product();
let mut acc = vec![init; output_size];
for (flat_idx, &x) in data.iter().enumerate() {
let mut remaining = flat_idx;
let mut out_flat = 0usize;
for (dim, &reduced) in is_reduced.iter().enumerate() {
let coord = remaining / input_strides[dim];
remaining %= input_strides[dim];
if !reduced {
out_flat += coord * output_strides[dim];
}
}
acc[out_flat] = combine(acc[out_flat], elem(x));
}
let result_data: Vec<T> = acc.into_iter().map(finalize).collect();
let final_shape = if keepdim {
output_shape_keepdim
} else {
input_shape
.into_iter()
.zip(is_reduced.iter())
.filter(|(_, &reduced)| !reduced)
.map(|(size, _)| size)
.collect::<Vec<_>>()
};
Self::from_data(result_data, final_shape, self.device())
}
}
impl<T: TensorElement + Copy> Tensor<T> {
pub fn matmul_scirs2(&self, other: &Self) -> Result<Self>
where
T: num_traits::Float + num_traits::Zero + num_traits::One + std::iter::Sum,
{
self.basic_matmul(other)
}
pub fn sum_scirs2(&self) -> Result<Self>
where
T: std::ops::Add<Output = T> + num_traits::Zero,
{
let data = self.data()?;
let sum_value = data
.iter()
.fold(<T as num_traits::Zero>::zero(), |acc, &x| acc + x);
Tensor::from_data(vec![sum_value], vec![], self.device())
}
pub fn mean_scirs2(&self) -> Result<Self>
where
T: std::ops::Add<Output = T>
+ std::ops::Div<Output = T>
+ num_traits::Zero
+ From<usize>
+ num_traits::FromPrimitive,
{
let data = self.data()?;
if data.is_empty() {
return Err(TorshError::InvalidArgument(
"Cannot compute mean of empty tensor".to_string(),
));
}
let sum_value = data
.iter()
.fold(<T as num_traits::Zero>::zero(), |acc, &x| acc + x);
let mean_value = sum_value / T::from(data.len());
Tensor::from_data(vec![mean_value], vec![], self.device())
}
pub fn relu_scirs2(&self) -> Result<Self>
where
T: PartialOrd + num_traits::Zero,
{
let zero = <T as num_traits::Zero>::zero();
let result = self.map(|x| if x > zero { x } else { zero })?;
Ok(self.record_unary(result, UnaryKind::Relu))
}
pub fn sigmoid_scirs2(&self) -> Result<Self>
where
T: num_traits::Float,
{
let result = self.map(|x| {
let one = <T as num_traits::One>::one();
one / (one + (-x).exp())
})?;
Ok(self.record_unary(result, UnaryKind::Sigmoid))
}
pub fn tanh_scirs2(&self) -> Result<Self>
where
T: num_traits::Float,
{
let result = self.map(|x| x.tanh())?;
Ok(self.record_unary(result, UnaryKind::Tanh))
}
pub fn softmax(&self, dim: i32) -> Result<Self>
where
T: torsh_core::dtype::FloatElement
+ Copy
+ std::ops::Sub<Output = T>
+ std::ops::Div<Output = T>,
{
let data = self.data()?;
let shape_binding = self.shape();
let shape = shape_binding.dims();
if data.is_empty() || shape.is_empty() {
return Err(TorshError::InvalidOperation(
"Cannot compute softmax on empty tensor".to_string(),
));
}
let actual_dim = if dim < 0 {
(shape.len() as i32 + dim) as usize
} else {
dim as usize
};
if actual_dim >= shape.len() {
return Err(TorshError::InvalidOperation(format!(
"Dimension {} out of range for {}-dimensional tensor",
actual_dim,
shape.len()
)));
}
let max_tensor = self.max(Some(actual_dim), true)?;
let expanded_max = max_tensor.expand(shape)?;
let shifted = self.sub(&expanded_max)?;
let exp_tensor = shifted.exp()?;
let sum_tensor = exp_tensor.sum_dim(&[actual_dim as i32], true)?;
let expanded_sum = sum_tensor.expand(shape)?;
exp_tensor.div(&expanded_sum)
}
pub fn log_softmax(&self, dim: i32) -> Result<Self>
where
T: torsh_core::dtype::FloatElement + Copy + std::ops::Sub<Output = T>,
{
let shape_binding = self.shape();
let shape = shape_binding.dims().to_vec();
if shape.is_empty() {
return Err(TorshError::InvalidOperation(
"Cannot compute log_softmax on empty tensor".to_string(),
));
}
let actual_dim = if dim < 0 {
(shape.len() as i32 + dim) as usize
} else {
dim as usize
};
if actual_dim >= shape.len() {
return Err(TorshError::InvalidOperation(format!(
"Dimension {} out of range for {}-dimensional tensor",
actual_dim,
shape.len()
)));
}
let max_tensor = self.max(Some(actual_dim), true)?;
let expanded_max = max_tensor.expand(&shape)?;
let shifted = self.sub(&expanded_max)?;
let sum_exp = shifted.exp()?.sum_dim(&[actual_dim as i32], true)?;
let log_sum_exp = sum_exp.log()?;
let expanded_log_sum = log_sum_exp.expand(&shape)?;
let mut result = shifted.sub(&expanded_log_sum)?;
if crate::should_record_grad(self.requires_grad) {
result.requires_grad = true;
result.operation = crate::core_ops::Operation::LogSoftmax {
input: Arc::new(self.clone()),
dim: actual_dim,
};
}
Ok(result)
}
pub fn topk(
&self,
k: usize,
dim: Option<i32>,
largest: bool,
sorted: bool,
) -> Result<(Self, Tensor<i64>)>
where
T: std::cmp::PartialOrd + Copy + num_traits::Zero,
{
let data = self.data()?;
let shape_binding = self.shape();
let shape = shape_binding.dims();
if shape.is_empty() {
return Err(TorshError::InvalidOperation(
"Cannot compute topk on empty tensor".to_string(),
));
}
if k == 0 {
return Err(TorshError::InvalidArgument(
"k must be greater than 0".to_string(),
));
}
let actual_dim = match dim {
Some(d) => {
let norm = if d < 0 {
(shape.len() as i32 + d) as usize
} else {
d as usize
};
if norm >= shape.len() {
return Err(TorshError::InvalidArgument(format!(
"Dimension {} out of range for {}-dimensional tensor",
d,
shape.len()
)));
}
norm
}
None => shape.len() - 1,
};
let dim_size = shape[actual_dim];
let effective_k = k.min(dim_size);
let outer_size: usize = shape[..actual_dim].iter().product();
let inner_size: usize = shape[actual_dim + 1..].iter().product();
let mut result_shape = shape.to_vec();
result_shape[actual_dim] = effective_k;
let mut values_data = Vec::with_capacity(outer_size * effective_k * inner_size);
let mut indices_data = Vec::with_capacity(outer_size * effective_k * inner_size);
for outer in 0..outer_size {
for inner in 0..inner_size {
let mut slice: Vec<(usize, T)> = (0..dim_size)
.map(|d| {
let src = outer * dim_size * inner_size + d * inner_size + inner;
(d, data[src])
})
.collect();
if largest {
slice
.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
} else {
slice
.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
}
let mut top_k: Vec<(usize, T)> = slice.into_iter().take(effective_k).collect();
if !sorted {
top_k.sort_by_key(|(idx, _)| *idx);
}
for (local_idx, val) in &top_k {
values_data.push(*val);
indices_data.push(*local_idx as i64);
}
}
}
let transposed_len = outer_size * effective_k * inner_size;
let mut values_transposed = Vec::with_capacity(transposed_len);
let mut indices_transposed = Vec::with_capacity(transposed_len);
for outer in 0..outer_size {
for k_idx in 0..effective_k {
for inner in 0..inner_size {
let src = outer * inner_size * effective_k + inner * effective_k + k_idx;
values_transposed.push(values_data[src]);
indices_transposed.push(indices_data[src]);
}
}
}
let values_tensor = Self::from_data(values_transposed, result_shape.clone(), self.device)?;
let indices_tensor =
Tensor::<i64>::from_data(indices_transposed, result_shape, self.device)?;
Ok((values_tensor, indices_tensor))
}
}
#[cfg(test)]
mod tests {
use super::*;
use torsh_core::device::DeviceType;
#[cfg(feature = "simd")]
#[test]
fn make_unique_isolates_shared_simd_storage() {
let mut tensor = Tensor::from_data(vec![1.0f32; 4096], vec![4096], DeviceType::Cpu)
.expect("tensor creation should succeed");
assert!(
matches!(tensor.storage, TensorStorage::SimdOptimized(_)),
"precondition: a 16 KiB tensor uses lock-free SIMD storage"
);
let sibling = tensor.clone();
tensor.make_unique().expect("make_unique should succeed");
assert!(
!matches!(tensor.storage, TensorStorage::SimdOptimized(_)),
"SimdOptimized storage must be promoted to single-buffer mutable storage"
);
tensor.apply_(|value| value + 1.0).expect("apply_");
assert_eq!(
tensor.to_vec().expect("to_vec")[0],
2.0,
"the writer must see its own update"
);
assert_eq!(
sibling.to_vec().expect("to_vec")[0],
1.0,
"the sibling must not observe the write"
);
}
#[cfg(feature = "simd")]
#[test]
fn make_unique_preserves_mutated_simd_contents() {
let source = Tensor::from_data(
(0..4096).map(|i| i as f32).collect::<Vec<_>>(),
vec![4096],
DeviceType::Cpu,
)
.expect("tensor creation should succeed");
source
.storage
.with_slice_mut(|slice| {
slice[0] = -1.0;
Ok(())
})
.expect("copy-on-write write should succeed");
let mut promoted = source.clone();
promoted.make_unique().expect("make_unique should succeed");
let data = promoted.to_vec().expect("to_vec");
assert_eq!(data[0], -1.0, "the promotion must copy the CoW buffer");
assert_eq!(data[4095], 4095.0);
}
#[test]
fn test_scalar_creation() {
let scalar = Tensor::<f32>::scalar(42.0).expect("operation should succeed");
assert_eq!(scalar.shape().dims(), &[] as &[usize]);
assert_eq!(scalar.item().expect("item extraction should succeed"), 42.0);
}
#[test]
fn test_max_reduction() {
let data = vec![1.0f32, 5.0, 3.0, 2.0];
let tensor =
Tensor::from_data(data, vec![4], DeviceType::Cpu).expect("operation should succeed");
let max_val = tensor.max(None, false).expect("operation should succeed");
assert_eq!(max_val.item().expect("item extraction should succeed"), 5.0);
}
#[test]
fn test_norm_computation() {
let data = vec![3.0f32, 4.0]; let tensor =
Tensor::from_data(data, vec![2], DeviceType::Cpu).expect("operation should succeed");
let norm = tensor.norm().expect("norm computation should succeed");
assert!((norm.item().expect("item extraction should succeed") - 5.0).abs() < 1e-6);
}
#[test]
fn test_apply_operations() {
let data = vec![1.0f32, 2.0, 3.0, 4.0];
let mut tensor =
Tensor::from_data(data, vec![4], DeviceType::Cpu).expect("operation should succeed");
tensor
.apply_(|x| x * 2.0)
.expect("operation should succeed");
assert_eq!(
tensor.data().expect("data retrieval should succeed"),
vec![2.0, 4.0, 6.0, 8.0]
);
let original = Tensor::from_data(vec![1.0f32, 2.0, 3.0], vec![3], DeviceType::Cpu)
.expect("operation should succeed");
let mapped = original.map(|x| x + 1.0).expect("operation should succeed");
assert_eq!(
mapped.data().expect("data retrieval should succeed"),
vec![2.0, 3.0, 4.0]
);
assert_eq!(
original.data().expect("data retrieval should succeed"),
vec![1.0, 2.0, 3.0]
); }
#[test]
fn test_activation_functions() {
let data = vec![-1.0f32, 0.0, 1.0, 2.0];
let tensor =
Tensor::from_data(data, vec![4], DeviceType::Cpu).expect("operation should succeed");
let relu_result = tensor.relu().expect("relu should succeed");
assert_eq!(
relu_result.data().expect("data retrieval should succeed"),
vec![0.0, 0.0, 1.0, 2.0]
);
let abs_result = tensor.abs().expect("abs computation should succeed");
assert_eq!(
abs_result.data().expect("data retrieval should succeed"),
vec![1.0, 0.0, 1.0, 2.0]
);
let clamped = tensor.clamp(-0.5, 1.5).expect("operation should succeed");
assert_eq!(
clamped.data().expect("data retrieval should succeed"),
vec![-0.5, 0.0, 1.0, 1.5]
);
}
#[test]
fn test_storage_sharing() {
let tensor1 =
Tensor::<f32>::zeros(&[2, 2], DeviceType::Cpu).expect("operation should succeed");
let tensor2 = tensor1.clone();
let tensor3 = tensor1.clone_data();
assert!(tensor1.shares_storage(&tensor2));
assert!(!tensor1.shares_storage(&tensor3));
}
#[test]
fn test_basic_matmul() {
let a = Tensor::from_data(vec![1.0f32, 2.0, 3.0, 4.0], vec![2, 2], DeviceType::Cpu)
.expect("operation should succeed");
let b = Tensor::from_data(vec![5.0f32, 6.0, 7.0, 8.0], vec![2, 2], DeviceType::Cpu)
.expect("operation should succeed");
let result = a.basic_matmul(&b).expect("operation should succeed");
assert_eq!(result.shape().dims(), &[2, 2]);
let expected = vec![19.0, 22.0, 43.0, 50.0];
assert_eq!(
result.data().expect("data retrieval should succeed"),
expected
);
}
#[test]
fn test_reductions() {
let data = vec![1.0f32, 2.0, 3.0, 4.0];
let tensor =
Tensor::from_data(data, vec![4], DeviceType::Cpu).expect("operation should succeed");
let sum = tensor.sum().expect("sum should succeed");
assert_eq!(sum.item().expect("item extraction should succeed"), 10.0);
let mean = tensor.mean(None, false).expect("operation should succeed");
assert_eq!(mean.item().expect("item extraction should succeed"), 2.5);
}
#[test]
fn test_copy_on_write() {
let mut tensor1 =
Tensor::<f32>::ones(&[2], DeviceType::Cpu).expect("operation should succeed");
let tensor2 = tensor1.clone();
assert!(tensor1.shares_storage(&tensor2));
tensor1.make_unique().expect("make_unique should succeed");
assert!(!tensor1.shares_storage(&tensor2));
}
#[test]
fn test_make_unique_does_not_realloc_unshared_storage() {
let mut tensor = Tensor::<f32>::from_data(vec![1.0; 4096], vec![4096], DeviceType::Cpu)
.expect("tensor creation should succeed");
assert_eq!(
tensor.storage_type(),
"simd_optimized",
"precondition: large f32 tensors start out immutable"
);
tensor.make_unique().expect("make_unique should succeed");
assert_eq!(tensor.storage_type(), "aligned_simd");
let ptr_after_first = tensor
.storage
.with_slice(|s| Ok(s.as_ptr() as usize))
.expect("with_slice should succeed");
tensor.make_unique().expect("make_unique should succeed");
let ptr_after_second = tensor
.storage
.with_slice(|s| Ok(s.as_ptr() as usize))
.expect("with_slice should succeed");
assert_eq!(
ptr_after_first, ptr_after_second,
"make_unique must not copy storage that is already unique and mutable"
);
}
#[test]
fn test_make_unique_cow_copy_stays_mutable() {
let mut tensor = Tensor::<f32>::from_data(vec![1.0; 4096], vec![4096], DeviceType::Cpu)
.expect("tensor creation should succeed");
tensor.make_unique().expect("make_unique should succeed");
assert_eq!(tensor.storage_type(), "aligned_simd");
let shared = tensor.clone();
tensor.make_unique().expect("make_unique should succeed");
assert_eq!(
tensor.storage_type(),
"aligned_simd",
"a copy-on-write split must not fall back to immutable storage"
);
tensor.apply_(|x| x + 1.0).expect("apply_ should succeed");
assert_eq!(tensor.to_vec().expect("to_vec")[0], 2.0);
assert_eq!(
shared.to_vec().expect("to_vec")[0],
1.0,
"the shared tensor must be unaffected"
);
}
#[test]
fn test_apply_mutates_storage_in_place() {
let mut tensor = Tensor::<f32>::from_data(vec![2.0; 4096], vec![4096], DeviceType::Cpu)
.expect("tensor creation should succeed");
tensor.make_unique().expect("make_unique should succeed");
let ptr_before = tensor
.storage
.with_slice(|s| Ok(s.as_ptr() as usize))
.expect("with_slice should succeed");
tensor.apply_(|x| x * 3.0).expect("apply_ should succeed");
let ptr_after = tensor
.storage
.with_slice(|s| Ok(s.as_ptr() as usize))
.expect("with_slice should succeed");
assert_eq!(ptr_before, ptr_after, "apply_ must not reallocate storage");
assert_eq!(tensor.to_vec().expect("to_vec")[0], 6.0);
}
#[test]
fn test_item_extraction() {
let scalar = Tensor::from_data(vec![42.0f32], vec![], DeviceType::Cpu)
.expect("operation should succeed");
assert_eq!(scalar.item().expect("item extraction should succeed"), 42.0);
let vector = Tensor::from_data(vec![1.0f32, 2.0], vec![2], DeviceType::Cpu)
.expect("operation should succeed");
assert!(vector.item().is_err()); }
#[test]
fn test_all_dim() {
let data = vec![1i32, 0, 1, 1, 1, 1];
let tensor = Tensor::from_data(data, vec![2, 3], DeviceType::Cpu)
.expect("tensor creation should succeed");
let result = tensor.all_dim(0, false).expect("all_dim should succeed");
assert_eq!(result.shape().dims(), &[3]);
assert_eq!(
result.to_vec().expect("to_vec should succeed"),
vec![true, false, true]
);
let result_row = tensor.all_dim(1, false).expect("all_dim should succeed");
assert_eq!(result_row.shape().dims(), &[2]);
assert_eq!(
result_row.to_vec().expect("to_vec should succeed"),
vec![false, true]
);
let result_kd = tensor.all_dim(1, true).expect("all_dim should succeed");
assert_eq!(result_kd.shape().dims(), &[2, 1]);
}
#[test]
fn test_any_dim() {
let data = vec![0i32, 0, 0, 0, 1, 0];
let tensor = Tensor::from_data(data, vec![2, 3], DeviceType::Cpu)
.expect("tensor creation should succeed");
let result = tensor.any_dim(0, false).expect("any_dim should succeed");
assert_eq!(result.shape().dims(), &[3]);
assert_eq!(
result.to_vec().expect("to_vec should succeed"),
vec![false, true, false]
);
let result_row = tensor.any_dim(1, false).expect("any_dim should succeed");
assert_eq!(result_row.shape().dims(), &[2]);
assert_eq!(
result_row.to_vec().expect("to_vec should succeed"),
vec![false, true]
);
}
#[test]
fn test_cat_multidim() {
let a = Tensor::from_data(vec![1.0f32, 2.0, 3.0, 4.0], vec![2, 2], DeviceType::Cpu)
.expect("tensor creation should succeed");
let b = Tensor::from_data(vec![5.0f32, 6.0], vec![1, 2], DeviceType::Cpu)
.expect("tensor creation should succeed");
let cat0 = Tensor::<f32>::cat(&[&a, &b], 0).expect("cat should succeed");
assert_eq!(cat0.shape().dims(), &[3, 2]);
assert_eq!(
cat0.to_vec().expect("to_vec should succeed"),
vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]
);
let c = Tensor::from_data(vec![1.0f32, 2.0, 3.0, 4.0], vec![2, 2], DeviceType::Cpu)
.expect("tensor creation should succeed");
let d = Tensor::from_data(vec![5.0f32, 6.0, 7.0, 8.0], vec![2, 2], DeviceType::Cpu)
.expect("tensor creation should succeed");
let cat1 = Tensor::<f32>::cat(&[&c, &d], 1).expect("cat should succeed");
assert_eq!(cat1.shape().dims(), &[2, 4]);
assert_eq!(
cat1.to_vec().expect("to_vec should succeed"),
vec![1.0, 2.0, 5.0, 6.0, 3.0, 4.0, 7.0, 8.0]
);
}
#[test]
fn test_topk_along_dim() {
let data = vec![3.0f32, 1.0, 4.0, 2.0, 5.0, 9.0, 2.0, 6.0];
let tensor = Tensor::from_data(data, vec![2, 4], DeviceType::Cpu)
.expect("tensor creation should succeed");
let (vals, idxs) = tensor
.topk(2, Some(1), true, true)
.expect("topk should succeed");
assert_eq!(vals.shape().dims(), &[2, 2]);
assert_eq!(idxs.shape().dims(), &[2, 2]);
let vals_data = vals.to_vec().expect("to_vec should succeed");
let idxs_data = idxs.to_vec().expect("to_vec should succeed");
assert_eq!(vals_data[0], 4.0);
assert_eq!(vals_data[1], 3.0);
assert_eq!(vals_data[2], 9.0);
assert_eq!(vals_data[3], 6.0);
assert_eq!(idxs_data[0], 2);
assert_eq!(idxs_data[1], 0);
assert_eq!(idxs_data[2], 1);
assert_eq!(idxs_data[3], 3);
}
#[test]
fn test_issue_43_mean_propagates_requires_grad() {
let input = Tensor::from_data(vec![1.0f32, 2.0, 3.0, 4.0], vec![4], DeviceType::Cpu)
.expect("tensor creation failed")
.requires_grad_(true);
let result = input.mean(None, false).expect("mean should succeed");
assert!(
result.requires_grad(),
"mean result must have requires_grad=true when input does"
);
}
#[test]
fn test_issue_43_mean_no_requires_grad_when_input_has_none() {
let input = Tensor::from_data(vec![1.0f32, 2.0, 3.0], vec![3], DeviceType::Cpu)
.expect("tensor creation failed");
let result = input.mean(None, false).expect("mean should succeed");
assert!(
!result.requires_grad(),
"mean result must not require grad when input does not"
);
}
#[test]
fn test_issue_43_mean_backward() {
let n = 4usize;
let input = Tensor::from_data(vec![2.0f32, 4.0, 6.0, 8.0], vec![n], DeviceType::Cpu)
.expect("tensor creation failed")
.requires_grad_(true);
let result = input.mean(None, false).expect("mean should succeed");
assert!(result.requires_grad(), "mean result must track gradients");
result.backward().expect("backward should succeed");
let grad = input
.grad()
.expect("input must have gradient after backward");
let grad_data = grad.data().expect("gradient data");
let expected = 1.0f32 / n as f32;
for &g in &grad_data {
assert!(
(g - expected).abs() < 1e-6,
"each element grad should be 1/n={expected}, got {g}"
);
}
}
#[test]
fn test_sum_backward() {
let x = Tensor::from_data(vec![1.0f32, 2.0, 3.0, 4.0], vec![4], DeviceType::Cpu)
.expect("tensor creation failed")
.requires_grad_(true);
let loss = x.sum().expect("sum should succeed");
assert!(loss.requires_grad(), "sum result must track gradients");
loss.backward().expect("backward should succeed");
let grad = x.grad().expect("x must have a gradient after backward");
let grad_data = grad.data().expect("gradient data");
assert_eq!(
grad_data,
vec![1.0f32, 1.0, 1.0, 1.0],
"d(sum)/dx must be all ones"
);
}
#[test]
fn test_matmul_backward() {
let a = Tensor::from_data(vec![1.0f32, 2.0, 3.0, 4.0], vec![2, 2], DeviceType::Cpu)
.expect("tensor creation failed")
.requires_grad_(true);
let b = Tensor::from_data(vec![5.0f32, 6.0, 7.0, 8.0], vec![2, 2], DeviceType::Cpu)
.expect("tensor creation failed")
.requires_grad_(true);
let c = a.matmul(&b).expect("matmul should succeed");
assert!(c.requires_grad(), "matmul result must track gradients");
let loss = c.sum().expect("sum should succeed");
loss.backward().expect("backward should succeed");
let grad_a = a
.grad()
.expect("A must have a gradient")
.data()
.expect("grad data");
let grad_b = b
.grad()
.expect("B must have a gradient")
.data()
.expect("grad data");
assert_eq!(
grad_a,
vec![11.0f32, 15.0, 11.0, 15.0],
"grad_A = ones @ Bᵀ"
);
assert_eq!(grad_b, vec![4.0f32, 4.0, 6.0, 6.0], "grad_B = Aᵀ @ ones");
}
}