diff --git a/.github/workflows/lints.yml b/.github/workflows/lints.yml index 0a84bda5..e3e765f1 100644 --- a/.github/workflows/lints.yml +++ b/.github/workflows/lints.yml @@ -21,3 +21,17 @@ jobs: - uses: pre-commit/action@v3.0.1 with: extra_args: --all-files + + clippy: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v7 + - uses: actions/setup-python@v6 + with: + python-version: "3.12" + - uses: dtolnay/rust-toolchain@stable + with: + components: clippy + - uses: Swatinem/rust-cache@v2 + - name: Run clippy (deny warnings) + run: cargo clippy --workspace --all-targets -- -D warnings diff --git a/bindings/src/debug.rs b/bindings/src/debug.rs index 71b6d161..f33117a3 100644 --- a/bindings/src/debug.rs +++ b/bindings/src/debug.rs @@ -30,7 +30,7 @@ impl PyTensorInfo { /// Get tensor device #[getter] - fn device(&self) -> String { + pub(crate) fn device(&self) -> String { format!("{:?}", self.inner.device) } @@ -48,7 +48,7 @@ impl PyTensorInfo { /// Check if is leaf node #[getter] - fn is_leaf(&self) -> bool { + pub(crate) fn is_leaf(&self) -> bool { self.inner.is_leaf } diff --git a/bindings/src/functional.rs b/bindings/src/functional.rs index acfa7beb..4a917fd3 100644 --- a/bindings/src/functional.rs +++ b/bindings/src/functional.rs @@ -15,7 +15,7 @@ use pyo3::prelude::*; use pyo3::types::{PyAny, PyList, PyTuple}; use std::sync::Arc; -fn borrow_tensor<'py>(value: &'py Bound<'py, PyAny>) -> PyResult> { +pub(crate) fn borrow_tensor<'py>(value: &'py Bound<'py, PyAny>) -> PyResult> { if let Ok(tensor) = value.extract::>() { return Ok(tensor); } diff --git a/bindings/src/lib.rs b/bindings/src/lib.rs index 5f58767b..5ca81bd8 100644 --- a/bindings/src/lib.rs +++ b/bindings/src/lib.rs @@ -66,10 +66,15 @@ fn _core(py: Python, m: &Bound) -> PyResult<()> { serialization::register_serialization_module(py, m)?; // Autograd helpers + m.add_class::()?; m.add_function(wrap_pyfunction!(get_gradient, m)?)?; m.add_function(wrap_pyfunction!(clear_autograd_graph, m)?)?; m.add_function(wrap_pyfunction!(is_autograd_graph_consumed, m)?)?; m.add_function(wrap_pyfunction!(mark_autograd_graph_consumed, m)?)?; + m.add_function(wrap_pyfunction!(no_grad, m)?)?; + m.add_function(wrap_pyfunction!(enable_grad, m)?)?; + m.add_function(wrap_pyfunction!(is_grad_enabled, m)?)?; + m.add_function(wrap_pyfunction!(set_grad_enabled, m)?)?; m.add_function(wrap_pyfunction!(get_default_dtype, m)?)?; m.add_function(wrap_pyfunction!(set_default_dtype, m)?)?; @@ -78,6 +83,68 @@ fn _core(py: Python, m: &Bound) -> PyResult<()> { Ok(()) } +/// Context manager that sets the thread-local autograd recording mode on +/// entry and restores the previous mode on exit. Re-entrant: each `with` +/// block restores whatever mode was active when it was entered. +#[pyclass(name = "GradMode")] +struct GradMode { + target: bool, + previous: Option, +} + +#[pymethods] +impl GradMode { + fn __enter__(mut slf: PyRefMut<'_, Self>) -> PyRefMut<'_, Self> { + let prev = engine::autograd::set_grad_enabled(slf.target); + slf.previous = Some(prev); + slf + } + + #[pyo3(signature = (*_args))] + fn __exit__(&mut self, _args: &Bound<'_, pyo3::types::PyTuple>) -> bool { + if let Some(prev) = self.previous.take() { + engine::autograd::set_grad_enabled(prev); + } + false + } +} + +/// Return a context manager that disables gradient recording. +/// +/// Inside the block, operation results do not require gradients, no autograd +/// nodes are recorded, and no operands are saved for backward — mirroring +/// `torch.no_grad()`. Tensors can still opt in explicitly via +/// `requires_grad_(True)`. +#[pyfunction] +fn no_grad() -> GradMode { + GradMode { + target: false, + previous: None, + } +} + +/// Return a context manager that re-enables gradient recording, e.g. inside +/// an outer `no_grad()` block. +#[pyfunction] +fn enable_grad() -> GradMode { + GradMode { + target: true, + previous: None, + } +} + +/// Query whether gradient recording is currently enabled on this thread. +#[pyfunction] +fn is_grad_enabled() -> bool { + engine::autograd::is_grad_enabled() +} + +/// Set the gradient recording mode, returning the previous mode. +#[pyfunction] +fn set_grad_enabled(enabled: bool) -> bool { + engine::autograd::set_grad_enabled(enabled) +} + #[pyfunction] fn get_gradient(tensor: &PyTensor) -> PyResult> { Ok(engine::autograd::get_gradient(tensor.tensor()).map(PyTensor::from_tensor)) @@ -143,6 +210,10 @@ mod tests { "clear_autograd_graph", "is_autograd_graph_consumed", "mark_autograd_graph_consumed", + "no_grad", + "enable_grad", + "is_grad_enabled", + "set_grad_enabled", "get_default_dtype", "set_default_dtype", "manual_seed", diff --git a/bindings/src/nn.rs b/bindings/src/nn.rs index e55366b2..5602cedb 100644 --- a/bindings/src/nn.rs +++ b/bindings/src/nn.rs @@ -4,5 +4,7 @@ // This source code is licensed under the Apache-style license found in the // LICENSE file in the root directory of this source tree. -include!("nn/module.rs"); -include!("nn/layers.rs"); +#[path = "nn/module.rs"] +mod module; + +pub use self::module::*; diff --git a/bindings/src/nn/layers.rs b/bindings/src/nn/layers.rs index 46b11c90..0eaa5362 100644 --- a/bindings/src/nn/layers.rs +++ b/bindings/src/nn/layers.rs @@ -1,885 +1,886 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -#[pymethods] -impl PyReLU { - /// Create a new ReLU layer - #[new] - fn new() -> PyClassInitializer { - let relu = ReLU::new(); - PyClassInitializer::from(PyModule::from_relu(relu)).add_subclass(Self) - } -} - -/// Sigmoid activation layer -#[pyclass(name = "Sigmoid", extends = PyModule)] -pub struct PySigmoid; - -#[pymethods] -impl PySigmoid { - /// Create a new Sigmoid layer - #[new] - fn new() -> PyClassInitializer { - let sigmoid = Sigmoid::new(); - PyClassInitializer::from(PyModule::from_sigmoid(sigmoid)).add_subclass(Self) - } -} - -/// Tanh activation layer -#[pyclass(name = "Tanh", extends = PyModule)] -pub struct PyTanh; - -#[pymethods] -impl PyTanh { - /// Create a new Tanh layer - #[new] - fn new() -> PyClassInitializer { - let tanh = Tanh::new(); - PyClassInitializer::from(PyModule::from_tanh(tanh)).add_subclass(Self) - } -} - -/// Softmax activation layer -#[pyclass(name = "Softmax", extends = PyModule)] -pub struct PySoftmax; - -#[pymethods] -impl PySoftmax { - /// Create a new Softmax layer - #[new] - #[pyo3(signature = (dim=None))] - fn new(dim: Option) -> PyClassInitializer { - let softmax = Softmax::new(dim); - PyClassInitializer::from(PyModule::from_softmax(softmax)).add_subclass(Self) - } - - /// Get the dimension along which softmax is computed - #[getter] - fn dim(slf: PyRef) -> PyResult> { - let module = slf.as_ref(); - if let ModuleType::Softmax(layer) = &module.inner { - Ok(layer.dim()) - } else { - Err(PyErr::new::( - "Invalid layer type", - )) - } - } -} - -/// LeakyReLU activation layer -#[pyclass(name = "LeakyReLU", extends = PyModule)] -pub struct PyLeakyReLU; - -#[pymethods] -impl PyLeakyReLU { - /// Create a new LeakyReLU layer - #[new] - #[pyo3(signature = (negative_slope=None))] - fn new(negative_slope: Option) -> PyClassInitializer { - let negative_slope = negative_slope.unwrap_or(0.01); - let leaky_relu = LeakyReLU::new(Some(negative_slope)); - PyClassInitializer::from(PyModule::from_leaky_relu(leaky_relu)).add_subclass(Self) - } - - /// Get the negative slope parameter - #[getter] - fn negative_slope(slf: PyRef) -> PyResult { - let module = slf.as_ref(); - if let ModuleType::LeakyReLU(layer) = &module.inner { - Ok(layer.negative_slope()) - } else { - Err(PyErr::new::( - "Invalid layer type", - )) - } - } -} - -/// ELU activation layer -#[pyclass(name = "ELU", extends = PyModule)] -pub struct PyELU; - -#[pymethods] -impl PyELU { - /// Create a new ELU layer - #[new] - #[pyo3(signature = (alpha=None))] - fn new(alpha: Option) -> PyClassInitializer { - let alpha = alpha.unwrap_or(1.0); - let elu = ELU::new(Some(alpha)); - PyClassInitializer::from(PyModule::from_elu(elu)).add_subclass(Self) - } - - /// Get the alpha parameter - #[getter] - fn alpha(slf: PyRef) -> PyResult { - let module = slf.as_ref(); - if let ModuleType::Elu(layer) = &module.inner { - Ok(layer.alpha()) - } else { - Err(PyErr::new::( - "Invalid layer type", - )) - } - } -} - -/// GELU activation layer -#[pyclass(name = "GELU", extends = PyModule)] -pub struct PyGELU; - -#[pymethods] -impl PyGELU { - /// Create a new GELU layer - #[new] - fn new() -> PyClassInitializer { - let gelu = GELU::new(); - PyClassInitializer::from(PyModule::from_gelu(gelu)).add_subclass(Self) - } -} - -/// Dropout layer -#[pyclass(name = "Dropout", extends = PyModule)] -pub struct PyDropout; - -#[pymethods] -impl PyDropout { - /// Create a new Dropout layer - #[new] - #[pyo3(signature = (p=None))] - fn new(p: Option) -> PyResult> { - let p = p.unwrap_or(0.5); - let dropout = Dropout::new(Some(p)).map_err(_convert_error)?; - Ok(PyClassInitializer::from(PyModule::from_dropout(dropout)).add_subclass(Self)) - } - - /// Get the dropout probability - #[getter] - fn p(slf: PyRef) -> PyResult { - let module = slf.as_ref(); - if let ModuleType::Dropout(layer) = &module.inner { - Ok(layer.p()) - } else { - Err(PyErr::new::( - "Invalid layer type", - )) - } - } -} - -/// 2D Dropout layer -#[pyclass(name = "Dropout2d", extends = PyModule)] -pub struct PyDropout2d; - -#[pymethods] -impl PyDropout2d { - /// Create a new Dropout2d layer - #[new] - #[pyo3(signature = (p=None))] - fn new(p: Option) -> PyResult> { - let p = p.unwrap_or(0.5); - let dropout = Dropout2d::new(Some(p)).map_err(_convert_error)?; - Ok(PyClassInitializer::from(PyModule::from_dropout2d(dropout)).add_subclass(Self)) - } - - /// Get the dropout probability - #[getter] - fn p(slf: PyRef) -> PyResult { - let module = slf.as_ref(); - if let ModuleType::Dropout2d(layer) = &module.inner { - Ok(layer.p()) - } else { - Err(PyErr::new::( - "Invalid layer type", - )) - } - } -} - -/// Conv2d layer -#[pyclass(name = "Conv2d", extends = PyModule)] -pub struct PyConv2d; - -#[pymethods] -impl PyConv2d { - /// Create a new Conv2d layer - #[new] - #[pyo3(signature = ( - in_channels, - out_channels, - kernel_size, - stride=None, - padding=None, - bias=None, - device=None, - dtype=None - ))] - #[allow(clippy::too_many_arguments)] - fn new( - in_channels: usize, - out_channels: usize, - kernel_size: &Bound, - stride: Option<&Bound>, - padding: Option<&Bound>, - bias: Option, - device: Option<&PyDevice>, - dtype: Option<&str>, - ) -> PyResult> { - let kernel_size = parse_tuple2(kernel_size)?; - let stride = match stride { - Some(s) => parse_tuple2(s)?, - None => (1, 1), - }; - let padding = match padding { - Some(p) => parse_tuple2(p)?, - None => (0, 0), - }; - let bias = bias.unwrap_or(true); - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let dtype = dtype::resolve_dtype_arg(dtype)?; - - let conv2d = Conv2d::new( - in_channels, - out_channels, - kernel_size, - Some(stride), - Some(padding), - bias, - device, - dtype, - ) - .map_err(_convert_error)?; - - Ok(PyClassInitializer::from(PyModule::from_conv2d(conv2d)).add_subclass(Self)) - } - - /// Get input channels count - #[getter] - fn in_channels(slf: PyRef) -> PyResult { - let module = slf.as_ref(); - if let ModuleType::Conv2d(layer) = &module.inner { - Ok(layer.in_channels()) - } else { - Err(PyErr::new::( - "Invalid layer type", - )) - } - } - - /// Get output channels count - #[getter] - fn out_channels(slf: PyRef) -> PyResult { - let module = slf.as_ref(); - if let ModuleType::Conv2d(layer) = &module.inner { - Ok(layer.out_channels()) - } else { - Err(PyErr::new::( - "Invalid layer type", - )) - } - } - - /// Get kernel size - #[getter] - fn kernel_size(slf: PyRef) -> PyResult<(usize, usize)> { - let module = slf.as_ref(); - if let ModuleType::Conv2d(layer) = &module.inner { - Ok(layer.kernel_size()) - } else { - Err(PyErr::new::( - "Invalid layer type", - )) - } - } -} - -/// BatchNorm1d layer -#[pyclass(name = "BatchNorm1d", extends = PyModule)] -pub struct PyBatchNorm1d; - -#[pymethods] -impl PyBatchNorm1d { - /// Create a new BatchNorm1d layer - #[new] - #[pyo3(signature = (num_features, eps=None, momentum=None, affine=None, device=None, dtype=None))] - fn new( - num_features: usize, - eps: Option, - momentum: Option, - affine: Option, - device: Option<&PyDevice>, - dtype: Option<&str>, - ) -> PyResult> { - let eps = eps.unwrap_or(1e-5); - let momentum = momentum.unwrap_or(0.1); - let _affine = affine.unwrap_or(true); - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let dtype = dtype::resolve_dtype_arg(dtype)?; - - let batch_norm = BatchNorm1d::new(num_features, Some(eps), Some(momentum), device, dtype) - .map_err(_convert_error)?; - - Ok(PyClassInitializer::from(PyModule::from_batch_norm1d(batch_norm)).add_subclass(Self)) - } - - /// Get number of features - #[getter] - fn num_features(slf: PyRef) -> PyResult { - let module = slf.as_ref(); - if let ModuleType::BatchNorm1d(layer) = &module.inner { - Ok(layer.num_features()) - } else { - Err(PyErr::new::( - "Invalid layer type", - )) - } - } -} - -/// BatchNorm2d layer -#[pyclass(name = "BatchNorm2d", extends = PyModule)] -pub struct PyBatchNorm2d; - -#[pymethods] -impl PyBatchNorm2d { - /// Create a new BatchNorm2d layer - #[new] - #[pyo3(signature = (num_features, eps=None, momentum=None, affine=None, device=None, dtype=None))] - fn new( - num_features: usize, - eps: Option, - momentum: Option, - affine: Option, - device: Option<&PyDevice>, - dtype: Option<&str>, - ) -> PyResult> { - let eps = eps.unwrap_or(1e-5); - let momentum = momentum.unwrap_or(0.1); - let _affine = affine.unwrap_or(true); - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let dtype = dtype::resolve_dtype_arg(dtype)?; - - let batch_norm = BatchNorm2d::new(num_features, Some(eps), Some(momentum), device, dtype) - .map_err(_convert_error)?; - - Ok(PyClassInitializer::from(PyModule::from_batch_norm2d(batch_norm)).add_subclass(Self)) - } - - /// Get number of features - #[getter] - fn num_features(slf: PyRef) -> PyResult { - let module = slf.as_ref(); - if let ModuleType::BatchNorm2d(layer) = &module.inner { - Ok(layer.num_features()) - } else { - Err(PyErr::new::( - "Invalid layer type", - )) - } - } -} - -/// Sequential container for layers -#[pyclass(name = "Sequential", extends = PyModule)] -pub struct PySequential; - -#[pymethods] -impl PySequential { - /// Create a new Sequential container - #[new] - #[pyo3(signature = (layers=None))] - fn new(layers: Option>>) -> PyResult> { - let sequential = if let Some(layers) = layers { - let mut layer_objects = Vec::with_capacity(layers.len()); - for layer in layers { - layer_objects.push(layer.to_layer()?); - } - Sequential::from_layers(layer_objects) - } else { - Sequential::new() - }; - - Ok(PyClassInitializer::from(PyModule::from_sequential(sequential)).add_subclass(Self)) - } - - /// Add a layer to the sequential container - fn add_module(mut slf: PyRefMut, _name: &str, module: PyRef) -> PyResult<()> { - if !matches!(slf.as_ref().inner, ModuleType::Sequential(_)) { - return Err(PyErr::new::( - "Invalid layer type", - )); - } - - let layer = module.to_layer()?; - - if let ModuleType::Sequential(seq) = &mut slf.as_mut().inner { - seq.add_layer(layer); - Ok(()) - } else { - Err(PyErr::new::( - "Invalid layer type", - )) - } - } -} - -/// Helper function to parse data type string -fn parse_tuple2(obj: &Bound) -> PyResult<(usize, usize)> { - if let Ok(val) = obj.extract::() { - Ok((val, val)) - } else { - obj.extract::<(usize, usize)>() - } -} - -/// MSE Loss function -#[pyclass(name = "MSELoss")] -pub struct PyMSELoss { - inner: MSELoss, -} - -#[pymethods] -impl PyMSELoss { - /// Create a new MSE loss - #[new] - #[pyo3(signature = (reduction=None))] - fn new(reduction: Option<&str>) -> Self { - let reduction = reduction.unwrap_or("mean"); - Self { - inner: MSELoss::new(reduction), - } - } - - /// Compute the MSE loss - fn forward(&self, predictions: &Bound, targets: &Bound) -> PyResult { - let predictions = borrow_tensor(predictions)?; - let targets = borrow_tensor(targets)?; - let result = self - .inner - .forward(predictions.tensor(), targets.tensor()) - .map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) - } - - #[pyo3(name = "__call__")] - fn call(&self, predictions: &Bound, targets: &Bound) -> PyResult { - self.forward(predictions, targets) - } - - /// Get the reduction mode - #[getter] - fn reduction(&self) -> &str { - self.inner.reduction() - } - - /// String representation - fn __repr__(&self) -> String { - format!("MSELoss(reduction='{}')", self.inner.reduction()) - } -} - -/// MAE Loss function -#[pyclass(name = "MAELoss")] -pub struct PyMAELoss { - inner: MAELoss, -} - -#[pymethods] -impl PyMAELoss { - /// Create a new MAE loss - #[new] - #[pyo3(signature = (reduction=None))] - fn new(reduction: Option<&str>) -> Self { - let reduction = reduction.unwrap_or("mean"); - Self { - inner: MAELoss::new(reduction), - } - } - - /// Compute the MAE loss - fn forward(&self, predictions: &Bound, targets: &Bound) -> PyResult { - let predictions = borrow_tensor(predictions)?; - let targets = borrow_tensor(targets)?; - let result = self - .inner - .forward(predictions.tensor(), targets.tensor()) - .map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) - } - - #[pyo3(name = "__call__")] - fn call(&self, predictions: &Bound, targets: &Bound) -> PyResult { - self.forward(predictions, targets) - } - - /// Get the reduction mode - #[getter] - fn reduction(&self) -> &str { - self.inner.reduction() - } - - /// String representation - fn __repr__(&self) -> String { - format!("MAELoss(reduction='{}')", self.inner.reduction()) - } -} - -/// Huber Loss function -#[pyclass(name = "HuberLoss")] -pub struct PyHuberLoss { - inner: HuberLoss, -} - -#[pymethods] -impl PyHuberLoss { - /// Create a new Huber loss - #[new] - #[pyo3(signature = (delta=None, reduction=None))] - fn new(delta: Option, reduction: Option<&str>) -> Self { - let delta = delta.unwrap_or(1.0); - let reduction = reduction.unwrap_or("mean"); - Self { - inner: HuberLoss::new(delta, reduction), - } - } - - /// Compute the Huber loss - fn forward(&self, predictions: &Bound, targets: &Bound) -> PyResult { - let predictions = borrow_tensor(predictions)?; - let targets = borrow_tensor(targets)?; - let result = self - .inner - .forward(predictions.tensor(), targets.tensor()) - .map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) - } - - #[pyo3(name = "__call__")] - fn call(&self, predictions: &Bound, targets: &Bound) -> PyResult { - self.forward(predictions, targets) - } - - /// Get the delta parameter - #[getter] - fn delta(&self) -> f64 { - self.inner.delta() - } - - /// Get the reduction mode - #[getter] - fn reduction(&self) -> &str { - self.inner.reduction() - } - - /// String representation - fn __repr__(&self) -> String { - format!( - "HuberLoss(delta={}, reduction='{}')", - self.inner.delta(), - self.inner.reduction() - ) - } -} - -/// Smooth L1 Loss function -#[pyclass(name = "SmoothL1Loss")] -pub struct PySmoothL1Loss { - inner: SmoothL1Loss, -} - -#[pymethods] -impl PySmoothL1Loss { - /// Create a new Smooth L1 loss - #[new] - #[pyo3(signature = (reduction=None))] - fn new(reduction: Option<&str>) -> Self { - let reduction = reduction.unwrap_or("mean"); - Self { - inner: SmoothL1Loss::new(reduction), - } - } - - /// Compute the Smooth L1 loss - fn forward(&self, predictions: &Bound, targets: &Bound) -> PyResult { - let predictions = borrow_tensor(predictions)?; - let targets = borrow_tensor(targets)?; - let result = self - .inner - .forward(predictions.tensor(), targets.tensor()) - .map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) - } - - #[pyo3(name = "__call__")] - fn call(&self, predictions: &Bound, targets: &Bound) -> PyResult { - self.forward(predictions, targets) - } - - /// Get the reduction mode - #[getter] - fn reduction(&self) -> &str { - self.inner.reduction() - } - - /// String representation - fn __repr__(&self) -> String { - format!("SmoothL1Loss(reduction='{}')", self.inner.reduction()) - } -} - -/// Log-cosh Loss function -#[pyclass(name = "LogCoshLoss")] -pub struct PyLogCoshLoss { - inner: LogCoshLoss, -} - -#[pymethods] -impl PyLogCoshLoss { - /// Create a new Log-cosh loss - #[new] - #[pyo3(signature = (reduction=None))] - fn new(reduction: Option<&str>) -> Self { - let reduction = reduction.unwrap_or("mean"); - Self { - inner: LogCoshLoss::new(reduction), - } - } - - /// Compute the Log-cosh loss - fn forward(&self, predictions: &Bound, targets: &Bound) -> PyResult { - let predictions = borrow_tensor(predictions)?; - let targets = borrow_tensor(targets)?; - let result = self - .inner - .forward(predictions.tensor(), targets.tensor()) - .map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) - } - - #[pyo3(name = "__call__")] - fn call(&self, predictions: &Bound, targets: &Bound) -> PyResult { - self.forward(predictions, targets) - } - - /// Get the reduction mode - #[getter] - fn reduction(&self) -> &str { - self.inner.reduction() - } - - /// String representation - fn __repr__(&self) -> String { - format!("LogCoshLoss(reduction='{}')", self.inner.reduction()) - } -} - -/// Cross Entropy Loss function -#[pyclass(name = "CrossEntropyLoss")] -pub struct PyCrossEntropyLoss { - inner: CrossEntropyLoss, -} - -#[pymethods] -impl PyCrossEntropyLoss { - /// Create a new Cross Entropy loss - #[new] - #[pyo3(signature = (reduction=None))] - fn new(reduction: Option<&str>) -> Self { - let reduction = reduction.unwrap_or("mean"); - Self { - inner: CrossEntropyLoss::new(reduction), - } - } - - /// Compute the Cross Entropy loss - fn forward(&self, predictions: &Bound, targets: &Bound) -> PyResult { - let predictions = borrow_tensor(predictions)?; - let targets = borrow_tensor(targets)?; - let result = self - .inner - .forward(predictions.tensor(), targets.tensor()) - .map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) - } - - #[pyo3(name = "__call__")] - fn call(&self, predictions: &Bound, targets: &Bound) -> PyResult { - self.forward(predictions, targets) - } - - /// Get the reduction mode - #[getter] - fn reduction(&self) -> &str { - self.inner.reduction() - } - - /// String representation - fn __repr__(&self) -> String { - format!("CrossEntropyLoss(reduction='{}')", self.inner.reduction()) - } -} - -/// Binary Cross Entropy Loss function -#[pyclass(name = "BCELoss")] -pub struct PyBCELoss { - inner: BCELoss, -} - -#[pymethods] -impl PyBCELoss { - /// Create a new BCE loss - #[new] - #[pyo3(signature = (reduction=None))] - fn new(reduction: Option<&str>) -> Self { - let reduction = reduction.unwrap_or("mean"); - Self { - inner: BCELoss::new(reduction), - } - } - - /// Compute the BCE loss - fn forward(&self, predictions: &Bound, targets: &Bound) -> PyResult { - let predictions = borrow_tensor(predictions)?; - let targets = borrow_tensor(targets)?; - let result = self - .inner - .forward(predictions.tensor(), targets.tensor()) - .map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) - } - - #[pyo3(name = "__call__")] - fn call(&self, predictions: &Bound, targets: &Bound) -> PyResult { - self.forward(predictions, targets) - } - - /// Get the reduction mode - #[getter] - fn reduction(&self) -> &str { - self.inner.reduction() - } - - /// String representation - fn __repr__(&self) -> String { - format!("BCELoss(reduction='{}')", self.inner.reduction()) - } -} - -/// Focal Loss function -#[pyclass(name = "FocalLoss")] -pub struct PyFocalLoss { - inner: FocalLoss, -} - -#[pymethods] -impl PyFocalLoss { - /// Create a new Focal loss - #[new] - #[pyo3(signature = (alpha=None, gamma=None, reduction=None))] - fn new(alpha: Option, gamma: Option, reduction: Option<&str>) -> Self { - let alpha = alpha.unwrap_or(0.25); - let gamma = gamma.unwrap_or(2.0); - let reduction = reduction.unwrap_or("mean"); - Self { - inner: FocalLoss::new(alpha, gamma, reduction), - } - } - - /// Compute the Focal loss - fn forward(&self, predictions: &Bound, targets: &Bound) -> PyResult { - let predictions = borrow_tensor(predictions)?; - let targets = borrow_tensor(targets)?; - let result = self - .inner - .forward(predictions.tensor(), targets.tensor()) - .map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) - } - - #[pyo3(name = "__call__")] - fn call(&self, predictions: &Bound, targets: &Bound) -> PyResult { - self.forward(predictions, targets) - } - - /// Get the alpha parameter - #[getter] - fn alpha(&self) -> f64 { - self.inner.alpha() - } - - /// Get the gamma parameter - #[getter] - fn gamma(&self) -> f64 { - self.inner.gamma() - } - - /// Get the reduction mode - #[getter] - fn reduction(&self) -> &str { - self.inner.reduction() - } - - /// String representation - fn __repr__(&self) -> String { - format!( - "FocalLoss(alpha={}, gamma={}, reduction='{}')", - self.inner.alpha(), - self.inner.gamma(), - self.inner.reduction() - ) - } -} - -/// Register neural network module with Python -pub fn register_nn_module(py: Python, parent_module: &Bound) -> PyResult<()> { - let nn_module = Pyo3Module::new(py, "nn")?; - - // Add layer classes - nn_module.add_class::()?; - nn_module.add_class::()?; - nn_module.add_class::()?; - nn_module.add_class::()?; - nn_module.add_class::()?; - nn_module.add_class::()?; - nn_module.add_class::()?; - nn_module.add_class::()?; - nn_module.add_class::()?; - nn_module.add_class::()?; - nn_module.add_class::()?; - nn_module.add_class::()?; - nn_module.add_class::()?; - nn_module.add_class::()?; - nn_module.add_class::()?; - - // Add functional APIs - nn_module.add_function(wrap_pyfunction!(dense_layer, &nn_module)?)?; - nn_module.add_function(wrap_pyfunction!(conv2d, &nn_module)?)?; - nn_module.add_function(wrap_pyfunction!(batch_norm, &nn_module)?)?; - nn_module.add_function(wrap_pyfunction!(cross_entropy, &nn_module)?)?; - nn_module.add_function(wrap_pyfunction!(dropout_functional, &nn_module)?)?; - nn_module.add_function(wrap_pyfunction!(dropout2d_functional, &nn_module)?)?; - nn_module.add_function(wrap_pyfunction!(mse_loss_functional, &nn_module)?)?; - nn_module.add_function(wrap_pyfunction!(smooth_l1_loss_functional, &nn_module)?)?; - nn_module.add_function(wrap_pyfunction!(log_cosh_loss_functional, &nn_module)?)?; - nn_module.add_function(wrap_pyfunction!( - binary_cross_entropy_functional, - &nn_module - )?)?; - - // Add loss function classes - nn_module.add_class::()?; - nn_module.add_class::()?; - nn_module.add_class::()?; - nn_module.add_class::()?; - nn_module.add_class::()?; - nn_module.add_class::()?; - nn_module.add_class::()?; - nn_module.add_class::()?; - - parent_module.add_submodule(&nn_module)?; - Ok(()) -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +#[pymethods] +impl PyReLU { + /// Create a new ReLU layer + #[new] + fn new() -> PyClassInitializer { + let relu = ReLU::new(); + PyClassInitializer::from(PyModule::from_relu(relu)).add_subclass(Self) + } +} + +/// Sigmoid activation layer +#[pyclass(name = "Sigmoid", extends = PyModule)] +pub struct PySigmoid; + +#[pymethods] +impl PySigmoid { + /// Create a new Sigmoid layer + #[new] + fn new() -> PyClassInitializer { + let sigmoid = Sigmoid::new(); + PyClassInitializer::from(PyModule::from_sigmoid(sigmoid)).add_subclass(Self) + } +} + +/// Tanh activation layer +#[pyclass(name = "Tanh", extends = PyModule)] +pub struct PyTanh; + +#[pymethods] +impl PyTanh { + /// Create a new Tanh layer + #[new] + fn new() -> PyClassInitializer { + let tanh = Tanh::new(); + PyClassInitializer::from(PyModule::from_tanh(tanh)).add_subclass(Self) + } +} + +/// Softmax activation layer +#[pyclass(name = "Softmax", extends = PyModule)] +pub struct PySoftmax; + +#[pymethods] +impl PySoftmax { + /// Create a new Softmax layer + #[new] + #[pyo3(signature = (dim=None))] + fn new(dim: Option) -> PyClassInitializer { + let softmax = Softmax::new(dim); + PyClassInitializer::from(PyModule::from_softmax(softmax)).add_subclass(Self) + } + + /// Get the dimension along which softmax is computed + #[getter] + fn dim(slf: PyRef) -> PyResult> { + let module = slf.as_ref(); + if let ModuleType::Softmax(layer) = &module.inner { + Ok(layer.dim()) + } else { + Err(PyErr::new::( + "Invalid layer type", + )) + } + } +} + +/// LeakyReLU activation layer +#[pyclass(name = "LeakyReLU", extends = PyModule)] +pub struct PyLeakyReLU; + +#[pymethods] +impl PyLeakyReLU { + /// Create a new LeakyReLU layer + #[new] + #[pyo3(signature = (negative_slope=None))] + fn new(negative_slope: Option) -> PyClassInitializer { + let negative_slope = negative_slope.unwrap_or(0.01); + let leaky_relu = LeakyReLU::new(Some(negative_slope)); + PyClassInitializer::from(PyModule::from_leaky_relu(leaky_relu)).add_subclass(Self) + } + + /// Get the negative slope parameter + #[getter] + fn negative_slope(slf: PyRef) -> PyResult { + let module = slf.as_ref(); + if let ModuleType::LeakyReLU(layer) = &module.inner { + Ok(layer.negative_slope()) + } else { + Err(PyErr::new::( + "Invalid layer type", + )) + } + } +} + +/// ELU activation layer +#[pyclass(name = "ELU", extends = PyModule)] +pub struct PyELU; + +#[pymethods] +impl PyELU { + /// Create a new ELU layer + #[new] + #[pyo3(signature = (alpha=None))] + fn new(alpha: Option) -> PyClassInitializer { + let alpha = alpha.unwrap_or(1.0); + let elu = ELU::new(Some(alpha)); + PyClassInitializer::from(PyModule::from_elu(elu)).add_subclass(Self) + } + + /// Get the alpha parameter + #[getter] + fn alpha(slf: PyRef) -> PyResult { + let module = slf.as_ref(); + if let ModuleType::Elu(layer) = &module.inner { + Ok(layer.alpha()) + } else { + Err(PyErr::new::( + "Invalid layer type", + )) + } + } +} + +/// GELU activation layer +#[pyclass(name = "GELU", extends = PyModule)] +pub struct PyGELU; + +#[pymethods] +impl PyGELU { + /// Create a new GELU layer + #[new] + fn new() -> PyClassInitializer { + let gelu = GELU::new(); + PyClassInitializer::from(PyModule::from_gelu(gelu)).add_subclass(Self) + } +} + +/// Dropout layer +#[pyclass(name = "Dropout", extends = PyModule)] +pub struct PyDropout; + +#[pymethods] +impl PyDropout { + /// Create a new Dropout layer + #[new] + #[pyo3(signature = (p=None))] + fn new(p: Option) -> PyResult> { + let p = p.unwrap_or(0.5); + let dropout = Dropout::new(Some(p)).map_err(_convert_error)?; + Ok(PyClassInitializer::from(PyModule::from_dropout(dropout)).add_subclass(Self)) + } + + /// Get the dropout probability + #[getter] + fn p(slf: PyRef) -> PyResult { + let module = slf.as_ref(); + if let ModuleType::Dropout(layer) = &module.inner { + Ok(layer.p()) + } else { + Err(PyErr::new::( + "Invalid layer type", + )) + } + } +} + +/// 2D Dropout layer +#[pyclass(name = "Dropout2d", extends = PyModule)] +pub struct PyDropout2d; + +#[pymethods] +impl PyDropout2d { + /// Create a new Dropout2d layer + #[new] + #[pyo3(signature = (p=None))] + fn new(p: Option) -> PyResult> { + let p = p.unwrap_or(0.5); + let dropout = Dropout2d::new(Some(p)).map_err(_convert_error)?; + Ok(PyClassInitializer::from(PyModule::from_dropout2d(dropout)).add_subclass(Self)) + } + + /// Get the dropout probability + #[getter] + fn p(slf: PyRef) -> PyResult { + let module = slf.as_ref(); + if let ModuleType::Dropout2d(layer) = &module.inner { + Ok(layer.p()) + } else { + Err(PyErr::new::( + "Invalid layer type", + )) + } + } +} + +/// Conv2d layer +#[pyclass(name = "Conv2d", extends = PyModule)] +pub struct PyConv2d; + +#[pymethods] +impl PyConv2d { + /// Create a new Conv2d layer + #[new] + #[pyo3(signature = ( + in_channels, + out_channels, + kernel_size, + stride=None, + padding=None, + bias=None, + device=None, + dtype=None + ))] + #[allow(clippy::too_many_arguments)] + fn new( + in_channels: usize, + out_channels: usize, + kernel_size: &Bound, + stride: Option<&Bound>, + padding: Option<&Bound>, + bias: Option, + device: Option<&PyDevice>, + dtype: Option<&str>, + ) -> PyResult> { + let kernel_size = parse_tuple2(kernel_size)?; + let stride = match stride { + Some(s) => parse_tuple2(s)?, + None => (1, 1), + }; + let padding = match padding { + Some(p) => parse_tuple2(p)?, + None => (0, 0), + }; + let bias = bias.unwrap_or(true); + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let dtype = dtype::resolve_dtype_arg(dtype)?; + + let conv2d = Conv2d::new( + in_channels, + out_channels, + kernel_size, + Some(stride), + Some(padding), + bias, + device, + dtype, + ) + .map_err(_convert_error)?; + + Ok(PyClassInitializer::from(PyModule::from_conv2d(conv2d)).add_subclass(Self)) + } + + /// Get input channels count + #[getter] + fn in_channels(slf: PyRef) -> PyResult { + let module = slf.as_ref(); + if let ModuleType::Conv2d(layer) = &module.inner { + Ok(layer.in_channels()) + } else { + Err(PyErr::new::( + "Invalid layer type", + )) + } + } + + /// Get output channels count + #[getter] + fn out_channels(slf: PyRef) -> PyResult { + let module = slf.as_ref(); + if let ModuleType::Conv2d(layer) = &module.inner { + Ok(layer.out_channels()) + } else { + Err(PyErr::new::( + "Invalid layer type", + )) + } + } + + /// Get kernel size + #[getter] + fn kernel_size(slf: PyRef) -> PyResult<(usize, usize)> { + let module = slf.as_ref(); + if let ModuleType::Conv2d(layer) = &module.inner { + Ok(layer.kernel_size()) + } else { + Err(PyErr::new::( + "Invalid layer type", + )) + } + } +} + +/// BatchNorm1d layer +#[pyclass(name = "BatchNorm1d", extends = PyModule)] +pub struct PyBatchNorm1d; + +#[pymethods] +impl PyBatchNorm1d { + /// Create a new BatchNorm1d layer + #[new] + #[pyo3(signature = (num_features, eps=None, momentum=None, affine=None, device=None, dtype=None))] + fn new( + num_features: usize, + eps: Option, + momentum: Option, + affine: Option, + device: Option<&PyDevice>, + dtype: Option<&str>, + ) -> PyResult> { + let eps = eps.unwrap_or(1e-5); + let momentum = momentum.unwrap_or(0.1); + let _affine = affine.unwrap_or(true); + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let dtype = dtype::resolve_dtype_arg(dtype)?; + + let batch_norm = BatchNorm1d::new(num_features, Some(eps), Some(momentum), device, dtype) + .map_err(_convert_error)?; + + Ok(PyClassInitializer::from(PyModule::from_batch_norm1d(batch_norm)).add_subclass(Self)) + } + + /// Get number of features + #[getter] + fn num_features(slf: PyRef) -> PyResult { + let module = slf.as_ref(); + if let ModuleType::BatchNorm1d(layer) = &module.inner { + Ok(layer.num_features()) + } else { + Err(PyErr::new::( + "Invalid layer type", + )) + } + } +} + +/// BatchNorm2d layer +#[pyclass(name = "BatchNorm2d", extends = PyModule)] +pub struct PyBatchNorm2d; + +#[pymethods] +impl PyBatchNorm2d { + /// Create a new BatchNorm2d layer + #[new] + #[pyo3(signature = (num_features, eps=None, momentum=None, affine=None, device=None, dtype=None))] + fn new( + num_features: usize, + eps: Option, + momentum: Option, + affine: Option, + device: Option<&PyDevice>, + dtype: Option<&str>, + ) -> PyResult> { + let eps = eps.unwrap_or(1e-5); + let momentum = momentum.unwrap_or(0.1); + let _affine = affine.unwrap_or(true); + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let dtype = dtype::resolve_dtype_arg(dtype)?; + + let batch_norm = BatchNorm2d::new(num_features, Some(eps), Some(momentum), device, dtype) + .map_err(_convert_error)?; + + Ok(PyClassInitializer::from(PyModule::from_batch_norm2d(batch_norm)).add_subclass(Self)) + } + + /// Get number of features + #[getter] + fn num_features(slf: PyRef) -> PyResult { + let module = slf.as_ref(); + if let ModuleType::BatchNorm2d(layer) = &module.inner { + Ok(layer.num_features()) + } else { + Err(PyErr::new::( + "Invalid layer type", + )) + } + } +} + +/// Sequential container for layers +#[pyclass(name = "Sequential", extends = PyModule)] +pub struct PySequential; + +#[pymethods] +impl PySequential { + /// Create a new Sequential container + #[new] + #[pyo3(signature = (layers=None))] + fn new(layers: Option>>) -> PyResult> { + let sequential = if let Some(layers) = layers { + let mut layer_objects = Vec::with_capacity(layers.len()); + for layer in layers { + layer_objects.push(layer.to_layer()?); + } + Sequential::from_layers(layer_objects) + } else { + Sequential::new() + }; + + Ok(PyClassInitializer::from(PyModule::from_sequential(sequential)).add_subclass(Self)) + } + + /// Add a layer to the sequential container + fn add_module(mut slf: PyRefMut, _name: &str, module: PyRef) -> PyResult<()> { + if !matches!(slf.as_ref().inner, ModuleType::Sequential(_)) { + return Err(PyErr::new::( + "Invalid layer type", + )); + } + + let layer = module.to_layer()?; + + if let ModuleType::Sequential(seq) = &mut slf.as_mut().inner { + seq.add_layer(layer); + Ok(()) + } else { + Err(PyErr::new::( + "Invalid layer type", + )) + } + } +} + +/// Helper function to parse data type string +fn parse_tuple2(obj: &Bound) -> PyResult<(usize, usize)> { + if let Ok(val) = obj.extract::() { + Ok((val, val)) + } else { + obj.extract::<(usize, usize)>() + } +} + +/// MSE Loss function +#[pyclass(name = "MSELoss")] +pub struct PyMSELoss { + inner: MSELoss, +} + +#[pymethods] +impl PyMSELoss { + /// Create a new MSE loss + #[new] + #[pyo3(signature = (reduction=None))] + fn new(reduction: Option<&str>) -> Self { + let reduction = reduction.unwrap_or("mean"); + Self { + inner: MSELoss::new(reduction), + } + } + + /// Compute the MSE loss + fn forward(&self, predictions: &Bound, targets: &Bound) -> PyResult { + let predictions = borrow_tensor(predictions)?; + let targets = borrow_tensor(targets)?; + let result = self + .inner + .forward(predictions.tensor(), targets.tensor()) + .map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) + } + + #[pyo3(name = "__call__")] + fn call(&self, predictions: &Bound, targets: &Bound) -> PyResult { + self.forward(predictions, targets) + } + + /// Get the reduction mode + #[getter] + fn reduction(&self) -> &str { + self.inner.reduction() + } + + /// String representation + fn __repr__(&self) -> String { + format!("MSELoss(reduction='{}')", self.inner.reduction()) + } +} + +/// MAE Loss function +#[pyclass(name = "MAELoss")] +pub struct PyMAELoss { + inner: MAELoss, +} + +#[pymethods] +impl PyMAELoss { + /// Create a new MAE loss + #[new] + #[pyo3(signature = (reduction=None))] + fn new(reduction: Option<&str>) -> Self { + let reduction = reduction.unwrap_or("mean"); + Self { + inner: MAELoss::new(reduction), + } + } + + /// Compute the MAE loss + fn forward(&self, predictions: &Bound, targets: &Bound) -> PyResult { + let predictions = borrow_tensor(predictions)?; + let targets = borrow_tensor(targets)?; + let result = self + .inner + .forward(predictions.tensor(), targets.tensor()) + .map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) + } + + #[pyo3(name = "__call__")] + fn call(&self, predictions: &Bound, targets: &Bound) -> PyResult { + self.forward(predictions, targets) + } + + /// Get the reduction mode + #[getter] + fn reduction(&self) -> &str { + self.inner.reduction() + } + + /// String representation + fn __repr__(&self) -> String { + format!("MAELoss(reduction='{}')", self.inner.reduction()) + } +} + +/// Huber Loss function +#[pyclass(name = "HuberLoss")] +pub struct PyHuberLoss { + inner: HuberLoss, +} + +#[pymethods] +impl PyHuberLoss { + /// Create a new Huber loss + #[new] + #[pyo3(signature = (delta=None, reduction=None))] + fn new(delta: Option, reduction: Option<&str>) -> Self { + let delta = delta.unwrap_or(1.0); + let reduction = reduction.unwrap_or("mean"); + Self { + inner: HuberLoss::new(delta, reduction), + } + } + + /// Compute the Huber loss + fn forward(&self, predictions: &Bound, targets: &Bound) -> PyResult { + let predictions = borrow_tensor(predictions)?; + let targets = borrow_tensor(targets)?; + let result = self + .inner + .forward(predictions.tensor(), targets.tensor()) + .map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) + } + + #[pyo3(name = "__call__")] + fn call(&self, predictions: &Bound, targets: &Bound) -> PyResult { + self.forward(predictions, targets) + } + + /// Get the delta parameter + #[getter] + fn delta(&self) -> f64 { + self.inner.delta() + } + + /// Get the reduction mode + #[getter] + fn reduction(&self) -> &str { + self.inner.reduction() + } + + /// String representation + fn __repr__(&self) -> String { + format!( + "HuberLoss(delta={}, reduction='{}')", + self.inner.delta(), + self.inner.reduction() + ) + } +} + +/// Smooth L1 Loss function +#[pyclass(name = "SmoothL1Loss")] +pub struct PySmoothL1Loss { + inner: SmoothL1Loss, +} + +#[pymethods] +impl PySmoothL1Loss { + /// Create a new Smooth L1 loss + #[new] + #[pyo3(signature = (reduction=None))] + fn new(reduction: Option<&str>) -> Self { + let reduction = reduction.unwrap_or("mean"); + Self { + inner: SmoothL1Loss::new(reduction), + } + } + + /// Compute the Smooth L1 loss + fn forward(&self, predictions: &Bound, targets: &Bound) -> PyResult { + let predictions = borrow_tensor(predictions)?; + let targets = borrow_tensor(targets)?; + let result = self + .inner + .forward(predictions.tensor(), targets.tensor()) + .map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) + } + + #[pyo3(name = "__call__")] + fn call(&self, predictions: &Bound, targets: &Bound) -> PyResult { + self.forward(predictions, targets) + } + + /// Get the reduction mode + #[getter] + fn reduction(&self) -> &str { + self.inner.reduction() + } + + /// String representation + fn __repr__(&self) -> String { + format!("SmoothL1Loss(reduction='{}')", self.inner.reduction()) + } +} + +/// Log-cosh Loss function +#[pyclass(name = "LogCoshLoss")] +pub struct PyLogCoshLoss { + inner: LogCoshLoss, +} + +#[pymethods] +impl PyLogCoshLoss { + /// Create a new Log-cosh loss + #[new] + #[pyo3(signature = (reduction=None))] + fn new(reduction: Option<&str>) -> Self { + let reduction = reduction.unwrap_or("mean"); + Self { + inner: LogCoshLoss::new(reduction), + } + } + + /// Compute the Log-cosh loss + fn forward(&self, predictions: &Bound, targets: &Bound) -> PyResult { + let predictions = borrow_tensor(predictions)?; + let targets = borrow_tensor(targets)?; + let result = self + .inner + .forward(predictions.tensor(), targets.tensor()) + .map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) + } + + #[pyo3(name = "__call__")] + fn call(&self, predictions: &Bound, targets: &Bound) -> PyResult { + self.forward(predictions, targets) + } + + /// Get the reduction mode + #[getter] + fn reduction(&self) -> &str { + self.inner.reduction() + } + + /// String representation + fn __repr__(&self) -> String { + format!("LogCoshLoss(reduction='{}')", self.inner.reduction()) + } +} + +/// Cross Entropy Loss function +#[pyclass(name = "CrossEntropyLoss")] +pub struct PyCrossEntropyLoss { + inner: CrossEntropyLoss, +} + +#[pymethods] +impl PyCrossEntropyLoss { + /// Create a new Cross Entropy loss + #[new] + #[pyo3(signature = (reduction=None))] + fn new(reduction: Option<&str>) -> Self { + let reduction = reduction.unwrap_or("mean"); + Self { + inner: CrossEntropyLoss::new(reduction), + } + } + + /// Compute the Cross Entropy loss + fn forward(&self, predictions: &Bound, targets: &Bound) -> PyResult { + let predictions = borrow_tensor(predictions)?; + let targets = borrow_tensor(targets)?; + let result = self + .inner + .forward(predictions.tensor(), targets.tensor()) + .map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) + } + + #[pyo3(name = "__call__")] + fn call(&self, predictions: &Bound, targets: &Bound) -> PyResult { + self.forward(predictions, targets) + } + + /// Get the reduction mode + #[getter] + fn reduction(&self) -> &str { + self.inner.reduction() + } + + /// String representation + fn __repr__(&self) -> String { + format!("CrossEntropyLoss(reduction='{}')", self.inner.reduction()) + } +} + +/// Binary Cross Entropy Loss function +#[pyclass(name = "BCELoss")] +pub struct PyBCELoss { + inner: BCELoss, +} + +#[pymethods] +impl PyBCELoss { + /// Create a new BCE loss + #[new] + #[pyo3(signature = (reduction=None))] + fn new(reduction: Option<&str>) -> Self { + let reduction = reduction.unwrap_or("mean"); + Self { + inner: BCELoss::new(reduction), + } + } + + /// Compute the BCE loss + fn forward(&self, predictions: &Bound, targets: &Bound) -> PyResult { + let predictions = borrow_tensor(predictions)?; + let targets = borrow_tensor(targets)?; + let result = self + .inner + .forward(predictions.tensor(), targets.tensor()) + .map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) + } + + #[pyo3(name = "__call__")] + fn call(&self, predictions: &Bound, targets: &Bound) -> PyResult { + self.forward(predictions, targets) + } + + /// Get the reduction mode + #[getter] + fn reduction(&self) -> &str { + self.inner.reduction() + } + + /// String representation + fn __repr__(&self) -> String { + format!("BCELoss(reduction='{}')", self.inner.reduction()) + } +} + +/// Focal Loss function +#[pyclass(name = "FocalLoss")] +pub struct PyFocalLoss { + inner: FocalLoss, +} + +#[pymethods] +impl PyFocalLoss { + /// Create a new Focal loss + #[new] + #[pyo3(signature = (alpha=None, gamma=None, reduction=None))] + fn new(alpha: Option, gamma: Option, reduction: Option<&str>) -> Self { + let alpha = alpha.unwrap_or(0.25); + let gamma = gamma.unwrap_or(2.0); + let reduction = reduction.unwrap_or("mean"); + Self { + inner: FocalLoss::new(alpha, gamma, reduction), + } + } + + /// Compute the Focal loss + fn forward(&self, predictions: &Bound, targets: &Bound) -> PyResult { + let predictions = borrow_tensor(predictions)?; + let targets = borrow_tensor(targets)?; + let result = self + .inner + .forward(predictions.tensor(), targets.tensor()) + .map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) + } + + #[pyo3(name = "__call__")] + fn call(&self, predictions: &Bound, targets: &Bound) -> PyResult { + self.forward(predictions, targets) + } + + /// Get the alpha parameter + #[getter] + fn alpha(&self) -> f64 { + self.inner.alpha() + } + + /// Get the gamma parameter + #[getter] + fn gamma(&self) -> f64 { + self.inner.gamma() + } + + /// Get the reduction mode + #[getter] + fn reduction(&self) -> &str { + self.inner.reduction() + } + + /// String representation + fn __repr__(&self) -> String { + format!( + "FocalLoss(alpha={}, gamma={}, reduction='{}')", + self.inner.alpha(), + self.inner.gamma(), + self.inner.reduction() + ) + } +} + +/// Register neural network module with Python +pub fn register_nn_module(py: Python, parent_module: &Bound) -> PyResult<()> { + let nn_module = Pyo3Module::new(py, "nn")?; + + // Add layer classes + nn_module.add_class::()?; + nn_module.add_class::()?; + nn_module.add_class::()?; + nn_module.add_class::()?; + nn_module.add_class::()?; + nn_module.add_class::()?; + nn_module.add_class::()?; + nn_module.add_class::()?; + nn_module.add_class::()?; + nn_module.add_class::()?; + nn_module.add_class::()?; + nn_module.add_class::()?; + nn_module.add_class::()?; + nn_module.add_class::()?; + nn_module.add_class::()?; + + // Add functional APIs + nn_module.add_function(wrap_pyfunction!(dense_layer, &nn_module)?)?; + nn_module.add_function(wrap_pyfunction!(conv2d, &nn_module)?)?; + nn_module.add_function(wrap_pyfunction!(batch_norm, &nn_module)?)?; + nn_module.add_function(wrap_pyfunction!(cross_entropy, &nn_module)?)?; + nn_module.add_function(wrap_pyfunction!(dropout_functional, &nn_module)?)?; + nn_module.add_function(wrap_pyfunction!(dropout2d_functional, &nn_module)?)?; + nn_module.add_function(wrap_pyfunction!(mse_loss_functional, &nn_module)?)?; + nn_module.add_function(wrap_pyfunction!(smooth_l1_loss_functional, &nn_module)?)?; + nn_module.add_function(wrap_pyfunction!(log_cosh_loss_functional, &nn_module)?)?; + nn_module.add_function(wrap_pyfunction!( + binary_cross_entropy_functional, + &nn_module + )?)?; + + // Add loss function classes + nn_module.add_class::()?; + nn_module.add_class::()?; + nn_module.add_class::()?; + nn_module.add_class::()?; + nn_module.add_class::()?; + nn_module.add_class::()?; + nn_module.add_class::()?; + nn_module.add_class::()?; + + parent_module.add_submodule(&nn_module)?; + Ok(()) +} diff --git a/bindings/src/nn/module.rs b/bindings/src/nn/module.rs index 407521bf..9c90b376 100644 --- a/bindings/src/nn/module.rs +++ b/bindings/src/nn/module.rs @@ -1,905 +1,913 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -use crate::device::PyDevice; -use crate::dtype; -use crate::error::_convert_error; -use crate::serialization::PyStateDict; -use crate::tensor::PyTensor; -use engine::Device; -use engine::nn::{ - BCELoss, CrossEntropyLoss, DenseLayer, FocalLoss, HuberLoss, Layer, LogCoshLoss, MAELoss, - MSELoss, ReLU, Sequential, Sigmoid, SmoothL1Loss, Softmax, Tanh, - activation::{ELU, GELU, LeakyReLU}, - conv::Conv2d, - dropout::{Dropout, Dropout2d}, - normalization::{BatchNorm1d, BatchNorm2d}, - utils::{LayerUtils, SequentialUtils}, -}; -use engine::operations::batch_norm as batch_norm_op; -use engine::operations::conv2d as conv2d_op; -use engine::operations::linalg::matmul as matmul_op; -use engine::operations::loss::cross_entropy as cross_entropy_op; -use engine::serialization::{ModelMetadata, ModelSerializer, SerializationFormat, SerializedModel}; -use pyo3::exceptions::{PyIndexError, PyTypeError, PyValueError}; -use pyo3::intern; -use pyo3::prelude::*; -use pyo3::PyClassInitializer; -use pyo3::types::{PyAny, PyDict, PyModule as Pyo3Module}; - -fn borrow_tensor<'py>(value: &'py Bound<'py, PyAny>) -> PyResult> { - if let Ok(tensor) = value.extract::>() { - return Ok(tensor); - } - - let py = value.py(); - let inner = value - .getattr(intern!(py, "_tensor")) - .map_err(|_| PyTypeError::new_err("expected a minitensor Tensor or core Tensor"))?; - Ok(inner.extract::>()?) -} - -fn borrow_optional_tensor<'py>( - value: Option<&'py Bound<'py, PyAny>>, -) -> PyResult>> { - value.map(borrow_tensor).transpose() -} - -fn borrow_tensor_mut<'py>(value: &'py Bound<'py, PyAny>) -> PyResult> { - if let Ok(tensor) = value.extract::>() { - return Ok(tensor); - } - - let py = value.py(); - let inner = value - .getattr(intern!(py, "_tensor")) - .map_err(|_| PyTypeError::new_err("expected a minitensor Tensor or core Tensor"))?; - Ok(inner.extract::>()?) -} - -fn borrow_optional_tensor_mut<'py>( - value: Option<&'py Bound<'py, PyAny>>, -) -> PyResult>> { - value.map(borrow_tensor_mut).transpose() -} - -#[pyfunction] -fn dense_layer( - input: &Bound, - weight: &Bound, - bias: Option<&Bound>, -) -> PyResult { - let input_tensor = borrow_tensor(input)?; - let weight_tensor = borrow_tensor(weight)?; - - if weight_tensor.tensor().ndim() != 2 { - return Err(PyValueError::new_err("weight tensor must be 2-dimensional")); - } - - let weight_t = weight_tensor - .tensor() - .transpose(0, 1) - .map_err(_convert_error)?; - let mut output = matmul_op(input_tensor.tensor(), &weight_t).map_err(_convert_error)?; - - let bias_tensor = borrow_optional_tensor(bias)?; - if let Some(bias_ref) = bias_tensor { - output = output.add(bias_ref.tensor()).map_err(_convert_error)?; - } - - Ok(PyTensor::from_tensor(output)) -} - -fn parse_pair_arg( - name: &str, - value: Option<&Bound>, - default: (usize, usize), -) -> PyResult<(usize, usize)> { - match value { - None => Ok(default), - Some(bound) => { - if let Ok(scalar) = bound.extract::() { - if scalar < 0 { - return Err(PyValueError::new_err(format!( - "{name} must be non-negative" - ))); - } - let scalar = scalar as usize; - return Ok((scalar, scalar)); - } - - if let Ok(pair) = bound.extract::<(isize, isize)>() { - if pair.0 < 0 || pair.1 < 0 { - return Err(PyValueError::new_err(format!( - "{name} values must be non-negative" - ))); - } - return Ok((pair.0 as usize, pair.1 as usize)); - } - - let seq = bound.extract::>()?; - if seq.len() != 2 { - return Err(PyTypeError::new_err(format!( - "{name} must be an int or a sequence of length 2" - ))); - } - if seq[0] < 0 || seq[1] < 0 { - return Err(PyValueError::new_err(format!( - "{name} values must be non-negative" - ))); - } - Ok((seq[0] as usize, seq[1] as usize)) - } - } -} - -#[pyfunction] -#[pyo3(signature = (input, weight, bias=None, stride=None, padding=None))] -fn conv2d( - input: &Bound, - weight: &Bound, - bias: Option<&Bound>, - stride: Option<&Bound>, - padding: Option<&Bound>, -) -> PyResult { - let input_tensor = borrow_tensor(input)?; - let weight_tensor = borrow_tensor(weight)?; - let bias_tensor = borrow_optional_tensor(bias)?; - let stride = parse_pair_arg("stride", stride, (1, 1))?; - let padding = parse_pair_arg("padding", padding, (0, 0))?; - let result = conv2d_op( - input_tensor.tensor(), - weight_tensor.tensor(), - bias_tensor.as_ref().map(|b| b.tensor()), - stride, - padding, - ) - .map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) -} - -#[pyfunction] -#[pyo3(signature = (input, running_mean=None, running_var=None, weight=None, bias=None, training=true, momentum=0.1, eps=1e-5))] -#[allow(clippy::too_many_arguments)] -fn batch_norm( - input: &Bound, - running_mean: Option<&Bound>, - running_var: Option<&Bound>, - weight: Option<&Bound>, - bias: Option<&Bound>, - training: bool, - momentum: f64, - eps: f64, -) -> PyResult { - let input_tensor = borrow_tensor(input)?; - let mut running_mean_tensor = borrow_optional_tensor_mut(running_mean)?; - let mut running_var_tensor = borrow_optional_tensor_mut(running_var)?; - let weight_tensor = borrow_optional_tensor(weight)?; - let bias_tensor = borrow_optional_tensor(bias)?; - - let rm_tensor = running_mean_tensor.as_mut().map(|t| t.tensor_mut()); - let rv_tensor = running_var_tensor.as_mut().map(|t| t.tensor_mut()); - let result = batch_norm_op( - input_tensor.tensor(), - rm_tensor, - rv_tensor, - weight_tensor.as_ref().map(|w| w.tensor()), - bias_tensor.as_ref().map(|b| b.tensor()), - training, - momentum, - eps, - ) - .map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) -} - -#[pyfunction] -#[pyo3(signature = (input, target, reduction="mean", dim=1))] -fn cross_entropy( - input: &Bound, - target: &Bound, - reduction: &str, - dim: isize, -) -> PyResult { - let input_tensor = borrow_tensor(input)?; - let target_tensor = borrow_tensor(target)?; - - let ndim = input_tensor.tensor().ndim() as isize; - let axis = if dim < 0 { ndim + dim } else { dim }; - if axis < 0 || axis as usize >= ndim as usize { - return Err(PyIndexError::new_err("dim out of range")); - } - let result = cross_entropy_op( - input_tensor.tensor(), - target_tensor.tensor(), - reduction, - axis as usize, - ) - .map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) -} - -#[pyfunction(name = "dropout")] -#[pyo3(signature = (input, p=0.5, training=true))] -fn dropout_functional(input: &Bound, p: f64, training: bool) -> PyResult { - let tensor = borrow_tensor(input)?; - let mut layer = Dropout::new(Some(p)).map_err(_convert_error)?; - if training { - layer.train(); - } else { - layer.eval(); - } - let result = layer.forward(tensor.tensor()).map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) -} - -#[pyfunction(name = "dropout2d")] -#[pyo3(signature = (input, p=0.5, training=true))] -fn dropout2d_functional(input: &Bound, p: f64, training: bool) -> PyResult { - let tensor = borrow_tensor(input)?; - let mut layer = Dropout2d::new(Some(p)).map_err(_convert_error)?; - if training { - layer.train(); - } else { - layer.eval(); - } - let result = layer.forward(tensor.tensor()).map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) -} - -#[pyfunction(name = "mse_loss")] -fn mse_loss_functional( - input: &Bound, - target: &Bound, - reduction: Option<&str>, -) -> PyResult { - let prediction = borrow_tensor(input)?; - let target_tensor = borrow_tensor(target)?; - let reduction = reduction.unwrap_or("mean"); - let loss = MSELoss::new(reduction); - let result = loss - .forward(prediction.tensor(), target_tensor.tensor()) - .map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) -} - -#[pyfunction(name = "smooth_l1_loss")] -fn smooth_l1_loss_functional( - input: &Bound, - target: &Bound, - reduction: Option<&str>, -) -> PyResult { - let prediction = borrow_tensor(input)?; - let target_tensor = borrow_tensor(target)?; - let reduction = reduction.unwrap_or("mean"); - let loss = SmoothL1Loss::new(reduction); - let result = loss - .forward(prediction.tensor(), target_tensor.tensor()) - .map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) -} - -#[pyfunction(name = "log_cosh_loss")] -fn log_cosh_loss_functional( - input: &Bound, - target: &Bound, - reduction: Option<&str>, -) -> PyResult { - let prediction = borrow_tensor(input)?; - let target_tensor = borrow_tensor(target)?; - let reduction = reduction.unwrap_or("mean"); - let loss = LogCoshLoss::new(reduction); - let result = loss - .forward(prediction.tensor(), target_tensor.tensor()) - .map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) -} - -#[pyfunction(name = "binary_cross_entropy")] -#[pyo3(signature = (input, target, reduction="mean"))] -fn binary_cross_entropy_functional( - input: &Bound, - target: &Bound, - reduction: &str, -) -> PyResult { - let prediction = borrow_tensor(input)?; - let target_tensor = borrow_tensor(target)?; - let loss = BCELoss::new(reduction); - let result = loss - .forward(prediction.tensor(), target_tensor.tensor()) - .map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) -} - -/// Base class for neural network modules -#[pyclass(name = "Module", subclass)] -pub struct PyModule { - // This will be a trait object in practice - // For now, we'll use an enum to handle different layer types - inner: ModuleType, -} - -enum ModuleType { - DenseLayer(DenseLayer), - ReLU(ReLU), - Sigmoid(Sigmoid), - Tanh(Tanh), - Softmax(Softmax), - LeakyReLU(LeakyReLU), - Elu(ELU), - Gelu(GELU), - Sequential(Sequential), - Conv2d(Conv2d), - BatchNorm1d(BatchNorm1d), - BatchNorm2d(BatchNorm2d), - Dropout(Dropout), - Dropout2d(Dropout2d), -} - -#[pymethods] -impl PyModule { - /// Forward pass through the module - fn forward(&mut self, input: &Bound) -> PyResult { - let input_tensor = borrow_tensor(input)?; - let result = match &mut self.inner { - ModuleType::DenseLayer(layer) => layer.forward(input_tensor.tensor()), - ModuleType::ReLU(layer) => layer.forward(input_tensor.tensor()), - ModuleType::Sigmoid(layer) => layer.forward(input_tensor.tensor()), - ModuleType::Tanh(layer) => layer.forward(input_tensor.tensor()), - ModuleType::Softmax(layer) => layer.forward(input_tensor.tensor()), - ModuleType::LeakyReLU(layer) => layer.forward(input_tensor.tensor()), - ModuleType::Elu(layer) => layer.forward(input_tensor.tensor()), - ModuleType::Gelu(layer) => layer.forward(input_tensor.tensor()), - ModuleType::Sequential(layer) => layer.forward(input_tensor.tensor()), - ModuleType::Conv2d(layer) => layer.forward(input_tensor.tensor()), - ModuleType::BatchNorm1d(layer) => layer.forward(input_tensor.tensor()), - ModuleType::BatchNorm2d(layer) => layer.forward(input_tensor.tensor()), - ModuleType::Dropout(layer) => layer.forward(input_tensor.tensor()), - ModuleType::Dropout2d(layer) => layer.forward(input_tensor.tensor()), - } - .map_err(_convert_error)?; - - Ok(PyTensor::from_tensor(result)) - } - - #[pyo3(name = "__call__")] - fn call(&mut self, input: &Bound) -> PyResult { - self.forward(input) - } - - /// Get all parameters of the module - fn parameters(&self) -> Vec { - let params = match &self.inner { - ModuleType::DenseLayer(layer) => layer.parameters(), - ModuleType::ReLU(layer) => layer.parameters(), - ModuleType::Sigmoid(layer) => layer.parameters(), - ModuleType::Tanh(layer) => layer.parameters(), - ModuleType::Softmax(layer) => layer.parameters(), - ModuleType::LeakyReLU(layer) => layer.parameters(), - ModuleType::Elu(layer) => layer.parameters(), - ModuleType::Gelu(layer) => layer.parameters(), - ModuleType::Sequential(layer) => layer.parameters(), - ModuleType::Conv2d(layer) => layer.parameters(), - ModuleType::BatchNorm1d(layer) => layer.parameters(), - ModuleType::BatchNorm2d(layer) => layer.parameters(), - ModuleType::Dropout(layer) => layer.parameters(), - ModuleType::Dropout2d(layer) => layer.parameters(), - }; - - params - .into_iter() - .map(|tensor| PyTensor::from_tensor(tensor.clone())) - .collect() - } - - /// Set module to training mode - fn train(&mut self) { - match &mut self.inner { - ModuleType::DenseLayer(layer) => layer.train(), - ModuleType::ReLU(layer) => layer.train(), - ModuleType::Sigmoid(layer) => layer.train(), - ModuleType::Tanh(layer) => layer.train(), - ModuleType::Softmax(layer) => layer.train(), - ModuleType::LeakyReLU(layer) => layer.train(), - ModuleType::Elu(layer) => layer.train(), - ModuleType::Gelu(layer) => layer.train(), - ModuleType::Sequential(layer) => layer.train(), - ModuleType::Conv2d(layer) => layer.train(), - ModuleType::BatchNorm1d(layer) => layer.train(), - ModuleType::BatchNorm2d(layer) => layer.train(), - ModuleType::Dropout(layer) => layer.train(), - ModuleType::Dropout2d(layer) => layer.train(), - } - } - - /// Set module to evaluation mode - fn eval(&mut self) { - match &mut self.inner { - ModuleType::DenseLayer(layer) => layer.eval(), - ModuleType::ReLU(layer) => layer.eval(), - ModuleType::Sigmoid(layer) => layer.eval(), - ModuleType::Tanh(layer) => layer.eval(), - ModuleType::Softmax(layer) => layer.eval(), - ModuleType::LeakyReLU(layer) => layer.eval(), - ModuleType::Elu(layer) => layer.eval(), - ModuleType::Gelu(layer) => layer.eval(), - ModuleType::Sequential(layer) => layer.eval(), - ModuleType::Conv2d(layer) => layer.eval(), - ModuleType::BatchNorm1d(layer) => layer.eval(), - ModuleType::BatchNorm2d(layer) => layer.eval(), - ModuleType::Dropout(layer) => layer.eval(), - ModuleType::Dropout2d(layer) => layer.eval(), - } - } - - /// Get number of parameters - fn num_parameters(&self) -> usize { - match &self.inner { - ModuleType::DenseLayer(layer) => layer.num_parameters(), - ModuleType::ReLU(layer) => layer.num_parameters(), - ModuleType::Sigmoid(layer) => layer.num_parameters(), - ModuleType::Tanh(layer) => layer.num_parameters(), - ModuleType::Softmax(layer) => layer.num_parameters(), - ModuleType::LeakyReLU(layer) => layer.num_parameters(), - ModuleType::Elu(layer) => layer.num_parameters(), - ModuleType::Gelu(layer) => layer.num_parameters(), - ModuleType::Sequential(layer) => layer.num_parameters(), - ModuleType::Conv2d(layer) => layer.num_parameters(), - ModuleType::BatchNorm1d(layer) => layer.num_parameters(), - ModuleType::BatchNorm2d(layer) => layer.num_parameters(), - ModuleType::Dropout(layer) => layer.num_parameters(), - ModuleType::Dropout2d(layer) => layer.num_parameters(), - } - } - - /// Get detailed parameter statistics - fn parameter_stats(&self, py: Python) -> PyResult> { - let layer: &dyn Layer = match &self.inner { - ModuleType::DenseLayer(layer) => layer, - ModuleType::ReLU(layer) => layer, - ModuleType::Sigmoid(layer) => layer, - ModuleType::Tanh(layer) => layer, - ModuleType::Softmax(layer) => layer, - ModuleType::LeakyReLU(layer) => layer, - ModuleType::Elu(layer) => layer, - ModuleType::Gelu(layer) => layer, - ModuleType::Sequential(layer) => layer, - ModuleType::Conv2d(layer) => layer, - ModuleType::BatchNorm1d(layer) => layer, - ModuleType::BatchNorm2d(layer) => layer, - ModuleType::Dropout(layer) => layer, - ModuleType::Dropout2d(layer) => layer, - }; - let stats = LayerUtils::parameter_stats(layer); - let dict = PyDict::new(py); - dict.set_item("total_parameters", stats.total_parameters)?; - dict.set_item("trainable_parameters", stats.trainable_parameters)?; - dict.set_item("non_trainable_parameters", stats.non_trainable_parameters)?; - dict.set_item("parameter_count_by_tensor", stats.parameter_count_by_tensor)?; - Ok(dict.into()) - } - - /// Get memory usage information - fn memory_usage(&self, py: Python) -> PyResult> { - let layer: &dyn Layer = match &self.inner { - ModuleType::DenseLayer(layer) => layer, - ModuleType::ReLU(layer) => layer, - ModuleType::Sigmoid(layer) => layer, - ModuleType::Tanh(layer) => layer, - ModuleType::Softmax(layer) => layer, - ModuleType::LeakyReLU(layer) => layer, - ModuleType::Elu(layer) => layer, - ModuleType::Gelu(layer) => layer, - ModuleType::Sequential(layer) => layer, - ModuleType::Conv2d(layer) => layer, - ModuleType::BatchNorm1d(layer) => layer, - ModuleType::BatchNorm2d(layer) => layer, - ModuleType::Dropout(layer) => layer, - ModuleType::Dropout2d(layer) => layer, - }; - let usage = LayerUtils::memory_usage(layer); - let dict = PyDict::new(py); - dict.set_item("total_bytes", usage.total_bytes)?; - let dtype_dict = PyDict::new(py); - for (dtype, bytes) in usage.bytes_by_dtype { - dtype_dict.set_item(format!("{:?}", dtype), bytes)?; - } - dict.set_item("bytes_by_dtype", dtype_dict)?; - Ok(dict.into()) - } - - /// Generate summary - #[pyo3(signature = (name=None))] - fn summary(&self, name: Option<&str>) -> PyResult { - match &self.inner { - ModuleType::Sequential(model) => Ok(SequentialUtils::model_summary(model, name)), - _ => { - let layer: &dyn Layer = match &self.inner { - ModuleType::DenseLayer(layer) => layer, - ModuleType::ReLU(layer) => layer, - ModuleType::Sigmoid(layer) => layer, - ModuleType::Tanh(layer) => layer, - ModuleType::Softmax(layer) => layer, - ModuleType::LeakyReLU(layer) => layer, - ModuleType::Elu(layer) => layer, - ModuleType::Gelu(layer) => layer, - ModuleType::Sequential(layer) => layer, - ModuleType::Conv2d(layer) => layer, - ModuleType::BatchNorm1d(layer) => layer, - ModuleType::BatchNorm2d(layer) => layer, - ModuleType::Dropout(layer) => layer, - ModuleType::Dropout2d(layer) => layer, - }; - let owned; - let layer_name = match name { - Some(n) => n, - None => { - owned = self.__repr__(); - &owned - } - }; - Ok(LayerUtils::layer_summary(layer, layer_name)) - } - } - } - - /// Estimate forward memory usage for Sequential models - fn forward_memory_estimate( - &self, - input_shape: Vec, - batch_size: usize, - py: Python, - ) -> PyResult> { - if let ModuleType::Sequential(model) = &self.inner { - let est = SequentialUtils::estimate_forward_memory(model, &input_shape, batch_size); - let dict = PyDict::new(py); - dict.set_item("parameter_memory", est.parameter_memory)?; - dict.set_item( - "estimated_activation_memory", - est.estimated_activation_memory, - )?; - dict.set_item("estimated_total_memory", est.estimated_total_memory)?; - dict.set_item("input_memory", est.input_memory)?; - Ok(dict.into()) - } else { - Err(PyErr::new::( - "forward_memory_estimate only valid for Sequential modules", - )) - } - } - - /// String representation - fn __repr__(&self) -> String { - match &self.inner { - ModuleType::DenseLayer(layer) => format!( - "DenseLayer(in_features={}, out_features={})", - layer.in_features(), - layer.out_features() - ), - ModuleType::ReLU(_) => "ReLU()".to_string(), - ModuleType::Sigmoid(_) => "Sigmoid()".to_string(), - ModuleType::Tanh(_) => "Tanh()".to_string(), - ModuleType::Softmax(layer) => format!("Softmax(dim={:?})", layer.dim()), - ModuleType::LeakyReLU(layer) => { - format!("LeakyReLU(negative_slope={})", layer.negative_slope()) - } - ModuleType::Elu(layer) => format!("ELU(alpha={})", layer.alpha()), - ModuleType::Gelu(_) => "GELU()".to_string(), - ModuleType::Sequential(_) => "Sequential(...)".to_string(), - ModuleType::Conv2d(layer) => format!( - "Conv2d(in_channels={}, out_channels={}, kernel_size={:?})", - layer.in_channels(), - layer.out_channels(), - layer.kernel_size() - ), - ModuleType::BatchNorm1d(layer) => { - format!("BatchNorm1d(num_features={})", layer.num_features()) - } - ModuleType::BatchNorm2d(layer) => { - format!("BatchNorm2d(num_features={})", layer.num_features()) - } - ModuleType::Dropout(layer) => format!("Dropout(p={})", layer.p()), - ModuleType::Dropout2d(layer) => format!("Dropout2d(p={})", layer.p()), - } - } - - /// Save module state to a file (basic implementation) - fn save(&self, path: &str, format: Option<&str>) -> PyResult<()> { - // Build a SerializedModel with metadata and engine state_dict - use engine::nn::Module as _; - let state = match &self.inner { - ModuleType::DenseLayer(layer) => layer.state_dict(), - ModuleType::ReLU(layer) => layer.state_dict(), - ModuleType::Sigmoid(layer) => layer.state_dict(), - ModuleType::Tanh(layer) => layer.state_dict(), - ModuleType::Softmax(layer) => layer.state_dict(), - ModuleType::LeakyReLU(layer) => layer.state_dict(), - ModuleType::Elu(layer) => layer.state_dict(), - ModuleType::Gelu(layer) => layer.state_dict(), - ModuleType::Sequential(layer) => layer.state_dict(), - ModuleType::Conv2d(layer) => layer.state_dict(), - ModuleType::BatchNorm1d(layer) => layer.state_dict(), - ModuleType::BatchNorm2d(layer) => layer.state_dict(), - ModuleType::Dropout(layer) => layer.state_dict(), - ModuleType::Dropout2d(layer) => layer.state_dict(), - }; - - let metadata = ModelMetadata::new("module".to_string(), "Module".to_string()); - let model = SerializedModel::new(metadata, state); - match format.map(|s| s.to_lowercase()) { - Some(ref s) if s == "json" => { - ModelSerializer::save(&model, path, SerializationFormat::Json) - } - Some(ref s) if s == "bin" || s == "binary" => { - ModelSerializer::save(&model, path, SerializationFormat::Binary) - } - Some(ref s) if s == "msgpack" || s == "messagepack" => { - ModelSerializer::save(&model, path, SerializationFormat::MessagePack) - } - _ => ModelSerializer::save_auto(&model, path), - } - .map_err(_convert_error) - } - - /// Load module state from a file (basic implementation) - #[staticmethod] - fn load_state_from(path: &str, format: Option<&str>) -> PyResult { - let model = match format.map(|s| s.to_lowercase()) { - Some(ref s) if s == "json" => ModelSerializer::load(path, SerializationFormat::Json), - Some(ref s) if s == "bin" || s == "binary" => { - ModelSerializer::load(path, SerializationFormat::Binary) - } - Some(ref s) if s == "msgpack" || s == "messagepack" => { - ModelSerializer::load(path, SerializationFormat::MessagePack) - } - _ => ModelSerializer::load_auto(path), - } - .map_err(_convert_error)?; - Ok(crate::serialization::PyStateDict::from_engine( - model.state_dict, - )) - } - - /// Return a StateDict snapshot of this module - fn state_dict(&self) -> PyStateDict { - use engine::nn::Module as _; - let state = match &self.inner { - ModuleType::DenseLayer(layer) => layer.state_dict(), - ModuleType::ReLU(layer) => layer.state_dict(), - ModuleType::Sigmoid(layer) => layer.state_dict(), - ModuleType::Tanh(layer) => layer.state_dict(), - ModuleType::Softmax(layer) => layer.state_dict(), - ModuleType::LeakyReLU(layer) => layer.state_dict(), - ModuleType::Elu(layer) => layer.state_dict(), - ModuleType::Gelu(layer) => layer.state_dict(), - ModuleType::Sequential(layer) => layer.state_dict(), - ModuleType::Conv2d(layer) => layer.state_dict(), - ModuleType::BatchNorm1d(layer) => layer.state_dict(), - ModuleType::BatchNorm2d(layer) => layer.state_dict(), - ModuleType::Dropout(layer) => layer.state_dict(), - ModuleType::Dropout2d(layer) => layer.state_dict(), - }; - crate::serialization::PyStateDict::from_engine(state) - } - - /// Load a provided StateDict into this module - fn load_state_dict(&mut self, state: &PyStateDict, device: Option<&PyDevice>) -> PyResult<()> { - use engine::nn::Module as _; - let dev = device.map(|d| d.device()); - let sd_ref = crate::serialization::PyStateDict::inner_ref(state); - let res = match &mut self.inner { - ModuleType::DenseLayer(layer) => layer.load_state_dict(sd_ref, dev), - ModuleType::ReLU(layer) => layer.load_state_dict(sd_ref, dev), - ModuleType::Sigmoid(layer) => layer.load_state_dict(sd_ref, dev), - ModuleType::Tanh(layer) => layer.load_state_dict(sd_ref, dev), - ModuleType::Softmax(layer) => layer.load_state_dict(sd_ref, dev), - ModuleType::LeakyReLU(layer) => layer.load_state_dict(sd_ref, dev), - ModuleType::Elu(layer) => layer.load_state_dict(sd_ref, dev), - ModuleType::Gelu(layer) => layer.load_state_dict(sd_ref, dev), - ModuleType::Sequential(layer) => layer.load_state_dict(sd_ref, dev), - ModuleType::Conv2d(layer) => layer.load_state_dict(sd_ref, dev), - ModuleType::BatchNorm1d(layer) => layer.load_state_dict(sd_ref, dev), - ModuleType::BatchNorm2d(layer) => layer.load_state_dict(sd_ref, dev), - ModuleType::Dropout(layer) => layer.load_state_dict(sd_ref, dev), - ModuleType::Dropout2d(layer) => layer.load_state_dict(sd_ref, dev), - }; - res.map_err(_convert_error) - } -} - -impl PyModule { - pub fn from_dense_layer(dense_layer: DenseLayer) -> Self { - Self { - inner: ModuleType::DenseLayer(dense_layer), - } - } - - pub fn from_relu(relu: ReLU) -> Self { - Self { - inner: ModuleType::ReLU(relu), - } - } - - pub fn from_sigmoid(sigmoid: Sigmoid) -> Self { - Self { - inner: ModuleType::Sigmoid(sigmoid), - } - } - - pub fn from_tanh(tanh: Tanh) -> Self { - Self { - inner: ModuleType::Tanh(tanh), - } - } - - pub fn from_softmax(softmax: Softmax) -> Self { - Self { - inner: ModuleType::Softmax(softmax), - } - } - - pub fn from_leaky_relu(leaky_relu: LeakyReLU) -> Self { - Self { - inner: ModuleType::LeakyReLU(leaky_relu), - } - } - - pub fn from_elu(elu: ELU) -> Self { - Self { - inner: ModuleType::Elu(elu), - } - } - - pub fn from_gelu(gelu: GELU) -> Self { - Self { - inner: ModuleType::Gelu(gelu), - } - } - - pub fn from_sequential(sequential: Sequential) -> Self { - Self { - inner: ModuleType::Sequential(sequential), - } - } - - pub fn from_conv2d(conv2d: Conv2d) -> Self { - Self { - inner: ModuleType::Conv2d(conv2d), - } - } - - pub fn from_batch_norm1d(batch_norm1d: BatchNorm1d) -> Self { - Self { - inner: ModuleType::BatchNorm1d(batch_norm1d), - } - } - - pub fn from_batch_norm2d(batch_norm2d: BatchNorm2d) -> Self { - Self { - inner: ModuleType::BatchNorm2d(batch_norm2d), - } - } - - pub fn from_dropout(dropout: Dropout) -> Self { - Self { - inner: ModuleType::Dropout(dropout), - } - } - - pub fn from_dropout2d(dropout: Dropout2d) -> Self { - Self { - inner: ModuleType::Dropout2d(dropout), - } - } - - pub fn to_layer(&self) -> PyResult> { - let layer: Box = match &self.inner { - ModuleType::DenseLayer(layer) => Box::new(layer.clone()), - ModuleType::ReLU(layer) => Box::new(layer.clone()), - ModuleType::Sigmoid(layer) => Box::new(layer.clone()), - ModuleType::Tanh(layer) => Box::new(layer.clone()), - ModuleType::Softmax(layer) => Box::new(layer.clone()), - ModuleType::LeakyReLU(layer) => Box::new(layer.clone()), - ModuleType::Elu(layer) => Box::new(layer.clone()), - ModuleType::Gelu(layer) => Box::new(layer.clone()), - ModuleType::Sequential(_) => { - return Err(PyTypeError::new_err( - "Nested Sequential modules are not supported", - )); - } - ModuleType::Conv2d(layer) => Box::new(layer.clone()), - ModuleType::BatchNorm1d(layer) => Box::new(layer.clone()), - ModuleType::BatchNorm2d(layer) => Box::new(layer.clone()), - ModuleType::Dropout(layer) => Box::new(layer.clone()), - ModuleType::Dropout2d(layer) => Box::new(layer.clone()), - }; - - Ok(layer) - } -} - -/// DenseLayer (fully connected) layer -#[pyclass(name = "DenseLayer", extends = PyModule)] -pub struct PyDenseLayer; - -#[pymethods] -impl PyDenseLayer { - /// Create a new dense layer - #[new] - #[pyo3(signature = (in_features, out_features, bias=None, device=None, dtype=None))] - fn new( - in_features: usize, - out_features: usize, - bias: Option, - device: Option<&PyDevice>, - dtype: Option<&str>, - ) -> PyResult> { - let bias = bias.unwrap_or(true); - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let dtype = dtype::resolve_dtype_arg(dtype)?; - - let dense_layer = DenseLayer::new(in_features, out_features, bias, device, dtype) - .map_err(_convert_error)?; - - Ok(PyClassInitializer::from(PyModule::from_dense_layer(dense_layer)).add_subclass(Self)) - } - - /// Get input features count - #[getter] - fn in_features(slf: PyRef) -> PyResult { - let module = slf.as_ref(); - if let ModuleType::DenseLayer(layer) = &module.inner { - Ok(layer.in_features()) - } else { - Err(PyErr::new::( - "Invalid layer type", - )) - } - } - - /// Get output features count - #[getter] - fn out_features(slf: PyRef) -> PyResult { - let module = slf.as_ref(); - if let ModuleType::DenseLayer(layer) = &module.inner { - Ok(layer.out_features()) - } else { - Err(PyErr::new::( - "Invalid layer type", - )) - } - } - - /// Get weight tensor - #[getter] - fn weight(slf: PyRef) -> PyResult { - let module = slf.as_ref(); - if let ModuleType::DenseLayer(layer) = &module.inner { - Ok(PyTensor::from_tensor(layer.weight().clone())) - } else { - Err(PyErr::new::( - "Invalid layer type", - )) - } - } - - /// Get bias tensor - #[getter] - fn bias(slf: PyRef) -> PyResult> { - let module = slf.as_ref(); - if let ModuleType::DenseLayer(layer) = &module.inner { - Ok(layer.bias().map(|b| PyTensor::from_tensor(b.clone()))) - } else { - Err(PyErr::new::( - "Invalid layer type", - )) - } - } -} - -/// ReLU activation layer -#[pyclass(name = "ReLU", extends = PyModule)] -pub struct PyReLU; +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +// `layers` hosts the PyClass wrappers and the module registration function. +// It is a child of this module so its `impl PyReLU`/`impl PyDenseLayer` +// blocks and its `wrap_pyfunction!` calls can reach the pyclass structs and +// `#[pyfunction]`s defined here. +#[path = "layers.rs"] +mod layers; +pub use self::layers::*; + +use crate::device::PyDevice; +use crate::dtype; +use crate::error::_convert_error; +use crate::serialization::PyStateDict; +use crate::tensor::PyTensor; +use engine::Device; +use engine::nn::{ + BCELoss, CrossEntropyLoss, DenseLayer, FocalLoss, HuberLoss, Layer, LogCoshLoss, MAELoss, + MSELoss, ReLU, Sequential, Sigmoid, SmoothL1Loss, Softmax, Tanh, + activation::{ELU, GELU, LeakyReLU}, + conv::Conv2d, + dropout::{Dropout, Dropout2d}, + normalization::{BatchNorm1d, BatchNorm2d}, + utils::{LayerUtils, SequentialUtils}, +}; +use engine::operations::batch_norm as batch_norm_op; +use engine::operations::conv2d as conv2d_op; +use engine::operations::linalg::matmul as matmul_op; +use engine::operations::loss::cross_entropy as cross_entropy_op; +use engine::serialization::{ModelMetadata, ModelSerializer, SerializationFormat, SerializedModel}; +use pyo3::PyClassInitializer; +use pyo3::exceptions::{PyIndexError, PyTypeError, PyValueError}; +use pyo3::intern; +use pyo3::prelude::*; +use pyo3::types::{PyAny, PyDict, PyModule as Pyo3Module}; + +fn borrow_tensor<'py>(value: &'py Bound<'py, PyAny>) -> PyResult> { + if let Ok(tensor) = value.extract::>() { + return Ok(tensor); + } + + let py = value.py(); + let inner = value + .getattr(intern!(py, "_tensor")) + .map_err(|_| PyTypeError::new_err("expected a minitensor Tensor or core Tensor"))?; + Ok(inner.extract::>()?) +} + +fn borrow_optional_tensor<'py>( + value: Option<&'py Bound<'py, PyAny>>, +) -> PyResult>> { + value.map(borrow_tensor).transpose() +} + +fn borrow_tensor_mut<'py>(value: &'py Bound<'py, PyAny>) -> PyResult> { + if let Ok(tensor) = value.extract::>() { + return Ok(tensor); + } + + let py = value.py(); + let inner = value + .getattr(intern!(py, "_tensor")) + .map_err(|_| PyTypeError::new_err("expected a minitensor Tensor or core Tensor"))?; + Ok(inner.extract::>()?) +} + +fn borrow_optional_tensor_mut<'py>( + value: Option<&'py Bound<'py, PyAny>>, +) -> PyResult>> { + value.map(borrow_tensor_mut).transpose() +} + +#[pyfunction] +fn dense_layer( + input: &Bound, + weight: &Bound, + bias: Option<&Bound>, +) -> PyResult { + let input_tensor = borrow_tensor(input)?; + let weight_tensor = borrow_tensor(weight)?; + + if weight_tensor.tensor().ndim() != 2 { + return Err(PyValueError::new_err("weight tensor must be 2-dimensional")); + } + + let weight_t = weight_tensor + .tensor() + .transpose(0, 1) + .map_err(_convert_error)?; + let mut output = matmul_op(input_tensor.tensor(), &weight_t).map_err(_convert_error)?; + + let bias_tensor = borrow_optional_tensor(bias)?; + if let Some(bias_ref) = bias_tensor { + output = output.add(bias_ref.tensor()).map_err(_convert_error)?; + } + + Ok(PyTensor::from_tensor(output)) +} + +fn parse_pair_arg( + name: &str, + value: Option<&Bound>, + default: (usize, usize), +) -> PyResult<(usize, usize)> { + match value { + None => Ok(default), + Some(bound) => { + if let Ok(scalar) = bound.extract::() { + if scalar < 0 { + return Err(PyValueError::new_err(format!( + "{name} must be non-negative" + ))); + } + let scalar = scalar as usize; + return Ok((scalar, scalar)); + } + + if let Ok(pair) = bound.extract::<(isize, isize)>() { + if pair.0 < 0 || pair.1 < 0 { + return Err(PyValueError::new_err(format!( + "{name} values must be non-negative" + ))); + } + return Ok((pair.0 as usize, pair.1 as usize)); + } + + let seq = bound.extract::>()?; + if seq.len() != 2 { + return Err(PyTypeError::new_err(format!( + "{name} must be an int or a sequence of length 2" + ))); + } + if seq[0] < 0 || seq[1] < 0 { + return Err(PyValueError::new_err(format!( + "{name} values must be non-negative" + ))); + } + Ok((seq[0] as usize, seq[1] as usize)) + } + } +} + +#[pyfunction] +#[pyo3(signature = (input, weight, bias=None, stride=None, padding=None))] +fn conv2d( + input: &Bound, + weight: &Bound, + bias: Option<&Bound>, + stride: Option<&Bound>, + padding: Option<&Bound>, +) -> PyResult { + let input_tensor = borrow_tensor(input)?; + let weight_tensor = borrow_tensor(weight)?; + let bias_tensor = borrow_optional_tensor(bias)?; + let stride = parse_pair_arg("stride", stride, (1, 1))?; + let padding = parse_pair_arg("padding", padding, (0, 0))?; + let result = conv2d_op( + input_tensor.tensor(), + weight_tensor.tensor(), + bias_tensor.as_ref().map(|b| b.tensor()), + stride, + padding, + ) + .map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) +} + +#[pyfunction] +#[pyo3(signature = (input, running_mean=None, running_var=None, weight=None, bias=None, training=true, momentum=0.1, eps=1e-5))] +#[allow(clippy::too_many_arguments)] +fn batch_norm( + input: &Bound, + running_mean: Option<&Bound>, + running_var: Option<&Bound>, + weight: Option<&Bound>, + bias: Option<&Bound>, + training: bool, + momentum: f64, + eps: f64, +) -> PyResult { + let input_tensor = borrow_tensor(input)?; + let mut running_mean_tensor = borrow_optional_tensor_mut(running_mean)?; + let mut running_var_tensor = borrow_optional_tensor_mut(running_var)?; + let weight_tensor = borrow_optional_tensor(weight)?; + let bias_tensor = borrow_optional_tensor(bias)?; + + let rm_tensor = running_mean_tensor.as_mut().map(|t| t.tensor_mut()); + let rv_tensor = running_var_tensor.as_mut().map(|t| t.tensor_mut()); + let result = batch_norm_op( + input_tensor.tensor(), + rm_tensor, + rv_tensor, + weight_tensor.as_ref().map(|w| w.tensor()), + bias_tensor.as_ref().map(|b| b.tensor()), + training, + momentum, + eps, + ) + .map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) +} + +#[pyfunction] +#[pyo3(signature = (input, target, reduction="mean", dim=1))] +fn cross_entropy( + input: &Bound, + target: &Bound, + reduction: &str, + dim: isize, +) -> PyResult { + let input_tensor = borrow_tensor(input)?; + let target_tensor = borrow_tensor(target)?; + + let ndim = input_tensor.tensor().ndim() as isize; + let axis = if dim < 0 { ndim + dim } else { dim }; + if axis < 0 || axis as usize >= ndim as usize { + return Err(PyIndexError::new_err("dim out of range")); + } + let result = cross_entropy_op( + input_tensor.tensor(), + target_tensor.tensor(), + reduction, + axis as usize, + ) + .map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) +} + +#[pyfunction(name = "dropout")] +#[pyo3(signature = (input, p=0.5, training=true))] +fn dropout_functional(input: &Bound, p: f64, training: bool) -> PyResult { + let tensor = borrow_tensor(input)?; + let mut layer = Dropout::new(Some(p)).map_err(_convert_error)?; + if training { + layer.train(); + } else { + layer.eval(); + } + let result = layer.forward(tensor.tensor()).map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) +} + +#[pyfunction(name = "dropout2d")] +#[pyo3(signature = (input, p=0.5, training=true))] +fn dropout2d_functional(input: &Bound, p: f64, training: bool) -> PyResult { + let tensor = borrow_tensor(input)?; + let mut layer = Dropout2d::new(Some(p)).map_err(_convert_error)?; + if training { + layer.train(); + } else { + layer.eval(); + } + let result = layer.forward(tensor.tensor()).map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) +} + +#[pyfunction(name = "mse_loss")] +fn mse_loss_functional( + input: &Bound, + target: &Bound, + reduction: Option<&str>, +) -> PyResult { + let prediction = borrow_tensor(input)?; + let target_tensor = borrow_tensor(target)?; + let reduction = reduction.unwrap_or("mean"); + let loss = MSELoss::new(reduction); + let result = loss + .forward(prediction.tensor(), target_tensor.tensor()) + .map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) +} + +#[pyfunction(name = "smooth_l1_loss")] +fn smooth_l1_loss_functional( + input: &Bound, + target: &Bound, + reduction: Option<&str>, +) -> PyResult { + let prediction = borrow_tensor(input)?; + let target_tensor = borrow_tensor(target)?; + let reduction = reduction.unwrap_or("mean"); + let loss = SmoothL1Loss::new(reduction); + let result = loss + .forward(prediction.tensor(), target_tensor.tensor()) + .map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) +} + +#[pyfunction(name = "log_cosh_loss")] +fn log_cosh_loss_functional( + input: &Bound, + target: &Bound, + reduction: Option<&str>, +) -> PyResult { + let prediction = borrow_tensor(input)?; + let target_tensor = borrow_tensor(target)?; + let reduction = reduction.unwrap_or("mean"); + let loss = LogCoshLoss::new(reduction); + let result = loss + .forward(prediction.tensor(), target_tensor.tensor()) + .map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) +} + +#[pyfunction(name = "binary_cross_entropy")] +#[pyo3(signature = (input, target, reduction="mean"))] +fn binary_cross_entropy_functional( + input: &Bound, + target: &Bound, + reduction: &str, +) -> PyResult { + let prediction = borrow_tensor(input)?; + let target_tensor = borrow_tensor(target)?; + let loss = BCELoss::new(reduction); + let result = loss + .forward(prediction.tensor(), target_tensor.tensor()) + .map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) +} + +/// Base class for neural network modules +#[pyclass(name = "Module", subclass)] +pub struct PyModule { + // This will be a trait object in practice + // For now, we'll use an enum to handle different layer types + inner: ModuleType, +} + +enum ModuleType { + DenseLayer(DenseLayer), + ReLU(ReLU), + Sigmoid(Sigmoid), + Tanh(Tanh), + Softmax(Softmax), + LeakyReLU(LeakyReLU), + Elu(ELU), + Gelu(GELU), + Sequential(Sequential), + Conv2d(Conv2d), + BatchNorm1d(BatchNorm1d), + BatchNorm2d(BatchNorm2d), + Dropout(Dropout), + Dropout2d(Dropout2d), +} + +#[pymethods] +impl PyModule { + /// Forward pass through the module + fn forward(&mut self, input: &Bound) -> PyResult { + let input_tensor = borrow_tensor(input)?; + let result = match &mut self.inner { + ModuleType::DenseLayer(layer) => layer.forward(input_tensor.tensor()), + ModuleType::ReLU(layer) => layer.forward(input_tensor.tensor()), + ModuleType::Sigmoid(layer) => layer.forward(input_tensor.tensor()), + ModuleType::Tanh(layer) => layer.forward(input_tensor.tensor()), + ModuleType::Softmax(layer) => layer.forward(input_tensor.tensor()), + ModuleType::LeakyReLU(layer) => layer.forward(input_tensor.tensor()), + ModuleType::Elu(layer) => layer.forward(input_tensor.tensor()), + ModuleType::Gelu(layer) => layer.forward(input_tensor.tensor()), + ModuleType::Sequential(layer) => layer.forward(input_tensor.tensor()), + ModuleType::Conv2d(layer) => layer.forward(input_tensor.tensor()), + ModuleType::BatchNorm1d(layer) => layer.forward(input_tensor.tensor()), + ModuleType::BatchNorm2d(layer) => layer.forward(input_tensor.tensor()), + ModuleType::Dropout(layer) => layer.forward(input_tensor.tensor()), + ModuleType::Dropout2d(layer) => layer.forward(input_tensor.tensor()), + } + .map_err(_convert_error)?; + + Ok(PyTensor::from_tensor(result)) + } + + #[pyo3(name = "__call__")] + fn call(&mut self, input: &Bound) -> PyResult { + self.forward(input) + } + + /// Get all parameters of the module + fn parameters(&self) -> Vec { + let params = match &self.inner { + ModuleType::DenseLayer(layer) => layer.parameters(), + ModuleType::ReLU(layer) => layer.parameters(), + ModuleType::Sigmoid(layer) => layer.parameters(), + ModuleType::Tanh(layer) => layer.parameters(), + ModuleType::Softmax(layer) => layer.parameters(), + ModuleType::LeakyReLU(layer) => layer.parameters(), + ModuleType::Elu(layer) => layer.parameters(), + ModuleType::Gelu(layer) => layer.parameters(), + ModuleType::Sequential(layer) => layer.parameters(), + ModuleType::Conv2d(layer) => layer.parameters(), + ModuleType::BatchNorm1d(layer) => layer.parameters(), + ModuleType::BatchNorm2d(layer) => layer.parameters(), + ModuleType::Dropout(layer) => layer.parameters(), + ModuleType::Dropout2d(layer) => layer.parameters(), + }; + + params + .into_iter() + .map(|tensor| PyTensor::from_tensor(tensor.clone())) + .collect() + } + + /// Set module to training mode + fn train(&mut self) { + match &mut self.inner { + ModuleType::DenseLayer(layer) => layer.train(), + ModuleType::ReLU(layer) => layer.train(), + ModuleType::Sigmoid(layer) => layer.train(), + ModuleType::Tanh(layer) => layer.train(), + ModuleType::Softmax(layer) => layer.train(), + ModuleType::LeakyReLU(layer) => layer.train(), + ModuleType::Elu(layer) => layer.train(), + ModuleType::Gelu(layer) => layer.train(), + ModuleType::Sequential(layer) => layer.train(), + ModuleType::Conv2d(layer) => layer.train(), + ModuleType::BatchNorm1d(layer) => layer.train(), + ModuleType::BatchNorm2d(layer) => layer.train(), + ModuleType::Dropout(layer) => layer.train(), + ModuleType::Dropout2d(layer) => layer.train(), + } + } + + /// Set module to evaluation mode + fn eval(&mut self) { + match &mut self.inner { + ModuleType::DenseLayer(layer) => layer.eval(), + ModuleType::ReLU(layer) => layer.eval(), + ModuleType::Sigmoid(layer) => layer.eval(), + ModuleType::Tanh(layer) => layer.eval(), + ModuleType::Softmax(layer) => layer.eval(), + ModuleType::LeakyReLU(layer) => layer.eval(), + ModuleType::Elu(layer) => layer.eval(), + ModuleType::Gelu(layer) => layer.eval(), + ModuleType::Sequential(layer) => layer.eval(), + ModuleType::Conv2d(layer) => layer.eval(), + ModuleType::BatchNorm1d(layer) => layer.eval(), + ModuleType::BatchNorm2d(layer) => layer.eval(), + ModuleType::Dropout(layer) => layer.eval(), + ModuleType::Dropout2d(layer) => layer.eval(), + } + } + + /// Get number of parameters + fn num_parameters(&self) -> usize { + match &self.inner { + ModuleType::DenseLayer(layer) => layer.num_parameters(), + ModuleType::ReLU(layer) => layer.num_parameters(), + ModuleType::Sigmoid(layer) => layer.num_parameters(), + ModuleType::Tanh(layer) => layer.num_parameters(), + ModuleType::Softmax(layer) => layer.num_parameters(), + ModuleType::LeakyReLU(layer) => layer.num_parameters(), + ModuleType::Elu(layer) => layer.num_parameters(), + ModuleType::Gelu(layer) => layer.num_parameters(), + ModuleType::Sequential(layer) => layer.num_parameters(), + ModuleType::Conv2d(layer) => layer.num_parameters(), + ModuleType::BatchNorm1d(layer) => layer.num_parameters(), + ModuleType::BatchNorm2d(layer) => layer.num_parameters(), + ModuleType::Dropout(layer) => layer.num_parameters(), + ModuleType::Dropout2d(layer) => layer.num_parameters(), + } + } + + /// Get detailed parameter statistics + fn parameter_stats(&self, py: Python) -> PyResult> { + let layer: &dyn Layer = match &self.inner { + ModuleType::DenseLayer(layer) => layer, + ModuleType::ReLU(layer) => layer, + ModuleType::Sigmoid(layer) => layer, + ModuleType::Tanh(layer) => layer, + ModuleType::Softmax(layer) => layer, + ModuleType::LeakyReLU(layer) => layer, + ModuleType::Elu(layer) => layer, + ModuleType::Gelu(layer) => layer, + ModuleType::Sequential(layer) => layer, + ModuleType::Conv2d(layer) => layer, + ModuleType::BatchNorm1d(layer) => layer, + ModuleType::BatchNorm2d(layer) => layer, + ModuleType::Dropout(layer) => layer, + ModuleType::Dropout2d(layer) => layer, + }; + let stats = LayerUtils::parameter_stats(layer); + let dict = PyDict::new(py); + dict.set_item("total_parameters", stats.total_parameters)?; + dict.set_item("trainable_parameters", stats.trainable_parameters)?; + dict.set_item("non_trainable_parameters", stats.non_trainable_parameters)?; + dict.set_item("parameter_count_by_tensor", stats.parameter_count_by_tensor)?; + Ok(dict.into()) + } + + /// Get memory usage information + fn memory_usage(&self, py: Python) -> PyResult> { + let layer: &dyn Layer = match &self.inner { + ModuleType::DenseLayer(layer) => layer, + ModuleType::ReLU(layer) => layer, + ModuleType::Sigmoid(layer) => layer, + ModuleType::Tanh(layer) => layer, + ModuleType::Softmax(layer) => layer, + ModuleType::LeakyReLU(layer) => layer, + ModuleType::Elu(layer) => layer, + ModuleType::Gelu(layer) => layer, + ModuleType::Sequential(layer) => layer, + ModuleType::Conv2d(layer) => layer, + ModuleType::BatchNorm1d(layer) => layer, + ModuleType::BatchNorm2d(layer) => layer, + ModuleType::Dropout(layer) => layer, + ModuleType::Dropout2d(layer) => layer, + }; + let usage = LayerUtils::memory_usage(layer); + let dict = PyDict::new(py); + dict.set_item("total_bytes", usage.total_bytes)?; + let dtype_dict = PyDict::new(py); + for (dtype, bytes) in usage.bytes_by_dtype { + dtype_dict.set_item(format!("{:?}", dtype), bytes)?; + } + dict.set_item("bytes_by_dtype", dtype_dict)?; + Ok(dict.into()) + } + + /// Generate summary + #[pyo3(signature = (name=None))] + fn summary(&self, name: Option<&str>) -> PyResult { + match &self.inner { + ModuleType::Sequential(model) => Ok(SequentialUtils::model_summary(model, name)), + _ => { + let layer: &dyn Layer = match &self.inner { + ModuleType::DenseLayer(layer) => layer, + ModuleType::ReLU(layer) => layer, + ModuleType::Sigmoid(layer) => layer, + ModuleType::Tanh(layer) => layer, + ModuleType::Softmax(layer) => layer, + ModuleType::LeakyReLU(layer) => layer, + ModuleType::Elu(layer) => layer, + ModuleType::Gelu(layer) => layer, + ModuleType::Sequential(layer) => layer, + ModuleType::Conv2d(layer) => layer, + ModuleType::BatchNorm1d(layer) => layer, + ModuleType::BatchNorm2d(layer) => layer, + ModuleType::Dropout(layer) => layer, + ModuleType::Dropout2d(layer) => layer, + }; + let owned; + let layer_name = match name { + Some(n) => n, + None => { + owned = self.__repr__(); + &owned + } + }; + Ok(LayerUtils::layer_summary(layer, layer_name)) + } + } + } + + /// Estimate forward memory usage for Sequential models + fn forward_memory_estimate( + &self, + input_shape: Vec, + batch_size: usize, + py: Python, + ) -> PyResult> { + if let ModuleType::Sequential(model) = &self.inner { + let est = SequentialUtils::estimate_forward_memory(model, &input_shape, batch_size); + let dict = PyDict::new(py); + dict.set_item("parameter_memory", est.parameter_memory)?; + dict.set_item( + "estimated_activation_memory", + est.estimated_activation_memory, + )?; + dict.set_item("estimated_total_memory", est.estimated_total_memory)?; + dict.set_item("input_memory", est.input_memory)?; + Ok(dict.into()) + } else { + Err(PyErr::new::( + "forward_memory_estimate only valid for Sequential modules", + )) + } + } + + /// String representation + fn __repr__(&self) -> String { + match &self.inner { + ModuleType::DenseLayer(layer) => format!( + "DenseLayer(in_features={}, out_features={})", + layer.in_features(), + layer.out_features() + ), + ModuleType::ReLU(_) => "ReLU()".to_string(), + ModuleType::Sigmoid(_) => "Sigmoid()".to_string(), + ModuleType::Tanh(_) => "Tanh()".to_string(), + ModuleType::Softmax(layer) => format!("Softmax(dim={:?})", layer.dim()), + ModuleType::LeakyReLU(layer) => { + format!("LeakyReLU(negative_slope={})", layer.negative_slope()) + } + ModuleType::Elu(layer) => format!("ELU(alpha={})", layer.alpha()), + ModuleType::Gelu(_) => "GELU()".to_string(), + ModuleType::Sequential(_) => "Sequential(...)".to_string(), + ModuleType::Conv2d(layer) => format!( + "Conv2d(in_channels={}, out_channels={}, kernel_size={:?})", + layer.in_channels(), + layer.out_channels(), + layer.kernel_size() + ), + ModuleType::BatchNorm1d(layer) => { + format!("BatchNorm1d(num_features={})", layer.num_features()) + } + ModuleType::BatchNorm2d(layer) => { + format!("BatchNorm2d(num_features={})", layer.num_features()) + } + ModuleType::Dropout(layer) => format!("Dropout(p={})", layer.p()), + ModuleType::Dropout2d(layer) => format!("Dropout2d(p={})", layer.p()), + } + } + + /// Save module state to a file (basic implementation) + fn save(&self, path: &str, format: Option<&str>) -> PyResult<()> { + // Build a SerializedModel with metadata and engine state_dict + use engine::nn::Module as _; + let state = match &self.inner { + ModuleType::DenseLayer(layer) => layer.state_dict(), + ModuleType::ReLU(layer) => layer.state_dict(), + ModuleType::Sigmoid(layer) => layer.state_dict(), + ModuleType::Tanh(layer) => layer.state_dict(), + ModuleType::Softmax(layer) => layer.state_dict(), + ModuleType::LeakyReLU(layer) => layer.state_dict(), + ModuleType::Elu(layer) => layer.state_dict(), + ModuleType::Gelu(layer) => layer.state_dict(), + ModuleType::Sequential(layer) => layer.state_dict(), + ModuleType::Conv2d(layer) => layer.state_dict(), + ModuleType::BatchNorm1d(layer) => layer.state_dict(), + ModuleType::BatchNorm2d(layer) => layer.state_dict(), + ModuleType::Dropout(layer) => layer.state_dict(), + ModuleType::Dropout2d(layer) => layer.state_dict(), + }; + + let metadata = ModelMetadata::new("module".to_string(), "Module".to_string()); + let model = SerializedModel::new(metadata, state); + match format.map(|s| s.to_lowercase()) { + Some(ref s) if s == "json" => { + ModelSerializer::save(&model, path, SerializationFormat::Json) + } + Some(ref s) if s == "bin" || s == "binary" => { + ModelSerializer::save(&model, path, SerializationFormat::Binary) + } + Some(ref s) if s == "msgpack" || s == "messagepack" => { + ModelSerializer::save(&model, path, SerializationFormat::MessagePack) + } + _ => ModelSerializer::save_auto(&model, path), + } + .map_err(_convert_error) + } + + /// Load module state from a file (basic implementation) + #[staticmethod] + fn load_state_from(path: &str, format: Option<&str>) -> PyResult { + let model = match format.map(|s| s.to_lowercase()) { + Some(ref s) if s == "json" => ModelSerializer::load(path, SerializationFormat::Json), + Some(ref s) if s == "bin" || s == "binary" => { + ModelSerializer::load(path, SerializationFormat::Binary) + } + Some(ref s) if s == "msgpack" || s == "messagepack" => { + ModelSerializer::load(path, SerializationFormat::MessagePack) + } + _ => ModelSerializer::load_auto(path), + } + .map_err(_convert_error)?; + Ok(crate::serialization::PyStateDict::from_engine( + model.state_dict, + )) + } + + /// Return a StateDict snapshot of this module + fn state_dict(&self) -> PyStateDict { + use engine::nn::Module as _; + let state = match &self.inner { + ModuleType::DenseLayer(layer) => layer.state_dict(), + ModuleType::ReLU(layer) => layer.state_dict(), + ModuleType::Sigmoid(layer) => layer.state_dict(), + ModuleType::Tanh(layer) => layer.state_dict(), + ModuleType::Softmax(layer) => layer.state_dict(), + ModuleType::LeakyReLU(layer) => layer.state_dict(), + ModuleType::Elu(layer) => layer.state_dict(), + ModuleType::Gelu(layer) => layer.state_dict(), + ModuleType::Sequential(layer) => layer.state_dict(), + ModuleType::Conv2d(layer) => layer.state_dict(), + ModuleType::BatchNorm1d(layer) => layer.state_dict(), + ModuleType::BatchNorm2d(layer) => layer.state_dict(), + ModuleType::Dropout(layer) => layer.state_dict(), + ModuleType::Dropout2d(layer) => layer.state_dict(), + }; + crate::serialization::PyStateDict::from_engine(state) + } + + /// Load a provided StateDict into this module + fn load_state_dict(&mut self, state: &PyStateDict, device: Option<&PyDevice>) -> PyResult<()> { + use engine::nn::Module as _; + let dev = device.map(|d| d.device()); + let sd_ref = crate::serialization::PyStateDict::inner_ref(state); + let res = match &mut self.inner { + ModuleType::DenseLayer(layer) => layer.load_state_dict(sd_ref, dev), + ModuleType::ReLU(layer) => layer.load_state_dict(sd_ref, dev), + ModuleType::Sigmoid(layer) => layer.load_state_dict(sd_ref, dev), + ModuleType::Tanh(layer) => layer.load_state_dict(sd_ref, dev), + ModuleType::Softmax(layer) => layer.load_state_dict(sd_ref, dev), + ModuleType::LeakyReLU(layer) => layer.load_state_dict(sd_ref, dev), + ModuleType::Elu(layer) => layer.load_state_dict(sd_ref, dev), + ModuleType::Gelu(layer) => layer.load_state_dict(sd_ref, dev), + ModuleType::Sequential(layer) => layer.load_state_dict(sd_ref, dev), + ModuleType::Conv2d(layer) => layer.load_state_dict(sd_ref, dev), + ModuleType::BatchNorm1d(layer) => layer.load_state_dict(sd_ref, dev), + ModuleType::BatchNorm2d(layer) => layer.load_state_dict(sd_ref, dev), + ModuleType::Dropout(layer) => layer.load_state_dict(sd_ref, dev), + ModuleType::Dropout2d(layer) => layer.load_state_dict(sd_ref, dev), + }; + res.map_err(_convert_error) + } +} + +impl PyModule { + pub fn from_dense_layer(dense_layer: DenseLayer) -> Self { + Self { + inner: ModuleType::DenseLayer(dense_layer), + } + } + + pub fn from_relu(relu: ReLU) -> Self { + Self { + inner: ModuleType::ReLU(relu), + } + } + + pub fn from_sigmoid(sigmoid: Sigmoid) -> Self { + Self { + inner: ModuleType::Sigmoid(sigmoid), + } + } + + pub fn from_tanh(tanh: Tanh) -> Self { + Self { + inner: ModuleType::Tanh(tanh), + } + } + + pub fn from_softmax(softmax: Softmax) -> Self { + Self { + inner: ModuleType::Softmax(softmax), + } + } + + pub fn from_leaky_relu(leaky_relu: LeakyReLU) -> Self { + Self { + inner: ModuleType::LeakyReLU(leaky_relu), + } + } + + pub fn from_elu(elu: ELU) -> Self { + Self { + inner: ModuleType::Elu(elu), + } + } + + pub fn from_gelu(gelu: GELU) -> Self { + Self { + inner: ModuleType::Gelu(gelu), + } + } + + pub fn from_sequential(sequential: Sequential) -> Self { + Self { + inner: ModuleType::Sequential(sequential), + } + } + + pub fn from_conv2d(conv2d: Conv2d) -> Self { + Self { + inner: ModuleType::Conv2d(conv2d), + } + } + + pub fn from_batch_norm1d(batch_norm1d: BatchNorm1d) -> Self { + Self { + inner: ModuleType::BatchNorm1d(batch_norm1d), + } + } + + pub fn from_batch_norm2d(batch_norm2d: BatchNorm2d) -> Self { + Self { + inner: ModuleType::BatchNorm2d(batch_norm2d), + } + } + + pub fn from_dropout(dropout: Dropout) -> Self { + Self { + inner: ModuleType::Dropout(dropout), + } + } + + pub fn from_dropout2d(dropout: Dropout2d) -> Self { + Self { + inner: ModuleType::Dropout2d(dropout), + } + } + + pub fn to_layer(&self) -> PyResult> { + let layer: Box = match &self.inner { + ModuleType::DenseLayer(layer) => Box::new(layer.clone()), + ModuleType::ReLU(layer) => Box::new(layer.clone()), + ModuleType::Sigmoid(layer) => Box::new(layer.clone()), + ModuleType::Tanh(layer) => Box::new(layer.clone()), + ModuleType::Softmax(layer) => Box::new(layer.clone()), + ModuleType::LeakyReLU(layer) => Box::new(layer.clone()), + ModuleType::Elu(layer) => Box::new(layer.clone()), + ModuleType::Gelu(layer) => Box::new(layer.clone()), + ModuleType::Sequential(_) => { + return Err(PyTypeError::new_err( + "Nested Sequential modules are not supported", + )); + } + ModuleType::Conv2d(layer) => Box::new(layer.clone()), + ModuleType::BatchNorm1d(layer) => Box::new(layer.clone()), + ModuleType::BatchNorm2d(layer) => Box::new(layer.clone()), + ModuleType::Dropout(layer) => Box::new(layer.clone()), + ModuleType::Dropout2d(layer) => Box::new(layer.clone()), + }; + + Ok(layer) + } +} + +/// DenseLayer (fully connected) layer +#[pyclass(name = "DenseLayer", extends = PyModule)] +pub struct PyDenseLayer; + +#[pymethods] +impl PyDenseLayer { + /// Create a new dense layer + #[new] + #[pyo3(signature = (in_features, out_features, bias=None, device=None, dtype=None))] + fn new( + in_features: usize, + out_features: usize, + bias: Option, + device: Option<&PyDevice>, + dtype: Option<&str>, + ) -> PyResult> { + let bias = bias.unwrap_or(true); + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let dtype = dtype::resolve_dtype_arg(dtype)?; + + let dense_layer = DenseLayer::new(in_features, out_features, bias, device, dtype) + .map_err(_convert_error)?; + + Ok(PyClassInitializer::from(PyModule::from_dense_layer(dense_layer)).add_subclass(Self)) + } + + /// Get input features count + #[getter] + fn in_features(slf: PyRef) -> PyResult { + let module = slf.as_ref(); + if let ModuleType::DenseLayer(layer) = &module.inner { + Ok(layer.in_features()) + } else { + Err(PyErr::new::( + "Invalid layer type", + )) + } + } + + /// Get output features count + #[getter] + fn out_features(slf: PyRef) -> PyResult { + let module = slf.as_ref(); + if let ModuleType::DenseLayer(layer) = &module.inner { + Ok(layer.out_features()) + } else { + Err(PyErr::new::( + "Invalid layer type", + )) + } + } + + /// Get weight tensor + #[getter] + fn weight(slf: PyRef) -> PyResult { + let module = slf.as_ref(); + if let ModuleType::DenseLayer(layer) = &module.inner { + Ok(PyTensor::from_tensor(layer.weight().clone())) + } else { + Err(PyErr::new::( + "Invalid layer type", + )) + } + } + + /// Get bias tensor + #[getter] + fn bias(slf: PyRef) -> PyResult> { + let module = slf.as_ref(); + if let ModuleType::DenseLayer(layer) = &module.inner { + Ok(layer.bias().map(|b| PyTensor::from_tensor(b.clone()))) + } else { + Err(PyErr::new::( + "Invalid layer type", + )) + } + } +} + +/// ReLU activation layer +#[pyclass(name = "ReLU", extends = PyModule)] +pub struct PyReLU; diff --git a/bindings/src/numpy_compat.rs b/bindings/src/numpy_compat.rs index 09d9cfc2..a39b5edd 100644 --- a/bindings/src/numpy_compat.rs +++ b/bindings/src/numpy_compat.rs @@ -104,7 +104,8 @@ fn asarray(data: &Bound, dtype: Option<&str>, requires_grad: bool) -> PyR } if tensor.requires_grad() != requires_grad { - tensor.requires_grad_(requires_grad)?; + let inner = tensor.tensor().clone().requires_grad_(requires_grad); + tensor = PyTensor::from_tensor(inner); } Ok(tensor) diff --git a/bindings/src/tensor.rs b/bindings/src/tensor.rs index 5701c6f2..9145e1b8 100644 --- a/bindings/src/tensor.rs +++ b/bindings/src/tensor.rs @@ -1,25 +1,16 @@ -// Copyright (c) Soumyadip Sarkar. +// Copyright (c) 2026 Soumyadip Sarkar. // All rights reserved. // // This source code is licensed under the Apache-style license found in the // LICENSE file in the root directory of this source tree. -include!("tensor/preamble.rs"); -include!("tensor/pytensor/properties.rs"); -include!("tensor/pytensor/operations.rs"); -include!("tensor/pytensor/grad.rs"); -include!("tensor/pytensor/arithmetic.rs"); -include!("tensor/pytensor/comparison_dunder.rs"); -include!("tensor/pytensor/comparison.rs"); -include!("tensor/pytensor/isclose.rs"); -include!("tensor/pytensor/reduction.rs"); -include!("tensor/pytensor/math.rs"); -include!("tensor/pytensor/numpy.rs"); -include!("tensor/pytensor/repr.rs"); -include!("tensor/pytensor/creation/basic.rs"); -include!("tensor/pytensor/creation/like.rs"); -include!("tensor/pytensor/creation/range.rs"); -include!("tensor/pytensor/concat_split.rs"); -include!("tensor/python/args.rs"); -include!("tensor/python/convert.rs"); -include!("tensor/python/numpy.rs"); +//! Python tensor bindings. +//! +//! `preamble` declares `PyTensor` and the shared conversion helpers; the +//! method impls and creation/interop functions live in its child modules so +//! they keep access to the private `inner` field and the shared imports. + +#[path = "tensor/preamble.rs"] +mod preamble; + +pub use self::preamble::*; diff --git a/bindings/src/tensor/preamble.rs b/bindings/src/tensor/preamble.rs index 2f95162a..00385886 100644 --- a/bindings/src/tensor/preamble.rs +++ b/bindings/src/tensor/preamble.rs @@ -1,358 +1,402 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -use crate::device::PyDevice; -use crate::dtype; -use crate::error::_convert_error; -use crate::numpy_compat::cross_impl; -use engine::nn; -use engine::operations::binary::{BinaryOpKind, coerce_binary_operands}; -use engine::operations::reduction::QuantileInterpolation; -use engine::operations::shape_ops::RepeatInterleaveSpec; -use engine::random; -use engine::tensor::{Shape, TensorData}; -use engine::{DataType, Device, MinitensorError, Tensor, TensorIndex}; -use numpy::{PyArray, PyArrayDyn, PyArrayMethods, PyUntypedArrayMethods}; -use once_cell::sync::OnceCell; -use pyo3::conversion::IntoPyObjectExt; -use pyo3::exceptions::{ - PyIndexError, PyNotImplementedError, PyRuntimeError, PyTypeError, PyValueError, -}; -use pyo3::intern; -use pyo3::prelude::*; -use pyo3::types::{ - PyAny, PyBool, PyDict, PyInt, PyList, PyModule, PySequence, PySequenceMethods, PySlice, - PyString, PyTuple, -}; -use pyo3::{Py, PyRefMut}; -use std::borrow::Cow; -use std::cmp::Ordering; -use std::convert::TryFrom; -use std::panic::{self, AssertUnwindSafe}; -use std::sync::Arc; - -fn register_leaf_tensor(tensor: &Tensor) { - if tensor.requires_grad() && tensor.grad_fn().is_none() { - let _ = engine::autograd::add_to_graph(tensor, None); - } -} - -fn extract_wrapped_pytensor(value: &Bound) -> Option { - if let Ok(py_tensor) = value.extract::() { - return Some(py_tensor); - } - - let attr_name = intern!(value.py(), "_tensor"); - if value.hasattr(attr_name).ok()? - && let Ok(inner_attr) = value.getattr(attr_name) - && let Ok(py_tensor) = inner_attr.extract::() - { - return Some(py_tensor); - } - - None -} - -/// Extract integer indices from either an integer tensor or any Python -/// sequence of ints (list, tuple, numpy array, ...). -fn extract_index_vector(indices: &Bound) -> PyResult> { - if let Some(py_tensor) = extract_wrapped_pytensor(indices) { - let tensor = py_tensor.inner.contiguous().map_err(_convert_error)?; - if tensor.ndim() > 1 { - return Err(PyValueError::new_err( - "index tensor must be 0-D or 1-D".to_string(), - )); - } - let values: Vec = match tensor.dtype() { - DataType::Int32 => tensor - .data() - .as_i32_slice() - .ok_or_else(|| PyRuntimeError::new_err("failed to read index tensor data"))? - .iter() - .map(|&v| v as i64) - .collect(), - DataType::Int64 => tensor - .data() - .as_i64_slice() - .ok_or_else(|| PyRuntimeError::new_err("failed to read index tensor data"))? - .to_vec(), - dtype => { - return Err(PyTypeError::new_err(format!( - "index tensor must have an integer dtype, got {dtype:?}", - ))); - } - }; - values - .into_iter() - .map(|v| { - usize::try_from(v).map_err(|_| { - PyValueError::new_err(format!("index {v} is negative; indices must be >= 0")) - }) - }) - .collect() - } else { - let seq = indices.extract::>()?; - seq.into_iter() - .map(|v| { - usize::try_from(v).map_err(|_| { - PyValueError::new_err(format!("index {v} is negative; indices must be >= 0")) - }) - }) - .collect() - } -} - -fn parse_clip_bound(value: Option<&Bound>, name: &str) -> PyResult> { - match value { - None => Ok(None), - Some(bound) => { - if bound.is_none() { - return Ok(None); - } - - if let Ok(val) = bound.extract::() { - Ok(Some(val)) - } else if let Ok(int_val) = bound.extract::() { - Ok(Some(int_val as f64)) - } else { - Err(PyTypeError::new_err(format!( - "{name} must be a real number or None", - ))) - } - } - } -} - -fn extract_real_scalar(value: &Bound, name: &str) -> PyResult { - if let Ok(boolean) = value.extract::() { - return Ok(if boolean { 1.0 } else { 0.0 }); - } - - if let Ok(int_val) = value.extract::() { - return Ok(int_val as f64); - } - - if let Ok(float_val) = value.extract::() { - return Ok(float_val); - } - - Err(PyTypeError::new_err(format!( - "{name} must be a real number or boolean", - ))) -} - -fn parse_quantile_interpolation(mode: Option<&str>) -> PyResult { - let mode = mode.unwrap_or("linear"); - if mode.eq_ignore_ascii_case("linear") { - Ok(QuantileInterpolation::Linear) - } else if mode.eq_ignore_ascii_case("lower") { - Ok(QuantileInterpolation::Lower) - } else if mode.eq_ignore_ascii_case("higher") { - Ok(QuantileInterpolation::Higher) - } else if mode.eq_ignore_ascii_case("midpoint") { - Ok(QuantileInterpolation::Midpoint) - } else if mode.eq_ignore_ascii_case("nearest") { - Ok(QuantileInterpolation::Nearest) - } else { - Err(PyValueError::new_err(format!( - "Invalid interpolation mode '{mode}'. Expected one of: linear, lower, higher, midpoint, nearest", - ))) - } -} - -enum QuantileArg { - Scalar(f64), - Multiple(Vec), -} - -fn parse_quantile_arg(q: &Bound) -> PyResult { - if let Ok(value) = q.extract::() { - return Ok(QuantileArg::Scalar(value)); - } - - if let Ok(values) = q.extract::>() { - if values.is_empty() { - Err(PyValueError::new_err( - "quantile() expected at least one probability value", - )) - } else { - Ok(QuantileArg::Multiple(values)) - } - } else { - Err(PyTypeError::new_err( - "q must be a float or a sequence of floats", - )) - } -} - -#[pyclass(name = "Shape", module = "minitensor._core", from_py_object)] -#[derive(Clone, Debug)] -pub struct ShapeSequence { - dims: Vec, -} - -impl ShapeSequence { - pub fn from_dims>>(dims: D) -> Self { - Self { dims: dims.into() } - } -} - -#[pymethods] -impl ShapeSequence { - #[new] - fn py_new(dims: Vec) -> Self { - Self { dims } - } - - fn __repr__(&self) -> PyResult { - Ok(format!("Shape({:?})", self.dims)) - } - - fn __len__(&self) -> usize { - self.dims.len() - } - - fn __getitem__(&self, index: &Bound) -> PyResult> { - let py = index.py(); - if let Ok(idx) = index.extract::() { - let len = self.dims.len() as isize; - let resolved = if idx < 0 { idx + len } else { idx }; - if resolved < 0 || resolved >= len { - Err(PyIndexError::new_err("Shape index out of range")) - } else { - let value = self.dims[resolved as usize]; - let py_value = i64::try_from(value) - .map_err(|_| PyValueError::new_err("Shape dimension too large"))?; - Ok(PyInt::new(py, py_value).into()) - } - } else if let Ok(slice) = index.cast::() { - let indices = slice.indices(self.dims.len() as isize)?; - let mut values = Vec::with_capacity(indices.slicelength); - let mut current = indices.start; - for _ in 0..indices.slicelength { - values.push(self.dims[current as usize]); - current += indices.step; - } - Ok(Py::new(py, ShapeSequence::from_dims(values))?.into()) - } else { - Err(PyTypeError::new_err( - "Shape indices must be integers or slices", - )) - } - } - - fn __eq__(&self, other: &Bound) -> PyResult { - if let Ok(other_shape) = other.extract::() { - return Ok(self.dims == other_shape.dims); - } - - if let Ok(other_vec) = other.extract::>() { - return Ok(self.dims == other_vec); - } - - Ok(false) - } - - fn to_list(&self) -> Vec { - self.dims.clone() - } - - fn to_tuple<'py>(&self, py: Python<'py>) -> PyResult> { - PyTuple::new(py, &self.dims) - } -} - -/// Python wrapper for Tensor -#[pyclass(name = "Tensor", module = "minitensor._core", from_py_object)] -#[derive(Clone)] -pub struct PyTensor { - inner: Tensor, -} - -impl PyTensor { - /// Get reference to inner tensor - pub fn tensor(&self) -> &Tensor { - &self.inner - } - - /// Get mutable reference to inner tensor - pub fn tensor_mut(&mut self) -> &mut Tensor { - &mut self.inner - } - - /// Create from inner tensor - pub fn from_tensor(tensor: Tensor) -> Self { - // The engine's kernels read tensor storage in contiguous logical - // order, so a non-contiguous view (today only `expand` produces one) - // must be materialised before it becomes visible to Python; otherwise - // every downstream operation would read the wrong elements. - let tensor = if tensor.is_contiguous() { - tensor - } else { - match tensor.contiguous() { - Ok(contiguous) => contiguous, - Err(_) => tensor, - } - }; - register_leaf_tensor(&tensor); - Self { inner: tensor } - } - - pub fn from_python_value(value: &Bound) -> PyResult { - Self::from_python_value_with_dtype(value, dtype::default_dtype()) - } - - pub fn from_python_value_with_dtype(value: &Bound, dtype: DataType) -> PyResult { - if let Some(py_tensor) = extract_wrapped_pytensor(value) { - return Ok(py_tensor); - } - - let tensor = convert_python_data_to_tensor(value, dtype, Device::cpu(), false)?; - Ok(Self::from_tensor(tensor)) - } - - pub fn infer_python_dtype(value: &Bound) -> Option { - infer_python_value_dtype(value) - } - - pub fn max_values(&self, dim: Option, keepdim: bool) -> PyResult { - let result = self.inner.max(dim, keepdim).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn nanmax_values(&self, dim: Option, keepdim: bool) -> PyResult { - let result = self.inner.nanmax(dim, keepdim).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn min_values(&self, dim: Option, keepdim: bool) -> PyResult { - let result = self.inner.min(dim, keepdim).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn nanmin_values(&self, dim: Option, keepdim: bool) -> PyResult { - let result = self.inner.nanmin(dim, keepdim).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn median_with_indices( - &self, - dim: Option, - keepdim: bool, - ) -> PyResult<(Self, Option)> { - match self.inner.median(dim, keepdim) { - Ok((values, indices_opt)) => { - let values_tensor = Self::from_tensor(values); - let indices_tensor = indices_opt.map(Self::from_tensor); - Ok((values_tensor, indices_tensor)) - } - Err(err @ MinitensorError::InvalidArgument { .. }) => { - Err(PyRuntimeError::new_err(err.detailed_message())) - } - Err(err) => Err(_convert_error(err)), - } - } -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +// Child modules hosting the PyTensor method impls and interop helpers. They +// are children (not siblings) so they retain access to the private `inner` +// field and inherit this file's imports via `use super::*`. +#[path = "pytensor/arithmetic.rs"] +mod arithmetic; +#[path = "pytensor/comparison.rs"] +mod comparison; +#[path = "pytensor/comparison_dunder.rs"] +mod comparison_dunder; +#[path = "pytensor/concat_split.rs"] +mod concat_split; +#[path = "pytensor/creation/basic.rs"] +mod creation_basic; +#[path = "pytensor/creation/like.rs"] +mod creation_like; +#[path = "pytensor/creation/range.rs"] +mod creation_range; +#[path = "pytensor/grad.rs"] +mod grad; +#[path = "pytensor/isclose.rs"] +mod isclose; +#[path = "pytensor/math.rs"] +mod math; +#[path = "pytensor/operations.rs"] +mod operations; +#[path = "pytensor/properties.rs"] +mod properties; +#[path = "python/args.rs"] +mod py_args; +#[path = "python/convert.rs"] +mod py_convert; +#[path = "pytensor/numpy.rs"] +mod py_numpy; +#[path = "python/numpy.rs"] +mod py_numpy_interop; +#[path = "pytensor/reduction.rs"] +mod reduction; +#[path = "pytensor/repr.rs"] +mod repr_impl; + +pub(crate) use self::py_args::*; +pub(crate) use self::py_convert::*; +pub(crate) use self::py_numpy_interop::*; + +use crate::device::PyDevice; +use crate::dtype; +use crate::error::_convert_error; +use crate::numpy_compat::cross_impl; +use engine::nn; +use engine::operations::binary::{BinaryOpKind, coerce_binary_operands}; +use engine::operations::reduction::QuantileInterpolation; +use engine::operations::shape_ops::RepeatInterleaveSpec; +use engine::random; +use engine::tensor::{Shape, TensorData}; +use engine::{DataType, Device, MinitensorError, Tensor, TensorIndex}; +use numpy::{PyArray, PyArrayDyn, PyArrayMethods, PyUntypedArrayMethods}; +use once_cell::sync::OnceCell; +use pyo3::conversion::IntoPyObjectExt; +use pyo3::exceptions::{ + PyIndexError, PyNotImplementedError, PyRuntimeError, PyTypeError, PyValueError, +}; +use pyo3::intern; +use pyo3::prelude::*; +use pyo3::types::{ + PyAny, PyBool, PyDict, PyInt, PyList, PyModule, PySequence, PySequenceMethods, PySlice, + PyString, PyTuple, +}; +use pyo3::{Py, PyRefMut}; +use std::borrow::Cow; +use std::cmp::Ordering; +use std::convert::TryFrom; +use std::panic::{self, AssertUnwindSafe}; +use std::sync::Arc; + +fn register_leaf_tensor(tensor: &Tensor) { + if tensor.requires_grad() && tensor.grad_fn().is_none() { + let _ = engine::autograd::add_to_graph(tensor, None); + } +} + +fn extract_wrapped_pytensor(value: &Bound) -> Option { + if let Ok(py_tensor) = value.extract::() { + return Some(py_tensor); + } + + let attr_name = intern!(value.py(), "_tensor"); + if value.hasattr(attr_name).ok()? + && let Ok(inner_attr) = value.getattr(attr_name) + && let Ok(py_tensor) = inner_attr.extract::() + { + return Some(py_tensor); + } + + None +} + +/// Extract integer indices from either an integer tensor or any Python +/// sequence of ints (list, tuple, numpy array, ...). +fn extract_index_vector(indices: &Bound) -> PyResult> { + if let Some(py_tensor) = extract_wrapped_pytensor(indices) { + let tensor = py_tensor.inner.contiguous().map_err(_convert_error)?; + if tensor.ndim() > 1 { + return Err(PyValueError::new_err( + "index tensor must be 0-D or 1-D".to_string(), + )); + } + let values: Vec = match tensor.dtype() { + DataType::Int32 => tensor + .data() + .as_i32_slice() + .ok_or_else(|| PyRuntimeError::new_err("failed to read index tensor data"))? + .iter() + .map(|&v| v as i64) + .collect(), + DataType::Int64 => tensor + .data() + .as_i64_slice() + .ok_or_else(|| PyRuntimeError::new_err("failed to read index tensor data"))? + .to_vec(), + dtype => { + return Err(PyTypeError::new_err(format!( + "index tensor must have an integer dtype, got {dtype:?}", + ))); + } + }; + values + .into_iter() + .map(|v| { + usize::try_from(v).map_err(|_| { + PyValueError::new_err(format!("index {v} is negative; indices must be >= 0")) + }) + }) + .collect() + } else { + let seq = indices.extract::>()?; + seq.into_iter() + .map(|v| { + usize::try_from(v).map_err(|_| { + PyValueError::new_err(format!("index {v} is negative; indices must be >= 0")) + }) + }) + .collect() + } +} + +fn parse_clip_bound(value: Option<&Bound>, name: &str) -> PyResult> { + match value { + None => Ok(None), + Some(bound) => { + if bound.is_none() { + return Ok(None); + } + + if let Ok(val) = bound.extract::() { + Ok(Some(val)) + } else if let Ok(int_val) = bound.extract::() { + Ok(Some(int_val as f64)) + } else { + Err(PyTypeError::new_err(format!( + "{name} must be a real number or None", + ))) + } + } + } +} + +fn extract_real_scalar(value: &Bound, name: &str) -> PyResult { + if let Ok(boolean) = value.extract::() { + return Ok(if boolean { 1.0 } else { 0.0 }); + } + + if let Ok(int_val) = value.extract::() { + return Ok(int_val as f64); + } + + if let Ok(float_val) = value.extract::() { + return Ok(float_val); + } + + Err(PyTypeError::new_err(format!( + "{name} must be a real number or boolean", + ))) +} + +fn parse_quantile_interpolation(mode: Option<&str>) -> PyResult { + let mode = mode.unwrap_or("linear"); + if mode.eq_ignore_ascii_case("linear") { + Ok(QuantileInterpolation::Linear) + } else if mode.eq_ignore_ascii_case("lower") { + Ok(QuantileInterpolation::Lower) + } else if mode.eq_ignore_ascii_case("higher") { + Ok(QuantileInterpolation::Higher) + } else if mode.eq_ignore_ascii_case("midpoint") { + Ok(QuantileInterpolation::Midpoint) + } else if mode.eq_ignore_ascii_case("nearest") { + Ok(QuantileInterpolation::Nearest) + } else { + Err(PyValueError::new_err(format!( + "Invalid interpolation mode '{mode}'. Expected one of: linear, lower, higher, midpoint, nearest", + ))) + } +} + +enum QuantileArg { + Scalar(f64), + Multiple(Vec), +} + +fn parse_quantile_arg(q: &Bound) -> PyResult { + if let Ok(value) = q.extract::() { + return Ok(QuantileArg::Scalar(value)); + } + + if let Ok(values) = q.extract::>() { + if values.is_empty() { + Err(PyValueError::new_err( + "quantile() expected at least one probability value", + )) + } else { + Ok(QuantileArg::Multiple(values)) + } + } else { + Err(PyTypeError::new_err( + "q must be a float or a sequence of floats", + )) + } +} + +#[pyclass(name = "Shape", module = "minitensor._core", from_py_object)] +#[derive(Clone, Debug)] +pub struct ShapeSequence { + dims: Vec, +} + +impl ShapeSequence { + pub fn from_dims>>(dims: D) -> Self { + Self { dims: dims.into() } + } +} + +#[pymethods] +impl ShapeSequence { + #[new] + fn py_new(dims: Vec) -> Self { + Self { dims } + } + + fn __repr__(&self) -> PyResult { + Ok(format!("Shape({:?})", self.dims)) + } + + fn __len__(&self) -> usize { + self.dims.len() + } + + fn __getitem__(&self, index: &Bound) -> PyResult> { + let py = index.py(); + if let Ok(idx) = index.extract::() { + let len = self.dims.len() as isize; + let resolved = if idx < 0 { idx + len } else { idx }; + if resolved < 0 || resolved >= len { + Err(PyIndexError::new_err("Shape index out of range")) + } else { + let value = self.dims[resolved as usize]; + let py_value = i64::try_from(value) + .map_err(|_| PyValueError::new_err("Shape dimension too large"))?; + Ok(PyInt::new(py, py_value).into()) + } + } else if let Ok(slice) = index.cast::() { + let indices = slice.indices(self.dims.len() as isize)?; + let mut values = Vec::with_capacity(indices.slicelength); + let mut current = indices.start; + for _ in 0..indices.slicelength { + values.push(self.dims[current as usize]); + current += indices.step; + } + Ok(Py::new(py, ShapeSequence::from_dims(values))?.into()) + } else { + Err(PyTypeError::new_err( + "Shape indices must be integers or slices", + )) + } + } + + fn __eq__(&self, other: &Bound) -> PyResult { + if let Ok(other_shape) = other.extract::() { + return Ok(self.dims == other_shape.dims); + } + + if let Ok(other_vec) = other.extract::>() { + return Ok(self.dims == other_vec); + } + + Ok(false) + } + + fn to_list(&self) -> Vec { + self.dims.clone() + } + + fn to_tuple<'py>(&self, py: Python<'py>) -> PyResult> { + PyTuple::new(py, &self.dims) + } +} + +/// Python wrapper for Tensor +#[pyclass(name = "Tensor", module = "minitensor._core", from_py_object)] +#[derive(Clone)] +pub struct PyTensor { + inner: Tensor, +} + +impl PyTensor { + /// Get reference to inner tensor + pub fn tensor(&self) -> &Tensor { + &self.inner + } + + /// Get mutable reference to inner tensor + pub fn tensor_mut(&mut self) -> &mut Tensor { + &mut self.inner + } + + /// Create from inner tensor + pub fn from_tensor(tensor: Tensor) -> Self { + // The engine's kernels read tensor storage in contiguous logical + // order, so a non-contiguous view (today only `expand` produces one) + // must be materialised before it becomes visible to Python; otherwise + // every downstream operation would read the wrong elements. + let tensor = if tensor.is_contiguous() { + tensor + } else { + match tensor.contiguous() { + Ok(contiguous) => contiguous, + Err(_) => tensor, + } + }; + register_leaf_tensor(&tensor); + Self { inner: tensor } + } + + pub fn from_python_value(value: &Bound) -> PyResult { + Self::from_python_value_with_dtype(value, dtype::default_dtype()) + } + + pub fn from_python_value_with_dtype(value: &Bound, dtype: DataType) -> PyResult { + if let Some(py_tensor) = extract_wrapped_pytensor(value) { + return Ok(py_tensor); + } + + let tensor = convert_python_data_to_tensor(value, dtype, Device::cpu(), false)?; + Ok(Self::from_tensor(tensor)) + } + + pub fn infer_python_dtype(value: &Bound) -> Option { + infer_python_value_dtype(value) + } + + pub fn max_values(&self, dim: Option, keepdim: bool) -> PyResult { + let result = self.inner.max(dim, keepdim).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn nanmax_values(&self, dim: Option, keepdim: bool) -> PyResult { + let result = self.inner.nanmax(dim, keepdim).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn min_values(&self, dim: Option, keepdim: bool) -> PyResult { + let result = self.inner.min(dim, keepdim).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn nanmin_values(&self, dim: Option, keepdim: bool) -> PyResult { + let result = self.inner.nanmin(dim, keepdim).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn median_with_indices( + &self, + dim: Option, + keepdim: bool, + ) -> PyResult<(Self, Option)> { + match self.inner.median(dim, keepdim) { + Ok((values, indices_opt)) => { + let values_tensor = Self::from_tensor(values); + let indices_tensor = indices_opt.map(Self::from_tensor); + Ok((values_tensor, indices_tensor)) + } + Err(err @ MinitensorError::InvalidArgument { .. }) => { + Err(PyRuntimeError::new_err(err.detailed_message())) + } + Err(err) => Err(_convert_error(err)), + } + } +} diff --git a/bindings/src/tensor/pytensor/arithmetic.rs b/bindings/src/tensor/pytensor/arithmetic.rs index cf7ce9f9..8ffcb033 100644 --- a/bindings/src/tensor/pytensor/arithmetic.rs +++ b/bindings/src/tensor/pytensor/arithmetic.rs @@ -1,78 +1,78 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -#[pymethods] -impl PyTensor { - // Arithmetic operations - fn __neg__(&self) -> PyResult { - use engine::operations::arithmetic::neg; - let result = neg(&self.inner).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn __add__(&self, other: &Bound) -> PyResult { - let (lhs, rhs) = - prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Add)?; - let result = lhs.add(&rhs).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn __radd__(&self, other: &Bound) -> PyResult { - let (lhs, rhs) = - prepare_binary_operands_from_py(&self.inner, other, true, BinaryOpKind::Add)?; - let result = lhs.add(&rhs).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn __sub__(&self, other: &Bound) -> PyResult { - use engine::operations::arithmetic::sub; - let (lhs, rhs) = - prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Sub)?; - let result = sub(&lhs, &rhs).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn __rsub__(&self, other: &Bound) -> PyResult { - use engine::operations::arithmetic::sub; - let (lhs, rhs) = - prepare_binary_operands_from_py(&self.inner, other, true, BinaryOpKind::Sub)?; - let result = sub(&lhs, &rhs).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn __mul__(&self, other: &Bound) -> PyResult { - use engine::operations::arithmetic::mul; - let (lhs, rhs) = - prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Mul)?; - let result = mul(&lhs, &rhs).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn __rmul__(&self, other: &Bound) -> PyResult { - use engine::operations::arithmetic::mul; - let (lhs, rhs) = - prepare_binary_operands_from_py(&self.inner, other, true, BinaryOpKind::Mul)?; - let result = mul(&lhs, &rhs).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn __truediv__(&self, other: &Bound) -> PyResult { - use engine::operations::arithmetic::div; - let (lhs, rhs) = - prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Div)?; - let result = div(&lhs, &rhs).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn __rtruediv__(&self, other: &Bound) -> PyResult { - use engine::operations::arithmetic::div; - let (lhs, rhs) = - prepare_binary_operands_from_py(&self.inner, other, true, BinaryOpKind::Div)?; - let result = div(&lhs, &rhs).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - -} +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +#[pymethods] +impl PyTensor { + // Arithmetic operations + fn __neg__(&self) -> PyResult { + use engine::operations::arithmetic::neg; + let result = neg(&self.inner).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + fn __add__(&self, other: &Bound) -> PyResult { + let (lhs, rhs) = + prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Add)?; + let result = lhs.add(&rhs).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + fn __radd__(&self, other: &Bound) -> PyResult { + let (lhs, rhs) = + prepare_binary_operands_from_py(&self.inner, other, true, BinaryOpKind::Add)?; + let result = lhs.add(&rhs).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + fn __sub__(&self, other: &Bound) -> PyResult { + use engine::operations::arithmetic::sub; + let (lhs, rhs) = + prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Sub)?; + let result = sub(&lhs, &rhs).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + fn __rsub__(&self, other: &Bound) -> PyResult { + use engine::operations::arithmetic::sub; + let (lhs, rhs) = + prepare_binary_operands_from_py(&self.inner, other, true, BinaryOpKind::Sub)?; + let result = sub(&lhs, &rhs).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn __mul__(&self, other: &Bound) -> PyResult { + use engine::operations::arithmetic::mul; + let (lhs, rhs) = + prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Mul)?; + let result = mul(&lhs, &rhs).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn __rmul__(&self, other: &Bound) -> PyResult { + use engine::operations::arithmetic::mul; + let (lhs, rhs) = + prepare_binary_operands_from_py(&self.inner, other, true, BinaryOpKind::Mul)?; + let result = mul(&lhs, &rhs).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + fn __truediv__(&self, other: &Bound) -> PyResult { + use engine::operations::arithmetic::div; + let (lhs, rhs) = + prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Div)?; + let result = div(&lhs, &rhs).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + fn __rtruediv__(&self, other: &Bound) -> PyResult { + use engine::operations::arithmetic::div; + let (lhs, rhs) = + prepare_binary_operands_from_py(&self.inner, other, true, BinaryOpKind::Div)?; + let result = div(&lhs, &rhs).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } +} diff --git a/bindings/src/tensor/pytensor/comparison.rs b/bindings/src/tensor/pytensor/comparison.rs index 8e2bf8a0..8cc6c502 100644 --- a/bindings/src/tensor/pytensor/comparison.rs +++ b/bindings/src/tensor/pytensor/comparison.rs @@ -1,106 +1,105 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -#[pymethods] -impl PyTensor { - // Comparison operations - pub fn eq(&self, other: &Bound) -> PyResult { - self.eq_from_py(other) - } - - pub fn ne(&self, other: &Bound) -> PyResult { - self.ne_from_py(other) - } - - pub fn lt(&self, other: &Bound) -> PyResult { - self.lt_from_py(other) - } - - pub fn le(&self, other: &Bound) -> PyResult { - self.le_from_py(other) - } - - pub fn gt(&self, other: &Bound) -> PyResult { - self.gt_from_py(other) - } - - pub fn ge(&self, other: &Bound) -> PyResult { - self.ge_from_py(other) - } - - fn eq_from_py(&self, other: &Bound) -> PyResult { - let (lhs, rhs) = - prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Add)?; - let result = lhs.eq(&rhs).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn ne_from_py(&self, other: &Bound) -> PyResult { - let (lhs, rhs) = - prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Add)?; - let result = lhs.ne(&rhs).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn lt_from_py(&self, other: &Bound) -> PyResult { - let (lhs, rhs) = - prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Add)?; - let result = lhs.lt(&rhs).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn le_from_py(&self, other: &Bound) -> PyResult { - let (lhs, rhs) = - prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Add)?; - let result = lhs.le(&rhs).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn gt_from_py(&self, other: &Bound) -> PyResult { - let (lhs, rhs) = - prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Add)?; - let result = lhs.gt(&rhs).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn ge_from_py(&self, other: &Bound) -> PyResult { - let (lhs, rhs) = - prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Add)?; - let result = lhs.ge(&rhs).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn array_equal(&self, other: &PyTensor) -> PyResult { - if self.inner.shape() != other.inner.shape() { - return Ok(false); - } - let (lhs, rhs, _) = - coerce_binary_operands(&self.inner, &other.inner, BinaryOpKind::Add) - .map_err(_convert_error)?; - Ok(lhs.array_equal(&rhs)) - } - - #[pyo3(signature = (other, rtol=None, atol=None, equal_nan=false))] - pub fn allclose( - &self, - other: &PyTensor, - rtol: Option, - atol: Option, - equal_nan: bool, - ) -> PyResult { - let rtol = rtol.unwrap_or(1e-5); - let atol = atol.unwrap_or(1e-8); - if !rtol.is_finite() || !atol.is_finite() || rtol < 0.0 || atol < 0.0 { - return Err(PyValueError::new_err( - "rtol and atol must be non-negative, finite values", - )); - } - let (lhs, rhs, _) = - coerce_binary_operands(&self.inner, &other.inner, BinaryOpKind::Add) - .map_err(_convert_error)?; - Ok(lhs.allclose_with_equal_nan(&rhs, rtol, atol, equal_nan)) - } -} +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +#[pymethods] +impl PyTensor { + // Comparison operations + pub fn eq(&self, other: &Bound) -> PyResult { + self.eq_from_py(other) + } + + pub fn ne(&self, other: &Bound) -> PyResult { + self.ne_from_py(other) + } + + pub fn lt(&self, other: &Bound) -> PyResult { + self.lt_from_py(other) + } + + pub fn le(&self, other: &Bound) -> PyResult { + self.le_from_py(other) + } + + pub fn gt(&self, other: &Bound) -> PyResult { + self.gt_from_py(other) + } + + pub fn ge(&self, other: &Bound) -> PyResult { + self.ge_from_py(other) + } + + pub(crate) fn eq_from_py(&self, other: &Bound) -> PyResult { + let (lhs, rhs) = + prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Add)?; + let result = lhs.eq(&rhs).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub(crate) fn ne_from_py(&self, other: &Bound) -> PyResult { + let (lhs, rhs) = + prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Add)?; + let result = lhs.ne(&rhs).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub(crate) fn lt_from_py(&self, other: &Bound) -> PyResult { + let (lhs, rhs) = + prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Add)?; + let result = lhs.lt(&rhs).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub(crate) fn le_from_py(&self, other: &Bound) -> PyResult { + let (lhs, rhs) = + prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Add)?; + let result = lhs.le(&rhs).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub(crate) fn gt_from_py(&self, other: &Bound) -> PyResult { + let (lhs, rhs) = + prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Add)?; + let result = lhs.gt(&rhs).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub(crate) fn ge_from_py(&self, other: &Bound) -> PyResult { + let (lhs, rhs) = + prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Add)?; + let result = lhs.ge(&rhs).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn array_equal(&self, other: &PyTensor) -> PyResult { + if self.inner.shape() != other.inner.shape() { + return Ok(false); + } + let (lhs, rhs, _) = coerce_binary_operands(&self.inner, &other.inner, BinaryOpKind::Add) + .map_err(_convert_error)?; + Ok(lhs.array_equal(&rhs)) + } + + #[pyo3(signature = (other, rtol=None, atol=None, equal_nan=false))] + pub fn allclose( + &self, + other: &PyTensor, + rtol: Option, + atol: Option, + equal_nan: bool, + ) -> PyResult { + let rtol = rtol.unwrap_or(1e-5); + let atol = atol.unwrap_or(1e-8); + if !rtol.is_finite() || !atol.is_finite() || rtol < 0.0 || atol < 0.0 { + return Err(PyValueError::new_err( + "rtol and atol must be non-negative, finite values", + )); + } + let (lhs, rhs, _) = coerce_binary_operands(&self.inner, &other.inner, BinaryOpKind::Add) + .map_err(_convert_error)?; + Ok(lhs.allclose_with_equal_nan(&rhs, rtol, atol, equal_nan)) + } +} diff --git a/bindings/src/tensor/pytensor/comparison_dunder.rs b/bindings/src/tensor/pytensor/comparison_dunder.rs index 8860709e..6105db75 100644 --- a/bindings/src/tensor/pytensor/comparison_dunder.rs +++ b/bindings/src/tensor/pytensor/comparison_dunder.rs @@ -1,237 +1,237 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -#[pymethods] -impl PyTensor { - // Comparison operators as Python dunder methods - fn __eq__(&self, other: &Bound) -> PyResult { - self.eq_from_py(other) - } - - fn __ne__(&self, other: &Bound) -> PyResult { - self.ne_from_py(other) - } - - fn __lt__(&self, other: &Bound) -> PyResult { - self.lt_from_py(other) - } - - fn __le__(&self, other: &Bound) -> PyResult { - self.le_from_py(other) - } - - fn __gt__(&self, other: &Bound) -> PyResult { - self.gt_from_py(other) - } - - fn __ge__(&self, other: &Bound) -> PyResult { - self.ge_from_py(other) - } - - pub fn matmul(&self, other: &Bound) -> PyResult { - let other_tensor = tensor_from_py_value(&self.inner, other)?; - let result = self.inner.matmul(&other_tensor).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn __matmul__(&self, other: &Bound) -> PyResult { - self.matmul(other) - } - - fn __rmatmul__(&self, other: &Bound) -> PyResult { - let other_tensor = tensor_from_py_value(&self.inner, other)?; - let result = other_tensor.matmul(&self.inner).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn solve(&self, rhs: &Bound) -> PyResult { - let rhs_tensor = tensor_from_py_value(&self.inner, rhs)?; - let result = self.inner.solve(&rhs_tensor).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn bmm(&self, other: &Bound) -> PyResult { - let other_tensor = tensor_from_py_value(&self.inner, other)?; - let result = self.inner.bmm(&other_tensor).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn dot(&self, other: &Bound) -> PyResult { - let other_tensor = tensor_from_py_value(&self.inner, other)?; - let result = self.inner.dot(&other_tensor).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (diagonal=0))] - pub fn triu(&self, diagonal: i64) -> PyResult { - let result = self.inner.triu(diagonal).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (diagonal=0))] - pub fn tril(&self, diagonal: i64) -> PyResult { - let result = self.inner.tril(diagonal).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (offset=0, dim1=-2, dim2=-1))] - pub fn diagonal(&self, offset: isize, dim1: isize, dim2: isize) -> PyResult { - let result = self - .inner - .diagonal(offset, dim1, dim2) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (offset=0, dim1=-2, dim2=-1))] - pub fn trace(&self, offset: isize, dim1: isize, dim2: isize) -> PyResult { - let result = self - .inner - .trace(offset, dim1, dim2) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(name = "where")] - pub fn where_method(&self, condition: &Bound, other: &Bound) -> PyResult { - let device = self.inner.device(); - let condition_tensor = tensor_bool_from_py(condition, device)?; - - let other_input = tensor_from_py_value(&self.inner, other)?; - let (input_cast, other_cast, _) = - coerce_binary_operands(&self.inner, &other_input, BinaryOpKind::Add) - .map_err(_convert_error)?; - - let input_tensor = match input_cast { - Cow::Borrowed(_) => self.inner.clone(), - Cow::Owned(tensor) => tensor, - }; - let other_tensor = match other_cast { - Cow::Borrowed(_) => other_input, - Cow::Owned(tensor) => tensor, - }; - - let result = input_tensor - .where_select(&condition_tensor, &other_tensor) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn masked_fill(&self, mask: &Bound, value: &Bound) -> PyResult { - let device = self.inner.device(); - let mask_tensor = tensor_bool_from_py(mask, device)?; - - let mut tensor_value = tensor_from_py_value(&self.inner, value).map_err(|_| { - PyTypeError::new_err("masked_fill value must be a Tensor or numeric scalar") - })?; - - if tensor_value.device() != device { - tensor_value = tensor_value.to(device).map_err(_convert_error)?; - } - - let (input_cast, value_cast, _) = - coerce_binary_operands(&self.inner, &tensor_value, BinaryOpKind::Add) - .map_err(_convert_error)?; - - let input_tensor = match input_cast { - Cow::Borrowed(_) => self.inner.clone(), - Cow::Owned(tensor) => tensor, - }; - let value_tensor = match value_cast { - Cow::Borrowed(_) => tensor_value, - Cow::Owned(tensor) => tensor, - }; - - let result = input_tensor - .masked_fill(&mask_tensor, &value_tensor) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (other, axis=None))] - pub fn cross(&self, other: &Bound, axis: Option) -> PyResult { - let py = other.py(); - - let maybe_tensor = if let Ok(tensor) = other.extract::() { - Some(tensor) - } else if let Ok(attr) = other.getattr(intern!(py, "_tensor")) { - attr.extract::().ok() - } else { - None - }; - - let other_tensor = if let Some(tensor) = maybe_tensor { - tensor - } else { - let dtype = self.inner.dtype(); - let device = self.inner.device(); - let converted = convert_python_data_to_tensor(other, dtype, device, false)?; - PyTensor::from_tensor(converted) - }; - - cross_impl(self, &other_tensor, axis) - } - - pub fn maximum(&self, other: &Bound) -> PyResult { - let (lhs, rhs) = - prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Maximum)?; - let result = lhs.maximum(&rhs).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn minimum(&self, other: &Bound) -> PyResult { - let (lhs, rhs) = - prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Minimum)?; - let result = lhs.minimum(&rhs).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn logaddexp(&self, other: &Bound) -> PyResult { - let (lhs, rhs) = - prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Add)?; - let result = lhs.logaddexp(&rhs).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn _coerce_binary_operands( - &self, - other: &PyTensor, - op: &str, - ) -> PyResult<(PyTensor, PyTensor)> { - let op_kind = match op { - "__add__" | "add" | "logaddexp" => BinaryOpKind::Add, - "__sub__" | "sub" => BinaryOpKind::Sub, - "__mul__" | "mul" => BinaryOpKind::Mul, - "__truediv__" | "div" => BinaryOpKind::Div, - "maximum" => BinaryOpKind::Maximum, - "minimum" => BinaryOpKind::Minimum, - _ => { - return Err(PyValueError::new_err(format!( - "Unsupported binary operation for dtype coercion: {op}" - ))); - } - }; - - let (lhs_cast, rhs_cast, _) = - coerce_binary_operands(self.tensor(), other.tensor(), op_kind) - .map_err(_convert_error)?; - - let lhs_tensor = match lhs_cast { - Cow::Borrowed(_) => self.inner.clone(), - Cow::Owned(tensor) => tensor, - }; - let rhs_tensor = match rhs_cast { - Cow::Borrowed(_) => other.inner.clone(), - Cow::Owned(tensor) => tensor, - }; - - Ok(( - PyTensor::from_tensor(lhs_tensor), - PyTensor::from_tensor(rhs_tensor), - )) - } - -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +#[pymethods] +impl PyTensor { + // Comparison operators as Python dunder methods + fn __eq__(&self, other: &Bound) -> PyResult { + self.eq_from_py(other) + } + + fn __ne__(&self, other: &Bound) -> PyResult { + self.ne_from_py(other) + } + + fn __lt__(&self, other: &Bound) -> PyResult { + self.lt_from_py(other) + } + + fn __le__(&self, other: &Bound) -> PyResult { + self.le_from_py(other) + } + + fn __gt__(&self, other: &Bound) -> PyResult { + self.gt_from_py(other) + } + + fn __ge__(&self, other: &Bound) -> PyResult { + self.ge_from_py(other) + } + + pub fn matmul(&self, other: &Bound) -> PyResult { + let other_tensor = tensor_from_py_value(&self.inner, other)?; + let result = self.inner.matmul(&other_tensor).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + fn __matmul__(&self, other: &Bound) -> PyResult { + self.matmul(other) + } + + fn __rmatmul__(&self, other: &Bound) -> PyResult { + let other_tensor = tensor_from_py_value(&self.inner, other)?; + let result = other_tensor.matmul(&self.inner).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn solve(&self, rhs: &Bound) -> PyResult { + let rhs_tensor = tensor_from_py_value(&self.inner, rhs)?; + let result = self.inner.solve(&rhs_tensor).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn bmm(&self, other: &Bound) -> PyResult { + let other_tensor = tensor_from_py_value(&self.inner, other)?; + let result = self.inner.bmm(&other_tensor).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn dot(&self, other: &Bound) -> PyResult { + let other_tensor = tensor_from_py_value(&self.inner, other)?; + let result = self.inner.dot(&other_tensor).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (diagonal=0))] + pub fn triu(&self, diagonal: i64) -> PyResult { + let result = self.inner.triu(diagonal).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (diagonal=0))] + pub fn tril(&self, diagonal: i64) -> PyResult { + let result = self.inner.tril(diagonal).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (offset=0, dim1=-2, dim2=-1))] + pub fn diagonal(&self, offset: isize, dim1: isize, dim2: isize) -> PyResult { + let result = self + .inner + .diagonal(offset, dim1, dim2) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (offset=0, dim1=-2, dim2=-1))] + pub fn trace(&self, offset: isize, dim1: isize, dim2: isize) -> PyResult { + let result = self + .inner + .trace(offset, dim1, dim2) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(name = "where")] + pub fn where_method(&self, condition: &Bound, other: &Bound) -> PyResult { + let device = self.inner.device(); + let condition_tensor = tensor_bool_from_py(condition, device)?; + + let other_input = tensor_from_py_value(&self.inner, other)?; + let (input_cast, other_cast, _) = + coerce_binary_operands(&self.inner, &other_input, BinaryOpKind::Add) + .map_err(_convert_error)?; + + let input_tensor = match input_cast { + Cow::Borrowed(_) => self.inner.clone(), + Cow::Owned(tensor) => tensor, + }; + let other_tensor = match other_cast { + Cow::Borrowed(_) => other_input, + Cow::Owned(tensor) => tensor, + }; + + let result = input_tensor + .where_select(&condition_tensor, &other_tensor) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn masked_fill(&self, mask: &Bound, value: &Bound) -> PyResult { + let device = self.inner.device(); + let mask_tensor = tensor_bool_from_py(mask, device)?; + + let mut tensor_value = tensor_from_py_value(&self.inner, value).map_err(|_| { + PyTypeError::new_err("masked_fill value must be a Tensor or numeric scalar") + })?; + + if tensor_value.device() != device { + tensor_value = tensor_value.to(device).map_err(_convert_error)?; + } + + let (input_cast, value_cast, _) = + coerce_binary_operands(&self.inner, &tensor_value, BinaryOpKind::Add) + .map_err(_convert_error)?; + + let input_tensor = match input_cast { + Cow::Borrowed(_) => self.inner.clone(), + Cow::Owned(tensor) => tensor, + }; + let value_tensor = match value_cast { + Cow::Borrowed(_) => tensor_value, + Cow::Owned(tensor) => tensor, + }; + + let result = input_tensor + .masked_fill(&mask_tensor, &value_tensor) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (other, axis=None))] + pub fn cross(&self, other: &Bound, axis: Option) -> PyResult { + let py = other.py(); + + let maybe_tensor = if let Ok(tensor) = other.extract::() { + Some(tensor) + } else if let Ok(attr) = other.getattr(intern!(py, "_tensor")) { + attr.extract::().ok() + } else { + None + }; + + let other_tensor = if let Some(tensor) = maybe_tensor { + tensor + } else { + let dtype = self.inner.dtype(); + let device = self.inner.device(); + let converted = convert_python_data_to_tensor(other, dtype, device, false)?; + PyTensor::from_tensor(converted) + }; + + cross_impl(self, &other_tensor, axis) + } + + pub fn maximum(&self, other: &Bound) -> PyResult { + let (lhs, rhs) = + prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Maximum)?; + let result = lhs.maximum(&rhs).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn minimum(&self, other: &Bound) -> PyResult { + let (lhs, rhs) = + prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Minimum)?; + let result = lhs.minimum(&rhs).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn logaddexp(&self, other: &Bound) -> PyResult { + let (lhs, rhs) = + prepare_binary_operands_from_py(&self.inner, other, false, BinaryOpKind::Add)?; + let result = lhs.logaddexp(&rhs).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn _coerce_binary_operands( + &self, + other: &PyTensor, + op: &str, + ) -> PyResult<(PyTensor, PyTensor)> { + let op_kind = match op { + "__add__" | "add" | "logaddexp" => BinaryOpKind::Add, + "__sub__" | "sub" => BinaryOpKind::Sub, + "__mul__" | "mul" => BinaryOpKind::Mul, + "__truediv__" | "div" => BinaryOpKind::Div, + "maximum" => BinaryOpKind::Maximum, + "minimum" => BinaryOpKind::Minimum, + _ => { + return Err(PyValueError::new_err(format!( + "Unsupported binary operation for dtype coercion: {op}" + ))); + } + }; + + let (lhs_cast, rhs_cast, _) = + coerce_binary_operands(self.tensor(), other.tensor(), op_kind) + .map_err(_convert_error)?; + + let lhs_tensor = match lhs_cast { + Cow::Borrowed(_) => self.inner.clone(), + Cow::Owned(tensor) => tensor, + }; + let rhs_tensor = match rhs_cast { + Cow::Borrowed(_) => other.inner.clone(), + Cow::Owned(tensor) => tensor, + }; + + Ok(( + PyTensor::from_tensor(lhs_tensor), + PyTensor::from_tensor(rhs_tensor), + )) + } +} diff --git a/bindings/src/tensor/pytensor/concat_split.rs b/bindings/src/tensor/pytensor/concat_split.rs index 9af375d6..6d58465d 100644 --- a/bindings/src/tensor/pytensor/concat_split.rs +++ b/bindings/src/tensor/pytensor/concat_split.rs @@ -1,189 +1,190 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -#[pymethods] -impl PyTensor { - /// Concatenate tensors along an axis - #[staticmethod] - pub fn concatenate(tensors: &Bound, _axis: Option) -> PyResult { - if tensors.is_empty() { - return Err(PyErr::new::( - "Cannot concatenate empty list of tensors", - )); - } - - let axis = _axis.unwrap_or(0); - - let tensor_vec: Vec = tensors - .iter() - .map(|obj| PyTensor::from_python_value(&obj).map(|t| t.inner.clone())) - .collect::>()?; - - let tensor_refs: Vec<&Tensor> = tensor_vec.iter().collect(); - let result = engine::operations::shape_ops::concatenate(&tensor_refs, axis) - .map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) - } - - /// Stack tensors along a new axis - #[staticmethod] - pub fn stack(tensors: &Bound, _axis: Option) -> PyResult { - if tensors.is_empty() { - return Err(PyErr::new::( - "Cannot stack empty list of tensors", - )); - } - - let axis = _axis.unwrap_or(0); - - let unsqueezed: Vec = tensors - .iter() - .map(|obj| { - let t = PyTensor::from_python_value(&obj)?; - engine::operations::shape_ops::unsqueeze(&t.inner, axis).map_err(_convert_error) - }) - .collect::>()?; - - let refs: Vec<&Tensor> = unsqueezed.iter().collect(); - let result = - engine::operations::shape_ops::concatenate(&refs, axis).map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) - } - - /// Select elements along a dimension using integer indices - /// (a Python sequence or an integer tensor) - pub fn index_select(&self, dim: isize, indices: &Bound) -> PyResult { - let idx_vec = extract_index_vector(indices)?; - let result = engine::operations::shape_ops::index_select(&self.inner, dim, &idx_vec) - .map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) - } - - /// Gather elements along a dimension using an index tensor - pub fn gather(&self, dim: isize, index: &PyTensor) -> PyResult { - let result = engine::operations::shape_ops::gather(&self.inner, dim, &index.inner) - .map_err(_convert_error)?; - Ok(PyTensor::from_tensor(result)) - } - - /// Split tensor into multiple sub-tensors of equal size (``chunk``) - #[pyo3(signature = (sections, dim=0))] - pub fn chunk(&self, sections: usize, dim: isize) -> PyResult> { - if sections == 0 { - return Err(PyErr::new::( - "Sections must be greater than zero", - )); - } - - let ndim = self.inner.ndim() as isize; - let axis = if dim < 0 { dim + ndim } else { dim }; - if axis < 0 || axis >= ndim { - return Err(PyErr::new::(format!( - "Dimension {} out of range", - axis - ))); - } - - let dim_size = self.inner.shape().dims()[axis as usize]; - if !dim_size.is_multiple_of(sections) { - return Err(PyErr::new::( - "Tensor cannot be evenly split along the given axis", - )); - } - - let chunk_size = dim_size / sections; - let section_vec = vec![chunk_size; sections]; - self.split_with_sections(section_vec, axis as usize) - } - - /// Split tensor by chunk size or explicit sections along an axis - #[pyo3(signature = (split_size_or_sections, dim=0))] - pub fn split( - &self, - split_size_or_sections: &Bound, - dim: Option, - ) -> PyResult> { - let dim = dim.unwrap_or(0); - let ndim = self.inner.ndim() as isize; - let dim = if dim < 0 { dim + ndim } else { dim }; - if dim < 0 || dim >= ndim { - return Err(PyErr::new::(format!( - "Dimension {} out of range", - dim - ))); - } - let axis = dim as usize; - let dim_size = self.inner.shape().dims()[axis]; - - let mut sections: Vec = Vec::new(); - - if let Ok(split_size) = split_size_or_sections.extract::() { - if split_size == 0 { - return Err(PyErr::new::( - "split_size must be greater than zero", - )); - } - let mut remaining = dim_size; - while remaining > 0 { - let chunk = split_size.min(remaining); - sections.push(chunk); - remaining -= chunk; - } - } else if let Ok(list) = split_size_or_sections.cast::() { - for obj in list.iter() { - let size: usize = obj.extract()?; - if size == 0 { - return Err(PyErr::new::( - "section size must be greater than zero", - )); - } - sections.push(size); - } - let total: usize = sections.iter().sum(); - if total != dim_size { - return Err(PyErr::new::( - "split sizes do not sum to dimension size", - )); - } - } else if let Ok(tuple) = split_size_or_sections.cast::() { - for obj in tuple.iter() { - let size: usize = obj.extract()?; - if size == 0 { - return Err(PyErr::new::( - "section size must be greater than zero", - )); - } - sections.push(size); - } - let total: usize = sections.iter().sum(); - if total != dim_size { - return Err(PyErr::new::( - "split sizes do not sum to dimension size", - )); - } - } else { - return Err(PyErr::new::( - "split_size_or_sections must be int or sequence", - )); - } - - self.split_with_sections(sections, axis) - } - - fn split_with_sections(&self, sections: Vec, axis: usize) -> PyResult> { - let mut outputs = Vec::with_capacity(sections.len()); - let mut start = 0; - for size in sections { - let end = start + size; - let slice = - engine::operations::shape_ops::slice(&self.inner, axis as isize, start, end, 1) - .map_err(_convert_error)?; - outputs.push(PyTensor::from_tensor(slice)); - start = end; - } - Ok(outputs) - } -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +#[pymethods] +impl PyTensor { + /// Concatenate tensors along an axis + #[staticmethod] + pub fn concatenate(tensors: &Bound, _axis: Option) -> PyResult { + if tensors.is_empty() { + return Err(PyErr::new::( + "Cannot concatenate empty list of tensors", + )); + } + + let axis = _axis.unwrap_or(0); + + let tensor_vec: Vec = tensors + .iter() + .map(|obj| PyTensor::from_python_value(&obj).map(|t| t.inner.clone())) + .collect::>()?; + + let tensor_refs: Vec<&Tensor> = tensor_vec.iter().collect(); + let result = engine::operations::shape_ops::concatenate(&tensor_refs, axis) + .map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) + } + + /// Stack tensors along a new axis + #[staticmethod] + pub fn stack(tensors: &Bound, _axis: Option) -> PyResult { + if tensors.is_empty() { + return Err(PyErr::new::( + "Cannot stack empty list of tensors", + )); + } + + let axis = _axis.unwrap_or(0); + + let unsqueezed: Vec = tensors + .iter() + .map(|obj| { + let t = PyTensor::from_python_value(&obj)?; + engine::operations::shape_ops::unsqueeze(&t.inner, axis).map_err(_convert_error) + }) + .collect::>()?; + + let refs: Vec<&Tensor> = unsqueezed.iter().collect(); + let result = + engine::operations::shape_ops::concatenate(&refs, axis).map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) + } + + /// Select elements along a dimension using integer indices + /// (a Python sequence or an integer tensor) + pub fn index_select(&self, dim: isize, indices: &Bound) -> PyResult { + let idx_vec = extract_index_vector(indices)?; + let result = engine::operations::shape_ops::index_select(&self.inner, dim, &idx_vec) + .map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) + } + + /// Gather elements along a dimension using an index tensor + pub fn gather(&self, dim: isize, index: &PyTensor) -> PyResult { + let result = engine::operations::shape_ops::gather(&self.inner, dim, &index.inner) + .map_err(_convert_error)?; + Ok(PyTensor::from_tensor(result)) + } + + /// Split tensor into multiple sub-tensors of equal size (``chunk``) + #[pyo3(signature = (sections, dim=0))] + pub fn chunk(&self, sections: usize, dim: isize) -> PyResult> { + if sections == 0 { + return Err(PyErr::new::( + "Sections must be greater than zero", + )); + } + + let ndim = self.inner.ndim() as isize; + let axis = if dim < 0 { dim + ndim } else { dim }; + if axis < 0 || axis >= ndim { + return Err(PyErr::new::(format!( + "Dimension {} out of range", + axis + ))); + } + + let dim_size = self.inner.shape().dims()[axis as usize]; + if !dim_size.is_multiple_of(sections) { + return Err(PyErr::new::( + "Tensor cannot be evenly split along the given axis", + )); + } + + let chunk_size = dim_size / sections; + let section_vec = vec![chunk_size; sections]; + self.split_with_sections(section_vec, axis as usize) + } + + /// Split tensor by chunk size or explicit sections along an axis + #[pyo3(signature = (split_size_or_sections, dim=0))] + pub fn split( + &self, + split_size_or_sections: &Bound, + dim: Option, + ) -> PyResult> { + let dim = dim.unwrap_or(0); + let ndim = self.inner.ndim() as isize; + let dim = if dim < 0 { dim + ndim } else { dim }; + if dim < 0 || dim >= ndim { + return Err(PyErr::new::(format!( + "Dimension {} out of range", + dim + ))); + } + let axis = dim as usize; + let dim_size = self.inner.shape().dims()[axis]; + + let mut sections: Vec = Vec::new(); + + if let Ok(split_size) = split_size_or_sections.extract::() { + if split_size == 0 { + return Err(PyErr::new::( + "split_size must be greater than zero", + )); + } + let mut remaining = dim_size; + while remaining > 0 { + let chunk = split_size.min(remaining); + sections.push(chunk); + remaining -= chunk; + } + } else if let Ok(list) = split_size_or_sections.cast::() { + for obj in list.iter() { + let size: usize = obj.extract()?; + if size == 0 { + return Err(PyErr::new::( + "section size must be greater than zero", + )); + } + sections.push(size); + } + let total: usize = sections.iter().sum(); + if total != dim_size { + return Err(PyErr::new::( + "split sizes do not sum to dimension size", + )); + } + } else if let Ok(tuple) = split_size_or_sections.cast::() { + for obj in tuple.iter() { + let size: usize = obj.extract()?; + if size == 0 { + return Err(PyErr::new::( + "section size must be greater than zero", + )); + } + sections.push(size); + } + let total: usize = sections.iter().sum(); + if total != dim_size { + return Err(PyErr::new::( + "split sizes do not sum to dimension size", + )); + } + } else { + return Err(PyErr::new::( + "split_size_or_sections must be int or sequence", + )); + } + + self.split_with_sections(sections, axis) + } + + fn split_with_sections(&self, sections: Vec, axis: usize) -> PyResult> { + let mut outputs = Vec::with_capacity(sections.len()); + let mut start = 0; + for size in sections { + let end = start + size; + let slice = + engine::operations::shape_ops::slice(&self.inner, axis as isize, start, end, 1) + .map_err(_convert_error)?; + outputs.push(PyTensor::from_tensor(slice)); + start = end; + } + Ok(outputs) + } +} diff --git a/bindings/src/tensor/pytensor/creation/basic.rs b/bindings/src/tensor/pytensor/creation/basic.rs index 4e4bcab7..d1929afd 100644 --- a/bindings/src/tensor/pytensor/creation/basic.rs +++ b/bindings/src/tensor/pytensor/creation/basic.rs @@ -1,522 +1,522 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -#[pymethods] -impl PyTensor { - // Static tensor creation methods - #[staticmethod] - #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] - pub fn empty( - shape: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dims = parse_shape_tuple(shape, "shape")?; - let dtype = dtype::resolve_dtype_arg(dtype)?; - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - - let shape = Shape::new(dims); - let tensor = Tensor::empty(shape, dtype, device, requires_grad); - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] - pub fn zeros( - shape: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dims = parse_shape_tuple(shape, "shape")?; - let dtype = dtype::resolve_dtype_arg(dtype)?; - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - - let shape = Shape::new(dims); - let tensor = Tensor::zeros(shape, dtype, device, requires_grad); - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] - pub fn ones( - shape: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dims = parse_shape_tuple(shape, "shape")?; - let dtype = dtype::resolve_dtype_arg(dtype)?; - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - - let shape = Shape::new(dims); - let tensor = Tensor::ones(shape, dtype, device, requires_grad); - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (*shape, low=0.0, high=1.0, dtype=None, device=None, requires_grad=false))] - fn uniform( - shape: &Bound, - low: f64, - high: f64, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dims = parse_shape_tuple(shape, "shape")?; - let dtype = dtype::resolve_dtype_arg(dtype)?; - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - - let shape = Shape::new(dims); - let tensor = create_uniform_tensor(shape, dtype, device, requires_grad, low, high)?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] - fn xavier_uniform( - shape: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dims = parse_shape_tuple(shape, "shape")?; - let dtype = dtype::resolve_dtype_arg(dtype)?; - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - - let shape = Shape::new(dims); - let tensor = create_fan_init_tensor( - shape, - dtype, - device, - requires_grad, - FanInitKind::XavierUniform, - "xavier_uniform", - )?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] - fn xavier_normal( - shape: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dims = parse_shape_tuple(shape, "shape")?; - let dtype = dtype::resolve_dtype_arg(dtype)?; - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - - let shape = Shape::new(dims); - let tensor = create_fan_init_tensor( - shape, - dtype, - device, - requires_grad, - FanInitKind::XavierNormal, - "xavier_normal", - )?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] - fn he_uniform( - shape: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dims = parse_shape_tuple(shape, "shape")?; - let dtype = dtype::resolve_dtype_arg(dtype)?; - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - - let shape = Shape::new(dims); - let tensor = create_fan_init_tensor( - shape, - dtype, - device, - requires_grad, - FanInitKind::HeUniform, - "he_uniform", - )?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] - fn he_normal( - shape: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dims = parse_shape_tuple(shape, "shape")?; - let dtype = dtype::resolve_dtype_arg(dtype)?; - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - - let shape = Shape::new(dims); - let tensor = create_fan_init_tensor( - shape, - dtype, - device, - requires_grad, - FanInitKind::HeNormal, - "he_normal", - )?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] - fn lecun_uniform( - shape: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dims = parse_shape_tuple(shape, "shape")?; - let dtype = dtype::resolve_dtype_arg(dtype)?; - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - - let shape = Shape::new(dims); - let tensor = create_fan_init_tensor( - shape, - dtype, - device, - requires_grad, - FanInitKind::LecunUniform, - "lecun_uniform", - )?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] - fn lecun_normal( - shape: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dims = parse_shape_tuple(shape, "shape")?; - let dtype = dtype::resolve_dtype_arg(dtype)?; - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - - let shape = Shape::new(dims); - let tensor = create_fan_init_tensor( - shape, - dtype, - device, - requires_grad, - FanInitKind::LecunNormal, - "lecun_normal", - )?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] - fn rand( - shape: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dims = parse_shape_tuple(shape, "shape")?; - let dtype = dtype::resolve_dtype_arg(dtype)?; - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - - let shape = Shape::new(dims); - let tensor = create_random_tensor(shape, dtype, device, requires_grad, false)?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] - fn randn( - shape: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dims = parse_shape_tuple(shape, "shape")?; - let dtype = dtype::resolve_dtype_arg(dtype)?; - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - - let shape = Shape::new(dims); - let tensor = create_random_tensor(shape, dtype, device, requires_grad, true)?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (*shape, mean=0.0, std=1.0, lower=None, upper=None, dtype=None, device=None, requires_grad=false))] - #[allow(clippy::too_many_arguments)] - fn truncated_normal( - shape: &Bound, - mean: f64, - std: f64, - lower: Option, - upper: Option, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dims = parse_shape_tuple(shape, "shape")?; - let dtype = dtype::resolve_dtype_arg(dtype)?; - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - - let shape = Shape::new(dims); - let tensor = create_truncated_normal_tensor( - shape, - dtype, - device, - requires_grad, - mean, - std, - lower, - upper, - "truncated_normal", - )?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (input, low=0.0, high=1.0, dtype=None, device=None, requires_grad=None))] - fn uniform_like( - input: &Bound, - low: f64, - high: f64, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let reference = PyTensor::from_python_value(input)?; - let reference_tensor = reference.tensor(); - - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => reference_tensor.dtype(), - }; - - let device = device - .map(|d| d.device()) - .unwrap_or_else(|| reference_tensor.device()); - let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); - let shape = Shape::new(reference.shape_vec()); - let tensor = create_uniform_tensor(shape, dtype, device, requires_grad, low, high)?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] - fn xavier_uniform_like( - input: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let reference = PyTensor::from_python_value(input)?; - let reference_tensor = reference.tensor(); - - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => reference_tensor.dtype(), - }; - - let device = device - .map(|d| d.device()) - .unwrap_or_else(|| reference_tensor.device()); - let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); - let shape = Shape::new(reference.shape_vec()); - let tensor = create_fan_init_tensor( - shape, - dtype, - device, - requires_grad, - FanInitKind::XavierUniform, - "xavier_uniform_like", - )?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] - fn xavier_normal_like( - input: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let reference = PyTensor::from_python_value(input)?; - let reference_tensor = reference.tensor(); - - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => reference_tensor.dtype(), - }; - - let device = device - .map(|d| d.device()) - .unwrap_or_else(|| reference_tensor.device()); - let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); - let shape = Shape::new(reference.shape_vec()); - let tensor = create_fan_init_tensor( - shape, - dtype, - device, - requires_grad, - FanInitKind::XavierNormal, - "xavier_normal_like", - )?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] - fn he_uniform_like( - input: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let reference = PyTensor::from_python_value(input)?; - let reference_tensor = reference.tensor(); - - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => reference_tensor.dtype(), - }; - - let device = device - .map(|d| d.device()) - .unwrap_or_else(|| reference_tensor.device()); - let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); - let shape = Shape::new(reference.shape_vec()); - let tensor = create_fan_init_tensor( - shape, - dtype, - device, - requires_grad, - FanInitKind::HeUniform, - "he_uniform_like", - )?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] - fn he_normal_like( - input: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let reference = PyTensor::from_python_value(input)?; - let reference_tensor = reference.tensor(); - - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => reference_tensor.dtype(), - }; - - let device = device - .map(|d| d.device()) - .unwrap_or_else(|| reference_tensor.device()); - let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); - let shape = Shape::new(reference.shape_vec()); - let tensor = create_fan_init_tensor( - shape, - dtype, - device, - requires_grad, - FanInitKind::HeNormal, - "he_normal_like", - )?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] - fn lecun_uniform_like( - input: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let reference = PyTensor::from_python_value(input)?; - let reference_tensor = reference.tensor(); - - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => reference_tensor.dtype(), - }; - - let device = device - .map(|d| d.device()) - .unwrap_or_else(|| reference_tensor.device()); - let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); - let shape = Shape::new(reference.shape_vec()); - let tensor = create_fan_init_tensor( - shape, - dtype, - device, - requires_grad, - FanInitKind::LecunUniform, - "lecun_uniform_like", - )?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] - fn lecun_normal_like( - input: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let reference = PyTensor::from_python_value(input)?; - let reference_tensor = reference.tensor(); - - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => reference_tensor.dtype(), - }; - - let device = device - .map(|d| d.device()) - .unwrap_or_else(|| reference_tensor.device()); - let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); - let shape = Shape::new(reference.shape_vec()); - let tensor = create_fan_init_tensor( - shape, - dtype, - device, - requires_grad, - FanInitKind::LecunNormal, - "lecun_normal_like", - )?; - Ok(Self::from_tensor(tensor)) - } - -} +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +#[pymethods] +impl PyTensor { + // Static tensor creation methods + #[staticmethod] + #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] + pub fn empty( + shape: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dims = parse_shape_tuple(shape, "shape")?; + let dtype = dtype::resolve_dtype_arg(dtype)?; + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + + let shape = Shape::new(dims); + let tensor = Tensor::empty(shape, dtype, device, requires_grad); + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] + pub fn zeros( + shape: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dims = parse_shape_tuple(shape, "shape")?; + let dtype = dtype::resolve_dtype_arg(dtype)?; + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + + let shape = Shape::new(dims); + let tensor = Tensor::zeros(shape, dtype, device, requires_grad); + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] + pub fn ones( + shape: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dims = parse_shape_tuple(shape, "shape")?; + let dtype = dtype::resolve_dtype_arg(dtype)?; + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + + let shape = Shape::new(dims); + let tensor = Tensor::ones(shape, dtype, device, requires_grad); + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (*shape, low=0.0, high=1.0, dtype=None, device=None, requires_grad=false))] + fn uniform( + shape: &Bound, + low: f64, + high: f64, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dims = parse_shape_tuple(shape, "shape")?; + let dtype = dtype::resolve_dtype_arg(dtype)?; + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + + let shape = Shape::new(dims); + let tensor = create_uniform_tensor(shape, dtype, device, requires_grad, low, high)?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] + fn xavier_uniform( + shape: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dims = parse_shape_tuple(shape, "shape")?; + let dtype = dtype::resolve_dtype_arg(dtype)?; + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + + let shape = Shape::new(dims); + let tensor = create_fan_init_tensor( + shape, + dtype, + device, + requires_grad, + FanInitKind::XavierUniform, + "xavier_uniform", + )?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] + fn xavier_normal( + shape: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dims = parse_shape_tuple(shape, "shape")?; + let dtype = dtype::resolve_dtype_arg(dtype)?; + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + + let shape = Shape::new(dims); + let tensor = create_fan_init_tensor( + shape, + dtype, + device, + requires_grad, + FanInitKind::XavierNormal, + "xavier_normal", + )?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] + fn he_uniform( + shape: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dims = parse_shape_tuple(shape, "shape")?; + let dtype = dtype::resolve_dtype_arg(dtype)?; + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + + let shape = Shape::new(dims); + let tensor = create_fan_init_tensor( + shape, + dtype, + device, + requires_grad, + FanInitKind::HeUniform, + "he_uniform", + )?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] + fn he_normal( + shape: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dims = parse_shape_tuple(shape, "shape")?; + let dtype = dtype::resolve_dtype_arg(dtype)?; + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + + let shape = Shape::new(dims); + let tensor = create_fan_init_tensor( + shape, + dtype, + device, + requires_grad, + FanInitKind::HeNormal, + "he_normal", + )?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] + fn lecun_uniform( + shape: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dims = parse_shape_tuple(shape, "shape")?; + let dtype = dtype::resolve_dtype_arg(dtype)?; + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + + let shape = Shape::new(dims); + let tensor = create_fan_init_tensor( + shape, + dtype, + device, + requires_grad, + FanInitKind::LecunUniform, + "lecun_uniform", + )?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] + fn lecun_normal( + shape: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dims = parse_shape_tuple(shape, "shape")?; + let dtype = dtype::resolve_dtype_arg(dtype)?; + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + + let shape = Shape::new(dims); + let tensor = create_fan_init_tensor( + shape, + dtype, + device, + requires_grad, + FanInitKind::LecunNormal, + "lecun_normal", + )?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] + fn rand( + shape: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dims = parse_shape_tuple(shape, "shape")?; + let dtype = dtype::resolve_dtype_arg(dtype)?; + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + + let shape = Shape::new(dims); + let tensor = create_random_tensor(shape, dtype, device, requires_grad, false)?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (*shape, dtype=None, device=None, requires_grad=false))] + fn randn( + shape: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dims = parse_shape_tuple(shape, "shape")?; + let dtype = dtype::resolve_dtype_arg(dtype)?; + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + + let shape = Shape::new(dims); + let tensor = create_random_tensor(shape, dtype, device, requires_grad, true)?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (*shape, mean=0.0, std=1.0, lower=None, upper=None, dtype=None, device=None, requires_grad=false))] + #[allow(clippy::too_many_arguments)] + fn truncated_normal( + shape: &Bound, + mean: f64, + std: f64, + lower: Option, + upper: Option, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dims = parse_shape_tuple(shape, "shape")?; + let dtype = dtype::resolve_dtype_arg(dtype)?; + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + + let shape = Shape::new(dims); + let tensor = create_truncated_normal_tensor( + shape, + dtype, + device, + requires_grad, + mean, + std, + lower, + upper, + "truncated_normal", + )?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (input, low=0.0, high=1.0, dtype=None, device=None, requires_grad=None))] + fn uniform_like( + input: &Bound, + low: f64, + high: f64, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let reference = PyTensor::from_python_value(input)?; + let reference_tensor = reference.tensor(); + + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => reference_tensor.dtype(), + }; + + let device = device + .map(|d| d.device()) + .unwrap_or_else(|| reference_tensor.device()); + let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); + let shape = Shape::new(reference.shape_vec()); + let tensor = create_uniform_tensor(shape, dtype, device, requires_grad, low, high)?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] + fn xavier_uniform_like( + input: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let reference = PyTensor::from_python_value(input)?; + let reference_tensor = reference.tensor(); + + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => reference_tensor.dtype(), + }; + + let device = device + .map(|d| d.device()) + .unwrap_or_else(|| reference_tensor.device()); + let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); + let shape = Shape::new(reference.shape_vec()); + let tensor = create_fan_init_tensor( + shape, + dtype, + device, + requires_grad, + FanInitKind::XavierUniform, + "xavier_uniform_like", + )?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] + fn xavier_normal_like( + input: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let reference = PyTensor::from_python_value(input)?; + let reference_tensor = reference.tensor(); + + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => reference_tensor.dtype(), + }; + + let device = device + .map(|d| d.device()) + .unwrap_or_else(|| reference_tensor.device()); + let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); + let shape = Shape::new(reference.shape_vec()); + let tensor = create_fan_init_tensor( + shape, + dtype, + device, + requires_grad, + FanInitKind::XavierNormal, + "xavier_normal_like", + )?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] + fn he_uniform_like( + input: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let reference = PyTensor::from_python_value(input)?; + let reference_tensor = reference.tensor(); + + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => reference_tensor.dtype(), + }; + + let device = device + .map(|d| d.device()) + .unwrap_or_else(|| reference_tensor.device()); + let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); + let shape = Shape::new(reference.shape_vec()); + let tensor = create_fan_init_tensor( + shape, + dtype, + device, + requires_grad, + FanInitKind::HeUniform, + "he_uniform_like", + )?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] + fn he_normal_like( + input: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let reference = PyTensor::from_python_value(input)?; + let reference_tensor = reference.tensor(); + + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => reference_tensor.dtype(), + }; + + let device = device + .map(|d| d.device()) + .unwrap_or_else(|| reference_tensor.device()); + let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); + let shape = Shape::new(reference.shape_vec()); + let tensor = create_fan_init_tensor( + shape, + dtype, + device, + requires_grad, + FanInitKind::HeNormal, + "he_normal_like", + )?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] + fn lecun_uniform_like( + input: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let reference = PyTensor::from_python_value(input)?; + let reference_tensor = reference.tensor(); + + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => reference_tensor.dtype(), + }; + + let device = device + .map(|d| d.device()) + .unwrap_or_else(|| reference_tensor.device()); + let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); + let shape = Shape::new(reference.shape_vec()); + let tensor = create_fan_init_tensor( + shape, + dtype, + device, + requires_grad, + FanInitKind::LecunUniform, + "lecun_uniform_like", + )?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] + fn lecun_normal_like( + input: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let reference = PyTensor::from_python_value(input)?; + let reference_tensor = reference.tensor(); + + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => reference_tensor.dtype(), + }; + + let device = device + .map(|d| d.device()) + .unwrap_or_else(|| reference_tensor.device()); + let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); + let shape = Shape::new(reference.shape_vec()); + let tensor = create_fan_init_tensor( + shape, + dtype, + device, + requires_grad, + FanInitKind::LecunNormal, + "lecun_normal_like", + )?; + Ok(Self::from_tensor(tensor)) + } +} diff --git a/bindings/src/tensor/pytensor/creation/like.rs b/bindings/src/tensor/pytensor/creation/like.rs index a8472ab7..d8b13b61 100644 --- a/bindings/src/tensor/pytensor/creation/like.rs +++ b/bindings/src/tensor/pytensor/creation/like.rs @@ -1,568 +1,568 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -#[pymethods] -impl PyTensor { - #[staticmethod] - #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] - fn rand_like( - input: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let reference = PyTensor::from_python_value(input)?; - let reference_tensor = reference.tensor(); - - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => match reference_tensor.dtype() { - DataType::Float32 | DataType::Float64 => reference_tensor.dtype(), - _ => dtype::default_float_dtype(), - }, - }; - - match dtype { - DataType::Float32 | DataType::Float64 => {} - _ => { - return Err(PyValueError::new_err( - "rand_like only supports float32 or float64 dtypes", - )); - } - } - - let device = device - .map(|d| d.device()) - .unwrap_or_else(|| reference_tensor.device()); - let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); - let shape = Shape::new(reference.shape_vec()); - let tensor = create_random_tensor(shape, dtype, device, requires_grad, false)?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] - fn randn_like( - input: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let reference = PyTensor::from_python_value(input)?; - let reference_tensor = reference.tensor(); - - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => match reference_tensor.dtype() { - DataType::Float32 | DataType::Float64 => reference_tensor.dtype(), - _ => dtype::default_float_dtype(), - }, - }; - - match dtype { - DataType::Float32 | DataType::Float64 => {} - _ => { - return Err(PyValueError::new_err( - "randn_like only supports float32 or float64 dtypes", - )); - } - } - - let device = device - .map(|d| d.device()) - .unwrap_or_else(|| reference_tensor.device()); - let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); - let shape = Shape::new(reference.shape_vec()); - let tensor = create_random_tensor(shape, dtype, device, requires_grad, true)?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (input, mean=0.0, std=1.0, lower=None, upper=None, dtype=None, device=None, requires_grad=None))] - #[allow(clippy::too_many_arguments)] - fn truncated_normal_like( - input: &Bound, - mean: f64, - std: f64, - lower: Option, - upper: Option, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let reference = PyTensor::from_python_value(input)?; - let reference_tensor = reference.tensor(); - - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => match reference_tensor.dtype() { - DataType::Float32 | DataType::Float64 => reference_tensor.dtype(), - _ => dtype::default_float_dtype(), - }, - }; - - let device = device - .map(|d| d.device()) - .unwrap_or_else(|| reference_tensor.device()); - let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); - let shape = Shape::new(reference.shape_vec()); - let tensor = create_truncated_normal_tensor( - shape, - dtype, - device, - requires_grad, - mean, - std, - lower, - upper, - "truncated_normal_like", - )?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] - fn empty_like( - input: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let reference = PyTensor::from_python_value(input)?; - let reference_tensor = reference.tensor(); - - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => reference_tensor.dtype(), - }; - let device = device - .map(|d| d.device()) - .unwrap_or_else(|| reference_tensor.device()); - let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); - let shape = Shape::new(reference.shape_vec()); - let tensor = Tensor::empty(shape, dtype, device, requires_grad); - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] - fn zeros_like( - input: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let reference = PyTensor::from_python_value(input)?; - let reference_tensor = reference.tensor(); - - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => reference_tensor.dtype(), - }; - let device = device - .map(|d| d.device()) - .unwrap_or_else(|| reference_tensor.device()); - let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); - let shape = Shape::new(reference.shape_vec()); - let tensor = Tensor::zeros(shape, dtype, device, requires_grad); - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] - fn ones_like( - input: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let reference = PyTensor::from_python_value(input)?; - let reference_tensor = reference.tensor(); - - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => reference_tensor.dtype(), - }; - let device = device - .map(|d| d.device()) - .unwrap_or_else(|| reference_tensor.device()); - let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); - let shape = Shape::new(reference.shape_vec()); - let tensor = Tensor::ones(shape, dtype, device, requires_grad); - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (input, fill_value, dtype=None, device=None, requires_grad=None))] - fn full_like( - input: &Bound, - fill_value: f64, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let reference = PyTensor::from_python_value(input)?; - let reference_tensor = reference.tensor(); - - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => reference_tensor.dtype(), - }; - - let device = device - .map(|d| d.device()) - .unwrap_or_else(|| reference_tensor.device()); - let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); - let shape = reference.shape_vec(); - let tensor = create_full_tensor(shape, fill_value, dtype, device, requires_grad)?; - Ok(Self::from_tensor(tensor)) - } - - #[pyo3(signature = (shape, dtype=None, device=None, requires_grad=None))] - fn new_empty( - &self, - shape: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dims = parse_shape_like(shape, "shape")?; - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => self.inner.dtype(), - }; - let device = device - .map(|d| d.device()) - .unwrap_or_else(|| self.inner.device()); - let requires_grad = requires_grad.unwrap_or(self.inner.requires_grad()); - let tensor = Tensor::empty(Shape::new(dims), dtype, device, requires_grad); - Ok(Self::from_tensor(tensor)) - } - - #[pyo3(signature = (shape, dtype=None, device=None, requires_grad=None))] - fn new_zeros( - &self, - shape: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dims = parse_shape_like(shape, "shape")?; - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => self.inner.dtype(), - }; - let device = device - .map(|d| d.device()) - .unwrap_or_else(|| self.inner.device()); - let requires_grad = requires_grad.unwrap_or(self.inner.requires_grad()); - let tensor = Tensor::zeros(Shape::new(dims), dtype, device, requires_grad); - Ok(Self::from_tensor(tensor)) - } - - #[pyo3(signature = (shape, dtype=None, device=None, requires_grad=None))] - fn new_ones( - &self, - shape: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dims = parse_shape_like(shape, "shape")?; - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => self.inner.dtype(), - }; - let device = device - .map(|d| d.device()) - .unwrap_or_else(|| self.inner.device()); - let requires_grad = requires_grad.unwrap_or(self.inner.requires_grad()); - let tensor = Tensor::ones(Shape::new(dims), dtype, device, requires_grad); - Ok(Self::from_tensor(tensor)) - } - - #[pyo3(signature = (shape, fill_value, dtype=None, device=None, requires_grad=None))] - fn new_full( - &self, - shape: &Bound, - fill_value: f64, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dims = parse_shape_like(shape, "shape")?; - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => self.inner.dtype(), - }; - let device = device - .map(|d| d.device()) - .unwrap_or_else(|| self.inner.device()); - let requires_grad = requires_grad.unwrap_or(self.inner.requires_grad()); - let tensor = create_full_tensor(dims, fill_value, dtype, device, requires_grad)?; - Ok(Self::from_tensor(tensor)) - } - - #[pyo3(signature = (data, dtype=None, device=None, requires_grad=None))] - fn new_tensor( - &self, - data: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => self.inner.dtype(), - }; - let device = device - .map(|d| d.device()) - .unwrap_or_else(|| self.inner.device()); - let requires_grad = requires_grad.unwrap_or(self.inner.requires_grad()); - - if let Ok(py_tensor) = data.extract::>() { - let tensor = - prepare_new_tensor_from_existing(py_tensor.tensor(), dtype, device, requires_grad)?; - return Ok(Self::from_tensor(tensor)); - } - - if let Ok(inner_attr) = data.getattr(intern!(data.py(), "_tensor")) - && let Ok(py_tensor) = inner_attr.extract::>() - { - let tensor = - prepare_new_tensor_from_existing(py_tensor.tensor(), dtype, device, requires_grad)?; - return Ok(Self::from_tensor(tensor)); - } - - let tensor = convert_python_data_to_tensor(data, dtype, device, requires_grad)?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (input, low, high=None, dtype=None, device=None, requires_grad=None))] - fn randint_like( - input: &Bound, - low: i64, - high: Option, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let reference = PyTensor::from_python_value(input)?; - let reference_tensor = reference.tensor(); - - let (low, high) = match high { - Some(high) => (low, high), - None => (0, low), - }; - - if low >= high { - return Err(PyValueError::new_err( - "randint_like requires that low < high", - )); - } - - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => match reference_tensor.dtype() { - DataType::Int32 => DataType::Int32, - DataType::Int64 => DataType::Int64, - _ => DataType::Int64, - }, - }; - - match dtype { - DataType::Int32 | DataType::Int64 => {} - _ => { - return Err(PyValueError::new_err( - "randint_like only supports int32 or int64 dtypes", - )); - } - } - - let device = device - .map(|d| d.device()) - .unwrap_or_else(|| reference_tensor.device()); - let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); - let shape = Shape::new(reference.shape_vec()); - let tensor = create_randint_tensor(shape, dtype, device, requires_grad, low, high)?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (low, high=None, *shape, dtype=None, device=None, requires_grad=false))] - fn randint( - low: i64, - high: Option, - shape: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let (low, high) = match high { - Some(high) => (low, high), - None => (0, low), - }; - - if low >= high { - return Err(PyValueError::new_err("randint requires that low < high")); - } - - let dims = parse_shape_tuple(shape, "shape")?; - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => DataType::Int64, - }; - - match dtype { - DataType::Int32 | DataType::Int64 => {} - _ => { - return Err(PyValueError::new_err( - "randint only supports int32 or int64 dtypes", - )); - } - } - - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - - let shape = Shape::new(dims); - let tensor = create_randint_tensor(shape, dtype, device, requires_grad, low, high)?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (n, dtype=None, device=None, requires_grad=false))] - fn randperm( - n: usize, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => DataType::Int64, - }; - - match dtype { - DataType::Int32 | DataType::Int64 => {} - _ => { - return Err(PyValueError::new_err( - "randperm only supports int32 or int64 dtypes", - )); - } - } - - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - - let tensor = create_randperm_tensor(n, dtype, device, requires_grad)?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (n, m=None, dtype=None, device=None, requires_grad=false))] - fn eye( - n: usize, - m: Option, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let m = m.unwrap_or(n); - let dtype = dtype::resolve_dtype_arg(dtype)?; - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - - let tensor = create_eye_tensor(n, m, dtype, device, requires_grad)?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (shape, fill_value, dtype=None, device=None, requires_grad=false))] - pub fn full( - shape: &Bound, - fill_value: f64, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dtype = dtype::resolve_dtype_arg(dtype)?; - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - - let dims = parse_shape_like(shape, "shape")?; - let tensor = create_full_tensor(dims, fill_value, dtype, device, requires_grad)?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (data, dtype=None, device=None, requires_grad=None, copy=false))] - fn as_tensor( - data: &Bound, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - copy: Option, - ) -> PyResult { - let copy = copy.unwrap_or(false); - - if let Ok(py_tensor) = data.extract::>() { - let source = py_tensor.tensor(); - let target_dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => source.dtype(), - }; - let target_device = device - .map(|d| d.device()) - .unwrap_or_else(|| source.device()); - let target_requires_grad = requires_grad.unwrap_or(source.requires_grad()); - let tensor = adapt_tensor_for_as_tensor( - source, - target_dtype, - target_device, - target_requires_grad, - copy, - )?; - return Ok(Self::from_tensor(tensor)); - } - - if let Ok(inner_attr) = data.getattr(intern!(data.py(), "_tensor")) - && let Ok(py_tensor) = inner_attr.extract::>() - { - let source = py_tensor.tensor(); - let target_dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => source.dtype(), - }; - let target_device = device - .map(|d| d.device()) - .unwrap_or_else(|| source.device()); - let target_requires_grad = requires_grad.unwrap_or(source.requires_grad()); - let tensor = adapt_tensor_for_as_tensor( - source, - target_dtype, - target_device, - target_requires_grad, - copy, - )?; - return Ok(Self::from_tensor(tensor)); - } - - let target_dtype = match dtype { - Some(name) => dtype::parse_dtype(name)?, - None => infer_python_value_dtype(data).unwrap_or_else(dtype::default_dtype), - }; - - let target_device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let target_requires_grad = requires_grad.unwrap_or(false); - - let tensor = - convert_python_data_to_tensor(data, target_dtype, target_device, target_requires_grad)?; - Ok(Self::from_tensor(tensor)) - } - -} +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +#[pymethods] +impl PyTensor { + #[staticmethod] + #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] + fn rand_like( + input: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let reference = PyTensor::from_python_value(input)?; + let reference_tensor = reference.tensor(); + + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => match reference_tensor.dtype() { + DataType::Float32 | DataType::Float64 => reference_tensor.dtype(), + _ => dtype::default_float_dtype(), + }, + }; + + match dtype { + DataType::Float32 | DataType::Float64 => {} + _ => { + return Err(PyValueError::new_err( + "rand_like only supports float32 or float64 dtypes", + )); + } + } + + let device = device + .map(|d| d.device()) + .unwrap_or_else(|| reference_tensor.device()); + let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); + let shape = Shape::new(reference.shape_vec()); + let tensor = create_random_tensor(shape, dtype, device, requires_grad, false)?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] + fn randn_like( + input: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let reference = PyTensor::from_python_value(input)?; + let reference_tensor = reference.tensor(); + + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => match reference_tensor.dtype() { + DataType::Float32 | DataType::Float64 => reference_tensor.dtype(), + _ => dtype::default_float_dtype(), + }, + }; + + match dtype { + DataType::Float32 | DataType::Float64 => {} + _ => { + return Err(PyValueError::new_err( + "randn_like only supports float32 or float64 dtypes", + )); + } + } + + let device = device + .map(|d| d.device()) + .unwrap_or_else(|| reference_tensor.device()); + let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); + let shape = Shape::new(reference.shape_vec()); + let tensor = create_random_tensor(shape, dtype, device, requires_grad, true)?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (input, mean=0.0, std=1.0, lower=None, upper=None, dtype=None, device=None, requires_grad=None))] + #[allow(clippy::too_many_arguments)] + fn truncated_normal_like( + input: &Bound, + mean: f64, + std: f64, + lower: Option, + upper: Option, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let reference = PyTensor::from_python_value(input)?; + let reference_tensor = reference.tensor(); + + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => match reference_tensor.dtype() { + DataType::Float32 | DataType::Float64 => reference_tensor.dtype(), + _ => dtype::default_float_dtype(), + }, + }; + + let device = device + .map(|d| d.device()) + .unwrap_or_else(|| reference_tensor.device()); + let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); + let shape = Shape::new(reference.shape_vec()); + let tensor = create_truncated_normal_tensor( + shape, + dtype, + device, + requires_grad, + mean, + std, + lower, + upper, + "truncated_normal_like", + )?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] + fn empty_like( + input: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let reference = PyTensor::from_python_value(input)?; + let reference_tensor = reference.tensor(); + + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => reference_tensor.dtype(), + }; + let device = device + .map(|d| d.device()) + .unwrap_or_else(|| reference_tensor.device()); + let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); + let shape = Shape::new(reference.shape_vec()); + let tensor = Tensor::empty(shape, dtype, device, requires_grad); + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] + fn zeros_like( + input: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let reference = PyTensor::from_python_value(input)?; + let reference_tensor = reference.tensor(); + + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => reference_tensor.dtype(), + }; + let device = device + .map(|d| d.device()) + .unwrap_or_else(|| reference_tensor.device()); + let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); + let shape = Shape::new(reference.shape_vec()); + let tensor = Tensor::zeros(shape, dtype, device, requires_grad); + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (input, dtype=None, device=None, requires_grad=None))] + fn ones_like( + input: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let reference = PyTensor::from_python_value(input)?; + let reference_tensor = reference.tensor(); + + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => reference_tensor.dtype(), + }; + let device = device + .map(|d| d.device()) + .unwrap_or_else(|| reference_tensor.device()); + let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); + let shape = Shape::new(reference.shape_vec()); + let tensor = Tensor::ones(shape, dtype, device, requires_grad); + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (input, fill_value, dtype=None, device=None, requires_grad=None))] + fn full_like( + input: &Bound, + fill_value: f64, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let reference = PyTensor::from_python_value(input)?; + let reference_tensor = reference.tensor(); + + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => reference_tensor.dtype(), + }; + + let device = device + .map(|d| d.device()) + .unwrap_or_else(|| reference_tensor.device()); + let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); + let shape = reference.shape_vec(); + let tensor = create_full_tensor(shape, fill_value, dtype, device, requires_grad)?; + Ok(Self::from_tensor(tensor)) + } + + #[pyo3(signature = (shape, dtype=None, device=None, requires_grad=None))] + fn new_empty( + &self, + shape: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dims = parse_shape_like(shape, "shape")?; + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => self.inner.dtype(), + }; + let device = device + .map(|d| d.device()) + .unwrap_or_else(|| self.inner.device()); + let requires_grad = requires_grad.unwrap_or(self.inner.requires_grad()); + let tensor = Tensor::empty(Shape::new(dims), dtype, device, requires_grad); + Ok(Self::from_tensor(tensor)) + } + + #[pyo3(signature = (shape, dtype=None, device=None, requires_grad=None))] + fn new_zeros( + &self, + shape: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dims = parse_shape_like(shape, "shape")?; + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => self.inner.dtype(), + }; + let device = device + .map(|d| d.device()) + .unwrap_or_else(|| self.inner.device()); + let requires_grad = requires_grad.unwrap_or(self.inner.requires_grad()); + let tensor = Tensor::zeros(Shape::new(dims), dtype, device, requires_grad); + Ok(Self::from_tensor(tensor)) + } + + #[pyo3(signature = (shape, dtype=None, device=None, requires_grad=None))] + fn new_ones( + &self, + shape: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dims = parse_shape_like(shape, "shape")?; + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => self.inner.dtype(), + }; + let device = device + .map(|d| d.device()) + .unwrap_or_else(|| self.inner.device()); + let requires_grad = requires_grad.unwrap_or(self.inner.requires_grad()); + let tensor = Tensor::ones(Shape::new(dims), dtype, device, requires_grad); + Ok(Self::from_tensor(tensor)) + } + + #[pyo3(signature = (shape, fill_value, dtype=None, device=None, requires_grad=None))] + fn new_full( + &self, + shape: &Bound, + fill_value: f64, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dims = parse_shape_like(shape, "shape")?; + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => self.inner.dtype(), + }; + let device = device + .map(|d| d.device()) + .unwrap_or_else(|| self.inner.device()); + let requires_grad = requires_grad.unwrap_or(self.inner.requires_grad()); + let tensor = create_full_tensor(dims, fill_value, dtype, device, requires_grad)?; + Ok(Self::from_tensor(tensor)) + } + + #[pyo3(signature = (data, dtype=None, device=None, requires_grad=None))] + fn new_tensor( + &self, + data: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => self.inner.dtype(), + }; + let device = device + .map(|d| d.device()) + .unwrap_or_else(|| self.inner.device()); + let requires_grad = requires_grad.unwrap_or(self.inner.requires_grad()); + + if let Ok(py_tensor) = data.extract::>() { + let tensor = + prepare_new_tensor_from_existing(py_tensor.tensor(), dtype, device, requires_grad)?; + return Ok(Self::from_tensor(tensor)); + } + + if let Ok(inner_attr) = data.getattr(intern!(data.py(), "_tensor")) + && let Ok(py_tensor) = inner_attr.extract::>() + { + let tensor = + prepare_new_tensor_from_existing(py_tensor.tensor(), dtype, device, requires_grad)?; + return Ok(Self::from_tensor(tensor)); + } + + let tensor = convert_python_data_to_tensor(data, dtype, device, requires_grad)?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (input, low, high=None, dtype=None, device=None, requires_grad=None))] + fn randint_like( + input: &Bound, + low: i64, + high: Option, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let reference = PyTensor::from_python_value(input)?; + let reference_tensor = reference.tensor(); + + let (low, high) = match high { + Some(high) => (low, high), + None => (0, low), + }; + + if low >= high { + return Err(PyValueError::new_err( + "randint_like requires that low < high", + )); + } + + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => match reference_tensor.dtype() { + DataType::Int32 => DataType::Int32, + DataType::Int64 => DataType::Int64, + _ => DataType::Int64, + }, + }; + + match dtype { + DataType::Int32 | DataType::Int64 => {} + _ => { + return Err(PyValueError::new_err( + "randint_like only supports int32 or int64 dtypes", + )); + } + } + + let device = device + .map(|d| d.device()) + .unwrap_or_else(|| reference_tensor.device()); + let requires_grad = requires_grad.unwrap_or(reference_tensor.requires_grad()); + let shape = Shape::new(reference.shape_vec()); + let tensor = create_randint_tensor(shape, dtype, device, requires_grad, low, high)?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (low, high=None, *shape, dtype=None, device=None, requires_grad=false))] + fn randint( + low: i64, + high: Option, + shape: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let (low, high) = match high { + Some(high) => (low, high), + None => (0, low), + }; + + if low >= high { + return Err(PyValueError::new_err("randint requires that low < high")); + } + + let dims = parse_shape_tuple(shape, "shape")?; + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => DataType::Int64, + }; + + match dtype { + DataType::Int32 | DataType::Int64 => {} + _ => { + return Err(PyValueError::new_err( + "randint only supports int32 or int64 dtypes", + )); + } + } + + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + + let shape = Shape::new(dims); + let tensor = create_randint_tensor(shape, dtype, device, requires_grad, low, high)?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (n, dtype=None, device=None, requires_grad=false))] + fn randperm( + n: usize, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => DataType::Int64, + }; + + match dtype { + DataType::Int32 | DataType::Int64 => {} + _ => { + return Err(PyValueError::new_err( + "randperm only supports int32 or int64 dtypes", + )); + } + } + + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + + let tensor = create_randperm_tensor(n, dtype, device, requires_grad)?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (n, m=None, dtype=None, device=None, requires_grad=false))] + fn eye( + n: usize, + m: Option, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let m = m.unwrap_or(n); + let dtype = dtype::resolve_dtype_arg(dtype)?; + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + + let tensor = create_eye_tensor(n, m, dtype, device, requires_grad)?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (shape, fill_value, dtype=None, device=None, requires_grad=false))] + pub fn full( + shape: &Bound, + fill_value: f64, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dtype = dtype::resolve_dtype_arg(dtype)?; + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + + let dims = parse_shape_like(shape, "shape")?; + let tensor = create_full_tensor(dims, fill_value, dtype, device, requires_grad)?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (data, dtype=None, device=None, requires_grad=None, copy=false))] + fn as_tensor( + data: &Bound, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + copy: Option, + ) -> PyResult { + let copy = copy.unwrap_or(false); + + if let Ok(py_tensor) = data.extract::>() { + let source = py_tensor.tensor(); + let target_dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => source.dtype(), + }; + let target_device = device + .map(|d| d.device()) + .unwrap_or_else(|| source.device()); + let target_requires_grad = requires_grad.unwrap_or(source.requires_grad()); + let tensor = adapt_tensor_for_as_tensor( + source, + target_dtype, + target_device, + target_requires_grad, + copy, + )?; + return Ok(Self::from_tensor(tensor)); + } + + if let Ok(inner_attr) = data.getattr(intern!(data.py(), "_tensor")) + && let Ok(py_tensor) = inner_attr.extract::>() + { + let source = py_tensor.tensor(); + let target_dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => source.dtype(), + }; + let target_device = device + .map(|d| d.device()) + .unwrap_or_else(|| source.device()); + let target_requires_grad = requires_grad.unwrap_or(source.requires_grad()); + let tensor = adapt_tensor_for_as_tensor( + source, + target_dtype, + target_device, + target_requires_grad, + copy, + )?; + return Ok(Self::from_tensor(tensor)); + } + + let target_dtype = match dtype { + Some(name) => dtype::parse_dtype(name)?, + None => infer_python_value_dtype(data).unwrap_or_else(dtype::default_dtype), + }; + + let target_device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let target_requires_grad = requires_grad.unwrap_or(false); + + let tensor = + convert_python_data_to_tensor(data, target_dtype, target_device, target_requires_grad)?; + Ok(Self::from_tensor(tensor)) + } +} diff --git a/bindings/src/tensor/pytensor/creation/range.rs b/bindings/src/tensor/pytensor/creation/range.rs index 25d9032b..9b79f59c 100644 --- a/bindings/src/tensor/pytensor/creation/range.rs +++ b/bindings/src/tensor/pytensor/creation/range.rs @@ -1,92 +1,92 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -#[pymethods] -impl PyTensor { - #[staticmethod] - #[pyo3(signature = (start, end=None, step=1.0, dtype=None, device=None, requires_grad=false))] - fn arange( - start: f64, - end: Option, - step: f64, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dtype = dtype::resolve_dtype_arg(dtype)?; - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - - let (start, end) = match end { - Some(value) => (start, value), - None => (0.0, start), - }; - - let tensor = create_arange_tensor(start, end, step, dtype, device, requires_grad)?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (start, end, steps, dtype=None, device=None, requires_grad=false))] - fn linspace( - start: f64, - end: f64, - steps: usize, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - if steps == 0 { - return Err(PyValueError::new_err("steps must be greater than zero")); - } - - let dtype = dtype::resolve_dtype_arg(dtype)?; - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - - let tensor = create_linspace_tensor(start, end, steps, dtype, device, requires_grad)?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (start, end, steps, base=None, dtype=None, device=None, requires_grad=false))] - fn logspace( - start: f64, - end: f64, - steps: usize, - base: Option, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - if steps == 0 { - return Err(PyValueError::new_err("steps must be greater than zero")); - } - - let dtype = dtype::resolve_dtype_arg(dtype)?; - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - let base = base.unwrap_or(10.0); - - let tensor = create_logspace_tensor(start, end, steps, base, dtype, device, requires_grad)?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (array, requires_grad=false))] - fn from_numpy(array: &Bound, requires_grad: bool) -> PyResult { - let tensor = convert_numpy_to_tensor(array, requires_grad)?; - Ok(Self::from_tensor(tensor)) - } - - #[staticmethod] - #[pyo3(signature = (array, requires_grad=false))] - fn from_numpy_shared(array: &Bound, requires_grad: bool) -> PyResult { +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +#[pymethods] +impl PyTensor { + #[staticmethod] + #[pyo3(signature = (start, end=None, step=1.0, dtype=None, device=None, requires_grad=false))] + fn arange( + start: f64, + end: Option, + step: f64, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dtype = dtype::resolve_dtype_arg(dtype)?; + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + + let (start, end) = match end { + Some(value) => (start, value), + None => (0.0, start), + }; + + let tensor = create_arange_tensor(start, end, step, dtype, device, requires_grad)?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (start, end, steps, dtype=None, device=None, requires_grad=false))] + fn linspace( + start: f64, + end: f64, + steps: usize, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + if steps == 0 { + return Err(PyValueError::new_err("steps must be greater than zero")); + } + + let dtype = dtype::resolve_dtype_arg(dtype)?; + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + + let tensor = create_linspace_tensor(start, end, steps, dtype, device, requires_grad)?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (start, end, steps, base=None, dtype=None, device=None, requires_grad=false))] + fn logspace( + start: f64, + end: f64, + steps: usize, + base: Option, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + if steps == 0 { + return Err(PyValueError::new_err("steps must be greater than zero")); + } + + let dtype = dtype::resolve_dtype_arg(dtype)?; + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + let base = base.unwrap_or(10.0); + + let tensor = create_logspace_tensor(start, end, steps, base, dtype, device, requires_grad)?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (array, requires_grad=false))] + fn from_numpy(array: &Bound, requires_grad: bool) -> PyResult { + let tensor = convert_numpy_to_tensor(array, requires_grad)?; + Ok(Self::from_tensor(tensor)) + } + + #[staticmethod] + #[pyo3(signature = (array, requires_grad=false))] + fn from_numpy_shared(array: &Bound, requires_grad: bool) -> PyResult { // Currently delegates to from_numpy; true zero-copy requires more complex memory management. - Self::from_numpy(array, requires_grad) - } - -} + Self::from_numpy(array, requires_grad) + } +} diff --git a/bindings/src/tensor/pytensor/grad.rs b/bindings/src/tensor/pytensor/grad.rs index 4735e536..28378ed9 100644 --- a/bindings/src/tensor/pytensor/grad.rs +++ b/bindings/src/tensor/pytensor/grad.rs @@ -1,100 +1,109 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -#[pymethods] -impl PyTensor { - // Gradient operations - #[pyo3(signature = (gradient=None, retain_graph=false, create_graph=false))] - fn backward( - &self, - gradient: Option<&Bound>, - retain_graph: bool, - create_graph: bool, - ) -> PyResult<()> { - if create_graph { - return Err(PyNotImplementedError::new_err( - "create_graph=True is not supported; all computations execute in the Rust backend", - )); - } - - if !self.requires_grad() && self.is_leaf() { - return Err(PyRuntimeError::new_err( - "element 0 of tensors does not require grad and does not have a grad_fn", - )); - } - - if !retain_graph && engine::autograd::is_graph_consumed() { - return Err(PyRuntimeError::new_err( - "Computation graph has been freed. Re-run the forward pass or call backward(retain_graph=True).", - )); - } - - let grad_tensor = if let Some(value) = gradient { - if value.is_none() { - None - } else if let Ok(py_tensor) = value.extract::() { - let mut tensor = py_tensor.inner.clone(); - ensure_backward_gradient_compatible(&self.inner, &mut tensor)?; - Some(tensor) - } else { - let mut tensor = tensor_from_py_value(&self.inner, value)?; - ensure_backward_gradient_compatible(&self.inner, &mut tensor)?; - Some(tensor) - } - } else { - None - }; - - self.inner.backward(grad_tensor).map_err(_convert_error)?; - - if !retain_graph { - engine::autograd::mark_graph_consumed(); - } - - Ok(()) - } - - pub fn requires_grad_(&mut self, requires_grad: bool) -> PyResult<()> { - self.inner = self.inner.clone().requires_grad_(requires_grad); - Ok(()) - } - - #[pyo3(signature = (source, *, non_blocking=false))] - fn copy_<'py>( - mut slf: PyRefMut<'py, Self>, - source: &Bound, - non_blocking: Option, - ) -> PyResult> { - if non_blocking.unwrap_or(false) { - return Err(PyNotImplementedError::new_err( - "non_blocking copy_ is not implemented", - )); - } - - let reference = PyTensor::from_python_value(source)?; - slf.inner - .copy_(reference.tensor()) - .map_err(_convert_error)?; - register_leaf_tensor(&slf.inner); - Ok(slf) - } - - fn fill_<'py>( - mut slf: PyRefMut<'py, Self>, - value: &Bound, - ) -> PyResult> { - let fill_value = extract_real_scalar(value, "value")?; - slf.inner.fill_(fill_value).map_err(_convert_error)?; - register_leaf_tensor(&slf.inner); - Ok(slf) - } - - #[pyo3(signature = (set_to_none=false))] - fn zero_grad(&mut self, set_to_none: bool) { - self.inner.zero_grad(set_to_none); - } - -} +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +#[pymethods] +impl PyTensor { + // Gradient operations + #[pyo3(signature = (gradient=None, retain_graph=false, create_graph=false))] + fn backward( + &self, + gradient: Option<&Bound>, + retain_graph: bool, + create_graph: bool, + ) -> PyResult<()> { + if create_graph { + return Err(PyNotImplementedError::new_err( + "create_graph=True is not supported; all computations execute in the Rust backend", + )); + } + + if !self.requires_grad() && self.is_leaf() { + return Err(PyRuntimeError::new_err( + "element 0 of tensors does not require grad and does not have a grad_fn", + )); + } + + if !retain_graph && engine::autograd::is_graph_consumed() { + return Err(PyRuntimeError::new_err( + "Computation graph has been freed. Re-run the forward pass or call backward(retain_graph=True).", + )); + } + + let grad_tensor = if let Some(value) = gradient { + if value.is_none() { + None + } else if let Ok(py_tensor) = value.extract::() { + let mut tensor = py_tensor.inner.clone(); + ensure_backward_gradient_compatible(&self.inner, &mut tensor)?; + Some(tensor) + } else { + let mut tensor = tensor_from_py_value(&self.inner, value)?; + ensure_backward_gradient_compatible(&self.inner, &mut tensor)?; + Some(tensor) + } + } else { + None + }; + + self.inner.backward(grad_tensor).map_err(_convert_error)?; + + if !retain_graph { + engine::autograd::mark_graph_consumed(); + // Free the tensors saved for backward (activations, operands) + // immediately instead of holding them until the next optimizer + // step. Gradients remain available via `.grad`/`get_gradient`. + engine::autograd::release_saved_subgraph(&self.inner); + } + + Ok(()) + } + + /// Set `requires_grad` in place and return `self`, so calls chain the + /// same way as in PyTorch: `x = mt.randn(2, 2).requires_grad_(True)`. + pub fn requires_grad_<'py>( + mut slf: PyRefMut<'py, Self>, + requires_grad: bool, + ) -> PyRefMut<'py, Self> { + slf.inner = slf.inner.clone().requires_grad_(requires_grad); + slf + } + + #[pyo3(signature = (source, *, non_blocking=false))] + fn copy_<'py>( + mut slf: PyRefMut<'py, Self>, + source: &Bound, + non_blocking: Option, + ) -> PyResult> { + if non_blocking.unwrap_or(false) { + return Err(PyNotImplementedError::new_err( + "non_blocking copy_ is not implemented", + )); + } + + let reference = PyTensor::from_python_value(source)?; + slf.inner + .copy_(reference.tensor()) + .map_err(_convert_error)?; + register_leaf_tensor(&slf.inner); + Ok(slf) + } + + fn fill_<'py>( + mut slf: PyRefMut<'py, Self>, + value: &Bound, + ) -> PyResult> { + let fill_value = extract_real_scalar(value, "value")?; + slf.inner.fill_(fill_value).map_err(_convert_error)?; + register_leaf_tensor(&slf.inner); + Ok(slf) + } + + #[pyo3(signature = (set_to_none=false))] + fn zero_grad(&mut self, set_to_none: bool) { + self.inner.zero_grad(set_to_none); + } +} diff --git a/bindings/src/tensor/pytensor/isclose.rs b/bindings/src/tensor/pytensor/isclose.rs index d32cc550..d719ba4e 100644 --- a/bindings/src/tensor/pytensor/isclose.rs +++ b/bindings/src/tensor/pytensor/isclose.rs @@ -4,6 +4,7 @@ // This source code is licensed under the Apache-style license found in the // LICENSE file in the root directory of this source tree. +use super::*; #[pymethods] impl PyTensor { #[pyo3(signature = (other, rtol=None, atol=None, equal_nan=false))] diff --git a/bindings/src/tensor/pytensor/math.rs b/bindings/src/tensor/pytensor/math.rs index fbd927c4..44c7c966 100644 --- a/bindings/src/tensor/pytensor/math.rs +++ b/bindings/src/tensor/pytensor/math.rs @@ -1,344 +1,339 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -#[pymethods] -impl PyTensor { - // Mathematical functions - fn abs(&self) -> PyResult { - let result = self.inner.abs().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn sqrt(&self) -> PyResult { - let result = self.inner.sqrt().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn rsqrt(&self) -> PyResult { - let result = self.inner.rsqrt().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn pow(&self, exponent: &Bound) -> PyResult { - if let Ok(exp_tensor) = exponent.extract::() { - let result = self.inner.pow(&exp_tensor.inner).map_err(_convert_error)?; - return Ok(Self::from_tensor(result)); - } - - if let Ok(exp) = exponent.extract::() { - let result = self.inner.powf(exp).map_err(_convert_error)?; - return Ok(Self::from_tensor(result)); - } - - let exp_tensor = tensor_from_py_value(&self.inner, exponent)?; - let result = self.inner.pow(&exp_tensor).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn exp(&self) -> PyResult { - let result = self.inner.exp().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn log(&self) -> PyResult { - let result = self.inner.log().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn log1p(&self) -> PyResult { - let result = self.inner.log1p().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn expm1(&self) -> PyResult { - let result = self.inner.expm1().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn sin(&self) -> PyResult { - let result = self.inner.sin().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn cos(&self) -> PyResult { - let result = self.inner.cos().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn tan(&self) -> PyResult { - let result = self.inner.tan().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn asin(&self) -> PyResult { - let result = self.inner.asin().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn acos(&self) -> PyResult { - let result = self.inner.acos().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn atan(&self) -> PyResult { - let result = self.inner.atan().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn sinh(&self) -> PyResult { - let result = self.inner.sinh().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn cosh(&self) -> PyResult { - let result = self.inner.cosh().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn asinh(&self) -> PyResult { - let result = self.inner.asinh().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn acosh(&self) -> PyResult { - let result = self.inner.acosh().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn atanh(&self) -> PyResult { - let result = self.inner.atanh().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (nan=0.0, posinf=None, neginf=None))] - pub fn nan_to_num( - &self, - nan: f64, - posinf: Option, - neginf: Option, - ) -> PyResult { - let result = self - .inner - .nan_to_num(nan, posinf, neginf) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn isnan(&self) -> PyResult { - let result = self.inner.isnan().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn isinf(&self) -> PyResult { - let result = self.inner.isinf().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn isfinite(&self) -> PyResult { - let result = self.inner.isfinite().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn __pow__(&self, exponent: &Bound, _mod: Option<&Bound>) -> PyResult { - self.pow(exponent) - } - - fn __rpow__(&self, base: &Bound, _mod: Option<&Bound>) -> PyResult { - let base_tensor = tensor_from_py_value(&self.inner, base)?; - let result = base_tensor.pow(&self.inner).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn relu(&self) -> PyResult { - let result = self.inner.relu().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn hardshrink(&self, lambd: Option) -> PyResult { - let result = self - .inner - .hardshrink(lambd.unwrap_or(0.5)) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (dim=None))] - pub fn softmax(&self, dim: Option) -> PyResult { - let resolved_dim = match dim { - Some(dim) => { - let ndim = self.inner.ndim() as isize; - let dim = if dim < 0 { dim + ndim } else { dim }; - if dim < 0 || dim >= ndim { - return Err(PyIndexError::new_err(format!( - "Dimension out of range (expected to be in range of [-{ndim}, {ndim}), but got {dim})" - ))); - } - Some(dim as usize) - } - None => None, - }; - - let result = self.inner.softmax(resolved_dim).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (dim=None))] - pub fn log_softmax(&self, dim: Option) -> PyResult { - let resolved_dim = match dim { - Some(dim) => { - let ndim = self.inner.ndim() as isize; - let dim = if dim < 0 { dim + ndim } else { dim }; - if dim < 0 || dim >= ndim { - return Err(PyIndexError::new_err(format!( - "Dimension out of range (expected to be in range of [-{ndim}, {ndim}), but got {dim})" - ))); - } - Some(dim as usize) - } - None => None, - }; - - let result = self - .inner - .log_softmax(resolved_dim) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (mask, dim=None))] - pub fn masked_softmax(&self, mask: &Bound, dim: Option) -> PyResult { - let mask_tensor = tensor_from_py_value(&self.inner, mask)?; - let resolved_dim = match dim { - Some(dim) => { - let ndim = self.inner.ndim() as isize; - let dim = if dim < 0 { dim + ndim } else { dim }; - if dim < 0 || dim >= ndim { - return Err(PyIndexError::new_err(format!( - "Dimension out of range (expected to be in range of [-{ndim}, {ndim}), but got {dim})" - ))); - } - Some(dim as usize) - } - None => None, - }; - - let result = self - .inner - .masked_softmax(&mask_tensor, resolved_dim) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (mask, dim=None))] - pub fn masked_log_softmax(&self, mask: &Bound, dim: Option) -> PyResult { - let mask_tensor = tensor_from_py_value(&self.inner, mask)?; - let resolved_dim = match dim { - Some(dim) => { - let ndim = self.inner.ndim() as isize; - let dim = if dim < 0 { dim + ndim } else { dim }; - if dim < 0 || dim >= ndim { - return Err(PyIndexError::new_err(format!( - "Dimension out of range (expected to be in range of [-{ndim}, {ndim}), but got {dim})" - ))); - } - Some(dim as usize) - } - None => None, - }; - - let result = self - .inner - .masked_log_softmax(&mask_tensor, resolved_dim) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (normalized_shape, weight=None, bias=None, eps=1e-5))] - pub fn layer_norm( - &self, - normalized_shape: Vec, - weight: Option<&PyTensor>, - bias: Option<&PyTensor>, - eps: Option, - ) -> PyResult { - if normalized_shape.is_empty() { - return Err(PyValueError::new_err( - "layer_norm requires normalized_shape to contain at least one dimension", - )); - } - - let weight_inner = weight.map(|w| &w.inner); - let bias_inner = bias.map(|b| &b.inner); - let result = self - .inner - .layer_norm( - &normalized_shape, - weight_inner, - bias_inner, - eps.unwrap_or(1e-5), - ) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn gelu(&self, approximate: Option<&str>) -> PyResult { - let approx_mode = approximate.unwrap_or("none"); - let approximate = if approx_mode.eq_ignore_ascii_case("none") { - false - } else if approx_mode.eq_ignore_ascii_case("tanh") { - true - } else { - return Err(PyValueError::new_err( - "approximate must be 'none' or 'tanh' for gelu", - )); - }; - - let result = self.inner.gelu(approximate).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn sigmoid(&self) -> PyResult { - let result = self.inner.sigmoid().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn softplus(&self, beta: Option, threshold: Option) -> PyResult { - let result = self - .inner - .softplus(beta.unwrap_or(1.0), threshold.unwrap_or(20.0)) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn elu(&self, alpha: Option) -> PyResult { - let result = self - .inner - .elu(alpha.unwrap_or(1.0)) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn selu(&self) -> PyResult { - let result = self.inner.selu().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn silu(&self) -> PyResult { - let result = self.inner.silu().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn softsign(&self) -> PyResult { - let result = self.inner.softsign().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn tanh(&self) -> PyResult { - let result = self.inner.tanh().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - -} +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +#[pymethods] +impl PyTensor { + // Mathematical functions + fn abs(&self) -> PyResult { + let result = self.inner.abs().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + fn sqrt(&self) -> PyResult { + let result = self.inner.sqrt().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn rsqrt(&self) -> PyResult { + let result = self.inner.rsqrt().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + fn pow(&self, exponent: &Bound) -> PyResult { + if let Ok(exp_tensor) = exponent.extract::() { + let result = self.inner.pow(&exp_tensor.inner).map_err(_convert_error)?; + return Ok(Self::from_tensor(result)); + } + + if let Ok(exp) = exponent.extract::() { + let result = self.inner.powf(exp).map_err(_convert_error)?; + return Ok(Self::from_tensor(result)); + } + + let exp_tensor = tensor_from_py_value(&self.inner, exponent)?; + let result = self.inner.pow(&exp_tensor).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + fn exp(&self) -> PyResult { + let result = self.inner.exp().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + fn log(&self) -> PyResult { + let result = self.inner.log().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn log1p(&self) -> PyResult { + let result = self.inner.log1p().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn expm1(&self) -> PyResult { + let result = self.inner.expm1().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn sin(&self) -> PyResult { + let result = self.inner.sin().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn cos(&self) -> PyResult { + let result = self.inner.cos().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn tan(&self) -> PyResult { + let result = self.inner.tan().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn asin(&self) -> PyResult { + let result = self.inner.asin().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn acos(&self) -> PyResult { + let result = self.inner.acos().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn atan(&self) -> PyResult { + let result = self.inner.atan().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn sinh(&self) -> PyResult { + let result = self.inner.sinh().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn cosh(&self) -> PyResult { + let result = self.inner.cosh().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn asinh(&self) -> PyResult { + let result = self.inner.asinh().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn acosh(&self) -> PyResult { + let result = self.inner.acosh().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn atanh(&self) -> PyResult { + let result = self.inner.atanh().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (nan=0.0, posinf=None, neginf=None))] + pub fn nan_to_num(&self, nan: f64, posinf: Option, neginf: Option) -> PyResult { + let result = self + .inner + .nan_to_num(nan, posinf, neginf) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + fn isnan(&self) -> PyResult { + let result = self.inner.isnan().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + fn isinf(&self) -> PyResult { + let result = self.inner.isinf().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + fn isfinite(&self) -> PyResult { + let result = self.inner.isfinite().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + fn __pow__(&self, exponent: &Bound, _mod: Option<&Bound>) -> PyResult { + self.pow(exponent) + } + + fn __rpow__(&self, base: &Bound, _mod: Option<&Bound>) -> PyResult { + let base_tensor = tensor_from_py_value(&self.inner, base)?; + let result = base_tensor.pow(&self.inner).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn relu(&self) -> PyResult { + let result = self.inner.relu().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn hardshrink(&self, lambd: Option) -> PyResult { + let result = self + .inner + .hardshrink(lambd.unwrap_or(0.5)) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (dim=None))] + pub fn softmax(&self, dim: Option) -> PyResult { + let resolved_dim = match dim { + Some(dim) => { + let ndim = self.inner.ndim() as isize; + let dim = if dim < 0 { dim + ndim } else { dim }; + if dim < 0 || dim >= ndim { + return Err(PyIndexError::new_err(format!( + "Dimension out of range (expected to be in range of [-{ndim}, {ndim}), but got {dim})" + ))); + } + Some(dim as usize) + } + None => None, + }; + + let result = self.inner.softmax(resolved_dim).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (dim=None))] + pub fn log_softmax(&self, dim: Option) -> PyResult { + let resolved_dim = match dim { + Some(dim) => { + let ndim = self.inner.ndim() as isize; + let dim = if dim < 0 { dim + ndim } else { dim }; + if dim < 0 || dim >= ndim { + return Err(PyIndexError::new_err(format!( + "Dimension out of range (expected to be in range of [-{ndim}, {ndim}), but got {dim})" + ))); + } + Some(dim as usize) + } + None => None, + }; + + let result = self + .inner + .log_softmax(resolved_dim) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (mask, dim=None))] + pub fn masked_softmax(&self, mask: &Bound, dim: Option) -> PyResult { + let mask_tensor = tensor_from_py_value(&self.inner, mask)?; + let resolved_dim = match dim { + Some(dim) => { + let ndim = self.inner.ndim() as isize; + let dim = if dim < 0 { dim + ndim } else { dim }; + if dim < 0 || dim >= ndim { + return Err(PyIndexError::new_err(format!( + "Dimension out of range (expected to be in range of [-{ndim}, {ndim}), but got {dim})" + ))); + } + Some(dim as usize) + } + None => None, + }; + + let result = self + .inner + .masked_softmax(&mask_tensor, resolved_dim) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (mask, dim=None))] + pub fn masked_log_softmax(&self, mask: &Bound, dim: Option) -> PyResult { + let mask_tensor = tensor_from_py_value(&self.inner, mask)?; + let resolved_dim = match dim { + Some(dim) => { + let ndim = self.inner.ndim() as isize; + let dim = if dim < 0 { dim + ndim } else { dim }; + if dim < 0 || dim >= ndim { + return Err(PyIndexError::new_err(format!( + "Dimension out of range (expected to be in range of [-{ndim}, {ndim}), but got {dim})" + ))); + } + Some(dim as usize) + } + None => None, + }; + + let result = self + .inner + .masked_log_softmax(&mask_tensor, resolved_dim) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (normalized_shape, weight=None, bias=None, eps=1e-5))] + pub fn layer_norm( + &self, + normalized_shape: Vec, + weight: Option<&PyTensor>, + bias: Option<&PyTensor>, + eps: Option, + ) -> PyResult { + if normalized_shape.is_empty() { + return Err(PyValueError::new_err( + "layer_norm requires normalized_shape to contain at least one dimension", + )); + } + + let weight_inner = weight.map(|w| &w.inner); + let bias_inner = bias.map(|b| &b.inner); + let result = self + .inner + .layer_norm( + &normalized_shape, + weight_inner, + bias_inner, + eps.unwrap_or(1e-5), + ) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn gelu(&self, approximate: Option<&str>) -> PyResult { + let approx_mode = approximate.unwrap_or("none"); + let approximate = if approx_mode.eq_ignore_ascii_case("none") { + false + } else if approx_mode.eq_ignore_ascii_case("tanh") { + true + } else { + return Err(PyValueError::new_err( + "approximate must be 'none' or 'tanh' for gelu", + )); + }; + + let result = self.inner.gelu(approximate).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn sigmoid(&self) -> PyResult { + let result = self.inner.sigmoid().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn softplus(&self, beta: Option, threshold: Option) -> PyResult { + let result = self + .inner + .softplus(beta.unwrap_or(1.0), threshold.unwrap_or(20.0)) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn elu(&self, alpha: Option) -> PyResult { + let result = self + .inner + .elu(alpha.unwrap_or(1.0)) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn selu(&self) -> PyResult { + let result = self.inner.selu().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn silu(&self) -> PyResult { + let result = self.inner.silu().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn softsign(&self) -> PyResult { + let result = self.inner.softsign().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn tanh(&self) -> PyResult { + let result = self.inner.tanh().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } +} diff --git a/bindings/src/tensor/pytensor/numpy.rs b/bindings/src/tensor/pytensor/numpy.rs index 61f9b7ae..826dc43e 100644 --- a/bindings/src/tensor/pytensor/numpy.rs +++ b/bindings/src/tensor/pytensor/numpy.rs @@ -1,119 +1,119 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -#[pymethods] -impl PyTensor { - // NumPy conversion methods - fn numpy(&self, py: Python) -> PyResult> { - convert_tensor_to_numpy(&self.inner, py, false) - } - - fn numpy_copy(&self, py: Python) -> PyResult> { - convert_tensor_to_numpy(&self.inner, py, true) - } - - #[pyo3(signature = (dtype=None))] - fn __array__(&self, py: Python, dtype: Option<&Bound>) -> PyResult> { - let array = self.numpy(py)?; - if let Some(dtype_obj) = dtype { - let array_bound = array.bind(py); - let kwargs = PyDict::new(py); - kwargs.set_item(intern!(py, "copy"), false)?; - let casted = - array_bound.call_method(intern!(py, "astype"), (dtype_obj,), Some(&kwargs))?; - Ok(casted.into()) - } else { - Ok(array) - } - } - - #[pyo3(signature = (ufunc, method, *inputs, **kwargs))] - fn __array_ufunc__( - &self, - py: Python, - ufunc: &Bound, - method: &str, - inputs: &Bound, - kwargs: Option<&Bound>, - ) -> PyResult> { - if method != "__call__" { - return py_not_implemented(py); - } - - if let Some(mapping) = kwargs - && let Some(out) = mapping.get_item("out")? - && !out.is_none() - { - return py_not_implemented(py); - } - - let mut operands: Vec = Vec::with_capacity(inputs.len()); - for value in inputs.iter() { - match tensor_from_py_value(&self.inner, &value) { - Ok(tensor) => operands.push(tensor), - Err(_) => return py_not_implemented(py), - } - } - - let Some(name_obj) = ufunc.getattr(intern!(py, "__name__")).ok() else { - return py_not_implemented(py); - }; - let name = name_obj.str()?.to_str()?.to_ascii_lowercase(); - - let result = match (name.as_str(), operands.len()) { - ("add", 2) => { - apply_binary_ufunc(&operands, BinaryOpKind::Add, |lhs, rhs| lhs.add(rhs))? - } - ("subtract", 2) => apply_binary_ufunc(&operands, BinaryOpKind::Sub, |lhs, rhs| { - engine::operations::arithmetic::sub(lhs, rhs) - })?, - ("multiply", 2) => apply_binary_ufunc(&operands, BinaryOpKind::Mul, |lhs, rhs| { - engine::operations::arithmetic::mul(lhs, rhs) - })?, - ("true_divide", 2) | ("divide", 2) => { - apply_binary_ufunc(&operands, BinaryOpKind::Div, |lhs, rhs| { - engine::operations::arithmetic::div(lhs, rhs) - })? - } - ("power", 2) => { - apply_binary_ufunc(&operands, BinaryOpKind::Mul, |lhs, rhs| lhs.pow(rhs))? - } - ("maximum", 2) => apply_binary_ufunc(&operands, BinaryOpKind::Maximum, |lhs, rhs| { - lhs.maximum(rhs) - })?, - ("minimum", 2) => apply_binary_ufunc(&operands, BinaryOpKind::Minimum, |lhs, rhs| { - lhs.minimum(rhs) - })?, - ("negative", 1) => apply_unary_ufunc(&operands, |tensor| { - engine::operations::arithmetic::neg(tensor) - })?, - ("absolute", 1) | ("abs", 1) => apply_unary_ufunc(&operands, |tensor| tensor.abs())?, - ("exp", 1) => apply_unary_ufunc(&operands, |tensor| tensor.exp())?, - ("log", 1) => apply_unary_ufunc(&operands, |tensor| tensor.log())?, - ("sin", 1) => apply_unary_ufunc(&operands, |tensor| tensor.sin())?, - ("cos", 1) => apply_unary_ufunc(&operands, |tensor| tensor.cos())?, - ("tan", 1) => apply_unary_ufunc(&operands, |tensor| tensor.tan())?, - ("sqrt", 1) => apply_unary_ufunc(&operands, |tensor| tensor.sqrt())?, - _ => return py_not_implemented(py), - }; - - let py_tensor = Py::new(py, PyTensor::from_tensor(result))?; - Ok(py_tensor.into_any()) - } - - fn tolist(&self) -> PyResult> { - if self.inner.ndim() == 0 { - Python::attach(|py| convert_tensor_to_python_scalar(&self.inner, py)) - } else { - Python::attach(|py| convert_tensor_to_python_list(&self.inner, py)) - } - } - - fn item(&self) -> PyResult> { - Python::attach(|py| convert_tensor_to_python_scalar(&self.inner, py)) - } - -} +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +#[pymethods] +impl PyTensor { + // NumPy conversion methods + fn numpy(&self, py: Python) -> PyResult> { + convert_tensor_to_numpy(&self.inner, py, false) + } + + fn numpy_copy(&self, py: Python) -> PyResult> { + convert_tensor_to_numpy(&self.inner, py, true) + } + + #[pyo3(signature = (dtype=None))] + fn __array__(&self, py: Python, dtype: Option<&Bound>) -> PyResult> { + let array = self.numpy(py)?; + if let Some(dtype_obj) = dtype { + let array_bound = array.bind(py); + let kwargs = PyDict::new(py); + kwargs.set_item(intern!(py, "copy"), false)?; + let casted = + array_bound.call_method(intern!(py, "astype"), (dtype_obj,), Some(&kwargs))?; + Ok(casted.into()) + } else { + Ok(array) + } + } + + #[pyo3(signature = (ufunc, method, *inputs, **kwargs))] + fn __array_ufunc__( + &self, + py: Python, + ufunc: &Bound, + method: &str, + inputs: &Bound, + kwargs: Option<&Bound>, + ) -> PyResult> { + if method != "__call__" { + return py_not_implemented(py); + } + + if let Some(mapping) = kwargs + && let Some(out) = mapping.get_item("out")? + && !out.is_none() + { + return py_not_implemented(py); + } + + let mut operands: Vec = Vec::with_capacity(inputs.len()); + for value in inputs.iter() { + match tensor_from_py_value(&self.inner, &value) { + Ok(tensor) => operands.push(tensor), + Err(_) => return py_not_implemented(py), + } + } + + let Some(name_obj) = ufunc.getattr(intern!(py, "__name__")).ok() else { + return py_not_implemented(py); + }; + let name = name_obj.str()?.to_str()?.to_ascii_lowercase(); + + let result = match (name.as_str(), operands.len()) { + ("add", 2) => { + apply_binary_ufunc(&operands, BinaryOpKind::Add, |lhs, rhs| lhs.add(rhs))? + } + ("subtract", 2) => apply_binary_ufunc(&operands, BinaryOpKind::Sub, |lhs, rhs| { + engine::operations::arithmetic::sub(lhs, rhs) + })?, + ("multiply", 2) => apply_binary_ufunc(&operands, BinaryOpKind::Mul, |lhs, rhs| { + engine::operations::arithmetic::mul(lhs, rhs) + })?, + ("true_divide", 2) | ("divide", 2) => { + apply_binary_ufunc(&operands, BinaryOpKind::Div, |lhs, rhs| { + engine::operations::arithmetic::div(lhs, rhs) + })? + } + ("power", 2) => { + apply_binary_ufunc(&operands, BinaryOpKind::Mul, |lhs, rhs| lhs.pow(rhs))? + } + ("maximum", 2) => apply_binary_ufunc(&operands, BinaryOpKind::Maximum, |lhs, rhs| { + lhs.maximum(rhs) + })?, + ("minimum", 2) => apply_binary_ufunc(&operands, BinaryOpKind::Minimum, |lhs, rhs| { + lhs.minimum(rhs) + })?, + ("negative", 1) => apply_unary_ufunc(&operands, |tensor| { + engine::operations::arithmetic::neg(tensor) + })?, + ("absolute", 1) | ("abs", 1) => apply_unary_ufunc(&operands, |tensor| tensor.abs())?, + ("exp", 1) => apply_unary_ufunc(&operands, |tensor| tensor.exp())?, + ("log", 1) => apply_unary_ufunc(&operands, |tensor| tensor.log())?, + ("sin", 1) => apply_unary_ufunc(&operands, |tensor| tensor.sin())?, + ("cos", 1) => apply_unary_ufunc(&operands, |tensor| tensor.cos())?, + ("tan", 1) => apply_unary_ufunc(&operands, |tensor| tensor.tan())?, + ("sqrt", 1) => apply_unary_ufunc(&operands, |tensor| tensor.sqrt())?, + _ => return py_not_implemented(py), + }; + + let py_tensor = Py::new(py, PyTensor::from_tensor(result))?; + Ok(py_tensor.into_any()) + } + + pub(crate) fn tolist(&self) -> PyResult> { + if self.inner.ndim() == 0 { + Python::attach(|py| convert_tensor_to_python_scalar(&self.inner, py)) + } else { + Python::attach(|py| convert_tensor_to_python_list(&self.inner, py)) + } + } + + fn item(&self) -> PyResult> { + Python::attach(|py| convert_tensor_to_python_scalar(&self.inner, py)) + } +} diff --git a/bindings/src/tensor/pytensor/operations.rs b/bindings/src/tensor/pytensor/operations.rs index 7aac5402..1e27aec1 100644 --- a/bindings/src/tensor/pytensor/operations.rs +++ b/bindings/src/tensor/pytensor/operations.rs @@ -1,198 +1,198 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -#[pymethods] -impl PyTensor { - // Tensor operations - fn clone(&self) -> PyResult { - let result = self.inner.deep_clone().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn detach(&self) -> Self { - Self { - inner: self.inner.detach(), - } - } - - fn detach_(&mut self) { - self.inner.detach_inplace(); - } - - fn contiguous(&self) -> PyResult { - let result = self.inner.contiguous().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (*args, **kwargs))] - fn to(&self, args: &Bound, kwargs: Option<&Bound>) -> PyResult { - let mut dtype_spec: Option = None; - let mut device_spec: Option = None; - - if let Some(mapping) = kwargs { - for (key, value) in mapping.iter() { - let key_string = key.str()?.to_str()?.to_owned(); - match key_string.as_str() { - "dtype" => { - if !value.is_none() { - dtype_spec = Some(parse_dtype_like(&value)?); - } - } - "device" => { - if !value.is_none() { - device_spec = Some(parse_device_like(&value)?); - } - } - _ => { - return Err(PyTypeError::new_err(format!( - "to() got an unexpected keyword argument '{key_string}'" - ))); - } - } - } - } - - if args.len() > 1 { - return Err(PyTypeError::new_err(format!( - "to() takes at most 1 positional argument but {} were given", - args.len() - ))); - } - - if args.len() == 1 { - let arg0 = args.get_item(0)?; - if arg0.is_none() { - // Explicit None does nothing - } else if let Ok(py_device) = arg0.extract::() { - if device_spec.is_some() { - return Err(PyTypeError::new_err( - "to() received multiple device specifications", - )); - } - device_spec = Some(py_device.device()); - } else if let Ok(string_value) = arg0.extract::() { - match dtype::parse_dtype(&string_value) { - Ok(dtype) => { - if let Some(existing) = dtype_spec - && existing != dtype - { - return Err(PyTypeError::new_err( - "dtype specified both positionally and via keyword", - )); - } - dtype_spec = Some(dtype); - } - Err(_) => { - let device = Device::from_str(&string_value).map_err(|err| { - PyValueError::new_err(format!( - "Unsupported device specification '{string_value}': {err}" - )) - })?; - if device_spec.is_some() { - return Err(PyTypeError::new_err( - "to() received multiple device specifications", - )); - } - device_spec = Some(device); - } - } - } else { - return Err(PyTypeError::new_err( - "to() expects dtype strings, device strings, or Device objects", - )); - } - } - - let mut result = self.inner.clone(); - let mut mutated = false; - - if let Some(dtype) = dtype_spec - && result.dtype() != dtype - { - result = result.astype(dtype).map_err(_convert_error)?; - mutated = true; - } - - if let Some(device) = device_spec - && result.device() != device - { - result = result.to(device).map_err(_convert_error)?; - mutated = true; - } - - if mutated { - Ok(Self::from_tensor(result)) - } else { - Ok(Self { - inner: self.inner.clone(), - }) - } - } - - #[pyo3(signature = (min=None, max=None))] - pub fn clip(&self, min: Option<&Bound>, max: Option<&Bound>) -> PyResult { - let min_val = parse_clip_bound(min, "min")?; - let max_val = parse_clip_bound(max, "max")?; - let result = self.inner.clip(min_val, max_val).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (min=None, max=None))] - pub fn clamp(&self, min: Option<&Bound>, max: Option<&Bound>) -> PyResult { - let min_val = parse_clip_bound(min, "min")?; - let max_val = parse_clip_bound(max, "max")?; - let result = self.inner.clamp(min_val, max_val).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn clamp_min(&self, min: f64) -> PyResult { - let result = self.inner.clamp_min(min).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn clamp_max(&self, max: f64) -> PyResult { - let result = self.inner.clamp_max(max).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (decimals=0))] - pub fn round(&self, decimals: i32) -> PyResult { - let result = self.inner.round(decimals).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn floor(&self) -> PyResult { - let result = self.inner.floor().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn ceil(&self) -> PyResult { - let result = self.inner.ceil().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn sign(&self) -> PyResult { - let result = self.inner.sign().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn reciprocal(&self) -> PyResult { - let result = self.inner.reciprocal().map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - fn cpu(&self) -> PyResult { - let result = self.inner.to(Device::cpu()).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn astype(&self, dtype: &str) -> PyResult { - let dtype = dtype::parse_dtype(dtype)?; - let result = self.inner.astype(dtype).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - -} +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +#[pymethods] +impl PyTensor { + // Tensor operations + fn clone(&self) -> PyResult { + let result = self.inner.deep_clone().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + fn detach(&self) -> Self { + Self { + inner: self.inner.detach(), + } + } + + fn detach_(&mut self) { + self.inner.detach_inplace(); + } + + fn contiguous(&self) -> PyResult { + let result = self.inner.contiguous().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (*args, **kwargs))] + fn to(&self, args: &Bound, kwargs: Option<&Bound>) -> PyResult { + let mut dtype_spec: Option = None; + let mut device_spec: Option = None; + + if let Some(mapping) = kwargs { + for (key, value) in mapping.iter() { + let key_string = key.str()?.to_str()?.to_owned(); + match key_string.as_str() { + "dtype" => { + if !value.is_none() { + dtype_spec = Some(parse_dtype_like(&value)?); + } + } + "device" => { + if !value.is_none() { + device_spec = Some(parse_device_like(&value)?); + } + } + _ => { + return Err(PyTypeError::new_err(format!( + "to() got an unexpected keyword argument '{key_string}'" + ))); + } + } + } + } + + if args.len() > 1 { + return Err(PyTypeError::new_err(format!( + "to() takes at most 1 positional argument but {} were given", + args.len() + ))); + } + + if args.len() == 1 { + let arg0 = args.get_item(0)?; + if arg0.is_none() { + // Explicit None does nothing + } else if let Ok(py_device) = arg0.extract::() { + if device_spec.is_some() { + return Err(PyTypeError::new_err( + "to() received multiple device specifications", + )); + } + device_spec = Some(py_device.device()); + } else if let Ok(string_value) = arg0.extract::() { + match dtype::parse_dtype(&string_value) { + Ok(dtype) => { + if let Some(existing) = dtype_spec + && existing != dtype + { + return Err(PyTypeError::new_err( + "dtype specified both positionally and via keyword", + )); + } + dtype_spec = Some(dtype); + } + Err(_) => { + let device = Device::from_str(&string_value).map_err(|err| { + PyValueError::new_err(format!( + "Unsupported device specification '{string_value}': {err}" + )) + })?; + if device_spec.is_some() { + return Err(PyTypeError::new_err( + "to() received multiple device specifications", + )); + } + device_spec = Some(device); + } + } + } else { + return Err(PyTypeError::new_err( + "to() expects dtype strings, device strings, or Device objects", + )); + } + } + + let mut result = self.inner.clone(); + let mut mutated = false; + + if let Some(dtype) = dtype_spec + && result.dtype() != dtype + { + result = result.astype(dtype).map_err(_convert_error)?; + mutated = true; + } + + if let Some(device) = device_spec + && result.device() != device + { + result = result.to(device).map_err(_convert_error)?; + mutated = true; + } + + if mutated { + Ok(Self::from_tensor(result)) + } else { + Ok(Self { + inner: self.inner.clone(), + }) + } + } + + #[pyo3(signature = (min=None, max=None))] + pub fn clip(&self, min: Option<&Bound>, max: Option<&Bound>) -> PyResult { + let min_val = parse_clip_bound(min, "min")?; + let max_val = parse_clip_bound(max, "max")?; + let result = self.inner.clip(min_val, max_val).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (min=None, max=None))] + pub fn clamp(&self, min: Option<&Bound>, max: Option<&Bound>) -> PyResult { + let min_val = parse_clip_bound(min, "min")?; + let max_val = parse_clip_bound(max, "max")?; + let result = self.inner.clamp(min_val, max_val).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn clamp_min(&self, min: f64) -> PyResult { + let result = self.inner.clamp_min(min).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn clamp_max(&self, max: f64) -> PyResult { + let result = self.inner.clamp_max(max).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (decimals=0))] + pub fn round(&self, decimals: i32) -> PyResult { + let result = self.inner.round(decimals).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn floor(&self) -> PyResult { + let result = self.inner.floor().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn ceil(&self) -> PyResult { + let result = self.inner.ceil().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn sign(&self) -> PyResult { + let result = self.inner.sign().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn reciprocal(&self) -> PyResult { + let result = self.inner.reciprocal().map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + fn cpu(&self) -> PyResult { + let result = self.inner.to(Device::cpu()).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn astype(&self, dtype: &str) -> PyResult { + let dtype = dtype::parse_dtype(dtype)?; + let result = self.inner.astype(dtype).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } +} diff --git a/bindings/src/tensor/pytensor/properties.rs b/bindings/src/tensor/pytensor/properties.rs index 79336f35..9e09a8fc 100644 --- a/bindings/src/tensor/pytensor/properties.rs +++ b/bindings/src/tensor/pytensor/properties.rs @@ -1,350 +1,351 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -#[pymethods] -impl PyTensor { - #[classattr] - fn __array_priority__() -> f64 { - 1000.0 - } - - /// Create a new tensor from Python data - #[new] - #[pyo3(signature = (data=None, dtype=None, device=None, requires_grad=false))] - fn new( - data: Option<&Bound>, - dtype: Option<&str>, - device: Option<&PyDevice>, - requires_grad: Option, - ) -> PyResult { - let dtype = dtype::resolve_dtype_arg(dtype)?; - let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); - let requires_grad = requires_grad.unwrap_or(false); - - if let Some(value) = data { - let tensor = convert_python_data_to_tensor(value, dtype, device, requires_grad)?; - Ok(Self::from_tensor(tensor)) - } else { - let tensor = Tensor::empty(Shape::new(Vec::new()), dtype, device, requires_grad); - Ok(Self::from_tensor(tensor)) - } - } - - // Properties - #[getter] - pub fn shape(&self) -> ShapeSequence { - ShapeSequence::from_dims(self.inner.shape().dims().to_vec()) - } - - pub fn shape_vec(&self) -> Vec { - self.inner.shape().dims().to_vec() - } - - #[getter] - pub fn dtype(&self) -> String { - dtype::dtype_to_python_string(self.inner.dtype()).to_string() - } - - #[getter] - fn device(&self) -> String { - self.inner.device().to_string() - } - - #[getter] - fn _tensor(&self) -> Self { - Self { - inner: self.inner.clone(), - } - } - - #[setter] - #[allow(non_snake_case)] - fn set__tensor(&mut self, value: &PyTensor) { - self.inner = value.inner.clone(); - } - - #[getter] - pub fn requires_grad(&self) -> bool { - self.inner.requires_grad() - } - - #[getter] - fn is_leaf(&self) -> bool { - self.inner.is_leaf() - } - - #[getter] - fn has_grad(&self) -> bool { - if engine::autograd::get_gradient(&self.inner).is_some() { - return true; - } - - self.inner.has_grad() || self.inner.grad().is_some() - } - - #[getter] - fn grad(&self) -> PyResult> { - if let Some(grad) = engine::autograd::get_gradient(&self.inner) { - return Ok(Some(Self::from_tensor(grad))); - } - - if let Some(stored) = self.inner.grad() { - return Ok(Some(Self::from_tensor((**stored).clone()))); - } - - Ok(None) - } - - #[getter] - fn size(&self) -> usize { - self.inner.numel() - } - - #[getter] - fn itemsize(&self) -> usize { - match self.inner.dtype() { - DataType::Float32 | DataType::Int32 => 4, - DataType::Float64 | DataType::Int64 => 8, - DataType::Bool => 1, - } - } - - #[getter] - fn nbytes(&self) -> usize { - self.size() * self.itemsize() - } - - /// Get memory usage in bytes - fn memory_usage_bytes(&self) -> usize { - self.inner.memory_usage_bytes() - } - - #[getter] - fn strides<'py>(&self, py: Python<'py>) -> PyResult> { - PyTuple::new(py, self.inner.strides().as_slice()) - } - - // Basic tensor info methods - pub fn ndim(&self) -> usize { - self.inner.ndim() - } - - fn numel(&self) -> usize { - self.inner.numel() - } - - fn is_contiguous(&self) -> bool { - self.inner.is_contiguous() - } - - // Tensor manipulation methods - #[pyo3(signature = (*shape))] - pub fn reshape(&self, shape: &Bound) -> PyResult { - let dims = normalize_variadic_isize_args(shape, "shape")?; - let reshaped = engine::operations::reshape_with_inference(&self.inner, dims) - .map_err(_convert_error)?; - Ok(Self::from_tensor(reshaped)) - } - - #[pyo3(signature = (*shape))] - pub fn view(&self, shape: &Bound) -> PyResult { - self.reshape(shape) - } - - #[pyo3(signature = (dim0=0, dim1=1))] - pub fn transpose(&self, dim0: Option, dim1: Option) -> PyResult { - let dim0 = dim0.unwrap_or(0); - let dim1 = dim1.unwrap_or(1); - let result = self.inner.transpose(dim0, dim1).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (*dims))] - pub fn permute(&self, dims: &Bound) -> PyResult { - let dims_vec = normalize_variadic_isize_args(dims, "dims")?; - let result = self.inner.permute(dims_vec).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn movedim(&self, source: &Bound, destination: &Bound) -> PyResult { - let src_vec: Vec = match source.extract::() { - Ok(v) => vec![v], - Err(_) => source.extract()?, - }; - let dst_vec: Vec = match destination.extract::() { - Ok(v) => vec![v], - Err(_) => destination.extract()?, - }; - let result = engine::operations::shape_ops::movedim(&self.inner, &src_vec, &dst_vec) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(name = "moveaxis")] - #[pyo3(signature = (source, destination))] - pub fn moveaxis_alias( - &self, - source: &Bound, - destination: &Bound, - ) -> PyResult { - self.movedim(source, destination) - } - - #[pyo3(name = "swapaxes")] - #[pyo3(signature = (dim0, dim1))] - pub fn swapaxes_alias(&self, dim0: isize, dim1: isize) -> PyResult { - self.transpose(Some(dim0), Some(dim1)) - } - - #[pyo3(name = "swapdims")] - #[pyo3(signature = (dim0, dim1))] - pub fn swapdims_alias(&self, dim0: isize, dim1: isize) -> PyResult { - self.transpose(Some(dim0), Some(dim1)) - } - - #[pyo3(signature = (dim=None))] - pub fn squeeze(&self, dim: Option) -> PyResult { - let result = - engine::operations::shape_ops::squeeze(&self.inner, dim).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (dim))] - pub fn unsqueeze(&self, dim: isize) -> PyResult { - let result = - engine::operations::shape_ops::unsqueeze(&self.inner, dim).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (*dims))] - pub fn expand(&self, dims: &Bound) -> PyResult { - let dims_vec = normalize_variadic_isize_args(dims, "shape")?; - let result = self.inner.expand(dims_vec).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (*repeats))] - pub fn repeat(&self, repeats: &Bound) -> PyResult { - let repeats_any = if repeats.len() == 1 { - let first = repeats.get_item(0)?; - if first.cast::().is_ok() { - first.clone().into_any() - } else { - repeats.clone().into_any() - } - } else { - repeats.clone().into_any() - }; - let repeat_vec = normalize_repeat_spec(&repeats_any)?; - let result = self.inner.repeat(repeat_vec).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn flip(&self, dims: &Bound) -> PyResult { - let dims_vec = normalize_required_axes(dims, "dims")?; - let result = - engine::operations::shape_ops::flip(&self.inner, &dims_vec).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (shifts, dims=None))] - pub fn roll(&self, shifts: &Bound, dims: Option<&Bound>) -> PyResult { - let shift_vec = normalize_roll_shifts(shifts)?; - let dims_vec = normalize_optional_axes(dims)?; - let dims_ref = dims_vec.as_deref(); - let result = engine::operations::shape_ops::roll(&self.inner, &shift_vec, dims_ref) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (repeats, dim=None, output_size=None))] - pub fn repeat_interleave( - &self, - repeats: &Bound, - dim: Option, - output_size: Option, - ) -> PyResult { - if let Ok(value) = repeats.extract::() { - let result = engine::operations::shape_ops::repeat_interleave( - &self.inner, - RepeatInterleaveSpec::Scalar(value), - dim, - output_size, - ) - .map_err(_convert_error)?; - return Ok(Self::from_tensor(result)); - } - - if let Ok(seq) = repeats.extract::>() { - let mut converted = Vec::with_capacity(seq.len()); - for value in seq { - if value < 0 { - return Err(PyValueError::new_err( - "repeat_interleave: repeats must be non-negative integers", - )); - } - let value = usize::try_from(value).map_err(|_| { - PyValueError::new_err("repeat_interleave: repeat value exceeds platform limits") - })?; - converted.push(value); - } - let result = engine::operations::shape_ops::repeat_interleave( - &self.inner, - RepeatInterleaveSpec::Slice(&converted), - dim, - output_size, - ) - .map_err(_convert_error)?; - return Ok(Self::from_tensor(result)); - } - - if let Ok(py_tensor) = repeats.extract::>() { - let result = engine::operations::shape_ops::repeat_interleave( - &self.inner, - RepeatInterleaveSpec::Tensor(py_tensor.tensor()), - dim, - output_size, - ) - .map_err(_convert_error)?; - return Ok(Self::from_tensor(result)); - } - - if let Ok(bound_attr) = repeats.getattr("_tensor") - && let Ok(py_tensor) = bound_attr.extract::>() - { - let result = engine::operations::shape_ops::repeat_interleave( - &self.inner, - RepeatInterleaveSpec::Tensor(py_tensor.tensor()), - dim, - output_size, - ) - .map_err(_convert_error)?; - return Ok(Self::from_tensor(result)); - } - - Err(PyTypeError::new_err( - "repeat_interleave: repeats must be an int, sequence of ints, or Tensor", - )) - } - - #[pyo3(signature = (dim, start, length))] - pub fn narrow(&self, dim: isize, start: usize, length: usize) -> PyResult { - let result = engine::operations::shape_ops::narrow(&self.inner, dim, start, length) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (start_dim=0, end_dim=-1))] - pub fn flatten(&self, start_dim: isize, end_dim: isize) -> PyResult { - let result = engine::operations::shape_ops::flatten(&self.inner, start_dim, end_dim) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - pub fn ravel(&self) -> PyResult { - self.flatten(0, -1) - } -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +#[pymethods] +impl PyTensor { + #[classattr] + fn __array_priority__() -> f64 { + 1000.0 + } + + /// Create a new tensor from Python data + #[new] + #[pyo3(signature = (data=None, dtype=None, device=None, requires_grad=false))] + fn new( + data: Option<&Bound>, + dtype: Option<&str>, + device: Option<&PyDevice>, + requires_grad: Option, + ) -> PyResult { + let dtype = dtype::resolve_dtype_arg(dtype)?; + let device = device.map(|d| d.device()).unwrap_or_else(Device::cpu); + let requires_grad = requires_grad.unwrap_or(false); + + if let Some(value) = data { + let tensor = convert_python_data_to_tensor(value, dtype, device, requires_grad)?; + Ok(Self::from_tensor(tensor)) + } else { + let tensor = Tensor::empty(Shape::new(Vec::new()), dtype, device, requires_grad); + Ok(Self::from_tensor(tensor)) + } + } + + // Properties + #[getter] + pub fn shape(&self) -> ShapeSequence { + ShapeSequence::from_dims(self.inner.shape().dims().to_vec()) + } + + pub fn shape_vec(&self) -> Vec { + self.inner.shape().dims().to_vec() + } + + #[getter] + pub fn dtype(&self) -> String { + dtype::dtype_to_python_string(self.inner.dtype()).to_string() + } + + #[getter] + pub(crate) fn device(&self) -> String { + self.inner.device().to_string() + } + + #[getter] + fn _tensor(&self) -> Self { + Self { + inner: self.inner.clone(), + } + } + + #[setter] + #[allow(non_snake_case)] + fn set__tensor(&mut self, value: &PyTensor) { + self.inner = value.inner.clone(); + } + + #[getter] + pub fn requires_grad(&self) -> bool { + self.inner.requires_grad() + } + + #[getter] + pub(crate) fn is_leaf(&self) -> bool { + self.inner.is_leaf() + } + + #[getter] + fn has_grad(&self) -> bool { + if engine::autograd::get_gradient(&self.inner).is_some() { + return true; + } + + self.inner.has_grad() || self.inner.grad().is_some() + } + + #[getter] + fn grad(&self) -> PyResult> { + if let Some(grad) = engine::autograd::get_gradient(&self.inner) { + return Ok(Some(Self::from_tensor(grad))); + } + + if let Some(stored) = self.inner.grad() { + return Ok(Some(Self::from_tensor((**stored).clone()))); + } + + Ok(None) + } + + #[getter] + fn size(&self) -> usize { + self.inner.numel() + } + + #[getter] + fn itemsize(&self) -> usize { + match self.inner.dtype() { + DataType::Float32 | DataType::Int32 => 4, + DataType::Float64 | DataType::Int64 => 8, + DataType::Bool => 1, + } + } + + #[getter] + fn nbytes(&self) -> usize { + self.size() * self.itemsize() + } + + /// Get memory usage in bytes + fn memory_usage_bytes(&self) -> usize { + self.inner.memory_usage_bytes() + } + + #[getter] + fn strides<'py>(&self, py: Python<'py>) -> PyResult> { + PyTuple::new(py, self.inner.strides().as_slice()) + } + + // Basic tensor info methods + pub fn ndim(&self) -> usize { + self.inner.ndim() + } + + fn numel(&self) -> usize { + self.inner.numel() + } + + fn is_contiguous(&self) -> bool { + self.inner.is_contiguous() + } + + // Tensor manipulation methods + #[pyo3(signature = (*shape))] + pub fn reshape(&self, shape: &Bound) -> PyResult { + let dims = normalize_variadic_isize_args(shape, "shape")?; + let reshaped = engine::operations::reshape_with_inference(&self.inner, dims) + .map_err(_convert_error)?; + Ok(Self::from_tensor(reshaped)) + } + + #[pyo3(signature = (*shape))] + pub fn view(&self, shape: &Bound) -> PyResult { + self.reshape(shape) + } + + #[pyo3(signature = (dim0=0, dim1=1))] + pub fn transpose(&self, dim0: Option, dim1: Option) -> PyResult { + let dim0 = dim0.unwrap_or(0); + let dim1 = dim1.unwrap_or(1); + let result = self.inner.transpose(dim0, dim1).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (*dims))] + pub fn permute(&self, dims: &Bound) -> PyResult { + let dims_vec = normalize_variadic_isize_args(dims, "dims")?; + let result = self.inner.permute(dims_vec).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn movedim(&self, source: &Bound, destination: &Bound) -> PyResult { + let src_vec: Vec = match source.extract::() { + Ok(v) => vec![v], + Err(_) => source.extract()?, + }; + let dst_vec: Vec = match destination.extract::() { + Ok(v) => vec![v], + Err(_) => destination.extract()?, + }; + let result = engine::operations::shape_ops::movedim(&self.inner, &src_vec, &dst_vec) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(name = "moveaxis")] + #[pyo3(signature = (source, destination))] + pub fn moveaxis_alias( + &self, + source: &Bound, + destination: &Bound, + ) -> PyResult { + self.movedim(source, destination) + } + + #[pyo3(name = "swapaxes")] + #[pyo3(signature = (dim0, dim1))] + pub fn swapaxes_alias(&self, dim0: isize, dim1: isize) -> PyResult { + self.transpose(Some(dim0), Some(dim1)) + } + + #[pyo3(name = "swapdims")] + #[pyo3(signature = (dim0, dim1))] + pub fn swapdims_alias(&self, dim0: isize, dim1: isize) -> PyResult { + self.transpose(Some(dim0), Some(dim1)) + } + + #[pyo3(signature = (dim=None))] + pub fn squeeze(&self, dim: Option) -> PyResult { + let result = + engine::operations::shape_ops::squeeze(&self.inner, dim).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (dim))] + pub fn unsqueeze(&self, dim: isize) -> PyResult { + let result = + engine::operations::shape_ops::unsqueeze(&self.inner, dim).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (*dims))] + pub fn expand(&self, dims: &Bound) -> PyResult { + let dims_vec = normalize_variadic_isize_args(dims, "shape")?; + let result = self.inner.expand(dims_vec).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (*repeats))] + pub fn repeat(&self, repeats: &Bound) -> PyResult { + let repeats_any = if repeats.len() == 1 { + let first = repeats.get_item(0)?; + if first.cast::().is_ok() { + first.clone().into_any() + } else { + repeats.clone().into_any() + } + } else { + repeats.clone().into_any() + }; + let repeat_vec = normalize_repeat_spec(&repeats_any)?; + let result = self.inner.repeat(repeat_vec).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn flip(&self, dims: &Bound) -> PyResult { + let dims_vec = normalize_required_axes(dims, "dims")?; + let result = + engine::operations::shape_ops::flip(&self.inner, &dims_vec).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (shifts, dims=None))] + pub fn roll(&self, shifts: &Bound, dims: Option<&Bound>) -> PyResult { + let shift_vec = normalize_roll_shifts(shifts)?; + let dims_vec = normalize_optional_axes(dims)?; + let dims_ref = dims_vec.as_deref(); + let result = engine::operations::shape_ops::roll(&self.inner, &shift_vec, dims_ref) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (repeats, dim=None, output_size=None))] + pub fn repeat_interleave( + &self, + repeats: &Bound, + dim: Option, + output_size: Option, + ) -> PyResult { + if let Ok(value) = repeats.extract::() { + let result = engine::operations::shape_ops::repeat_interleave( + &self.inner, + RepeatInterleaveSpec::Scalar(value), + dim, + output_size, + ) + .map_err(_convert_error)?; + return Ok(Self::from_tensor(result)); + } + + if let Ok(seq) = repeats.extract::>() { + let mut converted = Vec::with_capacity(seq.len()); + for value in seq { + if value < 0 { + return Err(PyValueError::new_err( + "repeat_interleave: repeats must be non-negative integers", + )); + } + let value = usize::try_from(value).map_err(|_| { + PyValueError::new_err("repeat_interleave: repeat value exceeds platform limits") + })?; + converted.push(value); + } + let result = engine::operations::shape_ops::repeat_interleave( + &self.inner, + RepeatInterleaveSpec::Slice(&converted), + dim, + output_size, + ) + .map_err(_convert_error)?; + return Ok(Self::from_tensor(result)); + } + + if let Ok(py_tensor) = repeats.extract::>() { + let result = engine::operations::shape_ops::repeat_interleave( + &self.inner, + RepeatInterleaveSpec::Tensor(py_tensor.tensor()), + dim, + output_size, + ) + .map_err(_convert_error)?; + return Ok(Self::from_tensor(result)); + } + + if let Ok(bound_attr) = repeats.getattr("_tensor") + && let Ok(py_tensor) = bound_attr.extract::>() + { + let result = engine::operations::shape_ops::repeat_interleave( + &self.inner, + RepeatInterleaveSpec::Tensor(py_tensor.tensor()), + dim, + output_size, + ) + .map_err(_convert_error)?; + return Ok(Self::from_tensor(result)); + } + + Err(PyTypeError::new_err( + "repeat_interleave: repeats must be an int, sequence of ints, or Tensor", + )) + } + + #[pyo3(signature = (dim, start, length))] + pub fn narrow(&self, dim: isize, start: usize, length: usize) -> PyResult { + let result = engine::operations::shape_ops::narrow(&self.inner, dim, start, length) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (start_dim=0, end_dim=-1))] + pub fn flatten(&self, start_dim: isize, end_dim: isize) -> PyResult { + let result = engine::operations::shape_ops::flatten(&self.inner, start_dim, end_dim) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + pub fn ravel(&self) -> PyResult { + self.flatten(0, -1) + } +} diff --git a/bindings/src/tensor/pytensor/reduction.rs b/bindings/src/tensor/pytensor/reduction.rs index 21634d0c..d13ce9b6 100644 --- a/bindings/src/tensor/pytensor/reduction.rs +++ b/bindings/src/tensor/pytensor/reduction.rs @@ -1,364 +1,365 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -#[pymethods] -impl PyTensor { - // Reduction operations - #[pyo3(signature = (dim=None, keepdim=false))] - pub fn sum(&self, dim: Option<&Bound>, keepdim: Option) -> PyResult { - let keepdim = keepdim.unwrap_or(false); - let dims = normalize_optional_axes(dim)?; - let result = self.inner.sum(dims, keepdim).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (dim=None, keepdim=false))] - pub fn nansum(&self, dim: Option<&Bound>, keepdim: Option) -> PyResult { - let keepdim = keepdim.unwrap_or(false); - let dims = normalize_optional_axes(dim)?; - let result = self.inner.nansum(dims, keepdim).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (dim=None, keepdim=false))] - pub fn logsumexp(&self, dim: Option<&Bound>, keepdim: Option) -> PyResult { - let keepdim = keepdim.unwrap_or(false); - let dims = normalize_optional_axes(dim)?; - match self.inner.logsumexp(dims, keepdim) { - Ok(result) => Ok(Self::from_tensor(result)), - Err(err @ MinitensorError::InvalidOperation { .. }) => { - Err(PyRuntimeError::new_err(err.detailed_message())) - } - Err(err) => Err(_convert_error(err)), - } - } - - #[pyo3(signature = (dim=None, keepdim=false))] - pub fn prod(&self, dim: Option<&Bound>, keepdim: Option) -> PyResult { - let keepdim = keepdim.unwrap_or(false); - let dims = normalize_optional_axes(dim)?; - let result = self.inner.prod(dims, keepdim).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (dim=None, keepdim=false))] - pub fn mean(&self, dim: Option<&Bound>, keepdim: Option) -> PyResult { - let keepdim = keepdim.unwrap_or(false); - let dims = normalize_optional_axes(dim)?; - let result = self.inner.mean(dims, keepdim).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (dim=None, keepdim=false))] - pub fn nanmean(&self, dim: Option<&Bound>, keepdim: Option) -> PyResult { - let keepdim = keepdim.unwrap_or(false); - let dims = normalize_optional_axes(dim)?; - let result = self.inner.nanmean(dims, keepdim).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (dim=None, keepdim=false))] - pub fn all(&self, dim: Option, keepdim: Option) -> PyResult { - let keepdim = keepdim.unwrap_or(false); - let result = self.inner.all(dim, keepdim).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (dim=None, keepdim=false))] - pub fn any(&self, dim: Option, keepdim: Option) -> PyResult { - let keepdim = keepdim.unwrap_or(false); - let result = self.inner.any(dim, keepdim).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (dim))] - pub fn cumsum(&self, dim: isize) -> PyResult { - let result = self.inner.cumsum(dim).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (dim))] - pub fn cumprod(&self, dim: isize) -> PyResult { - let result = self.inner.cumprod(dim).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (dim=None, keepdim=false))] - pub fn max<'py>( - &self, - py: Python<'py>, - dim: Option, - keepdim: Option, - ) -> PyResult> { - let keepdim = keepdim.unwrap_or(false); - if let Some(dim) = dim { - let (values, indices) = self - .inner - .max_with_indices(dim, keepdim) - .map_err(_convert_error)?; - let values = Py::new(py, PyTensor::from_tensor(values))?.into_any(); - let indices = Py::new(py, PyTensor::from_tensor(indices))?.into_any(); - let tuple = PyTuple::new(py, [values, indices])?; - Ok(tuple.into_any().unbind()) - } else { - Ok(Py::new(py, self.max_values(None, keepdim)?)?.into_any()) - } - } - - #[pyo3(signature = (dim=None, keepdim=false))] - pub fn nanmax<'py>( - &self, - py: Python<'py>, - dim: Option, - keepdim: Option, - ) -> PyResult> { - let keepdim = keepdim.unwrap_or(false); - if let Some(dim) = dim { - let (values, indices) = self - .inner - .nanmax_with_indices(dim, keepdim) - .map_err(_convert_error)?; - let values = Py::new(py, PyTensor::from_tensor(values))?.into_any(); - let indices = Py::new(py, PyTensor::from_tensor(indices))?.into_any(); - let tuple = PyTuple::new(py, [values, indices])?; - Ok(tuple.into_any().unbind()) - } else { - Ok(Py::new(py, self.nanmax_values(None, keepdim)?)?.into_any()) - } - } - - #[pyo3(signature = (dim=None, keepdim=false))] - pub fn min<'py>( - &self, - py: Python<'py>, - dim: Option, - keepdim: Option, - ) -> PyResult> { - let keepdim = keepdim.unwrap_or(false); - if let Some(dim) = dim { - let (values, indices) = self - .inner - .min_with_indices(dim, keepdim) - .map_err(_convert_error)?; - let values = Py::new(py, PyTensor::from_tensor(values))?.into_any(); - let indices = Py::new(py, PyTensor::from_tensor(indices))?.into_any(); - let tuple = PyTuple::new(py, [values, indices])?; - Ok(tuple.into_any().unbind()) - } else { - Ok(Py::new(py, self.min_values(None, keepdim)?)?.into_any()) - } - } - - #[pyo3(signature = (dim=None, keepdim=false))] - pub fn nanmin<'py>( - &self, - py: Python<'py>, - dim: Option, - keepdim: Option, - ) -> PyResult> { - let keepdim = keepdim.unwrap_or(false); - if let Some(dim) = dim { - let (values, indices) = self - .inner - .nanmin_with_indices(dim, keepdim) - .map_err(_convert_error)?; - let values = Py::new(py, PyTensor::from_tensor(values))?.into_any(); - let indices = Py::new(py, PyTensor::from_tensor(indices))?.into_any(); - let tuple = PyTuple::new(py, [values, indices])?; - Ok(tuple.into_any().unbind()) - } else { - Ok(Py::new(py, self.nanmin_values(None, keepdim)?)?.into_any()) - } - } - - #[pyo3(signature = (dim=None, keepdim=false))] - pub fn median<'py>( - &self, - py: Python<'py>, - dim: Option, - keepdim: Option, - ) -> PyResult> { - let keepdim = keepdim.unwrap_or(false); - let (values, indices) = self.median_with_indices(dim, keepdim)?; - if dim.is_some() { - let indices = indices.ok_or_else(|| { - PyRuntimeError::new_err("median returned no indices for the requested dimension") - })?; - let values = Py::new(py, values)?.into_any(); - let indices = Py::new(py, indices)?.into_any(); - let tuple = PyTuple::new(py, [values, indices])?; - Ok(tuple.into_any().unbind()) - } else { - Ok(Py::new(py, values)?.into_any()) - } - } - - #[pyo3(signature = (dim=None, keepdim=false))] - pub fn nanmedian(&self, dim: Option, keepdim: Option) -> PyResult { - let keepdim = keepdim.unwrap_or(false); - let result = self.inner.nanmedian(dim, keepdim).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (q, dim=None, keepdim=false, interpolation="linear"))] - pub fn quantile( - &self, - q: &Bound, - dim: Option, - keepdim: Option, - interpolation: Option<&str>, - ) -> PyResult { - let keepdim = keepdim.unwrap_or(false); - let interpolation = parse_quantile_interpolation(interpolation)?; - match parse_quantile_arg(q)? { - QuantileArg::Scalar(prob) => { - let result = self - .inner - .quantile(prob, dim, keepdim, interpolation) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - QuantileArg::Multiple(qs) => { - let result = self - .inner - .quantiles(&qs, dim, keepdim, interpolation) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - } - } - - #[pyo3(signature = (q, dim=None, keepdim=false, interpolation="linear"))] - pub fn nanquantile( - &self, - q: &Bound, - dim: Option, - keepdim: Option, - interpolation: Option<&str>, - ) -> PyResult { - let keepdim = keepdim.unwrap_or(false); - let interpolation = parse_quantile_interpolation(interpolation)?; - match parse_quantile_arg(q)? { - QuantileArg::Scalar(prob) => { - let result = self - .inner - .nanquantile(prob, dim, keepdim, interpolation) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - QuantileArg::Multiple(qs) => { - let result = self - .inner - .nanquantiles(&qs, dim, keepdim, interpolation) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - } - } - - #[pyo3(signature = (dim=None, keepdim=false))] - pub fn argmax(&self, dim: Option, keepdim: Option) -> PyResult { - let keepdim = keepdim.unwrap_or(false); - let result = self.inner.argmax(dim, keepdim).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (dim=None, keepdim=false))] - pub fn argmin(&self, dim: Option, keepdim: Option) -> PyResult { - let keepdim = keepdim.unwrap_or(false); - let result = self.inner.argmin(dim, keepdim).map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (k, dim=None, largest=true, sorted=true))] - pub fn topk( - &self, - k: usize, - dim: Option, - largest: Option, - sorted: Option, - ) -> PyResult<(Self, Self)> { - let largest = largest.unwrap_or(true); - let sorted = sorted.unwrap_or(true); - match self.inner.topk(k, dim, largest, sorted) { - Ok((values, indices)) => Ok((Self::from_tensor(values), Self::from_tensor(indices))), - Err(err @ MinitensorError::InvalidArgument { .. }) => { - Err(PyRuntimeError::new_err(err.detailed_message())) - } - Err(err) => Err(_convert_error(err)), - } - } - - #[pyo3(signature = (dim=None, descending=false, stable=false))] - pub fn sort( - &self, - dim: Option, - descending: Option, - stable: Option, - ) -> PyResult<(Self, Self)> { - let descending = descending.unwrap_or(false); - let stable = stable.unwrap_or(false); - match self.inner.sort(dim, descending, stable) { - Ok((values, indices)) => Ok((Self::from_tensor(values), Self::from_tensor(indices))), - Err(err @ MinitensorError::InvalidArgument { .. }) => { - Err(PyRuntimeError::new_err(err.detailed_message())) - } - Err(err) => Err(_convert_error(err)), - } - } - - #[pyo3(signature = (dim=None, descending=false, stable=false))] - pub fn argsort( - &self, - dim: Option, - descending: Option, - stable: Option, - ) -> PyResult { - let descending = descending.unwrap_or(false); - let stable = stable.unwrap_or(false); - match self.inner.argsort(dim, descending, stable) { - Ok(indices) => Ok(Self::from_tensor(indices)), - Err(err @ MinitensorError::InvalidArgument { .. }) => { - Err(PyRuntimeError::new_err(err.detailed_message())) - } - Err(err) => Err(_convert_error(err)), - } - } - - #[pyo3(signature = (dim=None, unbiased=true, keepdim=false))] - pub fn std( - &self, - dim: Option<&Bound>, - unbiased: Option, - keepdim: Option, - ) -> PyResult { - let keepdim = keepdim.unwrap_or(false); - let unbiased = unbiased.unwrap_or(true); - let dims = normalize_optional_axes(dim)?; - let result = self - .inner - .std(dims, keepdim, unbiased) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } - - #[pyo3(signature = (dim=None, unbiased=true, keepdim=false))] - pub fn var( - &self, - dim: Option<&Bound>, - unbiased: Option, - keepdim: Option, - ) -> PyResult { - let keepdim = keepdim.unwrap_or(false); - let unbiased = unbiased.unwrap_or(true); - let dims = normalize_optional_axes(dim)?; - let result = self - .inner - .var(dims, keepdim, unbiased) - .map_err(_convert_error)?; - Ok(Self::from_tensor(result)) - } -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +#[pymethods] +impl PyTensor { + // Reduction operations + #[pyo3(signature = (dim=None, keepdim=false))] + pub fn sum(&self, dim: Option<&Bound>, keepdim: Option) -> PyResult { + let keepdim = keepdim.unwrap_or(false); + let dims = normalize_optional_axes(dim)?; + let result = self.inner.sum(dims, keepdim).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (dim=None, keepdim=false))] + pub fn nansum(&self, dim: Option<&Bound>, keepdim: Option) -> PyResult { + let keepdim = keepdim.unwrap_or(false); + let dims = normalize_optional_axes(dim)?; + let result = self.inner.nansum(dims, keepdim).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (dim=None, keepdim=false))] + pub fn logsumexp(&self, dim: Option<&Bound>, keepdim: Option) -> PyResult { + let keepdim = keepdim.unwrap_or(false); + let dims = normalize_optional_axes(dim)?; + match self.inner.logsumexp(dims, keepdim) { + Ok(result) => Ok(Self::from_tensor(result)), + Err(err @ MinitensorError::InvalidOperation { .. }) => { + Err(PyRuntimeError::new_err(err.detailed_message())) + } + Err(err) => Err(_convert_error(err)), + } + } + + #[pyo3(signature = (dim=None, keepdim=false))] + pub fn prod(&self, dim: Option<&Bound>, keepdim: Option) -> PyResult { + let keepdim = keepdim.unwrap_or(false); + let dims = normalize_optional_axes(dim)?; + let result = self.inner.prod(dims, keepdim).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (dim=None, keepdim=false))] + pub fn mean(&self, dim: Option<&Bound>, keepdim: Option) -> PyResult { + let keepdim = keepdim.unwrap_or(false); + let dims = normalize_optional_axes(dim)?; + let result = self.inner.mean(dims, keepdim).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (dim=None, keepdim=false))] + pub fn nanmean(&self, dim: Option<&Bound>, keepdim: Option) -> PyResult { + let keepdim = keepdim.unwrap_or(false); + let dims = normalize_optional_axes(dim)?; + let result = self.inner.nanmean(dims, keepdim).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (dim=None, keepdim=false))] + pub fn all(&self, dim: Option, keepdim: Option) -> PyResult { + let keepdim = keepdim.unwrap_or(false); + let result = self.inner.all(dim, keepdim).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (dim=None, keepdim=false))] + pub fn any(&self, dim: Option, keepdim: Option) -> PyResult { + let keepdim = keepdim.unwrap_or(false); + let result = self.inner.any(dim, keepdim).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (dim))] + pub fn cumsum(&self, dim: isize) -> PyResult { + let result = self.inner.cumsum(dim).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (dim))] + pub fn cumprod(&self, dim: isize) -> PyResult { + let result = self.inner.cumprod(dim).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (dim=None, keepdim=false))] + pub fn max<'py>( + &self, + py: Python<'py>, + dim: Option, + keepdim: Option, + ) -> PyResult> { + let keepdim = keepdim.unwrap_or(false); + if let Some(dim) = dim { + let (values, indices) = self + .inner + .max_with_indices(dim, keepdim) + .map_err(_convert_error)?; + let values = Py::new(py, PyTensor::from_tensor(values))?.into_any(); + let indices = Py::new(py, PyTensor::from_tensor(indices))?.into_any(); + let tuple = PyTuple::new(py, [values, indices])?; + Ok(tuple.into_any().unbind()) + } else { + Ok(Py::new(py, self.max_values(None, keepdim)?)?.into_any()) + } + } + + #[pyo3(signature = (dim=None, keepdim=false))] + pub fn nanmax<'py>( + &self, + py: Python<'py>, + dim: Option, + keepdim: Option, + ) -> PyResult> { + let keepdim = keepdim.unwrap_or(false); + if let Some(dim) = dim { + let (values, indices) = self + .inner + .nanmax_with_indices(dim, keepdim) + .map_err(_convert_error)?; + let values = Py::new(py, PyTensor::from_tensor(values))?.into_any(); + let indices = Py::new(py, PyTensor::from_tensor(indices))?.into_any(); + let tuple = PyTuple::new(py, [values, indices])?; + Ok(tuple.into_any().unbind()) + } else { + Ok(Py::new(py, self.nanmax_values(None, keepdim)?)?.into_any()) + } + } + + #[pyo3(signature = (dim=None, keepdim=false))] + pub fn min<'py>( + &self, + py: Python<'py>, + dim: Option, + keepdim: Option, + ) -> PyResult> { + let keepdim = keepdim.unwrap_or(false); + if let Some(dim) = dim { + let (values, indices) = self + .inner + .min_with_indices(dim, keepdim) + .map_err(_convert_error)?; + let values = Py::new(py, PyTensor::from_tensor(values))?.into_any(); + let indices = Py::new(py, PyTensor::from_tensor(indices))?.into_any(); + let tuple = PyTuple::new(py, [values, indices])?; + Ok(tuple.into_any().unbind()) + } else { + Ok(Py::new(py, self.min_values(None, keepdim)?)?.into_any()) + } + } + + #[pyo3(signature = (dim=None, keepdim=false))] + pub fn nanmin<'py>( + &self, + py: Python<'py>, + dim: Option, + keepdim: Option, + ) -> PyResult> { + let keepdim = keepdim.unwrap_or(false); + if let Some(dim) = dim { + let (values, indices) = self + .inner + .nanmin_with_indices(dim, keepdim) + .map_err(_convert_error)?; + let values = Py::new(py, PyTensor::from_tensor(values))?.into_any(); + let indices = Py::new(py, PyTensor::from_tensor(indices))?.into_any(); + let tuple = PyTuple::new(py, [values, indices])?; + Ok(tuple.into_any().unbind()) + } else { + Ok(Py::new(py, self.nanmin_values(None, keepdim)?)?.into_any()) + } + } + + #[pyo3(signature = (dim=None, keepdim=false))] + pub fn median<'py>( + &self, + py: Python<'py>, + dim: Option, + keepdim: Option, + ) -> PyResult> { + let keepdim = keepdim.unwrap_or(false); + let (values, indices) = self.median_with_indices(dim, keepdim)?; + if dim.is_some() { + let indices = indices.ok_or_else(|| { + PyRuntimeError::new_err("median returned no indices for the requested dimension") + })?; + let values = Py::new(py, values)?.into_any(); + let indices = Py::new(py, indices)?.into_any(); + let tuple = PyTuple::new(py, [values, indices])?; + Ok(tuple.into_any().unbind()) + } else { + Ok(Py::new(py, values)?.into_any()) + } + } + + #[pyo3(signature = (dim=None, keepdim=false))] + pub fn nanmedian(&self, dim: Option, keepdim: Option) -> PyResult { + let keepdim = keepdim.unwrap_or(false); + let result = self.inner.nanmedian(dim, keepdim).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (q, dim=None, keepdim=false, interpolation="linear"))] + pub fn quantile( + &self, + q: &Bound, + dim: Option, + keepdim: Option, + interpolation: Option<&str>, + ) -> PyResult { + let keepdim = keepdim.unwrap_or(false); + let interpolation = parse_quantile_interpolation(interpolation)?; + match parse_quantile_arg(q)? { + QuantileArg::Scalar(prob) => { + let result = self + .inner + .quantile(prob, dim, keepdim, interpolation) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + QuantileArg::Multiple(qs) => { + let result = self + .inner + .quantiles(&qs, dim, keepdim, interpolation) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + } + } + + #[pyo3(signature = (q, dim=None, keepdim=false, interpolation="linear"))] + pub fn nanquantile( + &self, + q: &Bound, + dim: Option, + keepdim: Option, + interpolation: Option<&str>, + ) -> PyResult { + let keepdim = keepdim.unwrap_or(false); + let interpolation = parse_quantile_interpolation(interpolation)?; + match parse_quantile_arg(q)? { + QuantileArg::Scalar(prob) => { + let result = self + .inner + .nanquantile(prob, dim, keepdim, interpolation) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + QuantileArg::Multiple(qs) => { + let result = self + .inner + .nanquantiles(&qs, dim, keepdim, interpolation) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + } + } + + #[pyo3(signature = (dim=None, keepdim=false))] + pub fn argmax(&self, dim: Option, keepdim: Option) -> PyResult { + let keepdim = keepdim.unwrap_or(false); + let result = self.inner.argmax(dim, keepdim).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (dim=None, keepdim=false))] + pub fn argmin(&self, dim: Option, keepdim: Option) -> PyResult { + let keepdim = keepdim.unwrap_or(false); + let result = self.inner.argmin(dim, keepdim).map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (k, dim=None, largest=true, sorted=true))] + pub fn topk( + &self, + k: usize, + dim: Option, + largest: Option, + sorted: Option, + ) -> PyResult<(Self, Self)> { + let largest = largest.unwrap_or(true); + let sorted = sorted.unwrap_or(true); + match self.inner.topk(k, dim, largest, sorted) { + Ok((values, indices)) => Ok((Self::from_tensor(values), Self::from_tensor(indices))), + Err(err @ MinitensorError::InvalidArgument { .. }) => { + Err(PyRuntimeError::new_err(err.detailed_message())) + } + Err(err) => Err(_convert_error(err)), + } + } + + #[pyo3(signature = (dim=None, descending=false, stable=false))] + pub fn sort( + &self, + dim: Option, + descending: Option, + stable: Option, + ) -> PyResult<(Self, Self)> { + let descending = descending.unwrap_or(false); + let stable = stable.unwrap_or(false); + match self.inner.sort(dim, descending, stable) { + Ok((values, indices)) => Ok((Self::from_tensor(values), Self::from_tensor(indices))), + Err(err @ MinitensorError::InvalidArgument { .. }) => { + Err(PyRuntimeError::new_err(err.detailed_message())) + } + Err(err) => Err(_convert_error(err)), + } + } + + #[pyo3(signature = (dim=None, descending=false, stable=false))] + pub fn argsort( + &self, + dim: Option, + descending: Option, + stable: Option, + ) -> PyResult { + let descending = descending.unwrap_or(false); + let stable = stable.unwrap_or(false); + match self.inner.argsort(dim, descending, stable) { + Ok(indices) => Ok(Self::from_tensor(indices)), + Err(err @ MinitensorError::InvalidArgument { .. }) => { + Err(PyRuntimeError::new_err(err.detailed_message())) + } + Err(err) => Err(_convert_error(err)), + } + } + + #[pyo3(signature = (dim=None, unbiased=true, keepdim=false))] + pub fn std( + &self, + dim: Option<&Bound>, + unbiased: Option, + keepdim: Option, + ) -> PyResult { + let keepdim = keepdim.unwrap_or(false); + let unbiased = unbiased.unwrap_or(true); + let dims = normalize_optional_axes(dim)?; + let result = self + .inner + .std(dims, keepdim, unbiased) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } + + #[pyo3(signature = (dim=None, unbiased=true, keepdim=false))] + pub fn var( + &self, + dim: Option<&Bound>, + unbiased: Option, + keepdim: Option, + ) -> PyResult { + let keepdim = keepdim.unwrap_or(false); + let unbiased = unbiased.unwrap_or(true); + let dims = normalize_optional_axes(dim)?; + let result = self + .inner + .var(dims, keepdim, unbiased) + .map_err(_convert_error)?; + Ok(Self::from_tensor(result)) + } +} diff --git a/bindings/src/tensor/pytensor/repr.rs b/bindings/src/tensor/pytensor/repr.rs index 23b562d6..f658345a 100644 --- a/bindings/src/tensor/pytensor/repr.rs +++ b/bindings/src/tensor/pytensor/repr.rs @@ -1,105 +1,104 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -#[pymethods] -impl PyTensor { - // String representations - fn __repr__(&self) -> String { - format!( - "Tensor(shape={:?}, dtype={}, device={}, requires_grad={})", - self.inner.shape().dims(), - self.dtype(), - self.device(), - self.inner.requires_grad() - ) - } - - fn __str__(&self) -> String { - if self.inner.numel() <= 100 { - match self.tolist() { - Ok(data) => Python::attach(|py| format!("tensor({})", data.bind(py))), - Err(_) => self.__repr__(), - } - } else { - self.__repr__() - } - } - - fn __len__(&self) -> PyResult { - if self.inner.ndim() == 0 { - Err(PyErr::new::( - "len() of unsized object", - )) - } else { - Ok(self.inner.shape().dims()[0]) - } - } - - fn __bool__(&self) -> PyResult { - if self.inner.numel() != 1 { - return Err(PyErr::new::( - "The truth value of a tensor with more than one element is ambiguous", - )); - } - - match self.inner.dtype() { - DataType::Float32 => { - let data = self.inner.data().as_f32_slice().ok_or_else(|| { - PyErr::new::("Failed to get f32 data") - })?; - Ok(data[0] != 0.0) - } - DataType::Float64 => { - let data = self.inner.data().as_f64_slice().ok_or_else(|| { - PyErr::new::("Failed to get f64 data") - })?; - Ok(data[0] != 0.0) - } - DataType::Int32 => { - let data = self.inner.data().as_i32_slice().ok_or_else(|| { - PyErr::new::("Failed to get i32 data") - })?; - Ok(data[0] != 0) - } - DataType::Int64 => { - let data = self.inner.data().as_i64_slice().ok_or_else(|| { - PyErr::new::("Failed to get i64 data") - })?; - Ok(data[0] != 0) - } - DataType::Bool => { - let data = self.inner.data().as_bool_slice().ok_or_else(|| { - PyErr::new::("Failed to get bool data") - })?; - Ok(data[0]) - } - } - } - - fn __getitem__(&self, key: &Bound) -> PyResult { - let (indices, newaxis_positions) = - parse_getitem_indices(key, self.inner.shape().dims())?; - let mut result = self.inner.index(&indices).map_err(_convert_error)?; - for &pos in &newaxis_positions { - result = result.unsqueeze(pos as isize).map_err(_convert_error)?; - } - Ok(Self::from_tensor(result)) - } - - fn __setitem__(&mut self, key: &Bound, value: &Bound) -> PyResult<()> { - let indices = parse_indices(key, self.inner.shape().dims())?; - let val_tensor = if let Ok(t) = value.extract::() { - t.inner - } else { - convert_python_data_to_tensor(value, self.inner.dtype(), self.inner.device(), false)? - }; - self.inner - .index_assign(&indices, &val_tensor) - .map_err(_convert_error)?; - Ok(()) - } - -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +#[pymethods] +impl PyTensor { + // String representations + fn __repr__(&self) -> String { + format!( + "Tensor(shape={:?}, dtype={}, device={}, requires_grad={})", + self.inner.shape().dims(), + self.dtype(), + self.device(), + self.inner.requires_grad() + ) + } + + fn __str__(&self) -> String { + if self.inner.numel() <= 100 { + match self.tolist() { + Ok(data) => Python::attach(|py| format!("tensor({})", data.bind(py))), + Err(_) => self.__repr__(), + } + } else { + self.__repr__() + } + } + + fn __len__(&self) -> PyResult { + if self.inner.ndim() == 0 { + Err(PyErr::new::( + "len() of unsized object", + )) + } else { + Ok(self.inner.shape().dims()[0]) + } + } + + fn __bool__(&self) -> PyResult { + if self.inner.numel() != 1 { + return Err(PyErr::new::( + "The truth value of a tensor with more than one element is ambiguous", + )); + } + + match self.inner.dtype() { + DataType::Float32 => { + let data = self.inner.data().as_f32_slice().ok_or_else(|| { + PyErr::new::("Failed to get f32 data") + })?; + Ok(data[0] != 0.0) + } + DataType::Float64 => { + let data = self.inner.data().as_f64_slice().ok_or_else(|| { + PyErr::new::("Failed to get f64 data") + })?; + Ok(data[0] != 0.0) + } + DataType::Int32 => { + let data = self.inner.data().as_i32_slice().ok_or_else(|| { + PyErr::new::("Failed to get i32 data") + })?; + Ok(data[0] != 0) + } + DataType::Int64 => { + let data = self.inner.data().as_i64_slice().ok_or_else(|| { + PyErr::new::("Failed to get i64 data") + })?; + Ok(data[0] != 0) + } + DataType::Bool => { + let data = self.inner.data().as_bool_slice().ok_or_else(|| { + PyErr::new::("Failed to get bool data") + })?; + Ok(data[0]) + } + } + } + + fn __getitem__(&self, key: &Bound) -> PyResult { + let (indices, newaxis_positions) = parse_getitem_indices(key, self.inner.shape().dims())?; + let mut result = self.inner.index(&indices).map_err(_convert_error)?; + for &pos in &newaxis_positions { + result = result.unsqueeze(pos as isize).map_err(_convert_error)?; + } + Ok(Self::from_tensor(result)) + } + + fn __setitem__(&mut self, key: &Bound, value: &Bound) -> PyResult<()> { + let indices = parse_indices(key, self.inner.shape().dims())?; + let val_tensor = if let Ok(t) = value.extract::() { + t.inner + } else { + convert_python_data_to_tensor(value, self.inner.dtype(), self.inner.device(), false)? + }; + self.inner + .index_assign(&indices, &val_tensor) + .map_err(_convert_error)?; + Ok(()) + } +} diff --git a/bindings/src/tensor/python/args.rs b/bindings/src/tensor/python/args.rs index 4eb97e58..2e0f7f3b 100644 --- a/bindings/src/tensor/python/args.rs +++ b/bindings/src/tensor/python/args.rs @@ -1,268 +1,275 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -fn convert_dimension(value: isize, arg_name: &str) -> PyResult { - if value < 0 { - return Err(PyValueError::new_err(format!( - "{arg_name} must contain non-negative integers", - ))); - } - - usize::try_from(value).map_err(|_| { - PyValueError::new_err(format!("{arg_name} value is too large for this platform",)) - }) -} - -fn convert_dimensions(values: Vec, arg_name: &str) -> PyResult> { - let mut dims = Vec::with_capacity(values.len()); - for value in values { - dims.push(convert_dimension(value, arg_name)?); - } - Ok(dims) -} - -fn normalize_variadic_isize_args(tuple: &Bound, arg_name: &str) -> PyResult> { - if tuple.is_empty() { - return Ok(Vec::new()); - } - - if tuple.len() == 1 { - let first = tuple.get_item(0)?; - - if let Ok(nested) = first.cast::() { - return normalize_variadic_isize_args(nested, arg_name); - } - - if let Ok(list) = first.cast::() { - let mut dims = Vec::with_capacity(list.len()); - for item in list.iter() { - dims.push(item.extract::()?); - } - return Ok(dims); - } - - if let Ok(shape_sequence) = first.extract::() { - return convert_usize_list_to_isize(shape_sequence.to_list(), arg_name); - } - - if let Ok(values) = first.extract::>() { - return Ok(values); - } - - if let Ok(values) = first.extract::>() { - return convert_usize_list_to_isize(values, arg_name); - } - - if let Ok(value) = first.extract::() { - return Ok(vec![value]); - } - - if let Ok(value) = first.extract::() { - return Ok(vec![convert_usize_to_isize(value, arg_name)?]); - } - } - - let mut dims = Vec::with_capacity(tuple.len()); - for item in tuple.iter() { - dims.push(item.extract::()?); - } - Ok(dims) -} - -fn convert_usize_list_to_isize(values: Vec, arg_name: &str) -> PyResult> { - let mut converted = Vec::with_capacity(values.len()); - for value in values { - converted.push(convert_usize_to_isize(value, arg_name)?); - } - Ok(converted) -} - -fn convert_usize_to_isize(value: usize, arg_name: &str) -> PyResult { - isize::try_from(value).map_err(|_| { - PyValueError::new_err(format!( - "{arg_name} dimension {value} is too large for this platform" - )) - }) -} - -fn parse_shape_tuple(shape: &Bound, arg_name: &str) -> PyResult> { - if shape.is_empty() { - return Ok(Vec::new()); - } - - if shape.len() == 1 { - let first = shape.get_item(0)?; - if let Ok(tuple) = first.cast::() { - return parse_shape_tuple(tuple, arg_name); - } - if let Ok(list) = first.cast::() { - let mut dims = Vec::with_capacity(list.len()); - for item in list.iter() { - let value: isize = item.extract()?; - dims.push(convert_dimension(value, arg_name)?); - } - return Ok(dims); - } - if let Ok(shape_seq) = first.extract::() { - return Ok(shape_seq.to_list()); - } - if let Ok(values) = first.extract::>() { - return convert_dimensions(values, arg_name); - } - if let Ok(value) = first.extract::() { - return Ok(vec![convert_dimension(value, arg_name)?]); - } - } - - let mut dims = Vec::with_capacity(shape.len()); - for item in shape.iter() { - let value: isize = item.extract()?; - dims.push(convert_dimension(value, arg_name)?); - } - Ok(dims) -} - -fn parse_shape_like(obj: &Bound, arg_name: &str) -> PyResult> { - if let Ok(tuple) = obj.cast::() { - return parse_shape_tuple(tuple, arg_name); - } - - if let Ok(list) = obj.cast::() { - let mut dims = Vec::with_capacity(list.len()); - for item in list.iter() { - let value: isize = item.extract()?; - dims.push(convert_dimension(value, arg_name)?); - } - return Ok(dims); - } - - if let Ok(shape_seq) = obj.extract::() { - return Ok(shape_seq.to_list()); - } - - if let Ok(values) = obj.extract::>() { - return convert_dimensions(values, arg_name); - } - - if let Ok(value) = obj.extract::() { - return Ok(vec![convert_dimension(value, arg_name)?]); - } - - Err(PyTypeError::new_err(format!( - "{arg_name} must be an int or sequence of ints", - ))) -} - -fn normalize_roll_shifts(shifts: &Bound) -> PyResult> { - normalize_required_axes(shifts, "shifts") -} - -fn normalize_required_axes<'py>(dim: &'py Bound<'py, PyAny>, name: &str) -> PyResult> { - match normalize_optional_axes(Some(dim))? { - Some(values) => Ok(values), - None => Err(PyTypeError::new_err(format!( - "{} must be an int or a sequence of ints", - name - ))), - } -} - -fn normalize_optional_axes(dim: Option<&Bound>) -> PyResult>> { - let Some(obj) = dim else { - return Ok(None); - }; - - if obj.is_none() { - return Ok(None); - } - - if is_bool_axis(obj)? { - return Err(PyTypeError::new_err( - "dim must be an int or a sequence of ints", - )); - } - - if let Ok(value) = obj.extract::() { - return Ok(Some(vec![value])); - } - - if obj.is_instance_of::() { - return Err(PyTypeError::new_err( - "dim must be an int or a sequence of ints", - )); - } - - if let Ok(sequence) = obj.cast::() { - let length = sequence.len()?; - let mut axes = Vec::with_capacity(length); - for index in 0..length { - let item = sequence.get_item(index)?; - if is_bool_axis(&item)? { - return Err(PyTypeError::new_err( - "dim must be an int or a sequence of ints", - )); - } - let value: isize = item.extract()?; - axes.push(value); - } - return Ok(Some(axes)); - } - - Err(PyTypeError::new_err( - "dim must be an int or a sequence of ints", - )) -} - -fn is_bool_axis(obj: &Bound) -> PyResult { - if obj.is_instance_of::() { - return Ok(true); - } - - static NUMPY_BOOL_TYPE: OnceCell> = OnceCell::new(); - let py = obj.py(); - if let Ok(numpy_bool) = NUMPY_BOOL_TYPE.get_or_try_init(|| -> PyResult> { - let numpy = PyModule::import(py, "numpy")?; - let bool_obj = numpy.getattr("bool_")?; - Ok(bool_obj.unbind()) - }) && obj.is_instance(numpy_bool.bind(py))? - { - return Ok(true); - } - - Ok(false) -} - -fn normalize_repeat_spec(repeats: &Bound) -> PyResult> { - if repeats.is_instance_of::() { - return Ok(vec![extract_repeat_element(repeats)?]); - } - - if let Ok(sequence) = repeats.extract::>() { - let mut values = Vec::with_capacity(sequence.len()); - for repeat in sequence { - if repeat < 0 { - return Err(PyValueError::new_err( - "repeat expects non-negative integers", - )); - } - values.push(repeat as usize); - } - return Ok(values); - } - - Ok(vec![extract_repeat_element(repeats)?]) -} - -fn extract_repeat_element(value: &Bound) -> PyResult { - let repeat: i64 = value.extract()?; - if repeat < 0 { - Err(PyValueError::new_err( - "repeat expects non-negative integers", - )) - } else { - Ok(repeat as usize) - } -} +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +fn convert_dimension(value: isize, arg_name: &str) -> PyResult { + if value < 0 { + return Err(PyValueError::new_err(format!( + "{arg_name} must contain non-negative integers", + ))); + } + + usize::try_from(value).map_err(|_| { + PyValueError::new_err(format!("{arg_name} value is too large for this platform",)) + }) +} + +fn convert_dimensions(values: Vec, arg_name: &str) -> PyResult> { + let mut dims = Vec::with_capacity(values.len()); + for value in values { + dims.push(convert_dimension(value, arg_name)?); + } + Ok(dims) +} + +pub(crate) fn normalize_variadic_isize_args( + tuple: &Bound, + arg_name: &str, +) -> PyResult> { + if tuple.is_empty() { + return Ok(Vec::new()); + } + + if tuple.len() == 1 { + let first = tuple.get_item(0)?; + + if let Ok(nested) = first.cast::() { + return normalize_variadic_isize_args(nested, arg_name); + } + + if let Ok(list) = first.cast::() { + let mut dims = Vec::with_capacity(list.len()); + for item in list.iter() { + dims.push(item.extract::()?); + } + return Ok(dims); + } + + if let Ok(shape_sequence) = first.extract::() { + return convert_usize_list_to_isize(shape_sequence.to_list(), arg_name); + } + + if let Ok(values) = first.extract::>() { + return Ok(values); + } + + if let Ok(values) = first.extract::>() { + return convert_usize_list_to_isize(values, arg_name); + } + + if let Ok(value) = first.extract::() { + return Ok(vec![value]); + } + + if let Ok(value) = first.extract::() { + return Ok(vec![convert_usize_to_isize(value, arg_name)?]); + } + } + + let mut dims = Vec::with_capacity(tuple.len()); + for item in tuple.iter() { + dims.push(item.extract::()?); + } + Ok(dims) +} + +fn convert_usize_list_to_isize(values: Vec, arg_name: &str) -> PyResult> { + let mut converted = Vec::with_capacity(values.len()); + for value in values { + converted.push(convert_usize_to_isize(value, arg_name)?); + } + Ok(converted) +} + +fn convert_usize_to_isize(value: usize, arg_name: &str) -> PyResult { + isize::try_from(value).map_err(|_| { + PyValueError::new_err(format!( + "{arg_name} dimension {value} is too large for this platform" + )) + }) +} + +pub(crate) fn parse_shape_tuple(shape: &Bound, arg_name: &str) -> PyResult> { + if shape.is_empty() { + return Ok(Vec::new()); + } + + if shape.len() == 1 { + let first = shape.get_item(0)?; + if let Ok(tuple) = first.cast::() { + return parse_shape_tuple(tuple, arg_name); + } + if let Ok(list) = first.cast::() { + let mut dims = Vec::with_capacity(list.len()); + for item in list.iter() { + let value: isize = item.extract()?; + dims.push(convert_dimension(value, arg_name)?); + } + return Ok(dims); + } + if let Ok(shape_seq) = first.extract::() { + return Ok(shape_seq.to_list()); + } + if let Ok(values) = first.extract::>() { + return convert_dimensions(values, arg_name); + } + if let Ok(value) = first.extract::() { + return Ok(vec![convert_dimension(value, arg_name)?]); + } + } + + let mut dims = Vec::with_capacity(shape.len()); + for item in shape.iter() { + let value: isize = item.extract()?; + dims.push(convert_dimension(value, arg_name)?); + } + Ok(dims) +} + +pub(crate) fn parse_shape_like(obj: &Bound, arg_name: &str) -> PyResult> { + if let Ok(tuple) = obj.cast::() { + return parse_shape_tuple(tuple, arg_name); + } + + if let Ok(list) = obj.cast::() { + let mut dims = Vec::with_capacity(list.len()); + for item in list.iter() { + let value: isize = item.extract()?; + dims.push(convert_dimension(value, arg_name)?); + } + return Ok(dims); + } + + if let Ok(shape_seq) = obj.extract::() { + return Ok(shape_seq.to_list()); + } + + if let Ok(values) = obj.extract::>() { + return convert_dimensions(values, arg_name); + } + + if let Ok(value) = obj.extract::() { + return Ok(vec![convert_dimension(value, arg_name)?]); + } + + Err(PyTypeError::new_err(format!( + "{arg_name} must be an int or sequence of ints", + ))) +} + +pub(crate) fn normalize_roll_shifts(shifts: &Bound) -> PyResult> { + normalize_required_axes(shifts, "shifts") +} + +pub(crate) fn normalize_required_axes<'py>( + dim: &'py Bound<'py, PyAny>, + name: &str, +) -> PyResult> { + match normalize_optional_axes(Some(dim))? { + Some(values) => Ok(values), + None => Err(PyTypeError::new_err(format!( + "{} must be an int or a sequence of ints", + name + ))), + } +} + +pub(crate) fn normalize_optional_axes(dim: Option<&Bound>) -> PyResult>> { + let Some(obj) = dim else { + return Ok(None); + }; + + if obj.is_none() { + return Ok(None); + } + + if is_bool_axis(obj)? { + return Err(PyTypeError::new_err( + "dim must be an int or a sequence of ints", + )); + } + + if let Ok(value) = obj.extract::() { + return Ok(Some(vec![value])); + } + + if obj.is_instance_of::() { + return Err(PyTypeError::new_err( + "dim must be an int or a sequence of ints", + )); + } + + if let Ok(sequence) = obj.cast::() { + let length = sequence.len()?; + let mut axes = Vec::with_capacity(length); + for index in 0..length { + let item = sequence.get_item(index)?; + if is_bool_axis(&item)? { + return Err(PyTypeError::new_err( + "dim must be an int or a sequence of ints", + )); + } + let value: isize = item.extract()?; + axes.push(value); + } + return Ok(Some(axes)); + } + + Err(PyTypeError::new_err( + "dim must be an int or a sequence of ints", + )) +} + +fn is_bool_axis(obj: &Bound) -> PyResult { + if obj.is_instance_of::() { + return Ok(true); + } + + static NUMPY_BOOL_TYPE: OnceCell> = OnceCell::new(); + let py = obj.py(); + if let Ok(numpy_bool) = NUMPY_BOOL_TYPE.get_or_try_init(|| -> PyResult> { + let numpy = PyModule::import(py, "numpy")?; + let bool_obj = numpy.getattr("bool_")?; + Ok(bool_obj.unbind()) + }) && obj.is_instance(numpy_bool.bind(py))? + { + return Ok(true); + } + + Ok(false) +} + +pub(crate) fn normalize_repeat_spec(repeats: &Bound) -> PyResult> { + if repeats.is_instance_of::() { + return Ok(vec![extract_repeat_element(repeats)?]); + } + + if let Ok(sequence) = repeats.extract::>() { + let mut values = Vec::with_capacity(sequence.len()); + for repeat in sequence { + if repeat < 0 { + return Err(PyValueError::new_err( + "repeat expects non-negative integers", + )); + } + values.push(repeat as usize); + } + return Ok(values); + } + + Ok(vec![extract_repeat_element(repeats)?]) +} + +fn extract_repeat_element(value: &Bound) -> PyResult { + let repeat: i64 = value.extract()?; + if repeat < 0 { + Err(PyValueError::new_err( + "repeat expects non-negative integers", + )) + } else { + Ok(repeat as usize) + } +} diff --git a/bindings/src/tensor/python/convert.rs b/bindings/src/tensor/python/convert.rs index b13aed03..6b8e7cb9 100644 --- a/bindings/src/tensor/python/convert.rs +++ b/bindings/src/tensor/python/convert.rs @@ -1,814 +1,830 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -fn convert_python_data_to_tensor( - data: &Bound, - dtype: DataType, - device: Device, - requires_grad: bool, -) -> PyResult { - // First try NumPy array conversion for any supported dtype - if let Ok(numpy_module) = PyModule::import(data.py(), "numpy") - && let Ok(ndarray_type) = numpy_module.getattr("ndarray") - && data.is_instance(&ndarray_type)? - { - let maybe_tensor = panic::catch_unwind(AssertUnwindSafe(|| { - convert_numpy_to_tensor(data, requires_grad) - })); - - match maybe_tensor { - Ok(Ok(tensor)) => { - let tensor = if tensor.dtype() != dtype { - tensor.astype(dtype).map_err(_convert_error)? - } else { - tensor - }; - return Ok(tensor); - } - Ok(Err(err)) => { - return Err(err); - } - Err(_) => { - // Fall back to the slower Python list conversion path - // when the NumPy capsule isn't available. - } - } - } - - // Handle Python lists and tuples by flattening values into scalar variants - if let Ok(list) = data.cast::() { - let (shape, flat_data) = flatten_python_data(list)?; - let (base_tensor, base_dtype) = - tensor_from_flat_scalars(shape, flat_data, device, requires_grad)?; - - if base_dtype == dtype { - return Ok(base_tensor); - } - - return base_tensor.astype(dtype).map_err(_convert_error); - } - - if let Ok(tuple) = data.cast::() { - let list = tuple.to_list(); - return convert_python_data_to_tensor(list.as_any(), dtype, device, requires_grad); - } - - // Handle scalars - if let Ok(value_bool) = data.extract::() { - let shape = Shape::new(vec![]); - let base_data = Arc::new(TensorData::from_vec_bool(vec![value_bool], device)); - let mut tensor = Tensor::new(base_data, shape, DataType::Bool, device, requires_grad); - if dtype != DataType::Bool { - tensor = tensor.astype(dtype).map_err(_convert_error)?; - } - return Ok(tensor); - } - - if let Ok(value_int) = data.extract::() { - let shape = Shape::new(vec![]); - let base_data = Arc::new(TensorData::from_vec_i64(vec![value_int], device)); - let mut tensor = Tensor::new(base_data, shape, DataType::Int64, device, requires_grad); - if dtype != DataType::Int64 { - tensor = tensor.astype(dtype).map_err(_convert_error)?; - } - return Ok(tensor); - } - - if let Ok(value_float) = data.extract::() { - let shape = Shape::new(vec![]); - let base_data = Arc::new(TensorData::from_vec_f64(vec![value_float], device)); - let mut tensor = Tensor::new(base_data, shape, DataType::Float64, device, requires_grad); - if dtype != DataType::Float64 { - tensor = tensor.astype(dtype).map_err(_convert_error)?; - } - return Ok(tensor); - } - - let float_name = intern!(data.py(), "__float__"); - if data.hasattr(float_name)? { - let method = data.getattr(float_name)?; - if method.is_callable() { - let float_obj = method.call0()?; - let val = float_obj.extract::()?; - let shape = Shape::new(vec![]); - let base_data = Arc::new(TensorData::from_vec_f64(vec![val], device)); - let mut tensor = - Tensor::new(base_data, shape, DataType::Float64, device, requires_grad); - if dtype != DataType::Float64 { - tensor = tensor.astype(dtype).map_err(_convert_error)?; - } - return Ok(tensor); - } - } - - Err(PyErr::new::( - "Unsupported data type for tensor creation", - )) -} - -fn apply_binary_ufunc(operands: &[Tensor], kind: BinaryOpKind, op: F) -> PyResult -where - F: Fn(&Tensor, &Tensor) -> Result, -{ - if operands.len() != 2 { - return Err(PyValueError::new_err( - "Binary ufuncs require exactly two operands", - )); - } - - let (lhs_cast, rhs_cast, _) = - coerce_binary_operands(&operands[0], &operands[1], kind).map_err(_convert_error)?; - - let lhs_tensor = match lhs_cast { - Cow::Borrowed(tensor) => tensor.clone(), - Cow::Owned(tensor) => tensor, - }; - let rhs_tensor = match rhs_cast { - Cow::Borrowed(tensor) => tensor.clone(), - Cow::Owned(tensor) => tensor, - }; - - op(&lhs_tensor, &rhs_tensor).map_err(_convert_error) -} - -fn apply_unary_ufunc(operands: &[Tensor], op: F) -> PyResult -where - F: Fn(&Tensor) -> Result, -{ - if operands.len() != 1 { - return Err(PyValueError::new_err( - "Unary ufuncs require exactly one operand", - )); - } - - let tensor = operands[0].clone(); - op(&tensor).map_err(_convert_error) -} - -fn py_not_implemented(py: Python) -> PyResult> { - unsafe { Ok(pyo3::Bound::::from_borrowed_ptr(py, pyo3::ffi::Py_NotImplemented()).unbind()) } -} - -fn parse_dtype_like(value: &Bound) -> PyResult { - if let Ok(name) = value.extract::() { - dtype::parse_dtype(&name) - } else { - Err(PyTypeError::new_err( - "dtype must be specified as a string such as 'float32'", - )) - } -} - -fn parse_device_like(value: &Bound) -> PyResult { - if let Ok(device) = value.extract::() { - return Ok(device.device()); - } - - if let Ok(spec) = value.extract::() { - return Device::from_str(&spec).map_err(|err| { - PyValueError::new_err(format!("Unsupported device specification '{spec}': {err}")) - }); - } - - Err(PyTypeError::new_err( - "device must be specified as a Device object or string like 'cpu' or 'cuda:0'", - )) -} - -fn ensure_backward_gradient_compatible(reference: &Tensor, gradient: &mut Tensor) -> PyResult<()> { - let expected_shape = reference.shape().dims(); - let actual_shape = gradient.shape().dims(); - if expected_shape != actual_shape { - return Err(PyRuntimeError::new_err(format!( - "backward() expected gradient tensor with shape {:?}, but got {:?}", - expected_shape, actual_shape - ))); - } - - if gradient.device() != reference.device() { - *gradient = gradient.to(reference.device()).map_err(_convert_error)?; - } - - if gradient.dtype() != reference.dtype() { - *gradient = gradient.astype(reference.dtype()).map_err(_convert_error)?; - } - - if gradient.requires_grad() { - *gradient = gradient.detach(); - } - - Ok(()) -} - -fn tensor_from_py_value(reference: &Tensor, value: &Bound) -> PyResult { - if let Some(py_tensor) = extract_wrapped_pytensor(value) { - return Ok(py_tensor.inner.clone()); - } - - if let Ok(numpy_module) = PyModule::import(value.py(), "numpy") - && let Ok(ndarray_type) = numpy_module.getattr("ndarray") - && value.is_instance(&ndarray_type)? - { - if let Ok(dtype_obj) = value.getattr("dtype") { - let dtype_str = dtype_obj.str()?.to_str()?.to_ascii_lowercase(); - if let Ok(array_dtype) = dtype::parse_dtype(&dtype_str) { - return convert_python_data_to_tensor( - value, - array_dtype, - reference.device(), - false, - ); - } - } - return convert_python_data_to_tensor(value, reference.dtype(), reference.device(), false); - } - - if let Ok(py_tensor) = PyTensor::from_python_value(value) { - let mut tensor = py_tensor.inner; - if tensor.device() != reference.device() { - tensor = tensor.to(reference.device()).map_err(_convert_error)?; - } - - let target_dtype = dtype::resolve_scalar_dtype(value, reference.dtype()) - .ok() - .or_else(|| infer_python_value_dtype(value)) - .unwrap_or(reference.dtype()); - - if tensor.dtype() != target_dtype { - tensor = tensor.astype(target_dtype).map_err(_convert_error)?; - } - - return Ok(tensor); - } - - let index_name = intern!(value.py(), "__index__"); - if value.hasattr(index_name)? { - let method = value.getattr(index_name)?; - if method.is_callable() { - let result = method.call0()?; - if result.is_instance_of::() { - let dtype = match dtype::resolve_scalar_dtype(value, reference.dtype()) { - Ok(dt) => dt, - Err(_) => reference.dtype(), - }; - return convert_python_data_to_tensor( - result.as_any(), - dtype, - reference.device(), - false, - ); - } - } - } - - let dtype = match dtype::resolve_scalar_dtype(value, reference.dtype()) { - Ok(dt) => dt, - Err(_) => infer_python_value_dtype(value).unwrap_or(reference.dtype()), - }; - convert_python_data_to_tensor(value, dtype, reference.device(), false) -} - -fn tensor_bool_from_py(value: &Bound, device: Device) -> PyResult { - if let Some(py_tensor) = extract_wrapped_pytensor(value) { - let mut tensor = py_tensor.inner.clone(); - if tensor.dtype() != DataType::Bool { - return Err(PyTypeError::new_err("mask must be a bool tensor")); - } - if tensor.device() != device { - tensor = tensor.to(device).map_err(_convert_error)?; - } - return Ok(tensor); - } - - if let Ok(value_bool) = value.extract::() { - let data = Arc::new(TensorData::from_vec_bool(vec![value_bool], device)); - return Ok(Tensor::new( - data, - Shape::new(vec![]), - DataType::Bool, - device, - false, - )); - } - - convert_python_data_to_tensor(value, DataType::Bool, device, false) -} - -fn promote_dtypes(a: DataType, b: DataType) -> DataType { - use DataType::*; - - if a == b { - return a; - } - - match (a, b) { - (Float64, _) | (_, Float64) => Float64, - (Float32, _) | (_, Float32) => Float32, - (Int64, _) | (_, Int64) => Int64, - (Int32, _) | (_, Int32) => Int32, - _ => Bool, - } -} - -fn infer_python_value_dtype(value: &Bound) -> Option { - if let Some(py_tensor) = extract_wrapped_pytensor(value) { - return Some(py_tensor.inner.dtype()); - } - - if value.extract::().is_ok() { - return Some(DataType::Bool); - } - - if value.extract::().is_ok() { - return Some(DataType::Int64); - } - - if value.extract::().is_ok() { - return Some(dtype::default_dtype()); - } - - if let Ok(numpy_module) = PyModule::import(value.py(), "numpy") - && let Ok(ndarray_type) = numpy_module.getattr("ndarray") - && let Ok(true) = value.is_instance(&ndarray_type) - && let Ok(dtype_obj) = value.getattr("dtype") - && let Ok(dtype_str) = dtype_obj.str() - && let Ok(dtype) = dtype::parse_dtype(&dtype_str.to_str().ok()?.to_ascii_lowercase()) - { - return Some(dtype); - } - - if let Ok(list) = value.cast::() { - return infer_sequence_dtype(list.iter()); - } - - if let Ok(tuple) = value.cast::() { - return infer_sequence_dtype(tuple.iter()); - } - - None -} - -fn infer_sequence_dtype<'py, I>(iter: I) -> Option -where - I: Iterator>, -{ - let mut dtype: Option = None; - for item in iter { - let item_dtype = infer_python_value_dtype(&item)?; - dtype = Some(match dtype { - Some(current) => promote_dtypes(current, item_dtype), - None => item_dtype, - }); - } - dtype -} - -fn prepare_binary_operands_from_py( - reference: &Tensor, - other: &Bound, - reverse: bool, - kind: BinaryOpKind, -) -> PyResult<(Tensor, Tensor)> { - let lhs_input = if reverse { - tensor_from_py_value(reference, other)? - } else { - reference.clone() - }; - - let rhs_input = if reverse { - reference.clone() - } else { - tensor_from_py_value(reference, other)? - }; - - let (lhs_cast, rhs_cast, _) = - coerce_binary_operands(&lhs_input, &rhs_input, kind).map_err(_convert_error)?; - let lhs_tensor = match lhs_cast { - Cow::Borrowed(_) => lhs_input.clone(), - Cow::Owned(tensor) => tensor, - }; - let rhs_tensor = match rhs_cast { - Cow::Borrowed(_) => rhs_input.clone(), - Cow::Owned(tensor) => tensor, - }; - - Ok((lhs_tensor, rhs_tensor)) -} - -fn flatten_python_data(list: &Bound) -> PyResult<(Vec, Vec)> { - let mut shape = vec![list.len()]; - let mut flat_data = vec![]; - - fn process_nested( - item: &Bound, - depth: usize, - shape: &mut Vec, - flat_data: &mut Vec, - ) -> PyResult<()> { - if let Ok(nested_list) = item.cast::() { - let length = nested_list.len(); - if depth >= shape.len() { - shape.push(length); - } else if shape[depth] != length { - return Err(PyErr::new::( - "Inconsistent nested sequence lengths", - )); - } - for nested_item in nested_list.iter() { - process_nested(&nested_item, depth + 1, shape, flat_data)?; - } - return Ok(()); - } - - if let Ok(nested_tuple) = item.cast::() { - let list = nested_tuple.to_list(); - let length = list.len(); - if depth >= shape.len() { - shape.push(length); - } else if shape[depth] != length { - return Err(PyErr::new::( - "Inconsistent nested sequence lengths", - )); - } - for nested_item in list.iter() { - process_nested(&nested_item, depth + 1, shape, flat_data)?; - } - return Ok(()); - } - - if let Ok(value_bool) = item.extract::() { - flat_data.push(ScalarValue::Bool(value_bool)); - return Ok(()); - } - - if let Ok(value_int) = item.extract::() { - flat_data.push(ScalarValue::Int(value_int)); - return Ok(()); - } - - let index_name = intern!(item.py(), "__index__"); - if item.hasattr(index_name)? { - let method = item.getattr(index_name)?; - if method.is_callable() { - let result = method.call0()?; - if result.is_instance_of::() { - let value = result.extract::()?; - flat_data.push(ScalarValue::Int(value)); - return Ok(()); - } - } - } - - if let Ok(value_float) = item.extract::() { - flat_data.push(ScalarValue::Float(value_float)); - return Ok(()); - } - - let float_name = intern!(item.py(), "__float__"); - if item.hasattr(float_name)? { - let method = item.getattr(float_name)?; - if method.is_callable() { - let float_obj = method.call0()?; - let value = float_obj.extract::()?; - flat_data.push(ScalarValue::Float(value)); - return Ok(()); - } - } - - Err(PyErr::new::( - "Unsupported scalar type in nested sequence", - )) - } - - for item in list.iter() { - process_nested(&item, 1, &mut shape, &mut flat_data)?; - } - - Ok((shape, flat_data)) -} - -#[derive(Clone, Copy)] -enum ScalarValue { - Bool(bool), - Int(i64), - Float(f64), -} - -impl ScalarValue { - fn kind(&self) -> ScalarKind { - match self { - ScalarValue::Bool(_) => ScalarKind::Bool, - ScalarValue::Int(_) => ScalarKind::Int, - ScalarValue::Float(_) => ScalarKind::Float, - } - } - - fn to_bool(self) -> bool { - match self { - ScalarValue::Bool(value) => value, - ScalarValue::Int(value) => value != 0, - ScalarValue::Float(value) => value != 0.0, - } - } - - fn to_i64(self) -> i64 { - match self { - ScalarValue::Bool(value) => value as i64, - ScalarValue::Int(value) => value, - ScalarValue::Float(value) => value as i64, - } - } - - fn to_f64(self) -> f64 { - match self { - ScalarValue::Bool(value) => { - if value { - 1.0 - } else { - 0.0 - } - } - ScalarValue::Int(value) => value as f64, - ScalarValue::Float(value) => value, - } - } -} - -#[derive(Clone, Copy, PartialEq, Eq)] -enum ScalarKind { - Bool, - Int, - Float, -} - -impl ScalarKind { - fn combine(self, other: ScalarKind) -> ScalarKind { - use ScalarKind::*; - match (self, other) { - (Float, _) | (_, Float) => Float, - (Int, _) | (_, Int) => Int, - _ => Bool, - } - } -} - -fn tensor_from_flat_scalars( - shape: Vec, - values: Vec, - device: Device, - requires_grad: bool, -) -> PyResult<(Tensor, DataType)> { - let mut kind = ScalarKind::Bool; - for value in &values { - kind = kind.combine(value.kind()); - } - - let tensor = match kind { - ScalarKind::Bool => { - let data: Vec = values.into_iter().map(ScalarValue::to_bool).collect(); - Tensor::new( - Arc::new(TensorData::from_vec_bool(data, device)), - Shape::new(shape), - DataType::Bool, - device, - requires_grad, - ) - } - ScalarKind::Int => { - let data: Vec = values.into_iter().map(ScalarValue::to_i64).collect(); - Tensor::new( - Arc::new(TensorData::from_vec_i64(data, device)), - Shape::new(shape), - DataType::Int64, - device, - requires_grad, - ) - } - ScalarKind::Float => { - let data: Vec = values.into_iter().map(ScalarValue::to_f64).collect(); - Tensor::new( - Arc::new(TensorData::from_vec_f64(data, device)), - Shape::new(shape), - DataType::Float64, - device, - requires_grad, - ) - } - }; - - let dtype = tensor.dtype(); - Ok((tensor, dtype)) -} -fn parse_index(item: &Bound, dim_size: usize) -> PyResult { - if let Ok(i) = item.extract::() { - let mut idx = i; - if idx < 0 { - idx += dim_size as isize; - } - if idx < 0 || idx >= dim_size as isize { - return Err(PyIndexError::new_err("Index out of bounds")); - } - Ok(TensorIndex::Index(idx as usize)) - } else if let Ok(slice) = item.cast::() { - use std::convert::TryInto; - - let dim_size_isize: isize = dim_size - .try_into() - .map_err(|_| PyValueError::new_err("dim_size too large"))?; - let indices = slice.indices(dim_size_isize)?; - if indices.step <= 0 { - return Err(PyIndexError::new_err("slice step must be positive")); - } - Ok(TensorIndex::Slice { - start: indices.start.max(0) as usize, - end: indices.stop.max(0) as usize, - step: indices.step as usize, - }) - } else if item.is_none() { - Ok(TensorIndex::Slice { - start: 0, - end: dim_size, - step: 1, - }) - } else { - Err(PyTypeError::new_err("Invalid index type")) - } -} - -/// Parse a `__getitem__` key into per-dimension indices plus the output-axis -/// positions where a size-1 axis must be inserted for `None`/`np.newaxis`. -/// -/// Unlike [`parse_indices`], `None` does not consume an input dimension; it -/// inserts a new length-1 axis at the corresponding output position. -/// Integer indices drop their dimension, slices keep it. -fn parse_getitem_indices( - key: &Bound, - shape: &[usize], -) -> PyResult<(Vec, Vec)> { - let items: Vec> = if let Ok(tup) = key.cast::() { - tup.iter().collect() - } else { - vec![key.clone()] - }; - - // `None` entries add axes rather than selecting dimensions, so only the - // real (non-None) entries count against the tensor rank. - let real_count = items.iter().filter(|it| !it.is_none()).count(); - if real_count > shape.len() { - return Err(PyIndexError::new_err("Too many indices")); - } - - let mut real_indices: Vec = Vec::with_capacity(shape.len()); - let mut newaxis_positions: Vec = Vec::new(); - let mut input_dim = 0usize; - let mut out_dim = 0usize; - - for item in &items { - if item.is_none() { - newaxis_positions.push(out_dim); - out_dim += 1; - continue; - } - let idx = parse_index(item, shape[input_dim])?; - // Integer indices remove the dimension; slices keep it in the output. - let keeps_dim = matches!(idx, TensorIndex::Slice { .. }); - real_indices.push(idx); - input_dim += 1; - if keeps_dim { - out_dim += 1; - } - } - - // Any dimensions not addressed explicitly are taken in full. - for &dim in &shape[input_dim..] { - real_indices.push(TensorIndex::Slice { - start: 0, - end: dim, - step: 1, - }); - } - - Ok((real_indices, newaxis_positions)) -} - -fn parse_indices(key: &Bound, shape: &[usize]) -> PyResult> { - if let Ok(tup) = key.cast::() { - if tup.len() > shape.len() { - return Err(PyIndexError::new_err("Too many indices")); - } - let mut result = Vec::new(); - for (i, dim) in shape.iter().enumerate() { - if i < tup.len() { - result.push(parse_index(&tup.get_item(i)?, *dim)?); - } else { - result.push(TensorIndex::Slice { - start: 0, - end: *dim, - step: 1, - }); - } - } - Ok(result) - } else { - let mut result = vec![parse_index(key, shape[0])?]; - for dim in &shape[1..] { - result.push(TensorIndex::Slice { - start: 0, - end: *dim, - step: 1, - }); - } - Ok(result) - } -} - -fn convert_numpy_to_tensor(array: &Bound, requires_grad: bool) -> PyResult { - if let Ok(array_f32) = array.cast::>() { - let readonly = array_f32.readonly(); - let shape = Shape::new(readonly.shape().to_vec()); - let data_vec: Vec = readonly.as_slice()?.to_vec(); - let tensor_data = Arc::new(TensorData::from_vec( - data_vec, - DataType::Float32, - Device::cpu(), - )); - Ok(Tensor::new( - tensor_data, - shape, - DataType::Float32, - Device::cpu(), - requires_grad, - )) - } else if let Ok(array_f64) = array.cast::>() { - let readonly = array_f64.readonly(); - let shape = Shape::new(readonly.shape().to_vec()); - let data_vec: Vec = readonly.as_slice()?.to_vec(); - let tensor_data = Arc::new(TensorData::from_vec( - data_vec, - DataType::Float64, - Device::cpu(), - )); - Ok(Tensor::new( - tensor_data, - shape, - DataType::Float64, - Device::cpu(), - requires_grad, - )) - } else if let Ok(array_i32) = array.cast::>() { - let readonly = array_i32.readonly(); - let shape = Shape::new(readonly.shape().to_vec()); - let data_vec: Vec = readonly.as_slice()?.to_vec(); - let tensor_data = Arc::new(TensorData::from_vec( - data_vec, - DataType::Int32, - Device::cpu(), - )); - Ok(Tensor::new( - tensor_data, - shape, - DataType::Int32, - Device::cpu(), - requires_grad, - )) - } else if let Ok(array_i64) = array.cast::>() { - let readonly = array_i64.readonly(); - let shape = Shape::new(readonly.shape().to_vec()); - let data_vec: Vec = readonly.as_slice()?.to_vec(); - let tensor_data = Arc::new(TensorData::from_vec( - data_vec, - DataType::Int64, - Device::cpu(), - )); - Ok(Tensor::new( - tensor_data, - shape, - DataType::Int64, - Device::cpu(), - requires_grad, - )) - } else if let Ok(array_bool) = array.cast::>() { - let readonly = array_bool.readonly(); - let shape = Shape::new(readonly.shape().to_vec()); - let data_vec: Vec = readonly.as_slice()?.to_vec(); - let tensor_data = Arc::new(TensorData::from_vec( - data_vec, - DataType::Bool, - Device::cpu(), - )); - Ok(Tensor::new( - tensor_data, - shape, - DataType::Bool, - Device::cpu(), - requires_grad, - )) - } else { - Err(PyErr::new::( - "Unsupported NumPy array type", - )) - } -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +pub(crate) fn convert_python_data_to_tensor( + data: &Bound, + dtype: DataType, + device: Device, + requires_grad: bool, +) -> PyResult { + // First try NumPy array conversion for any supported dtype + if let Ok(numpy_module) = PyModule::import(data.py(), "numpy") + && let Ok(ndarray_type) = numpy_module.getattr("ndarray") + && data.is_instance(&ndarray_type)? + { + let maybe_tensor = panic::catch_unwind(AssertUnwindSafe(|| { + convert_numpy_to_tensor(data, requires_grad) + })); + + match maybe_tensor { + Ok(Ok(tensor)) => { + let tensor = if tensor.dtype() != dtype { + tensor.astype(dtype).map_err(_convert_error)? + } else { + tensor + }; + return Ok(tensor); + } + Ok(Err(err)) => { + return Err(err); + } + Err(_) => { + // Fall back to the slower Python list conversion path + // when the NumPy capsule isn't available. + } + } + } + + // Handle Python lists and tuples by flattening values into scalar variants + if let Ok(list) = data.cast::() { + let (shape, flat_data) = flatten_python_data(list)?; + let (base_tensor, base_dtype) = + tensor_from_flat_scalars(shape, flat_data, device, requires_grad)?; + + if base_dtype == dtype { + return Ok(base_tensor); + } + + return base_tensor.astype(dtype).map_err(_convert_error); + } + + if let Ok(tuple) = data.cast::() { + let list = tuple.to_list(); + return convert_python_data_to_tensor(list.as_any(), dtype, device, requires_grad); + } + + // Handle scalars + if let Ok(value_bool) = data.extract::() { + let shape = Shape::new(vec![]); + let base_data = Arc::new(TensorData::from_vec_bool(vec![value_bool], device)); + let mut tensor = Tensor::new(base_data, shape, DataType::Bool, device, requires_grad); + if dtype != DataType::Bool { + tensor = tensor.astype(dtype).map_err(_convert_error)?; + } + return Ok(tensor); + } + + if let Ok(value_int) = data.extract::() { + let shape = Shape::new(vec![]); + let base_data = Arc::new(TensorData::from_vec_i64(vec![value_int], device)); + let mut tensor = Tensor::new(base_data, shape, DataType::Int64, device, requires_grad); + if dtype != DataType::Int64 { + tensor = tensor.astype(dtype).map_err(_convert_error)?; + } + return Ok(tensor); + } + + if let Ok(value_float) = data.extract::() { + let shape = Shape::new(vec![]); + let base_data = Arc::new(TensorData::from_vec_f64(vec![value_float], device)); + let mut tensor = Tensor::new(base_data, shape, DataType::Float64, device, requires_grad); + if dtype != DataType::Float64 { + tensor = tensor.astype(dtype).map_err(_convert_error)?; + } + return Ok(tensor); + } + + let float_name = intern!(data.py(), "__float__"); + if data.hasattr(float_name)? { + let method = data.getattr(float_name)?; + if method.is_callable() { + let float_obj = method.call0()?; + let val = float_obj.extract::()?; + let shape = Shape::new(vec![]); + let base_data = Arc::new(TensorData::from_vec_f64(vec![val], device)); + let mut tensor = + Tensor::new(base_data, shape, DataType::Float64, device, requires_grad); + if dtype != DataType::Float64 { + tensor = tensor.astype(dtype).map_err(_convert_error)?; + } + return Ok(tensor); + } + } + + Err(PyErr::new::( + "Unsupported data type for tensor creation", + )) +} + +pub(crate) fn apply_binary_ufunc( + operands: &[Tensor], + kind: BinaryOpKind, + op: F, +) -> PyResult +where + F: Fn(&Tensor, &Tensor) -> Result, +{ + if operands.len() != 2 { + return Err(PyValueError::new_err( + "Binary ufuncs require exactly two operands", + )); + } + + let (lhs_cast, rhs_cast, _) = + coerce_binary_operands(&operands[0], &operands[1], kind).map_err(_convert_error)?; + + let lhs_tensor = match lhs_cast { + Cow::Borrowed(tensor) => tensor.clone(), + Cow::Owned(tensor) => tensor, + }; + let rhs_tensor = match rhs_cast { + Cow::Borrowed(tensor) => tensor.clone(), + Cow::Owned(tensor) => tensor, + }; + + op(&lhs_tensor, &rhs_tensor).map_err(_convert_error) +} + +pub(crate) fn apply_unary_ufunc(operands: &[Tensor], op: F) -> PyResult +where + F: Fn(&Tensor) -> Result, +{ + if operands.len() != 1 { + return Err(PyValueError::new_err( + "Unary ufuncs require exactly one operand", + )); + } + + let tensor = operands[0].clone(); + op(&tensor).map_err(_convert_error) +} + +pub(crate) fn py_not_implemented(py: Python) -> PyResult> { + unsafe { + Ok( + pyo3::Bound::::from_borrowed_ptr(py, pyo3::ffi::Py_NotImplemented()) + .unbind(), + ) + } +} + +pub(crate) fn parse_dtype_like(value: &Bound) -> PyResult { + if let Ok(name) = value.extract::() { + dtype::parse_dtype(&name) + } else { + Err(PyTypeError::new_err( + "dtype must be specified as a string such as 'float32'", + )) + } +} + +pub(crate) fn parse_device_like(value: &Bound) -> PyResult { + if let Ok(device) = value.extract::() { + return Ok(device.device()); + } + + if let Ok(spec) = value.extract::() { + return Device::from_str(&spec).map_err(|err| { + PyValueError::new_err(format!("Unsupported device specification '{spec}': {err}")) + }); + } + + Err(PyTypeError::new_err( + "device must be specified as a Device object or string like 'cpu' or 'cuda:0'", + )) +} + +pub(crate) fn ensure_backward_gradient_compatible( + reference: &Tensor, + gradient: &mut Tensor, +) -> PyResult<()> { + let expected_shape = reference.shape().dims(); + let actual_shape = gradient.shape().dims(); + if expected_shape != actual_shape { + return Err(PyRuntimeError::new_err(format!( + "backward() expected gradient tensor with shape {:?}, but got {:?}", + expected_shape, actual_shape + ))); + } + + if gradient.device() != reference.device() { + *gradient = gradient.to(reference.device()).map_err(_convert_error)?; + } + + if gradient.dtype() != reference.dtype() { + *gradient = gradient.astype(reference.dtype()).map_err(_convert_error)?; + } + + if gradient.requires_grad() { + *gradient = gradient.detach(); + } + + Ok(()) +} + +pub(crate) fn tensor_from_py_value(reference: &Tensor, value: &Bound) -> PyResult { + if let Some(py_tensor) = extract_wrapped_pytensor(value) { + return Ok(py_tensor.inner.clone()); + } + + if let Ok(numpy_module) = PyModule::import(value.py(), "numpy") + && let Ok(ndarray_type) = numpy_module.getattr("ndarray") + && value.is_instance(&ndarray_type)? + { + if let Ok(dtype_obj) = value.getattr("dtype") { + let dtype_str = dtype_obj.str()?.to_str()?.to_ascii_lowercase(); + if let Ok(array_dtype) = dtype::parse_dtype(&dtype_str) { + return convert_python_data_to_tensor( + value, + array_dtype, + reference.device(), + false, + ); + } + } + return convert_python_data_to_tensor(value, reference.dtype(), reference.device(), false); + } + + if let Ok(py_tensor) = PyTensor::from_python_value(value) { + let mut tensor = py_tensor.inner; + if tensor.device() != reference.device() { + tensor = tensor.to(reference.device()).map_err(_convert_error)?; + } + + let target_dtype = dtype::resolve_scalar_dtype(value, reference.dtype()) + .ok() + .or_else(|| infer_python_value_dtype(value)) + .unwrap_or(reference.dtype()); + + if tensor.dtype() != target_dtype { + tensor = tensor.astype(target_dtype).map_err(_convert_error)?; + } + + return Ok(tensor); + } + + let index_name = intern!(value.py(), "__index__"); + if value.hasattr(index_name)? { + let method = value.getattr(index_name)?; + if method.is_callable() { + let result = method.call0()?; + if result.is_instance_of::() { + let dtype = match dtype::resolve_scalar_dtype(value, reference.dtype()) { + Ok(dt) => dt, + Err(_) => reference.dtype(), + }; + return convert_python_data_to_tensor( + result.as_any(), + dtype, + reference.device(), + false, + ); + } + } + } + + let dtype = match dtype::resolve_scalar_dtype(value, reference.dtype()) { + Ok(dt) => dt, + Err(_) => infer_python_value_dtype(value).unwrap_or(reference.dtype()), + }; + convert_python_data_to_tensor(value, dtype, reference.device(), false) +} + +pub(crate) fn tensor_bool_from_py(value: &Bound, device: Device) -> PyResult { + if let Some(py_tensor) = extract_wrapped_pytensor(value) { + let mut tensor = py_tensor.inner.clone(); + if tensor.dtype() != DataType::Bool { + return Err(PyTypeError::new_err("mask must be a bool tensor")); + } + if tensor.device() != device { + tensor = tensor.to(device).map_err(_convert_error)?; + } + return Ok(tensor); + } + + if let Ok(value_bool) = value.extract::() { + let data = Arc::new(TensorData::from_vec_bool(vec![value_bool], device)); + return Ok(Tensor::new( + data, + Shape::new(vec![]), + DataType::Bool, + device, + false, + )); + } + + convert_python_data_to_tensor(value, DataType::Bool, device, false) +} + +fn promote_dtypes(a: DataType, b: DataType) -> DataType { + use DataType::*; + + if a == b { + return a; + } + + match (a, b) { + (Float64, _) | (_, Float64) => Float64, + (Float32, _) | (_, Float32) => Float32, + (Int64, _) | (_, Int64) => Int64, + (Int32, _) | (_, Int32) => Int32, + _ => Bool, + } +} + +pub(crate) fn infer_python_value_dtype(value: &Bound) -> Option { + if let Some(py_tensor) = extract_wrapped_pytensor(value) { + return Some(py_tensor.inner.dtype()); + } + + if value.extract::().is_ok() { + return Some(DataType::Bool); + } + + if value.extract::().is_ok() { + return Some(DataType::Int64); + } + + if value.extract::().is_ok() { + return Some(dtype::default_dtype()); + } + + if let Ok(numpy_module) = PyModule::import(value.py(), "numpy") + && let Ok(ndarray_type) = numpy_module.getattr("ndarray") + && let Ok(true) = value.is_instance(&ndarray_type) + && let Ok(dtype_obj) = value.getattr("dtype") + && let Ok(dtype_str) = dtype_obj.str() + && let Ok(dtype) = dtype::parse_dtype(&dtype_str.to_str().ok()?.to_ascii_lowercase()) + { + return Some(dtype); + } + + if let Ok(list) = value.cast::() { + return infer_sequence_dtype(list.iter()); + } + + if let Ok(tuple) = value.cast::() { + return infer_sequence_dtype(tuple.iter()); + } + + None +} + +fn infer_sequence_dtype<'py, I>(iter: I) -> Option +where + I: Iterator>, +{ + let mut dtype: Option = None; + for item in iter { + let item_dtype = infer_python_value_dtype(&item)?; + dtype = Some(match dtype { + Some(current) => promote_dtypes(current, item_dtype), + None => item_dtype, + }); + } + dtype +} + +pub(crate) fn prepare_binary_operands_from_py( + reference: &Tensor, + other: &Bound, + reverse: bool, + kind: BinaryOpKind, +) -> PyResult<(Tensor, Tensor)> { + let lhs_input = if reverse { + tensor_from_py_value(reference, other)? + } else { + reference.clone() + }; + + let rhs_input = if reverse { + reference.clone() + } else { + tensor_from_py_value(reference, other)? + }; + + let (lhs_cast, rhs_cast, _) = + coerce_binary_operands(&lhs_input, &rhs_input, kind).map_err(_convert_error)?; + let lhs_tensor = match lhs_cast { + Cow::Borrowed(_) => lhs_input.clone(), + Cow::Owned(tensor) => tensor, + }; + let rhs_tensor = match rhs_cast { + Cow::Borrowed(_) => rhs_input.clone(), + Cow::Owned(tensor) => tensor, + }; + + Ok((lhs_tensor, rhs_tensor)) +} + +fn flatten_python_data(list: &Bound) -> PyResult<(Vec, Vec)> { + let mut shape = vec![list.len()]; + let mut flat_data = vec![]; + + fn process_nested( + item: &Bound, + depth: usize, + shape: &mut Vec, + flat_data: &mut Vec, + ) -> PyResult<()> { + if let Ok(nested_list) = item.cast::() { + let length = nested_list.len(); + if depth >= shape.len() { + shape.push(length); + } else if shape[depth] != length { + return Err(PyErr::new::( + "Inconsistent nested sequence lengths", + )); + } + for nested_item in nested_list.iter() { + process_nested(&nested_item, depth + 1, shape, flat_data)?; + } + return Ok(()); + } + + if let Ok(nested_tuple) = item.cast::() { + let list = nested_tuple.to_list(); + let length = list.len(); + if depth >= shape.len() { + shape.push(length); + } else if shape[depth] != length { + return Err(PyErr::new::( + "Inconsistent nested sequence lengths", + )); + } + for nested_item in list.iter() { + process_nested(&nested_item, depth + 1, shape, flat_data)?; + } + return Ok(()); + } + + if let Ok(value_bool) = item.extract::() { + flat_data.push(ScalarValue::Bool(value_bool)); + return Ok(()); + } + + if let Ok(value_int) = item.extract::() { + flat_data.push(ScalarValue::Int(value_int)); + return Ok(()); + } + + let index_name = intern!(item.py(), "__index__"); + if item.hasattr(index_name)? { + let method = item.getattr(index_name)?; + if method.is_callable() { + let result = method.call0()?; + if result.is_instance_of::() { + let value = result.extract::()?; + flat_data.push(ScalarValue::Int(value)); + return Ok(()); + } + } + } + + if let Ok(value_float) = item.extract::() { + flat_data.push(ScalarValue::Float(value_float)); + return Ok(()); + } + + let float_name = intern!(item.py(), "__float__"); + if item.hasattr(float_name)? { + let method = item.getattr(float_name)?; + if method.is_callable() { + let float_obj = method.call0()?; + let value = float_obj.extract::()?; + flat_data.push(ScalarValue::Float(value)); + return Ok(()); + } + } + + Err(PyErr::new::( + "Unsupported scalar type in nested sequence", + )) + } + + for item in list.iter() { + process_nested(&item, 1, &mut shape, &mut flat_data)?; + } + + Ok((shape, flat_data)) +} + +#[derive(Clone, Copy)] +enum ScalarValue { + Bool(bool), + Int(i64), + Float(f64), +} + +impl ScalarValue { + fn kind(&self) -> ScalarKind { + match self { + ScalarValue::Bool(_) => ScalarKind::Bool, + ScalarValue::Int(_) => ScalarKind::Int, + ScalarValue::Float(_) => ScalarKind::Float, + } + } + + fn to_bool(self) -> bool { + match self { + ScalarValue::Bool(value) => value, + ScalarValue::Int(value) => value != 0, + ScalarValue::Float(value) => value != 0.0, + } + } + + fn to_i64(self) -> i64 { + match self { + ScalarValue::Bool(value) => value as i64, + ScalarValue::Int(value) => value, + ScalarValue::Float(value) => value as i64, + } + } + + fn to_f64(self) -> f64 { + match self { + ScalarValue::Bool(value) => { + if value { + 1.0 + } else { + 0.0 + } + } + ScalarValue::Int(value) => value as f64, + ScalarValue::Float(value) => value, + } + } +} + +#[derive(Clone, Copy, PartialEq, Eq)] +enum ScalarKind { + Bool, + Int, + Float, +} + +impl ScalarKind { + fn combine(self, other: ScalarKind) -> ScalarKind { + use ScalarKind::*; + match (self, other) { + (Float, _) | (_, Float) => Float, + (Int, _) | (_, Int) => Int, + _ => Bool, + } + } +} + +fn tensor_from_flat_scalars( + shape: Vec, + values: Vec, + device: Device, + requires_grad: bool, +) -> PyResult<(Tensor, DataType)> { + let mut kind = ScalarKind::Bool; + for value in &values { + kind = kind.combine(value.kind()); + } + + let tensor = match kind { + ScalarKind::Bool => { + let data: Vec = values.into_iter().map(ScalarValue::to_bool).collect(); + Tensor::new( + Arc::new(TensorData::from_vec_bool(data, device)), + Shape::new(shape), + DataType::Bool, + device, + requires_grad, + ) + } + ScalarKind::Int => { + let data: Vec = values.into_iter().map(ScalarValue::to_i64).collect(); + Tensor::new( + Arc::new(TensorData::from_vec_i64(data, device)), + Shape::new(shape), + DataType::Int64, + device, + requires_grad, + ) + } + ScalarKind::Float => { + let data: Vec = values.into_iter().map(ScalarValue::to_f64).collect(); + Tensor::new( + Arc::new(TensorData::from_vec_f64(data, device)), + Shape::new(shape), + DataType::Float64, + device, + requires_grad, + ) + } + }; + + let dtype = tensor.dtype(); + Ok((tensor, dtype)) +} +fn parse_index(item: &Bound, dim_size: usize) -> PyResult { + if let Ok(i) = item.extract::() { + let mut idx = i; + if idx < 0 { + idx += dim_size as isize; + } + if idx < 0 || idx >= dim_size as isize { + return Err(PyIndexError::new_err("Index out of bounds")); + } + Ok(TensorIndex::Index(idx as usize)) + } else if let Ok(slice) = item.cast::() { + use std::convert::TryInto; + + let dim_size_isize: isize = dim_size + .try_into() + .map_err(|_| PyValueError::new_err("dim_size too large"))?; + let indices = slice.indices(dim_size_isize)?; + if indices.step <= 0 { + return Err(PyIndexError::new_err("slice step must be positive")); + } + Ok(TensorIndex::Slice { + start: indices.start.max(0) as usize, + end: indices.stop.max(0) as usize, + step: indices.step as usize, + }) + } else if item.is_none() { + Ok(TensorIndex::Slice { + start: 0, + end: dim_size, + step: 1, + }) + } else { + Err(PyTypeError::new_err("Invalid index type")) + } +} + +/// Parse a `__getitem__` key into per-dimension indices plus the output-axis +/// positions where a size-1 axis must be inserted for `None`/`np.newaxis`. +/// +/// Unlike [`parse_indices`], `None` does not consume an input dimension; it +/// inserts a new length-1 axis at the corresponding output position. +/// Integer indices drop their dimension, slices keep it. +pub(crate) fn parse_getitem_indices( + key: &Bound, + shape: &[usize], +) -> PyResult<(Vec, Vec)> { + let items: Vec> = if let Ok(tup) = key.cast::() { + tup.iter().collect() + } else { + vec![key.clone()] + }; + + // `None` entries add axes rather than selecting dimensions, so only the + // real (non-None) entries count against the tensor rank. + let real_count = items.iter().filter(|it| !it.is_none()).count(); + if real_count > shape.len() { + return Err(PyIndexError::new_err("Too many indices")); + } + + let mut real_indices: Vec = Vec::with_capacity(shape.len()); + let mut newaxis_positions: Vec = Vec::new(); + let mut input_dim = 0usize; + let mut out_dim = 0usize; + + for item in &items { + if item.is_none() { + newaxis_positions.push(out_dim); + out_dim += 1; + continue; + } + let idx = parse_index(item, shape[input_dim])?; + // Integer indices remove the dimension; slices keep it in the output. + let keeps_dim = matches!(idx, TensorIndex::Slice { .. }); + real_indices.push(idx); + input_dim += 1; + if keeps_dim { + out_dim += 1; + } + } + + // Any dimensions not addressed explicitly are taken in full. + for &dim in &shape[input_dim..] { + real_indices.push(TensorIndex::Slice { + start: 0, + end: dim, + step: 1, + }); + } + + Ok((real_indices, newaxis_positions)) +} + +pub(crate) fn parse_indices(key: &Bound, shape: &[usize]) -> PyResult> { + if let Ok(tup) = key.cast::() { + if tup.len() > shape.len() { + return Err(PyIndexError::new_err("Too many indices")); + } + let mut result = Vec::new(); + for (i, dim) in shape.iter().enumerate() { + if i < tup.len() { + result.push(parse_index(&tup.get_item(i)?, *dim)?); + } else { + result.push(TensorIndex::Slice { + start: 0, + end: *dim, + step: 1, + }); + } + } + Ok(result) + } else { + let mut result = vec![parse_index(key, shape[0])?]; + for dim in &shape[1..] { + result.push(TensorIndex::Slice { + start: 0, + end: *dim, + step: 1, + }); + } + Ok(result) + } +} + +pub(crate) fn convert_numpy_to_tensor( + array: &Bound, + requires_grad: bool, +) -> PyResult { + if let Ok(array_f32) = array.cast::>() { + let readonly = array_f32.readonly(); + let shape = Shape::new(readonly.shape().to_vec()); + let data_vec: Vec = readonly.as_slice()?.to_vec(); + let tensor_data = Arc::new(TensorData::from_vec( + data_vec, + DataType::Float32, + Device::cpu(), + )); + Ok(Tensor::new( + tensor_data, + shape, + DataType::Float32, + Device::cpu(), + requires_grad, + )) + } else if let Ok(array_f64) = array.cast::>() { + let readonly = array_f64.readonly(); + let shape = Shape::new(readonly.shape().to_vec()); + let data_vec: Vec = readonly.as_slice()?.to_vec(); + let tensor_data = Arc::new(TensorData::from_vec( + data_vec, + DataType::Float64, + Device::cpu(), + )); + Ok(Tensor::new( + tensor_data, + shape, + DataType::Float64, + Device::cpu(), + requires_grad, + )) + } else if let Ok(array_i32) = array.cast::>() { + let readonly = array_i32.readonly(); + let shape = Shape::new(readonly.shape().to_vec()); + let data_vec: Vec = readonly.as_slice()?.to_vec(); + let tensor_data = Arc::new(TensorData::from_vec( + data_vec, + DataType::Int32, + Device::cpu(), + )); + Ok(Tensor::new( + tensor_data, + shape, + DataType::Int32, + Device::cpu(), + requires_grad, + )) + } else if let Ok(array_i64) = array.cast::>() { + let readonly = array_i64.readonly(); + let shape = Shape::new(readonly.shape().to_vec()); + let data_vec: Vec = readonly.as_slice()?.to_vec(); + let tensor_data = Arc::new(TensorData::from_vec( + data_vec, + DataType::Int64, + Device::cpu(), + )); + Ok(Tensor::new( + tensor_data, + shape, + DataType::Int64, + Device::cpu(), + requires_grad, + )) + } else if let Ok(array_bool) = array.cast::>() { + let readonly = array_bool.readonly(); + let shape = Shape::new(readonly.shape().to_vec()); + let data_vec: Vec = readonly.as_slice()?.to_vec(); + let tensor_data = Arc::new(TensorData::from_vec( + data_vec, + DataType::Bool, + Device::cpu(), + )); + Ok(Tensor::new( + tensor_data, + shape, + DataType::Bool, + Device::cpu(), + requires_grad, + )) + } else { + Err(PyErr::new::( + "Unsupported NumPy array type", + )) + } +} diff --git a/bindings/src/tensor/python/numpy.rs b/bindings/src/tensor/python/numpy.rs index 1a0382ee..f07adf91 100644 --- a/bindings/src/tensor/python/numpy.rs +++ b/bindings/src/tensor/python/numpy.rs @@ -1,970 +1,975 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -fn convert_tensor_to_numpy(tensor: &Tensor, py: Python, _force_copy: bool) -> PyResult> { - if tensor.device() != Device::cpu() { - return Err(PyErr::new::( - "Cannot convert GPU tensor to NumPy array. Use .cpu() first.", - )); - } - - let shape = tensor.shape().dims(); - let strides = tensor.strides().as_slice(); - let numel: usize = shape.iter().product(); - - macro_rules! to_numpy { - ($slice:expr, $ty:ty) => {{ - let data = $slice.ok_or_else(|| { - PyErr::new::("Failed to get tensor data") - })?; - let mut out = Vec::<$ty>::with_capacity(numel); - let mut indices = vec![0usize; shape.len()]; - for _ in 0..numel { - let mut offset = 0usize; - for (idx, stride) in indices.iter().zip(strides) { - offset += idx * stride; - } - out.push(data[offset]); - for axis in (0..indices.len()).rev() { - indices[axis] += 1; - if indices[axis] < shape[axis] { - break; - } - indices[axis] = 0; - } - } - let array = PyArray::from_vec(py, out).reshape(shape)?; - Ok(array.into_any().unbind()) - }}; - } - - let array: PyResult> = match tensor.dtype() { - DataType::Float32 => to_numpy!(tensor.data().as_f32_slice(), f32), - DataType::Float64 => to_numpy!(tensor.data().as_f64_slice(), f64), - DataType::Int32 => to_numpy!(tensor.data().as_i32_slice(), i32), - DataType::Int64 => to_numpy!(tensor.data().as_i64_slice(), i64), - DataType::Bool => to_numpy!(tensor.data().as_bool_slice(), bool), - }; - - array -} - -fn convert_tensor_to_python_list(tensor: &Tensor, py: Python) -> PyResult> { - let shape: Vec = tensor.shape().dims().to_vec(); - match tensor.dtype() { - DataType::Float32 => { - let data = tensor.data().as_f32_slice().ok_or_else(|| { - PyErr::new::("Failed to get f32 data") - })?; - nested_list_from_slice(py, data, &shape) - } - DataType::Float64 => { - let data = tensor.data().as_f64_slice().ok_or_else(|| { - PyErr::new::("Failed to get f64 data") - })?; - nested_list_from_slice(py, data, &shape) - } - DataType::Int32 => { - let data = tensor.data().as_i32_slice().ok_or_else(|| { - PyErr::new::("Failed to get i32 data") - })?; - nested_list_from_slice(py, data, &shape) - } - DataType::Int64 => { - let data = tensor.data().as_i64_slice().ok_or_else(|| { - PyErr::new::("Failed to get i64 data") - })?; - nested_list_from_slice(py, data, &shape) - } - DataType::Bool => { - let data = tensor.data().as_bool_slice().ok_or_else(|| { - PyErr::new::("Failed to get bool data") - })?; - nested_list_from_slice(py, data, &shape) - } - } -} - -fn convert_tensor_to_python_scalar(tensor: &Tensor, py: Python) -> PyResult> { - if tensor.numel() != 1 { - return Err(PyErr::new::(format!( - "a Tensor with {} elements cannot be converted to Scalar", - tensor.numel() - ))); - } - - match tensor.dtype() { - DataType::Float32 => { - let data = tensor.data().as_f32_slice().ok_or_else(|| { - PyErr::new::("Failed to get f32 data") - })?; - data[0].into_py_any(py) - } - DataType::Float64 => { - let data = tensor.data().as_f64_slice().ok_or_else(|| { - PyErr::new::("Failed to get f64 data") - })?; - data[0].into_py_any(py) - } - DataType::Int32 => { - let data = tensor.data().as_i32_slice().ok_or_else(|| { - PyErr::new::("Failed to get i32 data") - })?; - data[0].into_py_any(py) - } - DataType::Int64 => { - let data = tensor.data().as_i64_slice().ok_or_else(|| { - PyErr::new::("Failed to get i64 data") - })?; - data[0].into_py_any(py) - } - DataType::Bool => { - let data = tensor.data().as_bool_slice().ok_or_else(|| { - PyErr::new::("Failed to get bool data") - })?; - data[0].into_py_any(py) - } - } -} - -fn nested_list_from_slice<'py, T>( - py: Python<'py>, - data: &[T], - shape: &[usize], -) -> PyResult> -where - T: Copy + IntoPyObjectExt<'py>, -{ - if shape.is_empty() { - if let Some(value) = data.first() { - return (*value).into_py_any(py); - } - return PyList::empty(py).into_py_any(py); - } - - if shape.len() == 1 { - let mut elements: Vec> = Vec::with_capacity(data.len()); - for value in data.iter().copied() { - elements.push(value.into_py_any(py)?); - } - let list = PyList::new(py, elements)?; - return list.into_py_any(py); - } - - let chunk = shape[1..] - .iter() - .fold(1usize, |acc, &dim| acc.saturating_mul(dim)); - let mut parts: Vec> = Vec::with_capacity(shape[0]); - for index in 0..shape[0] { - let start = index * chunk; - let end = start + chunk; - let slice = if start <= end && end <= data.len() { - &data[start..end] - } else { - &[] - }; - parts.push(nested_list_from_slice(py, slice, &shape[1..])?); - } - - let list = PyList::new(py, parts)?; - list.into_py_any(py) -} - -fn create_random_tensor( - shape: Shape, - dtype: DataType, - device: Device, - requires_grad: bool, - normal: bool, -) -> PyResult { - let mut tensor_data = TensorData::uninitialized_on_device(shape.numel(), dtype, device); - - match dtype { - DataType::Float32 => { - if let Some(slice) = tensor_data.as_f32_slice_mut() { - use rand::RngExt; - random::with_rng(|rng| { - if normal { - use rand_distr::{Distribution, Normal}; - let normal_dist = Normal::new(0.0f32, 1.0f32).unwrap(); - for val in slice.iter_mut() { - *val = normal_dist.sample(rng); - } - } else { - for val in slice.iter_mut() { - *val = rng.random::(); - } - } - }); - } - } - DataType::Float64 => { - if let Some(slice) = tensor_data.as_f64_slice_mut() { - use rand::RngExt; - random::with_rng(|rng| { - if normal { - use rand_distr::{Distribution, Normal}; - let normal_dist = Normal::new(0.0f64, 1.0f64).unwrap(); - for val in slice.iter_mut() { - *val = normal_dist.sample(rng); - } - } else { - for val in slice.iter_mut() { - *val = rng.random::(); - } - } - }); - } - } - DataType::Int32 => { - if let Some(slice) = tensor_data.as_i32_slice_mut() { - use rand::RngExt; - random::with_rng(|rng| { - if normal { - use rand_distr::{Distribution, Normal}; - let normal_dist = Normal::new(0.0f32, 1.0f32).unwrap(); - for val in slice.iter_mut() { - *val = normal_dist.sample(rng) as i32; - } - } else { - for val in slice.iter_mut() { - *val = rng.random::(); - } - } - }); - } - } - DataType::Int64 => { - if let Some(slice) = tensor_data.as_i64_slice_mut() { - use rand::RngExt; - random::with_rng(|rng| { - if normal { - use rand_distr::{Distribution, Normal}; - let normal_dist = Normal::new(0.0f64, 1.0f64).unwrap(); - for val in slice.iter_mut() { - *val = normal_dist.sample(rng) as i64; - } - } else { - for val in slice.iter_mut() { - *val = rng.random::(); - } - } - }); - } - } - DataType::Bool => { - if let Some(slice) = tensor_data.as_bool_slice_mut() { - use rand::RngExt; - random::with_rng(|rng| { - for val in slice.iter_mut() { - *val = rng.random::(); - } - }); - } - } - } - - Ok(Tensor::new( - Arc::new(tensor_data), - shape, - dtype, - device, - requires_grad, - )) -} - -enum FanInitKind { - XavierUniform, - XavierNormal, - HeUniform, - HeNormal, - LecunUniform, - LecunNormal, -} - -impl FanInitKind { - fn apply( - &self, - shape: Shape, - dtype: DataType, - device: Device, - requires_grad: bool, - ) -> Result { - match self { - FanInitKind::XavierUniform => { - nn::init::xavier_uniform_init(shape, dtype, device, requires_grad) - } - FanInitKind::XavierNormal => { - nn::init::xavier_normal_init(shape, dtype, device, requires_grad) - } - FanInitKind::HeUniform => { - nn::init::he_uniform_init(shape, dtype, device, requires_grad) - } - FanInitKind::HeNormal => nn::init::he_normal_init(shape, dtype, device, requires_grad), - FanInitKind::LecunUniform => { - nn::init::lecun_uniform_init(shape, dtype, device, requires_grad) - } - FanInitKind::LecunNormal => { - nn::init::lecun_normal_init(shape, dtype, device, requires_grad) - } - } - } -} - -fn ensure_float_dtype(dtype: DataType, context: &str) -> PyResult<()> { - match dtype { - DataType::Float32 | DataType::Float64 => Ok(()), - _ => Err(PyValueError::new_err(format!( - "{context} only supports float32 or float64 dtypes", - ))), - } -} - -fn ensure_valid_fan_shape(shape: &Shape, context: &str) -> PyResult<()> { - if shape.dims().contains(&0) { - Err(PyValueError::new_err(format!( - "{context} requires all shape dimensions to be at least 1", - ))) - } else { - Ok(()) - } -} - -fn create_fan_init_tensor( - shape: Shape, - dtype: DataType, - device: Device, - requires_grad: bool, - kind: FanInitKind, - context: &str, -) -> PyResult { - ensure_float_dtype(dtype, context)?; - ensure_valid_fan_shape(&shape, context)?; - let tensor = kind.apply(shape, dtype, device, requires_grad); - tensor.map_err(_convert_error) -} - -fn create_uniform_tensor( - shape: Shape, - dtype: DataType, - device: Device, - requires_grad: bool, - low: f64, - high: f64, -) -> PyResult { - if !low.is_finite() || !high.is_finite() { - return Err(PyValueError::new_err( - "uniform requires finite low and high values", - )); - } - - if high.partial_cmp(&low) != Some(Ordering::Greater) { - return Err(PyValueError::new_err( - "uniform requires high to be greater than low", - )); - } - - let tensor = nn::init::init_uniform(shape, low, high, dtype, device, requires_grad) - .map_err(_convert_error)?; - Ok(tensor) -} - -#[allow(clippy::too_many_arguments)] -fn create_truncated_normal_tensor( - shape: Shape, - dtype: DataType, - device: Device, - requires_grad: bool, - mean: f64, - std: f64, - lower: Option, - upper: Option, - context: &str, -) -> PyResult { - ensure_float_dtype(dtype, context)?; - - if !mean.is_finite() { - return Err(PyValueError::new_err(format!( - "{context} requires a finite mean", - ))); - } - - if !std.is_finite() || std <= 0.0 { - return Err(PyValueError::new_err(format!( - "{context} requires std to be a positive finite value", - ))); - } - - let default_lower = mean - 2.0 * std; - let default_upper = mean + 2.0 * std; - let lower = lower.unwrap_or(default_lower); - let upper = upper.unwrap_or(default_upper); - - if lower.is_nan() || upper.is_nan() { - return Err(PyValueError::new_err(format!( - "{context} requires non-NaN bounds", - ))); - } - - if upper.partial_cmp(&lower) != Some(Ordering::Greater) { - return Err(PyValueError::new_err(format!( - "{context} requires upper bound to be greater than lower bound", - ))); - } - - let tensor = nn::init::truncated_normal_init( - shape, - mean, - std, - lower, - upper, - dtype, - device, - requires_grad, - ) - .map_err(_convert_error)?; - Ok(tensor) -} - -fn prepare_new_tensor_from_existing( - source: &Tensor, - dtype: DataType, - device: Device, - requires_grad: bool, -) -> PyResult { - let mut tensor = source.detach(); - - if tensor.device() != device { - tensor = tensor.to(device).map_err(_convert_error)?; - } - - if tensor.dtype() != dtype { - tensor = tensor.astype(dtype).map_err(_convert_error)?; - } - - tensor = tensor.deep_clone().map_err(_convert_error)?; - - if requires_grad { - tensor = tensor.requires_grad_(true); - } - - Ok(tensor) -} - -fn adapt_tensor_for_as_tensor( - source: &Tensor, - dtype: DataType, - device: Device, - requires_grad: bool, - copy: bool, -) -> PyResult { - if !copy - && source.dtype() == dtype - && source.device() == device - && source.requires_grad() == requires_grad - { - return Ok(source.clone()); - } - - let mut tensor = if copy || (source.requires_grad() && !requires_grad) { - source.detach() - } else { - source.clone() - }; - - if tensor.device() != device { - tensor = tensor.to(device).map_err(_convert_error)?; - } - - if tensor.dtype() != dtype { - tensor = tensor.astype(dtype).map_err(_convert_error)?; - } - - if copy { - tensor = tensor.deep_clone().map_err(_convert_error)?; - } - - if tensor.requires_grad() != requires_grad { - tensor = tensor.requires_grad_(requires_grad); - } - - Ok(tensor) -} - -fn create_randint_tensor( - shape: Shape, - dtype: DataType, - device: Device, - requires_grad: bool, - low: i64, - high: i64, -) -> PyResult { - let tensor_data = match dtype { - DataType::Int32 => { - let low_i32 = i32::try_from(low) - .map_err(|_| PyValueError::new_err("low is out of range for dtype int32"))?; - let high_i32 = i32::try_from(high) - .map_err(|_| PyValueError::new_err("high is out of range for dtype int32"))?; - if low_i32 >= high_i32 { - return Err(PyValueError::new_err( - "randint requires that low < high after casting to int32", - )); - } - let mut values = vec![0i32; shape.numel()]; - random::with_rng(|rng| { - use rand::RngExt; - for value in &mut values { - *value = rng.random_range(low_i32..high_i32); - } - }); - TensorData::from_vec_i32(values, device) - } - DataType::Int64 => { - if high <= low { - return Err(PyValueError::new_err("randint requires that low < high")); - } - let mut values = vec![0i64; shape.numel()]; - random::with_rng(|rng| { - use rand::RngExt; - for value in &mut values { - *value = rng.random_range(low..high); - } - }); - TensorData::from_vec_i64(values, device) - } - _ => { - return Err(PyValueError::new_err( - "randint only supports int32 or int64 dtypes", - )); - } - }; - - Ok(Tensor::new( - Arc::new(tensor_data), - shape, - dtype, - device, - requires_grad, - )) -} - -fn create_randperm_tensor( - n: usize, - dtype: DataType, - device: Device, - requires_grad: bool, -) -> PyResult { - let tensor_data = match dtype { - DataType::Int32 => { - let _ = i32::try_from(n).map_err(|_| { - PyValueError::new_err("randperm with dtype int32 requires n <= i32::MAX") - })?; - let mut values = Vec::with_capacity(n); - for idx in 0..n { - values.push(i32::try_from(idx).map_err(|_| { - PyValueError::new_err("randperm with dtype int32 requires n <= i32::MAX") - })?); - } - random::with_rng(|rng| { - use rand::seq::SliceRandom; - values.shuffle(rng); - }); - TensorData::from_vec_i32(values, device) - } - DataType::Int64 => { - let _ = i64::try_from(n).map_err(|_| { - PyValueError::new_err("randperm with dtype int64 requires n <= i64::MAX") - })?; - let mut values = Vec::with_capacity(n); - for idx in 0..n { - values.push(idx as i64); - } - random::with_rng(|rng| { - use rand::seq::SliceRandom; - values.shuffle(rng); - }); - TensorData::from_vec_i64(values, device) - } - _ => { - return Err(PyValueError::new_err( - "randperm only supports int32 or int64 dtypes", - )); - } - }; - - Ok(Tensor::new( - Arc::new(tensor_data), - Shape::new(vec![n]), - dtype, - device, - requires_grad, - )) -} - -fn create_eye_tensor( - n: usize, - m: usize, - dtype: DataType, - device: Device, - requires_grad: bool, -) -> PyResult { - let shape = Shape::new(vec![n, m]); - let mut tensor_data = TensorData::zeros_on_device(shape.numel(), dtype, device); - - match dtype { - DataType::Float32 => { - if let Some(slice) = tensor_data.as_f32_slice_mut() { - for i in 0..n.min(m) { - slice[i * m + i] = 1.0; - } - } - } - DataType::Float64 => { - if let Some(slice) = tensor_data.as_f64_slice_mut() { - for i in 0..n.min(m) { - slice[i * m + i] = 1.0; - } - } - } - DataType::Int32 => { - if let Some(slice) = tensor_data.as_i32_slice_mut() { - for i in 0..n.min(m) { - slice[i * m + i] = 1; - } - } - } - DataType::Int64 => { - if let Some(slice) = tensor_data.as_i64_slice_mut() { - for i in 0..n.min(m) { - slice[i * m + i] = 1; - } - } - } - DataType::Bool => { - if let Some(slice) = tensor_data.as_bool_slice_mut() { - for i in 0..n.min(m) { - slice[i * m + i] = true; - } - } - } - } - - Ok(Tensor::new( - Arc::new(tensor_data), - shape, - dtype, - device, - requires_grad, - )) -} - -fn create_full_tensor( - shape: Vec, - fill_value: f64, - dtype: DataType, - device: Device, - requires_grad: bool, -) -> PyResult { - let shape = Shape::new(shape); - let mut tensor_data = TensorData::uninitialized_on_device(shape.numel(), dtype, device); - - match dtype { - DataType::Float32 => { - if let Some(slice) = tensor_data.as_f32_slice_mut() { - slice.fill(fill_value as f32); - } - } - DataType::Float64 => { - if let Some(slice) = tensor_data.as_f64_slice_mut() { - slice.fill(fill_value); - } - } - DataType::Int32 => { - if let Some(slice) = tensor_data.as_i32_slice_mut() { - slice.fill(fill_value as i32); - } - } - DataType::Int64 => { - if let Some(slice) = tensor_data.as_i64_slice_mut() { - slice.fill(fill_value as i64); - } - } - DataType::Bool => { - if let Some(slice) = tensor_data.as_bool_slice_mut() { - slice.fill(fill_value != 0.0); - } - } - } - - Ok(Tensor::new( - Arc::new(tensor_data), - shape, - dtype, - device, - requires_grad, - )) -} -fn create_arange_tensor( - start: f64, - end: f64, - step: f64, - dtype: DataType, - device: Device, - requires_grad: bool, -) -> PyResult { - if step == 0.0 { - return Err(PyErr::new::( - "Step cannot be zero", - )); - } - - let num_elements = ((end - start) / step).ceil() as usize; - let shape = Shape::new(vec![num_elements]); - let mut tensor_data = TensorData::uninitialized_on_device(shape.numel(), dtype, device); - - match dtype { - DataType::Float32 => { - if let Some(slice) = tensor_data.as_f32_slice_mut() { - for (i, val) in slice.iter_mut().enumerate() { - *val = (start + i as f64 * step) as f32; - } - } - } - DataType::Float64 => { - if let Some(slice) = tensor_data.as_f64_slice_mut() { - for (i, val) in slice.iter_mut().enumerate() { - *val = start + i as f64 * step; - } - } - } - DataType::Int32 => { - if let Some(slice) = tensor_data.as_i32_slice_mut() { - for (i, val) in slice.iter_mut().enumerate() { - *val = (start + i as f64 * step) as i32; - } - } - } - DataType::Int64 => { - if let Some(slice) = tensor_data.as_i64_slice_mut() { - for (i, val) in slice.iter_mut().enumerate() { - *val = (start + i as f64 * step) as i64; - } - } - } - DataType::Bool => { - if let Some(slice) = tensor_data.as_bool_slice_mut() { - for (i, val) in slice.iter_mut().enumerate() { - *val = (start + i as f64 * step) != 0.0; - } - } - } - } - - Ok(Tensor::new( - Arc::new(tensor_data), - shape, - dtype, - device, - requires_grad, - )) -} - -fn create_linspace_tensor( - start: f64, - end: f64, - steps: usize, - dtype: DataType, - device: Device, - requires_grad: bool, -) -> PyResult { - if steps == 0 { - return Err(PyErr::new::( - "Number of steps must be positive", - )); - } - - let shape = Shape::new(vec![steps]); - let mut tensor_data = TensorData::uninitialized_on_device(shape.numel(), dtype, device); - let denom = if steps > 1 { (steps - 1) as f64 } else { 1.0 }; - let step = if steps > 1 { - (end - start) / denom - } else { - 0.0 - }; - - match dtype { - DataType::Float32 => { - if let Some(slice) = tensor_data.as_f32_slice_mut() { - for (i, val) in slice.iter_mut().enumerate() { - let value = if steps == 1 { - start - } else { - start + i as f64 * step - }; - *val = value as f32; - } - } - } - DataType::Float64 => { - if let Some(slice) = tensor_data.as_f64_slice_mut() { - for (i, val) in slice.iter_mut().enumerate() { - let value = if steps == 1 { - start - } else { - start + i as f64 * step - }; - *val = value; - } - } - } - DataType::Int32 => { - if let Some(slice) = tensor_data.as_i32_slice_mut() { - for (i, val) in slice.iter_mut().enumerate() { - let value = if steps == 1 { - start - } else { - start + i as f64 * step - }; - *val = value.round() as i32; - } - } - } - DataType::Int64 => { - if let Some(slice) = tensor_data.as_i64_slice_mut() { - for (i, val) in slice.iter_mut().enumerate() { - let value = if steps == 1 { - start - } else { - start + i as f64 * step - }; - *val = value.round() as i64; - } - } - } - DataType::Bool => { - if let Some(slice) = tensor_data.as_bool_slice_mut() { - for (i, val) in slice.iter_mut().enumerate() { - let value = if steps == 1 { - start - } else { - start + i as f64 * step - }; - *val = value != 0.0; - } - } - } - } - - Ok(Tensor::new( - Arc::new(tensor_data), - shape, - dtype, - device, - requires_grad, - )) -} - -fn create_logspace_tensor( - start: f64, - end: f64, - steps: usize, - base: f64, - dtype: DataType, - device: Device, - requires_grad: bool, -) -> PyResult { - if steps == 0 { - return Err(PyErr::new::( - "Number of steps must be positive", - )); - } - - if base <= 0.0 { - return Err(PyErr::new::( - "Base must be positive", - )); - } - - let shape = Shape::new(vec![steps]); - let mut tensor_data = TensorData::uninitialized_on_device(shape.numel(), dtype, device); - let denom = if steps > 1 { (steps - 1) as f64 } else { 1.0 }; - let step = if steps > 1 { - (end - start) / denom - } else { - 0.0 - }; - - match dtype { - DataType::Float32 => { - if let Some(slice) = tensor_data.as_f32_slice_mut() { - for (i, val) in slice.iter_mut().enumerate() { - let exponent = if steps == 1 { - start - } else { - start + i as f64 * step - }; - *val = base.powf(exponent) as f32; - } - } - } - DataType::Float64 => { - if let Some(slice) = tensor_data.as_f64_slice_mut() { - for (i, val) in slice.iter_mut().enumerate() { - let exponent = if steps == 1 { - start - } else { - start + i as f64 * step - }; - *val = base.powf(exponent); - } - } - } - DataType::Int32 => { - if let Some(slice) = tensor_data.as_i32_slice_mut() { - for (i, val) in slice.iter_mut().enumerate() { - let exponent = if steps == 1 { - start - } else { - start + i as f64 * step - }; - *val = base.powf(exponent).round() as i32; - } - } - } - DataType::Int64 => { - if let Some(slice) = tensor_data.as_i64_slice_mut() { - for (i, val) in slice.iter_mut().enumerate() { - let exponent = if steps == 1 { - start - } else { - start + i as f64 * step - }; - *val = base.powf(exponent).round() as i64; - } - } - } - DataType::Bool => { - if let Some(slice) = tensor_data.as_bool_slice_mut() { - for (i, val) in slice.iter_mut().enumerate() { - let exponent = if steps == 1 { - start - } else { - start + i as f64 * step - }; - *val = base.powf(exponent) != 0.0; - } - } - } - } - - Ok(Tensor::new( - Arc::new(tensor_data), - shape, - dtype, - device, - requires_grad, - )) -} +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +pub(crate) fn convert_tensor_to_numpy( + tensor: &Tensor, + py: Python, + _force_copy: bool, +) -> PyResult> { + if tensor.device() != Device::cpu() { + return Err(PyErr::new::( + "Cannot convert GPU tensor to NumPy array. Use .cpu() first.", + )); + } + + let shape = tensor.shape().dims(); + let strides = tensor.strides().as_slice(); + let numel: usize = shape.iter().product(); + + macro_rules! to_numpy { + ($slice:expr, $ty:ty) => {{ + let data = $slice.ok_or_else(|| { + PyErr::new::("Failed to get tensor data") + })?; + let mut out = Vec::<$ty>::with_capacity(numel); + let mut indices = vec![0usize; shape.len()]; + for _ in 0..numel { + let mut offset = 0usize; + for (idx, stride) in indices.iter().zip(strides) { + offset += idx * stride; + } + out.push(data[offset]); + for axis in (0..indices.len()).rev() { + indices[axis] += 1; + if indices[axis] < shape[axis] { + break; + } + indices[axis] = 0; + } + } + let array = PyArray::from_vec(py, out).reshape(shape)?; + Ok(array.into_any().unbind()) + }}; + } + + let array: PyResult> = match tensor.dtype() { + DataType::Float32 => to_numpy!(tensor.data().as_f32_slice(), f32), + DataType::Float64 => to_numpy!(tensor.data().as_f64_slice(), f64), + DataType::Int32 => to_numpy!(tensor.data().as_i32_slice(), i32), + DataType::Int64 => to_numpy!(tensor.data().as_i64_slice(), i64), + DataType::Bool => to_numpy!(tensor.data().as_bool_slice(), bool), + }; + + array +} + +pub(crate) fn convert_tensor_to_python_list(tensor: &Tensor, py: Python) -> PyResult> { + let shape: Vec = tensor.shape().dims().to_vec(); + match tensor.dtype() { + DataType::Float32 => { + let data = tensor.data().as_f32_slice().ok_or_else(|| { + PyErr::new::("Failed to get f32 data") + })?; + nested_list_from_slice(py, data, &shape) + } + DataType::Float64 => { + let data = tensor.data().as_f64_slice().ok_or_else(|| { + PyErr::new::("Failed to get f64 data") + })?; + nested_list_from_slice(py, data, &shape) + } + DataType::Int32 => { + let data = tensor.data().as_i32_slice().ok_or_else(|| { + PyErr::new::("Failed to get i32 data") + })?; + nested_list_from_slice(py, data, &shape) + } + DataType::Int64 => { + let data = tensor.data().as_i64_slice().ok_or_else(|| { + PyErr::new::("Failed to get i64 data") + })?; + nested_list_from_slice(py, data, &shape) + } + DataType::Bool => { + let data = tensor.data().as_bool_slice().ok_or_else(|| { + PyErr::new::("Failed to get bool data") + })?; + nested_list_from_slice(py, data, &shape) + } + } +} + +pub(crate) fn convert_tensor_to_python_scalar(tensor: &Tensor, py: Python) -> PyResult> { + if tensor.numel() != 1 { + return Err(PyErr::new::(format!( + "a Tensor with {} elements cannot be converted to Scalar", + tensor.numel() + ))); + } + + match tensor.dtype() { + DataType::Float32 => { + let data = tensor.data().as_f32_slice().ok_or_else(|| { + PyErr::new::("Failed to get f32 data") + })?; + data[0].into_py_any(py) + } + DataType::Float64 => { + let data = tensor.data().as_f64_slice().ok_or_else(|| { + PyErr::new::("Failed to get f64 data") + })?; + data[0].into_py_any(py) + } + DataType::Int32 => { + let data = tensor.data().as_i32_slice().ok_or_else(|| { + PyErr::new::("Failed to get i32 data") + })?; + data[0].into_py_any(py) + } + DataType::Int64 => { + let data = tensor.data().as_i64_slice().ok_or_else(|| { + PyErr::new::("Failed to get i64 data") + })?; + data[0].into_py_any(py) + } + DataType::Bool => { + let data = tensor.data().as_bool_slice().ok_or_else(|| { + PyErr::new::("Failed to get bool data") + })?; + data[0].into_py_any(py) + } + } +} + +fn nested_list_from_slice<'py, T>( + py: Python<'py>, + data: &[T], + shape: &[usize], +) -> PyResult> +where + T: Copy + IntoPyObjectExt<'py>, +{ + if shape.is_empty() { + if let Some(value) = data.first() { + return (*value).into_py_any(py); + } + return PyList::empty(py).into_py_any(py); + } + + if shape.len() == 1 { + let mut elements: Vec> = Vec::with_capacity(data.len()); + for value in data.iter().copied() { + elements.push(value.into_py_any(py)?); + } + let list = PyList::new(py, elements)?; + return list.into_py_any(py); + } + + let chunk = shape[1..] + .iter() + .fold(1usize, |acc, &dim| acc.saturating_mul(dim)); + let mut parts: Vec> = Vec::with_capacity(shape[0]); + for index in 0..shape[0] { + let start = index * chunk; + let end = start + chunk; + let slice = if start <= end && end <= data.len() { + &data[start..end] + } else { + &[] + }; + parts.push(nested_list_from_slice(py, slice, &shape[1..])?); + } + + let list = PyList::new(py, parts)?; + list.into_py_any(py) +} + +pub(crate) fn create_random_tensor( + shape: Shape, + dtype: DataType, + device: Device, + requires_grad: bool, + normal: bool, +) -> PyResult { + let mut tensor_data = TensorData::uninitialized_on_device(shape.numel(), dtype, device); + + match dtype { + DataType::Float32 => { + if let Some(slice) = tensor_data.as_f32_slice_mut() { + use rand::RngExt; + random::with_rng(|rng| { + if normal { + use rand_distr::{Distribution, Normal}; + let normal_dist = Normal::new(0.0f32, 1.0f32).unwrap(); + for val in slice.iter_mut() { + *val = normal_dist.sample(rng); + } + } else { + for val in slice.iter_mut() { + *val = rng.random::(); + } + } + }); + } + } + DataType::Float64 => { + if let Some(slice) = tensor_data.as_f64_slice_mut() { + use rand::RngExt; + random::with_rng(|rng| { + if normal { + use rand_distr::{Distribution, Normal}; + let normal_dist = Normal::new(0.0f64, 1.0f64).unwrap(); + for val in slice.iter_mut() { + *val = normal_dist.sample(rng); + } + } else { + for val in slice.iter_mut() { + *val = rng.random::(); + } + } + }); + } + } + DataType::Int32 => { + if let Some(slice) = tensor_data.as_i32_slice_mut() { + use rand::RngExt; + random::with_rng(|rng| { + if normal { + use rand_distr::{Distribution, Normal}; + let normal_dist = Normal::new(0.0f32, 1.0f32).unwrap(); + for val in slice.iter_mut() { + *val = normal_dist.sample(rng) as i32; + } + } else { + for val in slice.iter_mut() { + *val = rng.random::(); + } + } + }); + } + } + DataType::Int64 => { + if let Some(slice) = tensor_data.as_i64_slice_mut() { + use rand::RngExt; + random::with_rng(|rng| { + if normal { + use rand_distr::{Distribution, Normal}; + let normal_dist = Normal::new(0.0f64, 1.0f64).unwrap(); + for val in slice.iter_mut() { + *val = normal_dist.sample(rng) as i64; + } + } else { + for val in slice.iter_mut() { + *val = rng.random::(); + } + } + }); + } + } + DataType::Bool => { + if let Some(slice) = tensor_data.as_bool_slice_mut() { + use rand::RngExt; + random::with_rng(|rng| { + for val in slice.iter_mut() { + *val = rng.random::(); + } + }); + } + } + } + + Ok(Tensor::new( + Arc::new(tensor_data), + shape, + dtype, + device, + requires_grad, + )) +} + +pub(crate) enum FanInitKind { + XavierUniform, + XavierNormal, + HeUniform, + HeNormal, + LecunUniform, + LecunNormal, +} + +impl FanInitKind { + fn apply( + &self, + shape: Shape, + dtype: DataType, + device: Device, + requires_grad: bool, + ) -> Result { + match self { + FanInitKind::XavierUniform => { + nn::init::xavier_uniform_init(shape, dtype, device, requires_grad) + } + FanInitKind::XavierNormal => { + nn::init::xavier_normal_init(shape, dtype, device, requires_grad) + } + FanInitKind::HeUniform => { + nn::init::he_uniform_init(shape, dtype, device, requires_grad) + } + FanInitKind::HeNormal => nn::init::he_normal_init(shape, dtype, device, requires_grad), + FanInitKind::LecunUniform => { + nn::init::lecun_uniform_init(shape, dtype, device, requires_grad) + } + FanInitKind::LecunNormal => { + nn::init::lecun_normal_init(shape, dtype, device, requires_grad) + } + } + } +} + +fn ensure_float_dtype(dtype: DataType, context: &str) -> PyResult<()> { + match dtype { + DataType::Float32 | DataType::Float64 => Ok(()), + _ => Err(PyValueError::new_err(format!( + "{context} only supports float32 or float64 dtypes", + ))), + } +} + +fn ensure_valid_fan_shape(shape: &Shape, context: &str) -> PyResult<()> { + if shape.dims().contains(&0) { + Err(PyValueError::new_err(format!( + "{context} requires all shape dimensions to be at least 1", + ))) + } else { + Ok(()) + } +} + +pub(crate) fn create_fan_init_tensor( + shape: Shape, + dtype: DataType, + device: Device, + requires_grad: bool, + kind: FanInitKind, + context: &str, +) -> PyResult { + ensure_float_dtype(dtype, context)?; + ensure_valid_fan_shape(&shape, context)?; + let tensor = kind.apply(shape, dtype, device, requires_grad); + tensor.map_err(_convert_error) +} + +pub(crate) fn create_uniform_tensor( + shape: Shape, + dtype: DataType, + device: Device, + requires_grad: bool, + low: f64, + high: f64, +) -> PyResult { + if !low.is_finite() || !high.is_finite() { + return Err(PyValueError::new_err( + "uniform requires finite low and high values", + )); + } + + if high.partial_cmp(&low) != Some(Ordering::Greater) { + return Err(PyValueError::new_err( + "uniform requires high to be greater than low", + )); + } + + let tensor = nn::init::init_uniform(shape, low, high, dtype, device, requires_grad) + .map_err(_convert_error)?; + Ok(tensor) +} + +#[allow(clippy::too_many_arguments)] +pub(crate) fn create_truncated_normal_tensor( + shape: Shape, + dtype: DataType, + device: Device, + requires_grad: bool, + mean: f64, + std: f64, + lower: Option, + upper: Option, + context: &str, +) -> PyResult { + ensure_float_dtype(dtype, context)?; + + if !mean.is_finite() { + return Err(PyValueError::new_err(format!( + "{context} requires a finite mean", + ))); + } + + if !std.is_finite() || std <= 0.0 { + return Err(PyValueError::new_err(format!( + "{context} requires std to be a positive finite value", + ))); + } + + let default_lower = mean - 2.0 * std; + let default_upper = mean + 2.0 * std; + let lower = lower.unwrap_or(default_lower); + let upper = upper.unwrap_or(default_upper); + + if lower.is_nan() || upper.is_nan() { + return Err(PyValueError::new_err(format!( + "{context} requires non-NaN bounds", + ))); + } + + if upper.partial_cmp(&lower) != Some(Ordering::Greater) { + return Err(PyValueError::new_err(format!( + "{context} requires upper bound to be greater than lower bound", + ))); + } + + let tensor = nn::init::truncated_normal_init( + shape, + mean, + std, + lower, + upper, + dtype, + device, + requires_grad, + ) + .map_err(_convert_error)?; + Ok(tensor) +} + +pub(crate) fn prepare_new_tensor_from_existing( + source: &Tensor, + dtype: DataType, + device: Device, + requires_grad: bool, +) -> PyResult { + let mut tensor = source.detach(); + + if tensor.device() != device { + tensor = tensor.to(device).map_err(_convert_error)?; + } + + if tensor.dtype() != dtype { + tensor = tensor.astype(dtype).map_err(_convert_error)?; + } + + tensor = tensor.deep_clone().map_err(_convert_error)?; + + if requires_grad { + tensor = tensor.requires_grad_(true); + } + + Ok(tensor) +} + +pub(crate) fn adapt_tensor_for_as_tensor( + source: &Tensor, + dtype: DataType, + device: Device, + requires_grad: bool, + copy: bool, +) -> PyResult { + if !copy + && source.dtype() == dtype + && source.device() == device + && source.requires_grad() == requires_grad + { + return Ok(source.clone()); + } + + let mut tensor = if copy || (source.requires_grad() && !requires_grad) { + source.detach() + } else { + source.clone() + }; + + if tensor.device() != device { + tensor = tensor.to(device).map_err(_convert_error)?; + } + + if tensor.dtype() != dtype { + tensor = tensor.astype(dtype).map_err(_convert_error)?; + } + + if copy { + tensor = tensor.deep_clone().map_err(_convert_error)?; + } + + if tensor.requires_grad() != requires_grad { + tensor = tensor.requires_grad_(requires_grad); + } + + Ok(tensor) +} + +pub(crate) fn create_randint_tensor( + shape: Shape, + dtype: DataType, + device: Device, + requires_grad: bool, + low: i64, + high: i64, +) -> PyResult { + let tensor_data = match dtype { + DataType::Int32 => { + let low_i32 = i32::try_from(low) + .map_err(|_| PyValueError::new_err("low is out of range for dtype int32"))?; + let high_i32 = i32::try_from(high) + .map_err(|_| PyValueError::new_err("high is out of range for dtype int32"))?; + if low_i32 >= high_i32 { + return Err(PyValueError::new_err( + "randint requires that low < high after casting to int32", + )); + } + let mut values = vec![0i32; shape.numel()]; + random::with_rng(|rng| { + use rand::RngExt; + for value in &mut values { + *value = rng.random_range(low_i32..high_i32); + } + }); + TensorData::from_vec_i32(values, device) + } + DataType::Int64 => { + if high <= low { + return Err(PyValueError::new_err("randint requires that low < high")); + } + let mut values = vec![0i64; shape.numel()]; + random::with_rng(|rng| { + use rand::RngExt; + for value in &mut values { + *value = rng.random_range(low..high); + } + }); + TensorData::from_vec_i64(values, device) + } + _ => { + return Err(PyValueError::new_err( + "randint only supports int32 or int64 dtypes", + )); + } + }; + + Ok(Tensor::new( + Arc::new(tensor_data), + shape, + dtype, + device, + requires_grad, + )) +} + +pub(crate) fn create_randperm_tensor( + n: usize, + dtype: DataType, + device: Device, + requires_grad: bool, +) -> PyResult { + let tensor_data = match dtype { + DataType::Int32 => { + let _ = i32::try_from(n).map_err(|_| { + PyValueError::new_err("randperm with dtype int32 requires n <= i32::MAX") + })?; + let mut values = Vec::with_capacity(n); + for idx in 0..n { + values.push(i32::try_from(idx).map_err(|_| { + PyValueError::new_err("randperm with dtype int32 requires n <= i32::MAX") + })?); + } + random::with_rng(|rng| { + use rand::seq::SliceRandom; + values.shuffle(rng); + }); + TensorData::from_vec_i32(values, device) + } + DataType::Int64 => { + let _ = i64::try_from(n).map_err(|_| { + PyValueError::new_err("randperm with dtype int64 requires n <= i64::MAX") + })?; + let mut values = Vec::with_capacity(n); + for idx in 0..n { + values.push(idx as i64); + } + random::with_rng(|rng| { + use rand::seq::SliceRandom; + values.shuffle(rng); + }); + TensorData::from_vec_i64(values, device) + } + _ => { + return Err(PyValueError::new_err( + "randperm only supports int32 or int64 dtypes", + )); + } + }; + + Ok(Tensor::new( + Arc::new(tensor_data), + Shape::new(vec![n]), + dtype, + device, + requires_grad, + )) +} + +pub(crate) fn create_eye_tensor( + n: usize, + m: usize, + dtype: DataType, + device: Device, + requires_grad: bool, +) -> PyResult { + let shape = Shape::new(vec![n, m]); + let mut tensor_data = TensorData::zeros_on_device(shape.numel(), dtype, device); + + match dtype { + DataType::Float32 => { + if let Some(slice) = tensor_data.as_f32_slice_mut() { + for i in 0..n.min(m) { + slice[i * m + i] = 1.0; + } + } + } + DataType::Float64 => { + if let Some(slice) = tensor_data.as_f64_slice_mut() { + for i in 0..n.min(m) { + slice[i * m + i] = 1.0; + } + } + } + DataType::Int32 => { + if let Some(slice) = tensor_data.as_i32_slice_mut() { + for i in 0..n.min(m) { + slice[i * m + i] = 1; + } + } + } + DataType::Int64 => { + if let Some(slice) = tensor_data.as_i64_slice_mut() { + for i in 0..n.min(m) { + slice[i * m + i] = 1; + } + } + } + DataType::Bool => { + if let Some(slice) = tensor_data.as_bool_slice_mut() { + for i in 0..n.min(m) { + slice[i * m + i] = true; + } + } + } + } + + Ok(Tensor::new( + Arc::new(tensor_data), + shape, + dtype, + device, + requires_grad, + )) +} + +pub(crate) fn create_full_tensor( + shape: Vec, + fill_value: f64, + dtype: DataType, + device: Device, + requires_grad: bool, +) -> PyResult { + let shape = Shape::new(shape); + let mut tensor_data = TensorData::uninitialized_on_device(shape.numel(), dtype, device); + + match dtype { + DataType::Float32 => { + if let Some(slice) = tensor_data.as_f32_slice_mut() { + slice.fill(fill_value as f32); + } + } + DataType::Float64 => { + if let Some(slice) = tensor_data.as_f64_slice_mut() { + slice.fill(fill_value); + } + } + DataType::Int32 => { + if let Some(slice) = tensor_data.as_i32_slice_mut() { + slice.fill(fill_value as i32); + } + } + DataType::Int64 => { + if let Some(slice) = tensor_data.as_i64_slice_mut() { + slice.fill(fill_value as i64); + } + } + DataType::Bool => { + if let Some(slice) = tensor_data.as_bool_slice_mut() { + slice.fill(fill_value != 0.0); + } + } + } + + Ok(Tensor::new( + Arc::new(tensor_data), + shape, + dtype, + device, + requires_grad, + )) +} +pub(crate) fn create_arange_tensor( + start: f64, + end: f64, + step: f64, + dtype: DataType, + device: Device, + requires_grad: bool, +) -> PyResult { + if step == 0.0 { + return Err(PyErr::new::( + "Step cannot be zero", + )); + } + + let num_elements = ((end - start) / step).ceil() as usize; + let shape = Shape::new(vec![num_elements]); + let mut tensor_data = TensorData::uninitialized_on_device(shape.numel(), dtype, device); + + match dtype { + DataType::Float32 => { + if let Some(slice) = tensor_data.as_f32_slice_mut() { + for (i, val) in slice.iter_mut().enumerate() { + *val = (start + i as f64 * step) as f32; + } + } + } + DataType::Float64 => { + if let Some(slice) = tensor_data.as_f64_slice_mut() { + for (i, val) in slice.iter_mut().enumerate() { + *val = start + i as f64 * step; + } + } + } + DataType::Int32 => { + if let Some(slice) = tensor_data.as_i32_slice_mut() { + for (i, val) in slice.iter_mut().enumerate() { + *val = (start + i as f64 * step) as i32; + } + } + } + DataType::Int64 => { + if let Some(slice) = tensor_data.as_i64_slice_mut() { + for (i, val) in slice.iter_mut().enumerate() { + *val = (start + i as f64 * step) as i64; + } + } + } + DataType::Bool => { + if let Some(slice) = tensor_data.as_bool_slice_mut() { + for (i, val) in slice.iter_mut().enumerate() { + *val = (start + i as f64 * step) != 0.0; + } + } + } + } + + Ok(Tensor::new( + Arc::new(tensor_data), + shape, + dtype, + device, + requires_grad, + )) +} + +pub(crate) fn create_linspace_tensor( + start: f64, + end: f64, + steps: usize, + dtype: DataType, + device: Device, + requires_grad: bool, +) -> PyResult { + if steps == 0 { + return Err(PyErr::new::( + "Number of steps must be positive", + )); + } + + let shape = Shape::new(vec![steps]); + let mut tensor_data = TensorData::uninitialized_on_device(shape.numel(), dtype, device); + let denom = if steps > 1 { (steps - 1) as f64 } else { 1.0 }; + let step = if steps > 1 { + (end - start) / denom + } else { + 0.0 + }; + + match dtype { + DataType::Float32 => { + if let Some(slice) = tensor_data.as_f32_slice_mut() { + for (i, val) in slice.iter_mut().enumerate() { + let value = if steps == 1 { + start + } else { + start + i as f64 * step + }; + *val = value as f32; + } + } + } + DataType::Float64 => { + if let Some(slice) = tensor_data.as_f64_slice_mut() { + for (i, val) in slice.iter_mut().enumerate() { + let value = if steps == 1 { + start + } else { + start + i as f64 * step + }; + *val = value; + } + } + } + DataType::Int32 => { + if let Some(slice) = tensor_data.as_i32_slice_mut() { + for (i, val) in slice.iter_mut().enumerate() { + let value = if steps == 1 { + start + } else { + start + i as f64 * step + }; + *val = value.round() as i32; + } + } + } + DataType::Int64 => { + if let Some(slice) = tensor_data.as_i64_slice_mut() { + for (i, val) in slice.iter_mut().enumerate() { + let value = if steps == 1 { + start + } else { + start + i as f64 * step + }; + *val = value.round() as i64; + } + } + } + DataType::Bool => { + if let Some(slice) = tensor_data.as_bool_slice_mut() { + for (i, val) in slice.iter_mut().enumerate() { + let value = if steps == 1 { + start + } else { + start + i as f64 * step + }; + *val = value != 0.0; + } + } + } + } + + Ok(Tensor::new( + Arc::new(tensor_data), + shape, + dtype, + device, + requires_grad, + )) +} + +pub(crate) fn create_logspace_tensor( + start: f64, + end: f64, + steps: usize, + base: f64, + dtype: DataType, + device: Device, + requires_grad: bool, +) -> PyResult { + if steps == 0 { + return Err(PyErr::new::( + "Number of steps must be positive", + )); + } + + if base <= 0.0 { + return Err(PyErr::new::( + "Base must be positive", + )); + } + + let shape = Shape::new(vec![steps]); + let mut tensor_data = TensorData::uninitialized_on_device(shape.numel(), dtype, device); + let denom = if steps > 1 { (steps - 1) as f64 } else { 1.0 }; + let step = if steps > 1 { + (end - start) / denom + } else { + 0.0 + }; + + match dtype { + DataType::Float32 => { + if let Some(slice) = tensor_data.as_f32_slice_mut() { + for (i, val) in slice.iter_mut().enumerate() { + let exponent = if steps == 1 { + start + } else { + start + i as f64 * step + }; + *val = base.powf(exponent) as f32; + } + } + } + DataType::Float64 => { + if let Some(slice) = tensor_data.as_f64_slice_mut() { + for (i, val) in slice.iter_mut().enumerate() { + let exponent = if steps == 1 { + start + } else { + start + i as f64 * step + }; + *val = base.powf(exponent); + } + } + } + DataType::Int32 => { + if let Some(slice) = tensor_data.as_i32_slice_mut() { + for (i, val) in slice.iter_mut().enumerate() { + let exponent = if steps == 1 { + start + } else { + start + i as f64 * step + }; + *val = base.powf(exponent).round() as i32; + } + } + } + DataType::Int64 => { + if let Some(slice) = tensor_data.as_i64_slice_mut() { + for (i, val) in slice.iter_mut().enumerate() { + let exponent = if steps == 1 { + start + } else { + start + i as f64 * step + }; + *val = base.powf(exponent).round() as i64; + } + } + } + DataType::Bool => { + if let Some(slice) = tensor_data.as_bool_slice_mut() { + for (i, val) in slice.iter_mut().enumerate() { + let exponent = if steps == 1 { + start + } else { + start + i as f64 * step + }; + *val = base.powf(exponent) != 0.0; + } + } + } + } + + Ok(Tensor::new( + Arc::new(tensor_data), + shape, + dtype, + device, + requires_grad, + )) +} diff --git a/docs/api_reference.md b/docs/api_reference.md index 488f971a..17bfa4e2 100644 --- a/docs/api_reference.md +++ b/docs/api_reference.md @@ -44,6 +44,10 @@ of convenience aliases. | `clear_autograd_graph()` | Clear the global autograd graph. | | `is_autograd_graph_consumed()` | Inspect whether a graph has been consumed. | | `mark_autograd_graph_consumed()` | Mark the current graph as consumed. | +| `no_grad()` | Context manager: disable gradient recording (results are detached leaves; nothing is saved for backward). | +| `enable_grad()` | Context manager: re-enable gradient recording inside a `no_grad()` block. | +| `is_grad_enabled()` | Query the thread-local gradient recording mode. | +| `set_grad_enabled(enabled)` | Set the gradient recording mode, returning the previous mode. | | `available_submodules()` | Return availability of optional submodules. | | `list_public_api()` | Return public API symbol lists by module. | | `api_summary()` | Return version and API counts by module. | @@ -128,7 +132,7 @@ plotting conventions; `indexing="ij"` preserves matrix-indexing order for all axes. Dense outputs are materialized broadcast grids, while `sparse=True` returns only reshaped coordinate vectors that can still broadcast together lazily inside later operations. Set `copy=True` when callers need storage -independent of the returned grid objects. Calling `meshgrid()` with no inputs returns `()`. +independent of the returned grid objects. Calling `meshgrid()` with no inputs returns `()`. Validation and edge cases: @@ -435,7 +439,7 @@ Behavior and validation: - Python scalars, Python sequences, NumPy arrays, and MiniTensor tensors are accepted for `other`; `input` should be a MiniTensor tensor or tensor wrapper, matching the rest of the tensor-centric functional binary helpers. -- Boolean inputs use logical OR for `maximum` and logical AND for `minimum`. +- Boolean inputs use logical OR for `maximum` and logical AND for `minimum`. - Floating-point NaNs are propagated when either operand at an element is NaN. - Incompatible shapes raise the normal MiniTensor shape/broadcasting error. diff --git a/docs/architecture_review.md b/docs/architecture_review.md new file mode 100644 index 00000000..305e8af6 --- /dev/null +++ b/docs/architecture_review.md @@ -0,0 +1,425 @@ +# Architecture Review and Refactoring Report + +This document records a full architectural assessment of minitensor (engine, +bindings, and Python layer), the problems identified, the changes implemented, +and a prioritized plan for the remaining work. It is intended as the living +reference for architectural decisions; update it as the items below are +addressed. + +## 1. Current architecture (assessment) + +minitensor is a three-layer system: + +```text +minitensor/ Pure-Python facade: re-exports, introspection helpers, + broadcasting utilities (~1.1k LOC) +bindings/ (PyO3) Python <-> Rust boundary: PyTensor, functional API, + nn/optim wrappers, NumPy interop (~13k LOC) +engine/ (Rust) Storage, ops, autograd, nn, optim, serialization, + plugins, backends (~57k LOC) +``` + +Key mechanics discovered during the review: + +- **Storage**: `TensorData` owns either a `Vec` (CPU) or a raw + device pointer, and is shared between tensors via `Arc`. + `Tensor` carries shape/strides/dtype/device plus autograd metadata + (`grad_fn`, `grad`, `tensor_id`). +- **Autograd**: a *thread-local global* `ComputationGraph` maps `TensorId -> + GraphNode`. Every differentiable op allocates an output tensor, attaches a + `*Backward` gradient function (which stores cloned operand tensors), and + registers the node with the global graph. `backward()` walks the graph; + gradients live in a map inside the graph and are read back via + `autograd::get_gradient` (optimizers, `.grad` in Python). +- **Views**: almost every op materializes its output (`transpose` copies, + `reshape` re-strides). Only `expand` produces true strided (stride-0) + views. The FFI boundary (`PyTensor::from_tensor`) force-materializes any + non-contiguous tensor before it reaches Python. +- **Execution**: CPU kernels are hand-written per dtype with rayon + parallelism and a SIMD fast path for f32/f64. GPU backends (CUDA, Metal, + OpenCL) exist behind cargo features but are not compiled by default, and + allocation silently falls back to CPU. +- **Python layer**: thin re-export surface plus API introspection helpers; + `nn`/`optim`/`functional` are implemented in Rust and exposed as modules. + +## 2. Problems, risks, and technical debt found + +Ordered roughly by severity. Items marked **[FIXED]** were addressed in this +refactor; the rest are documented with a migration path in section 7. + +1. **[FIXED] `Tensor::view`/`reshape` corrupted non-contiguous tensors.** + `view` re-strided unconditionally, so `expand(...).reshape(...)` at the + engine level produced a tensor whose shape claimed N elements over storage + holding fewer (verified: shape `[12]` over 3 stored elements). The Python + surface was protected only by a blanket copy in `PyTensor::from_tensor`. + `view` now rejects non-contiguous tensors; `reshape` (both the `Tensor` + method and the `shape_ops` op) materializes a contiguous copy first. +2. **[FIXED] Backward pass scaled with the whole recorded graph, not the + traced subgraph.** The graph topologically sorted *every node ever + recorded on the thread* on each backward call and iterated all of them. + Backward now plans only the subgraph reachable from the loss tensor. +3. **[FIXED] The full gradient map was cloned on every backward call** and + immediately discarded by every production caller. `backward` now returns + `Result<()>`; gradients stay in the graph store and are read individually. + Tests that want the full map use the new `backward_collect`. +4. **[FIXED] Gradient-kernel registration was suppressed by accident.** + During backward, ops run inside gradient functions (e.g. `matmul` with a + non-detached saved operand) attempt to register new autograd nodes. This + was prevented only because the graph's `RefCell` happened to be borrowed — + `add_to_graph` used `try_borrow_mut` and *silently dropped* registrations. + There is now an explicit thread-local grad-recording mode (`NoGradGuard`, + `is_grad_enabled`): the backward executor disables recording, and + `add_to_graph` is a loud `borrow_mut` otherwise. This also provides the + primitive for a future user-facing `no_grad()`. +5. **[FIXED] Saved tensors were held until the next optimizer step.** The + graph (including every activation captured by `*Backward` structs) was + only freed by `optimizer.step()`. After a non-retaining `backward()`, the + bindings now release the reachable interior nodes immediately + (`release_saved_subgraph`), which both frees memory earlier and makes the + "graph has been freed" error truthful. +6. **[FIXED] `TensorData` carried a vestigial manual reference count.** + `inc_ref`/`dec_ref` were never called outside their own tests, while + `Drop` only deallocated raw device buffers when the counter happened to + equal 1 — a leak trap. Sharing is `Arc`'s job; the counter is gone and + `Drop` unconditionally returns raw buffers to the allocator. +7. **[FIXED] Unsafe per-element parallel loops in gradient accumulation.** + `add_inplace` erased slices to raw `usize` addresses and indexed them from + a parallel loop. Replaced with safe chunked `par_chunks_mut`/`zip` loops + (`binary_assign_slices`), which keep bounds information visible to the + compiler and remove 10 `unsafe` blocks from the hottest accumulation path. +8. **[FIXED] `DivBackward` allocated a ones tensor and ran 6 kernels** to + compute two gradients. It now runs 4 kernels with no scratch ones tensor. +9. **[FIXED] `Tensor::data_mut` mutated shared `Arc` storage through a + raw-pointer cast** for all leaf tensors that require gradients — undefined + behavior under Rust's aliasing rules whenever the storage was shared. The + cast is gone: storage moved to `UnsafeCell` and `data_mut` returns a + `DataMut` token whose `Shared` variant routes in-place parameter updates + through interior mutability with a documented aliasing contract + (section 7, third change set). +10. **Thread-affine autograd.** The graph is thread-local: tensors created on + one thread silently lose their history on another. Single-threaded + Python use is safe (GIL), but the engine API allows cross-thread misuse + with silent wrong results. (Documented; see section 7.) +11. **[FIXED] `t.zero_grad()` on one tensor cleared every gradient on the + thread** (`autograd::zero_gradients()` wiped the global map). It now + removes only that tensor's gradient. +12. **Per-dtype kernel duplication.** Nearly every op repeats the same body + five times (f32/f64/i32/i64/bool) — thousands of duplicated lines (e.g. + `elementwise.rs` ~880 lines for 4 ops). A dtype-dispatch macro or generic + kernels would cut the ops layer dramatically. (Documented.) +13. **[PARTIALLY FIXED] Speculative/dead subsystems.** `operations/fusion.rs` + (623 lines, never called) is removed — a breaking change for the + `engine` crate only, documented in section 7. `hardware/` + (profiler/optimizer, ~2k LOC, referenced only by an example and a + compatibility test) and the pooled GPU memory manager remain, flagged + for the maintainer as a semver decision; `debug.rs` is exposed through + the Python API and stays. +14. **[FIXED] `#![allow(clippy::all)]` on the engine crate and no clippy in + CI** (lints.yml only ran rustfmt via pre-commit). The blanket allow is + gone, the workspace is warning-free, and CI now runs + `cargo clippy --workspace --all-targets -- -D warnings`. This also + surfaced and fixed real issues: UB-producing initializer kernels in + `nn/init.rs` and two deny-level raw-pointer lints in the memory stack. +15. **`include!`-based module layout** (`autograd/mod/*.rs`, + `tensor/mod/*.rs`, bindings `pytensor/*.rs`) merges many files into one + module, defeating visibility boundaries, slowing incremental compiles, + and confusing tools. (Documented.) +16. **[FIXED] Uninitialized `Vec` buffers** (`set_len` before write) were + relied on throughout. `nn/init.rs` no longer creates `&mut` slices over + uninitialized values, `bool` storage is always zero-initialized, and + `clone_data` copies safely. The op *output* buffers — which were the + last uninitialized path — are now zero-initialized too: allocating them + uninitialized and handing them out as `&mut [f32]`/`&mut [i64]`/… (which + every kernel does) forms a reference to invalid values, i.e. undefined + behavior, even though every bit pattern is a representable float/int. + Per the project's stated priority order (correctness before + performance), soundness wins. The measured cost is confined to + allocation-bound cases: interleaved A/B on release wheels shows ~25–35% + on the pure element-wise microbenchmark (`add`+`sum` over 2000×2000, + where the extra `memset` raises output write traffic from ~3N to ~4N) + and noise on the matmul-dominated training step. The zero-cost + alternative (`MaybeUninit`-typed kernel writes) is retained as tracked + future work in section 7 for anyone who wants the throughput back + without the unsoundness. +17. **[PARTIALLY FIXED] API-inconsistencies at the Python layer.** + `requires_grad_()` now returns `self` (chaining works, matching + PyTorch). Still open: `expand` materializes at the FFI boundary, so it + is semantically `repeat` with extra steps, and `get_gradient` can return + gradients for tensors that do not require grad. +18. **[FIXED] Unchecked shape products.** `Shape::numel` computed + `dims.iter().product()`, which wraps in release builds: an absurd shape + could report a small element count, under-allocate storage, and turn + later stride-based indexing into out-of-bounds access. `numel` and + `Strides::from_shape` now use checked multiplication in every profile. +19. **[FIXED] No inference mode.** Repeated forward passes with + `requires_grad` parameters grew the thread-local graph (and its saved + operands) without bound; there was no way to turn recording off. + `no_grad()` / `enable_grad()` now exist end to end (see section 7). + +## 3. Target architecture + +The refactor converges on the following boundaries (largely realized for the +autograd path in this change set): + +- **Storage layer** (`tensor::data`): dumb, `Arc`-shared buffers; no + reference counting of its own; eventual interior mutability + (`UnsafeCell`-backed cells behind a safe API) so parameter updates do not + need aliasing casts. +- **Tensor layer** (`tensor::mod`): shape/stride bookkeeping with *enforced* + invariants — a `Tensor`'s shape must always describe its storage; `view` + is contiguous-only, `reshape` materializes, `expand` is the only stride + producer until strided kernels exist. +- **Autograd layer** (`autograd`): tape building (`add_to_graph`, gated by an + explicit grad-mode), pass planning (`plan_backward`, reachable subgraph + only), pass execution (`execute_backward_plan`, borrow-free, grad-mode + off), and storage (`set_gradients`/`get_gradient`), with node release + decoupled from optimizer stepping (`release_saved_subgraph`). +- **Ops layer** (`operations`): kernels stay dtype-specialized but should be + generated (macro/generics) rather than hand-copied; broadcasting via + shape-level utilities; no direct graph manipulation beyond `add_to_graph`. +- **Boundary layer** (`bindings`): conversion, validation, and Python + semantics only; no correctness-critical patching of engine behavior (the + `from_tensor` materialization stays as defense-in-depth but is no longer + load-bearing). + +Trade-offs taken: the thread-local graph is retained (a process-global graph +with locking would penalize the single-threaded common case and PyO3 usage); +per-op saved-tensor `detach()` was preferred over reworking every gradient +function's captures; eager materialization remains the norm because the +kernel layer indexes storage in logical order — introducing lazy views +engine-wide is a larger project listed below. + +## 4. Recommended folder/module structure + +Realized now: no file moves (churn would obscure the behavioral fixes). +Recommended next structure, in order of value: + +```text +engine/src/ + autograd/ + graph.rs # tape + planning + execution (done, in place) + ops/ # one file per Backward family, real submodules (replace include!) + tensor/ + storage.rs # TensorData (rename of data.rs) + tensor.rs # Tensor core (replace include!-merged mod/) + ops/ + kernels/ # dtype-generic kernel bodies (macro-generated) + ... + # delete or feature-gate: hardware/, operations/fusion.rs, debug.rs +``` + +## 5. Component responsibilities and boundaries (after this change) + +| Component | Responsibility | May touch | +|---|---|---| +| `TensorData` | own one buffer, expose typed slices | allocator | +| `Tensor` | shape/stride invariants, COW mutation policy | `TensorData`, autograd registration | +| `ComputationGraph` | record nodes, plan/execute/store backward | gradient functions | +| `NoGradGuard` / grad mode | decide whether ops record | thread-local flag | +| ops (`operations::*`) | math + broadcasting + attaching `*Backward` | tensors, `add_to_graph` | +| optimizers | read grads (`get_gradient`), update params in place | tensors | +| bindings | conversion + Python semantics | public engine API only | + +## 6. Data / control / dependency flow + +Forward: Python call → binding extracts `Tensor` → op validates devices, +coerces dtypes, broadcasts shapes → kernel writes a fresh output buffer → if +grad mode is on and an input requires grad, a `*Backward` (with saved +operands) is attached and the node registered in the thread-local graph. + +Backward: `loss.backward()` → plan = reachable reverse-topological subgraph +(single borrow) → executor runs gradient functions with grad recording +disabled, accumulating into a local map (no graph borrow held) → map stored +in the graph → bindings release interior nodes (non-retaining case) → +optimizer reads `get_gradient(param)` and updates parameters in place → +`optimizer.step()` clears the graph. + +Dependencies flow one way: bindings → engine ops → tensor/storage; autograd +sits beside ops (ops register, autograd executes ops through trait objects). + +## 7. Prioritized remaining migration plan + +Completed in the second change set: + +- **User-facing `no_grad()` / `enable_grad()`** — gradient recording is now + gated centrally (`Tensor::new`, `Tensor::set_grad_fn`, `add_to_graph`), so + results inside `no_grad()` are detached leaves, nothing is saved for + backward, and inference no longer grows the graph. Exposed in Python as + `mt.no_grad()`, `mt.enable_grad()`, `mt.is_grad_enabled()`, + `mt.set_grad_enabled()`. One documented divergence from PyTorch: factory + functions called inside `no_grad` also produce non-grad tensors; use + `requires_grad_(True)` (which expresses explicit intent and is never + gated) to opt back in. +- **Per-tensor `zero_grad`** — `Tensor::zero_grad` now removes only its own + entry from the global gradient map (`autograd::clear_gradient`) instead of + wiping every gradient on the thread. +- **`requires_grad_` returns `self`** in Python, enabling + `mt.randn(...).requires_grad_(True)` chaining. +- **Clippy enabled** — `#![allow(clippy::all)]` removed; the workspace is + clean under `cargo clippy --workspace --all-targets -- -D warnings`, which + now runs in CI (lints workflow). Three style lints are allowed crate-wide + with documented justification (`needless_range_loop`, + `too_many_arguments`, `items_after_test_module` — the last is an artifact + of the `include!` layout). +- **Overflow-safe shape arithmetic** — `Shape::numel` and + `Strides::from_shape` use checked multiplication in all build profiles; a + wrapped element count could previously under-allocate storage while + indexing code still trusted the dimensions (out-of-bounds writes via the + raw-pointer kernels). +- **Undefined behavior removed from initializer kernels** — `nn/init.rs` + built `&mut` slices over uninitialized memory (instant UB for `bool`); + all twelve sites now write through `Vec::extend`. `clone_data` uses + `Vec::clone` instead of `set_len` + copy. + +Completed in the third change set: + +- **Interior mutability for parameter storage.** `TensorData`'s buffer now + lives in an `UnsafeCell`, and `Tensor::data_mut` returns a `DataMut` + access token instead of fabricating `&mut TensorData` from a shared + `Arc` (which was undefined behavior whenever any other handle existed — + i.e. always). `DataMut::Unique` is ordinary exclusive access (after + copy-on-write); `DataMut::Shared` is the one documented exception — + in-place parameter updates that must stay visible through every handle — + and routes through UnsafeCell-backed accessors with an explicit aliasing + contract. Call sites did not change shape (`t.data_mut().as_…_slice_mut()` + still works; the token consumes itself and returns a slice borrowing from + the tensor). +- **dtype-dispatch macro, first application.** The ten hand-copied typed + slice accessors (plus the five new shared-mutation variants) are + generated by one `typed_slice_accessors!` macro — the template for + collapsing the per-dtype duplication in the ops layer. +- **Frozen inputs skip gradient work.** `AddBackward`/`SubBackward`/ + `MulBackward`/`DivBackward` now carry per-input `requires_grad` flags + (as the min/max/where/matmul functions already did) and skip the entire + gradient chain for inputs that do not require gradients. This both fixes + `get_gradient` returning gradients for non-grad tensors (root cause, for + arithmetic ops) and measurably speeds up graphs with frozen operands: + the wide fan-out benchmark (`x * 0.5` per node — a frozen scalar that + previously got a full-size multiply plus a reduce-to-scalar every + backward) improved 15–30% in interleaved A/B runs. + +Completed in the fourth change set: + +- **Arithmetic kernel dedup.** The eighteen hand-copied `*_direct` binary + kernels (add/sub/mul/div × dtypes) are generated by two macros + (`binary_kernel!` / `binary_kernel_simd!`), and `neg`'s five dtype arms + by one — ~440 lines of copy-paste removed with byte-identical behavior + (interleaved A/B benchmarks show exact parity). This is the pattern to + roll out across the rest of `operations/`. +- **Dead `fusion` subsystem removed** (623 lines). It was compiled and + re-exported (`engine::operations::fusion::*`) but never called by any + execution path, the bindings, or the Python API. **Breaking change for + the `engine` crate only** (the Python package surface is unchanged): + anything that depended on `engine::operations::fusion` should pin + `engine 0.2.1` or vendor the module from git history. + +Completed in the fifth change set: + +- **Loss functions no longer compute gradients for frozen targets.** All + seven loss backward functions (MSE, MAE, Huber, CrossEntropy, BCE, + KLDiv, Focal) carry per-input `requires_grad` flags. Previously + MSE/MAE/Huber allocated and negated a full-size target gradient on + *every training step*, and KLDiv ran an entire `log`/`sub`/`add`/`mul` + chain for the target distribution — all discarded work whenever targets + are constants (the overwhelmingly common case). Prediction gradients are + likewise skipped when predictions are frozen. Regression test: + `test_loss_targets_receive_no_gradient`. + +Completed in the sixth change set: + +- **Gating audit finished.** The last multi-input gradient functions + without per-input `requires_grad` flags — `LogAddExpBackward` (which ran + an `exp`/`sub`/`mul`/reduce chain per side unconditionally) and + `ConcatBackward` (which extracted a gradient slice per input) — are now + gated. `PowBackward` and `LayerNormBackward` already had flags; + single-input functions need none (they are only reached when their + input requires grad). Regression test: + `test_concat_frozen_inputs_receive_no_gradient`. +- **Comparison kernels deduplicated.** The five hand-copied `cmp_*` + helpers collapsed into one `cmp_kernel!` macro, continuing the pattern + from the storage accessors and arithmetic kernels. + +Completed in the seventh change set: + +- **Autograd converted from `include!` to real modules.** The seven + merged files under `autograd/mod/` are now proper submodules with their + own imports and explicit `pub(crate)` boundaries for shared helpers, + re-exported through `autograd/mod.rs` so every `crate::autograd::X` + path is unchanged. The conversion also surfaced layout artifacts the + merged namespace had been hiding: the shared broadcast-reduction helper + lived in the *tests* include file (now in `core`), and + `PowBackward`'s trait impl lives in a different file than its struct. + This is the template for converting the remaining `include!` clusters. + +Still open, in priority order: + +1. **dtype dispatch macro for the remaining ops files** — the uniform + float-unary activation kernels in `activation/hyperbolic.rs` are done + (34 fetch/`unary_apply` wrappers collapsed into one + `float_unary_kernel!` macro parameterised by the mapping closure, ~300 + lines removed; closures were extracted verbatim, so behavior is + byte-identical, confirmed by the suite plus a numerical spot-check). + Reduction kernels and the non-uniform activation kernels (softplus/ + gelu/elu, which take extra parameters) still carry per-dtype copies. +2. **`include!` layout migration — complete.** Every `include!` cluster in + the codebase is now real modules: `autograd`, `tensor`, all seven + `operations` clusters (27 files), the bindings (`tensor.rs`: 19 files + across `pytensor`/`python`/`creation`; `nn.rs`: 2), and the + feature-gated `backends/opencl` pair (validated under `--features + opencl`, which does build in this environment; the conversion also + fixed pre-existing latent test bugs — a private-field access and a + missing `CL_MEM_*` import — that had never compiled because CI only + builds default features). Both structural patterns are used depending on + privacy needs: + - **siblings + `pub(crate)` re-exports** where files only share + free functions (operations, autograd); + - **children-of-core** where files carry `impl` blocks that touch a + struct's private fields or private methods — `tensor`'s method files + are children of the `Tensor`-declaring module, `nn`'s `layers` + (pyclass impls + registration) is a child of `module`, and opencl's + `ops_impl` (`impl OpenCLOps`) is a child of `context` (the type + declarations). This preserved every field's privacy; only genuinely + cross-file helpers/methods were widened to `pub(crate)`. + The `items_after_test_module` crate-wide allow — an artifact of the old + layout — has been removed; clippy passes without it under `-D warnings`, + both for the default build and for `--features opencl`. +3. **Feature-gate or remove the remaining speculative subsystems** + (`hardware`, pooled allocator; `debug` is exposed to Python and stays) + — semver-major for the engine crate, maintainer's call. +4. **`MaybeUninit`-typed kernel output writes (optional throughput + recovery).** Op output buffers are now zero-initialized for soundness + (finding 16), which the A/B measured at ~25–35% on the allocation-bound + element-wise microbenchmark and noise on realistic training. A maintainer + who wants that throughput back without reintroducing the UB can thread + `MaybeUninit`-typed writes through the ~71 output-producing kernels; the + uniform write pattern makes this mechanical, but it is a large surface, + so it is left as opt-in future work rather than done under time pressure. + +## 8-11. Implementation, tests, validation, benchmarks + +See the commit(s) accompanying this document for the full implementation. +Validation performed: + +- Rust: full workspace suite (`cargo test --workspace --all-targets`), + including new regression tests for non-contiguous `view`/`reshape`, + reachable-subgraph planning, and post-backward node release. +- Python: full pytest suite (778 passed, 5 skipped) against a rebuilt + release wheel. +- Benchmarks: interleaved A/B on release wheels (baseline vs refactor + alternating on the same machine, 3 rounds, median of per-round best; + shared-cloud CPU, so treat small deltas as noise): + +| Benchmark | Baseline | Refactor | Delta | +|---|---|---|---| +| training step (3-layer MLP, batch 64) | 0.87 ms | 0.83 ms | ~ noise | +| deep chain backward (200 ops) | 36.0 ms | 36.6 ms | ~ noise | +| wide fan-out backward (50 reuses of one tensor) | 9.7 ms | 8.0 ms | **−15 %** | +| elementwise add+sum (2000×2000) | 2.44 ms | 2.35 ms | ~ noise | +| 30 forward passes, no backward | 6.26 ms | 5.87 ms | ~ noise | + +The fan-out case improves consistently in every round (9.3–10.0 ms vs +7.9–8.5 ms): gradient accumulation dominates there, which benefits from the +subgraph-scoped planning and the removal of the per-backward gradient-map +clone. The memory-side changes (saved tensors released right after +backward, no full-map clone) do not show up in wall-time microbenchmarks +but reduce peak footprint between `backward()` and `optimizer.step()`. diff --git a/docs/index.md b/docs/index.md index 28caf4b5..22d014ee 100644 --- a/docs/index.md +++ b/docs/index.md @@ -21,6 +21,7 @@ performance :caption: Contributor guides development +architecture_review ``` ## Start here @@ -33,6 +34,9 @@ development utilities. - [Development guide](development.md) -- repository layout, environment setup, validation commands, documentation workflow, and release checks. +- [Architecture review](architecture_review.md) -- assessment of the engine, + bindings, and Python layers, the autograd/storage refactors that came out of + it, and the prioritized plan for remaining architectural work. ## Feature guides diff --git a/engine/src/autograd/graph.rs b/engine/src/autograd/graph.rs index f1a79fc9..5917eebb 100644 --- a/engine/src/autograd/graph.rs +++ b/engine/src/autograd/graph.rs @@ -76,16 +76,69 @@ impl GraphNode { } } +/// A single step of a backward pass: the tensor to propagate from and the +/// gradient function (if any) that produces gradients for its inputs. +/// +/// Plans hold `Arc` clones of the gradient functions so they can be executed +/// without keeping the graph borrowed. This is what allows the thread-local +/// graph to stay accessible while user-visible gradient kernels run. +pub struct BackwardStep { + pub tensor_id: TensorId, + pub grad_fn: Option>, + pub requires_grad: bool, +} + +/// Execute a previously planned backward pass. +/// +/// `plan` must be in reverse-topological order (outputs before inputs), as +/// produced by [`ComputationGraph::plan_backward`]. The function is free of +/// any borrow of the graph itself, so it is safe to run while the global +/// graph is unlocked. +pub fn execute_backward_plan( + plan: &[BackwardStep], + start_tensor: TensorId, + gradient: Tensor, +) -> Result> { + let mut gradients: FxHashMap = FxHashMap::default(); + gradients.reserve(plan.len().max(1)); + gradients.insert(start_tensor, gradient); + + for step in plan { + // Skip nodes that never received a gradient (dead branches) and nodes + // that do not participate in differentiation. + if !step.requires_grad { + continue; + } + // Take the gradient for this node out of the map to avoid cloning. + if let Some(grad_output) = gradients.remove(&step.tensor_id) { + if let Some(grad_fn) = &step.grad_fn { + let input_grads = grad_fn.backward(&grad_output)?; + for (input_id, grad) in input_grads { + match gradients.entry(input_id) { + Entry::Occupied(mut e) => { + arithmetic::add_inplace(e.get_mut(), &grad)?; + } + Entry::Vacant(e) => { + e.insert(grad); + } + } + } + } + // Re-insert the gradient for this node so it remains available to + // callers via `get_gradient` after the pass completes. + gradients.insert(step.tensor_id, grad_output); + } + } + + Ok(gradients) +} + /// Computation graph for automatic differentiation pub struct ComputationGraph { /// Nodes in the graph nodes: FxHashMap, - /// Topological ordering for backward pass - topological_order: Vec, /// Gradients computed during backward pass gradients: FxHashMap, - /// Whether the cached topological order is stale - needs_order_update: bool, } impl ComputationGraph { @@ -93,9 +146,7 @@ impl ComputationGraph { pub fn new() -> Self { Self { nodes: FxHashMap::default(), - topological_order: Vec::new(), gradients: FxHashMap::default(), - needs_order_update: false, } } @@ -121,7 +172,6 @@ impl ComputationGraph { let node = GraphNode::new(tensor_id, grad_fn, requires_grad); self.nodes.insert(tensor_id, node); - self.needs_order_update = true; } /// Add a named tensor to the computation graph (for debugging) @@ -134,97 +184,130 @@ impl ComputationGraph { ) { let node = GraphNode::new(tensor_id, grad_fn, requires_grad).with_name(name); self.nodes.insert(tensor_id, node); - self.needs_order_update = true; } - /// Update the topological ordering of nodes using a non-recursive DFS. - /// This avoids building explicit dependent lists and minimizes allocations - /// compared to Kahn's algorithm while still detecting cycles. - fn update_topological_order(&mut self) { - self.topological_order.clear(); - self.topological_order.reserve(self.nodes.len()); + /// Visit the subgraph reachable from `start` through input edges, in + /// reverse-topological order (outputs before inputs), calling `visit` for + /// each reachable node id. Detects cycles and reports them as errors. + fn visit_reachable_reverse_topo( + &self, + start: TensorId, + mut visit: impl FnMut(&GraphNode), + ) -> Result<()> { + if !self.nodes.contains_key(&start) { + return Ok(()); + } // 0 = unvisited, 1 = visiting, 2 = visited let mut state: FxHashMap = FxHashMap::default(); - state.reserve(self.nodes.len()); - let mut stack: Vec<(TensorId, usize)> = Vec::with_capacity(self.nodes.len()); - let mut has_cycle = false; - - for &node_id in self.nodes.keys() { - if state.get(&node_id).copied().unwrap_or(0) != 0 { - continue; - } - stack.push((node_id, 0)); - while let Some((current_id, idx)) = stack.last_mut() { - match state.get(current_id).copied().unwrap_or(0) { - 0 => { - state.insert(*current_id, 1); - } - 1 => {} - 2 => { - stack.pop(); - continue; - } - _ => unreachable!(), + let mut stack: Vec<(TensorId, usize)> = Vec::new(); + // Post-order collection (inputs before outputs); reversed at the end. + let mut post_order: Vec = Vec::new(); + + stack.push((start, 0)); + while let Some((current_id, idx)) = stack.last_mut() { + match state.get(current_id).copied().unwrap_or(0) { + 0 => { + state.insert(*current_id, 1); } + 1 => {} + 2 => { + stack.pop(); + continue; + } + _ => unreachable!(), + } - let inputs = if let Some(node) = self.nodes.get(current_id) { - &node.inputs - } else { + let inputs = match self.nodes.get(current_id) { + Some(node) => &node.inputs, + None => { state.insert(*current_id, 2); stack.pop(); continue; - }; + } + }; - if *idx < inputs.len() { - let next = inputs[*idx]; - *idx += 1; - if !self.nodes.contains_key(&next) { - continue; - } - match state.get(&next).copied().unwrap_or(0) { - 0 => stack.push((next, 0)), - 1 => { - has_cycle = true; - break; - } - 2 => {} - _ => unreachable!(), + if *idx < inputs.len() { + let next = inputs[*idx]; + *idx += 1; + if !self.nodes.contains_key(&next) { + continue; + } + match state.get(&next).copied().unwrap_or(0) { + 0 => stack.push((next, 0)), + 1 => { + return Err(crate::error::MinitensorError::gradient_error( + "Computation graph contains cycles", + )); } - } else { - state.insert(*current_id, 2); - self.topological_order.push(*current_id); - stack.pop(); + 2 => {} + _ => unreachable!(), } - } - if has_cycle { - break; + } else { + state.insert(*current_id, 2); + post_order.push(*current_id); + stack.pop(); } } - if has_cycle { - self.topological_order.clear(); - } else { - // We collected nodes from inputs to outputs; reverse for backward pass - self.topological_order.reverse(); + for id in post_order.into_iter().rev() { + if let Some(node) = self.nodes.get(&id) { + visit(node); + } } - - self.needs_order_update = false; + Ok(()) } - fn ensure_topological_order(&mut self) { - if self.needs_order_update { - self.update_topological_order(); + /// Build an executable backward plan for the subgraph reachable from + /// `start_tensor`. Only reachable nodes are visited, so the cost of a + /// backward pass is proportional to the size of the traced subgraph, not + /// to every tensor ever recorded on this thread. + pub fn plan_backward(&self, start_tensor: TensorId) -> Result> { + let mut plan = Vec::new(); + self.visit_reachable_reverse_topo(start_tensor, |node| { + plan.push(BackwardStep { + tensor_id: node.tensor_id, + grad_fn: node.grad_fn.clone(), + requires_grad: node.requires_grad, + }); + })?; + Ok(plan) + } + + /// Replace the stored gradient map with the results of a backward pass. + pub fn set_gradients(&mut self, gradients: FxHashMap) { + self.gradients = gradients; + } + + /// Clone the stored gradient map. Intended for tests and diagnostics; the + /// training path reads individual gradients via [`Self::get_gradient`]. + pub fn gradients_snapshot(&self) -> FxHashMap { + self.gradients.clone() + } + + /// Drop every node reachable from `start` that carries a gradient + /// function. This releases the tensors captured for backward (activations, + /// saved operands) as soon as the pass is finished, instead of holding + /// them until the next optimizer step. Leaf nodes and stored gradients are + /// preserved so `get_gradient` keeps working. + pub fn release_saved_subgraph(&mut self, start: TensorId) { + let mut interior: Vec = Vec::new(); + // Ignore cycle errors here: releasing is best-effort cleanup. + let _ = self.visit_reachable_reverse_topo(start, |node| { + if node.grad_fn.is_some() { + interior.push(node.tensor_id); + } + }); + for id in interior { + self.nodes.remove(&id); } } - /// Perform backward pass from a given tensor - pub fn backward( - &mut self, - start_tensor: TensorId, - gradient: Option, - ) -> Result> { - self.ensure_topological_order(); + /// Perform backward pass from a given tensor. + /// + /// Gradients are stored in the graph and can be queried afterwards with + /// [`Self::get_gradient`]; nothing is cloned on the hot path. + pub fn backward(&mut self, start_tensor: TensorId, gradient: Option) -> Result<()> { if !self.nodes.contains_key(&start_tensor) && crate::autograd::is_graph_consumed() { return Err( crate::error::MinitensorError::gradient_error_with_suggestion( @@ -234,55 +317,21 @@ impl ComputationGraph { ), ); } - // Clear previous gradients and ensure sufficient capacity for accumulation - self.gradients.clear(); - self.gradients.reserve(self.nodes.len()); - // Set the initial gradient - if let Some(grad) = gradient { - self.gradients.insert(start_tensor, grad); - } else { - return Err(crate::error::MinitensorError::gradient_error( - "Initial gradient must be provided", - )); - } - - // Process nodes in topological order (outputs to inputs) - for &node_id in &self.topological_order { - if let Some(node) = self.nodes.get(&node_id) { - // Skip if this node doesn't require gradients - if !node.requires_grad { - continue; - } - - // Take the gradient for this node out of the map to avoid cloning - if let Some(grad_output) = self.gradients.remove(&node_id) { - // If this node has a gradient function, compute gradients for inputs - if let Some(grad_fn) = &node.grad_fn { - let input_grads = grad_fn.backward(&grad_output)?; - for (input_id, grad) in input_grads { - match self.gradients.entry(input_id) { - Entry::Occupied(mut e) => { - arithmetic::add_inplace(e.get_mut(), &grad)?; - } - Entry::Vacant(e) => { - e.insert(grad); - } - } - } - } - // Re-insert the gradient for this node so it is available in the - // final gradient map returned to callers - self.gradients.insert(node_id, grad_output); - } - } - } + let gradient = gradient.ok_or_else(|| { + crate::error::MinitensorError::gradient_error("Initial gradient must be provided") + })?; - Ok(self.gradients.clone()) + let plan = self.plan_backward(start_tensor)?; + // Gradient kernels must not record new autograd nodes; make that + // explicit instead of relying on the graph being borrowed. + let _guard = crate::autograd::NoGradGuard::new(); + self.gradients = execute_backward_plan(&plan, start_tensor, gradient)?; + Ok(()) } /// Validate the computation graph for correctness - pub fn validate(&mut self) -> Result<()> { + pub fn validate(&self) -> Result<()> { // Check for cycles if self.has_cycles() { return Err(crate::error::MinitensorError::gradient_error( @@ -315,6 +364,11 @@ impl ComputationGraph { self.gradients.clear(); } + /// Remove the stored gradient for a single tensor, if any. + pub fn remove_gradient(&mut self, tensor_id: TensorId) -> Option { + self.gradients.remove(&tensor_id) + } + /// Get the number of nodes in the graph pub fn num_nodes(&self) -> usize { self.nodes.len() @@ -325,10 +379,66 @@ impl ComputationGraph { self.nodes.contains_key(&tensor_id) } - /// Get the topological order of tensor IDs - pub fn topological_order(&mut self) -> &[TensorId] { - self.ensure_topological_order(); - &self.topological_order + /// Compute a reverse-topological order (outputs before inputs) of the + /// whole graph. Diagnostic API: the backward pass itself only orders the + /// reachable subgraph via [`Self::plan_backward`]. Returns an empty vector + /// if the graph contains cycles. + pub fn topological_order(&self) -> Vec { + // 0 = unvisited, 1 = visiting, 2 = visited + let mut state: FxHashMap = FxHashMap::default(); + state.reserve(self.nodes.len()); + let mut stack: Vec<(TensorId, usize)> = Vec::with_capacity(self.nodes.len()); + let mut post_order: Vec = Vec::with_capacity(self.nodes.len()); + + for &root in self.nodes.keys() { + if state.get(&root).copied().unwrap_or(0) != 0 { + continue; + } + stack.push((root, 0)); + while let Some((current_id, idx)) = stack.last_mut() { + match state.get(current_id).copied().unwrap_or(0) { + 0 => { + state.insert(*current_id, 1); + } + 1 => {} + 2 => { + stack.pop(); + continue; + } + _ => unreachable!(), + } + + let inputs = match self.nodes.get(current_id) { + Some(node) => &node.inputs, + None => { + state.insert(*current_id, 2); + stack.pop(); + continue; + } + }; + + if *idx < inputs.len() { + let next = inputs[*idx]; + *idx += 1; + if !self.nodes.contains_key(&next) { + continue; + } + match state.get(&next).copied().unwrap_or(0) { + 0 => stack.push((next, 0)), + 1 => return Vec::new(), // cycle + 2 => {} + _ => unreachable!(), + } + } else { + state.insert(*current_id, 2); + post_order.push(*current_id); + stack.pop(); + } + } + } + + post_order.reverse(); + post_order } /// Get a node by tensor ID @@ -346,27 +456,20 @@ impl ComputationGraph { } /// Check if there are cycles in the computation graph - pub fn has_cycles(&mut self) -> bool { - self.ensure_topological_order(); - self.topological_order.len() != self.nodes.len() + pub fn has_cycles(&self) -> bool { + self.topological_order().len() != self.nodes.len() } /// Remove a tensor and its dependencies from the graph pub fn remove_tensor(&mut self, tensor_id: TensorId) { if self.nodes.remove(&tensor_id).is_some() { - // Remove from topological order - self.topological_order.retain(|&id| id != tensor_id); - // Remove any gradients self.gradients.remove(&tensor_id); - - // Mark order as stale since dependencies changed - self.needs_order_update = true; } } /// Get statistics about the computation graph - pub fn stats(&mut self) -> GraphStats { + pub fn stats(&self) -> GraphStats { let leaf_nodes = self .nodes .values() @@ -442,6 +545,7 @@ mod tests { let add_fn = Arc::new(AddBackward { input_shapes: [vec![2, 3], vec![2, 3]], input_ids: [leaf1, leaf2], + input_requires_grad: [true, true], }); graph.add_tensor(result, Some(add_fn)); @@ -473,6 +577,7 @@ mod tests { let add_fn = Arc::new(AddBackward { input_shapes: [vec![2], vec![2]], input_ids: [a, b], + input_requires_grad: [true, true], }); graph.add_tensor(c, Some(add_fn)); @@ -508,6 +613,7 @@ mod tests { let add_fn = Arc::new(AddBackward { input_shapes: [vec![2], vec![2]], input_ids: [leaf1, leaf2], + input_requires_grad: [true, true], }); graph.add_tensor(result, Some(add_fn)); @@ -559,6 +665,7 @@ mod tests { let add_fn = Arc::new(AddBackward { input_shapes: [vec![2], vec![2]], input_ids: [leaf, leaf], // Self-dependency is ok + input_requires_grad: [true, true], }); graph.add_tensor(result, Some(add_fn)); assert!(graph.validate().is_ok()); @@ -573,10 +680,12 @@ mod tests { let add_a = Arc::new(AddBackward { input_shapes: [vec![1], vec![1]], input_ids: [b, b], + input_requires_grad: [true, true], }); let add_b = Arc::new(AddBackward { input_shapes: [vec![1], vec![1]], input_ids: [a, a], + input_requires_grad: [true, true], }); graph.add_tensor(a, Some(add_a)); @@ -606,6 +715,7 @@ mod tests { let add_fn = Arc::new(AddBackward { input_shapes: [vec![2], vec![2]], input_ids: [a, b], + input_requires_grad: [true, true], }); graph.add_tensor(c, Some(add_fn)); @@ -614,7 +724,8 @@ mod tests { let grad_c = Tensor::ones(grad_shape, DataType::Float32, Device::cpu(), false); // Perform backward pass - let gradients = graph.backward(c, Some(grad_c)).unwrap(); + graph.backward(c, Some(grad_c)).unwrap(); + let gradients = graph.gradients_snapshot(); // Should have gradients for a and b assert!(gradients.contains_key(&a)); @@ -643,6 +754,7 @@ mod tests { let add_fn1 = Arc::new(AddBackward { input_shapes: [vec![2], vec![2]], input_ids: [a, b], + input_requires_grad: [true, true], }); graph.add_tensor(temp1, Some(add_fn1)); @@ -651,6 +763,7 @@ mod tests { let add_fn2 = Arc::new(AddBackward { input_shapes: [vec![2], vec![2]], input_ids: [a, c], + input_requires_grad: [true, true], }); graph.add_tensor(temp2, Some(add_fn2)); @@ -659,6 +772,7 @@ mod tests { let add_fn3 = Arc::new(AddBackward { input_shapes: [vec![2], vec![2]], input_ids: [temp1, temp2], + input_requires_grad: [true, true], }); graph.add_tensor(d, Some(add_fn3)); @@ -667,7 +781,8 @@ mod tests { let grad_d = Tensor::ones(grad_shape, DataType::Float32, Device::cpu(), false); // Perform backward pass - let gradients = graph.backward(d, Some(grad_d)).unwrap(); + graph.backward(d, Some(grad_d)).unwrap(); + let gradients = graph.gradients_snapshot(); // Should have gradients for all tensors assert!(gradients.contains_key(&a)); @@ -675,7 +790,84 @@ mod tests { assert!(gradients.contains_key(&c)); // 'a' should appear in the gradients (accumulated from both paths) - assert!(gradients.get(&a).is_some()); + assert!(gradients.contains_key(&a)); + } + + #[test] + fn test_plan_backward_only_visits_reachable_subgraph() { + let mut graph = ComputationGraph::new(); + + // Graph 1: c = a + b + let a = TensorId::new(); + let b = TensorId::new(); + let c = TensorId::new(); + graph.add_tensor(a, None); + graph.add_tensor(b, None); + graph.add_tensor( + c, + Some(Arc::new(AddBackward { + input_shapes: [vec![2], vec![2]], + input_ids: [a, b], + input_requires_grad: [true, true], + })), + ); + + // Unrelated graph 2: z = x + y + let x = TensorId::new(); + let y = TensorId::new(); + let z = TensorId::new(); + graph.add_tensor(x, None); + graph.add_tensor(y, None); + graph.add_tensor( + z, + Some(Arc::new(AddBackward { + input_shapes: [vec![2], vec![2]], + input_ids: [x, y], + input_requires_grad: [true, true], + })), + ); + + let plan = graph.plan_backward(c).unwrap(); + let planned: Vec = plan.iter().map(|s| s.tensor_id).collect(); + assert_eq!(planned.len(), 3); + assert!(planned.contains(&a) && planned.contains(&b) && planned.contains(&c)); + assert!(!planned.contains(&x) && !planned.contains(&y) && !planned.contains(&z)); + // Reverse-topological: the output comes first. + assert_eq!(planned[0], c); + } + + #[test] + fn test_release_saved_subgraph_drops_interior_nodes_keeps_gradients() { + use crate::device::Device; + use crate::tensor::{DataType, Shape, Tensor}; + + let mut graph = ComputationGraph::new(); + let a = TensorId::new(); + let b = TensorId::new(); + let c = TensorId::new(); + graph.add_tensor(a, None); + graph.add_tensor(b, None); + graph.add_tensor( + c, + Some(Arc::new(AddBackward { + input_shapes: [vec![2], vec![2]], + input_ids: [a, b], + input_requires_grad: [true, true], + })), + ); + + let grad = Tensor::ones(Shape::new(vec![2]), DataType::Float32, Device::cpu(), false); + graph.backward(c, Some(grad)).unwrap(); + assert!(graph.get_gradient(a).is_some()); + + graph.release_saved_subgraph(c); + // Interior node (with grad_fn) removed, leaves preserved. + assert!(!graph.contains_tensor(c)); + assert!(graph.contains_tensor(a)); + assert!(graph.contains_tensor(b)); + // Gradients stay readable after release. + assert!(graph.get_gradient(a).is_some()); + assert!(graph.get_gradient(b).is_some()); } #[test] @@ -693,6 +885,7 @@ mod tests { let add_fn = Arc::new(AddBackward { input_shapes: [vec![2], vec![2]], input_ids: [a, b], + input_requires_grad: [true, true], }); graph.add_tensor(c, Some(add_fn)); diff --git a/engine/src/autograd/mod.rs b/engine/src/autograd/mod.rs index 6e351b97..61c35238 100644 --- a/engine/src/autograd/mod.rs +++ b/engine/src/autograd/mod.rs @@ -4,11 +4,33 @@ // This source code is licensed under the Apache-style license found in the // LICENSE file in the root directory of this source tree. +//! Automatic differentiation: tape building, backward planning/execution, +//! and the gradient functions recorded by tensor operations. +//! +//! The submodules group gradient functions by operation family; everything +//! public is re-exported here so callers keep using `crate::autograd::X`. + pub mod graph; -include!("mod/core.rs"); -include!("mod/arithmetic.rs"); -include!("mod/linalg.rs"); -include!("mod/shape.rs"); -include!("mod/reduction.rs"); -include!("mod/activation.rs"); -include!("mod/tests.rs"); + +#[path = "mod/activation.rs"] +mod activation; +#[path = "mod/arithmetic.rs"] +mod arithmetic; +#[path = "mod/core.rs"] +mod core; +#[path = "mod/linalg.rs"] +mod linalg; +#[path = "mod/reduction.rs"] +mod reduction; +#[path = "mod/shape.rs"] +mod shape; +#[cfg(test)] +#[path = "mod/tests.rs"] +mod tests; + +pub use self::activation::*; +pub use self::arithmetic::*; +pub use self::core::*; +pub use self::linalg::*; +pub use self::reduction::*; +pub use self::shape::*; diff --git a/engine/src/autograd/mod/activation.rs b/engine/src/autograd/mod/activation.rs index 703f186c..549d08ae 100644 --- a/engine/src/autograd/mod/activation.rs +++ b/engine/src/autograd/mod/activation.rs @@ -1,695 +1,751 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -fn repeat_interleave_backward_impl( - grad_output: &Tensor, - input_shape: &[usize], - repeats: &[usize], - dim: usize, -) -> Result { - if dim >= input_shape.len() { - return Err(MinitensorError::index_error( - dim as isize, - 0, - input_shape.len(), - )); - } - - let dim_size = input_shape[dim]; - if repeats.len() != dim_size { - return Err(MinitensorError::invalid_operation( - "repeat_interleave backward: repeats must match input dimension size".to_string(), - )); - } - - let grad_shape_vec = input_shape.to_vec(); - let grad_shape = Shape::new(grad_shape_vec.clone()); - let numel = grad_shape.numel(); - let dtype = grad_output.dtype(); - let device = grad_output.device(); - let total_repeats: usize = repeats.iter().sum(); - - let inner: usize = if dim + 1 >= input_shape.len() { - 1 - } else { - input_shape[dim + 1..].iter().product() - }; - let outer: usize = if dim == 0 { - 1 - } else { - input_shape[..dim].iter().product() - }; - - if numel == 0 || total_repeats == 0 || inner == 0 || outer == 0 { - return Ok(Tensor::zeros( - Shape::new(grad_shape_vec), - dtype, - device, - false, - )); - } - - let output_dims = grad_output.shape().dims(); - if output_dims.len() != input_shape.len() || output_dims[dim] != total_repeats { - return Err(MinitensorError::shape_mismatch( - input_shape.to_vec(), - output_dims.to_vec(), - )); - } - - macro_rules! repeat_interleave_backward_impl_inner { - ($ty:ty, $slice:ident, $from_vec:ident) => {{ - let src = grad_output.data().$slice().ok_or_else(|| { - MinitensorError::invalid_operation( - "repeat_interleave backward: gradient tensor must be contiguous".to_string(), - ) - })?; - let mut dst = vec![<$ty>::default(); numel]; - let chunk = total_repeats * inner; - dst.par_chunks_mut(dim_size * inner) - .enumerate() - .for_each(|(outer_idx, dst_chunk)| { - let mut src_offset = outer_idx * chunk; - for (i, &rep) in repeats.iter().enumerate() { - if rep == 0 { - continue; - } - let dst_start = i * inner; - let dst_slice = &mut dst_chunk[dst_start..dst_start + inner]; - for _ in 0..rep { - let src_slice = &src[src_offset..src_offset + inner]; - dst_slice.iter_mut().zip(src_slice.iter()).for_each( - |(dst_val, &src_val)| { - *dst_val += src_val; - }, - ); - src_offset += inner; - } - } - }); - TensorData::$from_vec(dst, device) - }}; - } - - let data = match dtype { - DataType::Float32 => { - repeat_interleave_backward_impl_inner!(f32, as_f32_slice, from_vec_f32) - } - DataType::Float64 => { - repeat_interleave_backward_impl_inner!(f64, as_f64_slice, from_vec_f64) - } - DataType::Int32 => repeat_interleave_backward_impl_inner!(i32, as_i32_slice, from_vec_i32), - DataType::Int64 => repeat_interleave_backward_impl_inner!(i64, as_i64_slice, from_vec_i64), - DataType::Bool => { - return Ok(Tensor::zeros(grad_shape, dtype, device, false)); - } - }; - - Ok(Tensor::new( - Arc::new(data), - grad_shape, - dtype, - device, - false, - )) -} - -/// Gradient function for expand operation which reduces broadcasted gradients -pub struct ExpandBackward { - pub input_shape: Vec, - pub input_id: TensorId, -} - -impl GradientFunction for ExpandBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - let shape = Shape::new(self.input_shape.clone()); - let grad_input = reduce_gradient_for_broadcasting(grad_output, &shape)?; - gradients.insert(self.input_id, grad_input); - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for MSE loss -pub struct MSELossBackward { - pub predictions_shape: Vec, - pub targets_shape: Vec, - pub input_ids: [TensorId; 2], - pub reduction: String, - pub diff: Tensor, -} - -impl GradientFunction for MSELossBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(2); - - // Base gradient: 2 * (predictions - targets) - let two = create_scalar_tensor(2.0, self.diff.dtype(), self.diff.device())?; - let mut base_grad = arithmetic::mul(&self.diff, &two)?; - - // Apply reduction scaling - match self.reduction.as_str() { - "mean" => { - let n = self.diff.numel() as f64; - let scale = create_scalar_tensor(1.0 / n, base_grad.dtype(), base_grad.device())?; - base_grad = arithmetic::mul(&base_grad, &scale)?; - } - "sum" | "none" => {} - _ => { - return Err(MinitensorError::gradient_error(format!( - "Unknown reduction mode: {}", - self.reduction - ))); - } - } - - // Multiply by upstream gradient - let pred_grad = arithmetic::mul(&base_grad, grad_output)?; - let target_grad = arithmetic::neg(&pred_grad)?; - - accumulate_grad(&mut gradients, self.input_ids[0], pred_grad)?; - accumulate_grad(&mut gradients, self.input_ids[1], target_grad)?; - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} - -/// Gradient function for MAE loss -pub struct MAELossBackward { - pub predictions_shape: Vec, - pub targets_shape: Vec, - pub input_ids: [TensorId; 2], - pub reduction: String, - pub sign: Tensor, -} - -impl GradientFunction for MAELossBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(2); - - let mut base_grad = self.sign.clone(); - match self.reduction.as_str() { - "mean" => { - let n = self.sign.numel() as f64; - let scale = create_scalar_tensor(1.0 / n, base_grad.dtype(), base_grad.device())?; - base_grad = arithmetic::mul(&base_grad, &scale)?; - } - "sum" | "none" => {} - _ => { - return Err(MinitensorError::gradient_error(format!( - "Unknown reduction mode: {}", - self.reduction - ))); - } - } - - let pred_grad = arithmetic::mul(&base_grad, grad_output)?; - let target_grad = arithmetic::neg(&pred_grad)?; - - accumulate_grad(&mut gradients, self.input_ids[0], pred_grad)?; - accumulate_grad(&mut gradients, self.input_ids[1], target_grad)?; - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} - -/// Gradient function for Huber loss -pub struct HuberLossBackward { - pub predictions_shape: Vec, - pub targets_shape: Vec, - pub input_ids: [TensorId; 2], - pub delta: f64, - pub reduction: String, - pub diff: Tensor, -} - -impl GradientFunction for HuberLossBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(2); - - let numel = self.diff.numel(); - let dtype = self.diff.dtype(); - let device = self.diff.device(); - let mut grad_data = TensorData::zeros_on_device(numel, dtype, device); - - match dtype { - DataType::Float32 => { - let diff_slice = self.diff.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from diff") - })?; - let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from grad") - })?; - let delta = self.delta as f32; - if numel < PAR_THRESHOLD { - for i in 0..numel { - let d = diff_slice[i]; - grad_slice[i] = if d.abs() <= delta { - d - } else { - delta * d.signum() - }; - } - } else { - let diff_ptr = diff_slice.as_ptr() as usize; - let grad_ptr = grad_slice.as_mut_ptr() as usize; - (0..numel).into_par_iter().for_each(|i| unsafe { - let diff_ptr = diff_ptr as *const f32; - let grad_ptr = grad_ptr as *mut f32; - let d = *diff_ptr.add(i); - *grad_ptr.add(i) = if d.abs() <= delta { - d - } else { - delta * d.signum() - }; - }); - } - } - DataType::Float64 => { - let diff_slice = self.diff.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from diff") - })?; - let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from grad") - })?; - if numel < PAR_THRESHOLD { - for i in 0..numel { - let d = diff_slice[i]; - grad_slice[i] = if d.abs() <= self.delta { - d - } else { - self.delta * d.signum() - }; - } - } else { - let diff_ptr = diff_slice.as_ptr() as usize; - let grad_ptr = grad_slice.as_mut_ptr() as usize; - let delta = self.delta; - (0..numel).into_par_iter().for_each(|i| unsafe { - let diff_ptr = diff_ptr as *const f64; - let grad_ptr = grad_ptr as *mut f64; - let d = *diff_ptr.add(i); - *grad_ptr.add(i) = if d.abs() <= delta { - d - } else { - delta * d.signum() - }; - }); - } - } - _ => { - return Err(MinitensorError::invalid_operation( - "Huber loss only supports floating point tensors", - )); - } - } - - let mut base_grad = Tensor::new( - Arc::new(grad_data), - Shape::new(self.predictions_shape.clone()), - dtype, - device, - false, - ); - - if self.reduction == "mean" { - let scale = create_scalar_tensor(1.0 / numel as f64, dtype, device)?; - base_grad = arithmetic::mul(&base_grad, &scale)?; - } - - let pred_grad = arithmetic::mul(&base_grad, grad_output)?; - let target_grad = arithmetic::neg(&pred_grad)?; - - accumulate_grad(&mut gradients, self.input_ids[0], pred_grad)?; - accumulate_grad(&mut gradients, self.input_ids[1], target_grad)?; - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} - -/// Gradient function for Cross Entropy loss -pub struct CrossEntropyLossBackward { - pub predictions_shape: Vec, - pub targets_shape: Vec, - pub input_ids: [TensorId; 2], - pub reduction: String, - pub softmax_predictions: Tensor, - pub targets: Tensor, -} - -impl GradientFunction for CrossEntropyLossBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // Compute base gradient: softmax(predictions) - targets - let mut base_grad = - arithmetic::sub(&self.softmax_predictions.detach(), &self.targets.detach())?; - - // Apply reduction scaling - match self.reduction.as_str() { - "mean" => { - let batch = self.targets_shape[0] as f64; - let mut scalar_data = - TensorData::zeros_on_device(1, base_grad.dtype(), base_grad.device()); - match base_grad.dtype() { - DataType::Float32 => { - let slice = scalar_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from scalar", - ) - })?; - slice[0] = (1.0 / batch) as f32; - } - DataType::Float64 => { - let slice = scalar_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from scalar", - ) - })?; - slice[0] = 1.0 / batch; - } - _ => { - return Err(MinitensorError::invalid_operation( - "CrossEntropy backward only supports floating point tensors", - )); - } - } - let scalar_tensor = Tensor::new( - Arc::new(scalar_data), - Shape::new(vec![1]), - base_grad.dtype(), - base_grad.device(), - false, - ); - base_grad = arithmetic::mul(&base_grad, &scalar_tensor)?; - } - "sum" | "none" => {} - _ => { - return Err(MinitensorError::gradient_error(format!( - "Unknown reduction mode: {}", - self.reduction - ))); - } - } - - // Multiply by upstream gradient (handles broadcasting) - let pred_grad = arithmetic::mul(&base_grad, grad_output)?; - - // Targets typically have no gradient - accumulate_grad(&mut gradients, self.input_ids[0], pred_grad)?; - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} - -/// Gradient function for Binary Cross Entropy loss -pub struct BCELossBackward { - pub predictions_shape: Vec, - pub targets_shape: Vec, - pub input_ids: [TensorId; 2], - pub reduction: String, - pub predictions: Tensor, - pub targets: Tensor, -} - -impl GradientFunction for BCELossBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // BCE gradient: (predictions - targets) / (predictions * (1 - predictions)) - let one = Tensor::ones( - Shape::new(self.predictions_shape.clone()), - self.predictions.dtype(), - self.predictions.device(), - false, - ); - let one_minus_pred = arithmetic::sub(&one, &self.predictions)?; - let numerator = arithmetic::sub(&self.predictions, &self.targets)?; - let denom = arithmetic::mul(&self.predictions, &one_minus_pred)?; - let mut base_grad = arithmetic::div(&numerator, &denom)?; - - if self.reduction == "mean" { - let n = self.predictions.numel() as f64; - let scale = create_scalar_tensor(1.0 / n, base_grad.dtype(), base_grad.device())?; - base_grad = arithmetic::mul(&base_grad, &scale)?; - } - - let pred_grad = arithmetic::mul(&base_grad, grad_output)?; - accumulate_grad(&mut gradients, self.input_ids[0], pred_grad)?; - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} - -/// Gradient function for KL Divergence loss -pub struct KLDivLossBackward { - pub predictions_shape: Vec, - pub targets_shape: Vec, - pub input_ids: [TensorId; 2], - pub reduction: String, - pub predictions: Tensor, - pub targets: Tensor, -} - -impl GradientFunction for KLDivLossBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(2); - - // Gradient w.r.t predictions: -(targets / predictions) - let mut pred_grad = arithmetic::div(&self.targets, &self.predictions)?; - pred_grad = arithmetic::neg(&pred_grad)?; - if self.reduction == "mean" { - let n = self.predictions.numel() as f64; - let scale = create_scalar_tensor(1.0 / n, pred_grad.dtype(), pred_grad.device())?; - pred_grad = arithmetic::mul(&pred_grad, &scale)?; - } - let pred_grad = arithmetic::mul(&pred_grad, grad_output)?; - accumulate_grad(&mut gradients, self.input_ids[0], pred_grad)?; - - // Gradient w.r.t targets: log(targets) - log(predictions) + 1 - let log_targets = activation::log(&self.targets)?; - let log_preds = activation::log(&self.predictions)?; - let diff = arithmetic::sub(&log_targets, &log_preds)?; - let one = Tensor::ones( - self.targets.shape().clone(), - self.targets.dtype(), - self.targets.device(), - false, - ); - let mut target_grad = arithmetic::add(&diff, &one)?; - if self.reduction == "mean" { - let n = self.predictions.numel() as f64; - let scale = create_scalar_tensor(1.0 / n, target_grad.dtype(), target_grad.device())?; - target_grad = arithmetic::mul(&target_grad, &scale)?; - } - let target_grad = arithmetic::mul(&target_grad, grad_output)?; - accumulate_grad(&mut gradients, self.input_ids[1], target_grad)?; - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} - -/// Gradient function for Focal loss -pub struct FocalLossBackward { - pub predictions_shape: Vec, - pub targets_shape: Vec, - pub input_ids: [TensorId; 2], - pub alpha: f64, - pub gamma: f64, - pub reduction: String, - pub softmax_predictions: Tensor, - pub targets: Tensor, -} - -impl GradientFunction for FocalLossBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // Exact gradient of FL = -alpha * (1 - p_t)^gamma * log(p_t) wrt the - // logits, where p_t is the true-class softmax probability: - // dFL/dz_j = alpha * (p_j - onehot_j) - // * (1 - p_t)^(gamma-1) * [ (1 - p_t) - gamma * p_t * ln(p_t) ] - // The modulating factor is a per-sample scalar (broadcast over classes). - let p = self.softmax_predictions.detach(); - let t = self.targets.detach(); - let dtype = p.dtype(); - let device = p.device(); - - // True-class probability per sample: p_t = sum(p * onehot) over classes. - let class_dim = (p.ndim() - 1) as isize; - let pt = reduction::sum(&arithmetic::mul(&p, &t)?, Some(vec![class_dim]), true)?; - - let one = create_scalar_tensor(1.0, dtype, device)?; - let one_minus_pt = arithmetic::sub(&one, &pt)?; - let log_pt = crate::operations::activation::log(&pt)?; - let gamma_scalar = create_scalar_tensor(self.gamma, dtype, device)?; - // bracket = (1 - p_t) - gamma * p_t * ln(p_t) - let bracket = arithmetic::sub( - &one_minus_pt, - &arithmetic::mul(&arithmetic::mul(&gamma_scalar, &pt)?, &log_pt)?, - )?; - let modulating = arithmetic::mul(&tensor_power(&one_minus_pt, self.gamma - 1.0)?, &bracket)?; - let alpha_tensor = create_scalar_tensor(self.alpha, dtype, device)?; - let weight = arithmetic::mul(&modulating, &alpha_tensor)?; // per-sample scalar - - let mut base_grad = arithmetic::mul(&arithmetic::sub(&p, &t)?, &weight)?; - - if self.reduction == "mean" { - let num_classes = *self.predictions_shape.last().unwrap_or(&1); - let num_samples = - (self.predictions_shape.iter().product::() / num_classes.max(1)) as f64; - let scale = create_scalar_tensor(1.0 / num_samples, dtype, device)?; - base_grad = arithmetic::mul(&base_grad, &scale)?; - } - - let pred_grad = arithmetic::mul(&base_grad, grad_output)?; - accumulate_grad(&mut gradients, self.input_ids[0], pred_grad)?; - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} - -/// Create a scalar tensor with the given value -fn create_scalar_tensor(value: f64, dtype: DataType, device: Device) -> Result { - let mut data = TensorData::zeros_on_device(1, dtype, device); - match dtype { - DataType::Float32 => { - let slice = data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from scalar") - })?; - slice[0] = value as f32; - } - DataType::Float64 => { - let slice = data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from scalar") - })?; - slice[0] = value; - } - _ => { - return Err(MinitensorError::invalid_operation( - "Scalar tensors only supported for floating point types", - )); - } - } - - Ok(Tensor::new( - Arc::new(data), - Shape::new(vec![1]), - dtype, - device, - false, - )) -} - -/// Raise each tensor element to the given power -fn tensor_power(tensor: &Tensor, exponent: f64) -> Result { - let mut output_data = - TensorData::zeros_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => { - let input = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from tensor") - })?; - let output = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output") - })?; - let exp = exponent as f32; - let len = input.len(); - debug_assert_eq!(len, output.len()); - if len < PAR_THRESHOLD { - for i in 0..len { - output[i] = input[i].powf(exp); - } - } else { - let in_ptr = input.as_ptr() as usize; - let out_ptr = output.as_mut_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let in_ptr = in_ptr as *const f32; - let out_ptr = out_ptr as *mut f32; - *out_ptr.add(i) = (*in_ptr.add(i)).powf(exp); - }); - } - } - DataType::Float64 => { - let input = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from tensor") - })?; - let output = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output") - })?; - let len = input.len(); - debug_assert_eq!(len, output.len()); - if len < PAR_THRESHOLD { - for i in 0..len { - output[i] = input[i].powf(exponent); - } - } else { - let in_ptr = input.as_ptr() as usize; - let out_ptr = output.as_mut_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let in_ptr = in_ptr as *const f64; - let out_ptr = out_ptr as *mut f64; - *out_ptr.add(i) = (*in_ptr.add(i)).powf(exponent); - }); - } - } - _ => { - return Err(MinitensorError::invalid_operation( - "Power operation only supported for floating point tensors", - )); - } - } - - Ok(Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - false, - )) -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::{ + device::Device, + error::{MinitensorError, Result}, + operations::{activation, arithmetic, reduction}, + tensor::{DataType, Shape, Tensor, TensorData}, +}; +use rayon::prelude::*; +use rustc_hash::FxHashMap; +use std::sync::Arc; + +pub(crate) fn repeat_interleave_backward_impl( + grad_output: &Tensor, + input_shape: &[usize], + repeats: &[usize], + dim: usize, +) -> Result { + if dim >= input_shape.len() { + return Err(MinitensorError::index_error( + dim as isize, + 0, + input_shape.len(), + )); + } + + let dim_size = input_shape[dim]; + if repeats.len() != dim_size { + return Err(MinitensorError::invalid_operation( + "repeat_interleave backward: repeats must match input dimension size".to_string(), + )); + } + + let grad_shape_vec = input_shape.to_vec(); + let grad_shape = Shape::new(grad_shape_vec.clone()); + let numel = grad_shape.numel(); + let dtype = grad_output.dtype(); + let device = grad_output.device(); + let total_repeats: usize = repeats.iter().sum(); + + let inner: usize = if dim + 1 >= input_shape.len() { + 1 + } else { + input_shape[dim + 1..].iter().product() + }; + let outer: usize = if dim == 0 { + 1 + } else { + input_shape[..dim].iter().product() + }; + + if numel == 0 || total_repeats == 0 || inner == 0 || outer == 0 { + return Ok(Tensor::zeros( + Shape::new(grad_shape_vec), + dtype, + device, + false, + )); + } + + let output_dims = grad_output.shape().dims(); + if output_dims.len() != input_shape.len() || output_dims[dim] != total_repeats { + return Err(MinitensorError::shape_mismatch( + input_shape.to_vec(), + output_dims.to_vec(), + )); + } + + macro_rules! repeat_interleave_backward_impl_inner { + ($ty:ty, $slice:ident, $from_vec:ident) => {{ + let src = grad_output.data().$slice().ok_or_else(|| { + MinitensorError::invalid_operation( + "repeat_interleave backward: gradient tensor must be contiguous".to_string(), + ) + })?; + let mut dst = vec![<$ty>::default(); numel]; + let chunk = total_repeats * inner; + dst.par_chunks_mut(dim_size * inner) + .enumerate() + .for_each(|(outer_idx, dst_chunk)| { + let mut src_offset = outer_idx * chunk; + for (i, &rep) in repeats.iter().enumerate() { + if rep == 0 { + continue; + } + let dst_start = i * inner; + let dst_slice = &mut dst_chunk[dst_start..dst_start + inner]; + for _ in 0..rep { + let src_slice = &src[src_offset..src_offset + inner]; + dst_slice.iter_mut().zip(src_slice.iter()).for_each( + |(dst_val, &src_val)| { + *dst_val += src_val; + }, + ); + src_offset += inner; + } + } + }); + TensorData::$from_vec(dst, device) + }}; + } + + let data = match dtype { + DataType::Float32 => { + repeat_interleave_backward_impl_inner!(f32, as_f32_slice, from_vec_f32) + } + DataType::Float64 => { + repeat_interleave_backward_impl_inner!(f64, as_f64_slice, from_vec_f64) + } + DataType::Int32 => repeat_interleave_backward_impl_inner!(i32, as_i32_slice, from_vec_i32), + DataType::Int64 => repeat_interleave_backward_impl_inner!(i64, as_i64_slice, from_vec_i64), + DataType::Bool => { + return Ok(Tensor::zeros(grad_shape, dtype, device, false)); + } + }; + + Ok(Tensor::new( + Arc::new(data), + grad_shape, + dtype, + device, + false, + )) +} + +/// Gradient function for expand operation which reduces broadcasted gradients +pub struct ExpandBackward { + pub input_shape: Vec, + pub input_id: TensorId, +} + +impl GradientFunction for ExpandBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + let shape = Shape::new(self.input_shape.clone()); + let grad_input = reduce_gradient_for_broadcasting(grad_output, &shape)?; + gradients.insert(self.input_id, grad_input); + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for MSE loss +pub struct MSELossBackward { + pub predictions_shape: Vec, + pub targets_shape: Vec, + pub input_ids: [TensorId; 2], + /// Which of [predictions, targets] actually need a gradient. Targets + /// almost never do, so their gradient chain is skipped entirely. + pub input_requires_grad: [bool; 2], + pub reduction: String, + pub diff: Tensor, +} + +impl GradientFunction for MSELossBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(2); + + // Base gradient: 2 * (predictions - targets) + let two = create_scalar_tensor(2.0, self.diff.dtype(), self.diff.device())?; + let mut base_grad = arithmetic::mul(&self.diff, &two)?; + + // Apply reduction scaling + match self.reduction.as_str() { + "mean" => { + let n = self.diff.numel() as f64; + let scale = create_scalar_tensor(1.0 / n, base_grad.dtype(), base_grad.device())?; + base_grad = arithmetic::mul(&base_grad, &scale)?; + } + "sum" | "none" => {} + _ => { + return Err(MinitensorError::gradient_error(format!( + "Unknown reduction mode: {}", + self.reduction + ))); + } + } + + // Multiply by upstream gradient + let pred_grad = arithmetic::mul(&base_grad, grad_output)?; + if self.input_requires_grad[1] { + let target_grad = arithmetic::neg(&pred_grad)?; + accumulate_grad(&mut gradients, self.input_ids[1], target_grad)?; + } + if self.input_requires_grad[0] { + accumulate_grad(&mut gradients, self.input_ids[0], pred_grad)?; + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} + +/// Gradient function for MAE loss +pub struct MAELossBackward { + pub predictions_shape: Vec, + pub targets_shape: Vec, + pub input_ids: [TensorId; 2], + /// Which of [predictions, targets] actually need a gradient. Targets + /// almost never do, so their gradient chain is skipped entirely. + pub input_requires_grad: [bool; 2], + pub reduction: String, + pub sign: Tensor, +} + +impl GradientFunction for MAELossBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(2); + + let mut base_grad = self.sign.clone(); + match self.reduction.as_str() { + "mean" => { + let n = self.sign.numel() as f64; + let scale = create_scalar_tensor(1.0 / n, base_grad.dtype(), base_grad.device())?; + base_grad = arithmetic::mul(&base_grad, &scale)?; + } + "sum" | "none" => {} + _ => { + return Err(MinitensorError::gradient_error(format!( + "Unknown reduction mode: {}", + self.reduction + ))); + } + } + + let pred_grad = arithmetic::mul(&base_grad, grad_output)?; + if self.input_requires_grad[1] { + let target_grad = arithmetic::neg(&pred_grad)?; + accumulate_grad(&mut gradients, self.input_ids[1], target_grad)?; + } + if self.input_requires_grad[0] { + accumulate_grad(&mut gradients, self.input_ids[0], pred_grad)?; + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} + +/// Gradient function for Huber loss +pub struct HuberLossBackward { + pub predictions_shape: Vec, + pub targets_shape: Vec, + pub input_ids: [TensorId; 2], + /// Which of [predictions, targets] actually need a gradient. Targets + /// almost never do, so their gradient chain is skipped entirely. + pub input_requires_grad: [bool; 2], + pub delta: f64, + pub reduction: String, + pub diff: Tensor, +} + +impl GradientFunction for HuberLossBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(2); + + let numel = self.diff.numel(); + let dtype = self.diff.dtype(); + let device = self.diff.device(); + let mut grad_data = TensorData::zeros_on_device(numel, dtype, device); + + match dtype { + DataType::Float32 => { + let diff_slice = self.diff.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from diff") + })?; + let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from grad") + })?; + let delta = self.delta as f32; + if numel < PAR_THRESHOLD { + for i in 0..numel { + let d = diff_slice[i]; + grad_slice[i] = if d.abs() <= delta { + d + } else { + delta * d.signum() + }; + } + } else { + let diff_ptr = diff_slice.as_ptr() as usize; + let grad_ptr = grad_slice.as_mut_ptr() as usize; + (0..numel).into_par_iter().for_each(|i| unsafe { + let diff_ptr = diff_ptr as *const f32; + let grad_ptr = grad_ptr as *mut f32; + let d = *diff_ptr.add(i); + *grad_ptr.add(i) = if d.abs() <= delta { + d + } else { + delta * d.signum() + }; + }); + } + } + DataType::Float64 => { + let diff_slice = self.diff.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from diff") + })?; + let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from grad") + })?; + if numel < PAR_THRESHOLD { + for i in 0..numel { + let d = diff_slice[i]; + grad_slice[i] = if d.abs() <= self.delta { + d + } else { + self.delta * d.signum() + }; + } + } else { + let diff_ptr = diff_slice.as_ptr() as usize; + let grad_ptr = grad_slice.as_mut_ptr() as usize; + let delta = self.delta; + (0..numel).into_par_iter().for_each(|i| unsafe { + let diff_ptr = diff_ptr as *const f64; + let grad_ptr = grad_ptr as *mut f64; + let d = *diff_ptr.add(i); + *grad_ptr.add(i) = if d.abs() <= delta { + d + } else { + delta * d.signum() + }; + }); + } + } + _ => { + return Err(MinitensorError::invalid_operation( + "Huber loss only supports floating point tensors", + )); + } + } + + let mut base_grad = Tensor::new( + Arc::new(grad_data), + Shape::new(self.predictions_shape.clone()), + dtype, + device, + false, + ); + + if self.reduction == "mean" { + let scale = create_scalar_tensor(1.0 / numel as f64, dtype, device)?; + base_grad = arithmetic::mul(&base_grad, &scale)?; + } + + let pred_grad = arithmetic::mul(&base_grad, grad_output)?; + if self.input_requires_grad[1] { + let target_grad = arithmetic::neg(&pred_grad)?; + accumulate_grad(&mut gradients, self.input_ids[1], target_grad)?; + } + if self.input_requires_grad[0] { + accumulate_grad(&mut gradients, self.input_ids[0], pred_grad)?; + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} + +/// Gradient function for Cross Entropy loss +pub struct CrossEntropyLossBackward { + pub predictions_shape: Vec, + pub targets_shape: Vec, + pub input_ids: [TensorId; 2], + /// Which of [predictions, targets] actually need a gradient. Only the + /// prediction gradient is ever produced; it is skipped when frozen. + pub input_requires_grad: [bool; 2], + pub reduction: String, + pub softmax_predictions: Tensor, + pub targets: Tensor, +} + +impl GradientFunction for CrossEntropyLossBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + if !self.input_requires_grad[0] { + return Ok(gradients); + } + gradients.reserve(1); + + // Compute base gradient: softmax(predictions) - targets + let mut base_grad = + arithmetic::sub(&self.softmax_predictions.detach(), &self.targets.detach())?; + + // Apply reduction scaling + match self.reduction.as_str() { + "mean" => { + let batch = self.targets_shape[0] as f64; + let mut scalar_data = + TensorData::zeros_on_device(1, base_grad.dtype(), base_grad.device()); + match base_grad.dtype() { + DataType::Float32 => { + let slice = scalar_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from scalar", + ) + })?; + slice[0] = (1.0 / batch) as f32; + } + DataType::Float64 => { + let slice = scalar_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from scalar", + ) + })?; + slice[0] = 1.0 / batch; + } + _ => { + return Err(MinitensorError::invalid_operation( + "CrossEntropy backward only supports floating point tensors", + )); + } + } + let scalar_tensor = Tensor::new( + Arc::new(scalar_data), + Shape::new(vec![1]), + base_grad.dtype(), + base_grad.device(), + false, + ); + base_grad = arithmetic::mul(&base_grad, &scalar_tensor)?; + } + "sum" | "none" => {} + _ => { + return Err(MinitensorError::gradient_error(format!( + "Unknown reduction mode: {}", + self.reduction + ))); + } + } + + // Multiply by upstream gradient (handles broadcasting) + let pred_grad = arithmetic::mul(&base_grad, grad_output)?; + + // Targets typically have no gradient + accumulate_grad(&mut gradients, self.input_ids[0], pred_grad)?; + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} + +/// Gradient function for Binary Cross Entropy loss +pub struct BCELossBackward { + pub predictions_shape: Vec, + pub targets_shape: Vec, + pub input_ids: [TensorId; 2], + /// Which of [predictions, targets] actually need a gradient. Only the + /// prediction gradient is ever produced; it is skipped when frozen. + pub input_requires_grad: [bool; 2], + pub reduction: String, + pub predictions: Tensor, + pub targets: Tensor, +} + +impl GradientFunction for BCELossBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + if !self.input_requires_grad[0] { + return Ok(gradients); + } + gradients.reserve(1); + + // BCE gradient: (predictions - targets) / (predictions * (1 - predictions)) + let one = Tensor::ones( + Shape::new(self.predictions_shape.clone()), + self.predictions.dtype(), + self.predictions.device(), + false, + ); + let one_minus_pred = arithmetic::sub(&one, &self.predictions)?; + let numerator = arithmetic::sub(&self.predictions, &self.targets)?; + let denom = arithmetic::mul(&self.predictions, &one_minus_pred)?; + let mut base_grad = arithmetic::div(&numerator, &denom)?; + + if self.reduction == "mean" { + let n = self.predictions.numel() as f64; + let scale = create_scalar_tensor(1.0 / n, base_grad.dtype(), base_grad.device())?; + base_grad = arithmetic::mul(&base_grad, &scale)?; + } + + let pred_grad = arithmetic::mul(&base_grad, grad_output)?; + accumulate_grad(&mut gradients, self.input_ids[0], pred_grad)?; + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} + +/// Gradient function for KL Divergence loss +pub struct KLDivLossBackward { + pub predictions_shape: Vec, + pub targets_shape: Vec, + pub input_ids: [TensorId; 2], + /// Which of [predictions, targets] actually need a gradient. Targets + /// almost never do, so their gradient chain is skipped entirely. + pub input_requires_grad: [bool; 2], + pub reduction: String, + pub predictions: Tensor, + pub targets: Tensor, +} + +impl GradientFunction for KLDivLossBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(2); + + // Gradient w.r.t predictions: -(targets / predictions) + if self.input_requires_grad[0] { + let mut pred_grad = arithmetic::div(&self.targets, &self.predictions)?; + pred_grad = arithmetic::neg(&pred_grad)?; + if self.reduction == "mean" { + let n = self.predictions.numel() as f64; + let scale = create_scalar_tensor(1.0 / n, pred_grad.dtype(), pred_grad.device())?; + pred_grad = arithmetic::mul(&pred_grad, &scale)?; + } + let pred_grad = arithmetic::mul(&pred_grad, grad_output)?; + accumulate_grad(&mut gradients, self.input_ids[0], pred_grad)?; + } + + // Gradient w.r.t targets: log(targets) - log(predictions) + 1 + if self.input_requires_grad[1] { + let log_targets = activation::log(&self.targets)?; + let log_preds = activation::log(&self.predictions)?; + let diff = arithmetic::sub(&log_targets, &log_preds)?; + let one = Tensor::ones( + self.targets.shape().clone(), + self.targets.dtype(), + self.targets.device(), + false, + ); + let mut target_grad = arithmetic::add(&diff, &one)?; + if self.reduction == "mean" { + let n = self.predictions.numel() as f64; + let scale = + create_scalar_tensor(1.0 / n, target_grad.dtype(), target_grad.device())?; + target_grad = arithmetic::mul(&target_grad, &scale)?; + } + let target_grad = arithmetic::mul(&target_grad, grad_output)?; + accumulate_grad(&mut gradients, self.input_ids[1], target_grad)?; + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} + +/// Gradient function for Focal loss +pub struct FocalLossBackward { + pub predictions_shape: Vec, + pub targets_shape: Vec, + pub input_ids: [TensorId; 2], + /// Which of [predictions, targets] actually need a gradient. Only the + /// prediction gradient is ever produced; it is skipped when frozen. + pub input_requires_grad: [bool; 2], + pub alpha: f64, + pub gamma: f64, + pub reduction: String, + pub softmax_predictions: Tensor, + pub targets: Tensor, +} + +impl GradientFunction for FocalLossBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + if !self.input_requires_grad[0] { + return Ok(gradients); + } + gradients.reserve(1); + + // Exact gradient of FL = -alpha * (1 - p_t)^gamma * log(p_t) wrt the + // logits, where p_t is the true-class softmax probability: + // dFL/dz_j = alpha * (p_j - onehot_j) + // * (1 - p_t)^(gamma-1) * [ (1 - p_t) - gamma * p_t * ln(p_t) ] + // The modulating factor is a per-sample scalar (broadcast over classes). + let p = self.softmax_predictions.detach(); + let t = self.targets.detach(); + let dtype = p.dtype(); + let device = p.device(); + + // True-class probability per sample: p_t = sum(p * onehot) over classes. + let class_dim = (p.ndim() - 1) as isize; + let pt = reduction::sum(&arithmetic::mul(&p, &t)?, Some(vec![class_dim]), true)?; + + let one = create_scalar_tensor(1.0, dtype, device)?; + let one_minus_pt = arithmetic::sub(&one, &pt)?; + let log_pt = crate::operations::activation::log(&pt)?; + let gamma_scalar = create_scalar_tensor(self.gamma, dtype, device)?; + // bracket = (1 - p_t) - gamma * p_t * ln(p_t) + let bracket = arithmetic::sub( + &one_minus_pt, + &arithmetic::mul(&arithmetic::mul(&gamma_scalar, &pt)?, &log_pt)?, + )?; + let modulating = + arithmetic::mul(&tensor_power(&one_minus_pt, self.gamma - 1.0)?, &bracket)?; + let alpha_tensor = create_scalar_tensor(self.alpha, dtype, device)?; + let weight = arithmetic::mul(&modulating, &alpha_tensor)?; // per-sample scalar + + let mut base_grad = arithmetic::mul(&arithmetic::sub(&p, &t)?, &weight)?; + + if self.reduction == "mean" { + let num_classes = *self.predictions_shape.last().unwrap_or(&1); + let num_samples = + (self.predictions_shape.iter().product::() / num_classes.max(1)) as f64; + let scale = create_scalar_tensor(1.0 / num_samples, dtype, device)?; + base_grad = arithmetic::mul(&base_grad, &scale)?; + } + + let pred_grad = arithmetic::mul(&base_grad, grad_output)?; + accumulate_grad(&mut gradients, self.input_ids[0], pred_grad)?; + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} + +/// Create a scalar tensor with the given value +pub(crate) fn create_scalar_tensor(value: f64, dtype: DataType, device: Device) -> Result { + let mut data = TensorData::zeros_on_device(1, dtype, device); + match dtype { + DataType::Float32 => { + let slice = data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from scalar") + })?; + slice[0] = value as f32; + } + DataType::Float64 => { + let slice = data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from scalar") + })?; + slice[0] = value; + } + _ => { + return Err(MinitensorError::invalid_operation( + "Scalar tensors only supported for floating point types", + )); + } + } + + Ok(Tensor::new( + Arc::new(data), + Shape::new(vec![1]), + dtype, + device, + false, + )) +} + +/// Raise each tensor element to the given power +fn tensor_power(tensor: &Tensor, exponent: f64) -> Result { + let mut output_data = + TensorData::zeros_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => { + let input = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from tensor") + })?; + let output = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output") + })?; + let exp = exponent as f32; + let len = input.len(); + debug_assert_eq!(len, output.len()); + if len < PAR_THRESHOLD { + for i in 0..len { + output[i] = input[i].powf(exp); + } + } else { + let in_ptr = input.as_ptr() as usize; + let out_ptr = output.as_mut_ptr() as usize; + (0..len).into_par_iter().for_each(|i| unsafe { + let in_ptr = in_ptr as *const f32; + let out_ptr = out_ptr as *mut f32; + *out_ptr.add(i) = (*in_ptr.add(i)).powf(exp); + }); + } + } + DataType::Float64 => { + let input = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from tensor") + })?; + let output = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output") + })?; + let len = input.len(); + debug_assert_eq!(len, output.len()); + if len < PAR_THRESHOLD { + for i in 0..len { + output[i] = input[i].powf(exponent); + } + } else { + let in_ptr = input.as_ptr() as usize; + let out_ptr = output.as_mut_ptr() as usize; + (0..len).into_par_iter().for_each(|i| unsafe { + let in_ptr = in_ptr as *const f64; + let out_ptr = out_ptr as *mut f64; + *out_ptr.add(i) = (*in_ptr.add(i)).powf(exponent); + }); + } + } + _ => { + return Err(MinitensorError::invalid_operation( + "Power operation only supported for floating point tensors", + )); + } + } + + Ok(Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + false, + )) +} diff --git a/engine/src/autograd/mod/arithmetic.rs b/engine/src/autograd/mod/arithmetic.rs index fd1ea154..1ff9c69d 100644 --- a/engine/src/autograd/mod/arithmetic.rs +++ b/engine/src/autograd/mod/arithmetic.rs @@ -1,845 +1,856 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -fn expand_reduction_grad( - grad_output: &Tensor, - input_shape: &[usize], - dims: &Option>, - keepdim: bool, -) -> Result { - if keepdim { - return Ok(grad_output.clone()); - } - - if let Some(dims) = dims { - let mut shape = grad_output.shape().dims().to_vec(); - let mut sorted = dims.clone(); - sorted.sort_unstable(); - for &d in &sorted { - shape.insert(d, 1); - } - shape_ops::reshape(grad_output, Shape::new(shape)) - } else { - shape_ops::reshape(grad_output, Shape::new(vec![1; input_shape.len()])) - } -} - -impl GradientFunction for SumBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - let grad = expand_reduction_grad(grad_output, &self.input_shape, &self.dims, self.keepdim)?; - - let ones = Tensor::ones( - Shape::new(self.input_shape.clone()), - grad_output.dtype(), - grad_output.device(), - false, - ); - let grad_input = arithmetic::mul(&ones, &grad)?; - gradients.insert(self.input_id, grad_input); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for NaN-aware sum reduction -pub struct NanSumBackward { - pub input_id: TensorId, - pub input_shape: Vec, - pub dims: Option>, - pub keepdim: bool, - pub mask: Tensor, -} - -impl GradientFunction for NanSumBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - let grad = expand_reduction_grad(grad_output, &self.input_shape, &self.dims, self.keepdim)?; - let mask = self.mask.astype(grad_output.dtype())?; - let grad_input = arithmetic::mul(&mask, &grad)?; - gradients.insert(self.input_id, grad_input); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for NaN-aware mean reduction -pub struct NanMeanBackward { - pub input_id: TensorId, - pub input_shape: Vec, - pub dims: Option>, - pub keepdim: bool, - pub mask: Tensor, - pub count: Tensor, -} - -impl GradientFunction for NanMeanBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - let grad = expand_reduction_grad(grad_output, &self.input_shape, &self.dims, self.keepdim)?; - let count = - expand_reduction_grad(&self.count, &self.input_shape, &self.dims, self.keepdim)?; - let grad = sanitize_grad_for_nanmean(&grad, &count)?; - let count = safe_count_for_nanmean(&count)?; - - let scaled = arithmetic::div(&grad, &count)?; - let mask = self.mask.astype(grad_output.dtype())?; - let grad_input = arithmetic::mul(&mask, &scaled)?; - gradients.insert(self.input_id, grad_input); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -fn sanitize_grad_for_nanmean(grad: &Tensor, count: &Tensor) -> Result { - if grad.dtype() != count.dtype() { - return Err(MinitensorError::invalid_operation( - "nanmean backward expected matching gradient and count dtypes", - )); - } - - let numel = grad.numel(); - let mut new_data = TensorData::zeros_on_device(numel, grad.dtype(), grad.device()); - - match grad.dtype() { - DataType::Float32 => { - let grad_src = grad - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let count_src = count - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let dst = new_data - .as_f32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - dst.par_iter_mut() - .zip(grad_src.par_iter().zip(count_src.par_iter())) - .for_each(|(out, (&g, &c))| { - *out = if c == 0.0 { 0.0 } else { g }; - }); - } - DataType::Float64 => { - let grad_src = grad - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let count_src = count - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let dst = new_data - .as_f64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - dst.par_iter_mut() - .zip(grad_src.par_iter().zip(count_src.par_iter())) - .for_each(|(out, (&g, &c))| { - *out = if c == 0.0 { 0.0 } else { g }; - }); - } - _ => { - return Err(MinitensorError::invalid_operation( - "nanmean backward only supports floating point tensors", - )); - } - } - - Ok(Tensor::new( - Arc::new(new_data), - grad.shape().clone(), - grad.dtype(), - grad.device(), - false, - )) -} - -fn safe_count_for_nanmean(count: &Tensor) -> Result { - let numel = count.numel(); - let mut new_data = TensorData::zeros_on_device(numel, count.dtype(), count.device()); - - match count.dtype() { - DataType::Float32 => { - let src = count - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let dst = new_data - .as_f32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - dst.par_iter_mut() - .zip(src.par_iter()) - .for_each(|(out, &c)| { - *out = if c == 0.0 { 1.0 } else { c }; - }); - } - DataType::Float64 => { - let src = count - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let dst = new_data - .as_f64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - dst.par_iter_mut() - .zip(src.par_iter()) - .for_each(|(out, &c)| { - *out = if c == 0.0 { 1.0 } else { c }; - }); - } - _ => { - return Err(MinitensorError::invalid_operation( - "nanmean backward only supports floating point tensors", - )); - } - } - - Ok(Tensor::new( - Arc::new(new_data), - count.shape().clone(), - count.dtype(), - count.device(), - false, - )) -} - -/// Gradient function for product reduction -pub struct ProdBackward { - pub input: Tensor, - pub result: Tensor, - pub input_id: TensorId, - pub dims: Option>, - pub keepdim: bool, -} - -impl GradientFunction for ProdBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - let input = &self.input; - let input_shape = input.shape().dims().to_vec(); - let dtype = input.dtype(); - let device = input.device(); - - // Broadcast the upstream gradient back over the reduced axes. - let grad = expand_reduction_grad(grad_output, &input_shape, &self.dims, self.keepdim)?; - - let reduce_dims: Option> = self - .dims - .as_ref() - .map(|dims| dims.iter().map(|&d| d as isize).collect()); - - // d(prod)/dx_i is the product of the *other* elements in the reduction - // group. Computing it as `total_product / x_i` breaks when the group - // contains zeros (0 / 0 = NaN), so handle zeros explicitly: - // - no zeros in the group: grad_i = P / x_i - // - exactly one zero: grad_i = product of the non-zero elements - // at the zero position, 0 elsewhere - // - two or more zeros: grad_i = 0 everywhere - let zero = create_scalar_tensor(0.0, dtype, device)?; - let is_zero = crate::operations::comparison::eq(input, &zero)?; // bool mask - let is_zero_f = is_zero.astype(dtype)?; - let ones = Tensor::ones(input.shape().clone(), dtype, device, false); - - // Per-group zero count and product of the non-zero elements. - let zero_count = reduction::sum(&is_zero_f, reduce_dims.clone(), true)?; - let safe_input = crate::operations::selection::where_op(&is_zero, &ones, input)?; - let prod_nonzero = reduction::prod(&safe_input, reduce_dims, true)?; - - let one_scalar = create_scalar_tensor(1.0, dtype, device)?; - let no_zero = crate::operations::comparison::eq(&zero_count, &zero)?.astype(dtype)?; - let one_zero = crate::operations::comparison::eq(&zero_count, &one_scalar)?.astype(dtype)?; - - // Contribution at the (unique) zero position: product of the others. - let zero_term = arithmetic::mul(&is_zero_f, &one_zero)?; - let zero_term = arithmetic::mul(&zero_term, &prod_nonzero)?; - - // Contribution at non-zero positions when the group has no zeros. - let nonzero_mask = arithmetic::sub(&ones, &is_zero_f)?; - let quotient = arithmetic::div(&prod_nonzero, &safe_input)?; - let nonzero_term = arithmetic::mul(&nonzero_mask, &no_zero)?; - let nonzero_term = arithmetic::mul(&nonzero_term, "ient)?; - - let per_element = arithmetic::add(&zero_term, &nonzero_term)?; - let grad_input = arithmetic::mul(&grad, &per_element)?; - gradients.insert(self.input_id, grad_input); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for cumulative sum operation -pub struct CumsumBackward { - pub input_id: TensorId, - pub dim: usize, -} - -/// Gradient function for cumulative product operation -pub struct CumprodBackward { - pub input_id: TensorId, - pub input: Tensor, - pub output: Tensor, - pub dim: usize, -} - -impl GradientFunction for CumprodBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - let grad_input = - reduction::cumprod_backward(&self.input, &self.output, grad_output, self.dim)?; - gradients.insert(self.input_id, grad_input); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -impl GradientFunction for CumsumBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - let grad_input = reduction::cumsum_backward(grad_output, self.dim)?; - gradients.insert(self.input_id, grad_input); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -// Gradient functions for activation functions - -/// Gradient function for exponential -pub struct ExpBackward { - pub input_id: TensorId, - pub output: Tensor, -} - -impl GradientFunction for ExpBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // d/dx(exp(x)) = exp(x) * grad_output - let grad = arithmetic::mul(&self.output, grad_output)?; - gradients.insert(self.input_id, grad); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for logarithm -pub struct LogBackward { - pub input_id: TensorId, - pub input: Tensor, -} - -impl GradientFunction for LogBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // d/dx(log(x)) = 1/x * grad_output - let ones = Tensor::ones( - self.input.shape().clone(), - self.input.dtype(), - self.input.device(), - false, - ); - let inv = arithmetic::div(&ones, &self.input.detach())?; - let grad = arithmetic::mul(&inv, grad_output)?; - gradients.insert(self.input_id, grad); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for log1p -pub struct Log1pBackward { - pub input_id: TensorId, - pub input: Tensor, -} - -impl GradientFunction for Log1pBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - let ones = Tensor::ones( - self.input.shape().clone(), - self.input.dtype(), - self.input.device(), - false, - ); - let denom = arithmetic::add(&ones, &self.input.detach())?; - let grad = arithmetic::div(grad_output, &denom)?; - gradients.insert(self.input_id, grad); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for expm1 -pub struct Expm1Backward { - pub input_id: TensorId, - pub output: Tensor, -} - -impl GradientFunction for Expm1Backward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - let ones = Tensor::ones( - self.output.shape().clone(), - self.output.dtype(), - self.output.device(), - false, - ); - let term = arithmetic::add(&self.output.detach(), &ones)?; - let grad = arithmetic::mul(&term, grad_output)?; - gradients.insert(self.input_id, grad); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for sine -pub struct SinBackward { - pub input_id: TensorId, - pub input: Tensor, -} - -impl GradientFunction for SinBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // d/dx(sin(x)) = cos(x) * grad_output - let cos_x = self.input.cos()?; - let grad = arithmetic::mul(&cos_x, grad_output)?; - gradients.insert(self.input_id, grad); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for cosine -pub struct CosBackward { - pub input_id: TensorId, - pub input: Tensor, -} - -impl GradientFunction for CosBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // d/dx(cos(x)) = -sin(x) * grad_output - let sin_x = self.input.sin()?; - let mul = arithmetic::mul(&sin_x, grad_output)?; - let grad = arithmetic::neg(&mul)?; - gradients.insert(self.input_id, grad); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for tangent -pub struct TanBackward { - pub input_id: TensorId, - pub output: Tensor, -} - -impl GradientFunction for TanBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // d/dx(tan(x)) = (1 + tan²(x)) * grad_output - let tan_sq = arithmetic::mul(&self.output, &self.output)?; - let ones = Tensor::ones( - self.output.shape().clone(), - self.output.dtype(), - self.output.device(), - false, - ); - let term = arithmetic::add(&ones, &tan_sq)?; - let grad = arithmetic::mul(&term, grad_output)?; - gradients.insert(self.input_id, grad); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for inverse sine -pub struct AsinBackward { - pub input_id: TensorId, - pub input: Tensor, -} - -impl GradientFunction for AsinBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // d/dx(asin(x)) = grad_output / sqrt(1 - x^2) - let square = arithmetic::mul(&self.input, &self.input)?; - let ones = Tensor::ones( - self.input.shape().clone(), - self.input.dtype(), - self.input.device(), - false, - ); - let denom = arithmetic::sub(&ones, &square)?; - let sqrt = denom.sqrt()?; - let grad = arithmetic::div(grad_output, &sqrt)?; - gradients.insert(self.input_id, grad); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for inverse cosine -pub struct AcosBackward { - pub input_id: TensorId, - pub input: Tensor, -} - -impl GradientFunction for AcosBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // d/dx(acos(x)) = -grad_output / sqrt(1 - x^2) - let square = arithmetic::mul(&self.input, &self.input)?; - let ones = Tensor::ones( - self.input.shape().clone(), - self.input.dtype(), - self.input.device(), - false, - ); - let denom = arithmetic::sub(&ones, &square)?; - let sqrt = denom.sqrt()?; - let frac = arithmetic::div(grad_output, &sqrt)?; - let grad = arithmetic::neg(&frac)?; - gradients.insert(self.input_id, grad); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for inverse tangent -pub struct AtanBackward { - pub input_id: TensorId, - pub input: Tensor, -} - -impl GradientFunction for AtanBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // d/dx(atan(x)) = grad_output / (1 + x^2) - let square = arithmetic::mul(&self.input, &self.input)?; - let ones = Tensor::ones( - self.input.shape().clone(), - self.input.dtype(), - self.input.device(), - false, - ); - let denom = arithmetic::add(&ones, &square)?; - let grad = arithmetic::div(grad_output, &denom)?; - gradients.insert(self.input_id, grad); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for hyperbolic sine -pub struct SinhBackward { - pub input_id: TensorId, - pub input: Tensor, -} - -impl GradientFunction for SinhBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // d/dx(sinh(x)) = cosh(x) * grad_output - let cosh_x = self.input.cosh()?; - let grad = arithmetic::mul(&cosh_x, grad_output)?; - gradients.insert(self.input_id, grad); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for hyperbolic cosine -pub struct CoshBackward { - pub input_id: TensorId, - pub input: Tensor, -} - -impl GradientFunction for CoshBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // d/dx(cosh(x)) = sinh(x) * grad_output - let sinh_x = self.input.sinh()?; - let grad = arithmetic::mul(&sinh_x, grad_output)?; - gradients.insert(self.input_id, grad); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for inverse hyperbolic sine -pub struct AsinhBackward { - pub input_id: TensorId, - pub input: Tensor, -} - -impl GradientFunction for AsinhBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // d/dx(asinh(x)) = grad_output / sqrt(1 + x^2) - let square = arithmetic::mul(&self.input, &self.input)?; - let ones = Tensor::ones( - self.input.shape().clone(), - self.input.dtype(), - self.input.device(), - false, - ); - let denom = arithmetic::add(&square, &ones)?; - let sqrt = denom.sqrt()?; - let grad = arithmetic::div(grad_output, &sqrt)?; - gradients.insert(self.input_id, grad); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for inverse hyperbolic cosine -pub struct AcoshBackward { - pub input_id: TensorId, - pub input: Tensor, -} - -impl GradientFunction for AcoshBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // d/dx(acosh(x)) = grad_output / sqrt((x - 1)(x + 1)) - let ones = Tensor::ones( - self.input.shape().clone(), - self.input.dtype(), - self.input.device(), - false, - ); - let x_minus_one = arithmetic::sub(&self.input, &ones)?; - let x_plus_one = arithmetic::add(&self.input, &ones)?; - let product = arithmetic::mul(&x_minus_one, &x_plus_one)?; - let sqrt = product.sqrt()?; - let grad = arithmetic::div(grad_output, &sqrt)?; - gradients.insert(self.input_id, grad); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for inverse hyperbolic tangent -pub struct AtanhBackward { - pub input_id: TensorId, - pub input: Tensor, -} - -impl GradientFunction for AtanhBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // d/dx(atanh(x)) = grad_output / (1 - x^2) - let square = arithmetic::mul(&self.input, &self.input)?; - let ones = Tensor::ones( - self.input.shape().clone(), - self.input.dtype(), - self.input.device(), - false, - ); - let denom = arithmetic::sub(&ones, &square)?; - let grad = arithmetic::div(grad_output, &denom)?; - gradients.insert(self.input_id, grad); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for tanh -pub struct TanhBackward { - pub input_id: TensorId, - pub output: Tensor, -} - -impl GradientFunction for TanhBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // d/dx(tanh(x)) = (1 - tanh²(x)) * grad_output - let y2 = arithmetic::mul(&self.output, &self.output)?; - let ones = Tensor::ones( - self.output.shape().clone(), - self.output.dtype(), - self.output.device(), - false, - ); - let term = arithmetic::sub(&ones, &y2)?; - let grad = arithmetic::mul(&term, grad_output)?; - gradients.insert(self.input_id, grad); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for sigmoid -pub struct SigmoidBackward { - pub input_id: TensorId, - pub output: Tensor, -} - -impl GradientFunction for SigmoidBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // d/dx(sigmoid(x)) = sigmoid(x) * (1 - sigmoid(x)) * grad_output - let ones = Tensor::ones( - self.output.shape().clone(), - self.output.dtype(), - self.output.device(), - false, - ); - let one_minus = arithmetic::sub(&ones, &self.output)?; - let term = arithmetic::mul(&self.output, &one_minus)?; - let grad = arithmetic::mul(&term, grad_output)?; - gradients.insert(self.input_id, grad); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for Softplus -pub struct SoftplusBackward { - pub input_id: TensorId, - pub input: Tensor, - pub beta: f64, - pub threshold: f64, -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::{ + error::{MinitensorError, Result}, + operations::{arithmetic, reduction, shape_ops}, + tensor::{DataType, Shape, Tensor, TensorData}, +}; +use rayon::prelude::*; +use rustc_hash::FxHashMap; +use std::sync::Arc; + +pub(crate) fn expand_reduction_grad( + grad_output: &Tensor, + input_shape: &[usize], + dims: &Option>, + keepdim: bool, +) -> Result { + if keepdim { + return Ok(grad_output.clone()); + } + + if let Some(dims) = dims { + let mut shape = grad_output.shape().dims().to_vec(); + let mut sorted = dims.clone(); + sorted.sort_unstable(); + for &d in &sorted { + shape.insert(d, 1); + } + shape_ops::reshape(grad_output, Shape::new(shape)) + } else { + shape_ops::reshape(grad_output, Shape::new(vec![1; input_shape.len()])) + } +} + +impl GradientFunction for SumBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + let grad = expand_reduction_grad(grad_output, &self.input_shape, &self.dims, self.keepdim)?; + + let ones = Tensor::ones( + Shape::new(self.input_shape.clone()), + grad_output.dtype(), + grad_output.device(), + false, + ); + let grad_input = arithmetic::mul(&ones, &grad)?; + gradients.insert(self.input_id, grad_input); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for NaN-aware sum reduction +pub struct NanSumBackward { + pub input_id: TensorId, + pub input_shape: Vec, + pub dims: Option>, + pub keepdim: bool, + pub mask: Tensor, +} + +impl GradientFunction for NanSumBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + let grad = expand_reduction_grad(grad_output, &self.input_shape, &self.dims, self.keepdim)?; + let mask = self.mask.astype(grad_output.dtype())?; + let grad_input = arithmetic::mul(&mask, &grad)?; + gradients.insert(self.input_id, grad_input); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for NaN-aware mean reduction +pub struct NanMeanBackward { + pub input_id: TensorId, + pub input_shape: Vec, + pub dims: Option>, + pub keepdim: bool, + pub mask: Tensor, + pub count: Tensor, +} + +impl GradientFunction for NanMeanBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + let grad = expand_reduction_grad(grad_output, &self.input_shape, &self.dims, self.keepdim)?; + let count = + expand_reduction_grad(&self.count, &self.input_shape, &self.dims, self.keepdim)?; + let grad = sanitize_grad_for_nanmean(&grad, &count)?; + let count = safe_count_for_nanmean(&count)?; + + let scaled = arithmetic::div(&grad, &count)?; + let mask = self.mask.astype(grad_output.dtype())?; + let grad_input = arithmetic::mul(&mask, &scaled)?; + gradients.insert(self.input_id, grad_input); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +fn sanitize_grad_for_nanmean(grad: &Tensor, count: &Tensor) -> Result { + if grad.dtype() != count.dtype() { + return Err(MinitensorError::invalid_operation( + "nanmean backward expected matching gradient and count dtypes", + )); + } + + let numel = grad.numel(); + let mut new_data = TensorData::zeros_on_device(numel, grad.dtype(), grad.device()); + + match grad.dtype() { + DataType::Float32 => { + let grad_src = grad + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let count_src = count + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let dst = new_data + .as_f32_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + dst.par_iter_mut() + .zip(grad_src.par_iter().zip(count_src.par_iter())) + .for_each(|(out, (&g, &c))| { + *out = if c == 0.0 { 0.0 } else { g }; + }); + } + DataType::Float64 => { + let grad_src = grad + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let count_src = count + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let dst = new_data + .as_f64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + dst.par_iter_mut() + .zip(grad_src.par_iter().zip(count_src.par_iter())) + .for_each(|(out, (&g, &c))| { + *out = if c == 0.0 { 0.0 } else { g }; + }); + } + _ => { + return Err(MinitensorError::invalid_operation( + "nanmean backward only supports floating point tensors", + )); + } + } + + Ok(Tensor::new( + Arc::new(new_data), + grad.shape().clone(), + grad.dtype(), + grad.device(), + false, + )) +} + +fn safe_count_for_nanmean(count: &Tensor) -> Result { + let numel = count.numel(); + let mut new_data = TensorData::zeros_on_device(numel, count.dtype(), count.device()); + + match count.dtype() { + DataType::Float32 => { + let src = count + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let dst = new_data + .as_f32_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + dst.par_iter_mut() + .zip(src.par_iter()) + .for_each(|(out, &c)| { + *out = if c == 0.0 { 1.0 } else { c }; + }); + } + DataType::Float64 => { + let src = count + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let dst = new_data + .as_f64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + dst.par_iter_mut() + .zip(src.par_iter()) + .for_each(|(out, &c)| { + *out = if c == 0.0 { 1.0 } else { c }; + }); + } + _ => { + return Err(MinitensorError::invalid_operation( + "nanmean backward only supports floating point tensors", + )); + } + } + + Ok(Tensor::new( + Arc::new(new_data), + count.shape().clone(), + count.dtype(), + count.device(), + false, + )) +} + +/// Gradient function for product reduction +pub struct ProdBackward { + pub input: Tensor, + pub result: Tensor, + pub input_id: TensorId, + pub dims: Option>, + pub keepdim: bool, +} + +impl GradientFunction for ProdBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + let input = &self.input; + let input_shape = input.shape().dims().to_vec(); + let dtype = input.dtype(); + let device = input.device(); + + // Broadcast the upstream gradient back over the reduced axes. + let grad = expand_reduction_grad(grad_output, &input_shape, &self.dims, self.keepdim)?; + + let reduce_dims: Option> = self + .dims + .as_ref() + .map(|dims| dims.iter().map(|&d| d as isize).collect()); + + // d(prod)/dx_i is the product of the *other* elements in the reduction + // group. Computing it as `total_product / x_i` breaks when the group + // contains zeros (0 / 0 = NaN), so handle zeros explicitly: + // - no zeros in the group: grad_i = P / x_i + // - exactly one zero: grad_i = product of the non-zero elements + // at the zero position, 0 elsewhere + // - two or more zeros: grad_i = 0 everywhere + let zero = create_scalar_tensor(0.0, dtype, device)?; + let is_zero = crate::operations::comparison::eq(input, &zero)?; // bool mask + let is_zero_f = is_zero.astype(dtype)?; + let ones = Tensor::ones(input.shape().clone(), dtype, device, false); + + // Per-group zero count and product of the non-zero elements. + let zero_count = reduction::sum(&is_zero_f, reduce_dims.clone(), true)?; + let safe_input = crate::operations::selection::where_op(&is_zero, &ones, input)?; + let prod_nonzero = reduction::prod(&safe_input, reduce_dims, true)?; + + let one_scalar = create_scalar_tensor(1.0, dtype, device)?; + let no_zero = crate::operations::comparison::eq(&zero_count, &zero)?.astype(dtype)?; + let one_zero = + crate::operations::comparison::eq(&zero_count, &one_scalar)?.astype(dtype)?; + + // Contribution at the (unique) zero position: product of the others. + let zero_term = arithmetic::mul(&is_zero_f, &one_zero)?; + let zero_term = arithmetic::mul(&zero_term, &prod_nonzero)?; + + // Contribution at non-zero positions when the group has no zeros. + let nonzero_mask = arithmetic::sub(&ones, &is_zero_f)?; + let quotient = arithmetic::div(&prod_nonzero, &safe_input)?; + let nonzero_term = arithmetic::mul(&nonzero_mask, &no_zero)?; + let nonzero_term = arithmetic::mul(&nonzero_term, "ient)?; + + let per_element = arithmetic::add(&zero_term, &nonzero_term)?; + let grad_input = arithmetic::mul(&grad, &per_element)?; + gradients.insert(self.input_id, grad_input); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for cumulative sum operation +pub struct CumsumBackward { + pub input_id: TensorId, + pub dim: usize, +} + +/// Gradient function for cumulative product operation +pub struct CumprodBackward { + pub input_id: TensorId, + pub input: Tensor, + pub output: Tensor, + pub dim: usize, +} + +impl GradientFunction for CumprodBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + let grad_input = + reduction::cumprod_backward(&self.input, &self.output, grad_output, self.dim)?; + gradients.insert(self.input_id, grad_input); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +impl GradientFunction for CumsumBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + let grad_input = reduction::cumsum_backward(grad_output, self.dim)?; + gradients.insert(self.input_id, grad_input); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +// Gradient functions for activation functions + +/// Gradient function for exponential +pub struct ExpBackward { + pub input_id: TensorId, + pub output: Tensor, +} + +impl GradientFunction for ExpBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + // d/dx(exp(x)) = exp(x) * grad_output + let grad = arithmetic::mul(&self.output, grad_output)?; + gradients.insert(self.input_id, grad); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for logarithm +pub struct LogBackward { + pub input_id: TensorId, + pub input: Tensor, +} + +impl GradientFunction for LogBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + // d/dx(log(x)) = 1/x * grad_output + let ones = Tensor::ones( + self.input.shape().clone(), + self.input.dtype(), + self.input.device(), + false, + ); + let inv = arithmetic::div(&ones, &self.input.detach())?; + let grad = arithmetic::mul(&inv, grad_output)?; + gradients.insert(self.input_id, grad); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for log1p +pub struct Log1pBackward { + pub input_id: TensorId, + pub input: Tensor, +} + +impl GradientFunction for Log1pBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + let ones = Tensor::ones( + self.input.shape().clone(), + self.input.dtype(), + self.input.device(), + false, + ); + let denom = arithmetic::add(&ones, &self.input.detach())?; + let grad = arithmetic::div(grad_output, &denom)?; + gradients.insert(self.input_id, grad); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for expm1 +pub struct Expm1Backward { + pub input_id: TensorId, + pub output: Tensor, +} + +impl GradientFunction for Expm1Backward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + let ones = Tensor::ones( + self.output.shape().clone(), + self.output.dtype(), + self.output.device(), + false, + ); + let term = arithmetic::add(&self.output.detach(), &ones)?; + let grad = arithmetic::mul(&term, grad_output)?; + gradients.insert(self.input_id, grad); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for sine +pub struct SinBackward { + pub input_id: TensorId, + pub input: Tensor, +} + +impl GradientFunction for SinBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + // d/dx(sin(x)) = cos(x) * grad_output + let cos_x = self.input.cos()?; + let grad = arithmetic::mul(&cos_x, grad_output)?; + gradients.insert(self.input_id, grad); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for cosine +pub struct CosBackward { + pub input_id: TensorId, + pub input: Tensor, +} + +impl GradientFunction for CosBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + // d/dx(cos(x)) = -sin(x) * grad_output + let sin_x = self.input.sin()?; + let mul = arithmetic::mul(&sin_x, grad_output)?; + let grad = arithmetic::neg(&mul)?; + gradients.insert(self.input_id, grad); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for tangent +pub struct TanBackward { + pub input_id: TensorId, + pub output: Tensor, +} + +impl GradientFunction for TanBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + // d/dx(tan(x)) = (1 + tan²(x)) * grad_output + let tan_sq = arithmetic::mul(&self.output, &self.output)?; + let ones = Tensor::ones( + self.output.shape().clone(), + self.output.dtype(), + self.output.device(), + false, + ); + let term = arithmetic::add(&ones, &tan_sq)?; + let grad = arithmetic::mul(&term, grad_output)?; + gradients.insert(self.input_id, grad); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for inverse sine +pub struct AsinBackward { + pub input_id: TensorId, + pub input: Tensor, +} + +impl GradientFunction for AsinBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + // d/dx(asin(x)) = grad_output / sqrt(1 - x^2) + let square = arithmetic::mul(&self.input, &self.input)?; + let ones = Tensor::ones( + self.input.shape().clone(), + self.input.dtype(), + self.input.device(), + false, + ); + let denom = arithmetic::sub(&ones, &square)?; + let sqrt = denom.sqrt()?; + let grad = arithmetic::div(grad_output, &sqrt)?; + gradients.insert(self.input_id, grad); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for inverse cosine +pub struct AcosBackward { + pub input_id: TensorId, + pub input: Tensor, +} + +impl GradientFunction for AcosBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + // d/dx(acos(x)) = -grad_output / sqrt(1 - x^2) + let square = arithmetic::mul(&self.input, &self.input)?; + let ones = Tensor::ones( + self.input.shape().clone(), + self.input.dtype(), + self.input.device(), + false, + ); + let denom = arithmetic::sub(&ones, &square)?; + let sqrt = denom.sqrt()?; + let frac = arithmetic::div(grad_output, &sqrt)?; + let grad = arithmetic::neg(&frac)?; + gradients.insert(self.input_id, grad); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for inverse tangent +pub struct AtanBackward { + pub input_id: TensorId, + pub input: Tensor, +} + +impl GradientFunction for AtanBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + // d/dx(atan(x)) = grad_output / (1 + x^2) + let square = arithmetic::mul(&self.input, &self.input)?; + let ones = Tensor::ones( + self.input.shape().clone(), + self.input.dtype(), + self.input.device(), + false, + ); + let denom = arithmetic::add(&ones, &square)?; + let grad = arithmetic::div(grad_output, &denom)?; + gradients.insert(self.input_id, grad); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for hyperbolic sine +pub struct SinhBackward { + pub input_id: TensorId, + pub input: Tensor, +} + +impl GradientFunction for SinhBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + // d/dx(sinh(x)) = cosh(x) * grad_output + let cosh_x = self.input.cosh()?; + let grad = arithmetic::mul(&cosh_x, grad_output)?; + gradients.insert(self.input_id, grad); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for hyperbolic cosine +pub struct CoshBackward { + pub input_id: TensorId, + pub input: Tensor, +} + +impl GradientFunction for CoshBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + // d/dx(cosh(x)) = sinh(x) * grad_output + let sinh_x = self.input.sinh()?; + let grad = arithmetic::mul(&sinh_x, grad_output)?; + gradients.insert(self.input_id, grad); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for inverse hyperbolic sine +pub struct AsinhBackward { + pub input_id: TensorId, + pub input: Tensor, +} + +impl GradientFunction for AsinhBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + // d/dx(asinh(x)) = grad_output / sqrt(1 + x^2) + let square = arithmetic::mul(&self.input, &self.input)?; + let ones = Tensor::ones( + self.input.shape().clone(), + self.input.dtype(), + self.input.device(), + false, + ); + let denom = arithmetic::add(&square, &ones)?; + let sqrt = denom.sqrt()?; + let grad = arithmetic::div(grad_output, &sqrt)?; + gradients.insert(self.input_id, grad); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for inverse hyperbolic cosine +pub struct AcoshBackward { + pub input_id: TensorId, + pub input: Tensor, +} + +impl GradientFunction for AcoshBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + // d/dx(acosh(x)) = grad_output / sqrt((x - 1)(x + 1)) + let ones = Tensor::ones( + self.input.shape().clone(), + self.input.dtype(), + self.input.device(), + false, + ); + let x_minus_one = arithmetic::sub(&self.input, &ones)?; + let x_plus_one = arithmetic::add(&self.input, &ones)?; + let product = arithmetic::mul(&x_minus_one, &x_plus_one)?; + let sqrt = product.sqrt()?; + let grad = arithmetic::div(grad_output, &sqrt)?; + gradients.insert(self.input_id, grad); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for inverse hyperbolic tangent +pub struct AtanhBackward { + pub input_id: TensorId, + pub input: Tensor, +} + +impl GradientFunction for AtanhBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + // d/dx(atanh(x)) = grad_output / (1 - x^2) + let square = arithmetic::mul(&self.input, &self.input)?; + let ones = Tensor::ones( + self.input.shape().clone(), + self.input.dtype(), + self.input.device(), + false, + ); + let denom = arithmetic::sub(&ones, &square)?; + let grad = arithmetic::div(grad_output, &denom)?; + gradients.insert(self.input_id, grad); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for tanh +pub struct TanhBackward { + pub input_id: TensorId, + pub output: Tensor, +} + +impl GradientFunction for TanhBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + // d/dx(tanh(x)) = (1 - tanh²(x)) * grad_output + let y2 = arithmetic::mul(&self.output, &self.output)?; + let ones = Tensor::ones( + self.output.shape().clone(), + self.output.dtype(), + self.output.device(), + false, + ); + let term = arithmetic::sub(&ones, &y2)?; + let grad = arithmetic::mul(&term, grad_output)?; + gradients.insert(self.input_id, grad); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for sigmoid +pub struct SigmoidBackward { + pub input_id: TensorId, + pub output: Tensor, +} + +impl GradientFunction for SigmoidBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + // d/dx(sigmoid(x)) = sigmoid(x) * (1 - sigmoid(x)) * grad_output + let ones = Tensor::ones( + self.output.shape().clone(), + self.output.dtype(), + self.output.device(), + false, + ); + let one_minus = arithmetic::sub(&ones, &self.output)?; + let term = arithmetic::mul(&self.output, &one_minus)?; + let grad = arithmetic::mul(&term, grad_output)?; + gradients.insert(self.input_id, grad); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for Softplus +pub struct SoftplusBackward { + pub input_id: TensorId, + pub input: Tensor, + pub beta: f64, + pub threshold: f64, +} diff --git a/engine/src/autograd/mod/core.rs b/engine/src/autograd/mod/core.rs index 6c5b913d..db146a4f 100644 --- a/engine/src/autograd/mod/core.rs +++ b/engine/src/autograd/mod/core.rs @@ -1,906 +1,1084 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -use crate::{ - device::Device, - error::{MinitensorError, Result}, - operations::{activation, arithmetic, linalg, minmax, reduction, selection, shape_ops}, - tensor::{DataType, Shape, Strides, Tensor, TensorData}, -}; -use libm::{erf, erff}; -use rayon::prelude::*; -use rustc_hash::FxHashMap; -use smallvec::SmallVec; -use std::cell::Cell; -use std::sync::Arc; -use std::sync::atomic::{AtomicUsize, Ordering}; - -const PAR_THRESHOLD: usize = 1 << 12; // 4096 elements - -/// Unique identifier for tensors in the computation graph -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] -pub struct TensorId(usize); - -impl TensorId { - /// Create a new unique tensor ID - pub fn new() -> Self { - static COUNTER: AtomicUsize = AtomicUsize::new(0); - Self(COUNTER.fetch_add(1, Ordering::Relaxed)) - } -} - -impl Default for TensorId { - fn default() -> Self { - Self::new() - } -} - -impl std::fmt::Display for TensorId { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!(f, "TensorId({})", self.0) - } -} - -/// Trait for gradient functions in the computation graph -pub trait GradientFunction: Send + Sync { - /// Compute gradients for inputs given the output gradient - fn backward(&self, grad_output: &Tensor) -> Result>; - - /// Get the input tensor IDs that this function depends on - fn input_ids(&self) -> &[TensorId]; - - /// Name of the gradient function used for debugging and introspection - fn name(&self) -> &'static str { - let full = std::any::type_name::(); - match full.rsplit("::").next() { - Some(name) => name, - None => full, - } - } -} - -pub use graph::ComputationGraph; - -// Thread-local computation graph to avoid cross-test interference -thread_local! { - static GLOBAL_GRAPH: std::cell::RefCell = - std::cell::RefCell::new(ComputationGraph::new()); -} - -thread_local! { - static GRAPH_CONSUMED: Cell = Cell::new(false); -} - -/// Add a tensor and its gradient function to the global computation graph -pub fn add_to_graph(tensor: &Tensor, grad_fn: Option>) -> Result<()> { - GLOBAL_GRAPH.with(|graph| { - if let Ok(mut g) = graph.try_borrow_mut() { - g.add_tensor_with_grad_req(tensor.id(), grad_fn, tensor.requires_grad()); - } - }); - reset_graph_consumed(); - Ok(()) -} - -/// Perform backward pass from the given tensor using the global computation graph -pub fn backward( - tensor: &Tensor, - grad_output: Option, -) -> Result> { - GLOBAL_GRAPH.with(|graph| { - let grad = match grad_output { - Some(g) => g, - None => { - if tensor.numel() != 1 { - return Err(MinitensorError::gradient_error( - "Gradient can only be implicitly created for scalar tensors", - )); - } - Tensor::ones( - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - false, - ) - } - }; - graph.borrow_mut().backward(tensor.id(), Some(grad)) - }) -} - -/// Get the gradient for a tensor from the last backward pass -pub fn get_gradient(tensor: &Tensor) -> Option { - GLOBAL_GRAPH.with(|graph| graph.borrow().get_gradient(tensor.id()).cloned()) -} - -/// Clear all stored gradients in the global computation graph -pub fn zero_gradients() { - GLOBAL_GRAPH.with(|graph| graph.borrow_mut().zero_grad()); -} - -/// Clear the global computation graph -pub fn clear_graph() -> Result<()> { - GLOBAL_GRAPH.with(|graph| { - *graph.borrow_mut() = ComputationGraph::new(); - }); - reset_graph_consumed(); - Ok(()) -} - -/// Mark the computation graph as consumed after a backward pass completes. -pub fn mark_graph_consumed() { - GRAPH_CONSUMED.with(|flag| flag.set(true)); -} - -/// Reset the consumed flag so that future backward passes are permitted. -pub fn reset_graph_consumed() { - GRAPH_CONSUMED.with(|flag| flag.set(false)); -} - -/// Query whether the active computation graph has already been consumed. -pub fn is_graph_consumed() -> bool { - GRAPH_CONSUMED.with(|flag| flag.get()) -} - -// Gradient function implementations for common operations - -/// Accumulate a gradient contribution for `input_id` into `gradients`. -/// -/// A single backward pass may produce more than one gradient for the same input -/// when a tensor is used as several operands of one operation (`x * x`, `x + x`, -/// `x.matmul(x)`, `pow(x, x)`, ...). The gradients returned by a -/// [`GradientFunction`] are keyed by [`TensorId`], so a plain `insert` would let -/// the later contribution silently overwrite the earlier one and halve (or worse) -/// the gradient. Summing on collision matches the mathematically correct result -/// and mirrors the cross-node accumulation performed by the graph itself. -#[inline] -fn accumulate_grad( - gradients: &mut FxHashMap, - input_id: TensorId, - grad: Tensor, -) -> Result<()> { - use std::collections::hash_map::Entry; - match gradients.entry(input_id) { - Entry::Occupied(mut existing) => { - arithmetic::add_inplace(existing.get_mut(), &grad)?; - } - Entry::Vacant(slot) => { - slot.insert(grad); - } - } - Ok(()) -} - -/// Gradient function for tensor cloning operation -pub struct CloneBackward { - pub input_id: TensorId, -} - -impl GradientFunction for CloneBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - gradients.insert(self.input_id, grad_output.deep_clone()?); - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for addition operation -pub struct AddBackward { - pub input_shapes: [Vec; 2], - pub input_ids: [TensorId; 2], -} - -impl GradientFunction for AddBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(2); - - // For addition, gradients flow through unchanged, but we need to handle broadcasting - let lhs_shape = Shape::new(self.input_shapes[0].clone()); - let rhs_shape = Shape::new(self.input_shapes[1].clone()); - - // Reduce gradients to match input shapes if broadcasting occurred - let lhs_grad = reduce_gradient_for_broadcasting(grad_output, &lhs_shape)?; - let rhs_grad = reduce_gradient_for_broadcasting(grad_output, &rhs_shape)?; - - accumulate_grad(&mut gradients, self.input_ids[0], lhs_grad)?; - accumulate_grad(&mut gradients, self.input_ids[1], rhs_grad)?; - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} - -/// Gradient function for subtraction operation -pub struct SubBackward { - pub input_shapes: [Vec; 2], - pub input_ids: [TensorId; 2], -} - -impl GradientFunction for SubBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(2); - - let lhs_shape = Shape::new(self.input_shapes[0].clone()); - let rhs_shape = Shape::new(self.input_shapes[1].clone()); - - let lhs_grad = reduce_gradient_for_broadcasting(grad_output, &lhs_shape)?; - let rhs_base = reduce_gradient_for_broadcasting(grad_output, &rhs_shape)?; - let rhs_grad = arithmetic::neg(&rhs_base)?; - - accumulate_grad(&mut gradients, self.input_ids[0], lhs_grad)?; - accumulate_grad(&mut gradients, self.input_ids[1], rhs_grad)?; - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} - -/// Gradient function for multiplication operation -pub struct MulBackward { - pub lhs: Tensor, - pub rhs: Tensor, - pub input_ids: [TensorId; 2], -} - -impl GradientFunction for MulBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(2); - - // d/dx(x*y) = y and d/dy(x*y) = x - let lhs_term = arithmetic::mul(grad_output, &self.rhs.detach())?; - let rhs_term = arithmetic::mul(grad_output, &self.lhs.detach())?; - - let lhs_grad = reduce_gradient_for_broadcasting(&lhs_term, self.lhs.shape())?; - let rhs_grad = reduce_gradient_for_broadcasting(&rhs_term, self.rhs.shape())?; - - accumulate_grad(&mut gradients, self.input_ids[0], lhs_grad)?; - accumulate_grad(&mut gradients, self.input_ids[1], rhs_grad)?; - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} - -/// Gradient function for division operation -pub struct DivBackward { - pub lhs: Tensor, - pub rhs: Tensor, - pub input_ids: [TensorId; 2], -} - -impl GradientFunction for DivBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(2); - - // d/dx(x/y) = 1 / y - let rhs_inv = arithmetic::div( - &Tensor::ones( - self.rhs.shape().clone(), - self.rhs.dtype(), - self.rhs.device(), - false, - ), - &self.rhs.detach(), - )?; - let lhs_term = arithmetic::mul(grad_output, &rhs_inv)?; - let lhs_grad = reduce_gradient_for_broadcasting(&lhs_term, self.lhs.shape())?; - - // d/dy(x/y) = -x / y^2 - let num = arithmetic::mul(grad_output, &self.lhs.detach())?; - let rhs_sq = arithmetic::mul(&self.rhs.detach(), &self.rhs.detach())?; - let rhs_term = arithmetic::div(&num, &rhs_sq)?; - let rhs_term = arithmetic::neg(&rhs_term)?; - let rhs_grad = reduce_gradient_for_broadcasting(&rhs_term, self.rhs.shape())?; - - accumulate_grad(&mut gradients, self.input_ids[0], lhs_grad)?; - accumulate_grad(&mut gradients, self.input_ids[1], rhs_grad)?; - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} - -/// Gradient function for where/select operation -pub struct WhereBackward { - pub condition: Tensor, - pub input_shape: Vec, - pub other_shape: Vec, - pub input_requires_grad: bool, - pub other_requires_grad: bool, - pub input_ids: [TensorId; 2], -} - -impl GradientFunction for WhereBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(self.input_requires_grad as usize + self.other_requires_grad as usize); - - let mut zero_tensor: Option = None; - - if self.input_requires_grad { - let zeros = zero_tensor.get_or_insert_with(|| { - Tensor::zeros( - grad_output.shape().clone(), - grad_output.dtype(), - grad_output.device(), - false, - ) - }); - let selected = selection::where_op(&self.condition, grad_output, zeros)?; - let reduced = - reduce_gradient_for_broadcasting(&selected, &Shape::new(self.input_shape.clone()))?; - accumulate_grad(&mut gradients, self.input_ids[0], reduced)?; - } - - if self.other_requires_grad { - let zeros = zero_tensor.get_or_insert_with(|| { - Tensor::zeros( - grad_output.shape().clone(), - grad_output.dtype(), - grad_output.device(), - false, - ) - }); - let selected = selection::where_op(&self.condition, zeros, grad_output)?; - let reduced = - reduce_gradient_for_broadcasting(&selected, &Shape::new(self.other_shape.clone()))?; - accumulate_grad(&mut gradients, self.input_ids[1], reduced)?; - } - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} - -/// Gradient function for diagonal extraction. -pub struct DiagonalBackward { - pub input_shape: Vec, - pub input_strides: Vec, - pub input_dtype: DataType, - pub dim1: usize, - pub dim2: usize, - pub offset: isize, - pub input_requires_grad: bool, - pub input_id: TensorId, -} - -impl GradientFunction for DiagonalBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - - if !self.input_requires_grad { - return Ok(gradients); - } - - if grad_output.dtype() != self.input_dtype { - return Err(MinitensorError::type_mismatch( - format!("{:?}", grad_output.dtype()), - format!("{:?}", self.input_dtype), - )); - } - - let spec = linalg::compute_diagonal_spec( - &self.input_shape, - &self.input_strides, - self.dim1, - self.dim2, - self.offset, - )?; - - if grad_output.shape().dims() != spec.output_dims { - return Err(MinitensorError::shape_mismatch( - grad_output.shape().dims().to_vec(), - spec.output_dims.clone(), - )); - } - - let numel = self.input_shape.iter().product(); - let mut grad_data = - TensorData::zeros_on_device(numel, self.input_dtype, grad_output.device()); - - match self.input_dtype { - DataType::Float32 => { - let grad_out = grad_output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice for diagonal backward") - })?; - let grad_in = grad_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice for diagonal backward", - ) - })?; - linalg::diagonal_scatter( - grad_out, - grad_in, - &self.input_shape, - &self.input_strides, - &spec, - ); - } - DataType::Float64 => { - let grad_out = grad_output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice for diagonal backward") - })?; - let grad_in = grad_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice for diagonal backward", - ) - })?; - linalg::diagonal_scatter( - grad_out, - grad_in, - &self.input_shape, - &self.input_strides, - &spec, - ); - } - DataType::Int32 => { - let grad_out = grad_output.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice for diagonal backward") - })?; - let grad_in = grad_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable i32 slice for diagonal backward", - ) - })?; - linalg::diagonal_scatter( - grad_out, - grad_in, - &self.input_shape, - &self.input_strides, - &spec, - ); - } - DataType::Int64 => { - let grad_out = grad_output.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice for diagonal backward") - })?; - let grad_in = grad_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable i64 slice for diagonal backward", - ) - })?; - linalg::diagonal_scatter( - grad_out, - grad_in, - &self.input_shape, - &self.input_strides, - &spec, - ); - } - DataType::Bool => { - return Err(MinitensorError::invalid_operation( - "diagonal backward is not defined for bool tensors", - )); - } - } - - let grad_tensor = Tensor::new( - Arc::new(grad_data), - Shape::new(self.input_shape.clone()), - self.input_dtype, - grad_output.device(), - false, - ); - gradients.insert(self.input_id, grad_tensor); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for triangular masking operations (triu/tril) -pub struct TriangularBackward { - pub input_shape: Vec, - pub diagonal: isize, - pub upper: bool, - pub input_requires_grad: bool, - pub input_id: TensorId, -} - -impl GradientFunction for TriangularBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - - if self.input_requires_grad { - if grad_output.shape().dims() != self.input_shape { - return Err(MinitensorError::shape_mismatch( - grad_output.shape().dims().to_vec(), - self.input_shape.clone(), - )); - } - - let mut grad_data = TensorData::uninitialized_on_device( - grad_output.numel(), - grad_output.dtype(), - grad_output.device(), - ); - linalg::apply_triangular_mask(grad_output, &mut grad_data, self.diagonal, self.upper)?; - let grad = Tensor::new( - Arc::new(grad_data), - grad_output.shape().clone(), - grad_output.dtype(), - grad_output.device(), - false, - ); - gradients.insert(self.input_id, grad); - } - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for element-wise maximum operation -pub struct MaximumBackward { - pub lhs: Tensor, - pub rhs: Tensor, - pub input_shapes: [Vec; 2], - pub input_requires_grad: [bool; 2], - pub input_ids: [TensorId; 2], -} - -impl GradientFunction for MaximumBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(self.input_requires_grad.iter().filter(|&&b| b).count()); - - if !self.input_requires_grad[0] && !self.input_requires_grad[1] { - return Ok(gradients); - } - - let mask = minmax::maximum_backward_mask(&self.lhs, &self.rhs)?; - let mut zeros: Option = None; - - if self.input_requires_grad[0] { - let zero = zeros.get_or_insert_with(|| { - Tensor::zeros( - grad_output.shape().clone(), - grad_output.dtype(), - grad_output.device(), - false, - ) - }); - let selected = minmax::select_with_mask(&mask, grad_output, zero)?; - let reduced = reduce_gradient_for_broadcasting( - &selected, - &Shape::new(self.input_shapes[0].clone()), - )?; - accumulate_grad(&mut gradients, self.input_ids[0], reduced)?; - } - - if self.input_requires_grad[1] { - let zero = zeros.get_or_insert_with(|| { - Tensor::zeros( - grad_output.shape().clone(), - grad_output.dtype(), - grad_output.device(), - false, - ) - }); - let selected = minmax::select_with_mask(&mask, zero, grad_output)?; - let reduced = reduce_gradient_for_broadcasting( - &selected, - &Shape::new(self.input_shapes[1].clone()), - )?; - accumulate_grad(&mut gradients, self.input_ids[1], reduced)?; - } - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} - -/// Gradient function for element-wise minimum operation -pub struct MinimumBackward { - pub lhs: Tensor, - pub rhs: Tensor, - pub input_shapes: [Vec; 2], - pub input_requires_grad: [bool; 2], - pub input_ids: [TensorId; 2], -} - -impl GradientFunction for MinimumBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(self.input_requires_grad.iter().filter(|&&b| b).count()); - - if !self.input_requires_grad[0] && !self.input_requires_grad[1] { - return Ok(gradients); - } - - let mask = minmax::minimum_backward_mask(&self.lhs, &self.rhs)?; - let mut zeros: Option = None; - - if self.input_requires_grad[0] { - let zero = zeros.get_or_insert_with(|| { - Tensor::zeros( - grad_output.shape().clone(), - grad_output.dtype(), - grad_output.device(), - false, - ) - }); - let selected = minmax::select_with_mask(&mask, grad_output, zero)?; - let reduced = reduce_gradient_for_broadcasting( - &selected, - &Shape::new(self.input_shapes[0].clone()), - )?; - accumulate_grad(&mut gradients, self.input_ids[0], reduced)?; - } - - if self.input_requires_grad[1] { - let zero = zeros.get_or_insert_with(|| { - Tensor::zeros( - grad_output.shape().clone(), - grad_output.dtype(), - grad_output.device(), - false, - ) - }); - let selected = minmax::select_with_mask(&mask, zero, grad_output)?; - let reduced = reduce_gradient_for_broadcasting( - &selected, - &Shape::new(self.input_shapes[1].clone()), - )?; - accumulate_grad(&mut gradients, self.input_ids[1], reduced)?; - } - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} - -/// Gradient function for dot product -pub struct DotBackward { - pub lhs: Tensor, - pub rhs: Tensor, - pub input_ids: [TensorId; 2], - pub lhs_requires_grad: bool, - pub rhs_requires_grad: bool, -} - -impl GradientFunction for DotBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve((self.lhs_requires_grad as usize) + (self.rhs_requires_grad as usize)); - - if self.lhs_requires_grad { - let grad = crate::operations::arithmetic::mul(&self.rhs, grad_output)?; - accumulate_grad(&mut gradients, self.input_ids[0], grad)?; - } - - if self.rhs_requires_grad { - let grad = crate::operations::arithmetic::mul(&self.lhs, grad_output)?; - accumulate_grad(&mut gradients, self.input_ids[1], grad)?; - } - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} - -/// Gradient function for negation -pub struct NegBackward { - pub input_id: TensorId, -} - -impl GradientFunction for NegBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - let grad = arithmetic::neg(grad_output)?; - gradients.insert(self.input_id, grad); - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for matrix multiplication -pub struct MatMulBackward { - pub lhs: Tensor, - pub rhs: Tensor, - pub input_ids: [TensorId; 2], - pub lhs_requires_grad: bool, - pub rhs_requires_grad: bool, -} - -impl GradientFunction for MatMulBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve((self.lhs_requires_grad as usize) + (self.rhs_requires_grad as usize)); - - if self.lhs.ndim() < 2 || self.rhs.ndim() < 2 { - return Err(MinitensorError::invalid_operation( - "MatMulBackward requires tensors with at least 2 dimensions", - )); - } - - if self.lhs_requires_grad { - let rhs_t = crate::operations::linalg::transpose( - &self.rhs, - (self.rhs.ndim() - 2) as isize, - (self.rhs.ndim() - 1) as isize, - )?; - let lhs_grad = crate::operations::linalg::matmul(grad_output, &rhs_t)?; - accumulate_grad(&mut gradients, self.input_ids[0], lhs_grad)?; - } - - if self.rhs_requires_grad { - let lhs_t = crate::operations::linalg::transpose( - &self.lhs, - (self.lhs.ndim() - 2) as isize, - (self.lhs.ndim() - 1) as isize, - )?; - let rhs_grad = crate::operations::linalg::matmul(&lhs_t, grad_output)?; - accumulate_grad(&mut gradients, self.input_ids[1], rhs_grad)?; - } - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} - -/// Gradient function for solving linear systems. -pub struct SolveBackward { - pub lhs: Tensor, - pub solution: Tensor, - pub input_ids: [TensorId; 2], - pub lhs_requires_grad: bool, - pub rhs_requires_grad: bool, -} - -impl GradientFunction for SolveBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve((self.lhs_requires_grad as usize) + (self.rhs_requires_grad as usize)); - - let lhs_t = crate::operations::linalg::transpose( - &self.lhs, - (self.lhs.ndim() - 2) as isize, - (self.lhs.ndim() - 1) as isize, - )?; - - if self.rhs_requires_grad { - let grad_rhs = crate::operations::linalg::solve(&lhs_t, grad_output)?; - accumulate_grad(&mut gradients, self.input_ids[1], grad_rhs)?; - } - - if self.lhs_requires_grad { - let solution_view = if self.solution.ndim() == self.lhs.ndim() - 1 { - crate::operations::shape_ops::unsqueeze( - &self.solution, - self.solution.ndim() as isize, - )? - } else { - self.solution.clone() - }; - - let grad_output_view = if grad_output.ndim() == self.lhs.ndim() - 1 { - crate::operations::shape_ops::unsqueeze(grad_output, grad_output.ndim() as isize)? - } else { - grad_output.clone() - }; - - let solution_t = crate::operations::linalg::transpose( - &solution_view, - (solution_view.ndim() - 2) as isize, - (solution_view.ndim() - 1) as isize, - )?; - let gram = crate::operations::linalg::matmul(&grad_output_view, &solution_t)?; - let lhs_grad = crate::operations::linalg::solve(&lhs_t, &gram)?; - let lhs_grad = crate::operations::arithmetic::neg(&lhs_grad)?; - accumulate_grad(&mut gradients, self.input_ids[0], lhs_grad)?; - } - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} - -/// Gradient function for transpose operation -pub struct TransposeBackward { - pub dims: Vec, - pub input_id: TensorId, -} - -impl GradientFunction for TransposeBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // Transpose gradient: transpose back. Support both simple swaps and - // arbitrary dimension permutations by applying the inverse permutation. - let grad_input = if self.dims.len() == 2 { - crate::operations::linalg::transpose( - grad_output, - self.dims[0] as isize, - self.dims[1] as isize, - )? - } else { - let mut inverse = vec![0; self.dims.len()]; - for (i, &d) in self.dims.iter().enumerate() { - inverse[d] = i; - } - let mut grad = grad_output.clone(); - let mut current: Vec = (0..inverse.len()).collect(); - for i in 0..inverse.len() { - let j = current - .iter() - .position(|&x| x == inverse[i]) - .expect("invalid permutation"); - if i != j { - grad = crate::operations::linalg::transpose(&grad, i as isize, j as isize)?; - current.swap(i, j); - } - } - grad - }; - - gradients.insert(self.input_id, grad_input); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for sum reduction -pub struct SumBackward { - pub input_id: TensorId, - pub input_shape: Vec, - pub dims: Option>, - pub keepdim: bool, -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use crate::{ + error::{MinitensorError, Result}, + operations::{arithmetic, linalg, minmax, reduction, selection}, + tensor::{DataType, Shape, Tensor, TensorData}, +}; +use rustc_hash::FxHashMap; +use smallvec::SmallVec; +use std::cell::Cell; +use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; + +pub(crate) const PAR_THRESHOLD: usize = 1 << 12; // 4096 elements + +/// Unique identifier for tensors in the computation graph +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub struct TensorId(usize); + +impl TensorId { + /// Create a new unique tensor ID + pub fn new() -> Self { + static COUNTER: AtomicUsize = AtomicUsize::new(0); + Self(COUNTER.fetch_add(1, Ordering::Relaxed)) + } +} + +impl Default for TensorId { + fn default() -> Self { + Self::new() + } +} + +impl std::fmt::Display for TensorId { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "TensorId({})", self.0) + } +} + +/// Trait for gradient functions in the computation graph +pub trait GradientFunction: Send + Sync { + /// Compute gradients for inputs given the output gradient + fn backward(&self, grad_output: &Tensor) -> Result>; + + /// Get the input tensor IDs that this function depends on + fn input_ids(&self) -> &[TensorId]; + + /// Name of the gradient function used for debugging and introspection + fn name(&self) -> &'static str { + let full = std::any::type_name::(); + match full.rsplit("::").next() { + Some(name) => name, + None => full, + } + } +} + +pub use super::graph::{BackwardStep, ComputationGraph, execute_backward_plan}; + +// Thread-local computation graph to avoid cross-test interference +thread_local! { + static GLOBAL_GRAPH: std::cell::RefCell = + std::cell::RefCell::new(ComputationGraph::new()); +} + +thread_local! { + static GRAPH_CONSUMED: Cell = const { Cell::new(false) }; +} + +// Thread-local gradient recording mode. While disabled, `add_to_graph` is a +// no-op, so tensor operations do not record autograd nodes. The backward pass +// disables recording while it executes gradient kernels; previously this was +// enforced only by the accident that the graph happened to be borrowed, which +// silently dropped registrations instead of making the policy explicit. +thread_local! { + static GRAD_ENABLED: Cell = const { Cell::new(true) }; +} + +/// Query whether autograd recording is currently enabled on this thread. +pub fn is_grad_enabled() -> bool { + GRAD_ENABLED.with(|flag| flag.get()) +} + +/// Enable or disable autograd recording on this thread, returning the +/// previous state. Building block for user-facing `no_grad()` / +/// `enable_grad()` context managers. +pub fn set_grad_enabled(enabled: bool) -> bool { + GRAD_ENABLED.with(|flag| flag.replace(enabled)) +} + +/// RAII guard that disables autograd recording for its lifetime. +/// +/// Used internally by the backward pass; also usable as a building block for a +/// user-facing `no_grad` mode. +pub struct NoGradGuard { + prev: bool, +} + +impl NoGradGuard { + pub fn new() -> Self { + let prev = GRAD_ENABLED.with(|flag| flag.replace(false)); + Self { prev } + } +} + +impl Default for NoGradGuard { + fn default() -> Self { + Self::new() + } +} + +impl Drop for NoGradGuard { + fn drop(&mut self) { + let prev = self.prev; + GRAD_ENABLED.with(|flag| flag.set(prev)); + } +} + +/// Add a tensor and its gradient function to the global computation graph +pub fn add_to_graph(tensor: &Tensor, grad_fn: Option>) -> Result<()> { + if !is_grad_enabled() { + return Ok(()); + } + GLOBAL_GRAPH.with(|graph| { + graph + .borrow_mut() + .add_tensor_with_grad_req(tensor.id(), grad_fn, tensor.requires_grad()); + }); + reset_graph_consumed(); + Ok(()) +} + +fn implicit_gradient(tensor: &Tensor, grad_output: Option) -> Result { + match grad_output { + Some(g) => Ok(g), + None => { + if tensor.numel() != 1 { + return Err(MinitensorError::gradient_error( + "Gradient can only be implicitly created for scalar tensors", + )); + } + Ok(Tensor::ones( + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + false, + )) + } + } +} + +/// Perform backward pass from the given tensor using the global computation +/// graph. Gradients are stored in the graph and can be read individually with +/// [`get_gradient`]; nothing is cloned on this path. +/// +/// The graph is only borrowed to plan the pass and to store the results, so +/// gradient kernels run without holding the thread-local borrow. +pub fn backward(tensor: &Tensor, grad_output: Option) -> Result<()> { + let grad = implicit_gradient(tensor, grad_output)?; + + let plan = GLOBAL_GRAPH.with(|graph| { + let graph = graph.borrow(); + if !graph.contains_tensor(tensor.id()) && is_graph_consumed() { + return Err(MinitensorError::gradient_error_with_suggestion( + "Computation graph for this tensor has already been freed", + "Re-run the forward pass or call backward(retain_graph=True)", + None, + )); + } + graph.plan_backward(tensor.id()) + })?; + + let gradients = { + // Gradient kernels must not record new autograd nodes. + let _guard = NoGradGuard::new(); + execute_backward_plan(&plan, tensor.id(), grad)? + }; + + GLOBAL_GRAPH.with(|graph| graph.borrow_mut().set_gradients(gradients)); + Ok(()) +} + +/// Perform a backward pass and return a snapshot of every gradient computed. +/// +/// This clones the full gradient map and exists for tests and diagnostics; +/// production code should call [`backward`] and read the gradients it needs +/// via [`get_gradient`]. +pub fn backward_collect( + tensor: &Tensor, + grad_output: Option, +) -> Result> { + backward(tensor, grad_output)?; + Ok(GLOBAL_GRAPH.with(|graph| graph.borrow().gradients_snapshot())) +} + +/// Release the autograd nodes (and the tensors they saved for backward) +/// reachable from `tensor`. Stored gradients remain available. Called by the +/// bindings after a non-retaining backward pass so saved activations are freed +/// immediately rather than at the next optimizer step. +pub fn release_saved_subgraph(tensor: &Tensor) { + GLOBAL_GRAPH.with(|graph| graph.borrow_mut().release_saved_subgraph(tensor.id())); +} + +/// Get the gradient for a tensor from the last backward pass +pub fn get_gradient(tensor: &Tensor) -> Option { + GLOBAL_GRAPH.with(|graph| graph.borrow().get_gradient(tensor.id()).cloned()) +} + +/// Clear all stored gradients in the global computation graph +pub fn zero_gradients() { + GLOBAL_GRAPH.with(|graph| graph.borrow_mut().zero_grad()); +} + +/// Remove the stored gradient for a single tensor from the global graph. +pub fn clear_gradient(tensor: &Tensor) -> Option { + GLOBAL_GRAPH.with(|graph| graph.borrow_mut().remove_gradient(tensor.id())) +} + +/// Clear the global computation graph +pub fn clear_graph() -> Result<()> { + GLOBAL_GRAPH.with(|graph| { + *graph.borrow_mut() = ComputationGraph::new(); + }); + reset_graph_consumed(); + Ok(()) +} + +/// Mark the computation graph as consumed after a backward pass completes. +pub fn mark_graph_consumed() { + GRAPH_CONSUMED.with(|flag| flag.set(true)); +} + +/// Reset the consumed flag so that future backward passes are permitted. +pub fn reset_graph_consumed() { + GRAPH_CONSUMED.with(|flag| flag.set(false)); +} + +/// Query whether the active computation graph has already been consumed. +pub fn is_graph_consumed() -> bool { + GRAPH_CONSUMED.with(|flag| flag.get()) +} + +/// Helper function to reduce gradients for broadcasting +pub(crate) fn reduce_gradient_for_broadcasting( + grad_output: &Tensor, + target_shape: &Shape, +) -> Result { + if grad_output.shape() == target_shape { + return Ok(grad_output.clone()); + } + + let grad_dims = grad_output.shape().dims(); + let target_dims = target_shape.dims(); + if target_dims.len() > grad_dims.len() { + return Err(MinitensorError::BroadcastError { + shape1: grad_dims.to_vec(), + shape2: target_dims.to_vec(), + suggestion: Some( + "Ensure the target shape has no more dimensions than the gradient output." + .to_string(), + ), + context: Some("reduce_gradient_for_broadcasting".to_string()), + }); + } + let extra = grad_dims.len() - target_dims.len(); + + // Use a stack-allocated small vector and pre-allocate enough capacity to + // hold all potential broadcast axes. This avoids repeated reallocations for + // higher dimensional tensors. + let mut axes_to_sum: SmallVec<[usize; 8]> = SmallVec::with_capacity(grad_dims.len()); + axes_to_sum.extend(0..extra); + for i in 0..target_dims.len() { + let gdim = grad_dims[extra + i]; + let tdim = target_dims[i]; + if tdim == 1 { + if gdim != 1 { + axes_to_sum.push(extra + i); + } + } else if gdim != tdim { + return Err(MinitensorError::BroadcastError { + shape1: grad_dims.to_vec(), + shape2: target_dims.to_vec(), + suggestion: Some( + "Ensure each target dimension is 1 or matches the gradient dimension." + .to_string(), + ), + context: Some("reduce_gradient_for_broadcasting".to_string()), + }); + } + } + + if axes_to_sum.is_empty() { + return Ok(grad_output.clone()); + } + + let mut axes = Vec::with_capacity(axes_to_sum.len()); + for axis in axes_to_sum { + axes.push(axis as isize); + } + let mut grad = reduction::sum(grad_output, Some(axes), true)?; + + if grad.shape() != target_shape { + grad = grad.view(target_shape.clone())?; + } + + Ok(grad) +} + +// Gradient function implementations for common operations + +/// Accumulate a gradient contribution for `input_id` into `gradients`. +/// +/// A single backward pass may produce more than one gradient for the same input +/// when a tensor is used as several operands of one operation (`x * x`, `x + x`, +/// `x.matmul(x)`, `pow(x, x)`, ...). The gradients returned by a +/// [`GradientFunction`] are keyed by [`TensorId`], so a plain `insert` would let +/// the later contribution silently overwrite the earlier one and halve (or worse) +/// the gradient. Summing on collision matches the mathematically correct result +/// and mirrors the cross-node accumulation performed by the graph itself. +#[inline] +pub(crate) fn accumulate_grad( + gradients: &mut FxHashMap, + input_id: TensorId, + grad: Tensor, +) -> Result<()> { + use std::collections::hash_map::Entry; + match gradients.entry(input_id) { + Entry::Occupied(mut existing) => { + arithmetic::add_inplace(existing.get_mut(), &grad)?; + } + Entry::Vacant(slot) => { + slot.insert(grad); + } + } + Ok(()) +} + +/// Gradient function for tensor cloning operation +pub struct CloneBackward { + pub input_id: TensorId, +} + +impl GradientFunction for CloneBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + gradients.insert(self.input_id, grad_output.deep_clone()?); + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for addition operation +pub struct AddBackward { + pub input_shapes: [Vec; 2], + pub input_ids: [TensorId; 2], + /// Which inputs actually need a gradient; contributions for frozen + /// inputs are skipped entirely (no broadcast reduction, no map entry). + pub input_requires_grad: [bool; 2], +} + +impl GradientFunction for AddBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(2); + + // For addition, gradients flow through unchanged, but broadcasting must + // be undone. Inputs that do not require a gradient are skipped. + if self.input_requires_grad[0] { + let lhs_shape = Shape::new(self.input_shapes[0].clone()); + let lhs_grad = reduce_gradient_for_broadcasting(grad_output, &lhs_shape)?; + accumulate_grad(&mut gradients, self.input_ids[0], lhs_grad)?; + } + if self.input_requires_grad[1] { + let rhs_shape = Shape::new(self.input_shapes[1].clone()); + let rhs_grad = reduce_gradient_for_broadcasting(grad_output, &rhs_shape)?; + accumulate_grad(&mut gradients, self.input_ids[1], rhs_grad)?; + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} + +/// Gradient function for subtraction operation +pub struct SubBackward { + pub input_shapes: [Vec; 2], + pub input_ids: [TensorId; 2], + /// Which inputs actually need a gradient (see [`AddBackward`]). + pub input_requires_grad: [bool; 2], +} + +impl GradientFunction for SubBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(2); + + if self.input_requires_grad[0] { + let lhs_shape = Shape::new(self.input_shapes[0].clone()); + let lhs_grad = reduce_gradient_for_broadcasting(grad_output, &lhs_shape)?; + accumulate_grad(&mut gradients, self.input_ids[0], lhs_grad)?; + } + if self.input_requires_grad[1] { + let rhs_shape = Shape::new(self.input_shapes[1].clone()); + let rhs_base = reduce_gradient_for_broadcasting(grad_output, &rhs_shape)?; + let rhs_grad = arithmetic::neg(&rhs_base)?; + accumulate_grad(&mut gradients, self.input_ids[1], rhs_grad)?; + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} + +/// Gradient function for multiplication operation +pub struct MulBackward { + pub lhs: Tensor, + pub rhs: Tensor, + pub input_ids: [TensorId; 2], + /// Which inputs actually need a gradient (see [`AddBackward`]). + pub input_requires_grad: [bool; 2], +} + +impl GradientFunction for MulBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(2); + + // d/dx(x*y) = y and d/dy(x*y) = x; skip frozen inputs entirely. + if self.input_requires_grad[0] { + let lhs_term = arithmetic::mul(grad_output, &self.rhs.detach())?; + let lhs_grad = reduce_gradient_for_broadcasting(&lhs_term, self.lhs.shape())?; + accumulate_grad(&mut gradients, self.input_ids[0], lhs_grad)?; + } + if self.input_requires_grad[1] { + let rhs_term = arithmetic::mul(grad_output, &self.lhs.detach())?; + let rhs_grad = reduce_gradient_for_broadcasting(&rhs_term, self.rhs.shape())?; + accumulate_grad(&mut gradients, self.input_ids[1], rhs_grad)?; + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} + +/// Gradient function for division operation +pub struct DivBackward { + pub lhs: Tensor, + pub rhs: Tensor, + pub input_ids: [TensorId; 2], + /// Which inputs actually need a gradient (see [`AddBackward`]). + pub input_requires_grad: [bool; 2], +} + +impl GradientFunction for DivBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(2); + + let rhs = self.rhs.detach(); + + // grad/y is needed by both branches; compute it once if either input + // requires a gradient. + let grad_over_rhs = arithmetic::div(grad_output, &rhs)?; + + // d/dx(x/y) = 1/y => grad_x = grad / y + if self.input_requires_grad[0] { + let lhs_grad = reduce_gradient_for_broadcasting(&grad_over_rhs, self.lhs.shape())?; + accumulate_grad(&mut gradients, self.input_ids[0], lhs_grad)?; + } + + // d/dy(x/y) = -x/y^2 => grad_y = -(grad / y) * (x / y) + if self.input_requires_grad[1] { + let lhs_over_rhs = arithmetic::div(&self.lhs.detach(), &rhs)?; + let rhs_term = arithmetic::mul(&grad_over_rhs, &lhs_over_rhs)?; + let rhs_term = arithmetic::neg(&rhs_term)?; + let rhs_grad = reduce_gradient_for_broadcasting(&rhs_term, self.rhs.shape())?; + accumulate_grad(&mut gradients, self.input_ids[1], rhs_grad)?; + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} + +/// Gradient function for where/select operation +pub struct WhereBackward { + pub condition: Tensor, + pub input_shape: Vec, + pub other_shape: Vec, + pub input_requires_grad: bool, + pub other_requires_grad: bool, + pub input_ids: [TensorId; 2], +} + +impl GradientFunction for WhereBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(self.input_requires_grad as usize + self.other_requires_grad as usize); + + let mut zero_tensor: Option = None; + + if self.input_requires_grad { + let zeros = zero_tensor.get_or_insert_with(|| { + Tensor::zeros( + grad_output.shape().clone(), + grad_output.dtype(), + grad_output.device(), + false, + ) + }); + let selected = selection::where_op(&self.condition, grad_output, zeros)?; + let reduced = + reduce_gradient_for_broadcasting(&selected, &Shape::new(self.input_shape.clone()))?; + accumulate_grad(&mut gradients, self.input_ids[0], reduced)?; + } + + if self.other_requires_grad { + let zeros = zero_tensor.get_or_insert_with(|| { + Tensor::zeros( + grad_output.shape().clone(), + grad_output.dtype(), + grad_output.device(), + false, + ) + }); + let selected = selection::where_op(&self.condition, zeros, grad_output)?; + let reduced = + reduce_gradient_for_broadcasting(&selected, &Shape::new(self.other_shape.clone()))?; + accumulate_grad(&mut gradients, self.input_ids[1], reduced)?; + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} + +/// Gradient function for diagonal extraction. +pub struct DiagonalBackward { + pub input_shape: Vec, + pub input_strides: Vec, + pub input_dtype: DataType, + pub dim1: usize, + pub dim2: usize, + pub offset: isize, + pub input_requires_grad: bool, + pub input_id: TensorId, +} + +impl GradientFunction for DiagonalBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + + if !self.input_requires_grad { + return Ok(gradients); + } + + if grad_output.dtype() != self.input_dtype { + return Err(MinitensorError::type_mismatch( + format!("{:?}", grad_output.dtype()), + format!("{:?}", self.input_dtype), + )); + } + + let spec = linalg::compute_diagonal_spec( + &self.input_shape, + &self.input_strides, + self.dim1, + self.dim2, + self.offset, + )?; + + if grad_output.shape().dims() != spec.output_dims { + return Err(MinitensorError::shape_mismatch( + grad_output.shape().dims().to_vec(), + spec.output_dims.clone(), + )); + } + + let numel = self.input_shape.iter().product(); + let mut grad_data = + TensorData::zeros_on_device(numel, self.input_dtype, grad_output.device()); + + match self.input_dtype { + DataType::Float32 => { + let grad_out = grad_output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice for diagonal backward") + })?; + let grad_in = grad_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice for diagonal backward", + ) + })?; + linalg::diagonal_scatter( + grad_out, + grad_in, + &self.input_shape, + &self.input_strides, + &spec, + ); + } + DataType::Float64 => { + let grad_out = grad_output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice for diagonal backward") + })?; + let grad_in = grad_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice for diagonal backward", + ) + })?; + linalg::diagonal_scatter( + grad_out, + grad_in, + &self.input_shape, + &self.input_strides, + &spec, + ); + } + DataType::Int32 => { + let grad_out = grad_output.data().as_i32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i32 slice for diagonal backward") + })?; + let grad_in = grad_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable i32 slice for diagonal backward", + ) + })?; + linalg::diagonal_scatter( + grad_out, + grad_in, + &self.input_shape, + &self.input_strides, + &spec, + ); + } + DataType::Int64 => { + let grad_out = grad_output.data().as_i64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i64 slice for diagonal backward") + })?; + let grad_in = grad_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable i64 slice for diagonal backward", + ) + })?; + linalg::diagonal_scatter( + grad_out, + grad_in, + &self.input_shape, + &self.input_strides, + &spec, + ); + } + DataType::Bool => { + return Err(MinitensorError::invalid_operation( + "diagonal backward is not defined for bool tensors", + )); + } + } + + let grad_tensor = Tensor::new( + Arc::new(grad_data), + Shape::new(self.input_shape.clone()), + self.input_dtype, + grad_output.device(), + false, + ); + gradients.insert(self.input_id, grad_tensor); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for triangular masking operations (triu/tril) +pub struct TriangularBackward { + pub input_shape: Vec, + pub diagonal: isize, + pub upper: bool, + pub input_requires_grad: bool, + pub input_id: TensorId, +} + +impl GradientFunction for TriangularBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + + if self.input_requires_grad { + if grad_output.shape().dims() != self.input_shape { + return Err(MinitensorError::shape_mismatch( + grad_output.shape().dims().to_vec(), + self.input_shape.clone(), + )); + } + + let mut grad_data = TensorData::uninitialized_on_device( + grad_output.numel(), + grad_output.dtype(), + grad_output.device(), + ); + linalg::apply_triangular_mask(grad_output, &mut grad_data, self.diagonal, self.upper)?; + let grad = Tensor::new( + Arc::new(grad_data), + grad_output.shape().clone(), + grad_output.dtype(), + grad_output.device(), + false, + ); + gradients.insert(self.input_id, grad); + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for element-wise maximum operation +pub struct MaximumBackward { + pub lhs: Tensor, + pub rhs: Tensor, + pub input_shapes: [Vec; 2], + pub input_requires_grad: [bool; 2], + pub input_ids: [TensorId; 2], +} + +impl GradientFunction for MaximumBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(self.input_requires_grad.iter().filter(|&&b| b).count()); + + if !self.input_requires_grad[0] && !self.input_requires_grad[1] { + return Ok(gradients); + } + + let mask = minmax::maximum_backward_mask(&self.lhs, &self.rhs)?; + let mut zeros: Option = None; + + if self.input_requires_grad[0] { + let zero = zeros.get_or_insert_with(|| { + Tensor::zeros( + grad_output.shape().clone(), + grad_output.dtype(), + grad_output.device(), + false, + ) + }); + let selected = minmax::select_with_mask(&mask, grad_output, zero)?; + let reduced = reduce_gradient_for_broadcasting( + &selected, + &Shape::new(self.input_shapes[0].clone()), + )?; + accumulate_grad(&mut gradients, self.input_ids[0], reduced)?; + } + + if self.input_requires_grad[1] { + let zero = zeros.get_or_insert_with(|| { + Tensor::zeros( + grad_output.shape().clone(), + grad_output.dtype(), + grad_output.device(), + false, + ) + }); + let selected = minmax::select_with_mask(&mask, zero, grad_output)?; + let reduced = reduce_gradient_for_broadcasting( + &selected, + &Shape::new(self.input_shapes[1].clone()), + )?; + accumulate_grad(&mut gradients, self.input_ids[1], reduced)?; + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} + +/// Gradient function for element-wise minimum operation +pub struct MinimumBackward { + pub lhs: Tensor, + pub rhs: Tensor, + pub input_shapes: [Vec; 2], + pub input_requires_grad: [bool; 2], + pub input_ids: [TensorId; 2], +} + +impl GradientFunction for MinimumBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(self.input_requires_grad.iter().filter(|&&b| b).count()); + + if !self.input_requires_grad[0] && !self.input_requires_grad[1] { + return Ok(gradients); + } + + let mask = minmax::minimum_backward_mask(&self.lhs, &self.rhs)?; + let mut zeros: Option = None; + + if self.input_requires_grad[0] { + let zero = zeros.get_or_insert_with(|| { + Tensor::zeros( + grad_output.shape().clone(), + grad_output.dtype(), + grad_output.device(), + false, + ) + }); + let selected = minmax::select_with_mask(&mask, grad_output, zero)?; + let reduced = reduce_gradient_for_broadcasting( + &selected, + &Shape::new(self.input_shapes[0].clone()), + )?; + accumulate_grad(&mut gradients, self.input_ids[0], reduced)?; + } + + if self.input_requires_grad[1] { + let zero = zeros.get_or_insert_with(|| { + Tensor::zeros( + grad_output.shape().clone(), + grad_output.dtype(), + grad_output.device(), + false, + ) + }); + let selected = minmax::select_with_mask(&mask, zero, grad_output)?; + let reduced = reduce_gradient_for_broadcasting( + &selected, + &Shape::new(self.input_shapes[1].clone()), + )?; + accumulate_grad(&mut gradients, self.input_ids[1], reduced)?; + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} + +/// Gradient function for dot product +pub struct DotBackward { + pub lhs: Tensor, + pub rhs: Tensor, + pub input_ids: [TensorId; 2], + pub lhs_requires_grad: bool, + pub rhs_requires_grad: bool, +} + +impl GradientFunction for DotBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve((self.lhs_requires_grad as usize) + (self.rhs_requires_grad as usize)); + + if self.lhs_requires_grad { + let grad = crate::operations::arithmetic::mul(&self.rhs, grad_output)?; + accumulate_grad(&mut gradients, self.input_ids[0], grad)?; + } + + if self.rhs_requires_grad { + let grad = crate::operations::arithmetic::mul(&self.lhs, grad_output)?; + accumulate_grad(&mut gradients, self.input_ids[1], grad)?; + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} + +/// Gradient function for negation +pub struct NegBackward { + pub input_id: TensorId, +} + +impl GradientFunction for NegBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + let grad = arithmetic::neg(grad_output)?; + gradients.insert(self.input_id, grad); + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for matrix multiplication +pub struct MatMulBackward { + pub lhs: Tensor, + pub rhs: Tensor, + pub input_ids: [TensorId; 2], + pub lhs_requires_grad: bool, + pub rhs_requires_grad: bool, +} + +impl GradientFunction for MatMulBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve((self.lhs_requires_grad as usize) + (self.rhs_requires_grad as usize)); + + if self.lhs.ndim() < 2 || self.rhs.ndim() < 2 { + return Err(MinitensorError::invalid_operation( + "MatMulBackward requires tensors with at least 2 dimensions", + )); + } + + if self.lhs_requires_grad { + let rhs_t = crate::operations::linalg::transpose( + &self.rhs, + (self.rhs.ndim() - 2) as isize, + (self.rhs.ndim() - 1) as isize, + )?; + let lhs_grad = crate::operations::linalg::matmul(grad_output, &rhs_t)?; + accumulate_grad(&mut gradients, self.input_ids[0], lhs_grad)?; + } + + if self.rhs_requires_grad { + let lhs_t = crate::operations::linalg::transpose( + &self.lhs, + (self.lhs.ndim() - 2) as isize, + (self.lhs.ndim() - 1) as isize, + )?; + let rhs_grad = crate::operations::linalg::matmul(&lhs_t, grad_output)?; + accumulate_grad(&mut gradients, self.input_ids[1], rhs_grad)?; + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} + +/// Gradient function for solving linear systems. +pub struct SolveBackward { + pub lhs: Tensor, + pub solution: Tensor, + pub input_ids: [TensorId; 2], + pub lhs_requires_grad: bool, + pub rhs_requires_grad: bool, +} + +impl GradientFunction for SolveBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve((self.lhs_requires_grad as usize) + (self.rhs_requires_grad as usize)); + + let lhs_t = crate::operations::linalg::transpose( + &self.lhs, + (self.lhs.ndim() - 2) as isize, + (self.lhs.ndim() - 1) as isize, + )?; + + if self.rhs_requires_grad { + let grad_rhs = crate::operations::linalg::solve(&lhs_t, grad_output)?; + accumulate_grad(&mut gradients, self.input_ids[1], grad_rhs)?; + } + + if self.lhs_requires_grad { + let solution_view = if self.solution.ndim() == self.lhs.ndim() - 1 { + crate::operations::shape_ops::unsqueeze( + &self.solution, + self.solution.ndim() as isize, + )? + } else { + self.solution.clone() + }; + + let grad_output_view = if grad_output.ndim() == self.lhs.ndim() - 1 { + crate::operations::shape_ops::unsqueeze(grad_output, grad_output.ndim() as isize)? + } else { + grad_output.clone() + }; + + let solution_t = crate::operations::linalg::transpose( + &solution_view, + (solution_view.ndim() - 2) as isize, + (solution_view.ndim() - 1) as isize, + )?; + let gram = crate::operations::linalg::matmul(&grad_output_view, &solution_t)?; + let lhs_grad = crate::operations::linalg::solve(&lhs_t, &gram)?; + let lhs_grad = crate::operations::arithmetic::neg(&lhs_grad)?; + accumulate_grad(&mut gradients, self.input_ids[0], lhs_grad)?; + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} + +/// Gradient function for transpose operation +pub struct TransposeBackward { + pub dims: Vec, + pub input_id: TensorId, +} + +impl GradientFunction for TransposeBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + // Transpose gradient: transpose back. Support both simple swaps and + // arbitrary dimension permutations by applying the inverse permutation. + let grad_input = if self.dims.len() == 2 { + crate::operations::linalg::transpose( + grad_output, + self.dims[0] as isize, + self.dims[1] as isize, + )? + } else { + let mut inverse = vec![0; self.dims.len()]; + for (i, &d) in self.dims.iter().enumerate() { + inverse[d] = i; + } + let mut grad = grad_output.clone(); + let mut current: Vec = (0..inverse.len()).collect(); + for i in 0..inverse.len() { + let j = current + .iter() + .position(|&x| x == inverse[i]) + .expect("invalid permutation"); + if i != j { + grad = crate::operations::linalg::transpose(&grad, i as isize, j as isize)?; + current.swap(i, j); + } + } + grad + }; + + gradients.insert(self.input_id, grad_input); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for sum reduction +pub struct SumBackward { + pub input_id: TensorId, + pub input_shape: Vec, + pub dims: Option>, + pub keepdim: bool, +} diff --git a/engine/src/autograd/mod/linalg.rs b/engine/src/autograd/mod/linalg.rs index 04b9e88f..4df37052 100644 --- a/engine/src/autograd/mod/linalg.rs +++ b/engine/src/autograd/mod/linalg.rs @@ -1,781 +1,802 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -impl GradientFunction for SoftplusBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - match self.input.dtype() { - DataType::Float32 => { - let input_slice = self.input.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - let grad_out_slice = grad_output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get f32 slice from grad_output tensor", - ) - })?; - - let mut grad_data = TensorData::uninitialized_on_device( - input_slice.len(), - DataType::Float32, - self.input.device(), - ); - let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from gradient tensor", - ) - })?; - - let beta = self.beta as f32; - let threshold = self.threshold as f32; - for ((grad_slot, &x), &gout) in grad_slice - .iter_mut() - .zip(input_slice.iter()) - .zip(grad_out_slice.iter()) - { - let scaled = beta * x; - *grad_slot = if scaled > threshold { - gout - } else { - gout / (1.0 + (-scaled).exp()) - }; - } - - let grad_tensor = Tensor::new( - Arc::new(grad_data), - self.input.shape().clone(), - DataType::Float32, - self.input.device(), - false, - ); - gradients.insert(self.input_id, grad_tensor); - } - DataType::Float64 => { - let input_slice = self.input.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - let grad_out_slice = grad_output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get f64 slice from grad_output tensor", - ) - })?; - - let mut grad_data = TensorData::uninitialized_on_device( - input_slice.len(), - DataType::Float64, - self.input.device(), - ); - let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from gradient tensor", - ) - })?; - - let beta = self.beta; - let threshold = self.threshold; - for ((grad_slot, &x), &gout) in grad_slice - .iter_mut() - .zip(input_slice.iter()) - .zip(grad_out_slice.iter()) - { - let scaled = beta * x; - *grad_slot = if scaled > threshold { - gout - } else { - gout / (1.0 + (-scaled).exp()) - }; - } - - let grad_tensor = Tensor::new( - Arc::new(grad_data), - self.input.shape().clone(), - DataType::Float64, - self.input.device(), - false, - ); - gradients.insert(self.input_id, grad_tensor); - } - _ => { - return Err(MinitensorError::invalid_operation( - "Softplus gradient only defined for floating point tensors", - )); - } - } - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for GELU activation -pub struct GeluBackward { - pub input_id: TensorId, - pub input: Tensor, - pub approximate: bool, -} - -impl GradientFunction for GeluBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - match self.input.dtype() { - DataType::Float32 => { - let input_slice = self.input.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - let grad_out_slice = grad_output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get f32 slice from grad_output tensor", - ) - })?; - - let mut grad_data = TensorData::uninitialized_on_device( - input_slice.len(), - DataType::Float32, - self.input.device(), - ); - let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from gradient tensor", - ) - })?; - - if self.approximate { - let coeff = (2.0f32 / std::f32::consts::PI).sqrt(); - for ((grad_slot, &x), &gout) in grad_slice - .iter_mut() - .zip(input_slice.iter()) - .zip(grad_out_slice.iter()) - { - let x2 = x * x; - let inner = coeff * (x + 0.044715f32 * x * x2); - let tanh_inner = inner.tanh(); - let sech2 = 1.0f32 - tanh_inner * tanh_inner; - let grad_val = 0.5f32 * (1.0f32 + tanh_inner) - + 0.5f32 * x * sech2 * coeff * (1.0f32 + 3.0f32 * 0.044715f32 * x2); - *grad_slot = gout * grad_val; - } - } else { - let inv_sqrt_2 = std::f32::consts::FRAC_1_SQRT_2; - let inv_sqrt_2pi = 1.0f32 / ((2.0f32 * std::f32::consts::PI).sqrt()); - for ((grad_slot, &x), &gout) in grad_slice - .iter_mut() - .zip(input_slice.iter()) - .zip(grad_out_slice.iter()) - { - let cdf = 0.5f32 * (1.0f32 + erff(x * inv_sqrt_2)); - let pdf = (-0.5f32 * x * x).exp() * inv_sqrt_2pi; - let grad_val = cdf + x * pdf; - *grad_slot = gout * grad_val; - } - } - - let grad_tensor = Tensor::new( - Arc::new(grad_data), - self.input.shape().clone(), - DataType::Float32, - self.input.device(), - false, - ); - gradients.insert(self.input_id, grad_tensor); - } - DataType::Float64 => { - let input_slice = self.input.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - let grad_out_slice = grad_output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get f64 slice from grad_output tensor", - ) - })?; - - let mut grad_data = TensorData::uninitialized_on_device( - input_slice.len(), - DataType::Float64, - self.input.device(), - ); - let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from gradient tensor", - ) - })?; - - if self.approximate { - let coeff = (2.0f64 / std::f64::consts::PI).sqrt(); - for ((grad_slot, &x), &gout) in grad_slice - .iter_mut() - .zip(input_slice.iter()) - .zip(grad_out_slice.iter()) - { - let x2 = x * x; - let inner = coeff * (x + 0.044715f64 * x * x2); - let tanh_inner = inner.tanh(); - let sech2 = 1.0f64 - tanh_inner * tanh_inner; - let grad_val = 0.5f64 * (1.0f64 + tanh_inner) - + 0.5f64 * x * sech2 * coeff * (1.0f64 + 3.0f64 * 0.044715f64 * x2); - *grad_slot = gout * grad_val; - } - } else { - let inv_sqrt_2 = std::f64::consts::FRAC_1_SQRT_2; - let inv_sqrt_2pi = 1.0f64 / ((2.0f64 * std::f64::consts::PI).sqrt()); - for ((grad_slot, &x), &gout) in grad_slice - .iter_mut() - .zip(input_slice.iter()) - .zip(grad_out_slice.iter()) - { - let cdf = 0.5f64 * (1.0f64 + erf(x * inv_sqrt_2)); - let pdf = (-0.5f64 * x * x).exp() * inv_sqrt_2pi; - let grad_val = cdf + x * pdf; - *grad_slot = gout * grad_val; - } - } - - let grad_tensor = Tensor::new( - Arc::new(grad_data), - self.input.shape().clone(), - DataType::Float64, - self.input.device(), - false, - ); - gradients.insert(self.input_id, grad_tensor); - } - _ => { - return Err(MinitensorError::invalid_operation( - "GELU backward only supports floating point tensors", - )); - } - } - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for ELU activation -pub struct EluBackward { - pub input_id: TensorId, - pub output: Tensor, - pub alpha: f64, -} - -impl GradientFunction for EluBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - match self.output.dtype() { - DataType::Float32 => { - let output_slice = self.output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from output tensor") - })?; - let grad_out_slice = grad_output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get f32 slice from grad_output tensor", - ) - })?; - - let mut grad_data = TensorData::uninitialized_on_device( - output_slice.len(), - DataType::Float32, - self.output.device(), - ); - let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from gradient tensor", - ) - })?; - - let alpha = self.alpha as f32; - for ((grad_slot, &out), &gout) in grad_slice - .iter_mut() - .zip(output_slice.iter()) - .zip(grad_out_slice.iter()) - { - let local_grad = if out > 0.0f32 { 1.0f32 } else { out + alpha }; - *grad_slot = gout * local_grad; - } - - let grad_tensor = Tensor::new( - Arc::new(grad_data), - self.output.shape().clone(), - DataType::Float32, - self.output.device(), - false, - ); - gradients.insert(self.input_id, grad_tensor); - } - DataType::Float64 => { - let output_slice = self.output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from output tensor") - })?; - let grad_out_slice = grad_output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get f64 slice from grad_output tensor", - ) - })?; - - let mut grad_data = TensorData::uninitialized_on_device( - output_slice.len(), - DataType::Float64, - self.output.device(), - ); - let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from gradient tensor", - ) - })?; - - for ((grad_slot, &out), &gout) in grad_slice - .iter_mut() - .zip(output_slice.iter()) - .zip(grad_out_slice.iter()) - { - let local_grad = if out > 0.0f64 { - 1.0f64 - } else { - out + self.alpha - }; - *grad_slot = gout * local_grad; - } - - let grad_tensor = Tensor::new( - Arc::new(grad_data), - self.output.shape().clone(), - DataType::Float64, - self.output.device(), - false, - ); - gradients.insert(self.input_id, grad_tensor); - } - _ => { - return Err(MinitensorError::invalid_operation( - "ELU backward only supports floating point tensors", - )); - } - } - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for SELU activation -pub struct SeluBackward { - pub input_id: TensorId, - pub output: Tensor, -} - -impl GradientFunction for SeluBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - match self.output.dtype() { - DataType::Float32 => { - let output_slice = self.output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from output tensor") - })?; - let grad_out_slice = grad_output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get f32 slice from grad_output tensor", - ) - })?; - - let mut grad_data = TensorData::uninitialized_on_device( - output_slice.len(), - DataType::Float32, - self.output.device(), - ); - let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from gradient tensor", - ) - })?; - - const SCALE: f32 = 1.050701; - const ALPHA: f32 = 1.6732632; - for ((grad_slot, &out), &gout) in grad_slice - .iter_mut() - .zip(output_slice.iter()) - .zip(grad_out_slice.iter()) - { - let local_grad = if out > 0.0f32 { - SCALE - } else { - out + SCALE * ALPHA - }; - *grad_slot = gout * local_grad; - } - - let grad_tensor = Tensor::new( - Arc::new(grad_data), - self.output.shape().clone(), - DataType::Float32, - self.output.device(), - false, - ); - gradients.insert(self.input_id, grad_tensor); - } - DataType::Float64 => { - let output_slice = self.output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from output tensor") - })?; - let grad_out_slice = grad_output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get f64 slice from grad_output tensor", - ) - })?; - - let mut grad_data = TensorData::uninitialized_on_device( - output_slice.len(), - DataType::Float64, - self.output.device(), - ); - let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from gradient tensor", - ) - })?; - - const SCALE: f64 = 1.0507009873554804934193349852946; - const ALPHA: f64 = 1.6732632423543772848170429916717; - for ((grad_slot, &out), &gout) in grad_slice - .iter_mut() - .zip(output_slice.iter()) - .zip(grad_out_slice.iter()) - { - let local_grad = if out > 0.0f64 { - SCALE - } else { - out + SCALE * ALPHA - }; - *grad_slot = gout * local_grad; - } - - let grad_tensor = Tensor::new( - Arc::new(grad_data), - self.output.shape().clone(), - DataType::Float64, - self.output.device(), - false, - ); - gradients.insert(self.input_id, grad_tensor); - } - _ => { - return Err(MinitensorError::invalid_operation( - "SELU backward only supports floating point tensors", - )); - } - } - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for SiLU activation -pub struct SiluBackward { - pub input_id: TensorId, - pub input: Tensor, -} - -impl GradientFunction for SiluBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - match self.input.dtype() { - DataType::Float32 => { - let input_slice = self.input.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - let grad_out_slice = grad_output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get f32 slice from grad_output tensor", - ) - })?; - - let mut grad_data = TensorData::uninitialized_on_device( - input_slice.len(), - DataType::Float32, - self.input.device(), - ); - let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from gradient tensor", - ) - })?; - - for ((grad_slot, &x), &gout) in grad_slice - .iter_mut() - .zip(input_slice.iter()) - .zip(grad_out_slice.iter()) - { - let sigmoid = stable_sigmoid_f32(x); - let grad_val = sigmoid * (1.0f32 + x * (1.0f32 - sigmoid)); - *grad_slot = gout * grad_val; - } - - let grad_tensor = Tensor::new( - Arc::new(grad_data), - self.input.shape().clone(), - DataType::Float32, - self.input.device(), - false, - ); - gradients.insert(self.input_id, grad_tensor); - } - DataType::Float64 => { - let input_slice = self.input.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - let grad_out_slice = grad_output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get f64 slice from grad_output tensor", - ) - })?; - - let mut grad_data = TensorData::uninitialized_on_device( - input_slice.len(), - DataType::Float64, - self.input.device(), - ); - let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from gradient tensor", - ) - })?; - - for ((grad_slot, &x), &gout) in grad_slice - .iter_mut() - .zip(input_slice.iter()) - .zip(grad_out_slice.iter()) - { - let sigmoid = stable_sigmoid_f64(x); - let grad_val = sigmoid * (1.0f64 + x * (1.0f64 - sigmoid)); - *grad_slot = gout * grad_val; - } - - let grad_tensor = Tensor::new( - Arc::new(grad_data), - self.input.shape().clone(), - DataType::Float64, - self.input.device(), - false, - ); - gradients.insert(self.input_id, grad_tensor); - } - _ => { - return Err(MinitensorError::invalid_operation( - "SiLU backward only supports floating point tensors", - )); - } - } - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -#[inline] -fn stable_sigmoid_f32(x: f32) -> f32 { - if x >= 0.0 { - let exp_neg = (-x).exp(); - 1.0 / (1.0 + exp_neg) - } else { - let exp_pos = x.exp(); - exp_pos / (1.0 + exp_pos) - } -} - -#[inline] -fn stable_sigmoid_f64(x: f64) -> f64 { - if x >= 0.0 { - let exp_neg = (-x).exp(); - 1.0 / (1.0 + exp_neg) - } else { - let exp_pos = x.exp(); - exp_pos / (1.0 + exp_pos) - } -} - -/// Gradient function for Softsign activation -pub struct SoftsignBackward { - pub input_id: TensorId, - pub input: Tensor, -} - -impl GradientFunction for SoftsignBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - match self.input.dtype() { - DataType::Float32 => { - let input_slice = self.input.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - let grad_out_slice = grad_output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get f32 slice from grad_output tensor", - ) - })?; - - let mut grad_data = TensorData::uninitialized_on_device( - input_slice.len(), - DataType::Float32, - self.input.device(), - ); - let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from gradient tensor", - ) - })?; - - for ((grad_slot, &x), &gout) in grad_slice - .iter_mut() - .zip(input_slice.iter()) - .zip(grad_out_slice.iter()) - { - let denom = 1.0f32 + x.abs(); - let local_grad = 1.0f32 / (denom * denom); - *grad_slot = gout * local_grad; - } - - let grad_tensor = Tensor::new( - Arc::new(grad_data), - self.input.shape().clone(), - DataType::Float32, - self.input.device(), - false, - ); - gradients.insert(self.input_id, grad_tensor); - } - DataType::Float64 => { - let input_slice = self.input.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - let grad_out_slice = grad_output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get f64 slice from grad_output tensor", - ) - })?; - - let mut grad_data = TensorData::uninitialized_on_device( - input_slice.len(), - DataType::Float64, - self.input.device(), - ); - let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from gradient tensor", - ) - })?; - - for ((grad_slot, &x), &gout) in grad_slice - .iter_mut() - .zip(input_slice.iter()) - .zip(grad_out_slice.iter()) - { - let denom = 1.0f64 + x.abs(); - let local_grad = 1.0f64 / (denom * denom); - *grad_slot = gout * local_grad; - } - - let grad_tensor = Tensor::new( - Arc::new(grad_data), - self.input.shape().clone(), - DataType::Float64, - self.input.device(), - false, - ); - gradients.insert(self.input_id, grad_tensor); - } - _ => { - return Err(MinitensorError::invalid_operation( - "Softsign backward only supports floating point tensors", - )); - } - } - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for power operation -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum PowBroadcast { - None, - BaseScalar, - ExponentScalar, -} -pub struct PowBackward { - pub base: Tensor, - pub exponent: Tensor, - pub output: Tensor, - pub input_ids: [TensorId; 2], - pub base_requires_grad: bool, - pub exp_requires_grad: bool, - pub broadcast: PowBroadcast, -} - -/// Gradient function for logaddexp -pub struct LogAddExpBackward { - pub lhs: Tensor, - pub rhs: Tensor, - pub output: Tensor, - pub input_ids: [TensorId; 2], - pub input_shapes: [Vec; 2], -} - -impl GradientFunction for LogAddExpBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(2); - - let lhs_diff = arithmetic::sub(&self.lhs.detach(), &self.output.detach())?; - let lhs_term = lhs_diff.exp()?; - let lhs_mul = arithmetic::mul(&lhs_term, grad_output)?; - let lhs_grad = - reduce_gradient_for_broadcasting(&lhs_mul, &Shape::new(self.input_shapes[0].clone()))?; - accumulate_grad(&mut gradients, self.input_ids[0], lhs_grad)?; - - let rhs_diff = arithmetic::sub(&self.rhs.detach(), &self.output.detach())?; - let rhs_term = rhs_diff.exp()?; - let rhs_mul = arithmetic::mul(&rhs_term, grad_output)?; - let rhs_grad = - reduce_gradient_for_broadcasting(&rhs_mul, &Shape::new(self.input_shapes[1].clone()))?; - accumulate_grad(&mut gradients, self.input_ids[1], rhs_grad)?; - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::{ + error::{MinitensorError, Result}, + operations::arithmetic, + tensor::{DataType, Shape, Tensor, TensorData}, +}; +use libm::{erf, erff}; +use rustc_hash::FxHashMap; +use std::sync::Arc; + +impl GradientFunction for SoftplusBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + match self.input.dtype() { + DataType::Float32 => { + let input_slice = self.input.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + let grad_out_slice = grad_output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get f32 slice from grad_output tensor", + ) + })?; + + let mut grad_data = TensorData::uninitialized_on_device( + input_slice.len(), + DataType::Float32, + self.input.device(), + ); + let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from gradient tensor", + ) + })?; + + let beta = self.beta as f32; + let threshold = self.threshold as f32; + for ((grad_slot, &x), &gout) in grad_slice + .iter_mut() + .zip(input_slice.iter()) + .zip(grad_out_slice.iter()) + { + let scaled = beta * x; + *grad_slot = if scaled > threshold { + gout + } else { + gout / (1.0 + (-scaled).exp()) + }; + } + + let grad_tensor = Tensor::new( + Arc::new(grad_data), + self.input.shape().clone(), + DataType::Float32, + self.input.device(), + false, + ); + gradients.insert(self.input_id, grad_tensor); + } + DataType::Float64 => { + let input_slice = self.input.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + let grad_out_slice = grad_output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get f64 slice from grad_output tensor", + ) + })?; + + let mut grad_data = TensorData::uninitialized_on_device( + input_slice.len(), + DataType::Float64, + self.input.device(), + ); + let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from gradient tensor", + ) + })?; + + let beta = self.beta; + let threshold = self.threshold; + for ((grad_slot, &x), &gout) in grad_slice + .iter_mut() + .zip(input_slice.iter()) + .zip(grad_out_slice.iter()) + { + let scaled = beta * x; + *grad_slot = if scaled > threshold { + gout + } else { + gout / (1.0 + (-scaled).exp()) + }; + } + + let grad_tensor = Tensor::new( + Arc::new(grad_data), + self.input.shape().clone(), + DataType::Float64, + self.input.device(), + false, + ); + gradients.insert(self.input_id, grad_tensor); + } + _ => { + return Err(MinitensorError::invalid_operation( + "Softplus gradient only defined for floating point tensors", + )); + } + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for GELU activation +pub struct GeluBackward { + pub input_id: TensorId, + pub input: Tensor, + pub approximate: bool, +} + +impl GradientFunction for GeluBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + match self.input.dtype() { + DataType::Float32 => { + let input_slice = self.input.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + let grad_out_slice = grad_output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get f32 slice from grad_output tensor", + ) + })?; + + let mut grad_data = TensorData::uninitialized_on_device( + input_slice.len(), + DataType::Float32, + self.input.device(), + ); + let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from gradient tensor", + ) + })?; + + if self.approximate { + let coeff = (2.0f32 / std::f32::consts::PI).sqrt(); + for ((grad_slot, &x), &gout) in grad_slice + .iter_mut() + .zip(input_slice.iter()) + .zip(grad_out_slice.iter()) + { + let x2 = x * x; + let inner = coeff * (x + 0.044715f32 * x * x2); + let tanh_inner = inner.tanh(); + let sech2 = 1.0f32 - tanh_inner * tanh_inner; + let grad_val = 0.5f32 * (1.0f32 + tanh_inner) + + 0.5f32 * x * sech2 * coeff * (1.0f32 + 3.0f32 * 0.044715f32 * x2); + *grad_slot = gout * grad_val; + } + } else { + let inv_sqrt_2 = std::f32::consts::FRAC_1_SQRT_2; + let inv_sqrt_2pi = 1.0f32 / ((2.0f32 * std::f32::consts::PI).sqrt()); + for ((grad_slot, &x), &gout) in grad_slice + .iter_mut() + .zip(input_slice.iter()) + .zip(grad_out_slice.iter()) + { + let cdf = 0.5f32 * (1.0f32 + erff(x * inv_sqrt_2)); + let pdf = (-0.5f32 * x * x).exp() * inv_sqrt_2pi; + let grad_val = cdf + x * pdf; + *grad_slot = gout * grad_val; + } + } + + let grad_tensor = Tensor::new( + Arc::new(grad_data), + self.input.shape().clone(), + DataType::Float32, + self.input.device(), + false, + ); + gradients.insert(self.input_id, grad_tensor); + } + DataType::Float64 => { + let input_slice = self.input.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + let grad_out_slice = grad_output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get f64 slice from grad_output tensor", + ) + })?; + + let mut grad_data = TensorData::uninitialized_on_device( + input_slice.len(), + DataType::Float64, + self.input.device(), + ); + let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from gradient tensor", + ) + })?; + + if self.approximate { + let coeff = (2.0f64 / std::f64::consts::PI).sqrt(); + for ((grad_slot, &x), &gout) in grad_slice + .iter_mut() + .zip(input_slice.iter()) + .zip(grad_out_slice.iter()) + { + let x2 = x * x; + let inner = coeff * (x + 0.044715f64 * x * x2); + let tanh_inner = inner.tanh(); + let sech2 = 1.0f64 - tanh_inner * tanh_inner; + let grad_val = 0.5f64 * (1.0f64 + tanh_inner) + + 0.5f64 * x * sech2 * coeff * (1.0f64 + 3.0f64 * 0.044715f64 * x2); + *grad_slot = gout * grad_val; + } + } else { + let inv_sqrt_2 = std::f64::consts::FRAC_1_SQRT_2; + let inv_sqrt_2pi = 1.0f64 / ((2.0f64 * std::f64::consts::PI).sqrt()); + for ((grad_slot, &x), &gout) in grad_slice + .iter_mut() + .zip(input_slice.iter()) + .zip(grad_out_slice.iter()) + { + let cdf = 0.5f64 * (1.0f64 + erf(x * inv_sqrt_2)); + let pdf = (-0.5f64 * x * x).exp() * inv_sqrt_2pi; + let grad_val = cdf + x * pdf; + *grad_slot = gout * grad_val; + } + } + + let grad_tensor = Tensor::new( + Arc::new(grad_data), + self.input.shape().clone(), + DataType::Float64, + self.input.device(), + false, + ); + gradients.insert(self.input_id, grad_tensor); + } + _ => { + return Err(MinitensorError::invalid_operation( + "GELU backward only supports floating point tensors", + )); + } + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for ELU activation +pub struct EluBackward { + pub input_id: TensorId, + pub output: Tensor, + pub alpha: f64, +} + +impl GradientFunction for EluBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + match self.output.dtype() { + DataType::Float32 => { + let output_slice = self.output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from output tensor") + })?; + let grad_out_slice = grad_output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get f32 slice from grad_output tensor", + ) + })?; + + let mut grad_data = TensorData::uninitialized_on_device( + output_slice.len(), + DataType::Float32, + self.output.device(), + ); + let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from gradient tensor", + ) + })?; + + let alpha = self.alpha as f32; + for ((grad_slot, &out), &gout) in grad_slice + .iter_mut() + .zip(output_slice.iter()) + .zip(grad_out_slice.iter()) + { + let local_grad = if out > 0.0f32 { 1.0f32 } else { out + alpha }; + *grad_slot = gout * local_grad; + } + + let grad_tensor = Tensor::new( + Arc::new(grad_data), + self.output.shape().clone(), + DataType::Float32, + self.output.device(), + false, + ); + gradients.insert(self.input_id, grad_tensor); + } + DataType::Float64 => { + let output_slice = self.output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from output tensor") + })?; + let grad_out_slice = grad_output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get f64 slice from grad_output tensor", + ) + })?; + + let mut grad_data = TensorData::uninitialized_on_device( + output_slice.len(), + DataType::Float64, + self.output.device(), + ); + let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from gradient tensor", + ) + })?; + + for ((grad_slot, &out), &gout) in grad_slice + .iter_mut() + .zip(output_slice.iter()) + .zip(grad_out_slice.iter()) + { + let local_grad = if out > 0.0f64 { + 1.0f64 + } else { + out + self.alpha + }; + *grad_slot = gout * local_grad; + } + + let grad_tensor = Tensor::new( + Arc::new(grad_data), + self.output.shape().clone(), + DataType::Float64, + self.output.device(), + false, + ); + gradients.insert(self.input_id, grad_tensor); + } + _ => { + return Err(MinitensorError::invalid_operation( + "ELU backward only supports floating point tensors", + )); + } + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for SELU activation +pub struct SeluBackward { + pub input_id: TensorId, + pub output: Tensor, +} + +impl GradientFunction for SeluBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + match self.output.dtype() { + DataType::Float32 => { + let output_slice = self.output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from output tensor") + })?; + let grad_out_slice = grad_output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get f32 slice from grad_output tensor", + ) + })?; + + let mut grad_data = TensorData::uninitialized_on_device( + output_slice.len(), + DataType::Float32, + self.output.device(), + ); + let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from gradient tensor", + ) + })?; + + const SCALE: f32 = 1.050701; + const ALPHA: f32 = 1.6732632; + for ((grad_slot, &out), &gout) in grad_slice + .iter_mut() + .zip(output_slice.iter()) + .zip(grad_out_slice.iter()) + { + let local_grad = if out > 0.0f32 { + SCALE + } else { + out + SCALE * ALPHA + }; + *grad_slot = gout * local_grad; + } + + let grad_tensor = Tensor::new( + Arc::new(grad_data), + self.output.shape().clone(), + DataType::Float32, + self.output.device(), + false, + ); + gradients.insert(self.input_id, grad_tensor); + } + DataType::Float64 => { + let output_slice = self.output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from output tensor") + })?; + let grad_out_slice = grad_output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get f64 slice from grad_output tensor", + ) + })?; + + let mut grad_data = TensorData::uninitialized_on_device( + output_slice.len(), + DataType::Float64, + self.output.device(), + ); + let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from gradient tensor", + ) + })?; + + const SCALE: f64 = 1.0507009873554804934193349852946; + const ALPHA: f64 = 1.6732632423543772848170429916717; + for ((grad_slot, &out), &gout) in grad_slice + .iter_mut() + .zip(output_slice.iter()) + .zip(grad_out_slice.iter()) + { + let local_grad = if out > 0.0f64 { + SCALE + } else { + out + SCALE * ALPHA + }; + *grad_slot = gout * local_grad; + } + + let grad_tensor = Tensor::new( + Arc::new(grad_data), + self.output.shape().clone(), + DataType::Float64, + self.output.device(), + false, + ); + gradients.insert(self.input_id, grad_tensor); + } + _ => { + return Err(MinitensorError::invalid_operation( + "SELU backward only supports floating point tensors", + )); + } + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for SiLU activation +pub struct SiluBackward { + pub input_id: TensorId, + pub input: Tensor, +} + +impl GradientFunction for SiluBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + match self.input.dtype() { + DataType::Float32 => { + let input_slice = self.input.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + let grad_out_slice = grad_output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get f32 slice from grad_output tensor", + ) + })?; + + let mut grad_data = TensorData::uninitialized_on_device( + input_slice.len(), + DataType::Float32, + self.input.device(), + ); + let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from gradient tensor", + ) + })?; + + for ((grad_slot, &x), &gout) in grad_slice + .iter_mut() + .zip(input_slice.iter()) + .zip(grad_out_slice.iter()) + { + let sigmoid = stable_sigmoid_f32(x); + let grad_val = sigmoid * (1.0f32 + x * (1.0f32 - sigmoid)); + *grad_slot = gout * grad_val; + } + + let grad_tensor = Tensor::new( + Arc::new(grad_data), + self.input.shape().clone(), + DataType::Float32, + self.input.device(), + false, + ); + gradients.insert(self.input_id, grad_tensor); + } + DataType::Float64 => { + let input_slice = self.input.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + let grad_out_slice = grad_output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get f64 slice from grad_output tensor", + ) + })?; + + let mut grad_data = TensorData::uninitialized_on_device( + input_slice.len(), + DataType::Float64, + self.input.device(), + ); + let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from gradient tensor", + ) + })?; + + for ((grad_slot, &x), &gout) in grad_slice + .iter_mut() + .zip(input_slice.iter()) + .zip(grad_out_slice.iter()) + { + let sigmoid = stable_sigmoid_f64(x); + let grad_val = sigmoid * (1.0f64 + x * (1.0f64 - sigmoid)); + *grad_slot = gout * grad_val; + } + + let grad_tensor = Tensor::new( + Arc::new(grad_data), + self.input.shape().clone(), + DataType::Float64, + self.input.device(), + false, + ); + gradients.insert(self.input_id, grad_tensor); + } + _ => { + return Err(MinitensorError::invalid_operation( + "SiLU backward only supports floating point tensors", + )); + } + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +#[inline] +fn stable_sigmoid_f32(x: f32) -> f32 { + if x >= 0.0 { + let exp_neg = (-x).exp(); + 1.0 / (1.0 + exp_neg) + } else { + let exp_pos = x.exp(); + exp_pos / (1.0 + exp_pos) + } +} + +#[inline] +fn stable_sigmoid_f64(x: f64) -> f64 { + if x >= 0.0 { + let exp_neg = (-x).exp(); + 1.0 / (1.0 + exp_neg) + } else { + let exp_pos = x.exp(); + exp_pos / (1.0 + exp_pos) + } +} + +/// Gradient function for Softsign activation +pub struct SoftsignBackward { + pub input_id: TensorId, + pub input: Tensor, +} + +impl GradientFunction for SoftsignBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + match self.input.dtype() { + DataType::Float32 => { + let input_slice = self.input.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + let grad_out_slice = grad_output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get f32 slice from grad_output tensor", + ) + })?; + + let mut grad_data = TensorData::uninitialized_on_device( + input_slice.len(), + DataType::Float32, + self.input.device(), + ); + let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from gradient tensor", + ) + })?; + + for ((grad_slot, &x), &gout) in grad_slice + .iter_mut() + .zip(input_slice.iter()) + .zip(grad_out_slice.iter()) + { + let denom = 1.0f32 + x.abs(); + let local_grad = 1.0f32 / (denom * denom); + *grad_slot = gout * local_grad; + } + + let grad_tensor = Tensor::new( + Arc::new(grad_data), + self.input.shape().clone(), + DataType::Float32, + self.input.device(), + false, + ); + gradients.insert(self.input_id, grad_tensor); + } + DataType::Float64 => { + let input_slice = self.input.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + let grad_out_slice = grad_output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get f64 slice from grad_output tensor", + ) + })?; + + let mut grad_data = TensorData::uninitialized_on_device( + input_slice.len(), + DataType::Float64, + self.input.device(), + ); + let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from gradient tensor", + ) + })?; + + for ((grad_slot, &x), &gout) in grad_slice + .iter_mut() + .zip(input_slice.iter()) + .zip(grad_out_slice.iter()) + { + let denom = 1.0f64 + x.abs(); + let local_grad = 1.0f64 / (denom * denom); + *grad_slot = gout * local_grad; + } + + let grad_tensor = Tensor::new( + Arc::new(grad_data), + self.input.shape().clone(), + DataType::Float64, + self.input.device(), + false, + ); + gradients.insert(self.input_id, grad_tensor); + } + _ => { + return Err(MinitensorError::invalid_operation( + "Softsign backward only supports floating point tensors", + )); + } + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for power operation +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum PowBroadcast { + None, + BaseScalar, + ExponentScalar, +} +pub struct PowBackward { + pub base: Tensor, + pub exponent: Tensor, + pub output: Tensor, + pub input_ids: [TensorId; 2], + pub base_requires_grad: bool, + pub exp_requires_grad: bool, + pub broadcast: PowBroadcast, +} + +/// Gradient function for logaddexp +pub struct LogAddExpBackward { + pub lhs: Tensor, + pub rhs: Tensor, + pub output: Tensor, + pub input_ids: [TensorId; 2], + pub input_shapes: [Vec; 2], + /// Which inputs actually need a gradient; frozen inputs skip their + /// exp/sub/mul/reduce chain entirely. + pub input_requires_grad: [bool; 2], +} + +impl GradientFunction for LogAddExpBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(2); + + if self.input_requires_grad[0] { + let lhs_diff = arithmetic::sub(&self.lhs.detach(), &self.output.detach())?; + let lhs_term = lhs_diff.exp()?; + let lhs_mul = arithmetic::mul(&lhs_term, grad_output)?; + let lhs_grad = reduce_gradient_for_broadcasting( + &lhs_mul, + &Shape::new(self.input_shapes[0].clone()), + )?; + accumulate_grad(&mut gradients, self.input_ids[0], lhs_grad)?; + } + + if self.input_requires_grad[1] { + let rhs_diff = arithmetic::sub(&self.rhs.detach(), &self.output.detach())?; + let rhs_term = rhs_diff.exp()?; + let rhs_mul = arithmetic::mul(&rhs_term, grad_output)?; + let rhs_grad = reduce_gradient_for_broadcasting( + &rhs_mul, + &Shape::new(self.input_shapes[1].clone()), + )?; + accumulate_grad(&mut gradients, self.input_ids[1], rhs_grad)?; + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} diff --git a/engine/src/autograd/mod/reduction.rs b/engine/src/autograd/mod/reduction.rs index b6e400d2..b0ad7249 100644 --- a/engine/src/autograd/mod/reduction.rs +++ b/engine/src/autograd/mod/reduction.rs @@ -1,1143 +1,1157 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -impl GradientFunction for MaskedLogSoftmaxBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - let mut grad_data = TensorData::zeros_on_device( - self.output.numel(), - self.output.dtype(), - self.output.device(), - ); - - let output_dims = self.output.shape().dims(); - let mask_dims = self.mask.shape().dims(); - let same_shape = output_dims == mask_dims; - let output_strides = if same_shape { - None - } else { - Some(Strides::from_shape(self.output.shape())) - }; - let mask_strides = if same_shape { - None - } else { - Some(Strides::from_shape(self.mask.shape())) - }; - - match grad_output.dtype() { - DataType::Float32 => { - let go = grad_output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from grad_output") - })?; - let log_y = self.output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get f32 slice from masked log_softmax output", - ) - })?; - let mask_data = self.mask.data().as_bool_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from mask tensor") - })?; - let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from grad_data", - ) - })?; - masked_log_softmax_backward_f32( - go, - log_y, - mask_data, - grad_slice, - output_dims, - self.dim, - mask_dims, - output_strides.as_ref().map(Strides::as_slice), - mask_strides.as_ref().map(Strides::as_slice), - ); - } - DataType::Float64 => { - let go = grad_output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from grad_output") - })?; - let log_y = self.output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get f64 slice from masked log_softmax output", - ) - })?; - let mask_data = self.mask.data().as_bool_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from mask tensor") - })?; - let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from grad_data", - ) - })?; - masked_log_softmax_backward_f64( - go, - log_y, - mask_data, - grad_slice, - output_dims, - self.dim, - mask_dims, - output_strides.as_ref().map(Strides::as_slice), - mask_strides.as_ref().map(Strides::as_slice), - ); - } - _ => { - return Err(MinitensorError::invalid_operation( - "Masked log_softmax backward only supported for floating point tensors", - )); - } - } - - let grad_input = Tensor::new( - Arc::new(grad_data), - self.output.shape().clone(), - self.output.dtype(), - self.output.device(), - grad_output.requires_grad(), - ); - - accumulate_grad(&mut gradients, self.input_id, grad_input)?; - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for layer normalization -pub struct LayerNormBackward { - pub input_ids: SmallVec<[TensorId; 3]>, - pub input_id: TensorId, - pub weight_id: Option, - pub bias_id: Option, - pub normalized: Tensor, - pub inv_std: Tensor, - pub weight_broadcast: Option, - pub normalized_shape: Vec, - pub axis_start: usize, - pub element_count: usize, - pub input_requires_grad: bool, - pub weight_requires_grad: bool, - pub bias_requires_grad: bool, -} - -impl GradientFunction for LayerNormBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - - let grad_output_detached = grad_output.detach(); - let normalized = self.normalized.detach(); - - if self.element_count == 0 { - if self.input_requires_grad { - let zero = Tensor::zeros( - grad_output.shape().clone(), - grad_output.dtype(), - grad_output.device(), - false, - ); - accumulate_grad(&mut gradients, self.input_id, zero)?; - } - if self.weight_requires_grad { - if let Some(weight_id) = self.weight_id { - let zero = Tensor::zeros( - Shape::new(self.normalized_shape.clone()), - grad_output.dtype(), - grad_output.device(), - false, - ); - accumulate_grad(&mut gradients, weight_id, zero)?; - } - } - if self.bias_requires_grad { - if let Some(bias_id) = self.bias_id { - let zero = Tensor::zeros( - Shape::new(self.normalized_shape.clone()), - grad_output.dtype(), - grad_output.device(), - false, - ); - accumulate_grad(&mut gradients, bias_id, zero)?; - } - } - - return Ok(gradients); - } - - if self.input_requires_grad { - let mut grad_output_hat = if let Some(weight) = &self.weight_broadcast { - arithmetic::mul(&grad_output_detached, weight)? - } else { - grad_output_detached.clone() - }; - - let axes: Vec = (self.axis_start..grad_output_hat.ndim()) - .map(|d| d as isize) - .collect(); - let sum_grad = reduction::sum(&grad_output_hat, Some(axes.clone()), true)?; - let grad_norm_mul = arithmetic::mul(&grad_output_hat, &normalized)?; - let sum_grad_norm = reduction::sum(&grad_norm_mul, Some(axes), true)?; - - let count = self.element_count as f64; - let m_tensor = create_scalar_tensor(count, grad_output.dtype(), grad_output.device())?; - let inv_m_tensor = - create_scalar_tensor(1.0 / count, grad_output.dtype(), grad_output.device())?; - grad_output_hat = arithmetic::mul(&grad_output_hat, &m_tensor)?; - let tmp = arithmetic::sub(&grad_output_hat, &sum_grad)?; - let norm_term = arithmetic::mul(&normalized, &sum_grad_norm)?; - let numerator = arithmetic::sub(&tmp, &norm_term)?; - let grad_input = arithmetic::mul(&numerator, &self.inv_std)?; - let grad_input = arithmetic::mul(&grad_input, &inv_m_tensor)?; - accumulate_grad(&mut gradients, self.input_id, grad_input)?; - } - - if self.weight_requires_grad { - if let Some(weight_id) = self.weight_id { - let mut grad_weight = arithmetic::mul(&grad_output_detached, &normalized)?; - if self.axis_start > 0 { - let axes: Vec = (0..self.axis_start).map(|d| d as isize).collect(); - grad_weight = reduction::sum(&grad_weight, Some(axes), false)?; - } - if grad_weight.shape().dims() != self.normalized_shape.as_slice() { - grad_weight = grad_weight.view(Shape::new(self.normalized_shape.clone()))?; - } - accumulate_grad(&mut gradients, weight_id, grad_weight)?; - } - } - - if self.bias_requires_grad { - if let Some(bias_id) = self.bias_id { - let mut grad_bias = grad_output_detached.clone(); - if self.axis_start > 0 { - let axes: Vec = (0..self.axis_start).map(|d| d as isize).collect(); - grad_bias = reduction::sum(&grad_bias, Some(axes), false)?; - } - if grad_bias.shape().dims() != self.normalized_shape.as_slice() { - grad_bias = grad_bias.view(Shape::new(self.normalized_shape.clone()))?; - } - accumulate_grad(&mut gradients, bias_id, grad_bias)?; - } - } - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} - -fn softmax_backward_f32( - grad_output: &[f32], - y: &[f32], - grad_input: &mut [f32], - dims: &[usize], - dim: usize, -) { - if dims.is_empty() { - if let Some(first) = grad_input.first_mut() { - *first = 0.0; - } - return; - } - - let dim_size = dims[dim]; - if dim_size == 0 { - return; - } - let after: usize = if dim + 1 >= dims.len() { - 1 - } else { - dims[dim + 1..].iter().product() - }; - let group = dim_size * after; - if grad_output.len() < PAR_THRESHOLD { - for ((go_block, y_block), out_block) in grad_output - .chunks(group) - .zip(y.chunks(group)) - .zip(grad_input.chunks_mut(group)) - { - for a in 0..after { - let base = a; - let mut dot = 0.0f32; - for k in 0..dim_size { - let idx = base + k * after; - dot += go_block[idx] * y_block[idx]; - } - for k in 0..dim_size { - let idx = base + k * after; - out_block[idx] = y_block[idx] * (go_block[idx] - dot); - } - } - } - } else { - grad_output - .par_chunks(group) - .zip(y.par_chunks(group)) - .zip(grad_input.par_chunks_mut(group)) - .for_each(|((go_block, y_block), out_block)| { - for a in 0..after { - let base = a; - let mut dot = 0.0f32; - for k in 0..dim_size { - let idx = base + k * after; - dot += go_block[idx] * y_block[idx]; - } - for k in 0..dim_size { - let idx = base + k * after; - out_block[idx] = y_block[idx] * (go_block[idx] - dot); - } - } - }); - } -} - -fn softmax_backward_f64( - grad_output: &[f64], - y: &[f64], - grad_input: &mut [f64], - dims: &[usize], - dim: usize, -) { - if dims.is_empty() { - if let Some(first) = grad_input.first_mut() { - *first = 0.0; - } - return; - } - - let dim_size = dims[dim]; - if dim_size == 0 { - return; - } - let after: usize = if dim + 1 >= dims.len() { - 1 - } else { - dims[dim + 1..].iter().product() - }; - let group = dim_size * after; - if grad_output.len() < PAR_THRESHOLD { - for ((go_block, y_block), out_block) in grad_output - .chunks(group) - .zip(y.chunks(group)) - .zip(grad_input.chunks_mut(group)) - { - for a in 0..after { - let base = a; - let mut dot = 0.0f64; - for k in 0..dim_size { - let idx = base + k * after; - dot += go_block[idx] * y_block[idx]; - } - for k in 0..dim_size { - let idx = base + k * after; - out_block[idx] = y_block[idx] * (go_block[idx] - dot); - } - } - } - } else { - grad_output - .par_chunks(group) - .zip(y.par_chunks(group)) - .zip(grad_input.par_chunks_mut(group)) - .for_each(|((go_block, y_block), out_block)| { - for a in 0..after { - let base = a; - let mut dot = 0.0f64; - for k in 0..dim_size { - let idx = base + k * after; - dot += go_block[idx] * y_block[idx]; - } - for k in 0..dim_size { - let idx = base + k * after; - out_block[idx] = y_block[idx] * (go_block[idx] - dot); - } - } - }); - } -} - -fn log_softmax_backward_f32( - grad_output: &[f32], - log_y: &[f32], - grad_input: &mut [f32], - dims: &[usize], - dim: usize, -) { - if dims.is_empty() { - if let Some(first) = grad_input.first_mut() { - *first = 0.0; - } - return; - } - - let dim_size = dims[dim]; - if dim_size == 0 { - return; - } - let after: usize = if dim + 1 >= dims.len() { - 1 - } else { - dims[dim + 1..].iter().product() - }; - let group = dim_size * after; - - if grad_output.len() < PAR_THRESHOLD { - for ((go_block, log_block), out_block) in grad_output - .chunks(group) - .zip(log_y.chunks(group)) - .zip(grad_input.chunks_mut(group)) - { - for a in 0..after { - let base = a; - let mut sum = 0.0f32; - for k in 0..dim_size { - let idx = base + k * after; - sum += go_block[idx]; - } - for k in 0..dim_size { - let idx = base + k * after; - let prob = log_block[idx].exp(); - out_block[idx] = go_block[idx] - prob * sum; - } - } - } - } else { - grad_output - .par_chunks(group) - .zip(log_y.par_chunks(group)) - .zip(grad_input.par_chunks_mut(group)) - .for_each(|((go_block, log_block), out_block)| { - for a in 0..after { - let base = a; - let mut sum = 0.0f32; - for k in 0..dim_size { - let idx = base + k * after; - sum += go_block[idx]; - } - for k in 0..dim_size { - let idx = base + k * after; - let prob = log_block[idx].exp(); - out_block[idx] = go_block[idx] - prob * sum; - } - } - }); - } -} - -fn log_softmax_backward_f64( - grad_output: &[f64], - log_y: &[f64], - grad_input: &mut [f64], - dims: &[usize], - dim: usize, -) { - if dims.is_empty() { - if let Some(first) = grad_input.first_mut() { - *first = 0.0; - } - return; - } - - let dim_size = dims[dim]; - if dim_size == 0 { - return; - } - let after: usize = if dim + 1 >= dims.len() { - 1 - } else { - dims[dim + 1..].iter().product() - }; - let group = dim_size * after; - - if grad_output.len() < PAR_THRESHOLD { - for ((go_block, log_block), out_block) in grad_output - .chunks(group) - .zip(log_y.chunks(group)) - .zip(grad_input.chunks_mut(group)) - { - for a in 0..after { - let base = a; - let mut sum = 0.0f64; - for k in 0..dim_size { - let idx = base + k * after; - sum += go_block[idx]; - } - for k in 0..dim_size { - let idx = base + k * after; - let prob = log_block[idx].exp(); - out_block[idx] = go_block[idx] - prob * sum; - } - } - } - } else { - grad_output - .par_chunks(group) - .zip(log_y.par_chunks(group)) - .zip(grad_input.par_chunks_mut(group)) - .for_each(|((go_block, log_block), out_block)| { - for a in 0..after { - let base = a; - let mut sum = 0.0f64; - for k in 0..dim_size { - let idx = base + k * after; - sum += go_block[idx]; - } - for k in 0..dim_size { - let idx = base + k * after; - let prob = log_block[idx].exp(); - out_block[idx] = go_block[idx] - prob * sum; - } - } - }); - } -} - -fn broadcast_mask_index( - linear_idx: usize, - output_dims: &[usize], - output_strides: &[usize], - mask_dims: &[usize], - mask_strides: &[usize], -) -> usize { - if mask_dims.is_empty() { - return 0; - } - - let output_ndim = output_dims.len(); - let mask_ndim = mask_dims.len(); - let mut mask_index = 0usize; - - for i in 0..mask_ndim { - let output_dim_idx = output_ndim - 1 - i; - let mask_dim_idx = mask_ndim - 1 - i; - let stride = output_strides[output_dim_idx]; - let coord = if stride == 0 { - 0 - } else { - (linear_idx / stride) % output_dims[output_dim_idx] - }; - let mask_dim = mask_dims[mask_dim_idx]; - let mask_coord = if mask_dim == 1 { 0 } else { coord }; - mask_index += mask_coord * mask_strides[mask_dim_idx]; - } - - mask_index -} - -fn masked_log_softmax_backward_f32( - grad_output: &[f32], - log_y: &[f32], - mask: &[bool], - grad_input: &mut [f32], - dims: &[usize], - dim: usize, - mask_dims: &[usize], - output_strides: Option<&[usize]>, - mask_strides: Option<&[usize]>, -) { - if dims.is_empty() { - if let Some(first) = grad_input.first_mut() { - *first = 0.0; - } - return; - } - - let dim_size = dims[dim]; - if dim_size == 0 { - return; - } - - let after: usize = if dim + 1 >= dims.len() { - 1 - } else { - dims[dim + 1..].iter().product() - }; - let group = dim_size * after; - let same_shape = output_strides.is_none(); - - if grad_output.len() < PAR_THRESHOLD { - for (((go_block, log_block), out_block), block_idx) in grad_output - .chunks(group) - .zip(log_y.chunks(group)) - .zip(grad_input.chunks_mut(group)) - .zip(0..) - { - let block_offset = block_idx * group; - for a in 0..after { - let base = a; - let mut sum = 0.0f32; - for k in 0..dim_size { - let idx = base + k * after; - let linear_idx = block_offset + idx; - let masked = if same_shape { - mask[linear_idx] - } else { - let mask_index = broadcast_mask_index( - linear_idx, - dims, - output_strides.unwrap(), - mask_dims, - mask_strides.unwrap(), - ); - mask[mask_index] - }; - if !masked { - sum += go_block[idx]; - } - } - for k in 0..dim_size { - let idx = base + k * after; - let linear_idx = block_offset + idx; - let masked = if same_shape { - mask[linear_idx] - } else { - let mask_index = broadcast_mask_index( - linear_idx, - dims, - output_strides.unwrap(), - mask_dims, - mask_strides.unwrap(), - ); - mask[mask_index] - }; - if masked { - out_block[idx] = 0.0; - } else { - let prob = log_block[idx].exp(); - out_block[idx] = go_block[idx] - prob * sum; - } - } - } - } - } else { - grad_output - .par_chunks(group) - .zip(log_y.par_chunks(group)) - .zip(grad_input.par_chunks_mut(group)) - .enumerate() - .for_each(|(block_idx, ((go_block, log_block), out_block))| { - let block_offset = block_idx * group; - for a in 0..after { - let base = a; - let mut sum = 0.0f32; - for k in 0..dim_size { - let idx = base + k * after; - let linear_idx = block_offset + idx; - let masked = if same_shape { - mask[linear_idx] - } else { - let mask_index = broadcast_mask_index( - linear_idx, - dims, - output_strides.unwrap(), - mask_dims, - mask_strides.unwrap(), - ); - mask[mask_index] - }; - if !masked { - sum += go_block[idx]; - } - } - for k in 0..dim_size { - let idx = base + k * after; - let linear_idx = block_offset + idx; - let masked = if same_shape { - mask[linear_idx] - } else { - let mask_index = broadcast_mask_index( - linear_idx, - dims, - output_strides.unwrap(), - mask_dims, - mask_strides.unwrap(), - ); - mask[mask_index] - }; - if masked { - out_block[idx] = 0.0; - } else { - let prob = log_block[idx].exp(); - out_block[idx] = go_block[idx] - prob * sum; - } - } - } - }); - } -} - -fn masked_log_softmax_backward_f64( - grad_output: &[f64], - log_y: &[f64], - mask: &[bool], - grad_input: &mut [f64], - dims: &[usize], - dim: usize, - mask_dims: &[usize], - output_strides: Option<&[usize]>, - mask_strides: Option<&[usize]>, -) { - if dims.is_empty() { - if let Some(first) = grad_input.first_mut() { - *first = 0.0; - } - return; - } - - let dim_size = dims[dim]; - if dim_size == 0 { - return; - } - - let after: usize = if dim + 1 >= dims.len() { - 1 - } else { - dims[dim + 1..].iter().product() - }; - let group = dim_size * after; - let same_shape = output_strides.is_none(); - - if grad_output.len() < PAR_THRESHOLD { - for (((go_block, log_block), out_block), block_idx) in grad_output - .chunks(group) - .zip(log_y.chunks(group)) - .zip(grad_input.chunks_mut(group)) - .zip(0..) - { - let block_offset = block_idx * group; - for a in 0..after { - let base = a; - let mut sum = 0.0f64; - for k in 0..dim_size { - let idx = base + k * after; - let linear_idx = block_offset + idx; - let masked = if same_shape { - mask[linear_idx] - } else { - let mask_index = broadcast_mask_index( - linear_idx, - dims, - output_strides.unwrap(), - mask_dims, - mask_strides.unwrap(), - ); - mask[mask_index] - }; - if !masked { - sum += go_block[idx]; - } - } - for k in 0..dim_size { - let idx = base + k * after; - let linear_idx = block_offset + idx; - let masked = if same_shape { - mask[linear_idx] - } else { - let mask_index = broadcast_mask_index( - linear_idx, - dims, - output_strides.unwrap(), - mask_dims, - mask_strides.unwrap(), - ); - mask[mask_index] - }; - if masked { - out_block[idx] = 0.0; - } else { - let prob = log_block[idx].exp(); - out_block[idx] = go_block[idx] - prob * sum; - } - } - } - } - } else { - grad_output - .par_chunks(group) - .zip(log_y.par_chunks(group)) - .zip(grad_input.par_chunks_mut(group)) - .enumerate() - .for_each(|(block_idx, ((go_block, log_block), out_block))| { - let block_offset = block_idx * group; - for a in 0..after { - let base = a; - let mut sum = 0.0f64; - for k in 0..dim_size { - let idx = base + k * after; - let linear_idx = block_offset + idx; - let masked = if same_shape { - mask[linear_idx] - } else { - let mask_index = broadcast_mask_index( - linear_idx, - dims, - output_strides.unwrap(), - mask_dims, - mask_strides.unwrap(), - ); - mask[mask_index] - }; - if !masked { - sum += go_block[idx]; - } - } - for k in 0..dim_size { - let idx = base + k * after; - let linear_idx = block_offset + idx; - let masked = if same_shape { - mask[linear_idx] - } else { - let mask_index = broadcast_mask_index( - linear_idx, - dims, - output_strides.unwrap(), - mask_dims, - mask_strides.unwrap(), - ); - mask[mask_index] - }; - if masked { - out_block[idx] = 0.0; - } else { - let prob = log_block[idx].exp(); - out_block[idx] = go_block[idx] - prob * sum; - } - } - } - }); - } -} - -/// Gradient function for reshape operation -pub struct ReshapeBackward { - pub input_shape: Vec, - pub input_id: TensorId, -} - -impl GradientFunction for ReshapeBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // Reshape gradient: reshape back to original shape - let original_shape = Shape::new(self.input_shape.clone()); - let grad_input = crate::operations::shape_ops::reshape(grad_output, original_shape)?; - accumulate_grad(&mut gradients, self.input_id, grad_input)?; - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for repeat_interleave operation -pub struct RepeatInterleaveBackward { - pub input_shape: Vec, - pub repeats: Vec, - pub input_id: TensorId, - pub dim: usize, -} - -impl GradientFunction for RepeatInterleaveBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let grad_input = repeat_interleave_backward_impl( - grad_output, - &self.input_shape, - &self.repeats, - self.dim, - )?; - - let mut gradients = FxHashMap::default(); - accumulate_grad(&mut gradients, self.input_id, grad_input)?; - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for `min`/`max` reductions (global with `dim == None`, or -/// along a single `dim`). -/// -/// The gradient flows to every input element equal to the reduced extremum, -/// split equally among ties so the contributions sum to the upstream gradient. The -/// extremum, its selection mask and the tie count are recomputed from the stored -/// (detached) input, so nothing beyond the input needs to be retained. -pub struct MinMaxBackward { - pub input_id: TensorId, - pub input: Tensor, - pub dim: Option, - pub keepdim: bool, - pub is_max: bool, - pub nan_aware: bool, -} - -/// Route `grad_output` to every input element equal to the selected reduction -/// value (`reduced`, recomputed with keepdim so it broadcasts), splitting equally -/// among ties. Shared by min/max and median value reductions. -fn distribute_selection_grad( - input: &Tensor, - reduced: &Tensor, - grad_output: &Tensor, - dim: Option, - keepdim: bool, -) -> Result { - let input_shape = input.shape().dims().to_vec(); - let mask = crate::operations::comparison::eq(input, reduced)?; - let mask_f = mask.astype(input.dtype())?; - - let sum_dims = dim.map(|d| vec![d as isize]); - let count = reduction::sum(&mask_f, sum_dims, true)?; - - let dims_vec = dim.map(|d| vec![d]); - let grad_kd = expand_reduction_grad(grad_output, &input_shape, &dims_vec, keepdim)?; - let scaled = arithmetic::div(&grad_kd, &count)?; - arithmetic::mul(&mask_f, &scaled) -} - -impl GradientFunction for MinMaxBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - let input = &self.input; - - if input.numel() == 0 { - let zero = Tensor::zeros(input.shape().clone(), input.dtype(), input.device(), false); - accumulate_grad(&mut gradients, self.input_id, zero)?; - return Ok(gradients); - } - - let dim_isize = self.dim.map(|d| d as isize); - // Recompute the extremum with keepdim so it broadcasts against the input. - // NaN-aware reductions must recompute with the matching op, otherwise the - // propagated NaN would fail the equality mask and zero every gradient. - let reduced = match (self.is_max, self.nan_aware) { - (true, false) => reduction::max(input, dim_isize, true)?, - (false, false) => reduction::min(input, dim_isize, true)?, - (true, true) => reduction::nanmax(input, dim_isize, true)?, - (false, true) => reduction::nanmin(input, dim_isize, true)?, - }; - let grad_input = - distribute_selection_grad(input, &reduced, grad_output, self.dim, self.keepdim)?; - accumulate_grad(&mut gradients, self.input_id, grad_input)?; - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for `median`/`nanmedian` value reductions. The median is one -/// of the input elements, so the gradient flows to every element equal to it, -/// split over ties (a valid subgradient, matching the min/max convention). -pub struct MedianBackward { - pub input_id: TensorId, - pub input: Tensor, - pub dim: Option, - pub keepdim: bool, - pub nan_aware: bool, -} - -impl GradientFunction for MedianBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - let input = &self.input; - - if input.numel() == 0 { - let zero = Tensor::zeros(input.shape().clone(), input.dtype(), input.device(), false); - accumulate_grad(&mut gradients, self.input_id, zero)?; - return Ok(gradients); - } - - let dim_isize = self.dim.map(|d| d as isize); - let reduced = if self.nan_aware { - reduction::nanmedian(input, dim_isize, true)? - } else { - reduction::median(input, dim_isize, true)?.0 - }; - let grad_input = - distribute_selection_grad(input, &reduced, grad_output, self.dim, self.keepdim)?; - accumulate_grad(&mut gradients, self.input_id, grad_input)?; - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for the `quantile` reduction (global with `dim == None`, or -/// along a single `dim`). -/// -/// A quantile is a fixed linear combination of the two order statistics that -/// bracket the requested position, so the gradient routes back to the two -/// original elements occupying those sorted ranks with the interpolation weights -/// (`Lower`/`Higher`/`Nearest` collapse to a single element; `Midpoint` splits -/// evenly). Groups containing NaN produced NaN and receive no gradient. -pub struct QuantileBackward { - pub input_id: TensorId, - pub input: Tensor, - pub dim: Option, - pub q: f64, - pub interpolation: crate::operations::reduction::QuantileInterpolation, - pub nan_aware: bool, -} - -/// Sorted-rank indices and their gradient weights for a group of length `len`. -fn quantile_grad_coeffs( - len: usize, - q: f64, - interp: crate::operations::reduction::QuantileInterpolation, -) -> (usize, usize, f64, f64) { - use crate::operations::reduction::QuantileInterpolation as Qi; - if len <= 1 { - return (0, 0, 1.0, 0.0); - } - let pos = q * (len - 1) as f64; - let lower = pos.floor() as usize; - let upper = pos.ceil() as usize; - let weight = (pos - lower as f64).clamp(0.0, 1.0); - match interp { - Qi::Linear => (lower, upper, 1.0 - weight, weight), - Qi::Lower => (lower, upper, 1.0, 0.0), - Qi::Higher => (lower, upper, 0.0, 1.0), - Qi::Midpoint => (lower, upper, 0.5, 0.5), - Qi::Nearest => { +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::{ + error::{MinitensorError, Result}, + operations::{arithmetic, reduction}, + tensor::{DataType, Shape, Strides, Tensor, TensorData}, +}; +use rayon::prelude::*; +use rustc_hash::FxHashMap; +use smallvec::SmallVec; +use std::sync::Arc; + +impl GradientFunction for MaskedLogSoftmaxBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + let mut grad_data = TensorData::zeros_on_device( + self.output.numel(), + self.output.dtype(), + self.output.device(), + ); + + let output_dims = self.output.shape().dims(); + let mask_dims = self.mask.shape().dims(); + let same_shape = output_dims == mask_dims; + let output_strides = if same_shape { + None + } else { + Some(Strides::from_shape(self.output.shape())) + }; + let mask_strides = if same_shape { + None + } else { + Some(Strides::from_shape(self.mask.shape())) + }; + + match grad_output.dtype() { + DataType::Float32 => { + let go = grad_output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from grad_output") + })?; + let log_y = self.output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get f32 slice from masked log_softmax output", + ) + })?; + let mask_data = self.mask.data().as_bool_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get bool slice from mask tensor") + })?; + let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from grad_data", + ) + })?; + masked_log_softmax_backward_f32( + go, + log_y, + mask_data, + grad_slice, + output_dims, + self.dim, + mask_dims, + output_strides.as_ref().map(Strides::as_slice), + mask_strides.as_ref().map(Strides::as_slice), + ); + } + DataType::Float64 => { + let go = grad_output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from grad_output") + })?; + let log_y = self.output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get f64 slice from masked log_softmax output", + ) + })?; + let mask_data = self.mask.data().as_bool_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get bool slice from mask tensor") + })?; + let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from grad_data", + ) + })?; + masked_log_softmax_backward_f64( + go, + log_y, + mask_data, + grad_slice, + output_dims, + self.dim, + mask_dims, + output_strides.as_ref().map(Strides::as_slice), + mask_strides.as_ref().map(Strides::as_slice), + ); + } + _ => { + return Err(MinitensorError::invalid_operation( + "Masked log_softmax backward only supported for floating point tensors", + )); + } + } + + let grad_input = Tensor::new( + Arc::new(grad_data), + self.output.shape().clone(), + self.output.dtype(), + self.output.device(), + grad_output.requires_grad(), + ); + + accumulate_grad(&mut gradients, self.input_id, grad_input)?; + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for layer normalization +pub struct LayerNormBackward { + pub input_ids: SmallVec<[TensorId; 3]>, + pub input_id: TensorId, + pub weight_id: Option, + pub bias_id: Option, + pub normalized: Tensor, + pub inv_std: Tensor, + pub weight_broadcast: Option, + pub normalized_shape: Vec, + pub axis_start: usize, + pub element_count: usize, + pub input_requires_grad: bool, + pub weight_requires_grad: bool, + pub bias_requires_grad: bool, +} + +impl GradientFunction for LayerNormBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + + let grad_output_detached = grad_output.detach(); + let normalized = self.normalized.detach(); + + if self.element_count == 0 { + if self.input_requires_grad { + let zero = Tensor::zeros( + grad_output.shape().clone(), + grad_output.dtype(), + grad_output.device(), + false, + ); + accumulate_grad(&mut gradients, self.input_id, zero)?; + } + if self.weight_requires_grad + && let Some(weight_id) = self.weight_id + { + let zero = Tensor::zeros( + Shape::new(self.normalized_shape.clone()), + grad_output.dtype(), + grad_output.device(), + false, + ); + accumulate_grad(&mut gradients, weight_id, zero)?; + } + if self.bias_requires_grad + && let Some(bias_id) = self.bias_id + { + let zero = Tensor::zeros( + Shape::new(self.normalized_shape.clone()), + grad_output.dtype(), + grad_output.device(), + false, + ); + accumulate_grad(&mut gradients, bias_id, zero)?; + } + + return Ok(gradients); + } + + if self.input_requires_grad { + let mut grad_output_hat = if let Some(weight) = &self.weight_broadcast { + arithmetic::mul(&grad_output_detached, weight)? + } else { + grad_output_detached.clone() + }; + + let axes: Vec = (self.axis_start..grad_output_hat.ndim()) + .map(|d| d as isize) + .collect(); + let sum_grad = reduction::sum(&grad_output_hat, Some(axes.clone()), true)?; + let grad_norm_mul = arithmetic::mul(&grad_output_hat, &normalized)?; + let sum_grad_norm = reduction::sum(&grad_norm_mul, Some(axes), true)?; + + let count = self.element_count as f64; + let m_tensor = create_scalar_tensor(count, grad_output.dtype(), grad_output.device())?; + let inv_m_tensor = + create_scalar_tensor(1.0 / count, grad_output.dtype(), grad_output.device())?; + grad_output_hat = arithmetic::mul(&grad_output_hat, &m_tensor)?; + let tmp = arithmetic::sub(&grad_output_hat, &sum_grad)?; + let norm_term = arithmetic::mul(&normalized, &sum_grad_norm)?; + let numerator = arithmetic::sub(&tmp, &norm_term)?; + let grad_input = arithmetic::mul(&numerator, &self.inv_std)?; + let grad_input = arithmetic::mul(&grad_input, &inv_m_tensor)?; + accumulate_grad(&mut gradients, self.input_id, grad_input)?; + } + + if self.weight_requires_grad + && let Some(weight_id) = self.weight_id + { + let mut grad_weight = arithmetic::mul(&grad_output_detached, &normalized)?; + if self.axis_start > 0 { + let axes: Vec = (0..self.axis_start).map(|d| d as isize).collect(); + grad_weight = reduction::sum(&grad_weight, Some(axes), false)?; + } + if grad_weight.shape().dims() != self.normalized_shape.as_slice() { + grad_weight = grad_weight.view(Shape::new(self.normalized_shape.clone()))?; + } + accumulate_grad(&mut gradients, weight_id, grad_weight)?; + } + + if self.bias_requires_grad + && let Some(bias_id) = self.bias_id + { + let mut grad_bias = grad_output_detached.clone(); + if self.axis_start > 0 { + let axes: Vec = (0..self.axis_start).map(|d| d as isize).collect(); + grad_bias = reduction::sum(&grad_bias, Some(axes), false)?; + } + if grad_bias.shape().dims() != self.normalized_shape.as_slice() { + grad_bias = grad_bias.view(Shape::new(self.normalized_shape.clone()))?; + } + accumulate_grad(&mut gradients, bias_id, grad_bias)?; + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} + +pub(crate) fn softmax_backward_f32( + grad_output: &[f32], + y: &[f32], + grad_input: &mut [f32], + dims: &[usize], + dim: usize, +) { + if dims.is_empty() { + if let Some(first) = grad_input.first_mut() { + *first = 0.0; + } + return; + } + + let dim_size = dims[dim]; + if dim_size == 0 { + return; + } + let after: usize = if dim + 1 >= dims.len() { + 1 + } else { + dims[dim + 1..].iter().product() + }; + let group = dim_size * after; + if grad_output.len() < PAR_THRESHOLD { + for ((go_block, y_block), out_block) in grad_output + .chunks(group) + .zip(y.chunks(group)) + .zip(grad_input.chunks_mut(group)) + { + for a in 0..after { + let base = a; + let mut dot = 0.0f32; + for k in 0..dim_size { + let idx = base + k * after; + dot += go_block[idx] * y_block[idx]; + } + for k in 0..dim_size { + let idx = base + k * after; + out_block[idx] = y_block[idx] * (go_block[idx] - dot); + } + } + } + } else { + grad_output + .par_chunks(group) + .zip(y.par_chunks(group)) + .zip(grad_input.par_chunks_mut(group)) + .for_each(|((go_block, y_block), out_block)| { + for a in 0..after { + let base = a; + let mut dot = 0.0f32; + for k in 0..dim_size { + let idx = base + k * after; + dot += go_block[idx] * y_block[idx]; + } + for k in 0..dim_size { + let idx = base + k * after; + out_block[idx] = y_block[idx] * (go_block[idx] - dot); + } + } + }); + } +} + +pub(crate) fn softmax_backward_f64( + grad_output: &[f64], + y: &[f64], + grad_input: &mut [f64], + dims: &[usize], + dim: usize, +) { + if dims.is_empty() { + if let Some(first) = grad_input.first_mut() { + *first = 0.0; + } + return; + } + + let dim_size = dims[dim]; + if dim_size == 0 { + return; + } + let after: usize = if dim + 1 >= dims.len() { + 1 + } else { + dims[dim + 1..].iter().product() + }; + let group = dim_size * after; + if grad_output.len() < PAR_THRESHOLD { + for ((go_block, y_block), out_block) in grad_output + .chunks(group) + .zip(y.chunks(group)) + .zip(grad_input.chunks_mut(group)) + { + for a in 0..after { + let base = a; + let mut dot = 0.0f64; + for k in 0..dim_size { + let idx = base + k * after; + dot += go_block[idx] * y_block[idx]; + } + for k in 0..dim_size { + let idx = base + k * after; + out_block[idx] = y_block[idx] * (go_block[idx] - dot); + } + } + } + } else { + grad_output + .par_chunks(group) + .zip(y.par_chunks(group)) + .zip(grad_input.par_chunks_mut(group)) + .for_each(|((go_block, y_block), out_block)| { + for a in 0..after { + let base = a; + let mut dot = 0.0f64; + for k in 0..dim_size { + let idx = base + k * after; + dot += go_block[idx] * y_block[idx]; + } + for k in 0..dim_size { + let idx = base + k * after; + out_block[idx] = y_block[idx] * (go_block[idx] - dot); + } + } + }); + } +} + +pub(crate) fn log_softmax_backward_f32( + grad_output: &[f32], + log_y: &[f32], + grad_input: &mut [f32], + dims: &[usize], + dim: usize, +) { + if dims.is_empty() { + if let Some(first) = grad_input.first_mut() { + *first = 0.0; + } + return; + } + + let dim_size = dims[dim]; + if dim_size == 0 { + return; + } + let after: usize = if dim + 1 >= dims.len() { + 1 + } else { + dims[dim + 1..].iter().product() + }; + let group = dim_size * after; + + if grad_output.len() < PAR_THRESHOLD { + for ((go_block, log_block), out_block) in grad_output + .chunks(group) + .zip(log_y.chunks(group)) + .zip(grad_input.chunks_mut(group)) + { + for a in 0..after { + let base = a; + let mut sum = 0.0f32; + for k in 0..dim_size { + let idx = base + k * after; + sum += go_block[idx]; + } + for k in 0..dim_size { + let idx = base + k * after; + let prob = log_block[idx].exp(); + out_block[idx] = go_block[idx] - prob * sum; + } + } + } + } else { + grad_output + .par_chunks(group) + .zip(log_y.par_chunks(group)) + .zip(grad_input.par_chunks_mut(group)) + .for_each(|((go_block, log_block), out_block)| { + for a in 0..after { + let base = a; + let mut sum = 0.0f32; + for k in 0..dim_size { + let idx = base + k * after; + sum += go_block[idx]; + } + for k in 0..dim_size { + let idx = base + k * after; + let prob = log_block[idx].exp(); + out_block[idx] = go_block[idx] - prob * sum; + } + } + }); + } +} + +pub(crate) fn log_softmax_backward_f64( + grad_output: &[f64], + log_y: &[f64], + grad_input: &mut [f64], + dims: &[usize], + dim: usize, +) { + if dims.is_empty() { + if let Some(first) = grad_input.first_mut() { + *first = 0.0; + } + return; + } + + let dim_size = dims[dim]; + if dim_size == 0 { + return; + } + let after: usize = if dim + 1 >= dims.len() { + 1 + } else { + dims[dim + 1..].iter().product() + }; + let group = dim_size * after; + + if grad_output.len() < PAR_THRESHOLD { + for ((go_block, log_block), out_block) in grad_output + .chunks(group) + .zip(log_y.chunks(group)) + .zip(grad_input.chunks_mut(group)) + { + for a in 0..after { + let base = a; + let mut sum = 0.0f64; + for k in 0..dim_size { + let idx = base + k * after; + sum += go_block[idx]; + } + for k in 0..dim_size { + let idx = base + k * after; + let prob = log_block[idx].exp(); + out_block[idx] = go_block[idx] - prob * sum; + } + } + } + } else { + grad_output + .par_chunks(group) + .zip(log_y.par_chunks(group)) + .zip(grad_input.par_chunks_mut(group)) + .for_each(|((go_block, log_block), out_block)| { + for a in 0..after { + let base = a; + let mut sum = 0.0f64; + for k in 0..dim_size { + let idx = base + k * after; + sum += go_block[idx]; + } + for k in 0..dim_size { + let idx = base + k * after; + let prob = log_block[idx].exp(); + out_block[idx] = go_block[idx] - prob * sum; + } + } + }); + } +} + +fn broadcast_mask_index( + linear_idx: usize, + output_dims: &[usize], + output_strides: &[usize], + mask_dims: &[usize], + mask_strides: &[usize], +) -> usize { + if mask_dims.is_empty() { + return 0; + } + + let output_ndim = output_dims.len(); + let mask_ndim = mask_dims.len(); + let mut mask_index = 0usize; + + for i in 0..mask_ndim { + let output_dim_idx = output_ndim - 1 - i; + let mask_dim_idx = mask_ndim - 1 - i; + let stride = output_strides[output_dim_idx]; + let coord = if stride == 0 { + 0 + } else { + (linear_idx / stride) % output_dims[output_dim_idx] + }; + let mask_dim = mask_dims[mask_dim_idx]; + let mask_coord = if mask_dim == 1 { 0 } else { coord }; + mask_index += mask_coord * mask_strides[mask_dim_idx]; + } + + mask_index +} + +fn masked_log_softmax_backward_f32( + grad_output: &[f32], + log_y: &[f32], + mask: &[bool], + grad_input: &mut [f32], + dims: &[usize], + dim: usize, + mask_dims: &[usize], + output_strides: Option<&[usize]>, + mask_strides: Option<&[usize]>, +) { + if dims.is_empty() { + if let Some(first) = grad_input.first_mut() { + *first = 0.0; + } + return; + } + + let dim_size = dims[dim]; + if dim_size == 0 { + return; + } + + let after: usize = if dim + 1 >= dims.len() { + 1 + } else { + dims[dim + 1..].iter().product() + }; + let group = dim_size * after; + let same_shape = output_strides.is_none(); + + if grad_output.len() < PAR_THRESHOLD { + for (((go_block, log_block), out_block), block_idx) in grad_output + .chunks(group) + .zip(log_y.chunks(group)) + .zip(grad_input.chunks_mut(group)) + .zip(0..) + { + let block_offset = block_idx * group; + for a in 0..after { + let base = a; + let mut sum = 0.0f32; + for k in 0..dim_size { + let idx = base + k * after; + let linear_idx = block_offset + idx; + let masked = if same_shape { + mask[linear_idx] + } else { + let mask_index = broadcast_mask_index( + linear_idx, + dims, + output_strides.unwrap(), + mask_dims, + mask_strides.unwrap(), + ); + mask[mask_index] + }; + if !masked { + sum += go_block[idx]; + } + } + for k in 0..dim_size { + let idx = base + k * after; + let linear_idx = block_offset + idx; + let masked = if same_shape { + mask[linear_idx] + } else { + let mask_index = broadcast_mask_index( + linear_idx, + dims, + output_strides.unwrap(), + mask_dims, + mask_strides.unwrap(), + ); + mask[mask_index] + }; + if masked { + out_block[idx] = 0.0; + } else { + let prob = log_block[idx].exp(); + out_block[idx] = go_block[idx] - prob * sum; + } + } + } + } + } else { + grad_output + .par_chunks(group) + .zip(log_y.par_chunks(group)) + .zip(grad_input.par_chunks_mut(group)) + .enumerate() + .for_each(|(block_idx, ((go_block, log_block), out_block))| { + let block_offset = block_idx * group; + for a in 0..after { + let base = a; + let mut sum = 0.0f32; + for k in 0..dim_size { + let idx = base + k * after; + let linear_idx = block_offset + idx; + let masked = if same_shape { + mask[linear_idx] + } else { + let mask_index = broadcast_mask_index( + linear_idx, + dims, + output_strides.unwrap(), + mask_dims, + mask_strides.unwrap(), + ); + mask[mask_index] + }; + if !masked { + sum += go_block[idx]; + } + } + for k in 0..dim_size { + let idx = base + k * after; + let linear_idx = block_offset + idx; + let masked = if same_shape { + mask[linear_idx] + } else { + let mask_index = broadcast_mask_index( + linear_idx, + dims, + output_strides.unwrap(), + mask_dims, + mask_strides.unwrap(), + ); + mask[mask_index] + }; + if masked { + out_block[idx] = 0.0; + } else { + let prob = log_block[idx].exp(); + out_block[idx] = go_block[idx] - prob * sum; + } + } + } + }); + } +} + +fn masked_log_softmax_backward_f64( + grad_output: &[f64], + log_y: &[f64], + mask: &[bool], + grad_input: &mut [f64], + dims: &[usize], + dim: usize, + mask_dims: &[usize], + output_strides: Option<&[usize]>, + mask_strides: Option<&[usize]>, +) { + if dims.is_empty() { + if let Some(first) = grad_input.first_mut() { + *first = 0.0; + } + return; + } + + let dim_size = dims[dim]; + if dim_size == 0 { + return; + } + + let after: usize = if dim + 1 >= dims.len() { + 1 + } else { + dims[dim + 1..].iter().product() + }; + let group = dim_size * after; + let same_shape = output_strides.is_none(); + + if grad_output.len() < PAR_THRESHOLD { + for (((go_block, log_block), out_block), block_idx) in grad_output + .chunks(group) + .zip(log_y.chunks(group)) + .zip(grad_input.chunks_mut(group)) + .zip(0..) + { + let block_offset = block_idx * group; + for a in 0..after { + let base = a; + let mut sum = 0.0f64; + for k in 0..dim_size { + let idx = base + k * after; + let linear_idx = block_offset + idx; + let masked = if same_shape { + mask[linear_idx] + } else { + let mask_index = broadcast_mask_index( + linear_idx, + dims, + output_strides.unwrap(), + mask_dims, + mask_strides.unwrap(), + ); + mask[mask_index] + }; + if !masked { + sum += go_block[idx]; + } + } + for k in 0..dim_size { + let idx = base + k * after; + let linear_idx = block_offset + idx; + let masked = if same_shape { + mask[linear_idx] + } else { + let mask_index = broadcast_mask_index( + linear_idx, + dims, + output_strides.unwrap(), + mask_dims, + mask_strides.unwrap(), + ); + mask[mask_index] + }; + if masked { + out_block[idx] = 0.0; + } else { + let prob = log_block[idx].exp(); + out_block[idx] = go_block[idx] - prob * sum; + } + } + } + } + } else { + grad_output + .par_chunks(group) + .zip(log_y.par_chunks(group)) + .zip(grad_input.par_chunks_mut(group)) + .enumerate() + .for_each(|(block_idx, ((go_block, log_block), out_block))| { + let block_offset = block_idx * group; + for a in 0..after { + let base = a; + let mut sum = 0.0f64; + for k in 0..dim_size { + let idx = base + k * after; + let linear_idx = block_offset + idx; + let masked = if same_shape { + mask[linear_idx] + } else { + let mask_index = broadcast_mask_index( + linear_idx, + dims, + output_strides.unwrap(), + mask_dims, + mask_strides.unwrap(), + ); + mask[mask_index] + }; + if !masked { + sum += go_block[idx]; + } + } + for k in 0..dim_size { + let idx = base + k * after; + let linear_idx = block_offset + idx; + let masked = if same_shape { + mask[linear_idx] + } else { + let mask_index = broadcast_mask_index( + linear_idx, + dims, + output_strides.unwrap(), + mask_dims, + mask_strides.unwrap(), + ); + mask[mask_index] + }; + if masked { + out_block[idx] = 0.0; + } else { + let prob = log_block[idx].exp(); + out_block[idx] = go_block[idx] - prob * sum; + } + } + } + }); + } +} + +/// Gradient function for reshape operation +pub struct ReshapeBackward { + pub input_shape: Vec, + pub input_id: TensorId, +} + +impl GradientFunction for ReshapeBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + // Reshape gradient: reshape back to original shape + let original_shape = Shape::new(self.input_shape.clone()); + let grad_input = crate::operations::shape_ops::reshape(grad_output, original_shape)?; + accumulate_grad(&mut gradients, self.input_id, grad_input)?; + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for repeat_interleave operation +pub struct RepeatInterleaveBackward { + pub input_shape: Vec, + pub repeats: Vec, + pub input_id: TensorId, + pub dim: usize, +} + +impl GradientFunction for RepeatInterleaveBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let grad_input = repeat_interleave_backward_impl( + grad_output, + &self.input_shape, + &self.repeats, + self.dim, + )?; + + let mut gradients = FxHashMap::default(); + accumulate_grad(&mut gradients, self.input_id, grad_input)?; + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for `min`/`max` reductions (global with `dim == None`, or +/// along a single `dim`). +/// +/// The gradient flows to every input element equal to the reduced extremum, +/// split equally among ties so the contributions sum to the upstream gradient. The +/// extremum, its selection mask and the tie count are recomputed from the stored +/// (detached) input, so nothing beyond the input needs to be retained. +pub struct MinMaxBackward { + pub input_id: TensorId, + pub input: Tensor, + pub dim: Option, + pub keepdim: bool, + pub is_max: bool, + pub nan_aware: bool, +} + +/// Route `grad_output` to every input element equal to the selected reduction +/// value (`reduced`, recomputed with keepdim so it broadcasts), splitting equally +/// among ties. Shared by min/max and median value reductions. +fn distribute_selection_grad( + input: &Tensor, + reduced: &Tensor, + grad_output: &Tensor, + dim: Option, + keepdim: bool, +) -> Result { + let input_shape = input.shape().dims().to_vec(); + let mask = crate::operations::comparison::eq(input, reduced)?; + let mask_f = mask.astype(input.dtype())?; + + let sum_dims = dim.map(|d| vec![d as isize]); + let count = reduction::sum(&mask_f, sum_dims, true)?; + + let dims_vec = dim.map(|d| vec![d]); + let grad_kd = expand_reduction_grad(grad_output, &input_shape, &dims_vec, keepdim)?; + let scaled = arithmetic::div(&grad_kd, &count)?; + arithmetic::mul(&mask_f, &scaled) +} + +impl GradientFunction for MinMaxBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + let input = &self.input; + + if input.numel() == 0 { + let zero = Tensor::zeros(input.shape().clone(), input.dtype(), input.device(), false); + accumulate_grad(&mut gradients, self.input_id, zero)?; + return Ok(gradients); + } + + let dim_isize = self.dim.map(|d| d as isize); + // Recompute the extremum with keepdim so it broadcasts against the input. + // NaN-aware reductions must recompute with the matching op, otherwise the + // propagated NaN would fail the equality mask and zero every gradient. + let reduced = match (self.is_max, self.nan_aware) { + (true, false) => reduction::max(input, dim_isize, true)?, + (false, false) => reduction::min(input, dim_isize, true)?, + (true, true) => reduction::nanmax(input, dim_isize, true)?, + (false, true) => reduction::nanmin(input, dim_isize, true)?, + }; + let grad_input = + distribute_selection_grad(input, &reduced, grad_output, self.dim, self.keepdim)?; + accumulate_grad(&mut gradients, self.input_id, grad_input)?; + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for `median`/`nanmedian` value reductions. The median is one +/// of the input elements, so the gradient flows to every element equal to it, +/// split over ties (a valid subgradient, matching the min/max convention). +pub struct MedianBackward { + pub input_id: TensorId, + pub input: Tensor, + pub dim: Option, + pub keepdim: bool, + pub nan_aware: bool, +} + +impl GradientFunction for MedianBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + let input = &self.input; + + if input.numel() == 0 { + let zero = Tensor::zeros(input.shape().clone(), input.dtype(), input.device(), false); + accumulate_grad(&mut gradients, self.input_id, zero)?; + return Ok(gradients); + } + + let dim_isize = self.dim.map(|d| d as isize); + let reduced = if self.nan_aware { + reduction::nanmedian(input, dim_isize, true)? + } else { + reduction::median(input, dim_isize, true)?.0 + }; + let grad_input = + distribute_selection_grad(input, &reduced, grad_output, self.dim, self.keepdim)?; + accumulate_grad(&mut gradients, self.input_id, grad_input)?; + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for the `quantile` reduction (global with `dim == None`, or +/// along a single `dim`). +/// +/// A quantile is a fixed linear combination of the two order statistics that +/// bracket the requested position, so the gradient routes back to the two +/// original elements occupying those sorted ranks with the interpolation weights +/// (`Lower`/`Higher`/`Nearest` collapse to a single element; `Midpoint` splits +/// evenly). Groups containing NaN produced NaN and receive no gradient. +pub struct QuantileBackward { + pub input_id: TensorId, + pub input: Tensor, + pub dim: Option, + pub q: f64, + pub interpolation: crate::operations::reduction::QuantileInterpolation, + pub nan_aware: bool, +} + +/// Sorted-rank indices and their gradient weights for a group of length `len`. +fn quantile_grad_coeffs( + len: usize, + q: f64, + interp: crate::operations::reduction::QuantileInterpolation, +) -> (usize, usize, f64, f64) { + use crate::operations::reduction::QuantileInterpolation as Qi; + if len <= 1 { + return (0, 0, 1.0, 0.0); + } + let pos = q * (len - 1) as f64; + let lower = pos.floor() as usize; + let upper = pos.ceil() as usize; + let weight = (pos - lower as f64).clamp(0.0, 1.0); + match interp { + Qi::Linear => (lower, upper, 1.0 - weight, weight), + Qi::Lower => (lower, upper, 1.0, 0.0), + Qi::Higher => (lower, upper, 0.0, 1.0), + Qi::Midpoint => (lower, upper, 0.5, 0.5), + Qi::Nearest => { // Ties at weight == 0.5 round to the even index. - let nearest = if weight < 0.5 { - lower - } else if weight > 0.5 { - upper - } else { - lower + (lower & 1) - }; - if nearest == lower { - (lower, upper, 1.0, 0.0) - } else { - (lower, upper, 0.0, 1.0) - } - } - } -} - -impl GradientFunction for QuantileBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - let input = &self.input; - let numel = input.numel(); - let mut grad_data = TensorData::zeros_on_device(numel, input.dtype(), input.device()); - - if numel != 0 { - // Treat the global reduction as one group over the flattened tensor. - let dims = input.shape().dims(); - let (outer, inner, dim_size) = match self.dim { - None => (1usize, 1usize, numel), - Some(d) => { - let outer: usize = dims[..d].iter().product(); - let inner: usize = dims[d + 1..].iter().product(); - (outer, inner, dims[d]) - } - }; - let outer_stride = dim_size * inner; - - macro_rules! scatter { - ($slice:ident, $mut_slice:ident, $ty:ty) => {{ - let x = input.data().$slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to read input for quantile backward") - })?; - let go = grad_output.data().$slice().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to read grad_output for quantile backward", - ) - })?; - let gi = grad_data.$mut_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to write grad for quantile backward") - })?; - let mut buffer: Vec<(usize, $ty)> = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - buffer.clear(); - let mut skip_group = false; - for d in 0..dim_size { - let v = x[o * outer_stride + d * inner + r]; - if v.is_nan() { - if self.nan_aware { - // nanquantile ignores NaN entries. - continue; - } - // quantile propagates NaN, so the whole group's - // output is NaN and gets no gradient. - skip_group = true; - break; - } - buffer.push((d, v)); - } - if skip_group || buffer.is_empty() { - continue; - } - let (lo, up, c_lo, c_up) = - quantile_grad_coeffs(buffer.len(), self.q, self.interpolation); - // Only the elements at sorted ranks `lo` and `up` - // (adjacent) are needed, so select them in O(n) rather - // than fully sorting. NaN is already filtered, so the - // comparator never sees an incomparable value. - let cmp = |a: &(usize, $ty), b: &(usize, $ty)| { - a.1.partial_cmp(&b.1).unwrap() - }; - buffer.select_nth_unstable_by(up, cmp); - let d_up = buffer[up].0; - let d_lo = if lo == up { - d_up - } else { - // `lo == up - 1`: the lo-th order statistic is the - // largest of the elements left of `up`. - buffer[..up].select_nth_unstable_by(lo, cmp); - buffer[lo].0 - }; - let g = go[o * inner + r]; - gi[o * outer_stride + d_lo * inner + r] += g * c_lo as $ty; - gi[o * outer_stride + d_up * inner + r] += g * c_up as $ty; - } - } - }}; - } - - match input.dtype() { - DataType::Float32 => scatter!(as_f32_slice, as_f32_slice_mut, f32), - DataType::Float64 => scatter!(as_f64_slice, as_f64_slice_mut, f64), - _ => { - return Err(MinitensorError::invalid_operation( - "quantile backward only supported for floating point tensors", - )); - } - } - } - - let grad_input = Tensor::new( - Arc::new(grad_data), - input.shape().clone(), - input.dtype(), - input.device(), - false, - ); - accumulate_grad(&mut gradients, self.input_id, grad_input)?; - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} + let nearest = if weight < 0.5 { + lower + } else if weight > 0.5 { + upper + } else { + lower + (lower & 1) + }; + if nearest == lower { + (lower, upper, 1.0, 0.0) + } else { + (lower, upper, 0.0, 1.0) + } + } + } +} + +impl GradientFunction for QuantileBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + let input = &self.input; + let numel = input.numel(); + let mut grad_data = TensorData::zeros_on_device(numel, input.dtype(), input.device()); + + if numel != 0 { + // Treat the global reduction as one group over the flattened tensor. + let dims = input.shape().dims(); + let (outer, inner, dim_size) = match self.dim { + None => (1usize, 1usize, numel), + Some(d) => { + let outer: usize = dims[..d].iter().product(); + let inner: usize = dims[d + 1..].iter().product(); + (outer, inner, dims[d]) + } + }; + let outer_stride = dim_size * inner; + + macro_rules! scatter { + ($slice:ident, $mut_slice:ident, $ty:ty) => {{ + let x = input.data().$slice().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to read input for quantile backward", + ) + })?; + let go = grad_output.data().$slice().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to read grad_output for quantile backward", + ) + })?; + let gi = grad_data.$mut_slice().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to write grad for quantile backward", + ) + })?; + let mut buffer: Vec<(usize, $ty)> = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + buffer.clear(); + let mut skip_group = false; + for d in 0..dim_size { + let v = x[o * outer_stride + d * inner + r]; + if v.is_nan() { + if self.nan_aware { + // nanquantile ignores NaN entries. + continue; + } + // quantile propagates NaN, so the whole group's + // output is NaN and gets no gradient. + skip_group = true; + break; + } + buffer.push((d, v)); + } + if skip_group || buffer.is_empty() { + continue; + } + let (lo, up, c_lo, c_up) = + quantile_grad_coeffs(buffer.len(), self.q, self.interpolation); + // Only the elements at sorted ranks `lo` and `up` + // (adjacent) are needed, so select them in O(n) rather + // than fully sorting. NaN is already filtered, so the + // comparator never sees an incomparable value. + let cmp = + |a: &(usize, $ty), b: &(usize, $ty)| a.1.partial_cmp(&b.1).unwrap(); + buffer.select_nth_unstable_by(up, cmp); + let d_up = buffer[up].0; + let d_lo = if lo == up { + d_up + } else { + // `lo == up - 1`: the lo-th order statistic is the + // largest of the elements left of `up`. + buffer[..up].select_nth_unstable_by(lo, cmp); + buffer[lo].0 + }; + let g = go[o * inner + r]; + gi[o * outer_stride + d_lo * inner + r] += g * c_lo as $ty; + gi[o * outer_stride + d_up * inner + r] += g * c_up as $ty; + } + } + }}; + } + + match input.dtype() { + DataType::Float32 => scatter!(as_f32_slice, as_f32_slice_mut, f32), + DataType::Float64 => scatter!(as_f64_slice, as_f64_slice_mut, f64), + _ => { + return Err(MinitensorError::invalid_operation( + "quantile backward only supported for floating point tensors", + )); + } + } + } + + let grad_input = Tensor::new( + Arc::new(grad_data), + input.shape().clone(), + input.dtype(), + input.device(), + false, + ); + accumulate_grad(&mut gradients, self.input_id, grad_input)?; + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} diff --git a/engine/src/autograd/mod/shape.rs b/engine/src/autograd/mod/shape.rs index 4966a79f..c53373b5 100644 --- a/engine/src/autograd/mod/shape.rs +++ b/engine/src/autograd/mod/shape.rs @@ -1,1601 +1,1622 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -/// Map an output tap `(oh, ow, kh, kw)` to the input coordinate it reads, or -/// `None` when it lands in the zero padding. The padding bound already forces -/// `0 <= ih < in_h` (and likewise for `iw`), so no second range check is needed. -#[inline(always)] -fn conv_input_coord( - oh: usize, - ow: usize, - kh: usize, - kw: usize, - stride: (usize, usize), - padding: (usize, usize), - in_h: usize, - in_w: usize, -) -> Option<(usize, usize)> { - let h_in = oh * stride.0 + kh; - let w_in = ow * stride.1 + kw; - if h_in < padding.0 || w_in < padding.1 || h_in >= in_h + padding.0 || w_in >= in_w + padding.1 { - None - } else { - Some((h_in - padding.0, w_in - padding.1)) - } -} - -/// Gradient function for 2D convolution (`operations::conv2d`). -/// -/// Given `grad_output` of shape `[N, C_out, OH, OW]`, produces: -/// * `grad_input[n, ic, ih, iw] = Σ grad_output[n, oc, oh, ow] · weight[oc, ic, kh, kw]` -/// * `grad_weight[oc, ic, kh, kw] = Σ grad_output[n, oc, oh, ow] · input[n, ic, ih, iw]` -/// * `grad_bias[oc] = Σ grad_output[n, oc, oh, ow]` -/// -/// with the same padding/stride index mapping as the forward pass. Each gradient -/// is only computed when its operand requires it, and each is parallelised over a -/// race-free axis: `grad_input` over the batch (disjoint output slices), -/// `grad_weight`/`grad_bias` over the output channel (disjoint kernel/bias -/// slices). The padding/stride coordinate is hoisted out of the input-channel -/// loop since it does not depend on it. -pub struct Conv2dBackward { - pub input: Tensor, - pub weight: Tensor, - pub input_id: TensorId, - pub weight_id: TensorId, - pub bias_id: Option, - pub input_requires_grad: bool, - pub weight_requires_grad: bool, - pub bias_requires_grad: bool, - pub stride: (usize, usize), - pub padding: (usize, usize), - pub deps: SmallVec<[TensorId; 3]>, -} - -impl GradientFunction for Conv2dBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let in_dims = self.input.shape().dims(); - let w_dims = self.weight.shape().dims(); - let (batch, in_channels, in_h, in_w) = (in_dims[0], in_dims[1], in_dims[2], in_dims[3]); - let (out_channels, kernel_h, kernel_w) = (w_dims[0], w_dims[2], w_dims[3]); - let go_dims = grad_output.shape().dims(); - let (out_h, out_w) = (go_dims[2], go_dims[3]); - let stride = self.stride; - let padding = self.padding; - - let input = self - .input - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("conv2d backward expects f32 input"))?; - let weight = self - .weight - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("conv2d backward expects f32 weight"))?; - let go = grad_output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("conv2d backward expects f32 grad_output") - })?; - - let device = self.input.device(); - let mut gradients = FxHashMap::default(); - - // grad_input: parallel over batches, which own disjoint output regions. - if self.input_requires_grad { - let in_stride = in_channels * in_h * in_w; - let mut grad_input = vec![0f32; batch * in_stride]; - grad_input - .par_chunks_mut(in_stride) - .enumerate() - .for_each(|(n, gi)| { - for oc in 0..out_channels { - let w_base = oc * in_channels * kernel_h * kernel_w; - for oh in 0..out_h { - for ow in 0..out_w { - let g = - go[((n * out_channels + oc) * out_h + oh) * out_w + ow]; - for kh in 0..kernel_h { - for kw in 0..kernel_w { - if let Some((ih, iw)) = conv_input_coord( - oh, ow, kh, kw, stride, padding, in_h, in_w, - ) { - let spatial = ih * in_w + iw; - for ic in 0..in_channels { - let w_idx = w_base - + (ic * kernel_h + kh) * kernel_w - + kw; - gi[ic * in_h * in_w + spatial] += - g * weight[w_idx]; - } - } - } - } - } - } - } - }); - let grad = Tensor::new( - Arc::new(TensorData::from_vec_f32(grad_input, device)), - self.input.shape().clone(), - DataType::Float32, - device, - false, - ); - accumulate_grad(&mut gradients, self.input_id, grad)?; - } - - // grad_weight: parallel over output channels, which own disjoint slices. - if self.weight_requires_grad { - let w_stride = in_channels * kernel_h * kernel_w; - let mut grad_weight = vec![0f32; out_channels * w_stride]; - grad_weight - .par_chunks_mut(w_stride) - .enumerate() - .for_each(|(oc, gw)| { - for n in 0..batch { - for oh in 0..out_h { - for ow in 0..out_w { - let g = - go[((n * out_channels + oc) * out_h + oh) * out_w + ow]; - for kh in 0..kernel_h { - for kw in 0..kernel_w { - if let Some((ih, iw)) = conv_input_coord( - oh, ow, kh, kw, stride, padding, in_h, in_w, - ) { - let spatial = ih * in_w + iw; - for ic in 0..in_channels { - let in_idx = (n * in_channels + ic) * in_h * in_w - + spatial; - gw[(ic * kernel_h + kh) * kernel_w + kw] += - g * input[in_idx]; - } - } - } - } - } - } - } - }); - let grad = Tensor::new( - Arc::new(TensorData::from_vec_f32(grad_weight, device)), - self.weight.shape().clone(), - DataType::Float32, - device, - false, - ); - accumulate_grad(&mut gradients, self.weight_id, grad)?; - } - - // grad_bias: parallel over output channels. - if self.bias_requires_grad { - if let Some(bias_id) = self.bias_id { - let mut grad_bias = vec![0f32; out_channels]; - grad_bias.par_iter_mut().enumerate().for_each(|(oc, gb)| { - let mut sum = 0f32; - for n in 0..batch { - let base = (n * out_channels + oc) * out_h * out_w; - for k in 0..out_h * out_w { - sum += go[base + k]; - } - } - *gb = sum; - }); - let grad = Tensor::new( - Arc::new(TensorData::from_vec_f32(grad_bias, device)), - Shape::new(vec![out_channels]), - DataType::Float32, - device, - false, - ); - accumulate_grad(&mut gradients, bias_id, grad)?; - } - } - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.deps - } -} - -impl GradientFunction for PowBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(2); - - match self.output.dtype() { - DataType::Float32 => { - let base_slice = self.base.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from base tensor") - })?; - let exp_slice = self.exponent.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from exponent tensor") - })?; - let out_slice = self.output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from output tensor") - })?; - let grad_out = grad_output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from grad_output") - })?; - - if self.base_requires_grad { - let mut grad_data = TensorData::zeros_on_device( - self.base.numel(), - self.base.dtype(), - self.base.device(), - ); - let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from grad_data", - ) - })?; - - match self.broadcast { - PowBroadcast::None => { - let len = base_slice.len(); - if len < PAR_THRESHOLD { - for i in 0..len { - grad_slice[i] = exp_slice[i] - * base_slice[i].powf(exp_slice[i] - 1.0) - * grad_out[i]; - } - } else { - let base_ptr = base_slice.as_ptr() as usize; - let exp_ptr = exp_slice.as_ptr() as usize; - let go_ptr = grad_out.as_ptr() as usize; - let grad_ptr = grad_slice.as_mut_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let base_ptr = base_ptr as *const f32; - let exp_ptr = exp_ptr as *const f32; - let go_ptr = go_ptr as *const f32; - let grad_ptr = grad_ptr as *mut f32; - *grad_ptr.add(i) = *exp_ptr.add(i) - * (*base_ptr.add(i)).powf(*exp_ptr.add(i) - 1.0) - * *go_ptr.add(i); - }); - } - } - PowBroadcast::BaseScalar => { - let base_val = base_slice[0]; - let mut accum = 0.0_f32; - for i in 0..grad_out.len() { - accum += - exp_slice[i] * base_val.powf(exp_slice[i] - 1.0) * grad_out[i]; - } - grad_slice[0] = accum; - } - PowBroadcast::ExponentScalar => { - let exp_val = exp_slice[0]; - let len = base_slice.len(); - if len < PAR_THRESHOLD { - for i in 0..len { - grad_slice[i] = - exp_val * base_slice[i].powf(exp_val - 1.0) * grad_out[i]; - } - } else { - let base_ptr = base_slice.as_ptr() as usize; - let go_ptr = grad_out.as_ptr() as usize; - let grad_ptr = grad_slice.as_mut_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let base_ptr = base_ptr as *const f32; - let go_ptr = go_ptr as *const f32; - let grad_ptr = grad_ptr as *mut f32; - *grad_ptr.add(i) = exp_val - * (*base_ptr.add(i)).powf(exp_val - 1.0) - * *go_ptr.add(i); - }); - } - } - } - - let grad_tensor = Tensor::new( - Arc::new(grad_data), - self.base.shape().clone(), - self.base.dtype(), - self.base.device(), - false, - ); - accumulate_grad(&mut gradients, self.input_ids[0], grad_tensor)?; - } - - if self.exp_requires_grad { - let mut grad_data = TensorData::zeros_on_device( - self.exponent.numel(), - self.exponent.dtype(), - self.exponent.device(), - ); - let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from grad_data", - ) - })?; - - match self.broadcast { - PowBroadcast::None => { - let len = exp_slice.len(); - if len < PAR_THRESHOLD { - for i in 0..len { - grad_slice[i] = out_slice[i] * base_slice[i].ln() * grad_out[i]; - } - } else { - let out_ptr = out_slice.as_ptr() as usize; - let base_ptr = base_slice.as_ptr() as usize; - let go_ptr = grad_out.as_ptr() as usize; - let grad_ptr = grad_slice.as_mut_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let out_ptr = out_ptr as *const f32; - let base_ptr = base_ptr as *const f32; - let go_ptr = go_ptr as *const f32; - let grad_ptr = grad_ptr as *mut f32; - *grad_ptr.add(i) = - *out_ptr.add(i) * (*base_ptr.add(i)).ln() * *go_ptr.add(i); - }); - } - } - PowBroadcast::BaseScalar => { - let base_val = base_slice[0]; - for i in 0..grad_out.len() { - grad_slice[i] = out_slice[i] * base_val.ln() * grad_out[i]; - } - } - PowBroadcast::ExponentScalar => { - let mut accum = 0.0_f32; - for i in 0..grad_out.len() { - accum += out_slice[i] * base_slice[i].ln() * grad_out[i]; - } - grad_slice[0] = accum; - } - } - - let grad_tensor = Tensor::new( - Arc::new(grad_data), - self.exponent.shape().clone(), - self.exponent.dtype(), - self.exponent.device(), - false, - ); - accumulate_grad(&mut gradients, self.input_ids[1], grad_tensor)?; - } - } - DataType::Float64 => { - let base_slice = self.base.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from base tensor") - })?; - let exp_slice = self.exponent.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from exponent tensor") - })?; - let out_slice = self.output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from output tensor") - })?; - let grad_out = grad_output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from grad_output") - })?; - - if self.base_requires_grad { - let mut grad_data = TensorData::zeros_on_device( - self.base.numel(), - self.base.dtype(), - self.base.device(), - ); - let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from grad_data", - ) - })?; - - match self.broadcast { - PowBroadcast::None => { - let len = base_slice.len(); - if len < PAR_THRESHOLD { - for i in 0..len { - grad_slice[i] = exp_slice[i] - * base_slice[i].powf(exp_slice[i] - 1.0) - * grad_out[i]; - } - } else { - let base_ptr = base_slice.as_ptr() as usize; - let exp_ptr = exp_slice.as_ptr() as usize; - let go_ptr = grad_out.as_ptr() as usize; - let grad_ptr = grad_slice.as_mut_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let base_ptr = base_ptr as *const f64; - let exp_ptr = exp_ptr as *const f64; - let go_ptr = go_ptr as *const f64; - let grad_ptr = grad_ptr as *mut f64; - *grad_ptr.add(i) = *exp_ptr.add(i) - * (*base_ptr.add(i)).powf(*exp_ptr.add(i) - 1.0) - * *go_ptr.add(i); - }); - } - } - PowBroadcast::BaseScalar => { - let base_val = base_slice[0]; - let mut accum = 0.0_f64; - for i in 0..grad_out.len() { - accum += - exp_slice[i] * base_val.powf(exp_slice[i] - 1.0) * grad_out[i]; - } - grad_slice[0] = accum; - } - PowBroadcast::ExponentScalar => { - let exp_val = exp_slice[0]; - let len = base_slice.len(); - if len < PAR_THRESHOLD { - for i in 0..len { - grad_slice[i] = - exp_val * base_slice[i].powf(exp_val - 1.0) * grad_out[i]; - } - } else { - let base_ptr = base_slice.as_ptr() as usize; - let go_ptr = grad_out.as_ptr() as usize; - let grad_ptr = grad_slice.as_mut_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let base_ptr = base_ptr as *const f64; - let go_ptr = go_ptr as *const f64; - let grad_ptr = grad_ptr as *mut f64; - *grad_ptr.add(i) = exp_val - * (*base_ptr.add(i)).powf(exp_val - 1.0) - * *go_ptr.add(i); - }); - } - } - } - - let grad_tensor = Tensor::new( - Arc::new(grad_data), - self.base.shape().clone(), - self.base.dtype(), - self.base.device(), - false, - ); - accumulate_grad(&mut gradients, self.input_ids[0], grad_tensor)?; - } - - if self.exp_requires_grad { - let mut grad_data = TensorData::zeros_on_device( - self.exponent.numel(), - self.exponent.dtype(), - self.exponent.device(), - ); - let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from grad_data", - ) - })?; - - match self.broadcast { - PowBroadcast::None => { - let len = exp_slice.len(); - if len < PAR_THRESHOLD { - for i in 0..len { - grad_slice[i] = out_slice[i] * base_slice[i].ln() * grad_out[i]; - } - } else { - let out_ptr = out_slice.as_ptr() as usize; - let base_ptr = base_slice.as_ptr() as usize; - let go_ptr = grad_out.as_ptr() as usize; - let grad_ptr = grad_slice.as_mut_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let out_ptr = out_ptr as *const f64; - let base_ptr = base_ptr as *const f64; - let go_ptr = go_ptr as *const f64; - let grad_ptr = grad_ptr as *mut f64; - *grad_ptr.add(i) = - *out_ptr.add(i) * (*base_ptr.add(i)).ln() * *go_ptr.add(i); - }); - } - } - PowBroadcast::BaseScalar => { - let base_val = base_slice[0]; - for i in 0..grad_out.len() { - grad_slice[i] = out_slice[i] * base_val.ln() * grad_out[i]; - } - } - PowBroadcast::ExponentScalar => { - let mut accum = 0.0_f64; - for i in 0..grad_out.len() { - accum += out_slice[i] * base_slice[i].ln() * grad_out[i]; - } - grad_slice[0] = accum; - } - } - - let grad_tensor = Tensor::new( - Arc::new(grad_data), - self.exponent.shape().clone(), - self.exponent.dtype(), - self.exponent.device(), - false, - ); - accumulate_grad(&mut gradients, self.input_ids[1], grad_tensor)?; - } - } - _ => { - return Err(MinitensorError::invalid_operation( - "Power backward only supported for floating point tensors", - )); - } - } - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} - -/// Gradient function for Hardshrink -pub struct HardshrinkBackward { - pub input_id: TensorId, - pub mask: Vec, -} - -impl GradientFunction for HardshrinkBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - let mut grad_data = TensorData::zeros_on_device( - grad_output.numel(), - grad_output.dtype(), - grad_output.device(), - ); - - match grad_output.dtype() { - DataType::Float32 => { - let go = grad_output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from grad_output") - })?; - let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from grad_data", - ) - })?; - let len = go.len(); - if len < PAR_THRESHOLD { - for i in 0..len { - grad_slice[i] = if self.mask[i] { go[i] } else { 0.0 }; - } - } else { - let mask = &self.mask; - let go_ptr = go.as_ptr() as usize; - let grad_ptr = grad_slice.as_mut_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let go_ptr = go_ptr as *const f32; - let grad_ptr = grad_ptr as *mut f32; - if *mask.get_unchecked(i) { - *grad_ptr.add(i) = *go_ptr.add(i); - } else { - *grad_ptr.add(i) = 0.0; - } - }); - } - } - DataType::Float64 => { - let go = grad_output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from grad_output") - })?; - let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from grad_data", - ) - })?; - let len = go.len(); - if len < PAR_THRESHOLD { - for i in 0..len { - grad_slice[i] = if self.mask[i] { go[i] } else { 0.0 }; - } - } else { - let mask = &self.mask; - let go_ptr = go.as_ptr() as usize; - let grad_ptr = grad_slice.as_mut_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let go_ptr = go_ptr as *const f64; - let grad_ptr = grad_ptr as *mut f64; - if *mask.get_unchecked(i) { - *grad_ptr.add(i) = *go_ptr.add(i); - } else { - *grad_ptr.add(i) = 0.0; - } - }); - } - } - _ => { - return Err(MinitensorError::invalid_operation( - "hardshrink backward only supported for floating point tensors", - )); - } - } - - let grad_input = Tensor::new( - Arc::new(grad_data), - grad_output.shape().clone(), - grad_output.dtype(), - grad_output.device(), - grad_output.requires_grad(), - ); - gradients.insert(self.input_id, grad_input); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for nan_to_num. -pub struct NanToNumBackward { - pub input_id: TensorId, - pub finite_mask: Vec, -} - -impl GradientFunction for NanToNumBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - if self.finite_mask.len() != grad_output.numel() { - return Err(MinitensorError::gradient_error( - "nan_to_num backward mask length does not match gradient size", - )); - } - - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - let mut grad_data = TensorData::zeros_on_device( - grad_output.numel(), - grad_output.dtype(), - grad_output.device(), - ); - - match grad_output.dtype() { - DataType::Float32 => { - let grad = grad_output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from grad_output") - })?; - let out = grad_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from grad_data", - ) - })?; - apply_finite_mask(grad, out, &self.finite_mask); - } - DataType::Float64 => { - let grad = grad_output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from grad_output") - })?; - let out = grad_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from grad_data", - ) - })?; - apply_finite_mask(grad, out, &self.finite_mask); - } - _ => { - return Err(MinitensorError::invalid_operation( - "nan_to_num backward only supported for floating point tensors", - )); - } - } - - let grad_input = Tensor::new( - Arc::new(grad_data), - grad_output.shape().clone(), - grad_output.dtype(), - grad_output.device(), - false, - ); - gradients.insert(self.input_id, grad_input); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -#[inline(always)] -fn apply_finite_mask(grad: &[T], output: &mut [T], finite_mask: &[bool]) -where - T: Copy + Default + Send + Sync, -{ - debug_assert_eq!(grad.len(), output.len()); - debug_assert_eq!(grad.len(), finite_mask.len()); - - let len = grad.len(); - if len < PAR_THRESHOLD { - for i in 0..len { - output[i] = if finite_mask[i] { grad[i] } else { T::default() }; - } - } else { - grad.par_iter() - .zip(output.par_iter_mut()) - .zip(finite_mask.par_iter()) - .for_each(|((g, out), is_finite)| { - *out = if *is_finite { *g } else { T::default() }; - }); - } -} - -/// Gradient function for ReLU -pub struct ReluBackward { - pub input_id: TensorId, - pub mask: Vec, -} - -impl GradientFunction for ReluBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - let mut grad_data = TensorData::zeros_on_device( - grad_output.numel(), - grad_output.dtype(), - grad_output.device(), - ); - - match grad_output.dtype() { - DataType::Float32 => { - let go = grad_output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from grad_output") - })?; - let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from grad_data", - ) - })?; - let len = go.len(); - if len < PAR_THRESHOLD { - for i in 0..len { - grad_slice[i] = go[i] * if self.mask[i] { 1.0 } else { 0.0 }; - } - } else { - let mask = &self.mask; - let go_ptr = go.as_ptr() as usize; - let grad_ptr = grad_slice.as_mut_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let go_ptr = go_ptr as *const f32; - let grad_ptr = grad_ptr as *mut f32; - let m = if *mask.get_unchecked(i) { 1.0 } else { 0.0 }; - *grad_ptr.add(i) = *go_ptr.add(i) * m; - }); - } - } - DataType::Float64 => { - let go = grad_output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from grad_output") - })?; - let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from grad_data", - ) - })?; - let len = go.len(); - if len < PAR_THRESHOLD { - for i in 0..len { - grad_slice[i] = go[i] * if self.mask[i] { 1.0 } else { 0.0 }; - } - } else { - let mask = &self.mask; - let go_ptr = go.as_ptr() as usize; - let grad_ptr = grad_slice.as_mut_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let go_ptr = go_ptr as *const f64; - let grad_ptr = grad_ptr as *mut f64; - let m = if *mask.get_unchecked(i) { 1.0 } else { 0.0 }; - *grad_ptr.add(i) = *go_ptr.add(i) * m; - }); - } - } - _ => { - return Err(MinitensorError::invalid_operation( - "ReLU backward only supported for floating point tensors", - )); - } - } - - let grad_input = Tensor::new( - Arc::new(grad_data), - grad_output.shape().clone(), - grad_output.dtype(), - grad_output.device(), - grad_output.requires_grad(), - ); - gradients.insert(self.input_id, grad_input); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for LeakyReLU -pub struct LeakyReluBackward { - pub input_id: TensorId, - pub negative_slope: f64, - pub mask: Vec, -} - -impl GradientFunction for LeakyReluBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - let mut grad_data = TensorData::zeros_on_device( - grad_output.numel(), - grad_output.dtype(), - grad_output.device(), - ); - - match grad_output.dtype() { - DataType::Float32 => { - let go = grad_output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from grad_output") - })?; - let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from grad_data", - ) - })?; - let len = go.len(); - let slope = self.negative_slope as f32; - if len < PAR_THRESHOLD { - for i in 0..len { - grad_slice[i] = if self.mask[i] { go[i] } else { go[i] * slope }; - } - } else { - let mask = &self.mask; - let go_ptr = go.as_ptr() as usize; - let grad_ptr = grad_slice.as_mut_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let go_ptr = go_ptr as *const f32; - let grad_ptr = grad_ptr as *mut f32; - let val = if *mask.get_unchecked(i) { - *go_ptr.add(i) - } else { - *go_ptr.add(i) * slope - }; - *grad_ptr.add(i) = val; - }); - } - } - DataType::Float64 => { - let go = grad_output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from grad_output") - })?; - let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from grad_data", - ) - })?; - let len = go.len(); - let slope = self.negative_slope; - if len < PAR_THRESHOLD { - for i in 0..len { - grad_slice[i] = if self.mask[i] { go[i] } else { go[i] * slope }; - } - } else { - let mask = &self.mask; - let go_ptr = go.as_ptr() as usize; - let grad_ptr = grad_slice.as_mut_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let go_ptr = go_ptr as *const f64; - let grad_ptr = grad_ptr as *mut f64; - let val = if *mask.get_unchecked(i) { - *go_ptr.add(i) - } else { - *go_ptr.add(i) * slope - }; - *grad_ptr.add(i) = val; - }); - } - } - _ => { - return Err(MinitensorError::invalid_operation( - "LeakyReLU backward only supported for floating point tensors", - )); - } - } - - let grad_input = Tensor::new( - Arc::new(grad_data), - grad_output.shape().clone(), - grad_output.dtype(), - grad_output.device(), - grad_output.requires_grad(), - ); - gradients.insert(self.input_id, grad_input); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for the element-wise absolute value. -/// -/// `d/dx |x| = sign(x)` with the sub-gradient at `x == 0` taken as `0`. -/// The stored input shares storage with the forward input (a detached -/// clone), so no data is copied. -pub struct AbsBackward { - pub input_id: TensorId, - pub input: Tensor, -} - -impl GradientFunction for AbsBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - let mut grad_data = TensorData::zeros_on_device( - grad_output.numel(), - grad_output.dtype(), - grad_output.device(), - ); - - macro_rules! abs_grad { - ($slice:ident, $mut_slice:ident, $ty:ty) => {{ - let x = self.input.data().$slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to read input for abs backward") - })?; - let go = grad_output.data().$slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to read grad_output for abs backward") - })?; - let gi = grad_data.$mut_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to write grad for abs backward") - })?; - let sign = |v: $ty| -> $ty { - if v > 0.0 { - 1.0 - } else if v < 0.0 { - -1.0 - } else { - 0.0 - } - }; - if gi.len() < PAR_THRESHOLD { - for i in 0..gi.len() { - gi[i] = go[i] * sign(x[i]); - } - } else { - gi.par_iter_mut() - .zip(go.par_iter()) - .zip(x.par_iter()) - .for_each(|((g, &o), &v)| *g = o * sign(v)); - } - }}; - } - - match grad_output.dtype() { - DataType::Float32 => abs_grad!(as_f32_slice, as_f32_slice_mut, f32), - DataType::Float64 => abs_grad!(as_f64_slice, as_f64_slice_mut, f64), - _ => { - return Err(MinitensorError::invalid_operation( - "abs backward only supported for floating point tensors", - )); - } - } - - let grad_input = Tensor::new( - Arc::new(grad_data), - grad_output.shape().clone(), - grad_output.dtype(), - grad_output.device(), - false, - ); - gradients.insert(self.input_id, grad_input); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for `clamp`/`clip`. -/// -/// The gradient is passed through where the input lies inside the (inclusive) +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::{ + error::{MinitensorError, Result}, + operations::reduction, + tensor::{DataType, Shape, Strides, Tensor, TensorData}, +}; +use rayon::prelude::*; +use rustc_hash::FxHashMap; +use smallvec::SmallVec; +use std::sync::Arc; + +/// Map an output tap `(oh, ow, kh, kw)` to the input coordinate it reads, or +/// `None` when it lands in the zero padding. The padding bound already forces +/// `0 <= ih < in_h` (and likewise for `iw`), so no second range check is needed. +#[inline(always)] +fn conv_input_coord( + oh: usize, + ow: usize, + kh: usize, + kw: usize, + stride: (usize, usize), + padding: (usize, usize), + in_h: usize, + in_w: usize, +) -> Option<(usize, usize)> { + let h_in = oh * stride.0 + kh; + let w_in = ow * stride.1 + kw; + if h_in < padding.0 || w_in < padding.1 || h_in >= in_h + padding.0 || w_in >= in_w + padding.1 + { + None + } else { + Some((h_in - padding.0, w_in - padding.1)) + } +} + +/// Gradient function for 2D convolution (`operations::conv2d`). +/// +/// Given `grad_output` of shape `[N, C_out, OH, OW]`, produces: +/// * `grad_input[n, ic, ih, iw] = Σ grad_output[n, oc, oh, ow] · weight[oc, ic, kh, kw]` +/// * `grad_weight[oc, ic, kh, kw] = Σ grad_output[n, oc, oh, ow] · input[n, ic, ih, iw]` +/// * `grad_bias[oc] = Σ grad_output[n, oc, oh, ow]` +/// +/// with the same padding/stride index mapping as the forward pass. Each gradient +/// is only computed when its operand requires it, and each is parallelised over a +/// race-free axis: `grad_input` over the batch (disjoint output slices), +/// `grad_weight`/`grad_bias` over the output channel (disjoint kernel/bias +/// slices). The padding/stride coordinate is hoisted out of the input-channel +/// loop since it does not depend on it. +pub struct Conv2dBackward { + pub input: Tensor, + pub weight: Tensor, + pub input_id: TensorId, + pub weight_id: TensorId, + pub bias_id: Option, + pub input_requires_grad: bool, + pub weight_requires_grad: bool, + pub bias_requires_grad: bool, + pub stride: (usize, usize), + pub padding: (usize, usize), + pub deps: SmallVec<[TensorId; 3]>, +} + +impl GradientFunction for Conv2dBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let in_dims = self.input.shape().dims(); + let w_dims = self.weight.shape().dims(); + let (batch, in_channels, in_h, in_w) = (in_dims[0], in_dims[1], in_dims[2], in_dims[3]); + let (out_channels, kernel_h, kernel_w) = (w_dims[0], w_dims[2], w_dims[3]); + let go_dims = grad_output.shape().dims(); + let (out_h, out_w) = (go_dims[2], go_dims[3]); + let stride = self.stride; + let padding = self.padding; + + let input = + self.input.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("conv2d backward expects f32 input") + })?; + let weight = + self.weight.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("conv2d backward expects f32 weight") + })?; + let go = grad_output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("conv2d backward expects f32 grad_output") + })?; + + let device = self.input.device(); + let mut gradients = FxHashMap::default(); + + // grad_input: parallel over batches, which own disjoint output regions. + if self.input_requires_grad { + let in_stride = in_channels * in_h * in_w; + let mut grad_input = vec![0f32; batch * in_stride]; + grad_input + .par_chunks_mut(in_stride) + .enumerate() + .for_each(|(n, gi)| { + for oc in 0..out_channels { + let w_base = oc * in_channels * kernel_h * kernel_w; + for oh in 0..out_h { + for ow in 0..out_w { + let g = go[((n * out_channels + oc) * out_h + oh) * out_w + ow]; + for kh in 0..kernel_h { + for kw in 0..kernel_w { + if let Some((ih, iw)) = conv_input_coord( + oh, ow, kh, kw, stride, padding, in_h, in_w, + ) { + let spatial = ih * in_w + iw; + for ic in 0..in_channels { + let w_idx = + w_base + (ic * kernel_h + kh) * kernel_w + kw; + gi[ic * in_h * in_w + spatial] += g * weight[w_idx]; + } + } + } + } + } + } + } + }); + let grad = Tensor::new( + Arc::new(TensorData::from_vec_f32(grad_input, device)), + self.input.shape().clone(), + DataType::Float32, + device, + false, + ); + accumulate_grad(&mut gradients, self.input_id, grad)?; + } + + // grad_weight: parallel over output channels, which own disjoint slices. + if self.weight_requires_grad { + let w_stride = in_channels * kernel_h * kernel_w; + let mut grad_weight = vec![0f32; out_channels * w_stride]; + grad_weight + .par_chunks_mut(w_stride) + .enumerate() + .for_each(|(oc, gw)| { + for n in 0..batch { + for oh in 0..out_h { + for ow in 0..out_w { + let g = go[((n * out_channels + oc) * out_h + oh) * out_w + ow]; + for kh in 0..kernel_h { + for kw in 0..kernel_w { + if let Some((ih, iw)) = conv_input_coord( + oh, ow, kh, kw, stride, padding, in_h, in_w, + ) { + let spatial = ih * in_w + iw; + for ic in 0..in_channels { + let in_idx = + (n * in_channels + ic) * in_h * in_w + spatial; + gw[(ic * kernel_h + kh) * kernel_w + kw] += + g * input[in_idx]; + } + } + } + } + } + } + } + }); + let grad = Tensor::new( + Arc::new(TensorData::from_vec_f32(grad_weight, device)), + self.weight.shape().clone(), + DataType::Float32, + device, + false, + ); + accumulate_grad(&mut gradients, self.weight_id, grad)?; + } + + // grad_bias: parallel over output channels. + if self.bias_requires_grad + && let Some(bias_id) = self.bias_id + { + let mut grad_bias = vec![0f32; out_channels]; + grad_bias.par_iter_mut().enumerate().for_each(|(oc, gb)| { + let mut sum = 0f32; + for n in 0..batch { + let base = (n * out_channels + oc) * out_h * out_w; + for k in 0..out_h * out_w { + sum += go[base + k]; + } + } + *gb = sum; + }); + let grad = Tensor::new( + Arc::new(TensorData::from_vec_f32(grad_bias, device)), + Shape::new(vec![out_channels]), + DataType::Float32, + device, + false, + ); + accumulate_grad(&mut gradients, bias_id, grad)?; + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.deps + } +} + +impl GradientFunction for PowBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(2); + + match self.output.dtype() { + DataType::Float32 => { + let base_slice = self.base.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from base tensor") + })?; + let exp_slice = self.exponent.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from exponent tensor") + })?; + let out_slice = self.output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from output tensor") + })?; + let grad_out = grad_output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from grad_output") + })?; + + if self.base_requires_grad { + let mut grad_data = TensorData::zeros_on_device( + self.base.numel(), + self.base.dtype(), + self.base.device(), + ); + let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from grad_data", + ) + })?; + + match self.broadcast { + PowBroadcast::None => { + let len = base_slice.len(); + if len < PAR_THRESHOLD { + for i in 0..len { + grad_slice[i] = exp_slice[i] + * base_slice[i].powf(exp_slice[i] - 1.0) + * grad_out[i]; + } + } else { + let base_ptr = base_slice.as_ptr() as usize; + let exp_ptr = exp_slice.as_ptr() as usize; + let go_ptr = grad_out.as_ptr() as usize; + let grad_ptr = grad_slice.as_mut_ptr() as usize; + (0..len).into_par_iter().for_each(|i| unsafe { + let base_ptr = base_ptr as *const f32; + let exp_ptr = exp_ptr as *const f32; + let go_ptr = go_ptr as *const f32; + let grad_ptr = grad_ptr as *mut f32; + *grad_ptr.add(i) = *exp_ptr.add(i) + * (*base_ptr.add(i)).powf(*exp_ptr.add(i) - 1.0) + * *go_ptr.add(i); + }); + } + } + PowBroadcast::BaseScalar => { + let base_val = base_slice[0]; + let mut accum = 0.0_f32; + for i in 0..grad_out.len() { + accum += + exp_slice[i] * base_val.powf(exp_slice[i] - 1.0) * grad_out[i]; + } + grad_slice[0] = accum; + } + PowBroadcast::ExponentScalar => { + let exp_val = exp_slice[0]; + let len = base_slice.len(); + if len < PAR_THRESHOLD { + for i in 0..len { + grad_slice[i] = + exp_val * base_slice[i].powf(exp_val - 1.0) * grad_out[i]; + } + } else { + let base_ptr = base_slice.as_ptr() as usize; + let go_ptr = grad_out.as_ptr() as usize; + let grad_ptr = grad_slice.as_mut_ptr() as usize; + (0..len).into_par_iter().for_each(|i| unsafe { + let base_ptr = base_ptr as *const f32; + let go_ptr = go_ptr as *const f32; + let grad_ptr = grad_ptr as *mut f32; + *grad_ptr.add(i) = exp_val + * (*base_ptr.add(i)).powf(exp_val - 1.0) + * *go_ptr.add(i); + }); + } + } + } + + let grad_tensor = Tensor::new( + Arc::new(grad_data), + self.base.shape().clone(), + self.base.dtype(), + self.base.device(), + false, + ); + accumulate_grad(&mut gradients, self.input_ids[0], grad_tensor)?; + } + + if self.exp_requires_grad { + let mut grad_data = TensorData::zeros_on_device( + self.exponent.numel(), + self.exponent.dtype(), + self.exponent.device(), + ); + let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from grad_data", + ) + })?; + + match self.broadcast { + PowBroadcast::None => { + let len = exp_slice.len(); + if len < PAR_THRESHOLD { + for i in 0..len { + grad_slice[i] = out_slice[i] * base_slice[i].ln() * grad_out[i]; + } + } else { + let out_ptr = out_slice.as_ptr() as usize; + let base_ptr = base_slice.as_ptr() as usize; + let go_ptr = grad_out.as_ptr() as usize; + let grad_ptr = grad_slice.as_mut_ptr() as usize; + (0..len).into_par_iter().for_each(|i| unsafe { + let out_ptr = out_ptr as *const f32; + let base_ptr = base_ptr as *const f32; + let go_ptr = go_ptr as *const f32; + let grad_ptr = grad_ptr as *mut f32; + *grad_ptr.add(i) = + *out_ptr.add(i) * (*base_ptr.add(i)).ln() * *go_ptr.add(i); + }); + } + } + PowBroadcast::BaseScalar => { + let base_val = base_slice[0]; + for i in 0..grad_out.len() { + grad_slice[i] = out_slice[i] * base_val.ln() * grad_out[i]; + } + } + PowBroadcast::ExponentScalar => { + let mut accum = 0.0_f32; + for i in 0..grad_out.len() { + accum += out_slice[i] * base_slice[i].ln() * grad_out[i]; + } + grad_slice[0] = accum; + } + } + + let grad_tensor = Tensor::new( + Arc::new(grad_data), + self.exponent.shape().clone(), + self.exponent.dtype(), + self.exponent.device(), + false, + ); + accumulate_grad(&mut gradients, self.input_ids[1], grad_tensor)?; + } + } + DataType::Float64 => { + let base_slice = self.base.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from base tensor") + })?; + let exp_slice = self.exponent.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from exponent tensor") + })?; + let out_slice = self.output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from output tensor") + })?; + let grad_out = grad_output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from grad_output") + })?; + + if self.base_requires_grad { + let mut grad_data = TensorData::zeros_on_device( + self.base.numel(), + self.base.dtype(), + self.base.device(), + ); + let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from grad_data", + ) + })?; + + match self.broadcast { + PowBroadcast::None => { + let len = base_slice.len(); + if len < PAR_THRESHOLD { + for i in 0..len { + grad_slice[i] = exp_slice[i] + * base_slice[i].powf(exp_slice[i] - 1.0) + * grad_out[i]; + } + } else { + let base_ptr = base_slice.as_ptr() as usize; + let exp_ptr = exp_slice.as_ptr() as usize; + let go_ptr = grad_out.as_ptr() as usize; + let grad_ptr = grad_slice.as_mut_ptr() as usize; + (0..len).into_par_iter().for_each(|i| unsafe { + let base_ptr = base_ptr as *const f64; + let exp_ptr = exp_ptr as *const f64; + let go_ptr = go_ptr as *const f64; + let grad_ptr = grad_ptr as *mut f64; + *grad_ptr.add(i) = *exp_ptr.add(i) + * (*base_ptr.add(i)).powf(*exp_ptr.add(i) - 1.0) + * *go_ptr.add(i); + }); + } + } + PowBroadcast::BaseScalar => { + let base_val = base_slice[0]; + let mut accum = 0.0_f64; + for i in 0..grad_out.len() { + accum += + exp_slice[i] * base_val.powf(exp_slice[i] - 1.0) * grad_out[i]; + } + grad_slice[0] = accum; + } + PowBroadcast::ExponentScalar => { + let exp_val = exp_slice[0]; + let len = base_slice.len(); + if len < PAR_THRESHOLD { + for i in 0..len { + grad_slice[i] = + exp_val * base_slice[i].powf(exp_val - 1.0) * grad_out[i]; + } + } else { + let base_ptr = base_slice.as_ptr() as usize; + let go_ptr = grad_out.as_ptr() as usize; + let grad_ptr = grad_slice.as_mut_ptr() as usize; + (0..len).into_par_iter().for_each(|i| unsafe { + let base_ptr = base_ptr as *const f64; + let go_ptr = go_ptr as *const f64; + let grad_ptr = grad_ptr as *mut f64; + *grad_ptr.add(i) = exp_val + * (*base_ptr.add(i)).powf(exp_val - 1.0) + * *go_ptr.add(i); + }); + } + } + } + + let grad_tensor = Tensor::new( + Arc::new(grad_data), + self.base.shape().clone(), + self.base.dtype(), + self.base.device(), + false, + ); + accumulate_grad(&mut gradients, self.input_ids[0], grad_tensor)?; + } + + if self.exp_requires_grad { + let mut grad_data = TensorData::zeros_on_device( + self.exponent.numel(), + self.exponent.dtype(), + self.exponent.device(), + ); + let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from grad_data", + ) + })?; + + match self.broadcast { + PowBroadcast::None => { + let len = exp_slice.len(); + if len < PAR_THRESHOLD { + for i in 0..len { + grad_slice[i] = out_slice[i] * base_slice[i].ln() * grad_out[i]; + } + } else { + let out_ptr = out_slice.as_ptr() as usize; + let base_ptr = base_slice.as_ptr() as usize; + let go_ptr = grad_out.as_ptr() as usize; + let grad_ptr = grad_slice.as_mut_ptr() as usize; + (0..len).into_par_iter().for_each(|i| unsafe { + let out_ptr = out_ptr as *const f64; + let base_ptr = base_ptr as *const f64; + let go_ptr = go_ptr as *const f64; + let grad_ptr = grad_ptr as *mut f64; + *grad_ptr.add(i) = + *out_ptr.add(i) * (*base_ptr.add(i)).ln() * *go_ptr.add(i); + }); + } + } + PowBroadcast::BaseScalar => { + let base_val = base_slice[0]; + for i in 0..grad_out.len() { + grad_slice[i] = out_slice[i] * base_val.ln() * grad_out[i]; + } + } + PowBroadcast::ExponentScalar => { + let mut accum = 0.0_f64; + for i in 0..grad_out.len() { + accum += out_slice[i] * base_slice[i].ln() * grad_out[i]; + } + grad_slice[0] = accum; + } + } + + let grad_tensor = Tensor::new( + Arc::new(grad_data), + self.exponent.shape().clone(), + self.exponent.dtype(), + self.exponent.device(), + false, + ); + accumulate_grad(&mut gradients, self.input_ids[1], grad_tensor)?; + } + } + _ => { + return Err(MinitensorError::invalid_operation( + "Power backward only supported for floating point tensors", + )); + } + } + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} + +/// Gradient function for Hardshrink +pub struct HardshrinkBackward { + pub input_id: TensorId, + pub mask: Vec, +} + +impl GradientFunction for HardshrinkBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + let mut grad_data = TensorData::zeros_on_device( + grad_output.numel(), + grad_output.dtype(), + grad_output.device(), + ); + + match grad_output.dtype() { + DataType::Float32 => { + let go = grad_output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from grad_output") + })?; + let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from grad_data", + ) + })?; + let len = go.len(); + if len < PAR_THRESHOLD { + for i in 0..len { + grad_slice[i] = if self.mask[i] { go[i] } else { 0.0 }; + } + } else { + let mask = &self.mask; + let go_ptr = go.as_ptr() as usize; + let grad_ptr = grad_slice.as_mut_ptr() as usize; + (0..len).into_par_iter().for_each(|i| unsafe { + let go_ptr = go_ptr as *const f32; + let grad_ptr = grad_ptr as *mut f32; + if *mask.get_unchecked(i) { + *grad_ptr.add(i) = *go_ptr.add(i); + } else { + *grad_ptr.add(i) = 0.0; + } + }); + } + } + DataType::Float64 => { + let go = grad_output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from grad_output") + })?; + let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from grad_data", + ) + })?; + let len = go.len(); + if len < PAR_THRESHOLD { + for i in 0..len { + grad_slice[i] = if self.mask[i] { go[i] } else { 0.0 }; + } + } else { + let mask = &self.mask; + let go_ptr = go.as_ptr() as usize; + let grad_ptr = grad_slice.as_mut_ptr() as usize; + (0..len).into_par_iter().for_each(|i| unsafe { + let go_ptr = go_ptr as *const f64; + let grad_ptr = grad_ptr as *mut f64; + if *mask.get_unchecked(i) { + *grad_ptr.add(i) = *go_ptr.add(i); + } else { + *grad_ptr.add(i) = 0.0; + } + }); + } + } + _ => { + return Err(MinitensorError::invalid_operation( + "hardshrink backward only supported for floating point tensors", + )); + } + } + + let grad_input = Tensor::new( + Arc::new(grad_data), + grad_output.shape().clone(), + grad_output.dtype(), + grad_output.device(), + grad_output.requires_grad(), + ); + gradients.insert(self.input_id, grad_input); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for nan_to_num. +pub struct NanToNumBackward { + pub input_id: TensorId, + pub finite_mask: Vec, +} + +impl GradientFunction for NanToNumBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + if self.finite_mask.len() != grad_output.numel() { + return Err(MinitensorError::gradient_error( + "nan_to_num backward mask length does not match gradient size", + )); + } + + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + let mut grad_data = TensorData::zeros_on_device( + grad_output.numel(), + grad_output.dtype(), + grad_output.device(), + ); + + match grad_output.dtype() { + DataType::Float32 => { + let grad = grad_output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from grad_output") + })?; + let out = grad_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from grad_data", + ) + })?; + apply_finite_mask(grad, out, &self.finite_mask); + } + DataType::Float64 => { + let grad = grad_output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from grad_output") + })?; + let out = grad_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from grad_data", + ) + })?; + apply_finite_mask(grad, out, &self.finite_mask); + } + _ => { + return Err(MinitensorError::invalid_operation( + "nan_to_num backward only supported for floating point tensors", + )); + } + } + + let grad_input = Tensor::new( + Arc::new(grad_data), + grad_output.shape().clone(), + grad_output.dtype(), + grad_output.device(), + false, + ); + gradients.insert(self.input_id, grad_input); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +#[inline(always)] +fn apply_finite_mask(grad: &[T], output: &mut [T], finite_mask: &[bool]) +where + T: Copy + Default + Send + Sync, +{ + debug_assert_eq!(grad.len(), output.len()); + debug_assert_eq!(grad.len(), finite_mask.len()); + + let len = grad.len(); + if len < PAR_THRESHOLD { + for i in 0..len { + output[i] = if finite_mask[i] { + grad[i] + } else { + T::default() + }; + } + } else { + grad.par_iter() + .zip(output.par_iter_mut()) + .zip(finite_mask.par_iter()) + .for_each(|((g, out), is_finite)| { + *out = if *is_finite { *g } else { T::default() }; + }); + } +} + +/// Gradient function for ReLU +pub struct ReluBackward { + pub input_id: TensorId, + pub mask: Vec, +} + +impl GradientFunction for ReluBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + let mut grad_data = TensorData::zeros_on_device( + grad_output.numel(), + grad_output.dtype(), + grad_output.device(), + ); + + match grad_output.dtype() { + DataType::Float32 => { + let go = grad_output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from grad_output") + })?; + let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from grad_data", + ) + })?; + let len = go.len(); + if len < PAR_THRESHOLD { + for i in 0..len { + grad_slice[i] = go[i] * if self.mask[i] { 1.0 } else { 0.0 }; + } + } else { + let mask = &self.mask; + let go_ptr = go.as_ptr() as usize; + let grad_ptr = grad_slice.as_mut_ptr() as usize; + (0..len).into_par_iter().for_each(|i| unsafe { + let go_ptr = go_ptr as *const f32; + let grad_ptr = grad_ptr as *mut f32; + let m = if *mask.get_unchecked(i) { 1.0 } else { 0.0 }; + *grad_ptr.add(i) = *go_ptr.add(i) * m; + }); + } + } + DataType::Float64 => { + let go = grad_output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from grad_output") + })?; + let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from grad_data", + ) + })?; + let len = go.len(); + if len < PAR_THRESHOLD { + for i in 0..len { + grad_slice[i] = go[i] * if self.mask[i] { 1.0 } else { 0.0 }; + } + } else { + let mask = &self.mask; + let go_ptr = go.as_ptr() as usize; + let grad_ptr = grad_slice.as_mut_ptr() as usize; + (0..len).into_par_iter().for_each(|i| unsafe { + let go_ptr = go_ptr as *const f64; + let grad_ptr = grad_ptr as *mut f64; + let m = if *mask.get_unchecked(i) { 1.0 } else { 0.0 }; + *grad_ptr.add(i) = *go_ptr.add(i) * m; + }); + } + } + _ => { + return Err(MinitensorError::invalid_operation( + "ReLU backward only supported for floating point tensors", + )); + } + } + + let grad_input = Tensor::new( + Arc::new(grad_data), + grad_output.shape().clone(), + grad_output.dtype(), + grad_output.device(), + grad_output.requires_grad(), + ); + gradients.insert(self.input_id, grad_input); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for LeakyReLU +pub struct LeakyReluBackward { + pub input_id: TensorId, + pub negative_slope: f64, + pub mask: Vec, +} + +impl GradientFunction for LeakyReluBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + let mut grad_data = TensorData::zeros_on_device( + grad_output.numel(), + grad_output.dtype(), + grad_output.device(), + ); + + match grad_output.dtype() { + DataType::Float32 => { + let go = grad_output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from grad_output") + })?; + let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from grad_data", + ) + })?; + let len = go.len(); + let slope = self.negative_slope as f32; + if len < PAR_THRESHOLD { + for i in 0..len { + grad_slice[i] = if self.mask[i] { go[i] } else { go[i] * slope }; + } + } else { + let mask = &self.mask; + let go_ptr = go.as_ptr() as usize; + let grad_ptr = grad_slice.as_mut_ptr() as usize; + (0..len).into_par_iter().for_each(|i| unsafe { + let go_ptr = go_ptr as *const f32; + let grad_ptr = grad_ptr as *mut f32; + let val = if *mask.get_unchecked(i) { + *go_ptr.add(i) + } else { + *go_ptr.add(i) * slope + }; + *grad_ptr.add(i) = val; + }); + } + } + DataType::Float64 => { + let go = grad_output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from grad_output") + })?; + let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from grad_data", + ) + })?; + let len = go.len(); + let slope = self.negative_slope; + if len < PAR_THRESHOLD { + for i in 0..len { + grad_slice[i] = if self.mask[i] { go[i] } else { go[i] * slope }; + } + } else { + let mask = &self.mask; + let go_ptr = go.as_ptr() as usize; + let grad_ptr = grad_slice.as_mut_ptr() as usize; + (0..len).into_par_iter().for_each(|i| unsafe { + let go_ptr = go_ptr as *const f64; + let grad_ptr = grad_ptr as *mut f64; + let val = if *mask.get_unchecked(i) { + *go_ptr.add(i) + } else { + *go_ptr.add(i) * slope + }; + *grad_ptr.add(i) = val; + }); + } + } + _ => { + return Err(MinitensorError::invalid_operation( + "LeakyReLU backward only supported for floating point tensors", + )); + } + } + + let grad_input = Tensor::new( + Arc::new(grad_data), + grad_output.shape().clone(), + grad_output.dtype(), + grad_output.device(), + grad_output.requires_grad(), + ); + gradients.insert(self.input_id, grad_input); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for the element-wise absolute value. +/// +/// `d/dx |x| = sign(x)` with the sub-gradient at `x == 0` taken as `0`. +/// The stored input shares storage with the forward input (a detached +/// clone), so no data is copied. +pub struct AbsBackward { + pub input_id: TensorId, + pub input: Tensor, +} + +impl GradientFunction for AbsBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + let mut grad_data = TensorData::zeros_on_device( + grad_output.numel(), + grad_output.dtype(), + grad_output.device(), + ); + + macro_rules! abs_grad { + ($slice:ident, $mut_slice:ident, $ty:ty) => {{ + let x = self.input.data().$slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to read input for abs backward") + })?; + let go = grad_output.data().$slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to read grad_output for abs backward") + })?; + let gi = grad_data.$mut_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to write grad for abs backward") + })?; + let sign = |v: $ty| -> $ty { + if v > 0.0 { + 1.0 + } else if v < 0.0 { + -1.0 + } else { + 0.0 + } + }; + if gi.len() < PAR_THRESHOLD { + for i in 0..gi.len() { + gi[i] = go[i] * sign(x[i]); + } + } else { + gi.par_iter_mut() + .zip(go.par_iter()) + .zip(x.par_iter()) + .for_each(|((g, &o), &v)| *g = o * sign(v)); + } + }}; + } + + match grad_output.dtype() { + DataType::Float32 => abs_grad!(as_f32_slice, as_f32_slice_mut, f32), + DataType::Float64 => abs_grad!(as_f64_slice, as_f64_slice_mut, f64), + _ => { + return Err(MinitensorError::invalid_operation( + "abs backward only supported for floating point tensors", + )); + } + } + + let grad_input = Tensor::new( + Arc::new(grad_data), + grad_output.shape().clone(), + grad_output.dtype(), + grad_output.device(), + false, + ); + gradients.insert(self.input_id, grad_input); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for `clamp`/`clip`. +/// +/// The gradient is passed through where the input lies inside the (inclusive) /// clamp bounds and zeroed where it was saturated. Either -/// bound may be absent (`clamp_min`/`clamp_max`). -pub struct ClampBackward { - pub input_id: TensorId, - pub input: Tensor, - pub min: Option, - pub max: Option, -} - -impl GradientFunction for ClampBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - let mut grad_data = TensorData::zeros_on_device( - grad_output.numel(), - grad_output.dtype(), - grad_output.device(), - ); - - macro_rules! clamp_grad { - ($slice:ident, $mut_slice:ident, $ty:ty) => {{ - let x = self.input.data().$slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to read input for clamp backward") - })?; - let go = grad_output.data().$slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to read grad_output for clamp backward") - })?; - let gi = grad_data.$mut_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to write grad for clamp backward") - })?; - let min = self.min.map(|m| m as $ty); - let max = self.max.map(|m| m as $ty); - let passes = move |v: $ty| -> bool { - min.map_or(true, |m| v >= m) && max.map_or(true, |m| v <= m) - }; - if gi.len() < PAR_THRESHOLD { - for i in 0..gi.len() { - gi[i] = if passes(x[i]) { go[i] } else { 0.0 }; - } - } else { - gi.par_iter_mut() - .zip(go.par_iter()) - .zip(x.par_iter()) - .for_each(|((g, &o), &v)| *g = if passes(v) { o } else { 0.0 }); - } - }}; - } - - match grad_output.dtype() { - DataType::Float32 => clamp_grad!(as_f32_slice, as_f32_slice_mut, f32), - DataType::Float64 => clamp_grad!(as_f64_slice, as_f64_slice_mut, f64), - _ => { - return Err(MinitensorError::invalid_operation( - "clamp backward only supported for floating point tensors", - )); - } - } - - let grad_input = Tensor::new( - Arc::new(grad_data), - grad_output.shape().clone(), - grad_output.dtype(), - grad_output.device(), - false, - ); - gradients.insert(self.input_id, grad_input); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Scatter-add `grad_output` back to the source positions selected along `dim`. -/// -/// `indices[i]` is the source position (along `dim`) that produced output row `i` -/// for every outer/inner coordinate. This is the shared backward for -/// `index_select` and `slice` (and, transitively, `narrow`/`flip`/`roll`). -/// Duplicated source indices accumulate, matching the forward gather semantics. -fn index_select_backward_grad( - grad_output: &Tensor, - input_shape: &[usize], - dim: usize, - indices: &[usize], -) -> Result { - let numel: usize = input_shape.iter().product(); - let mut grad_data = - TensorData::zeros_on_device(numel, grad_output.dtype(), grad_output.device()); - - let dim_size = input_shape[dim]; - let inner: usize = input_shape[dim + 1..].iter().product(); - let out_dim = indices.len(); - - if numel != 0 && out_dim != 0 && inner != 0 { - let in_chunk = dim_size * inner; - let out_chunk = out_dim * inner; - - macro_rules! fill { - ($slice:ident, $mut_slice:ident) => {{ - let go = grad_output.data().$slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to read grad_output for index backward") - })?; - let gi = grad_data.$mut_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to write grad for index backward") - })?; - gi.par_chunks_mut(in_chunk) - .enumerate() - .for_each(|(o, gi_chunk)| { - let go_chunk = &go[o * out_chunk..(o + 1) * out_chunk]; - for (i, &idx) in indices.iter().enumerate() { - let dst = idx * inner; - let src = i * inner; - for j in 0..inner { - gi_chunk[dst + j] += go_chunk[src + j]; - } - } - }); - }}; - } - - match grad_output.dtype() { - DataType::Float32 => fill!(as_f32_slice, as_f32_slice_mut), - DataType::Float64 => fill!(as_f64_slice, as_f64_slice_mut), - _ => { - return Err(MinitensorError::invalid_operation( - "index/slice backward only supported for floating point tensors", - )); - } - } - } - - Ok(Tensor::new( - Arc::new(grad_data), - Shape::new(input_shape.to_vec()), - grad_output.dtype(), - grad_output.device(), - false, - )) -} - -/// Scatter-add `grad_output` back to the input positions named by a full `index` -/// tensor (`gather` backward, also reused by min/max/sort/topk along a dim). The -/// `index` slice is laid out identically to `grad_output`; entry `index[..]` is -/// the source coordinate along `dim`. Colliding indices accumulate. -fn gather_backward_grad( - grad_output: &Tensor, - input_shape: &[usize], - dim: usize, - index: &[i64], -) -> Result { - let numel: usize = input_shape.iter().product(); - let mut grad_data = - TensorData::zeros_on_device(numel, grad_output.dtype(), grad_output.device()); - - let dim_size = input_shape[dim]; - let inner: usize = input_shape[dim + 1..].iter().product(); - // The index tensor shares `grad_output`'s shape, so the output extent along - // `dim` is read directly from it. - let out_dim = grad_output.shape().dims()[dim]; - - if numel != 0 && !index.is_empty() && inner != 0 { - let in_chunk = dim_size * inner; - let out_chunk = out_dim * inner; - - macro_rules! fill { - ($slice:ident, $mut_slice:ident) => {{ - let go = grad_output.data().$slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to read grad_output for gather backward") - })?; - let gi = grad_data.$mut_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to write grad for gather backward") - })?; - gi.par_chunks_mut(in_chunk) - .enumerate() - .for_each(|(o, gi_chunk)| { - let go_chunk = &go[o * out_chunk..(o + 1) * out_chunk]; - let idx_chunk = &index[o * out_chunk..(o + 1) * out_chunk]; - for i in 0..out_dim { - for j in 0..inner { - let pos = i * inner + j; - let src_idx = idx_chunk[pos] as usize; - gi_chunk[src_idx * inner + j] += go_chunk[pos]; - } - } - }); - }}; - } - - match grad_output.dtype() { - DataType::Float32 => fill!(as_f32_slice, as_f32_slice_mut), - DataType::Float64 => fill!(as_f64_slice, as_f64_slice_mut), - _ => { - return Err(MinitensorError::invalid_operation( - "gather backward only supported for floating point tensors", - )); - } - } - } - - Ok(Tensor::new( - Arc::new(grad_data), - Shape::new(input_shape.to_vec()), - grad_output.dtype(), - grad_output.device(), - false, - )) -} - -/// Gradient function for `index_select` and `slice` (source indices along `dim`). -pub struct IndexSelectBackward { - pub input_id: TensorId, - pub input_shape: Vec, - pub dim: usize, - pub indices: Vec, -} - -impl GradientFunction for IndexSelectBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let grad_input = - index_select_backward_grad(grad_output, &self.input_shape, self.dim, &self.indices)?; - let mut gradients = FxHashMap::default(); - accumulate_grad(&mut gradients, self.input_id, grad_input)?; - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for `gather` (and, reused, min/max/sort/topk along a dim). -pub struct GatherBackward { - pub input_id: TensorId, - pub input_shape: Vec, - pub dim: usize, - pub index: Vec, -} - -impl GradientFunction for GatherBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let grad_input = - gather_backward_grad(grad_output, &self.input_shape, self.dim, &self.index)?; - let mut gradients = FxHashMap::default(); - accumulate_grad(&mut gradients, self.input_id, grad_input)?; - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for `concatenate` (and, transitively, `cat`/`stack`/`roll`). -pub struct ConcatBackward { - pub input_ids: SmallVec<[TensorId; 4]>, - pub sizes: SmallVec<[usize; 4]>, - pub dim: usize, -} - -impl GradientFunction for ConcatBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - let mut offset = 0usize; - for (&id, &size) in self.input_ids.iter().zip(self.sizes.iter()) { - let grad_slice = crate::operations::shape_ops::narrow( - grad_output, - self.dim as isize, - offset, - size, - )?; - accumulate_grad(&mut gradients, id, grad_slice)?; - offset += size; - } - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - &self.input_ids - } -} - -/// Gradient function for `roll`: rolling is a bijection, so the gradient is the -/// input rolled back by the negated shifts. Computed with a dedicated node rather -/// than by composing `slice`/`concatenate`, because `roll`'s flatten path builds -/// a storage-sharing view whose gradient edges cannot be composed safely. -pub struct RollBackward { - pub input_id: TensorId, - pub shifts: Vec, - pub dims: Option>, -} - -impl GradientFunction for RollBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let neg: Vec = self.shifts.iter().map(|s| -s).collect(); - let grad_input = - crate::operations::shape_ops::roll(grad_output, &neg, self.dims.as_deref())?; - let mut gradients = FxHashMap::default(); - accumulate_grad(&mut gradients, self.input_id, grad_input)?; - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for `repeat` (tiling): sum the gradient over the tiled copies. -pub struct RepeatBackward { - pub input_id: TensorId, - pub input_shape: Vec, - pub repeats: Vec, -} - -impl GradientFunction for RepeatBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - // `repeat` may prepend leading singleton axes; align the input rank to the - // repeat/output rank, tile every axis, then sum the tiled copies back down. - let out_ndim = self.repeats.len(); - let pad = out_ndim - self.input_shape.len(); - let mut aligned = vec![1usize; pad]; - aligned.extend_from_slice(&self.input_shape); - - // View grad_output as (rep_0, in_0, rep_1, in_1, ...) then sum the rep axes. - let mut split_shape = Vec::with_capacity(2 * out_ndim); - for axis in 0..out_ndim { - split_shape.push(self.repeats[axis]); - split_shape.push(aligned[axis]); - } - let reshaped = - crate::operations::shape_ops::reshape(grad_output, Shape::new(split_shape))?; - let rep_axes: Vec = (0..out_ndim).map(|axis| (2 * axis) as isize).collect(); - let summed = reduction::sum(&reshaped, Some(rep_axes), false)?; - let grad_input = - crate::operations::shape_ops::reshape(&summed, Shape::new(self.input_shape.clone()))?; - - let mut gradients = FxHashMap::default(); - accumulate_grad(&mut gradients, self.input_id, grad_input)?; - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for basic indexing (`tensor[...]` via [`Tensor::index`]). -/// -/// The forward gathers input element `offset + Σ_j (start_j + coord_j·step_j)· -/// input_stride_{dim_j}` for each output coordinate; the backward scatters the -/// gradient straight back to those positions. Assumes contiguous input storage, -/// which always holds at the Python boundary where indexing is applied. -pub struct IndexBackward { - pub input_id: TensorId, - pub input_shape: Vec, - pub input_strides: Vec, - pub offset: usize, - pub out_dims: Vec, - pub orig_dim_map: Vec, - pub starts: Vec, - pub steps: Vec, -} - -impl GradientFunction for IndexBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let numel: usize = self.input_shape.iter().product(); - let mut grad_data = - TensorData::zeros_on_device(numel, grad_output.dtype(), grad_output.device()); - let out_strides = Strides::from_shape(&Shape::new(self.out_dims.clone())); - let out_strides = out_strides.as_slice(); - - macro_rules! scatter { - ($slice:ident, $mut_slice:ident) => {{ - let go = grad_output.data().$slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to read grad_output for index backward") - })?; - let gi = grad_data.$mut_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to write grad for index backward") - })?; - if self.out_dims.is_empty() { - // Scalar result: a single collapsed element. - gi[self.offset] += go[0]; - } else { - for (idx, &g) in go.iter().enumerate() { - let mut rem = idx; - let mut src = self.offset; - for (j, &ostride) in out_strides.iter().enumerate() { - let coord = rem / ostride; - rem %= ostride; - src += (self.starts[j] + coord * self.steps[j]) - * self.input_strides[self.orig_dim_map[j]]; - } - gi[src] += g; - } - } - }}; - } - - match grad_output.dtype() { - DataType::Float32 => scatter!(as_f32_slice, as_f32_slice_mut), - DataType::Float64 => scatter!(as_f64_slice, as_f64_slice_mut), - _ => { - return Err(MinitensorError::invalid_operation( - "index backward only supported for floating point tensors", - )); - } - } - - let grad_input = Tensor::new( - Arc::new(grad_data), - Shape::new(self.input_shape.clone()), - grad_output.dtype(), - grad_output.device(), - false, - ); - let mut gradients = FxHashMap::default(); - accumulate_grad(&mut gradients, self.input_id, grad_input)?; - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for softmax -pub struct SoftmaxBackward { - pub input_id: TensorId, - pub output: Tensor, - pub dim: usize, -} - -impl GradientFunction for SoftmaxBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - // Allocate gradient buffer - let mut grad_data = TensorData::zeros_on_device( - self.output.numel(), - self.output.dtype(), - self.output.device(), - ); - - match grad_output.dtype() { - DataType::Float32 => { - let go = grad_output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from grad_output") - })?; - let y = self.output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from softmax output") - })?; - let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from grad_data", - ) - })?; - softmax_backward_f32(go, y, grad_slice, self.output.shape().dims(), self.dim); - } - DataType::Float64 => { - let go = grad_output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from grad_output") - })?; - let y = self.output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from softmax output") - })?; - let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from grad_data", - ) - })?; - softmax_backward_f64(go, y, grad_slice, self.output.shape().dims(), self.dim); - } - _ => { - return Err(MinitensorError::invalid_operation( - "Softmax backward only supported for floating point tensors", - )); - } - } - - let grad_input = Tensor::new( - Arc::new(grad_data), - self.output.shape().clone(), - self.output.dtype(), - self.output.device(), - grad_output.requires_grad(), - ); - - gradients.insert(self.input_id, grad_input); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for log-softmax -pub struct LogSoftmaxBackward { - pub input_id: TensorId, - pub output: Tensor, - pub dim: usize, -} - -impl GradientFunction for LogSoftmaxBackward { - fn backward(&self, grad_output: &Tensor) -> Result> { - let mut gradients = FxHashMap::default(); - gradients.reserve(1); - - let mut grad_data = TensorData::zeros_on_device( - self.output.numel(), - self.output.dtype(), - self.output.device(), - ); - - match grad_output.dtype() { - DataType::Float32 => { - let go = grad_output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from grad_output") - })?; - let log_y = self.output.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get f32 slice from log_softmax output", - ) - })?; - let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from grad_data", - ) - })?; - log_softmax_backward_f32( - go, - log_y, - grad_slice, - self.output.shape().dims(), - self.dim, - ); - } - DataType::Float64 => { - let go = grad_output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from grad_output") - })?; - let log_y = self.output.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get f64 slice from log_softmax output", - ) - })?; - let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from grad_data", - ) - })?; - log_softmax_backward_f64( - go, - log_y, - grad_slice, - self.output.shape().dims(), - self.dim, - ); - } - _ => { - return Err(MinitensorError::invalid_operation( - "LogSoftmax backward only supported for floating point tensors", - )); - } - } - - let grad_input = Tensor::new( - Arc::new(grad_data), - self.output.shape().clone(), - self.output.dtype(), - self.output.device(), - grad_output.requires_grad(), - ); - - gradients.insert(self.input_id, grad_input); - - Ok(gradients) - } - - fn input_ids(&self) -> &[TensorId] { - std::slice::from_ref(&self.input_id) - } -} - -/// Gradient function for masked log-softmax -pub struct MaskedLogSoftmaxBackward { - pub input_id: TensorId, - pub output: Tensor, - pub mask: Tensor, - pub dim: usize, -} +/// bound may be absent (`clamp_min`/`clamp_max`). +pub struct ClampBackward { + pub input_id: TensorId, + pub input: Tensor, + pub min: Option, + pub max: Option, +} + +impl GradientFunction for ClampBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + let mut grad_data = TensorData::zeros_on_device( + grad_output.numel(), + grad_output.dtype(), + grad_output.device(), + ); + + macro_rules! clamp_grad { + ($slice:ident, $mut_slice:ident, $ty:ty) => {{ + let x = self.input.data().$slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to read input for clamp backward") + })?; + let go = grad_output.data().$slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to read grad_output for clamp backward") + })?; + let gi = grad_data.$mut_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to write grad for clamp backward") + })?; + let min = self.min.map(|m| m as $ty); + let max = self.max.map(|m| m as $ty); + let passes = move |v: $ty| -> bool { + min.map_or(true, |m| v >= m) && max.map_or(true, |m| v <= m) + }; + if gi.len() < PAR_THRESHOLD { + for i in 0..gi.len() { + gi[i] = if passes(x[i]) { go[i] } else { 0.0 }; + } + } else { + gi.par_iter_mut() + .zip(go.par_iter()) + .zip(x.par_iter()) + .for_each(|((g, &o), &v)| *g = if passes(v) { o } else { 0.0 }); + } + }}; + } + + match grad_output.dtype() { + DataType::Float32 => clamp_grad!(as_f32_slice, as_f32_slice_mut, f32), + DataType::Float64 => clamp_grad!(as_f64_slice, as_f64_slice_mut, f64), + _ => { + return Err(MinitensorError::invalid_operation( + "clamp backward only supported for floating point tensors", + )); + } + } + + let grad_input = Tensor::new( + Arc::new(grad_data), + grad_output.shape().clone(), + grad_output.dtype(), + grad_output.device(), + false, + ); + gradients.insert(self.input_id, grad_input); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Scatter-add `grad_output` back to the source positions selected along `dim`. +/// +/// `indices[i]` is the source position (along `dim`) that produced output row `i` +/// for every outer/inner coordinate. This is the shared backward for +/// `index_select` and `slice` (and, transitively, `narrow`/`flip`/`roll`). +/// Duplicated source indices accumulate, matching the forward gather semantics. +fn index_select_backward_grad( + grad_output: &Tensor, + input_shape: &[usize], + dim: usize, + indices: &[usize], +) -> Result { + let numel: usize = input_shape.iter().product(); + let mut grad_data = + TensorData::zeros_on_device(numel, grad_output.dtype(), grad_output.device()); + + let dim_size = input_shape[dim]; + let inner: usize = input_shape[dim + 1..].iter().product(); + let out_dim = indices.len(); + + if numel != 0 && out_dim != 0 && inner != 0 { + let in_chunk = dim_size * inner; + let out_chunk = out_dim * inner; + + macro_rules! fill { + ($slice:ident, $mut_slice:ident) => {{ + let go = grad_output.data().$slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to read grad_output for index backward") + })?; + let gi = grad_data.$mut_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to write grad for index backward") + })?; + gi.par_chunks_mut(in_chunk) + .enumerate() + .for_each(|(o, gi_chunk)| { + let go_chunk = &go[o * out_chunk..(o + 1) * out_chunk]; + for (i, &idx) in indices.iter().enumerate() { + let dst = idx * inner; + let src = i * inner; + for j in 0..inner { + gi_chunk[dst + j] += go_chunk[src + j]; + } + } + }); + }}; + } + + match grad_output.dtype() { + DataType::Float32 => fill!(as_f32_slice, as_f32_slice_mut), + DataType::Float64 => fill!(as_f64_slice, as_f64_slice_mut), + _ => { + return Err(MinitensorError::invalid_operation( + "index/slice backward only supported for floating point tensors", + )); + } + } + } + + Ok(Tensor::new( + Arc::new(grad_data), + Shape::new(input_shape.to_vec()), + grad_output.dtype(), + grad_output.device(), + false, + )) +} + +/// Scatter-add `grad_output` back to the input positions named by a full `index` +/// tensor (`gather` backward, also reused by min/max/sort/topk along a dim). The +/// `index` slice is laid out identically to `grad_output`; entry `index[..]` is +/// the source coordinate along `dim`. Colliding indices accumulate. +fn gather_backward_grad( + grad_output: &Tensor, + input_shape: &[usize], + dim: usize, + index: &[i64], +) -> Result { + let numel: usize = input_shape.iter().product(); + let mut grad_data = + TensorData::zeros_on_device(numel, grad_output.dtype(), grad_output.device()); + + let dim_size = input_shape[dim]; + let inner: usize = input_shape[dim + 1..].iter().product(); + // The index tensor shares `grad_output`'s shape, so the output extent along + // `dim` is read directly from it. + let out_dim = grad_output.shape().dims()[dim]; + + if numel != 0 && !index.is_empty() && inner != 0 { + let in_chunk = dim_size * inner; + let out_chunk = out_dim * inner; + + macro_rules! fill { + ($slice:ident, $mut_slice:ident) => {{ + let go = grad_output.data().$slice().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to read grad_output for gather backward", + ) + })?; + let gi = grad_data.$mut_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to write grad for gather backward") + })?; + gi.par_chunks_mut(in_chunk) + .enumerate() + .for_each(|(o, gi_chunk)| { + let go_chunk = &go[o * out_chunk..(o + 1) * out_chunk]; + let idx_chunk = &index[o * out_chunk..(o + 1) * out_chunk]; + for i in 0..out_dim { + for j in 0..inner { + let pos = i * inner + j; + let src_idx = idx_chunk[pos] as usize; + gi_chunk[src_idx * inner + j] += go_chunk[pos]; + } + } + }); + }}; + } + + match grad_output.dtype() { + DataType::Float32 => fill!(as_f32_slice, as_f32_slice_mut), + DataType::Float64 => fill!(as_f64_slice, as_f64_slice_mut), + _ => { + return Err(MinitensorError::invalid_operation( + "gather backward only supported for floating point tensors", + )); + } + } + } + + Ok(Tensor::new( + Arc::new(grad_data), + Shape::new(input_shape.to_vec()), + grad_output.dtype(), + grad_output.device(), + false, + )) +} + +/// Gradient function for `index_select` and `slice` (source indices along `dim`). +pub struct IndexSelectBackward { + pub input_id: TensorId, + pub input_shape: Vec, + pub dim: usize, + pub indices: Vec, +} + +impl GradientFunction for IndexSelectBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let grad_input = + index_select_backward_grad(grad_output, &self.input_shape, self.dim, &self.indices)?; + let mut gradients = FxHashMap::default(); + accumulate_grad(&mut gradients, self.input_id, grad_input)?; + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for `gather` (and, reused, min/max/sort/topk along a dim). +pub struct GatherBackward { + pub input_id: TensorId, + pub input_shape: Vec, + pub dim: usize, + pub index: Vec, +} + +impl GradientFunction for GatherBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let grad_input = + gather_backward_grad(grad_output, &self.input_shape, self.dim, &self.index)?; + let mut gradients = FxHashMap::default(); + accumulate_grad(&mut gradients, self.input_id, grad_input)?; + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for `concatenate` (and, transitively, `cat`/`stack`/`roll`). +pub struct ConcatBackward { + pub input_ids: SmallVec<[TensorId; 4]>, + pub sizes: SmallVec<[usize; 4]>, + pub dim: usize, + /// Which inputs actually need a gradient; frozen inputs skip their + /// slice extraction. + pub input_requires_grad: SmallVec<[bool; 4]>, +} + +impl GradientFunction for ConcatBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + let mut offset = 0usize; + for ((&id, &size), &needs_grad) in self + .input_ids + .iter() + .zip(self.sizes.iter()) + .zip(self.input_requires_grad.iter()) + { + if needs_grad { + let grad_slice = crate::operations::shape_ops::narrow( + grad_output, + self.dim as isize, + offset, + size, + )?; + accumulate_grad(&mut gradients, id, grad_slice)?; + } + offset += size; + } + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + &self.input_ids + } +} + +/// Gradient function for `roll`: rolling is a bijection, so the gradient is the +/// input rolled back by the negated shifts. Computed with a dedicated node rather +/// than by composing `slice`/`concatenate`, because `roll`'s flatten path builds +/// a storage-sharing view whose gradient edges cannot be composed safely. +pub struct RollBackward { + pub input_id: TensorId, + pub shifts: Vec, + pub dims: Option>, +} + +impl GradientFunction for RollBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let neg: Vec = self.shifts.iter().map(|s| -s).collect(); + let grad_input = + crate::operations::shape_ops::roll(grad_output, &neg, self.dims.as_deref())?; + let mut gradients = FxHashMap::default(); + accumulate_grad(&mut gradients, self.input_id, grad_input)?; + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for `repeat` (tiling): sum the gradient over the tiled copies. +pub struct RepeatBackward { + pub input_id: TensorId, + pub input_shape: Vec, + pub repeats: Vec, +} + +impl GradientFunction for RepeatBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + // `repeat` may prepend leading singleton axes; align the input rank to the + // repeat/output rank, tile every axis, then sum the tiled copies back down. + let out_ndim = self.repeats.len(); + let pad = out_ndim - self.input_shape.len(); + let mut aligned = vec![1usize; pad]; + aligned.extend_from_slice(&self.input_shape); + + // View grad_output as (rep_0, in_0, rep_1, in_1, ...) then sum the rep axes. + let mut split_shape = Vec::with_capacity(2 * out_ndim); + for axis in 0..out_ndim { + split_shape.push(self.repeats[axis]); + split_shape.push(aligned[axis]); + } + let reshaped = crate::operations::shape_ops::reshape(grad_output, Shape::new(split_shape))?; + let rep_axes: Vec = (0..out_ndim).map(|axis| (2 * axis) as isize).collect(); + let summed = reduction::sum(&reshaped, Some(rep_axes), false)?; + let grad_input = + crate::operations::shape_ops::reshape(&summed, Shape::new(self.input_shape.clone()))?; + + let mut gradients = FxHashMap::default(); + accumulate_grad(&mut gradients, self.input_id, grad_input)?; + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for basic indexing (`tensor[...]` via [`Tensor::index`]). +/// +/// The forward gathers input element `offset + Σ_j (start_j + coord_j·step_j)· +/// input_stride_{dim_j}` for each output coordinate; the backward scatters the +/// gradient straight back to those positions. Assumes contiguous input storage, +/// which always holds at the Python boundary where indexing is applied. +pub struct IndexBackward { + pub input_id: TensorId, + pub input_shape: Vec, + pub input_strides: Vec, + pub offset: usize, + pub out_dims: Vec, + pub orig_dim_map: Vec, + pub starts: Vec, + pub steps: Vec, +} + +impl GradientFunction for IndexBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let numel: usize = self.input_shape.iter().product(); + let mut grad_data = + TensorData::zeros_on_device(numel, grad_output.dtype(), grad_output.device()); + let out_strides = Strides::from_shape(&Shape::new(self.out_dims.clone())); + let out_strides = out_strides.as_slice(); + + macro_rules! scatter { + ($slice:ident, $mut_slice:ident) => {{ + let go = grad_output.data().$slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to read grad_output for index backward") + })?; + let gi = grad_data.$mut_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to write grad for index backward") + })?; + if self.out_dims.is_empty() { + // Scalar result: a single collapsed element. + gi[self.offset] += go[0]; + } else { + for (idx, &g) in go.iter().enumerate() { + let mut rem = idx; + let mut src = self.offset; + for (j, &ostride) in out_strides.iter().enumerate() { + let coord = rem / ostride; + rem %= ostride; + src += (self.starts[j] + coord * self.steps[j]) + * self.input_strides[self.orig_dim_map[j]]; + } + gi[src] += g; + } + } + }}; + } + + match grad_output.dtype() { + DataType::Float32 => scatter!(as_f32_slice, as_f32_slice_mut), + DataType::Float64 => scatter!(as_f64_slice, as_f64_slice_mut), + _ => { + return Err(MinitensorError::invalid_operation( + "index backward only supported for floating point tensors", + )); + } + } + + let grad_input = Tensor::new( + Arc::new(grad_data), + Shape::new(self.input_shape.clone()), + grad_output.dtype(), + grad_output.device(), + false, + ); + let mut gradients = FxHashMap::default(); + accumulate_grad(&mut gradients, self.input_id, grad_input)?; + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for softmax +pub struct SoftmaxBackward { + pub input_id: TensorId, + pub output: Tensor, + pub dim: usize, +} + +impl GradientFunction for SoftmaxBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + // Allocate gradient buffer + let mut grad_data = TensorData::zeros_on_device( + self.output.numel(), + self.output.dtype(), + self.output.device(), + ); + + match grad_output.dtype() { + DataType::Float32 => { + let go = grad_output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from grad_output") + })?; + let y = self.output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from softmax output") + })?; + let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from grad_data", + ) + })?; + softmax_backward_f32(go, y, grad_slice, self.output.shape().dims(), self.dim); + } + DataType::Float64 => { + let go = grad_output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from grad_output") + })?; + let y = self.output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from softmax output") + })?; + let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from grad_data", + ) + })?; + softmax_backward_f64(go, y, grad_slice, self.output.shape().dims(), self.dim); + } + _ => { + return Err(MinitensorError::invalid_operation( + "Softmax backward only supported for floating point tensors", + )); + } + } + + let grad_input = Tensor::new( + Arc::new(grad_data), + self.output.shape().clone(), + self.output.dtype(), + self.output.device(), + grad_output.requires_grad(), + ); + + gradients.insert(self.input_id, grad_input); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for log-softmax +pub struct LogSoftmaxBackward { + pub input_id: TensorId, + pub output: Tensor, + pub dim: usize, +} + +impl GradientFunction for LogSoftmaxBackward { + fn backward(&self, grad_output: &Tensor) -> Result> { + let mut gradients = FxHashMap::default(); + gradients.reserve(1); + + let mut grad_data = TensorData::zeros_on_device( + self.output.numel(), + self.output.dtype(), + self.output.device(), + ); + + match grad_output.dtype() { + DataType::Float32 => { + let go = grad_output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from grad_output") + })?; + let log_y = self.output.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get f32 slice from log_softmax output", + ) + })?; + let grad_slice = grad_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from grad_data", + ) + })?; + log_softmax_backward_f32( + go, + log_y, + grad_slice, + self.output.shape().dims(), + self.dim, + ); + } + DataType::Float64 => { + let go = grad_output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from grad_output") + })?; + let log_y = self.output.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get f64 slice from log_softmax output", + ) + })?; + let grad_slice = grad_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from grad_data", + ) + })?; + log_softmax_backward_f64( + go, + log_y, + grad_slice, + self.output.shape().dims(), + self.dim, + ); + } + _ => { + return Err(MinitensorError::invalid_operation( + "LogSoftmax backward only supported for floating point tensors", + )); + } + } + + let grad_input = Tensor::new( + Arc::new(grad_data), + self.output.shape().clone(), + self.output.dtype(), + self.output.device(), + grad_output.requires_grad(), + ); + + gradients.insert(self.input_id, grad_input); + + Ok(gradients) + } + + fn input_ids(&self) -> &[TensorId] { + std::slice::from_ref(&self.input_id) + } +} + +/// Gradient function for masked log-softmax +pub struct MaskedLogSoftmaxBackward { + pub input_id: TensorId, + pub output: Tensor, + pub mask: Tensor, + pub dim: usize, +} diff --git a/engine/src/autograd/mod/tests.rs b/engine/src/autograd/mod/tests.rs index ce49edb0..0dbeab54 100644 --- a/engine/src/autograd/mod/tests.rs +++ b/engine/src/autograd/mod/tests.rs @@ -1,538 +1,652 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -/// Helper function to reduce gradients for broadcasting -fn reduce_gradient_for_broadcasting(grad_output: &Tensor, target_shape: &Shape) -> Result { - if grad_output.shape() == target_shape { - return Ok(grad_output.clone()); - } - - let grad_dims = grad_output.shape().dims(); - let target_dims = target_shape.dims(); - if target_dims.len() > grad_dims.len() { - return Err(MinitensorError::BroadcastError { - shape1: grad_dims.to_vec(), - shape2: target_dims.to_vec(), - suggestion: Some( - "Ensure the target shape has no more dimensions than the gradient output." - .to_string(), - ), - context: Some("reduce_gradient_for_broadcasting".to_string()), - }); - } - let extra = grad_dims.len() - target_dims.len(); - - // Use a stack-allocated small vector and pre-allocate enough capacity to - // hold all potential broadcast axes. This avoids repeated reallocations for - // higher dimensional tensors. - let mut axes_to_sum: SmallVec<[usize; 8]> = SmallVec::with_capacity(grad_dims.len()); - axes_to_sum.extend(0..extra); - for i in 0..target_dims.len() { - let gdim = grad_dims[extra + i]; - let tdim = target_dims[i]; - if tdim == 1 { - if gdim != 1 { - axes_to_sum.push(extra + i); - } - } else if gdim != tdim { - return Err(MinitensorError::BroadcastError { - shape1: grad_dims.to_vec(), - shape2: target_dims.to_vec(), - suggestion: Some( - "Ensure each target dimension is 1 or matches the gradient dimension." - .to_string(), - ), - context: Some("reduce_gradient_for_broadcasting".to_string()), - }); - } - } - - if axes_to_sum.is_empty() { - return Ok(grad_output.clone()); - } - - let mut axes = Vec::with_capacity(axes_to_sum.len()); - for axis in axes_to_sum { - axes.push(axis as isize); - } - let mut grad = reduction::sum(grad_output, Some(axes), true)?; - - if grad.shape() != target_shape { - grad = grad.view(target_shape.clone())?; - } - - Ok(grad) -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::device::Device; - use crate::tensor::DataType; - - #[test] - fn test_tensor_id_generation() { - let id1 = TensorId::new(); - let id2 = TensorId::new(); - assert_ne!(id1, id2); - } - - #[test] - fn test_computation_graph() { - let mut graph = ComputationGraph::new(); - let tensor_id = TensorId::new(); - - let grad_fn = Arc::new(AddBackward { - input_shapes: [vec![2, 2], vec![2, 2]], - input_ids: [TensorId::new(), TensorId::new()], - }); - - graph.add_tensor_with_grad_req(tensor_id, Some(grad_fn), true); - assert!(graph.nodes().contains_key(&tensor_id)); - } - - #[test] - fn test_add_backward() { - let grad_fn = AddBackward { - input_shapes: [vec![2, 2], vec![2, 2]], - input_ids: [TensorId::new(), TensorId::new()], - }; - - let grad_output = Tensor::ones( - Shape::new(vec![2, 2]), - crate::tensor::DataType::Float32, - Device::cpu(), - false, - ); - let gradients = grad_fn.backward(&grad_output).unwrap(); - - assert_eq!(gradients.len(), 2); - } - - #[test] - fn test_add_backward_same_input_accumulates() { - // `x + x`: both operands are the same tensor. The two gradient - // contributions must be summed (2), not overwritten (1). - let shared = TensorId::new(); - let grad_fn = AddBackward { - input_shapes: [vec![3], vec![3]], - input_ids: [shared, shared], - }; - - let grad_output = Tensor::ones(Shape::new(vec![3]), DataType::Float32, Device::cpu(), false); - let gradients = grad_fn.backward(&grad_output).unwrap(); - - assert_eq!(gradients.len(), 1); - let grad = gradients.get(&shared).unwrap(); - for &v in grad.data().as_f32_slice().unwrap() { - assert_eq!(v, 2.0); - } - } - - #[test] - fn test_mul_backward_same_input_accumulates() { - // `x * x`: d/dx(x^2) = 2x. With shared operands the backward emits the - // same id twice (grad*rhs and grad*lhs); both must accumulate. - let shared = TensorId::new(); - let x = Tensor::new( - Arc::new(TensorData::from_vec_f32(vec![1.0, 2.0, 3.0], Device::cpu())), - Shape::new(vec![3]), - DataType::Float32, - Device::cpu(), - false, - ); - let grad_fn = MulBackward { - lhs: x.clone(), - rhs: x.clone(), - input_ids: [shared, shared], - }; - - let grad_output = Tensor::ones(Shape::new(vec![3]), DataType::Float32, Device::cpu(), false); - let gradients = grad_fn.backward(&grad_output).unwrap(); - - assert_eq!(gradients.len(), 1); - let grad = gradients.get(&shared).unwrap(); - let expected = [2.0f32, 4.0, 6.0]; - for (&got, &want) in grad - .data() - .as_f32_slice() - .unwrap() - .iter() - .zip(expected.iter()) - { - assert_eq!(got, want); - } - } - - #[test] - fn test_reduce_gradient_for_broadcasting() { - let grad_output = Tensor::ones( - Shape::new(vec![2, 3]), - DataType::Float32, - Device::cpu(), - false, - ); - let target_shape = Shape::new(vec![2, 1]); - let reduced = reduce_gradient_for_broadcasting(&grad_output, &target_shape).unwrap(); - assert_eq!(reduced.shape().dims(), &[2, 1]); - let slice = reduced.data().as_f32_slice().unwrap(); - assert!(slice.iter().all(|&x| (x - 3.0).abs() < 1e-6)); - } - - #[test] - fn test_reduce_gradient_multiple_axes() { - let grad_output = Tensor::ones( - Shape::new(vec![2, 3]), - DataType::Float32, - Device::cpu(), - false, - ); - let target_shape = Shape::new(vec![1, 1]); - let reduced = reduce_gradient_for_broadcasting(&grad_output, &target_shape).unwrap(); - assert_eq!(reduced.shape().dims(), &[1, 1]); - let slice = reduced.data().as_f32_slice().unwrap(); - assert!((slice[0] - 6.0).abs() < 1e-6); - } - - #[test] - fn test_reduce_gradient_with_leading_and_inner_axes() { - let grad_output = Tensor::ones( - Shape::new(vec![4, 2, 3, 5]), - DataType::Float32, - Device::cpu(), - false, - ); - let target_shape = Shape::new(vec![2, 1, 5]); - let reduced = reduce_gradient_for_broadcasting(&grad_output, &target_shape).unwrap(); - assert_eq!(reduced.shape().dims(), &[2, 1, 5]); - let slice = reduced.data().as_f32_slice().unwrap(); - assert!(slice.iter().all(|&x| (x - 12.0).abs() < 1e-6)); - } - - #[test] - fn test_reduce_gradient_noop_for_same_shape() { - let grad_output = Tensor::ones( - Shape::new(vec![2, 3, 4]), - DataType::Float32, - Device::cpu(), - false, - ); - let target_shape = grad_output.shape().clone(); - let reduced = reduce_gradient_for_broadcasting(&grad_output, &target_shape).unwrap(); - assert_eq!(reduced.shape().dims(), &[2, 3, 4]); - assert!(reduced.allclose(&grad_output, 1e-6, 1e-6)); - } - - #[test] - fn test_reduce_gradient_invalid_broadcast() { - let grad_output = Tensor::ones( - Shape::new(vec![2, 1]), - DataType::Float32, - Device::cpu(), - false, - ); - let target_shape = Shape::new(vec![2, 2]); - let err = reduce_gradient_for_broadcasting(&grad_output, &target_shape) - .expect_err("expected invalid broadcast error"); - assert!(matches!(err, MinitensorError::BroadcastError { .. })); - } - - #[test] - fn test_reduce_gradient_zero_dim_broadcast() { - let grad_output = Tensor::ones( - Shape::new(vec![0, 2]), - DataType::Float32, - Device::cpu(), - false, - ); - let target_shape = Shape::new(vec![1, 2]); - let reduced = reduce_gradient_for_broadcasting(&grad_output, &target_shape).unwrap(); - assert_eq!(reduced.shape().dims(), &[1, 2]); - let slice = reduced.data().as_f32_slice().unwrap(); - assert!(slice.iter().all(|&x| x == 0.0)); - } - - #[test] - fn test_softmax_backward_dim1() { - let input = Tensor::new( - Arc::new(TensorData::from_vec_f32( - vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], - Device::cpu(), - )), - Shape::new(vec![2, 3]), - DataType::Float32, - Device::cpu(), - false, - ); - let softmax_out = activation::softmax(&input, Some(1)).unwrap(); - let grad_output = Tensor::ones( - Shape::new(vec![2, 3]), - DataType::Float32, - Device::cpu(), - false, - ); - let grad_y = arithmetic::mul(&grad_output, &softmax_out).unwrap(); - let sum = reduction::sum(&grad_y, Some(vec![1]), true).unwrap(); - let sub = arithmetic::sub(&grad_output, &sum).unwrap(); - let expected = arithmetic::mul(&softmax_out, &sub).unwrap(); - - let grad_fn = SoftmaxBackward { - input_id: TensorId::new(), - output: softmax_out.clone(), - dim: 1, - }; - let grads = grad_fn.backward(&grad_output).unwrap(); - let grad_input = grads.values().next().unwrap(); - assert!(grad_input.allclose(&expected, 1e-6, 1e-6)); - } - - #[test] - fn test_softmax_backward_dim0_f64() { - let data: Vec = (1..=6).map(|v| v as f64).collect(); - let input = Tensor::new( - Arc::new(TensorData::from_vec_f64(data, Device::cpu())), - Shape::new(vec![2, 3]), - DataType::Float64, - Device::cpu(), - false, - ); - let softmax_out = activation::softmax(&input, Some(0)).unwrap(); - let grad_output = Tensor::ones( - Shape::new(vec![2, 3]), - DataType::Float64, - Device::cpu(), - false, - ); - let grad_y = arithmetic::mul(&grad_output, &softmax_out).unwrap(); - let sum = reduction::sum(&grad_y, Some(vec![0]), true).unwrap(); - let sub = arithmetic::sub(&grad_output, &sum).unwrap(); - let expected = arithmetic::mul(&softmax_out, &sub).unwrap(); - - let grad_fn = SoftmaxBackward { - input_id: TensorId::new(), - output: softmax_out.clone(), - dim: 0, - }; - let grads = grad_fn.backward(&grad_output).unwrap(); - let grad_input = grads.values().next().unwrap(); - assert!(grad_input.allclose(&expected, 1e-6, 1e-6)); - } - - #[test] - fn test_backward_broadcast_addition() { - clear_graph().unwrap(); - - let a = Tensor::ones( - Shape::new(vec![2, 3]), - DataType::Float32, - Device::cpu(), - true, - ); - let b = Tensor::ones(Shape::new(vec![3]), DataType::Float32, Device::cpu(), true); - let out = arithmetic::add(&a, &b).unwrap(); - - let grad = Tensor::ones(out.shape().clone(), out.dtype(), out.device(), false); - let grads = backward(&out, Some(grad)).unwrap(); - - let grad_a = grads.get(&a.id()).unwrap(); - let grad_b = grads.get(&b.id()).unwrap(); - assert_eq!(grad_a.shape().dims(), &[2, 3]); - assert_eq!(grad_b.shape().dims(), &[3]); - let slice_b = grad_b.data().as_f32_slice().unwrap(); - assert!(slice_b.iter().all(|&x| (x - 2.0).abs() < 1e-6)); - } - - #[test] - fn test_matmul_backward_gradients() { - let lhs = Tensor::new( - Arc::new(TensorData::from_vec_f32( - vec![1.0, 2.0, 3.0, 4.0], - Device::cpu(), - )), - Shape::new(vec![2, 2]), - DataType::Float32, - Device::cpu(), - false, - ); - let rhs = Tensor::new( - Arc::new(TensorData::from_vec_f32( - vec![5.0, 6.0, 7.0, 8.0], - Device::cpu(), - )), - Shape::new(vec![2, 2]), - DataType::Float32, - Device::cpu(), - false, - ); - let input_ids = [TensorId::new(), TensorId::new()]; - let grad_fn = MatMulBackward { - lhs: lhs.clone(), - rhs: rhs.clone(), - input_ids, - lhs_requires_grad: true, - rhs_requires_grad: true, - }; - let grad_output = Tensor::ones( - Shape::new(vec![2, 2]), - DataType::Float32, - Device::cpu(), - false, - ); - let grads = grad_fn.backward(&grad_output).unwrap(); - let rhs_t = crate::operations::linalg::transpose(&rhs, 0, 1).unwrap(); - let expected_lhs = crate::operations::linalg::matmul(&grad_output, &rhs_t).unwrap(); - let lhs_grad = grads.get(&input_ids[0]).unwrap(); - assert!(lhs_grad.allclose(&expected_lhs, 1e-6, 1e-6)); - } - - #[test] - fn test_matmul_backward_batched() { - let lhs = Tensor::new( - Arc::new(TensorData::from_vec_f32( - (0..12).map(|x| x as f32).collect(), - Device::cpu(), - )), - Shape::new(vec![2, 2, 3]), - DataType::Float32, - Device::cpu(), - false, - ); - let rhs = Tensor::new( - Arc::new(TensorData::from_vec_f32( - (0..24).map(|x| x as f32).collect(), - Device::cpu(), - )), - Shape::new(vec![2, 3, 4]), - DataType::Float32, - Device::cpu(), - false, - ); - let input_ids = [TensorId::new(), TensorId::new()]; - let grad_fn = MatMulBackward { - lhs: lhs.clone(), - rhs: rhs.clone(), - input_ids, - lhs_requires_grad: true, - rhs_requires_grad: true, - }; - let grad_output = Tensor::ones( - Shape::new(vec![2, 2, 4]), - DataType::Float32, - Device::cpu(), - false, - ); - let grads = grad_fn.backward(&grad_output).unwrap(); - let rhs_t = crate::operations::linalg::transpose( - &rhs, - (rhs.ndim() - 2) as isize, - (rhs.ndim() - 1) as isize, - ) - .unwrap(); - let expected_lhs = crate::operations::linalg::matmul(&grad_output, &rhs_t).unwrap(); - assert!( - grads - .get(&input_ids[0]) - .unwrap() - .allclose(&expected_lhs, 1e-6, 1e-6) - ); - let lhs_t = crate::operations::linalg::transpose( - &lhs, - (lhs.ndim() - 2) as isize, - (lhs.ndim() - 1) as isize, - ) - .unwrap(); - let expected_rhs = crate::operations::linalg::matmul(&lhs_t, &grad_output).unwrap(); - assert!( - grads - .get(&input_ids[1]) - .unwrap() - .allclose(&expected_rhs, 1e-6, 1e-6) - ); - } - - #[test] - fn test_matmul_backward_requires_grad_flags() { - let lhs = Tensor::new( - Arc::new(TensorData::from_vec_f32( - vec![1.0, 2.0, 3.0, 4.0], - Device::cpu(), - )), - Shape::new(vec![2, 2]), - DataType::Float32, - Device::cpu(), - false, - ); - let rhs = Tensor::new( - Arc::new(TensorData::from_vec_f32( - vec![5.0, 6.0, 7.0, 8.0], - Device::cpu(), - )), - Shape::new(vec![2, 2]), - DataType::Float32, - Device::cpu(), - false, - ); - let ids = [TensorId::new(), TensorId::new()]; - let grad_fn = MatMulBackward { - lhs: lhs.clone(), - rhs: rhs.clone(), - input_ids: ids, - lhs_requires_grad: true, - rhs_requires_grad: false, - }; - let grad_output = Tensor::ones( - Shape::new(vec![2, 2]), - DataType::Float32, - Device::cpu(), - false, - ); - let grads = grad_fn.backward(&grad_output).unwrap(); - assert!(grads.contains_key(&ids[0])); - assert!(!grads.contains_key(&ids[1])); - } - - #[test] - fn test_transpose_backward_permutation() { - let _input = Tensor::new( - Arc::new(TensorData::from_vec_f32( - (0..24).map(|x| x as f32).collect(), - Device::cpu(), - )), - Shape::new(vec![2, 3, 4]), - DataType::Float32, - Device::cpu(), - false, - ); - let dims = vec![1, 2, 0]; - let grad_fn = TransposeBackward { - dims: dims.clone(), - input_id: TensorId::new(), - }; - let grad_output = Tensor::ones( - Shape::new(vec![3, 4, 2]), - DataType::Float32, - Device::cpu(), - false, - ); - let grads = grad_fn.backward(&grad_output).unwrap(); - let grad_input = grads.values().next().unwrap(); - let mut inverse = vec![0; dims.len()]; - for (i, &d) in dims.iter().enumerate() { - inverse[d] = i; - } - let mut expected = grad_output.clone(); - let mut current: Vec = (0..inverse.len()).collect(); - for i in 0..inverse.len() { - let j = current.iter().position(|&x| x == inverse[i]).unwrap(); - if i != j { - expected = crate::operations::linalg::transpose(&expected, i as isize, j as isize) - .unwrap(); - current.swap(i, j); - } - } - assert!(grad_input.allclose(&expected, 1e-6, 1e-6)); - } -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +#![allow(clippy::module_inception)] + +mod tests { + use crate::autograd::*; + use crate::device::Device; + use crate::error::{MinitensorError, Result}; + use crate::operations::{activation, arithmetic, reduction}; + use crate::tensor::DataType; + use crate::tensor::{Shape, Tensor, TensorData}; + use std::sync::Arc; + + #[test] + fn test_tensor_id_generation() { + let id1 = TensorId::new(); + let id2 = TensorId::new(); + assert_ne!(id1, id2); + } + + #[test] + fn test_computation_graph() { + let mut graph = ComputationGraph::new(); + let tensor_id = TensorId::new(); + + let grad_fn = Arc::new(AddBackward { + input_shapes: [vec![2, 2], vec![2, 2]], + input_ids: [TensorId::new(), TensorId::new()], + input_requires_grad: [true, true], + }); + + graph.add_tensor_with_grad_req(tensor_id, Some(grad_fn), true); + assert!(graph.nodes().contains_key(&tensor_id)); + } + + #[test] + fn test_add_backward() { + let grad_fn = AddBackward { + input_shapes: [vec![2, 2], vec![2, 2]], + input_ids: [TensorId::new(), TensorId::new()], + input_requires_grad: [true, true], + }; + + let grad_output = Tensor::ones( + Shape::new(vec![2, 2]), + crate::tensor::DataType::Float32, + Device::cpu(), + false, + ); + let gradients = grad_fn.backward(&grad_output).unwrap(); + + assert_eq!(gradients.len(), 2); + } + + #[test] + fn test_add_backward_same_input_accumulates() { + // `x + x`: both operands are the same tensor. The two gradient + // contributions must be summed (2), not overwritten (1). + let shared = TensorId::new(); + let grad_fn = AddBackward { + input_shapes: [vec![3], vec![3]], + input_ids: [shared, shared], + input_requires_grad: [true, true], + }; + + let grad_output = + Tensor::ones(Shape::new(vec![3]), DataType::Float32, Device::cpu(), false); + let gradients = grad_fn.backward(&grad_output).unwrap(); + + assert_eq!(gradients.len(), 1); + let grad = gradients.get(&shared).unwrap(); + for &v in grad.data().as_f32_slice().unwrap() { + assert_eq!(v, 2.0); + } + } + + #[test] + fn test_mul_backward_same_input_accumulates() { + // `x * x`: d/dx(x^2) = 2x. With shared operands the backward emits the + // same id twice (grad*rhs and grad*lhs); both must accumulate. + let shared = TensorId::new(); + let x = Tensor::new( + Arc::new(TensorData::from_vec_f32(vec![1.0, 2.0, 3.0], Device::cpu())), + Shape::new(vec![3]), + DataType::Float32, + Device::cpu(), + false, + ); + let grad_fn = MulBackward { + lhs: x.clone(), + rhs: x.clone(), + input_ids: [shared, shared], + input_requires_grad: [true, true], + }; + + let grad_output = + Tensor::ones(Shape::new(vec![3]), DataType::Float32, Device::cpu(), false); + let gradients = grad_fn.backward(&grad_output).unwrap(); + + assert_eq!(gradients.len(), 1); + let grad = gradients.get(&shared).unwrap(); + let expected = [2.0f32, 4.0, 6.0]; + for (&got, &want) in grad + .data() + .as_f32_slice() + .unwrap() + .iter() + .zip(expected.iter()) + { + assert_eq!(got, want); + } + } + + #[test] + fn test_reduce_gradient_for_broadcasting() { + let grad_output = Tensor::ones( + Shape::new(vec![2, 3]), + DataType::Float32, + Device::cpu(), + false, + ); + let target_shape = Shape::new(vec![2, 1]); + let reduced = reduce_gradient_for_broadcasting(&grad_output, &target_shape).unwrap(); + assert_eq!(reduced.shape().dims(), &[2, 1]); + let slice = reduced.data().as_f32_slice().unwrap(); + assert!(slice.iter().all(|&x| (x - 3.0).abs() < 1e-6)); + } + + #[test] + fn test_reduce_gradient_multiple_axes() { + let grad_output = Tensor::ones( + Shape::new(vec![2, 3]), + DataType::Float32, + Device::cpu(), + false, + ); + let target_shape = Shape::new(vec![1, 1]); + let reduced = reduce_gradient_for_broadcasting(&grad_output, &target_shape).unwrap(); + assert_eq!(reduced.shape().dims(), &[1, 1]); + let slice = reduced.data().as_f32_slice().unwrap(); + assert!((slice[0] - 6.0).abs() < 1e-6); + } + + #[test] + fn test_reduce_gradient_with_leading_and_inner_axes() { + let grad_output = Tensor::ones( + Shape::new(vec![4, 2, 3, 5]), + DataType::Float32, + Device::cpu(), + false, + ); + let target_shape = Shape::new(vec![2, 1, 5]); + let reduced = reduce_gradient_for_broadcasting(&grad_output, &target_shape).unwrap(); + assert_eq!(reduced.shape().dims(), &[2, 1, 5]); + let slice = reduced.data().as_f32_slice().unwrap(); + assert!(slice.iter().all(|&x| (x - 12.0).abs() < 1e-6)); + } + + #[test] + fn test_reduce_gradient_noop_for_same_shape() { + let grad_output = Tensor::ones( + Shape::new(vec![2, 3, 4]), + DataType::Float32, + Device::cpu(), + false, + ); + let target_shape = grad_output.shape().clone(); + let reduced = reduce_gradient_for_broadcasting(&grad_output, &target_shape).unwrap(); + assert_eq!(reduced.shape().dims(), &[2, 3, 4]); + assert!(reduced.allclose(&grad_output, 1e-6, 1e-6)); + } + + #[test] + fn test_reduce_gradient_invalid_broadcast() { + let grad_output = Tensor::ones( + Shape::new(vec![2, 1]), + DataType::Float32, + Device::cpu(), + false, + ); + let target_shape = Shape::new(vec![2, 2]); + let err = reduce_gradient_for_broadcasting(&grad_output, &target_shape) + .expect_err("expected invalid broadcast error"); + assert!(matches!(err, MinitensorError::BroadcastError { .. })); + } + + #[test] + fn test_reduce_gradient_zero_dim_broadcast() { + let grad_output = Tensor::ones( + Shape::new(vec![0, 2]), + DataType::Float32, + Device::cpu(), + false, + ); + let target_shape = Shape::new(vec![1, 2]); + let reduced = reduce_gradient_for_broadcasting(&grad_output, &target_shape).unwrap(); + assert_eq!(reduced.shape().dims(), &[1, 2]); + let slice = reduced.data().as_f32_slice().unwrap(); + assert!(slice.iter().all(|&x| x == 0.0)); + } + + #[test] + fn test_softmax_backward_dim1() { + let input = Tensor::new( + Arc::new(TensorData::from_vec_f32( + vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], + Device::cpu(), + )), + Shape::new(vec![2, 3]), + DataType::Float32, + Device::cpu(), + false, + ); + let softmax_out = activation::softmax(&input, Some(1)).unwrap(); + let grad_output = Tensor::ones( + Shape::new(vec![2, 3]), + DataType::Float32, + Device::cpu(), + false, + ); + let grad_y = arithmetic::mul(&grad_output, &softmax_out).unwrap(); + let sum = reduction::sum(&grad_y, Some(vec![1]), true).unwrap(); + let sub = arithmetic::sub(&grad_output, &sum).unwrap(); + let expected = arithmetic::mul(&softmax_out, &sub).unwrap(); + + let grad_fn = SoftmaxBackward { + input_id: TensorId::new(), + output: softmax_out.clone(), + dim: 1, + }; + let grads = grad_fn.backward(&grad_output).unwrap(); + let grad_input = grads.values().next().unwrap(); + assert!(grad_input.allclose(&expected, 1e-6, 1e-6)); + } + + #[test] + fn test_softmax_backward_dim0_f64() { + let data: Vec = (1..=6).map(|v| v as f64).collect(); + let input = Tensor::new( + Arc::new(TensorData::from_vec_f64(data, Device::cpu())), + Shape::new(vec![2, 3]), + DataType::Float64, + Device::cpu(), + false, + ); + let softmax_out = activation::softmax(&input, Some(0)).unwrap(); + let grad_output = Tensor::ones( + Shape::new(vec![2, 3]), + DataType::Float64, + Device::cpu(), + false, + ); + let grad_y = arithmetic::mul(&grad_output, &softmax_out).unwrap(); + let sum = reduction::sum(&grad_y, Some(vec![0]), true).unwrap(); + let sub = arithmetic::sub(&grad_output, &sum).unwrap(); + let expected = arithmetic::mul(&softmax_out, &sub).unwrap(); + + let grad_fn = SoftmaxBackward { + input_id: TensorId::new(), + output: softmax_out.clone(), + dim: 0, + }; + let grads = grad_fn.backward(&grad_output).unwrap(); + let grad_input = grads.values().next().unwrap(); + assert!(grad_input.allclose(&expected, 1e-6, 1e-6)); + } + + #[test] + fn test_backward_broadcast_addition() { + clear_graph().unwrap(); + + let a = Tensor::ones( + Shape::new(vec![2, 3]), + DataType::Float32, + Device::cpu(), + true, + ); + let b = Tensor::ones(Shape::new(vec![3]), DataType::Float32, Device::cpu(), true); + let out = arithmetic::add(&a, &b).unwrap(); + + let grad = Tensor::ones(out.shape().clone(), out.dtype(), out.device(), false); + let grads = backward_collect(&out, Some(grad)).unwrap(); + + let grad_a = grads.get(&a.id()).unwrap(); + let grad_b = grads.get(&b.id()).unwrap(); + assert_eq!(grad_a.shape().dims(), &[2, 3]); + assert_eq!(grad_b.shape().dims(), &[3]); + let slice_b = grad_b.data().as_f32_slice().unwrap(); + assert!(slice_b.iter().all(|&x| (x - 2.0).abs() < 1e-6)); + } + + #[test] + fn test_matmul_backward_gradients() { + let lhs = Tensor::new( + Arc::new(TensorData::from_vec_f32( + vec![1.0, 2.0, 3.0, 4.0], + Device::cpu(), + )), + Shape::new(vec![2, 2]), + DataType::Float32, + Device::cpu(), + false, + ); + let rhs = Tensor::new( + Arc::new(TensorData::from_vec_f32( + vec![5.0, 6.0, 7.0, 8.0], + Device::cpu(), + )), + Shape::new(vec![2, 2]), + DataType::Float32, + Device::cpu(), + false, + ); + let input_ids = [TensorId::new(), TensorId::new()]; + let grad_fn = MatMulBackward { + lhs: lhs.clone(), + rhs: rhs.clone(), + input_ids, + lhs_requires_grad: true, + rhs_requires_grad: true, + }; + let grad_output = Tensor::ones( + Shape::new(vec![2, 2]), + DataType::Float32, + Device::cpu(), + false, + ); + let grads = grad_fn.backward(&grad_output).unwrap(); + let rhs_t = crate::operations::linalg::transpose(&rhs, 0, 1).unwrap(); + let expected_lhs = crate::operations::linalg::matmul(&grad_output, &rhs_t).unwrap(); + let lhs_grad = grads.get(&input_ids[0]).unwrap(); + assert!(lhs_grad.allclose(&expected_lhs, 1e-6, 1e-6)); + } + + #[test] + fn test_matmul_backward_batched() { + let lhs = Tensor::new( + Arc::new(TensorData::from_vec_f32( + (0..12).map(|x| x as f32).collect(), + Device::cpu(), + )), + Shape::new(vec![2, 2, 3]), + DataType::Float32, + Device::cpu(), + false, + ); + let rhs = Tensor::new( + Arc::new(TensorData::from_vec_f32( + (0..24).map(|x| x as f32).collect(), + Device::cpu(), + )), + Shape::new(vec![2, 3, 4]), + DataType::Float32, + Device::cpu(), + false, + ); + let input_ids = [TensorId::new(), TensorId::new()]; + let grad_fn = MatMulBackward { + lhs: lhs.clone(), + rhs: rhs.clone(), + input_ids, + lhs_requires_grad: true, + rhs_requires_grad: true, + }; + let grad_output = Tensor::ones( + Shape::new(vec![2, 2, 4]), + DataType::Float32, + Device::cpu(), + false, + ); + let grads = grad_fn.backward(&grad_output).unwrap(); + let rhs_t = crate::operations::linalg::transpose( + &rhs, + (rhs.ndim() - 2) as isize, + (rhs.ndim() - 1) as isize, + ) + .unwrap(); + let expected_lhs = crate::operations::linalg::matmul(&grad_output, &rhs_t).unwrap(); + assert!( + grads + .get(&input_ids[0]) + .unwrap() + .allclose(&expected_lhs, 1e-6, 1e-6) + ); + let lhs_t = crate::operations::linalg::transpose( + &lhs, + (lhs.ndim() - 2) as isize, + (lhs.ndim() - 1) as isize, + ) + .unwrap(); + let expected_rhs = crate::operations::linalg::matmul(&lhs_t, &grad_output).unwrap(); + assert!( + grads + .get(&input_ids[1]) + .unwrap() + .allclose(&expected_rhs, 1e-6, 1e-6) + ); + } + + #[test] + fn test_matmul_backward_requires_grad_flags() { + let lhs = Tensor::new( + Arc::new(TensorData::from_vec_f32( + vec![1.0, 2.0, 3.0, 4.0], + Device::cpu(), + )), + Shape::new(vec![2, 2]), + DataType::Float32, + Device::cpu(), + false, + ); + let rhs = Tensor::new( + Arc::new(TensorData::from_vec_f32( + vec![5.0, 6.0, 7.0, 8.0], + Device::cpu(), + )), + Shape::new(vec![2, 2]), + DataType::Float32, + Device::cpu(), + false, + ); + let ids = [TensorId::new(), TensorId::new()]; + let grad_fn = MatMulBackward { + lhs: lhs.clone(), + rhs: rhs.clone(), + input_ids: ids, + lhs_requires_grad: true, + rhs_requires_grad: false, + }; + let grad_output = Tensor::ones( + Shape::new(vec![2, 2]), + DataType::Float32, + Device::cpu(), + false, + ); + let grads = grad_fn.backward(&grad_output).unwrap(); + assert!(grads.contains_key(&ids[0])); + assert!(!grads.contains_key(&ids[1])); + } + + #[test] + fn test_transpose_backward_permutation() { + let _input = Tensor::new( + Arc::new(TensorData::from_vec_f32( + (0..24).map(|x| x as f32).collect(), + Device::cpu(), + )), + Shape::new(vec![2, 3, 4]), + DataType::Float32, + Device::cpu(), + false, + ); + let dims = vec![1, 2, 0]; + let grad_fn = TransposeBackward { + dims: dims.clone(), + input_id: TensorId::new(), + }; + let grad_output = Tensor::ones( + Shape::new(vec![3, 4, 2]), + DataType::Float32, + Device::cpu(), + false, + ); + let grads = grad_fn.backward(&grad_output).unwrap(); + let grad_input = grads.values().next().unwrap(); + let mut inverse = vec![0; dims.len()]; + for (i, &d) in dims.iter().enumerate() { + inverse[d] = i; + } + let mut expected = grad_output.clone(); + let mut current: Vec = (0..inverse.len()).collect(); + for i in 0..inverse.len() { + let j = current.iter().position(|&x| x == inverse[i]).unwrap(); + if i != j { + expected = crate::operations::linalg::transpose(&expected, i as isize, j as isize) + .unwrap(); + current.swap(i, j); + } + } + assert!(grad_input.allclose(&expected, 1e-6, 1e-6)); + } + + #[test] + fn test_no_grad_guard_disables_recording() { + clear_graph().unwrap(); + + let a = Tensor::ones( + Shape::new(vec![2, 2]), + DataType::Float32, + Device::cpu(), + true, + ); + let b = Tensor::ones( + Shape::new(vec![2, 2]), + DataType::Float32, + Device::cpu(), + true, + ); + + { + let _guard = NoGradGuard::new(); + assert!(!is_grad_enabled()); + + // New tensors created inside the scope never require gradients. + let c = Tensor::ones(Shape::new(vec![2]), DataType::Float32, Device::cpu(), true); + assert!(!c.requires_grad()); + + // Operation results are detached leaves. + let out = arithmetic::add(&a, &b).unwrap(); + assert!(!out.requires_grad()); + assert!(out.grad_fn().is_none()); + + // Explicit opt-in still works. + let opted = c.requires_grad_(true); + assert!(opted.requires_grad()); + } + + // Mode restored: recording works again. + assert!(is_grad_enabled()); + let out = arithmetic::add(&a, &b).unwrap(); + assert!(out.requires_grad()); + assert!(out.grad_fn().is_some()); + + clear_graph().unwrap(); + } + + #[test] + fn test_set_grad_enabled_round_trip() { + assert!(is_grad_enabled()); + let prev = set_grad_enabled(false); + assert!(prev); + assert!(!is_grad_enabled()); + let prev = set_grad_enabled(true); + assert!(!prev); + assert!(is_grad_enabled()); + } + + #[test] + fn test_frozen_inputs_receive_no_gradient() { + clear_graph().unwrap(); + + let trainable = Tensor::ones( + Shape::new(vec![2, 2]), + DataType::Float32, + Device::cpu(), + true, + ); + let frozen = Tensor::ones( + Shape::new(vec![2, 2]), + DataType::Float32, + Device::cpu(), + false, + ); + + type BinOp = fn(&Tensor, &Tensor) -> Result; + let ops: [BinOp; 4] = [ + arithmetic::add, + arithmetic::sub, + arithmetic::mul, + arithmetic::div, + ]; + for op in ops { + let out = op(&trainable, &frozen).unwrap(); + let grad = Tensor::ones(out.shape().clone(), out.dtype(), out.device(), false); + let grads = backward_collect(&out, Some(grad)).unwrap(); + assert!( + grads.contains_key(&trainable.id()), + "trainable input must receive a gradient" + ); + assert!( + !grads.contains_key(&frozen.id()), + "frozen input must not receive a gradient" + ); + clear_graph().unwrap(); + } + } + + #[test] + fn test_loss_targets_receive_no_gradient() { + clear_graph().unwrap(); + + let predictions = Tensor::ones( + Shape::new(vec![4, 3]), + DataType::Float32, + Device::cpu(), + true, + ); + let targets = Tensor::ones( + Shape::new(vec![4, 3]), + DataType::Float32, + Device::cpu(), + false, + ); + + let loss = crate::operations::loss::mse_loss(&predictions, &targets, "mean").unwrap(); + let grads = backward_collect(&loss, None).unwrap(); + assert!( + grads.contains_key(&predictions.id()), + "predictions must receive a gradient" + ); + assert!( + !grads.contains_key(&targets.id()), + "frozen targets must not receive a gradient" + ); + clear_graph().unwrap(); + + // When targets DO require grad, both gradients flow. + let targets_rg = Tensor::ones( + Shape::new(vec![4, 3]), + DataType::Float32, + Device::cpu(), + true, + ); + let loss = crate::operations::loss::mse_loss(&predictions, &targets_rg, "mean").unwrap(); + let grads = backward_collect(&loss, None).unwrap(); + assert!(grads.contains_key(&predictions.id())); + assert!(grads.contains_key(&targets_rg.id())); + clear_graph().unwrap(); + } + + #[test] + fn test_concat_frozen_inputs_receive_no_gradient() { + clear_graph().unwrap(); + + let trainable = Tensor::ones( + Shape::new(vec![2, 2]), + DataType::Float32, + Device::cpu(), + true, + ); + let frozen = Tensor::ones( + Shape::new(vec![2, 2]), + DataType::Float32, + Device::cpu(), + false, + ); + + let out = crate::operations::shape_ops::concatenate(&[&trainable, &frozen], 0).unwrap(); + let grad = Tensor::ones(out.shape().clone(), out.dtype(), out.device(), false); + let grads = backward_collect(&out, Some(grad)).unwrap(); + + let trainable_grad = grads.get(&trainable.id()).expect("trainable grad"); + assert_eq!(trainable_grad.shape().dims(), &[2, 2]); + assert!(!grads.contains_key(&frozen.id())); + + clear_graph().unwrap(); + } +} diff --git a/engine/src/backends/opencl/integration_test.rs b/engine/src/backends/opencl/integration_test.rs index 88d1faf7..d47da867 100644 --- a/engine/src/backends/opencl/integration_test.rs +++ b/engine/src/backends/opencl/integration_test.rs @@ -23,7 +23,7 @@ mod tests { let backend = backend.unwrap(); assert!(backend.device().is_gpu()); assert_eq!( - backend.device().device_type, + backend.device().device_type(), crate::device::DeviceType::OpenCL ); } @@ -96,7 +96,7 @@ mod tests { let mut result_data = vec![0.0f32; 4]; backend.read_buffer(&c_buffer, &mut result_data).unwrap(); - let expected = vec![6.0f32, 8.0, 10.0, 12.0]; + let expected = [6.0f32, 8.0, 10.0, 12.0]; for (r, e) in result_data.iter().zip(expected.iter()) { assert!( (r - e).abs() < 1e-6, @@ -147,7 +147,7 @@ mod tests { backend.read_buffer(&c_buffer, &mut result_data).unwrap(); // Expected result: [1*5+2*7, 1*6+2*8, 3*5+4*7, 3*6+4*8] = [19, 22, 43, 50] - let expected = vec![19.0f32, 22.0, 43.0, 50.0]; + let expected = [19.0f32, 22.0, 43.0, 50.0]; for (r, e) in result_data.iter().zip(expected.iter()) { assert!( (r - e).abs() < 1e-6, @@ -192,7 +192,7 @@ mod tests { .read_buffer(&output_buffer, &mut result_data) .unwrap(); - let expected = vec![0.0f32, 0.0, 0.0, 1.0, 2.0]; + let expected = [0.0f32, 0.0, 0.0, 1.0, 2.0]; for (r, e) in result_data.iter().zip(expected.iter()) { assert!((r - e).abs() < 1e-6, "ReLU result mismatch: {} != {}", r, e); } diff --git a/engine/src/backends/opencl/mod.rs b/engine/src/backends/opencl/mod.rs index cd35310c..9a8dcf45 100644 --- a/engine/src/backends/opencl/mod.rs +++ b/engine/src/backends/opencl/mod.rs @@ -4,5 +4,7 @@ // This source code is licensed under the Apache-style license found in the // LICENSE file in the root directory of this source tree. -include!("mod/context.rs"); -include!("mod/kernels.rs"); +#[path = "mod/context.rs"] +mod context; + +pub use self::context::*; diff --git a/engine/src/backends/opencl/mod/context.rs b/engine/src/backends/opencl/mod/context.rs index b26d5c30..4f23bf85 100644 --- a/engine/src/backends/opencl/mod/context.rs +++ b/engine/src/backends/opencl/mod/context.rs @@ -1,660 +1,670 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -use super::Backend; -use crate::{device::Device, error::Result}; -use opencl3::command_queue::{CL_QUEUE_PROFILING_ENABLE, CommandQueue}; -use opencl3::context::Context; -use opencl3::device::{CL_DEVICE_TYPE_GPU, Device as OpenCLDevice}; -use opencl3::kernel::{ExecuteKernel, Kernel}; -use opencl3::memory::{Buffer, CL_MEM_READ_WRITE}; -use opencl3::platform::get_platforms; -use opencl3::program::Program; -use opencl3::types::{CL_BLOCKING, cl_float}; -use parking_lot::RwLock; -use rustc_hash::FxHashMap; -use std::ptr; -use std::sync::Arc; -use std::sync::atomic::{AtomicUsize, Ordering}; - -/// OpenCL buffer wrapper for memory management -pub struct OpenCLBuffer { - buffer: Buffer, - size_bytes: usize, -} - -unsafe impl Send for OpenCLBuffer {} -unsafe impl Sync for OpenCLBuffer {} - -/// OpenCL backend for cross-platform GPU tensor operations -pub struct OpenCLBackend { - device: Device, - opencl_device: OpenCLDevice, - context: Context, - command_queue: CommandQueue, - programs: Arc>>, - kernels: Arc>>, - buffers: Arc>>, - buffer_pool: Arc>>>>, - next_buffer_id: AtomicUsize, -} - -unsafe impl Send for OpenCLBackend {} -unsafe impl Sync for OpenCLBackend {} - -impl OpenCLBackend { - /// Get the OpenCL device - #[inline(always)] - pub fn opencl_device(&self) -> &OpenCLDevice { - &self.opencl_device - } - - /// Get the OpenCL context - #[inline(always)] - pub fn context(&self) -> &Context { - &self.context - } - - /// Get the command queue - #[inline(always)] - pub fn command_queue(&self) -> &CommandQueue { - &self.command_queue - } - - /// Create an OpenCL buffer - #[inline(always)] - pub fn create_buffer(&self, size: usize, flags: u64) -> Result> { - let buffer = - unsafe { Buffer::::create(&self.context, flags, size, ptr::null_mut()) } - .map_err(|e| { - crate::error::MinitensorError::memory_error(format!( - "Failed to create OpenCL buffer: {}", - e - )) - })?; - - Ok(buffer) - } - - /// Create an OpenCL buffer with data - #[inline(always)] - pub fn create_buffer_with_data(&self, data: &[f32], flags: u64) -> Result> { - let mut buffer = unsafe { - Buffer::::create(&self.context, flags, data.len(), ptr::null_mut()) - } - .map_err(|e| { - crate::error::MinitensorError::memory_error(format!( - "Failed to create OpenCL buffer: {}", - e - )) - })?; - - // Write data to buffer - unsafe { - self.command_queue - .enqueue_write_buffer(&mut buffer, CL_BLOCKING, 0, data, &[]) - .map_err(|e| { - crate::error::MinitensorError::memory_error(format!( - "Failed to write to OpenCL buffer: {}", - e - )) - })?; - } - - Ok(buffer) - } - - /// Build an OpenCL program - #[inline(always)] - pub fn build_program(&self, name: &str, source: &str) -> Result<()> { - let program = - Program::create_and_build_from_source(&self.context, source, "").map_err(|e| { - crate::error::MinitensorError::backend_error( - "OpenCL", - format!("Failed to build OpenCL program: {}", e), - ) - })?; - - let mut programs = self.programs.write(); - programs.insert(name.to_string(), program); - - Ok(()) - } - - /// Create a kernel from a program - #[inline(always)] - pub fn create_kernel(&self, program_name: &str, kernel_name: &str) -> Result<()> { - let programs = self.programs.read(); - let program = programs.get(program_name).ok_or_else(|| { - crate::error::MinitensorError::backend_error( - "OpenCL", - format!("Program '{}' not found", program_name), - ) - })?; - - let kernel = Kernel::create(program, kernel_name).map_err(|e| { - crate::error::MinitensorError::backend_error( - "OpenCL", - format!("Failed to create kernel '{}': {}", kernel_name, e), - ) - })?; - - let mut kernels = self.kernels.write(); - kernels.insert(kernel_name.to_string(), kernel); - - Ok(()) - } - - /// Get a kernel (creates a new kernel instance to avoid borrowing issues) - #[inline(always)] - pub fn get_kernel(&self, kernel_name: &str) -> Option { - let programs = self.programs.read(); - if let Some(program) = programs.get("tensor_ops") { - Kernel::create(program, kernel_name).ok() - } else { - None - } - } - - /// Execute a kernel - #[inline(always)] - pub fn execute_kernel( - &self, - kernel_name: &str, - global_work_size: &[usize], - local_work_size: Option<&[usize]>, - ) -> Result<()> { - let kernel = self.get_kernel(kernel_name).ok_or_else(|| { - crate::error::MinitensorError::backend_error( - "OpenCL", - format!("Kernel '{}' not found", kernel_name), - ) - })?; - - let kernel_event = unsafe { - ExecuteKernel::new(&kernel) - .set_global_work_sizes(global_work_size) - .set_local_work_sizes(local_work_size.unwrap_or(&[])) - .enqueue_nd_range(&self.command_queue) - } - .map_err(|e| { - crate::error::MinitensorError::backend_error( - "OpenCL", - format!("Failed to execute kernel: {}", e), - ) - })?; - - kernel_event.wait().map_err(|e| { - crate::error::MinitensorError::backend_error( - "OpenCL", - format!("Failed to wait for kernel completion: {}", e), - ) - })?; - - Ok(()) - } - - /// Read data from buffer - #[inline(always)] - pub fn read_buffer(&self, buffer: &Buffer, data: &mut [f32]) -> Result<()> { - unsafe { - self.command_queue - .enqueue_read_buffer(buffer, CL_BLOCKING, 0, data, &[]) - .map_err(|e| { - crate::error::MinitensorError::memory_error(format!( - "Failed to read from OpenCL buffer: {}", - e - )) - })?; - } - - Ok(()) - } - - /// Write data to buffer - #[inline(always)] - pub fn write_buffer(&self, buffer: &mut Buffer, data: &[f32]) -> Result<()> { - unsafe { - self.command_queue - .enqueue_write_buffer(buffer, CL_BLOCKING, 0, data, &[]) - .map_err(|e| { - crate::error::MinitensorError::memory_error(format!( - "Failed to write to OpenCL buffer: {}", - e - )) - })?; - } - - Ok(()) - } - - /// Execute operation on buffers by pointer - #[inline(always)] - pub fn execute_buffer_operation(&self, ptr: *const u8, operation: F) -> Result - where - F: FnOnce(&Buffer) -> Result, - { - let buffer_id = ptr as usize; - let buffers = self.buffers.read(); - if let Some(opencl_buffer) = buffers.get(&buffer_id) { - operation(&opencl_buffer.buffer) - } else { - Err(crate::error::MinitensorError::memory_error( - "OpenCL buffer not found for pointer", - )) - } - } - - /// Get buffer information for debugging - #[inline(always)] - pub fn get_buffer_info(&self, ptr: *const u8) -> Option<(usize, usize)> { - let buffer_id = ptr as usize; - let buffers = self.buffers.read(); - buffers - .get(&buffer_id) - .map(|buf| (buffer_id, buf.size_bytes)) - } - - /// Get total number of tracked buffers - #[inline(always)] - pub fn buffer_count(&self) -> usize { - self.buffers.read().len() - } - - /// Finish all operations in the command queue - #[inline(always)] - pub fn finish(&self) -> Result<()> { - self.command_queue.finish().map_err(|e| { - crate::error::MinitensorError::backend_error( - "OpenCL", - format!("Failed to finish OpenCL operations: {}", e), - ) - }) - } -} - -impl Backend for OpenCLBackend { - #[inline(always)] - fn device(&self) -> Device { - self.device - } - - #[inline(always)] - fn is_available() -> bool { - // Check if OpenCL platforms and GPU devices are available - if let Ok(platforms) = get_platforms() { - for platform in platforms { - if let Ok(devices) = - opencl3::device::get_device_ids(platform.id(), CL_DEVICE_TYPE_GPU) - { - if !devices.is_empty() { - return true; - } - } - } - } - false - } - - #[inline(always)] - fn initialize() -> Result { - // Get the first available GPU device - let platforms = get_platforms().map_err(|e| { - crate::error::MinitensorError::backend_error( - "OpenCL", - format!("Failed to get OpenCL platforms: {}", e), - ) - })?; - - let mut all_devices = Vec::new(); - for platform in platforms { - if let Ok(platform_devices) = - opencl3::device::get_device_ids(platform.id(), CL_DEVICE_TYPE_GPU) - { - all_devices.extend(platform_devices); - } - } - let devices = all_devices; - - if devices.is_empty() { - return Err(crate::error::MinitensorError::backend_error( - "OpenCL", - "No OpenCL GPU device found", - )); - } - - let opencl_device_id = devices[0]; - let opencl_device = opencl3::device::Device::new(opencl_device_id); - - // Create context and command queue - let context = Context::from_device(&opencl_device).map_err(|e| { - crate::error::MinitensorError::backend_error( - "OpenCL", - format!("Failed to create OpenCL context: {}", e), - ) - })?; - - #[allow(deprecated)] - let command_queue = CommandQueue::create_default(&context, CL_QUEUE_PROFILING_ENABLE) - .map_err(|e| { - crate::error::MinitensorError::backend_error( - "OpenCL", - format!("Failed to create OpenCL command queue: {}", e), - ) - })?; - - Ok(Self { - device: Device::opencl(Some(0)), - opencl_device, - context, - command_queue, - programs: Arc::new(RwLock::new(FxHashMap::default())), - kernels: Arc::new(RwLock::new(FxHashMap::default())), - buffers: Arc::new(RwLock::new(FxHashMap::default())), - buffer_pool: Arc::new(RwLock::new(FxHashMap::default())), - next_buffer_id: AtomicUsize::new(1), - }) - } - - #[inline(always)] - fn allocate(&self, size_bytes: usize) -> Result<*mut u8> { - if size_bytes == 0 { - return Ok(std::ptr::null_mut()); - } - - let buffer = { - let mut pool = self.buffer_pool.write(); - if let Some(buf) = pool.get_mut(&size_bytes).and_then(|v| v.pop()) { - buf - } else { - let size_floats = - (size_bytes + std::mem::size_of::() - 1) / std::mem::size_of::(); - drop(pool); - self.create_buffer(size_floats, CL_MEM_READ_WRITE)? - } - }; - - // Create a unique ID to track this buffer - let buffer_id = self.next_buffer_id.fetch_add(1, Ordering::Relaxed); - - let opencl_buffer = OpenCLBuffer { buffer, size_bytes }; - - // Store the buffer for tracking - let mut buffers = self.buffers.write(); - buffers.insert(buffer_id, opencl_buffer); - - // Return the buffer ID as a pointer - Ok(buffer_id as *mut u8) - } - - #[inline(always)] - fn deallocate(&self, ptr: *mut u8, _size_bytes: usize) -> Result<()> { - if ptr.is_null() { - return Ok(()); - } - - // Remove the buffer from tracking and return to pool - let buffer_id = ptr as usize; - let mut buffers = self.buffers.write(); - if let Some(opencl_buffer) = buffers.remove(&buffer_id) { - let mut pool = self.buffer_pool.write(); - pool.entry(opencl_buffer.size_bytes) - .or_default() - .push(opencl_buffer.buffer); - } - - Ok(()) - } - - #[inline(always)] - fn copy_from_host(&self, dst: *mut u8, src: &[u8]) -> Result<()> { - if src.is_empty() { - return Ok(()); - } - if dst.is_null() { - return Err(crate::error::MinitensorError::memory_error( - "Null destination pointer", - )); - } - - // Find the OpenCL buffer corresponding to this pointer - let buffer_id = dst as usize; - let mut buffers = self.buffers.write(); - if let Some(opencl_buffer) = buffers.get_mut(&buffer_id) { - // Convert bytes to f32 for OpenCL buffer - let src_floats = unsafe { - std::slice::from_raw_parts( - src.as_ptr() as *const f32, - src.len() / std::mem::size_of::(), - ) - }; - - unsafe { - self.command_queue.enqueue_write_buffer( - &mut opencl_buffer.buffer, - CL_BLOCKING, - 0, - src_floats, - &[], - ) - } - .map_err(|e| { - crate::error::MinitensorError::memory_error(format!( - "Failed to copy data to OpenCL buffer: {}", - e - )) - })?; - } else { - return Err(crate::error::MinitensorError::memory_error( - "OpenCL buffer not found for pointer", - )); - } - - Ok(()) - } - - #[inline(always)] - fn copy_to_host(&self, dst: &mut [u8], src: *const u8) -> Result<()> { - if dst.is_empty() { - return Ok(()); - } - if src.is_null() { - return Err(crate::error::MinitensorError::memory_error( - "Null source pointer", - )); - } - - // Find the OpenCL buffer corresponding to this pointer - let buffer_id = src as usize; - let buffers = self.buffers.read(); - if let Some(opencl_buffer) = buffers.get(&buffer_id) { - // Convert bytes to f32 for OpenCL buffer - let dst_floats = unsafe { - std::slice::from_raw_parts_mut( - dst.as_mut_ptr() as *mut f32, - dst.len() / std::mem::size_of::(), - ) - }; - - unsafe { - self.command_queue.enqueue_read_buffer( - &opencl_buffer.buffer, - CL_BLOCKING, - 0, - dst_floats, - &[], - ) - } - .map_err(|e| { - crate::error::MinitensorError::memory_error(format!( - "Failed to copy data from OpenCL buffer: {}", - e - )) - })?; - } else { - return Err(crate::error::MinitensorError::memory_error( - "OpenCL buffer not found for pointer", - )); - } - - Ok(()) - } -} - -impl Drop for OpenCLBackend { - fn drop(&mut self) { - { - let mut buffers = self.buffers.write(); - for (_, buf) in buffers.drain() { - drop(buf); - } - } - let mut pool = self.buffer_pool.write(); - for (_, mut vec) in pool.drain() { - for buf in vec.drain(..) { - drop(buf); - } - } - } -} - -/// OpenCL kernel source code for basic tensor operations -pub mod kernels { - /// Element-wise addition kernel - pub const ADD_KERNEL: &str = r#" -__kernel void add_kernel(__global const float* a, - __global const float* b, - __global float* c, - const unsigned int n) { - int gid = get_global_id(0); - if (gid < n) { - c[gid] = a[gid] + b[gid]; - } -} -"#; - - /// Element-wise multiplication kernel - pub const MUL_KERNEL: &str = r#" -__kernel void mul_kernel(__global const float* a, - __global const float* b, - __global float* c, - const unsigned int n) { - int gid = get_global_id(0); - if (gid < n) { - c[gid] = a[gid] * b[gid]; - } -} -"#; - - /// Matrix multiplication kernel - pub const MATMUL_KERNEL: &str = r#" -__kernel void matmul_kernel(__global const float* a, - __global const float* b, - __global float* c, - const unsigned int m, - const unsigned int n, - const unsigned int k) { - int row = get_global_id(1); - int col = get_global_id(0); - - if (row < m && col < n) { - float sum = 0.0f; - for (int i = 0; i < k; i++) { - sum += a[row * k + i] * b[i * n + col]; - } - c[row * n + col] = sum; - } -} -"#; - - /// ReLU activation kernel - pub const RELU_KERNEL: &str = r#" -__kernel void relu_kernel(__global const float* input, - __global float* output, - const unsigned int n) { - int gid = get_global_id(0); - if (gid < n) { - output[gid] = fmax(0.0f, input[gid]); - } -} -"#; - - /// Sigmoid activation kernel - pub const SIGMOID_KERNEL: &str = r#" -__kernel void sigmoid_kernel(__global const float* input, - __global float* output, - const unsigned int n) { - int gid = get_global_id(0); - if (gid < n) { - output[gid] = 1.0f / (1.0f + exp(-input[gid])); - } -} -"#; - - /// Combined kernel source - pub const ALL_KERNELS: &str = r#" -__kernel void add_kernel(__global const float* a, - __global const float* b, - __global float* c, - const unsigned int n) { - int gid = get_global_id(0); - if (gid < n) { - c[gid] = a[gid] + b[gid]; - } -} - -__kernel void mul_kernel(__global const float* a, - __global const float* b, - __global float* c, - const unsigned int n) { - int gid = get_global_id(0); - if (gid < n) { - c[gid] = a[gid] * b[gid]; - } -} - -__kernel void matmul_kernel(__global const float* a, - __global const float* b, - __global float* c, - const unsigned int m, - const unsigned int n, - const unsigned int k) { - int row = get_global_id(1); - int col = get_global_id(0); - - if (row < m && col < n) { - float sum = 0.0f; - for (int i = 0; i < k; i++) { - sum += a[row * k + i] * b[i * n + col]; - } - c[row * n + col] = sum; - } -} - -__kernel void relu_kernel(__global const float* input, - __global float* output, - const unsigned int n) { - int gid = get_global_id(0); - if (gid < n) { - output[gid] = fmax(0.0f, input[gid]); - } -} - -__kernel void sigmoid_kernel(__global const float* input, - __global float* output, - const unsigned int n) { - int gid = get_global_id(0); - if (gid < n) { - output[gid] = 1.0f / (1.0f + exp(-input[gid])); - } -} -"#; -} - -/// OpenCL operations for tensor computations -pub struct OpenCLOps { - backend: Arc, -} +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +// `ops_impl` holds the `impl OpenCLOps` block (kept out of the inline +// `kernels` module below, which holds OpenCL source strings); it is a child of this module so +// it keeps access to the OpenCL types declared here. +#[path = "kernels.rs"] +mod ops_impl; // impl-only module; nothing to re-export + +use crate::backends::Backend; +use crate::{device::Device, error::Result}; +use opencl3::command_queue::{CL_QUEUE_PROFILING_ENABLE, CommandQueue}; +use opencl3::context::Context; +use opencl3::device::{CL_DEVICE_TYPE_GPU, Device as OpenCLDevice}; +use opencl3::kernel::{ExecuteKernel, Kernel}; +use opencl3::memory::{Buffer, CL_MEM_READ_WRITE}; +use opencl3::platform::get_platforms; +use opencl3::program::Program; +use opencl3::types::{CL_BLOCKING, cl_float}; +use parking_lot::RwLock; +use rustc_hash::FxHashMap; +use std::ptr; +use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; + +/// OpenCL buffer wrapper for memory management +pub struct OpenCLBuffer { + buffer: Buffer, + size_bytes: usize, +} + +unsafe impl Send for OpenCLBuffer {} +unsafe impl Sync for OpenCLBuffer {} + +/// OpenCL backend for cross-platform GPU tensor operations +pub struct OpenCLBackend { + device: Device, + opencl_device: OpenCLDevice, + context: Context, + command_queue: CommandQueue, + programs: Arc>>, + kernels: Arc>>, + buffers: Arc>>, + buffer_pool: Arc>>>>, + next_buffer_id: AtomicUsize, +} + +unsafe impl Send for OpenCLBackend {} +unsafe impl Sync for OpenCLBackend {} + +impl OpenCLBackend { + /// Get the OpenCL device + #[inline(always)] + pub fn opencl_device(&self) -> &OpenCLDevice { + &self.opencl_device + } + + /// Get the OpenCL context + #[inline(always)] + pub fn context(&self) -> &Context { + &self.context + } + + /// Get the command queue + #[inline(always)] + pub fn command_queue(&self) -> &CommandQueue { + &self.command_queue + } + + /// Create an OpenCL buffer + #[inline(always)] + pub fn create_buffer(&self, size: usize, flags: u64) -> Result> { + let buffer = + unsafe { Buffer::::create(&self.context, flags, size, ptr::null_mut()) } + .map_err(|e| { + crate::error::MinitensorError::memory_error(format!( + "Failed to create OpenCL buffer: {}", + e + )) + })?; + + Ok(buffer) + } + + /// Create an OpenCL buffer with data + #[inline(always)] + pub fn create_buffer_with_data(&self, data: &[f32], flags: u64) -> Result> { + let mut buffer = unsafe { + Buffer::::create(&self.context, flags, data.len(), ptr::null_mut()) + } + .map_err(|e| { + crate::error::MinitensorError::memory_error(format!( + "Failed to create OpenCL buffer: {}", + e + )) + })?; + + // Write data to buffer + unsafe { + self.command_queue + .enqueue_write_buffer(&mut buffer, CL_BLOCKING, 0, data, &[]) + .map_err(|e| { + crate::error::MinitensorError::memory_error(format!( + "Failed to write to OpenCL buffer: {}", + e + )) + })?; + } + + Ok(buffer) + } + + /// Build an OpenCL program + #[inline(always)] + pub fn build_program(&self, name: &str, source: &str) -> Result<()> { + let program = + Program::create_and_build_from_source(&self.context, source, "").map_err(|e| { + crate::error::MinitensorError::backend_error( + "OpenCL", + format!("Failed to build OpenCL program: {}", e), + ) + })?; + + let mut programs = self.programs.write(); + programs.insert(name.to_string(), program); + + Ok(()) + } + + /// Create a kernel from a program + #[inline(always)] + pub fn create_kernel(&self, program_name: &str, kernel_name: &str) -> Result<()> { + let programs = self.programs.read(); + let program = programs.get(program_name).ok_or_else(|| { + crate::error::MinitensorError::backend_error( + "OpenCL", + format!("Program '{}' not found", program_name), + ) + })?; + + let kernel = Kernel::create(program, kernel_name).map_err(|e| { + crate::error::MinitensorError::backend_error( + "OpenCL", + format!("Failed to create kernel '{}': {}", kernel_name, e), + ) + })?; + + let mut kernels = self.kernels.write(); + kernels.insert(kernel_name.to_string(), kernel); + + Ok(()) + } + + /// Get a kernel (creates a new kernel instance to avoid borrowing issues) + #[inline(always)] + pub fn get_kernel(&self, kernel_name: &str) -> Option { + let programs = self.programs.read(); + if let Some(program) = programs.get("tensor_ops") { + Kernel::create(program, kernel_name).ok() + } else { + None + } + } + + /// Execute a kernel + #[inline(always)] + pub fn execute_kernel( + &self, + kernel_name: &str, + global_work_size: &[usize], + local_work_size: Option<&[usize]>, + ) -> Result<()> { + let kernel = self.get_kernel(kernel_name).ok_or_else(|| { + crate::error::MinitensorError::backend_error( + "OpenCL", + format!("Kernel '{}' not found", kernel_name), + ) + })?; + + let kernel_event = unsafe { + ExecuteKernel::new(&kernel) + .set_global_work_sizes(global_work_size) + .set_local_work_sizes(local_work_size.unwrap_or(&[])) + .enqueue_nd_range(&self.command_queue) + } + .map_err(|e| { + crate::error::MinitensorError::backend_error( + "OpenCL", + format!("Failed to execute kernel: {}", e), + ) + })?; + + kernel_event.wait().map_err(|e| { + crate::error::MinitensorError::backend_error( + "OpenCL", + format!("Failed to wait for kernel completion: {}", e), + ) + })?; + + Ok(()) + } + + /// Read data from buffer + #[inline(always)] + pub fn read_buffer(&self, buffer: &Buffer, data: &mut [f32]) -> Result<()> { + unsafe { + self.command_queue + .enqueue_read_buffer(buffer, CL_BLOCKING, 0, data, &[]) + .map_err(|e| { + crate::error::MinitensorError::memory_error(format!( + "Failed to read from OpenCL buffer: {}", + e + )) + })?; + } + + Ok(()) + } + + /// Write data to buffer + #[inline(always)] + pub fn write_buffer(&self, buffer: &mut Buffer, data: &[f32]) -> Result<()> { + unsafe { + self.command_queue + .enqueue_write_buffer(buffer, CL_BLOCKING, 0, data, &[]) + .map_err(|e| { + crate::error::MinitensorError::memory_error(format!( + "Failed to write to OpenCL buffer: {}", + e + )) + })?; + } + + Ok(()) + } + + /// Execute operation on buffers by pointer + #[inline(always)] + pub fn execute_buffer_operation(&self, ptr: *const u8, operation: F) -> Result + where + F: FnOnce(&Buffer) -> Result, + { + let buffer_id = ptr as usize; + let buffers = self.buffers.read(); + if let Some(opencl_buffer) = buffers.get(&buffer_id) { + operation(&opencl_buffer.buffer) + } else { + Err(crate::error::MinitensorError::memory_error( + "OpenCL buffer not found for pointer", + )) + } + } + + /// Get buffer information for debugging + #[inline(always)] + pub fn get_buffer_info(&self, ptr: *const u8) -> Option<(usize, usize)> { + let buffer_id = ptr as usize; + let buffers = self.buffers.read(); + buffers + .get(&buffer_id) + .map(|buf| (buffer_id, buf.size_bytes)) + } + + /// Get total number of tracked buffers + #[inline(always)] + pub fn buffer_count(&self) -> usize { + self.buffers.read().len() + } + + /// Finish all operations in the command queue + #[inline(always)] + pub fn finish(&self) -> Result<()> { + self.command_queue.finish().map_err(|e| { + crate::error::MinitensorError::backend_error( + "OpenCL", + format!("Failed to finish OpenCL operations: {}", e), + ) + }) + } +} + +impl Backend for OpenCLBackend { + #[inline(always)] + fn device(&self) -> Device { + self.device + } + + #[inline(always)] + fn is_available() -> bool { + // Check if OpenCL platforms and GPU devices are available + if let Ok(platforms) = get_platforms() { + for platform in platforms { + if let Ok(devices) = + opencl3::device::get_device_ids(platform.id(), CL_DEVICE_TYPE_GPU) + && !devices.is_empty() + { + return true; + } + } + } + false + } + + // The OpenCL handles cached in these `Arc>` fields are not + // `Send`/`Sync` — the OpenCL backend is single-threaded by design (all + // access is serialized through the owning backend). The `Arc` is for + // shared ownership within that single thread, so the lint's concern + // (cross-thread sharing of a non-thread-safe value) does not apply. + #[allow(clippy::arc_with_non_send_sync)] + #[inline(always)] + fn initialize() -> Result { + // Get the first available GPU device + let platforms = get_platforms().map_err(|e| { + crate::error::MinitensorError::backend_error( + "OpenCL", + format!("Failed to get OpenCL platforms: {}", e), + ) + })?; + + let mut all_devices = Vec::new(); + for platform in platforms { + if let Ok(platform_devices) = + opencl3::device::get_device_ids(platform.id(), CL_DEVICE_TYPE_GPU) + { + all_devices.extend(platform_devices); + } + } + let devices = all_devices; + + if devices.is_empty() { + return Err(crate::error::MinitensorError::backend_error( + "OpenCL", + "No OpenCL GPU device found", + )); + } + + let opencl_device_id = devices[0]; + let opencl_device = opencl3::device::Device::new(opencl_device_id); + + // Create context and command queue + let context = Context::from_device(&opencl_device).map_err(|e| { + crate::error::MinitensorError::backend_error( + "OpenCL", + format!("Failed to create OpenCL context: {}", e), + ) + })?; + + #[allow(deprecated)] + let command_queue = CommandQueue::create_default(&context, CL_QUEUE_PROFILING_ENABLE) + .map_err(|e| { + crate::error::MinitensorError::backend_error( + "OpenCL", + format!("Failed to create OpenCL command queue: {}", e), + ) + })?; + + Ok(Self { + device: Device::opencl(Some(0)), + opencl_device, + context, + command_queue, + programs: Arc::new(RwLock::new(FxHashMap::default())), + kernels: Arc::new(RwLock::new(FxHashMap::default())), + buffers: Arc::new(RwLock::new(FxHashMap::default())), + buffer_pool: Arc::new(RwLock::new(FxHashMap::default())), + next_buffer_id: AtomicUsize::new(1), + }) + } + + #[inline(always)] + fn allocate(&self, size_bytes: usize) -> Result<*mut u8> { + if size_bytes == 0 { + return Ok(std::ptr::null_mut()); + } + + let buffer = { + let mut pool = self.buffer_pool.write(); + if let Some(buf) = pool.get_mut(&size_bytes).and_then(|v| v.pop()) { + buf + } else { + let size_floats = size_bytes.div_ceil(std::mem::size_of::()); + drop(pool); + self.create_buffer(size_floats, CL_MEM_READ_WRITE)? + } + }; + + // Create a unique ID to track this buffer + let buffer_id = self.next_buffer_id.fetch_add(1, Ordering::Relaxed); + + let opencl_buffer = OpenCLBuffer { buffer, size_bytes }; + + // Store the buffer for tracking + let mut buffers = self.buffers.write(); + buffers.insert(buffer_id, opencl_buffer); + + // Return the buffer ID as a pointer + Ok(buffer_id as *mut u8) + } + + #[inline(always)] + fn deallocate(&self, ptr: *mut u8, _size_bytes: usize) -> Result<()> { + if ptr.is_null() { + return Ok(()); + } + + // Remove the buffer from tracking and return to pool + let buffer_id = ptr as usize; + let mut buffers = self.buffers.write(); + if let Some(opencl_buffer) = buffers.remove(&buffer_id) { + let mut pool = self.buffer_pool.write(); + pool.entry(opencl_buffer.size_bytes) + .or_default() + .push(opencl_buffer.buffer); + } + + Ok(()) + } + + #[inline(always)] + fn copy_from_host(&self, dst: *mut u8, src: &[u8]) -> Result<()> { + if src.is_empty() { + return Ok(()); + } + if dst.is_null() { + return Err(crate::error::MinitensorError::memory_error( + "Null destination pointer", + )); + } + + // Find the OpenCL buffer corresponding to this pointer + let buffer_id = dst as usize; + let mut buffers = self.buffers.write(); + if let Some(opencl_buffer) = buffers.get_mut(&buffer_id) { + // Convert bytes to f32 for OpenCL buffer + let src_floats = unsafe { + std::slice::from_raw_parts( + src.as_ptr() as *const f32, + src.len() / std::mem::size_of::(), + ) + }; + + unsafe { + self.command_queue.enqueue_write_buffer( + &mut opencl_buffer.buffer, + CL_BLOCKING, + 0, + src_floats, + &[], + ) + } + .map_err(|e| { + crate::error::MinitensorError::memory_error(format!( + "Failed to copy data to OpenCL buffer: {}", + e + )) + })?; + } else { + return Err(crate::error::MinitensorError::memory_error( + "OpenCL buffer not found for pointer", + )); + } + + Ok(()) + } + + #[inline(always)] + fn copy_to_host(&self, dst: &mut [u8], src: *const u8) -> Result<()> { + if dst.is_empty() { + return Ok(()); + } + if src.is_null() { + return Err(crate::error::MinitensorError::memory_error( + "Null source pointer", + )); + } + + // Find the OpenCL buffer corresponding to this pointer + let buffer_id = src as usize; + let buffers = self.buffers.read(); + if let Some(opencl_buffer) = buffers.get(&buffer_id) { + // Convert bytes to f32 for OpenCL buffer + let dst_floats = unsafe { + std::slice::from_raw_parts_mut( + dst.as_mut_ptr() as *mut f32, + dst.len() / std::mem::size_of::(), + ) + }; + + unsafe { + self.command_queue.enqueue_read_buffer( + &opencl_buffer.buffer, + CL_BLOCKING, + 0, + dst_floats, + &[], + ) + } + .map_err(|e| { + crate::error::MinitensorError::memory_error(format!( + "Failed to copy data from OpenCL buffer: {}", + e + )) + })?; + } else { + return Err(crate::error::MinitensorError::memory_error( + "OpenCL buffer not found for pointer", + )); + } + + Ok(()) + } +} + +impl Drop for OpenCLBackend { + fn drop(&mut self) { + { + let mut buffers = self.buffers.write(); + for (_, buf) in buffers.drain() { + drop(buf); + } + } + let mut pool = self.buffer_pool.write(); + for (_, mut vec) in pool.drain() { + for buf in vec.drain(..) { + drop(buf); + } + } + } +} + +/// OpenCL kernel source code for basic tensor operations +pub mod kernels { + /// Element-wise addition kernel + pub const ADD_KERNEL: &str = r#" +__kernel void add_kernel(__global const float* a, + __global const float* b, + __global float* c, + const unsigned int n) { + int gid = get_global_id(0); + if (gid < n) { + c[gid] = a[gid] + b[gid]; + } +} +"#; + + /// Element-wise multiplication kernel + pub const MUL_KERNEL: &str = r#" +__kernel void mul_kernel(__global const float* a, + __global const float* b, + __global float* c, + const unsigned int n) { + int gid = get_global_id(0); + if (gid < n) { + c[gid] = a[gid] * b[gid]; + } +} +"#; + + /// Matrix multiplication kernel + pub const MATMUL_KERNEL: &str = r#" +__kernel void matmul_kernel(__global const float* a, + __global const float* b, + __global float* c, + const unsigned int m, + const unsigned int n, + const unsigned int k) { + int row = get_global_id(1); + int col = get_global_id(0); + + if (row < m && col < n) { + float sum = 0.0f; + for (int i = 0; i < k; i++) { + sum += a[row * k + i] * b[i * n + col]; + } + c[row * n + col] = sum; + } +} +"#; + + /// ReLU activation kernel + pub const RELU_KERNEL: &str = r#" +__kernel void relu_kernel(__global const float* input, + __global float* output, + const unsigned int n) { + int gid = get_global_id(0); + if (gid < n) { + output[gid] = fmax(0.0f, input[gid]); + } +} +"#; + + /// Sigmoid activation kernel + pub const SIGMOID_KERNEL: &str = r#" +__kernel void sigmoid_kernel(__global const float* input, + __global float* output, + const unsigned int n) { + int gid = get_global_id(0); + if (gid < n) { + output[gid] = 1.0f / (1.0f + exp(-input[gid])); + } +} +"#; + + /// Combined kernel source + pub const ALL_KERNELS: &str = r#" +__kernel void add_kernel(__global const float* a, + __global const float* b, + __global float* c, + const unsigned int n) { + int gid = get_global_id(0); + if (gid < n) { + c[gid] = a[gid] + b[gid]; + } +} + +__kernel void mul_kernel(__global const float* a, + __global const float* b, + __global float* c, + const unsigned int n) { + int gid = get_global_id(0); + if (gid < n) { + c[gid] = a[gid] * b[gid]; + } +} + +__kernel void matmul_kernel(__global const float* a, + __global const float* b, + __global float* c, + const unsigned int m, + const unsigned int n, + const unsigned int k) { + int row = get_global_id(1); + int col = get_global_id(0); + + if (row < m && col < n) { + float sum = 0.0f; + for (int i = 0; i < k; i++) { + sum += a[row * k + i] * b[i * n + col]; + } + c[row * n + col] = sum; + } +} + +__kernel void relu_kernel(__global const float* input, + __global float* output, + const unsigned int n) { + int gid = get_global_id(0); + if (gid < n) { + output[gid] = fmax(0.0f, input[gid]); + } +} + +__kernel void sigmoid_kernel(__global const float* input, + __global float* output, + const unsigned int n) { + int gid = get_global_id(0); + if (gid < n) { + output[gid] = 1.0f / (1.0f + exp(-input[gid])); + } +} +"#; +} + +/// OpenCL operations for tensor computations +pub struct OpenCLOps { + backend: Arc, +} diff --git a/engine/src/backends/opencl/mod/kernels.rs b/engine/src/backends/opencl/mod/kernels.rs index acd43a2e..bfc485d0 100644 --- a/engine/src/backends/opencl/mod/kernels.rs +++ b/engine/src/backends/opencl/mod/kernels.rs @@ -1,398 +1,401 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -impl OpenCLOps { - /// Create new OpenCL operations instance - pub fn new(backend: Arc) -> Result { - // Build the kernel program - backend.build_program("tensor_ops", kernels::ALL_KERNELS)?; - - // Create all kernels - backend.create_kernel("tensor_ops", "add_kernel")?; - backend.create_kernel("tensor_ops", "mul_kernel")?; - backend.create_kernel("tensor_ops", "matmul_kernel")?; - backend.create_kernel("tensor_ops", "relu_kernel")?; - backend.create_kernel("tensor_ops", "sigmoid_kernel")?; - - Ok(Self { backend }) - } - - /// Element-wise addition on GPU - pub fn add( - &self, - a: &Buffer, - b: &Buffer, - c: &Buffer, - n: u32, - ) -> Result<()> { - let kernel = self.backend.get_kernel("add_kernel").ok_or_else(|| { - crate::error::MinitensorError::backend_error("OpenCL", "Add kernel not found") - })?; - - unsafe { - ExecuteKernel::new(&kernel) - .set_arg(a) - .set_arg(b) - .set_arg(c) - .set_arg(&n) - .set_global_work_size(n as usize) - .enqueue_nd_range(&self.backend.command_queue) - } - .map_err(|e| { - crate::error::MinitensorError::backend_error( - "OpenCL", - format!("Failed to execute add kernel: {}", e), - ) - })? - .wait() - .map_err(|e| { - crate::error::MinitensorError::backend_error( - "OpenCL", - format!("Failed to wait for add kernel: {}", e), - ) - })?; - - Ok(()) - } - - /// Element-wise multiplication on GPU - pub fn mul( - &self, - a: &Buffer, - b: &Buffer, - c: &Buffer, - n: u32, - ) -> Result<()> { - let kernel = self.backend.get_kernel("mul_kernel").ok_or_else(|| { - crate::error::MinitensorError::backend_error("OpenCL", "Mul kernel not found") - })?; - - unsafe { - ExecuteKernel::new(&kernel) - .set_arg(a) - .set_arg(b) - .set_arg(c) - .set_arg(&n) - .set_global_work_size(n as usize) - .enqueue_nd_range(&self.backend.command_queue) - } - .map_err(|e| { - crate::error::MinitensorError::backend_error( - "OpenCL", - format!("Failed to execute mul kernel: {}", e), - ) - })? - .wait() - .map_err(|e| { - crate::error::MinitensorError::backend_error( - "OpenCL", - format!("Failed to wait for mul kernel: {}", e), - ) - })?; - - Ok(()) - } - - /// Matrix multiplication on GPU - pub fn matmul( - &self, - a: &Buffer, - b: &Buffer, - c: &Buffer, - m: u32, - n: u32, - k: u32, - ) -> Result<()> { - let kernel = self.backend.get_kernel("matmul_kernel").ok_or_else(|| { - crate::error::MinitensorError::backend_error("OpenCL", "Matmul kernel not found") - })?; - - unsafe { - ExecuteKernel::new(&kernel) - .set_arg(a) - .set_arg(b) - .set_arg(c) - .set_arg(&m) - .set_arg(&n) - .set_arg(&k) - .set_global_work_sizes(&[n as usize, m as usize]) - .enqueue_nd_range(&self.backend.command_queue) - } - .map_err(|e| { - crate::error::MinitensorError::backend_error( - "OpenCL", - format!("Failed to execute matmul kernel: {}", e), - ) - })? - .wait() - .map_err(|e| { - crate::error::MinitensorError::backend_error( - "OpenCL", - format!("Failed to wait for matmul kernel: {}", e), - ) - })?; - - Ok(()) - } - - /// ReLU activation on GPU - pub fn relu(&self, input: &Buffer, output: &Buffer, n: u32) -> Result<()> { - let kernel = self.backend.get_kernel("relu_kernel").ok_or_else(|| { - crate::error::MinitensorError::backend_error("OpenCL", "ReLU kernel not found") - })?; - - unsafe { - ExecuteKernel::new(&kernel) - .set_arg(input) - .set_arg(output) - .set_arg(&n) - .set_global_work_size(n as usize) - .enqueue_nd_range(&self.backend.command_queue) - } - .map_err(|e| { - crate::error::MinitensorError::backend_error( - "OpenCL", - format!("Failed to execute relu kernel: {}", e), - ) - })? - .wait() - .map_err(|e| { - crate::error::MinitensorError::backend_error( - "OpenCL", - format!("Failed to wait for relu kernel: {}", e), - ) - })?; - - Ok(()) - } - - /// Sigmoid activation on GPU - pub fn sigmoid( - &self, - input: &Buffer, - output: &Buffer, - n: u32, - ) -> Result<()> { - let kernel = self.backend.get_kernel("sigmoid_kernel").ok_or_else(|| { - crate::error::MinitensorError::backend_error("OpenCL", "Sigmoid kernel not found") - })?; - - unsafe { - ExecuteKernel::new(&kernel) - .set_arg(input) - .set_arg(output) - .set_arg(&n) - .set_global_work_size(n as usize) - .enqueue_nd_range(&self.backend.command_queue) - } - .map_err(|e| { - crate::error::MinitensorError::backend_error( - "OpenCL", - format!("Failed to execute sigmoid kernel: {}", e), - ) - })? - .wait() - .map_err(|e| { - crate::error::MinitensorError::backend_error( - "OpenCL", - format!("Failed to wait for sigmoid kernel: {}", e), - ) - })?; - - Ok(()) - } - - /// Execute element-wise addition using pointers - pub fn add_ptr( - &self, - a_ptr: *const u8, - b_ptr: *const u8, - c_ptr: *mut u8, - n: u32, - ) -> Result<()> { - // Get buffers from the backend's buffer tracking system - let a_buffer_id = a_ptr as usize; - let b_buffer_id = b_ptr as usize; - let c_buffer_id = c_ptr as usize; - - let buffers = self.backend.buffers.read(); - - if let (Some(a_buf), Some(b_buf), Some(c_buf)) = ( - buffers.get(&a_buffer_id), - buffers.get(&b_buffer_id), - buffers.get(&c_buffer_id), - ) { - self.add(&a_buf.buffer, &b_buf.buffer, &c_buf.buffer, n) - } else { - Err(crate::error::MinitensorError::memory_error( - "OpenCL buffer not found for operation", - )) - } - } - - /// Execute element-wise multiplication using pointers - pub fn mul_ptr( - &self, - a_ptr: *const u8, - b_ptr: *const u8, - c_ptr: *mut u8, - n: u32, - ) -> Result<()> { - let a_buffer_id = a_ptr as usize; - let b_buffer_id = b_ptr as usize; - let c_buffer_id = c_ptr as usize; - - let buffers = self.backend.buffers.read(); - - if let (Some(a_buf), Some(b_buf), Some(c_buf)) = ( - buffers.get(&a_buffer_id), - buffers.get(&b_buffer_id), - buffers.get(&c_buffer_id), - ) { - self.mul(&a_buf.buffer, &b_buf.buffer, &c_buf.buffer, n) - } else { - Err(crate::error::MinitensorError::memory_error( - "OpenCL buffer not found for operation", - )) - } - } - - /// Execute matrix multiplication using pointers - pub fn matmul_ptr( - &self, - a_ptr: *const u8, - b_ptr: *const u8, - c_ptr: *mut u8, - m: u32, - n: u32, - k: u32, - ) -> Result<()> { - let a_buffer_id = a_ptr as usize; - let b_buffer_id = b_ptr as usize; - let c_buffer_id = c_ptr as usize; - - let buffers = self.backend.buffers.read(); - - if let (Some(a_buf), Some(b_buf), Some(c_buf)) = ( - buffers.get(&a_buffer_id), - buffers.get(&b_buffer_id), - buffers.get(&c_buffer_id), - ) { - self.matmul(&a_buf.buffer, &b_buf.buffer, &c_buf.buffer, m, n, k) - } else { - Err(crate::error::MinitensorError::memory_error( - "OpenCL buffer not found for operation", - )) - } - } - - /// Execute ReLU activation using pointers - pub fn relu_ptr(&self, input_ptr: *const u8, output_ptr: *mut u8, n: u32) -> Result<()> { - let input_buffer_id = input_ptr as usize; - let output_buffer_id = output_ptr as usize; - - let buffers = self.backend.buffers.read(); - - if let (Some(input_buf), Some(output_buf)) = ( - buffers.get(&input_buffer_id), - buffers.get(&output_buffer_id), - ) { - self.relu(&input_buf.buffer, &output_buf.buffer, n) - } else { - Err(crate::error::MinitensorError::memory_error( - "OpenCL buffer not found for operation", - )) - } - } - - /// Execute Sigmoid activation using pointers - pub fn sigmoid_ptr(&self, input_ptr: *const u8, output_ptr: *mut u8, n: u32) -> Result<()> { - let input_buffer_id = input_ptr as usize; - let output_buffer_id = output_ptr as usize; - - let buffers = self.backend.buffers.read(); - - if let (Some(input_buf), Some(output_buf)) = ( - buffers.get(&input_buffer_id), - buffers.get(&output_buffer_id), - ) { - self.sigmoid(&input_buf.buffer, &output_buf.buffer, n) - } else { - Err(crate::error::MinitensorError::memory_error( - "OpenCL buffer not found for operation", - )) - } - } -} - -#[cfg(test)] -mod integration_test; - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_opencl_availability() { - // This test will only pass if OpenCL is available - if OpenCLBackend::is_available() { - let backend = OpenCLBackend::initialize().unwrap(); - assert!(backend.device().is_gpu()); - } - } - - #[test] - fn test_opencl_buffer_operations() { - if !OpenCLBackend::is_available() { - return; // Skip test if OpenCL not available - } - - let backend = OpenCLBackend::initialize().unwrap(); - - // Test buffer creation and data transfer - let data = vec![1.0f32, 2.0, 3.0, 4.0, 5.0]; - let buffer = backend - .create_buffer_with_data(&data, CL_MEM_READ_WRITE) - .unwrap(); - - let mut result = vec![0.0f32; 5]; - backend.read_buffer(&buffer, &mut result).unwrap(); - - assert_eq!(data, result); - } - - #[test] - fn test_opencl_operations() { - if !OpenCLBackend::is_available() { - return; // Skip test if OpenCL not available - } - - let backend = Arc::new(OpenCLBackend::initialize().unwrap()); - let ops = OpenCLOps::new(backend.clone()).unwrap(); - - // Test addition - let a_data = vec![1.0f32, 2.0, 3.0, 4.0]; - let b_data = vec![5.0f32, 6.0, 7.0, 8.0]; - - let a_buffer = backend - .create_buffer_with_data(&a_data, CL_MEM_READ_ONLY) - .unwrap(); - let b_buffer = backend - .create_buffer_with_data(&b_data, CL_MEM_READ_ONLY) - .unwrap(); - let c_buffer = backend.create_buffer(4, CL_MEM_WRITE_ONLY).unwrap(); - - ops.add(&a_buffer, &b_buffer, &c_buffer, 4).unwrap(); - - let mut result = vec![0.0f32; 4]; - backend.read_buffer(&c_buffer, &mut result).unwrap(); - - let expected = vec![6.0f32, 8.0, 10.0, 12.0]; - for (r, e) in result.iter().zip(expected.iter()) { - assert!((r - e).abs() < 1e-6); - } - } -} +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +impl OpenCLOps { + /// Create new OpenCL operations instance + pub fn new(backend: Arc) -> Result { + // Build the kernel program + backend.build_program("tensor_ops", kernels::ALL_KERNELS)?; + + // Create all kernels + backend.create_kernel("tensor_ops", "add_kernel")?; + backend.create_kernel("tensor_ops", "mul_kernel")?; + backend.create_kernel("tensor_ops", "matmul_kernel")?; + backend.create_kernel("tensor_ops", "relu_kernel")?; + backend.create_kernel("tensor_ops", "sigmoid_kernel")?; + + Ok(Self { backend }) + } + + /// Element-wise addition on GPU + pub fn add( + &self, + a: &Buffer, + b: &Buffer, + c: &Buffer, + n: u32, + ) -> Result<()> { + let kernel = self.backend.get_kernel("add_kernel").ok_or_else(|| { + crate::error::MinitensorError::backend_error("OpenCL", "Add kernel not found") + })?; + + unsafe { + ExecuteKernel::new(&kernel) + .set_arg(a) + .set_arg(b) + .set_arg(c) + .set_arg(&n) + .set_global_work_size(n as usize) + .enqueue_nd_range(&self.backend.command_queue) + } + .map_err(|e| { + crate::error::MinitensorError::backend_error( + "OpenCL", + format!("Failed to execute add kernel: {}", e), + ) + })? + .wait() + .map_err(|e| { + crate::error::MinitensorError::backend_error( + "OpenCL", + format!("Failed to wait for add kernel: {}", e), + ) + })?; + + Ok(()) + } + + /// Element-wise multiplication on GPU + pub fn mul( + &self, + a: &Buffer, + b: &Buffer, + c: &Buffer, + n: u32, + ) -> Result<()> { + let kernel = self.backend.get_kernel("mul_kernel").ok_or_else(|| { + crate::error::MinitensorError::backend_error("OpenCL", "Mul kernel not found") + })?; + + unsafe { + ExecuteKernel::new(&kernel) + .set_arg(a) + .set_arg(b) + .set_arg(c) + .set_arg(&n) + .set_global_work_size(n as usize) + .enqueue_nd_range(&self.backend.command_queue) + } + .map_err(|e| { + crate::error::MinitensorError::backend_error( + "OpenCL", + format!("Failed to execute mul kernel: {}", e), + ) + })? + .wait() + .map_err(|e| { + crate::error::MinitensorError::backend_error( + "OpenCL", + format!("Failed to wait for mul kernel: {}", e), + ) + })?; + + Ok(()) + } + + /// Matrix multiplication on GPU + pub fn matmul( + &self, + a: &Buffer, + b: &Buffer, + c: &Buffer, + m: u32, + n: u32, + k: u32, + ) -> Result<()> { + let kernel = self.backend.get_kernel("matmul_kernel").ok_or_else(|| { + crate::error::MinitensorError::backend_error("OpenCL", "Matmul kernel not found") + })?; + + unsafe { + ExecuteKernel::new(&kernel) + .set_arg(a) + .set_arg(b) + .set_arg(c) + .set_arg(&m) + .set_arg(&n) + .set_arg(&k) + .set_global_work_sizes(&[n as usize, m as usize]) + .enqueue_nd_range(&self.backend.command_queue) + } + .map_err(|e| { + crate::error::MinitensorError::backend_error( + "OpenCL", + format!("Failed to execute matmul kernel: {}", e), + ) + })? + .wait() + .map_err(|e| { + crate::error::MinitensorError::backend_error( + "OpenCL", + format!("Failed to wait for matmul kernel: {}", e), + ) + })?; + + Ok(()) + } + + /// ReLU activation on GPU + pub fn relu(&self, input: &Buffer, output: &Buffer, n: u32) -> Result<()> { + let kernel = self.backend.get_kernel("relu_kernel").ok_or_else(|| { + crate::error::MinitensorError::backend_error("OpenCL", "ReLU kernel not found") + })?; + + unsafe { + ExecuteKernel::new(&kernel) + .set_arg(input) + .set_arg(output) + .set_arg(&n) + .set_global_work_size(n as usize) + .enqueue_nd_range(&self.backend.command_queue) + } + .map_err(|e| { + crate::error::MinitensorError::backend_error( + "OpenCL", + format!("Failed to execute relu kernel: {}", e), + ) + })? + .wait() + .map_err(|e| { + crate::error::MinitensorError::backend_error( + "OpenCL", + format!("Failed to wait for relu kernel: {}", e), + ) + })?; + + Ok(()) + } + + /// Sigmoid activation on GPU + pub fn sigmoid( + &self, + input: &Buffer, + output: &Buffer, + n: u32, + ) -> Result<()> { + let kernel = self.backend.get_kernel("sigmoid_kernel").ok_or_else(|| { + crate::error::MinitensorError::backend_error("OpenCL", "Sigmoid kernel not found") + })?; + + unsafe { + ExecuteKernel::new(&kernel) + .set_arg(input) + .set_arg(output) + .set_arg(&n) + .set_global_work_size(n as usize) + .enqueue_nd_range(&self.backend.command_queue) + } + .map_err(|e| { + crate::error::MinitensorError::backend_error( + "OpenCL", + format!("Failed to execute sigmoid kernel: {}", e), + ) + })? + .wait() + .map_err(|e| { + crate::error::MinitensorError::backend_error( + "OpenCL", + format!("Failed to wait for sigmoid kernel: {}", e), + ) + })?; + + Ok(()) + } + + /// Execute element-wise addition using pointers + pub fn add_ptr( + &self, + a_ptr: *const u8, + b_ptr: *const u8, + c_ptr: *mut u8, + n: u32, + ) -> Result<()> { + // Get buffers from the backend's buffer tracking system + let a_buffer_id = a_ptr as usize; + let b_buffer_id = b_ptr as usize; + let c_buffer_id = c_ptr as usize; + + let buffers = self.backend.buffers.read(); + + if let (Some(a_buf), Some(b_buf), Some(c_buf)) = ( + buffers.get(&a_buffer_id), + buffers.get(&b_buffer_id), + buffers.get(&c_buffer_id), + ) { + self.add(&a_buf.buffer, &b_buf.buffer, &c_buf.buffer, n) + } else { + Err(crate::error::MinitensorError::memory_error( + "OpenCL buffer not found for operation", + )) + } + } + + /// Execute element-wise multiplication using pointers + pub fn mul_ptr( + &self, + a_ptr: *const u8, + b_ptr: *const u8, + c_ptr: *mut u8, + n: u32, + ) -> Result<()> { + let a_buffer_id = a_ptr as usize; + let b_buffer_id = b_ptr as usize; + let c_buffer_id = c_ptr as usize; + + let buffers = self.backend.buffers.read(); + + if let (Some(a_buf), Some(b_buf), Some(c_buf)) = ( + buffers.get(&a_buffer_id), + buffers.get(&b_buffer_id), + buffers.get(&c_buffer_id), + ) { + self.mul(&a_buf.buffer, &b_buf.buffer, &c_buf.buffer, n) + } else { + Err(crate::error::MinitensorError::memory_error( + "OpenCL buffer not found for operation", + )) + } + } + + /// Execute matrix multiplication using pointers + pub fn matmul_ptr( + &self, + a_ptr: *const u8, + b_ptr: *const u8, + c_ptr: *mut u8, + m: u32, + n: u32, + k: u32, + ) -> Result<()> { + let a_buffer_id = a_ptr as usize; + let b_buffer_id = b_ptr as usize; + let c_buffer_id = c_ptr as usize; + + let buffers = self.backend.buffers.read(); + + if let (Some(a_buf), Some(b_buf), Some(c_buf)) = ( + buffers.get(&a_buffer_id), + buffers.get(&b_buffer_id), + buffers.get(&c_buffer_id), + ) { + self.matmul(&a_buf.buffer, &b_buf.buffer, &c_buf.buffer, m, n, k) + } else { + Err(crate::error::MinitensorError::memory_error( + "OpenCL buffer not found for operation", + )) + } + } + + /// Execute ReLU activation using pointers + pub fn relu_ptr(&self, input_ptr: *const u8, output_ptr: *mut u8, n: u32) -> Result<()> { + let input_buffer_id = input_ptr as usize; + let output_buffer_id = output_ptr as usize; + + let buffers = self.backend.buffers.read(); + + if let (Some(input_buf), Some(output_buf)) = ( + buffers.get(&input_buffer_id), + buffers.get(&output_buffer_id), + ) { + self.relu(&input_buf.buffer, &output_buf.buffer, n) + } else { + Err(crate::error::MinitensorError::memory_error( + "OpenCL buffer not found for operation", + )) + } + } + + /// Execute Sigmoid activation using pointers + pub fn sigmoid_ptr(&self, input_ptr: *const u8, output_ptr: *mut u8, n: u32) -> Result<()> { + let input_buffer_id = input_ptr as usize; + let output_buffer_id = output_ptr as usize; + + let buffers = self.backend.buffers.read(); + + if let (Some(input_buf), Some(output_buf)) = ( + buffers.get(&input_buffer_id), + buffers.get(&output_buffer_id), + ) { + self.sigmoid(&input_buf.buffer, &output_buf.buffer, n) + } else { + Err(crate::error::MinitensorError::memory_error( + "OpenCL buffer not found for operation", + )) + } + } +} + +#[cfg(test)] +#[path = "../integration_test.rs"] +mod integration_test; + +#[cfg(test)] +mod tests { + use super::*; + use opencl3::memory::{CL_MEM_READ_ONLY, CL_MEM_READ_WRITE, CL_MEM_WRITE_ONLY}; + + #[test] + fn test_opencl_availability() { + // This test will only pass if OpenCL is available + if OpenCLBackend::is_available() { + let backend = OpenCLBackend::initialize().unwrap(); + assert!(backend.device().is_gpu()); + } + } + + #[test] + fn test_opencl_buffer_operations() { + if !OpenCLBackend::is_available() { + return; // Skip test if OpenCL not available + } + + let backend = OpenCLBackend::initialize().unwrap(); + + // Test buffer creation and data transfer + let data = vec![1.0f32, 2.0, 3.0, 4.0, 5.0]; + let buffer = backend + .create_buffer_with_data(&data, CL_MEM_READ_WRITE) + .unwrap(); + + let mut result = vec![0.0f32; 5]; + backend.read_buffer(&buffer, &mut result).unwrap(); + + assert_eq!(data, result); + } + + #[test] + fn test_opencl_operations() { + if !OpenCLBackend::is_available() { + return; // Skip test if OpenCL not available + } + + let backend = Arc::new(OpenCLBackend::initialize().unwrap()); + let ops = OpenCLOps::new(backend.clone()).unwrap(); + + // Test addition + let a_data = vec![1.0f32, 2.0, 3.0, 4.0]; + let b_data = vec![5.0f32, 6.0, 7.0, 8.0]; + + let a_buffer = backend + .create_buffer_with_data(&a_data, CL_MEM_READ_ONLY) + .unwrap(); + let b_buffer = backend + .create_buffer_with_data(&b_data, CL_MEM_READ_ONLY) + .unwrap(); + let c_buffer = backend.create_buffer(4, CL_MEM_WRITE_ONLY).unwrap(); + + ops.add(&a_buffer, &b_buffer, &c_buffer, 4).unwrap(); + + let mut result = vec![0.0f32; 4]; + backend.read_buffer(&c_buffer, &mut result).unwrap(); + + let expected = [6.0f32, 8.0, 10.0, 12.0]; + for (r, e) in result.iter().zip(expected.iter()) { + assert!((r - e).abs() < 1e-6); + } + } +} diff --git a/engine/src/custom_ops.rs b/engine/src/custom_ops.rs index 75560c68..d6a64f29 100644 --- a/engine/src/custom_ops.rs +++ b/engine/src/custom_ops.rs @@ -165,10 +165,8 @@ impl CustomOpRegistry { // Set up gradient tracking if any input requires gradients let requires_grad = inputs.iter().any(|t| t.requires_grad()); - if requires_grad { - if let Some(grad_fn) = op.create_gradient_function(inputs, &output) { - add_to_graph(&output, Some(grad_fn))?; - } + if requires_grad && let Some(grad_fn) = op.create_gradient_function(inputs, &output) { + add_to_graph(&output, Some(grad_fn))?; } Ok(output) @@ -454,7 +452,7 @@ impl CustomOp for BuiltCustomOp { "No input devices provided", )) } else { - Ok(input_devices[0].clone()) + Ok(*input_devices[0]) } } } diff --git a/engine/src/debug.rs b/engine/src/debug.rs index 53b232b0..c1baeb5f 100644 --- a/engine/src/debug.rs +++ b/engine/src/debug.rs @@ -430,27 +430,24 @@ impl TensorDebugger { } // Check for very large values - if let Some(max_val) = tensor.max_value() { - if max_val > 1e6 { - issues.push(format!( - "⚠️ Tensor has very large values (max: {:.2e})", - max_val - )); - } + if let Some(max_val) = tensor.max_value() + && max_val > 1e6 + { + issues.push(format!( + "⚠️ Tensor has very large values (max: {:.2e})", + max_val + )); } // Check for very small gradients - if tensor.requires_grad() { - if let Some(grad) = tensor.grad() { - if let Some(grad_max) = grad.max_value() { - if grad_max < 1e-8 { - issues.push( - "⚠️ Gradients are very small, may indicate vanishing gradient problem" - .to_string(), - ); - } - } - } + if tensor.requires_grad() + && let Some(grad) = tensor.grad() + && let Some(grad_max) = grad.max_value() + && grad_max < 1e-8 + { + issues.push( + "⚠️ Gradients are very small, may indicate vanishing gradient problem".to_string(), + ); } // Check memory usage diff --git a/engine/src/device.rs b/engine/src/device.rs index 246a1f4c..75921544 100644 --- a/engine/src/device.rs +++ b/engine/src/device.rs @@ -115,8 +115,21 @@ impl Device { } } - /// Parse device from string + /// Parse device from string. + /// + /// Inherent convenience wrapper kept for API compatibility; the canonical + /// implementation is the [`std::str::FromStr`] impl, so `"cuda:1".parse()` + /// also works. + #[allow(clippy::should_implement_trait)] pub fn from_str(device_str: &str) -> Result { + device_str.parse() + } +} + +impl std::str::FromStr for Device { + type Err = String; + + fn from_str(device_str: &str) -> Result { match device_str.to_lowercase().as_str() { "cpu" => Ok(Self::cpu()), "cuda" => Ok(Self::cuda(Some(0))), diff --git a/engine/src/error.rs b/engine/src/error.rs index bf6e430d..4253ddc0 100644 --- a/engine/src/error.rs +++ b/engine/src/error.rs @@ -449,13 +449,9 @@ impl MinitensorError { let suggestion = match (expected_dims, actual_dims) { (Some(expected), Some(actual)) => { if expected > actual { - Some(format!( - "Use .unsqueeze() to add dimensions or .view() to reshape" - )) + Some("Use .unsqueeze() to add dimensions or .view() to reshape".to_string()) } else { - Some(format!( - "Use .squeeze() to remove dimensions or .view() to reshape" - )) + Some("Use .squeeze() to remove dimensions or .view() to reshape".to_string()) } } _ => Some("Check tensor dimensions and use reshape operations if needed".to_string()), diff --git a/engine/src/hardware/cpu.rs b/engine/src/hardware/cpu.rs index 15a1d4a9..618f5942 100644 --- a/engine/src/hardware/cpu.rs +++ b/engine/src/hardware/cpu.rs @@ -222,61 +222,57 @@ impl CpuInfo { // L1 data cache if let Ok(size_str) = std::fs::read_to_string("/sys/devices/system/cpu/cpu0/cache/index0/size") + && let Some(size) = Self::parse_cache_size(&size_str) { - if let Some(size) = Self::parse_cache_size(&size_str) { - cache_levels.push(CacheLevel { - level: 1, - cache_type: CacheType::Data, - size, - line_size: 64, // Common default - associativity: 8, // Common default - }); - } + cache_levels.push(CacheLevel { + level: 1, + cache_type: CacheType::Data, + size, + line_size: 64, // Common default + associativity: 8, // Common default + }); } // L1 instruction cache if let Ok(size_str) = std::fs::read_to_string("/sys/devices/system/cpu/cpu0/cache/index1/size") + && let Some(size) = Self::parse_cache_size(&size_str) { - if let Some(size) = Self::parse_cache_size(&size_str) { - cache_levels.push(CacheLevel { - level: 1, - cache_type: CacheType::Instruction, - size, - line_size: 64, - associativity: 8, - }); - } + cache_levels.push(CacheLevel { + level: 1, + cache_type: CacheType::Instruction, + size, + line_size: 64, + associativity: 8, + }); } // L2 cache if let Ok(size_str) = std::fs::read_to_string("/sys/devices/system/cpu/cpu0/cache/index2/size") + && let Some(size) = Self::parse_cache_size(&size_str) { - if let Some(size) = Self::parse_cache_size(&size_str) { - cache_levels.push(CacheLevel { - level: 2, - cache_type: CacheType::Unified, - size, - line_size: 64, - associativity: 8, - }); - } + cache_levels.push(CacheLevel { + level: 2, + cache_type: CacheType::Unified, + size, + line_size: 64, + associativity: 8, + }); } // L3 cache if let Ok(size_str) = std::fs::read_to_string("/sys/devices/system/cpu/cpu0/cache/index3/size") + && let Some(size) = Self::parse_cache_size(&size_str) { - if let Some(size) = Self::parse_cache_size(&size_str) { - cache_levels.push(CacheLevel { - level: 3, - cache_type: CacheType::Unified, - size, - line_size: 64, - associativity: 16, - }); - } + cache_levels.push(CacheLevel { + level: 3, + cache_type: CacheType::Unified, + size, + line_size: 64, + associativity: 16, + }); } } diff --git a/engine/src/hardware/gpu.rs b/engine/src/hardware/gpu.rs index 83e06e09..a1b6bcd4 100644 --- a/engine/src/hardware/gpu.rs +++ b/engine/src/hardware/gpu.rs @@ -262,7 +262,7 @@ impl GpuDevice { pub fn memory_bandwidth(&self) -> f64 { self.capabilities .memory_bandwidth - .unwrap_or_else(|| match self.device_type { + .unwrap_or(match self.device_type { DeviceType::Cuda => { if self.memory_size > 32 * 1024 * 1024 * 1024 { 900.0 diff --git a/engine/src/hardware/memory.rs b/engine/src/hardware/memory.rs index ea53e7d9..ff011de9 100644 --- a/engine/src/hardware/memory.rs +++ b/engine/src/hardware/memory.rs @@ -58,12 +58,11 @@ impl MemoryInfo { { if let Ok(content) = std::fs::read_to_string("/proc/meminfo") { for line in content.lines() { - if line.starts_with("MemTotal:") { - if let Some(kb_str) = line.split_whitespace().nth(1) { - if let Ok(kb) = kb_str.parse::() { - return kb * 1024; // Convert KB to bytes - } - } + if line.starts_with("MemTotal:") + && let Some(kb_str) = line.split_whitespace().nth(1) + && let Ok(kb) = kb_str.parse::() + { + return kb * 1024; // Convert KB to bytes } } } @@ -100,12 +99,11 @@ impl MemoryInfo { { if let Ok(content) = std::fs::read_to_string("/proc/meminfo") { for line in content.lines() { - if line.starts_with("MemAvailable:") { - if let Some(kb_str) = line.split_whitespace().nth(1) { - if let Ok(kb) = kb_str.parse::() { - return kb * 1024; // Convert KB to bytes - } - } + if line.starts_with("MemAvailable:") + && let Some(kb_str) = line.split_whitespace().nth(1) + && let Ok(kb) = kb_str.parse::() + { + return kb * 1024; // Convert KB to bytes } } } @@ -120,12 +118,11 @@ impl MemoryInfo { { if let Ok(content) = std::fs::read_to_string("/proc/meminfo") { for line in content.lines() { - if line.starts_with("SwapTotal:") { - if let Some(kb_str) = line.split_whitespace().nth(1) { - if let Ok(kb) = kb_str.parse::() { - return kb * 1024; // Convert KB to bytes - } - } + if line.starts_with("SwapTotal:") + && let Some(kb_str) = line.split_whitespace().nth(1) + && let Ok(kb) = kb_str.parse::() + { + return kb * 1024; // Convert KB to bytes } } } @@ -139,12 +136,11 @@ impl MemoryInfo { { if let Ok(content) = std::fs::read_to_string("/proc/meminfo") { for line in content.lines() { - if line.starts_with("SwapFree:") { - if let Some(kb_str) = line.split_whitespace().nth(1) { - if let Ok(kb) = kb_str.parse::() { - return kb * 1024; // Convert KB to bytes - } - } + if line.starts_with("SwapFree:") + && let Some(kb_str) = line.split_whitespace().nth(1) + && let Ok(kb) = kb_str.parse::() + { + return kb * 1024; // Convert KB to bytes } } } @@ -176,29 +172,25 @@ impl MemoryInfo { if let (Ok(size_str), Ok(line_size_str)) = ( std::fs::read_to_string(format!("{}/size", cache_path)), std::fs::read_to_string(format!("{}/coherency_line_size", cache_path)), + ) && let (Some(size), Ok(line_size)) = ( + Self::parse_cache_size(&size_str), + line_size_str.trim().parse::(), ) { - if let (Some(size), Ok(line_size)) = ( - Self::parse_cache_size(&size_str), - line_size_str.trim().parse::(), - ) { - // Try to read associativity - let associativity = std::fs::read_to_string(format!( - "{}/ways_of_associativity", - cache_path - )) - .ok() - .and_then(|s| s.trim().parse::().ok()) - .unwrap_or(8); // Default associativity - - cache_info.push(CacheInfo { - level: level as u8, - size, - line_size, - associativity, - latency_cycles: Self::estimate_cache_latency(level as u8), - bandwidth: None, // Will be benchmarked separately if needed - }); - } + // Try to read associativity + let associativity = + std::fs::read_to_string(format!("{}/ways_of_associativity", cache_path)) + .ok() + .and_then(|s| s.trim().parse::().ok()) + .unwrap_or(8); // Default associativity + + cache_info.push(CacheInfo { + level: level as u8, + size, + line_size, + associativity, + latency_cycles: Self::estimate_cache_latency(level as u8), + bandwidth: None, // Will be benchmarked separately if needed + }); } } } @@ -370,7 +362,7 @@ impl MemoryBandwidth { } let avg_time = total_time.as_secs_f64() / iterations as f64; - let accesses = (size + 63) / 64; + let accesses = size.div_ceil(64); let bytes_per_second = (accesses * 64) as f64 / avg_time; bytes_per_second / (1024.0 * 1024.0 * 1024.0) } @@ -393,7 +385,7 @@ impl MemoryBandwidth { } let avg_time = total_time.as_secs_f64() / iterations as f64; - let accesses = (size + 63) / 64; + let accesses = size.div_ceil(64); let bytes_per_second = (accesses * 64) as f64 / avg_time; bytes_per_second / (1024.0 * 1024.0 * 1024.0) } @@ -463,7 +455,7 @@ impl CacheInfo { /// Estimate access latency for this cache level #[inline] pub fn access_latency(&self) -> Duration { - let cycles = self.latency_cycles.unwrap_or_else(|| match self.level { + let cycles = self.latency_cycles.unwrap_or(match self.level { 1 => 4, 2 => 12, 3 => 40, diff --git a/engine/src/hardware/profiler.rs b/engine/src/hardware/profiler.rs index 7b004268..44255bed 100644 --- a/engine/src/hardware/profiler.rs +++ b/engine/src/hardware/profiler.rs @@ -498,12 +498,11 @@ impl SystemInfo { fn get_uptime() -> Duration { #[cfg(target_os = "linux")] { - if let Ok(content) = std::fs::read_to_string("/proc/uptime") { - if let Some(uptime_str) = content.split_whitespace().next() { - if let Ok(uptime_secs) = uptime_str.parse::() { - return Duration::from_secs_f64(uptime_secs); - } - } + if let Ok(content) = std::fs::read_to_string("/proc/uptime") + && let Some(uptime_str) = content.split_whitespace().next() + && let Ok(uptime_secs) = uptime_str.parse::() + { + return Duration::from_secs_f64(uptime_secs); } } @@ -515,14 +514,14 @@ impl SystemInfo { { if let Ok(content) = std::fs::read_to_string("/proc/loadavg") { let parts: Vec<&str> = content.split_whitespace().collect(); - if parts.len() >= 3 { - if let (Ok(load1), Ok(load5), Ok(load15)) = ( + if parts.len() >= 3 + && let (Ok(load1), Ok(load5), Ok(load15)) = ( parts[0].parse::(), parts[1].parse::(), parts[2].parse::(), - ) { - return Some((load1, load5, load15)); - } + ) + { + return Some((load1, load5, load15)); } } } @@ -548,10 +547,10 @@ impl ThermalInfo { // Try to read from thermal zones for i in 0..10 { let temp_path = format!("/sys/class/thermal/thermal_zone{}/temp", i); - if let Ok(temp_str) = std::fs::read_to_string(&temp_path) { - if let Ok(temp_millicelsius) = temp_str.trim().parse::() { - return Some(temp_millicelsius as f64 / 1000.0); - } + if let Ok(temp_str) = std::fs::read_to_string(&temp_path) + && let Ok(temp_millicelsius) = temp_str.trim().parse::() + { + return Some(temp_millicelsius as f64 / 1000.0); } } } diff --git a/engine/src/lib.rs b/engine/src/lib.rs index ba14eaa7..4e001543 100644 --- a/engine/src/lib.rs +++ b/engine/src/lib.rs @@ -4,7 +4,13 @@ // This source code is licensed under the Apache-style license found in the // LICENSE file in the root directory of this source tree. -#![allow(clippy::all)] +// Two style lints are allowed crate-wide because they fight numeric-kernel +// idioms: `needless_range_loop` (stencil loops whose index feeds stride +// arithmetic, not just slice access) and `too_many_arguments` (kernel helpers +// taking many scalar parameters). Every other clippy lint is enforced; CI runs +// `cargo clippy -- -D warnings`. +#![allow(clippy::needless_range_loop)] +#![allow(clippy::too_many_arguments)] pub mod autograd; pub mod backends; diff --git a/engine/src/memory/allocator.rs b/engine/src/memory/allocator.rs index be29ba4b..ac27fbc9 100644 --- a/engine/src/memory/allocator.rs +++ b/engine/src/memory/allocator.rs @@ -172,7 +172,7 @@ impl OpenCLAllocator { #[cfg(feature = "opencl")] impl Allocator for OpenCLAllocator { - fn allocate(&mut self, size: usize) -> Result<*mut u8> { + fn allocate(&mut self, _size: usize) -> Result<*mut u8> { Err(crate::error::MinitensorError::backend_error( "OpenCL", "OpenCL allocator not yet implemented", @@ -218,6 +218,10 @@ impl Allocator for CpuAllocator { } } + // The `Allocator` trait keeps `deallocate` a safe fn for API-compat; the + // caller contract (pointer must come from `allocate` with the same size) + // is documented on the trait. Same pattern as `backends::cpu`. + #[allow(clippy::not_unsafe_ptr_arg_deref)] #[inline(always)] fn deallocate(&mut self, ptr: *mut u8, size: usize) -> Result<()> { if ptr.is_null() || size == 0 { diff --git a/engine/src/memory/pool.rs b/engine/src/memory/pool.rs index 159612e0..78df3d55 100644 --- a/engine/src/memory/pool.rs +++ b/engine/src/memory/pool.rs @@ -98,6 +98,11 @@ impl MemoryPool { } /// Return memory to the pool for future reuse. + /// + /// The pointer must have been produced by [`Self::allocate`] with the + /// same `size`; the fn stays safe for API-compat with the allocator + /// stack (see `Allocator::deallocate`). + #[allow(clippy::not_unsafe_ptr_arg_deref)] #[inline] pub fn deallocate(&mut self, ptr: *mut u8, size: usize) -> Result<()> { if ptr.is_null() || size == 0 { diff --git a/engine/src/nn/activation.rs b/engine/src/nn/activation.rs index bec7b27e..9e339e69 100644 --- a/engine/src/nn/activation.rs +++ b/engine/src/nn/activation.rs @@ -49,12 +49,12 @@ fn scalar_tensor( } } -fn cached_scalar<'a>( - cache: &'a mut Option, +fn cached_scalar( + cache: &mut Option, value: f64, dtype: DataType, device: Device, -) -> Result<&'a Tensor> { +) -> Result<&Tensor> { let needs_update = match cache { Some(t) => t.dtype() != dtype || t.device() != device, None => true, @@ -181,8 +181,9 @@ impl Softmax { /// Create a new Softmax activation layer /// /// # Arguments - /// * `dim` - A dimension along which Softmax will be computed (so every slice - /// along dim will sum to 1). Default: None (applies to the last dimension) + /// * `dim` - A dimension along which Softmax will be computed (so every + /// slice along dim will sum to 1). Default: None (applies to the last + /// dimension) pub fn new(dim: Option) -> Self { Self { dim } } @@ -480,7 +481,7 @@ mod tests { ); let output = elu.forward(&input).unwrap(); let out_slice = output.data().as_f32_slice().unwrap(); - let expected = vec![(-1f32).exp() - 1.0, 0.0, 1.0]; + let expected = [(-1f32).exp() - 1.0, 0.0, 1.0]; for (o, e) in out_slice.iter().zip(expected.iter()) { assert!((o - e).abs() < 1e-4); } diff --git a/engine/src/nn/init.rs b/engine/src/nn/init.rs index 4e022e26..f8f12c02 100644 --- a/engine/src/nn/init.rs +++ b/engine/src/nn/init.rs @@ -87,7 +87,7 @@ pub fn init_constant( TensorData::from_vec_f32(vec, device) } DataType::Float64 => { - let vec = vec![value as f64; numel]; + let vec = vec![value; numel]; TensorData::from_vec_f64(vec, device) } DataType::Int32 => { @@ -126,65 +126,40 @@ pub fn init_uniform( DataType::Float32 => { let dist = Uniform::new(a as f32, b as f32).unwrap(); let mut vec = Vec::with_capacity(numel); - unsafe { - vec.set_len(numel); - } random::with_rng(|rng| { - for v in &mut vec { - *v = dist.sample(rng); - } + vec.extend((0..numel).map(|_| dist.sample(rng))); }); TensorData::from_vec_f32(vec, device) } DataType::Float64 => { let dist = Uniform::new(a, b).unwrap(); let mut vec = Vec::with_capacity(numel); - unsafe { - vec.set_len(numel); - } random::with_rng(|rng| { - for v in &mut vec { - *v = dist.sample(rng); - } + vec.extend((0..numel).map(|_| dist.sample(rng))); }); TensorData::from_vec_f64(vec, device) } DataType::Int32 => { let dist = Uniform::new(a as i32, b as i32).unwrap(); let mut vec = Vec::with_capacity(numel); - unsafe { - vec.set_len(numel); - } random::with_rng(|rng| { - for v in &mut vec { - *v = dist.sample(rng); - } + vec.extend((0..numel).map(|_| dist.sample(rng))); }); TensorData::from_vec_i32(vec, device) } DataType::Int64 => { let dist = Uniform::new(a as i64, b as i64).unwrap(); let mut vec = Vec::with_capacity(numel); - unsafe { - vec.set_len(numel); - } random::with_rng(|rng| { - for v in &mut vec { - *v = dist.sample(rng); - } + vec.extend((0..numel).map(|_| dist.sample(rng))); }); TensorData::from_vec_i64(vec, device) } DataType::Bool => { let dist = Uniform::new(0.0, 1.0).unwrap(); let mut vec = Vec::with_capacity(numel); - unsafe { - vec.set_len(numel); - } random::with_rng(|rng| { - for v in &mut vec { - *v = dist.sample(rng) > 0.5; - } + vec.extend((0..numel).map(|_| dist.sample(rng) > 0.5)); }); TensorData::from_vec_bool(vec, device) } @@ -212,65 +187,40 @@ pub fn init_normal( DataType::Float32 => { let dist = Normal::new(mean as f32, std as f32).unwrap(); let mut vec = Vec::with_capacity(numel); - unsafe { - vec.set_len(numel); - } random::with_rng(|rng| { - for v in &mut vec { - *v = dist.sample(rng); - } + vec.extend((0..numel).map(|_| dist.sample(rng))); }); TensorData::from_vec_f32(vec, device) } DataType::Float64 => { let dist = Normal::new(mean, std).unwrap(); let mut vec = Vec::with_capacity(numel); - unsafe { - vec.set_len(numel); - } random::with_rng(|rng| { - for v in &mut vec { - *v = dist.sample(rng); - } + vec.extend((0..numel).map(|_| dist.sample(rng))); }); TensorData::from_vec_f64(vec, device) } DataType::Int32 => { let dist = Normal::new(mean, std).unwrap(); let mut vec = Vec::with_capacity(numel); - unsafe { - vec.set_len(numel); - } random::with_rng(|rng| { - for v in &mut vec { - *v = dist.sample(rng).round() as i32; - } + vec.extend((0..numel).map(|_| dist.sample(rng).round() as i32)); }); TensorData::from_vec_i32(vec, device) } DataType::Int64 => { let dist = Normal::new(mean, std).unwrap(); let mut vec = Vec::with_capacity(numel); - unsafe { - vec.set_len(numel); - } random::with_rng(|rng| { - for v in &mut vec { - *v = dist.sample(rng).round() as i64; - } + vec.extend((0..numel).map(|_| dist.sample(rng).round() as i64)); }); TensorData::from_vec_i64(vec, device) } DataType::Bool => { let dist = Normal::new(mean, std).unwrap(); let mut vec = Vec::with_capacity(numel); - unsafe { - vec.set_len(numel); - } random::with_rng(|rng| { - for v in &mut vec { - *v = dist.sample(rng) > 0.0; - } + vec.extend((0..numel).map(|_| dist.sample(rng) > 0.0)); }); TensorData::from_vec_bool(vec, device) } @@ -313,7 +263,7 @@ pub fn truncated_normal_init( )); } - if !(upper > lower) { + if upper <= lower { return Err(MinitensorError::invalid_argument( "truncated_normal requires upper bound to be greater than lower bound", )); @@ -327,7 +277,7 @@ pub fn truncated_normal_init( let lower_cdf = normal.cdf(lower); let upper_cdf = normal.cdf(upper); - if !(upper_cdf > lower_cdf) { + if upper_cdf <= lower_cdf { return Err(MinitensorError::invalid_argument( "truncated_normal bounds must span non-zero probability mass", )); @@ -337,40 +287,33 @@ pub fn truncated_normal_init( let data = match dtype { DataType::Float32 => { let mut vec = Vec::with_capacity(numel); - unsafe { - vec.set_len(numel); - } random::with_rng(|rng| { let uniform = Uniform::new(lower_cdf, upper_cdf).unwrap(); - for value in &mut vec { + vec.extend((0..numel).map(|_| { let mut sample_cdf = uniform.sample(rng); if sample_cdf <= 0.0 { sample_cdf = f64::EPSILON; } else if sample_cdf >= 1.0 { sample_cdf = 1.0 - f64::EPSILON; } - let sample = normal.inverse_cdf(sample_cdf); - *value = sample as f32; - } + normal.inverse_cdf(sample_cdf) as f32 + })); }); TensorData::from_vec_f32(vec, device) } DataType::Float64 => { let mut vec = Vec::with_capacity(numel); - unsafe { - vec.set_len(numel); - } random::with_rng(|rng| { let uniform = Uniform::new(lower_cdf, upper_cdf).unwrap(); - for value in &mut vec { + vec.extend((0..numel).map(|_| { let mut sample_cdf = uniform.sample(rng); if sample_cdf <= 0.0 { sample_cdf = f64::EPSILON; } else if sample_cdf >= 1.0 { sample_cdf = 1.0 - f64::EPSILON; } - *value = normal.inverse_cdf(sample_cdf); - } + normal.inverse_cdf(sample_cdf) + })); }); TensorData::from_vec_f64(vec, device) } @@ -579,7 +522,7 @@ mod tests { .unwrap(); let slice = tensor.data().as_f32_slice().unwrap(); for &v in slice { - assert!(v >= -0.5 && v <= 0.5); + assert!((-0.5..=0.5).contains(&v)); } } diff --git a/engine/src/nn/normalization.rs b/engine/src/nn/normalization.rs index 7ba7063e..3d57d948 100644 --- a/engine/src/nn/normalization.rs +++ b/engine/src/nn/normalization.rs @@ -352,8 +352,8 @@ impl BatchNorm2d { Ok(Self { weight, bias, - running_mean: running_mean, - running_var: running_var, + running_mean, + running_var, num_features, _eps: eps, _momentum: momentum, diff --git a/engine/src/nn/sequential.rs b/engine/src/nn/sequential.rs index e3e1c448..44b5b71c 100644 --- a/engine/src/nn/sequential.rs +++ b/engine/src/nn/sequential.rs @@ -171,6 +171,8 @@ impl SequentialBuilder { } /// Add a layer to the builder + // Builder-pattern `add`, not arithmetic; renaming would break the public API. + #[allow(clippy::should_implement_trait)] pub fn add(mut self, layer: Box) -> Self { self.layers.push(layer); self diff --git a/engine/src/operations/activation.rs b/engine/src/operations/activation.rs index 1c02a975..f239f702 100644 --- a/engine/src/operations/activation.rs +++ b/engine/src/operations/activation.rs @@ -4,9 +4,22 @@ // This source code is licensed under the Apache-style license found in the // LICENSE file in the root directory of this source tree. -include!("activation/elementwise.rs"); -include!("activation/trigonometry.rs"); -include!("activation/hyperbolic.rs"); -include!("activation/softmax.rs"); -include!("activation/power.rs"); -include!("activation/advanced.rs"); +#[path = "activation/advanced.rs"] +mod advanced_impl; +#[path = "activation/elementwise.rs"] +mod elementwise_impl; +#[path = "activation/hyperbolic.rs"] +mod hyperbolic_impl; +#[path = "activation/power.rs"] +mod power_impl; +#[path = "activation/softmax.rs"] +mod softmax_impl; +#[path = "activation/trigonometry.rs"] +mod trigonometry_impl; + +pub(crate) use self::advanced_impl::*; +pub use self::elementwise_impl::*; +pub use self::hyperbolic_impl::*; +pub use self::power_impl::*; +pub(crate) use self::softmax_impl::*; +pub use self::trigonometry_impl::*; diff --git a/engine/src/operations/activation/advanced.rs b/engine/src/operations/activation/advanced.rs index cc953e91..1b79dfea 100644 --- a/engine/src/operations/activation/advanced.rs +++ b/engine/src/operations/activation/advanced.rs @@ -1,593 +1,601 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -fn ceil_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - unary_apply(input_data, output_slice, f64::ceil); - Ok(()) -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::{ - autograd, - device::Device, - tensor::{Shape, Tensor, TensorData}, - }; - - fn create_test_tensor_f32(data: Vec, shape: Vec, requires_grad: bool) -> Tensor { - let shape_obj = Shape::new(shape); - let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Float32); - - if let Some(slice) = tensor_data.as_f32_slice_mut() { - slice.copy_from_slice(&data); - } - - Tensor::new( - Arc::new(tensor_data), - shape_obj, - DataType::Float32, - Device::cpu(), - requires_grad, - ) - } - - fn create_test_tensor_bool(data: Vec, shape: Vec) -> Tensor { - let shape_obj = Shape::new(shape); - let tensor_data = TensorData::from_vec_bool(data, Device::cpu()); - Tensor::new( - Arc::new(tensor_data), - shape_obj, - DataType::Bool, - Device::cpu(), - false, - ) - } - - #[test] - fn test_nan_to_num_custom_replacements() { - let tensor = create_test_tensor_f32( - vec![f32::NAN, f32::INFINITY, f32::NEG_INFINITY, 4.0], - vec![4], - false, - ); - let result = nan_to_num(&tensor, -1.0, Some(9.0), Some(-9.0)).unwrap(); - let data = result.data().as_f32_slice().unwrap(); - assert_eq!(data, &[-1.0, 9.0, -9.0, 4.0]); - } - - #[test] - fn test_nan_to_num_backward_masks_replaced_entries() { - let tensor = create_test_tensor_f32( - vec![f32::NAN, f32::INFINITY, f32::NEG_INFINITY, -2.0, 3.0], - vec![5], - true, - ); - let result = nan_to_num(&tensor, 0.0, Some(10.0), Some(-10.0)).unwrap(); - let summed = crate::operations::reduction::sum(&result, None, false).unwrap(); - let grads = autograd::backward(&summed, None).unwrap(); - let grad = grads.get(&tensor.id()).unwrap(); - assert_eq!( - grad.data().as_f32_slice().unwrap(), - &[0.0, 0.0, 0.0, 1.0, 1.0] - ); - } - - #[test] - fn test_nan_to_num_integer_is_unchanged() { - let shape = Shape::new(vec![3]); - let data = TensorData::from_vec_i32(vec![1, -2, 3], Device::cpu()); - let tensor = Tensor::new(Arc::new(data), shape, DataType::Int32, Device::cpu(), false); - let result = nan_to_num(&tensor, 99.0, None, None).unwrap(); - assert_eq!(result.data().as_i32_slice().unwrap(), &[1, -2, 3]); - } - - #[test] - fn test_exp() { - let tensor = create_test_tensor_f32(vec![0.0, 1.0, 2.0], vec![3], false); - let result = exp(&tensor).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - - assert!((result_data[0] - 1.0).abs() < 1e-6); - assert!((result_data[1] - std::f32::consts::E).abs() < 1e-6); - assert!((result_data[2] - (std::f32::consts::E * std::f32::consts::E)).abs() < 1e-5); - } - - #[test] - fn test_exp_invalid_dtype() { - let shape = Shape::new(vec![3]); - let data = TensorData::from_vec_i32(vec![1, 2, 3], Device::cpu()); - let tensor = Tensor::new(Arc::new(data), shape, DataType::Int32, Device::cpu(), false); - assert!(exp(&tensor).is_err()); - } - - #[test] - fn test_log() { - let tensor = create_test_tensor_f32( - vec![ - 1.0, - std::f32::consts::E, - std::f32::consts::E * std::f32::consts::E, - ], - vec![3], - false, - ); - let result = log(&tensor).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - - assert!((result_data[0] - 0.0).abs() < 1e-6); - assert!((result_data[1] - 1.0).abs() < 1e-6); - assert!((result_data[2] - 2.0).abs() < 1e-5); - } - - #[test] - fn test_sin() { - let tensor = create_test_tensor_f32( - vec![0.0, std::f32::consts::PI / 2.0, std::f32::consts::PI], - vec![3], - false, - ); - let result = sin(&tensor).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - - assert!((result_data[0] - 0.0).abs() < 1e-6); - assert!((result_data[1] - 1.0).abs() < 1e-6); - assert!(result_data[2].abs() < 1e-6); // sin(π) ≈ 0 - } - - #[test] - fn test_cos() { - let tensor = create_test_tensor_f32( - vec![0.0, std::f32::consts::PI / 2.0, std::f32::consts::PI], - vec![3], - false, - ); - let result = cos(&tensor).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - - assert!((result_data[0] - 1.0).abs() < 1e-6); - assert!(result_data[1].abs() < 1e-6); // cos(π/2) ≈ 0 - assert!((result_data[2] + 1.0).abs() < 1e-6); // cos(π) ≈ -1 - } - - #[test] - fn test_tan() { - let tensor = create_test_tensor_f32( - vec![0.0, std::f32::consts::PI / 4.0, -std::f32::consts::PI / 4.0], - vec![3], - false, - ); - let result = tan(&tensor).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - - assert!((result_data[0] - 0.0).abs() < 1e-6); - assert!((result_data[1] - (std::f32::consts::PI / 4.0).tan()).abs() < 1e-6); - assert!((result_data[2] - (-std::f32::consts::PI / 4.0).tan()).abs() < 1e-6); - } - - #[test] - fn test_tanh() { - let tensor = create_test_tensor_f32(vec![0.0, 1.0, -1.0], vec![3], false); - let result = tanh(&tensor).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - - assert!((result_data[0] - 0.0).abs() < 1e-6); - assert!((result_data[1] - 1.0_f32.tanh()).abs() < 1e-6); - assert!((result_data[2] - (-1.0_f32).tanh()).abs() < 1e-6); - } - - #[test] - fn test_sigmoid() { - let tensor = create_test_tensor_f32(vec![0.0, 1.0, -1.0], vec![3], false); - let result = sigmoid(&tensor).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - - assert!((result_data[0] - 0.5).abs() < 1e-6); - assert!((result_data[1] - (1.0 / (1.0 + (-1.0_f32).exp()))).abs() < 1e-6); - assert!((result_data[2] - (1.0 / (1.0 + 1.0_f32.exp()))).abs() < 1e-6); - } - - #[test] - fn test_relu() { - let tensor = create_test_tensor_f32(vec![-2.0, -1.0, 0.0, 1.0, 2.0], vec![5], false); - let result = relu(&tensor).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - - assert_eq!(result_data, &[0.0, 0.0, 0.0, 1.0, 2.0]); - } - - #[test] - fn test_hardshrink_forward() { - let tensor = create_test_tensor_f32(vec![-1.2, -0.2, 0.0, 0.45, 0.9], vec![5], false); - let result = hardshrink(&tensor, 0.3).unwrap(); - let data = result.data().as_f32_slice().unwrap(); - assert_eq!(data, &[-1.2, 0.0, 0.0, 0.45, 0.9]); - } - - #[test] - fn test_hardshrink_backward() { - let tensor = create_test_tensor_f32(vec![-1.2, -0.25, 0.0, 0.35, 0.8], vec![5], true); - let result = hardshrink(&tensor, 0.3).unwrap(); - let ones = Tensor::ones( - result.shape().clone(), - result.dtype(), - result.device(), - false, - ); - let grads = autograd::backward(&result, Some(ones)).unwrap(); - let grad = grads.get(&tensor.id()).unwrap(); - let grad_vals = grad.data().as_f32_slice().unwrap(); - assert_eq!(grad_vals, &[1.0, 0.0, 0.0, 1.0, 1.0]); - } - - #[test] - fn test_hardshrink_invalid_lambda() { - let tensor = create_test_tensor_f32(vec![1.0], vec![1], false); - assert!(hardshrink(&tensor, -0.1).is_err()); - } - - #[test] - fn test_leaky_relu() { - let tensor = create_test_tensor_f32(vec![-2.0, -1.0, 0.0, 1.0, 2.0], vec![5], false); - let result = leaky_relu(&tensor, 0.1).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - - assert_eq!(result_data, &[-0.2, -0.1, 0.0, 1.0, 2.0]); - } - - #[test] - fn test_softmax() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], false); - let result = softmax(&tensor, None).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - - // Check that probabilities sum to 1 - let sum: f32 = result_data.iter().sum(); - assert!((sum - 1.0).abs() < 1e-6); - - // Check that all values are positive - for &val in result_data { - assert!(val > 0.0); - } - - // Check that larger input values produce larger probabilities - assert!(result_data[2] > result_data[1]); - assert!(result_data[1] > result_data[0]); - } - - #[test] - fn test_masked_softmax_respects_mask() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); - let mask = create_test_tensor_bool(vec![true, false, false, true], vec![2, 2]); - let result = masked_softmax(&tensor, &mask, Some(1)).unwrap(); - let data = result.data().as_f32_slice().unwrap(); - assert_eq!(data[0], 0.0); - assert_eq!(data[3], 0.0); - assert!((data[1] - 1.0).abs() < 1e-6); - assert!((data[2] - 1.0).abs() < 1e-6); - } - - #[test] - fn test_masked_softmax_all_masked_returns_zero() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - let mask = create_test_tensor_bool(vec![true, true], vec![2]); - let result = masked_softmax(&tensor, &mask, Some(0)).unwrap(); - let data = result.data().as_f32_slice().unwrap(); - assert!(data.iter().all(|v| *v == 0.0)); - } - - #[test] - fn test_masked_softmax_all_negative_infinity_unmasked() { - let tensor = - create_test_tensor_f32(vec![f32::NEG_INFINITY, f32::NEG_INFINITY], vec![2], false); - let mask = create_test_tensor_bool(vec![false, false], vec![2]); - let result = masked_softmax(&tensor, &mask, Some(0)).unwrap(); - let data = result.data().as_f32_slice().unwrap(); - assert!(data.iter().all(|v| *v == 0.0)); - } - - #[test] - fn test_masked_log_softmax_broadcast_mask() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); - let mask = create_test_tensor_bool(vec![true, false], vec![2, 1]); - let result = masked_log_softmax(&tensor, &mask, Some(1)).unwrap(); - let data = result.data().as_f32_slice().unwrap(); - assert!(data[0].is_infinite() && data[0].is_sign_negative()); - assert!(data[1].is_infinite() && data[1].is_sign_negative()); - let row1_prob0 = data[2].exp(); - let row1_prob1 = data[3].exp(); - assert!((row1_prob0 + row1_prob1 - 1.0).abs() < 1e-6); - assert!(data[3] > data[2]); - } - - #[test] - fn test_masked_log_softmax_all_masked_returns_negative_infinity() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - let mask = create_test_tensor_bool(vec![true, true], vec![2]); - let result = masked_log_softmax(&tensor, &mask, Some(0)).unwrap(); - let data = result.data().as_f32_slice().unwrap(); - assert!(data.iter().all(|v| v.is_infinite() && v.is_sign_negative())); - } - - #[test] - fn test_log_softmax_large_negative_values() { - let tensor = create_test_tensor_f32(vec![-1000.0, 0.0], vec![2], false); - let result = log_softmax(&tensor, None).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - - assert!(result_data[0].is_finite()); - assert!((result_data[0] + 1000.0).abs() < 1e-3); - assert!(result_data[1].abs() < 1e-6); - } - - #[test] - fn test_log_softmax_all_negative_infinity() { - let tensor = - create_test_tensor_f32(vec![f32::NEG_INFINITY, f32::NEG_INFINITY], vec![2], false); - let result = log_softmax(&tensor, None).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - - assert!( - result_data - .iter() - .all(|v| v.is_infinite() && v.is_sign_negative()) - ); - } - - #[test] - fn test_powf_scalar() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], false); - let result = powf(&tensor, 2.0).unwrap(); - let data = result.data().as_f32_slice().unwrap(); - assert_eq!(data, &[1.0, 4.0, 9.0]); - } - - #[test] - fn test_pow_tensor() { - let base = create_test_tensor_f32(vec![2.0, 3.0, 4.0], vec![3], false); - let exp = create_test_tensor_f32(vec![1.0, 2.0, 0.5], vec![3], false); - let result = pow(&base, &exp).unwrap(); - let data = result.data().as_f32_slice().unwrap(); - assert!((data[0] - 2.0).abs() < 1e-6); - assert!((data[1] - 9.0).abs() < 1e-6); - assert!((data[2] - 2.0).abs() < 1e-6); - } - - #[test] - fn test_pow_shape_mismatch_error() { - let base = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - let exp = create_test_tensor_f32(vec![3.0, 4.0, 5.0], vec![3], false); - assert!(pow(&base, &exp).is_err()); - } - - #[test] - fn test_pow_dtype_mismatch_error() { - let base = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - let shape = Shape::new(vec![2]); - let data = TensorData::from_vec_f64(vec![1.0, 2.0], Device::cpu()); - let exp = Tensor::new( - Arc::new(data), - shape, - DataType::Float64, - Device::cpu(), - false, - ); - assert!(pow(&base, &exp).is_err()); - } - - #[test] - fn test_pow_device_mismatch_error() { - let base = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - let shape = Shape::new(vec![2]); - let data = TensorData::from_vec_f32(vec![1.0, 2.0], Device::cuda(Some(0))); - let exp = Tensor::new( - Arc::new(data), - shape, - DataType::Float32, - Device::cuda(Some(0)), - false, - ); - assert!(pow(&base, &exp).is_err()); - } - - #[test] - fn test_powf_gradient() { - let tensor = create_test_tensor_f32(vec![2.0, 3.0], vec![2], true); - let result = powf(&tensor, 3.0).unwrap(); - let ones = Tensor::ones( - result.shape().clone(), - result.dtype(), - result.device(), - false, - ); - let grads = autograd::backward(&result, Some(ones)).unwrap(); - let grad = grads.get(&tensor.id()).unwrap(); - let g = grad.data().as_f32_slice().unwrap(); - assert!((g[0] - 3.0 * 2.0_f32.powf(2.0)).abs() < 1e-6); - assert!((g[1] - 3.0 * 3.0_f32.powf(2.0)).abs() < 1e-6); - } - - #[test] - fn test_pow_base_scalar_tensor_exponent() { - let base = create_test_tensor_f32(vec![2.0], vec![], false); - let exp = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], false); - let result = pow(&base, &exp).unwrap(); - let data = result.data().as_f32_slice().unwrap(); - assert!((data[0] - 2.0).abs() < 1e-6); - assert!((data[1] - 4.0).abs() < 1e-6); - assert!((data[2] - 8.0).abs() < 1e-6); - } - - #[test] - fn test_pow_exponent_scalar_tensor_base() { - let base = create_test_tensor_f32(vec![2.0, 3.0, 4.0], vec![3], false); - let exp = create_test_tensor_f32(vec![2.0], vec![1], false); - let result = pow(&base, &exp).unwrap(); - let data = result.data().as_f32_slice().unwrap(); - assert!((data[0] - 4.0).abs() < 1e-6); - assert!((data[1] - 9.0).abs() < 1e-6); - assert!((data[2] - 16.0).abs() < 1e-6); - } - - #[test] - fn test_pow_base_scalar_gradient() { - let base = create_test_tensor_f32(vec![2.0], vec![], true); - let exp = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - let result = pow(&base, &exp).unwrap(); - let ones = Tensor::ones( - result.shape().clone(), - result.dtype(), - result.device(), - false, - ); - let grads = autograd::backward(&result, Some(ones)).unwrap(); - let grad = grads.get(&base.id()).unwrap(); - let g = grad.data().as_f32_slice().unwrap(); - let base_val = base.data().as_f32_slice().unwrap()[0]; - let exp_vals = exp.data().as_f32_slice().unwrap(); - let expected = exp_vals - .iter() - .map(|&e| e * base_val.powf(e - 1.0)) - .sum::(); - assert!((g[0] - expected).abs() < 1e-6); - } - - #[test] - fn test_pow_exponent_scalar_gradient() { - let base = create_test_tensor_f32(vec![2.0, 3.0], vec![2], false); - let exp = create_test_tensor_f32(vec![1.5], vec![1], true); - let result = pow(&base, &exp).unwrap(); - let ones = Tensor::ones( - result.shape().clone(), - result.dtype(), - result.device(), - false, - ); - let grads = autograd::backward(&result, Some(ones)).unwrap(); - let grad = grads.get(&exp.id()).unwrap(); - let g = grad.data().as_f32_slice().unwrap(); - let exp_val = exp.data().as_f32_slice().unwrap()[0]; - let base_vals = base.data().as_f32_slice().unwrap(); - let expected = base_vals - .iter() - .map(|&b| b.powf(exp_val) * b.ln()) - .sum::(); - assert!((g[0] - expected).abs() < 1e-6); - } - - #[test] - fn test_sqrt() { - let tensor = create_test_tensor_f32(vec![1.0, 4.0, 9.0], vec![3], false); - let result = sqrt(&tensor).unwrap(); - let data = result.data().as_f32_slice().unwrap(); - assert_eq!(data, &[1.0, 2.0, 3.0]); - } - - #[test] - fn test_sqrt_gradient() { - let tensor = create_test_tensor_f32(vec![4.0, 9.0], vec![2], true); - let result = sqrt(&tensor).unwrap(); - let ones = Tensor::ones( - result.shape().clone(), - result.dtype(), - result.device(), - false, - ); - let grads = autograd::backward(&result, Some(ones)).unwrap(); - let grad = grads.get(&tensor.id()).unwrap(); - let g = grad.data().as_f32_slice().unwrap(); - assert!((g[0] - 0.25).abs() < 1e-6); - assert!((g[1] - (1.0 / 6.0)).abs() < 1e-6); - } - - #[test] - fn test_rsqrt() { - let tensor = create_test_tensor_f32(vec![0.25, 1.0, 4.0], vec![3], false); - let result = rsqrt(&tensor).unwrap(); - let data = result.data().as_f32_slice().unwrap(); - assert!((data[0] - 2.0).abs() < 1e-6); - assert!((data[1] - 1.0).abs() < 1e-6); - assert!((data[2] - 0.5).abs() < 1e-6); - } - - #[test] - fn test_rsqrt_gradient() { - let tensor = create_test_tensor_f32(vec![0.25, 4.0], vec![2], true); - let result = rsqrt(&tensor).unwrap(); - let ones = Tensor::ones( - result.shape().clone(), - result.dtype(), - result.device(), - false, - ); - let grads = autograd::backward(&result, Some(ones)).unwrap(); - let grad = grads.get(&tensor.id()).unwrap(); - let g = grad.data().as_f32_slice().unwrap(); - assert!((g[0] - (-0.5 * 0.25_f32.powf(-1.5))).abs() < 1e-5); - assert!((g[1] - (-0.5 * 4.0_f32.powf(-1.5))).abs() < 1e-6); - } - - #[test] - fn test_softsign_forward() { - let data = vec![-2.5f32, -0.5, 0.0, 0.25, 4.0]; - let tensor = create_test_tensor_f32(data.clone(), vec![5], false); - - let result = softsign(&tensor).unwrap(); - let values = result.data().as_f32_slice().unwrap(); - - for (out, &x) in values.iter().zip(data.iter()) { - let denom = 1.0 + x.abs(); - let expected = x / denom; - assert!((out - expected).abs() < 1e-6); - } - } - - #[test] - fn test_softsign_gradient() { - let data = vec![-1.5f32, -0.25, 0.5, 3.0]; - let tensor = create_test_tensor_f32(data.clone(), vec![4], true); - - let result = softsign(&tensor).unwrap(); - let ones = Tensor::ones( - result.shape().clone(), - result.dtype(), - result.device(), - false, - ); - let grads = autograd::backward(&result, Some(ones)).unwrap(); - let grad_tensor = grads.get(&tensor.id()).unwrap(); - let grad_data = grad_tensor.data().as_f32_slice().unwrap(); - - for ((&grad, &x), idx) in grad_data.iter().zip(data.iter()).zip(0..) { - let denom = 1.0 + x.abs(); - let expected = 1.0 / (denom * denom); - assert!( - (grad - expected).abs() < 1e-5, - "gradient mismatch at index {idx}: got {grad}, expected {expected}" - ); - } - } - - #[test] - fn test_gradient_tracking() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], true); - - let result = relu(&tensor).unwrap(); - assert!(result.requires_grad()); - assert!(result.grad_fn().is_some()); - - let result2 = sigmoid(&tensor).unwrap(); - assert!(result2.requires_grad()); - assert!(result2.grad_fn().is_some()); - } -} +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::{ + error::{MinitensorError, Result}, + tensor::{Tensor, TensorData}, +}; + +pub(crate) fn ceil_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + unary_apply(input_data, output_slice, f64::ceil); + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::tensor::DataType; + use crate::{ + autograd, + device::Device, + tensor::{Shape, Tensor, TensorData}, + }; + use std::sync::Arc; + + fn create_test_tensor_f32(data: Vec, shape: Vec, requires_grad: bool) -> Tensor { + let shape_obj = Shape::new(shape); + let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Float32); + + if let Some(slice) = tensor_data.as_f32_slice_mut() { + slice.copy_from_slice(&data); + } + + Tensor::new( + Arc::new(tensor_data), + shape_obj, + DataType::Float32, + Device::cpu(), + requires_grad, + ) + } + + fn create_test_tensor_bool(data: Vec, shape: Vec) -> Tensor { + let shape_obj = Shape::new(shape); + let tensor_data = TensorData::from_vec_bool(data, Device::cpu()); + Tensor::new( + Arc::new(tensor_data), + shape_obj, + DataType::Bool, + Device::cpu(), + false, + ) + } + + #[test] + fn test_nan_to_num_custom_replacements() { + let tensor = create_test_tensor_f32( + vec![f32::NAN, f32::INFINITY, f32::NEG_INFINITY, 4.0], + vec![4], + false, + ); + let result = nan_to_num(&tensor, -1.0, Some(9.0), Some(-9.0)).unwrap(); + let data = result.data().as_f32_slice().unwrap(); + assert_eq!(data, &[-1.0, 9.0, -9.0, 4.0]); + } + + #[test] + fn test_nan_to_num_backward_masks_replaced_entries() { + let tensor = create_test_tensor_f32( + vec![f32::NAN, f32::INFINITY, f32::NEG_INFINITY, -2.0, 3.0], + vec![5], + true, + ); + let result = nan_to_num(&tensor, 0.0, Some(10.0), Some(-10.0)).unwrap(); + let summed = crate::operations::reduction::sum(&result, None, false).unwrap(); + let grads = autograd::backward_collect(&summed, None).unwrap(); + let grad = grads.get(&tensor.id()).unwrap(); + assert_eq!( + grad.data().as_f32_slice().unwrap(), + &[0.0, 0.0, 0.0, 1.0, 1.0] + ); + } + + #[test] + fn test_nan_to_num_integer_is_unchanged() { + let shape = Shape::new(vec![3]); + let data = TensorData::from_vec_i32(vec![1, -2, 3], Device::cpu()); + let tensor = Tensor::new(Arc::new(data), shape, DataType::Int32, Device::cpu(), false); + let result = nan_to_num(&tensor, 99.0, None, None).unwrap(); + assert_eq!(result.data().as_i32_slice().unwrap(), &[1, -2, 3]); + } + + #[test] + fn test_exp() { + let tensor = create_test_tensor_f32(vec![0.0, 1.0, 2.0], vec![3], false); + let result = exp(&tensor).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + + assert!((result_data[0] - 1.0).abs() < 1e-6); + assert!((result_data[1] - std::f32::consts::E).abs() < 1e-6); + assert!((result_data[2] - (std::f32::consts::E * std::f32::consts::E)).abs() < 1e-5); + } + + #[test] + fn test_exp_invalid_dtype() { + let shape = Shape::new(vec![3]); + let data = TensorData::from_vec_i32(vec![1, 2, 3], Device::cpu()); + let tensor = Tensor::new(Arc::new(data), shape, DataType::Int32, Device::cpu(), false); + assert!(exp(&tensor).is_err()); + } + + #[test] + fn test_log() { + let tensor = create_test_tensor_f32( + vec![ + 1.0, + std::f32::consts::E, + std::f32::consts::E * std::f32::consts::E, + ], + vec![3], + false, + ); + let result = log(&tensor).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + + assert!((result_data[0] - 0.0).abs() < 1e-6); + assert!((result_data[1] - 1.0).abs() < 1e-6); + assert!((result_data[2] - 2.0).abs() < 1e-5); + } + + #[test] + fn test_sin() { + let tensor = create_test_tensor_f32( + vec![0.0, std::f32::consts::PI / 2.0, std::f32::consts::PI], + vec![3], + false, + ); + let result = sin(&tensor).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + + assert!((result_data[0] - 0.0).abs() < 1e-6); + assert!((result_data[1] - 1.0).abs() < 1e-6); + assert!(result_data[2].abs() < 1e-6); // sin(π) ≈ 0 + } + + #[test] + fn test_cos() { + let tensor = create_test_tensor_f32( + vec![0.0, std::f32::consts::PI / 2.0, std::f32::consts::PI], + vec![3], + false, + ); + let result = cos(&tensor).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + + assert!((result_data[0] - 1.0).abs() < 1e-6); + assert!(result_data[1].abs() < 1e-6); // cos(π/2) ≈ 0 + assert!((result_data[2] + 1.0).abs() < 1e-6); // cos(π) ≈ -1 + } + + #[test] + fn test_tan() { + let tensor = create_test_tensor_f32( + vec![0.0, std::f32::consts::PI / 4.0, -std::f32::consts::PI / 4.0], + vec![3], + false, + ); + let result = tan(&tensor).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + + assert!((result_data[0] - 0.0).abs() < 1e-6); + assert!((result_data[1] - (std::f32::consts::PI / 4.0).tan()).abs() < 1e-6); + assert!((result_data[2] - (-std::f32::consts::PI / 4.0).tan()).abs() < 1e-6); + } + + #[test] + fn test_tanh() { + let tensor = create_test_tensor_f32(vec![0.0, 1.0, -1.0], vec![3], false); + let result = tanh(&tensor).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + + assert!((result_data[0] - 0.0).abs() < 1e-6); + assert!((result_data[1] - 1.0_f32.tanh()).abs() < 1e-6); + assert!((result_data[2] - (-1.0_f32).tanh()).abs() < 1e-6); + } + + #[test] + fn test_sigmoid() { + let tensor = create_test_tensor_f32(vec![0.0, 1.0, -1.0], vec![3], false); + let result = sigmoid(&tensor).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + + assert!((result_data[0] - 0.5).abs() < 1e-6); + assert!((result_data[1] - (1.0 / (1.0 + (-1.0_f32).exp()))).abs() < 1e-6); + assert!((result_data[2] - (1.0 / (1.0 + 1.0_f32.exp()))).abs() < 1e-6); + } + + #[test] + fn test_relu() { + let tensor = create_test_tensor_f32(vec![-2.0, -1.0, 0.0, 1.0, 2.0], vec![5], false); + let result = relu(&tensor).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + + assert_eq!(result_data, &[0.0, 0.0, 0.0, 1.0, 2.0]); + } + + #[test] + fn test_hardshrink_forward() { + let tensor = create_test_tensor_f32(vec![-1.2, -0.2, 0.0, 0.45, 0.9], vec![5], false); + let result = hardshrink(&tensor, 0.3).unwrap(); + let data = result.data().as_f32_slice().unwrap(); + assert_eq!(data, &[-1.2, 0.0, 0.0, 0.45, 0.9]); + } + + #[test] + fn test_hardshrink_backward() { + let tensor = create_test_tensor_f32(vec![-1.2, -0.25, 0.0, 0.35, 0.8], vec![5], true); + let result = hardshrink(&tensor, 0.3).unwrap(); + let ones = Tensor::ones( + result.shape().clone(), + result.dtype(), + result.device(), + false, + ); + let grads = autograd::backward_collect(&result, Some(ones)).unwrap(); + let grad = grads.get(&tensor.id()).unwrap(); + let grad_vals = grad.data().as_f32_slice().unwrap(); + assert_eq!(grad_vals, &[1.0, 0.0, 0.0, 1.0, 1.0]); + } + + #[test] + fn test_hardshrink_invalid_lambda() { + let tensor = create_test_tensor_f32(vec![1.0], vec![1], false); + assert!(hardshrink(&tensor, -0.1).is_err()); + } + + #[test] + fn test_leaky_relu() { + let tensor = create_test_tensor_f32(vec![-2.0, -1.0, 0.0, 1.0, 2.0], vec![5], false); + let result = leaky_relu(&tensor, 0.1).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + + assert_eq!(result_data, &[-0.2, -0.1, 0.0, 1.0, 2.0]); + } + + #[test] + fn test_softmax() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], false); + let result = softmax(&tensor, None).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + + // Check that probabilities sum to 1 + let sum: f32 = result_data.iter().sum(); + assert!((sum - 1.0).abs() < 1e-6); + + // Check that all values are positive + for &val in result_data { + assert!(val > 0.0); + } + + // Check that larger input values produce larger probabilities + assert!(result_data[2] > result_data[1]); + assert!(result_data[1] > result_data[0]); + } + + #[test] + fn test_masked_softmax_respects_mask() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); + let mask = create_test_tensor_bool(vec![true, false, false, true], vec![2, 2]); + let result = masked_softmax(&tensor, &mask, Some(1)).unwrap(); + let data = result.data().as_f32_slice().unwrap(); + assert_eq!(data[0], 0.0); + assert_eq!(data[3], 0.0); + assert!((data[1] - 1.0).abs() < 1e-6); + assert!((data[2] - 1.0).abs() < 1e-6); + } + + #[test] + fn test_masked_softmax_all_masked_returns_zero() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + let mask = create_test_tensor_bool(vec![true, true], vec![2]); + let result = masked_softmax(&tensor, &mask, Some(0)).unwrap(); + let data = result.data().as_f32_slice().unwrap(); + assert!(data.iter().all(|v| *v == 0.0)); + } + + #[test] + fn test_masked_softmax_all_negative_infinity_unmasked() { + let tensor = + create_test_tensor_f32(vec![f32::NEG_INFINITY, f32::NEG_INFINITY], vec![2], false); + let mask = create_test_tensor_bool(vec![false, false], vec![2]); + let result = masked_softmax(&tensor, &mask, Some(0)).unwrap(); + let data = result.data().as_f32_slice().unwrap(); + assert!(data.iter().all(|v| *v == 0.0)); + } + + #[test] + fn test_masked_log_softmax_broadcast_mask() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); + let mask = create_test_tensor_bool(vec![true, false], vec![2, 1]); + let result = masked_log_softmax(&tensor, &mask, Some(1)).unwrap(); + let data = result.data().as_f32_slice().unwrap(); + assert!(data[0].is_infinite() && data[0].is_sign_negative()); + assert!(data[1].is_infinite() && data[1].is_sign_negative()); + let row1_prob0 = data[2].exp(); + let row1_prob1 = data[3].exp(); + assert!((row1_prob0 + row1_prob1 - 1.0).abs() < 1e-6); + assert!(data[3] > data[2]); + } + + #[test] + fn test_masked_log_softmax_all_masked_returns_negative_infinity() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + let mask = create_test_tensor_bool(vec![true, true], vec![2]); + let result = masked_log_softmax(&tensor, &mask, Some(0)).unwrap(); + let data = result.data().as_f32_slice().unwrap(); + assert!(data.iter().all(|v| v.is_infinite() && v.is_sign_negative())); + } + + #[test] + fn test_log_softmax_large_negative_values() { + let tensor = create_test_tensor_f32(vec![-1000.0, 0.0], vec![2], false); + let result = log_softmax(&tensor, None).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + + assert!(result_data[0].is_finite()); + assert!((result_data[0] + 1000.0).abs() < 1e-3); + assert!(result_data[1].abs() < 1e-6); + } + + #[test] + fn test_log_softmax_all_negative_infinity() { + let tensor = + create_test_tensor_f32(vec![f32::NEG_INFINITY, f32::NEG_INFINITY], vec![2], false); + let result = log_softmax(&tensor, None).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + + assert!( + result_data + .iter() + .all(|v| v.is_infinite() && v.is_sign_negative()) + ); + } + + #[test] + fn test_powf_scalar() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], false); + let result = powf(&tensor, 2.0).unwrap(); + let data = result.data().as_f32_slice().unwrap(); + assert_eq!(data, &[1.0, 4.0, 9.0]); + } + + #[test] + fn test_pow_tensor() { + let base = create_test_tensor_f32(vec![2.0, 3.0, 4.0], vec![3], false); + let exp = create_test_tensor_f32(vec![1.0, 2.0, 0.5], vec![3], false); + let result = pow(&base, &exp).unwrap(); + let data = result.data().as_f32_slice().unwrap(); + assert!((data[0] - 2.0).abs() < 1e-6); + assert!((data[1] - 9.0).abs() < 1e-6); + assert!((data[2] - 2.0).abs() < 1e-6); + } + + #[test] + fn test_pow_shape_mismatch_error() { + let base = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + let exp = create_test_tensor_f32(vec![3.0, 4.0, 5.0], vec![3], false); + assert!(pow(&base, &exp).is_err()); + } + + #[test] + fn test_pow_dtype_mismatch_error() { + let base = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + let shape = Shape::new(vec![2]); + let data = TensorData::from_vec_f64(vec![1.0, 2.0], Device::cpu()); + let exp = Tensor::new( + Arc::new(data), + shape, + DataType::Float64, + Device::cpu(), + false, + ); + assert!(pow(&base, &exp).is_err()); + } + + #[test] + fn test_pow_device_mismatch_error() { + let base = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + let shape = Shape::new(vec![2]); + let data = TensorData::from_vec_f32(vec![1.0, 2.0], Device::cuda(Some(0))); + let exp = Tensor::new( + Arc::new(data), + shape, + DataType::Float32, + Device::cuda(Some(0)), + false, + ); + assert!(pow(&base, &exp).is_err()); + } + + #[test] + fn test_powf_gradient() { + let tensor = create_test_tensor_f32(vec![2.0, 3.0], vec![2], true); + let result = powf(&tensor, 3.0).unwrap(); + let ones = Tensor::ones( + result.shape().clone(), + result.dtype(), + result.device(), + false, + ); + let grads = autograd::backward_collect(&result, Some(ones)).unwrap(); + let grad = grads.get(&tensor.id()).unwrap(); + let g = grad.data().as_f32_slice().unwrap(); + assert!((g[0] - 3.0 * 2.0_f32.powf(2.0)).abs() < 1e-6); + assert!((g[1] - 3.0 * 3.0_f32.powf(2.0)).abs() < 1e-6); + } + + #[test] + fn test_pow_base_scalar_tensor_exponent() { + let base = create_test_tensor_f32(vec![2.0], vec![], false); + let exp = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], false); + let result = pow(&base, &exp).unwrap(); + let data = result.data().as_f32_slice().unwrap(); + assert!((data[0] - 2.0).abs() < 1e-6); + assert!((data[1] - 4.0).abs() < 1e-6); + assert!((data[2] - 8.0).abs() < 1e-6); + } + + #[test] + fn test_pow_exponent_scalar_tensor_base() { + let base = create_test_tensor_f32(vec![2.0, 3.0, 4.0], vec![3], false); + let exp = create_test_tensor_f32(vec![2.0], vec![1], false); + let result = pow(&base, &exp).unwrap(); + let data = result.data().as_f32_slice().unwrap(); + assert!((data[0] - 4.0).abs() < 1e-6); + assert!((data[1] - 9.0).abs() < 1e-6); + assert!((data[2] - 16.0).abs() < 1e-6); + } + + #[test] + fn test_pow_base_scalar_gradient() { + let base = create_test_tensor_f32(vec![2.0], vec![], true); + let exp = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + let result = pow(&base, &exp).unwrap(); + let ones = Tensor::ones( + result.shape().clone(), + result.dtype(), + result.device(), + false, + ); + let grads = autograd::backward_collect(&result, Some(ones)).unwrap(); + let grad = grads.get(&base.id()).unwrap(); + let g = grad.data().as_f32_slice().unwrap(); + let base_val = base.data().as_f32_slice().unwrap()[0]; + let exp_vals = exp.data().as_f32_slice().unwrap(); + let expected = exp_vals + .iter() + .map(|&e| e * base_val.powf(e - 1.0)) + .sum::(); + assert!((g[0] - expected).abs() < 1e-6); + } + + #[test] + fn test_pow_exponent_scalar_gradient() { + let base = create_test_tensor_f32(vec![2.0, 3.0], vec![2], false); + let exp = create_test_tensor_f32(vec![1.5], vec![1], true); + let result = pow(&base, &exp).unwrap(); + let ones = Tensor::ones( + result.shape().clone(), + result.dtype(), + result.device(), + false, + ); + let grads = autograd::backward_collect(&result, Some(ones)).unwrap(); + let grad = grads.get(&exp.id()).unwrap(); + let g = grad.data().as_f32_slice().unwrap(); + let exp_val = exp.data().as_f32_slice().unwrap()[0]; + let base_vals = base.data().as_f32_slice().unwrap(); + let expected = base_vals + .iter() + .map(|&b| b.powf(exp_val) * b.ln()) + .sum::(); + assert!((g[0] - expected).abs() < 1e-6); + } + + #[test] + fn test_sqrt() { + let tensor = create_test_tensor_f32(vec![1.0, 4.0, 9.0], vec![3], false); + let result = sqrt(&tensor).unwrap(); + let data = result.data().as_f32_slice().unwrap(); + assert_eq!(data, &[1.0, 2.0, 3.0]); + } + + #[test] + fn test_sqrt_gradient() { + let tensor = create_test_tensor_f32(vec![4.0, 9.0], vec![2], true); + let result = sqrt(&tensor).unwrap(); + let ones = Tensor::ones( + result.shape().clone(), + result.dtype(), + result.device(), + false, + ); + let grads = autograd::backward_collect(&result, Some(ones)).unwrap(); + let grad = grads.get(&tensor.id()).unwrap(); + let g = grad.data().as_f32_slice().unwrap(); + assert!((g[0] - 0.25).abs() < 1e-6); + assert!((g[1] - (1.0 / 6.0)).abs() < 1e-6); + } + + #[test] + fn test_rsqrt() { + let tensor = create_test_tensor_f32(vec![0.25, 1.0, 4.0], vec![3], false); + let result = rsqrt(&tensor).unwrap(); + let data = result.data().as_f32_slice().unwrap(); + assert!((data[0] - 2.0).abs() < 1e-6); + assert!((data[1] - 1.0).abs() < 1e-6); + assert!((data[2] - 0.5).abs() < 1e-6); + } + + #[test] + fn test_rsqrt_gradient() { + let tensor = create_test_tensor_f32(vec![0.25, 4.0], vec![2], true); + let result = rsqrt(&tensor).unwrap(); + let ones = Tensor::ones( + result.shape().clone(), + result.dtype(), + result.device(), + false, + ); + let grads = autograd::backward_collect(&result, Some(ones)).unwrap(); + let grad = grads.get(&tensor.id()).unwrap(); + let g = grad.data().as_f32_slice().unwrap(); + assert!((g[0] - (-0.5 * 0.25_f32.powf(-1.5))).abs() < 1e-5); + assert!((g[1] - (-0.5 * 4.0_f32.powf(-1.5))).abs() < 1e-6); + } + + #[test] + fn test_softsign_forward() { + let data = vec![-2.5f32, -0.5, 0.0, 0.25, 4.0]; + let tensor = create_test_tensor_f32(data.clone(), vec![5], false); + + let result = softsign(&tensor).unwrap(); + let values = result.data().as_f32_slice().unwrap(); + + for (out, &x) in values.iter().zip(data.iter()) { + let denom = 1.0 + x.abs(); + let expected = x / denom; + assert!((out - expected).abs() < 1e-6); + } + } + + #[test] + fn test_softsign_gradient() { + let data = vec![-1.5f32, -0.25, 0.5, 3.0]; + let tensor = create_test_tensor_f32(data.clone(), vec![4], true); + + let result = softsign(&tensor).unwrap(); + let ones = Tensor::ones( + result.shape().clone(), + result.dtype(), + result.device(), + false, + ); + let grads = autograd::backward_collect(&result, Some(ones)).unwrap(); + let grad_tensor = grads.get(&tensor.id()).unwrap(); + let grad_data = grad_tensor.data().as_f32_slice().unwrap(); + + for ((&grad, &x), idx) in grad_data.iter().zip(data.iter()).zip(0..) { + let denom = 1.0 + x.abs(); + let expected = 1.0 / (denom * denom); + assert!( + (grad - expected).abs() < 1e-5, + "gradient mismatch at index {idx}: got {grad}, expected {expected}" + ); + } + } + + #[test] + fn test_gradient_tracking() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], true); + + let result = relu(&tensor).unwrap(); + assert!(result.requires_grad()); + assert!(result.grad_fn().is_some()); + + let result2 = sigmoid(&tensor).unwrap(); + assert!(result2.requires_grad()); + assert!(result2.grad_fn().is_some()); + } +} diff --git a/engine/src/operations/activation/elementwise.rs b/engine/src/operations/activation/elementwise.rs index 9fa61dc3..86dc8c2d 100644 --- a/engine/src/operations/activation/elementwise.rs +++ b/engine/src/operations/activation/elementwise.rs @@ -1,767 +1,764 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -use crate::{ - autograd::{ - AbsBackward, AcosBackward, AcoshBackward, AsinBackward, AsinhBackward, AtanBackward, - AtanhBackward, ClampBackward, CosBackward, CoshBackward, EluBackward, ExpBackward, - Expm1Backward, GeluBackward, HardshrinkBackward, LeakyReluBackward, Log1pBackward, - LogAddExpBackward, LogBackward, LogSoftmaxBackward, MaskedLogSoftmaxBackward, - NanToNumBackward, PowBackward, PowBroadcast, ReluBackward, SeluBackward, SigmoidBackward, - SiluBackward, SinBackward, SinhBackward, SoftmaxBackward, SoftplusBackward, SoftsignBackward, - TanBackward, TanhBackward, add_to_graph, - }, - error::{MinitensorError, Result}, - tensor::{DataType, Shape, Strides, Tensor, TensorData}, -}; -use libm::{erf, erff}; -use rayon::prelude::*; -use std::sync::Arc; - -const PAR_THRESHOLD: usize = 1 << 12; // 4096 elements - -#[inline(always)] -fn unary_apply(input: &[T], output: &mut [T], op: F) -where - T: Copy + Send + Sync, - F: Fn(T) -> T + Sync + Send, -{ - #[inline(always)] - fn apply_chunk(input: &[T], output: &mut [T], op: &F) - where - T: Copy, - F: Fn(T) -> T, - { - let len = input.len(); - let mut i = 0usize; - let n = len.saturating_sub(len % 8); - while i < n { - unsafe { - *output.get_unchecked_mut(i) = op(*input.get_unchecked(i)); - *output.get_unchecked_mut(i + 1) = op(*input.get_unchecked(i + 1)); - *output.get_unchecked_mut(i + 2) = op(*input.get_unchecked(i + 2)); - *output.get_unchecked_mut(i + 3) = op(*input.get_unchecked(i + 3)); - *output.get_unchecked_mut(i + 4) = op(*input.get_unchecked(i + 4)); - *output.get_unchecked_mut(i + 5) = op(*input.get_unchecked(i + 5)); - *output.get_unchecked_mut(i + 6) = op(*input.get_unchecked(i + 6)); - *output.get_unchecked_mut(i + 7) = op(*input.get_unchecked(i + 7)); - } - i += 8; - } - for j in i..len { - unsafe { - *output.get_unchecked_mut(j) = op(*input.get_unchecked(j)); - } - } - } - - let len = input.len(); - debug_assert_eq!(len, output.len()); - if len < PAR_THRESHOLD { - apply_chunk(input, output, &op); - } else { - const CHUNK: usize = 1024; - input - .par_chunks(CHUNK) - .zip(output.par_chunks_mut(CHUNK)) - .for_each(|(in_chunk, out_chunk)| apply_chunk(in_chunk, out_chunk, &op)); - } -} - -/// Exponential function with gradient support -pub fn exp(tensor: &Tensor) -> Result { - // Create output tensor data - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - // Perform exponential based on data type - match tensor.dtype() { - DataType::Float32 => exp_f32(tensor, &mut output_data)?, - DataType::Float64 => exp_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "Exponential function only supported for floating point tensors", - )); - } - } - - // Create output tensor - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - // Set up gradient function if needed - if output.requires_grad() { - let grad_fn = Arc::new(ExpBackward { - input_id: tensor.id(), - output: output.clone().detach(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&output_with_grad, Some(grad_fn))?; - - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Natural logarithm function with gradient support -pub fn log(tensor: &Tensor) -> Result { - // Create output tensor data - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - // Perform logarithm based on data type - match tensor.dtype() { - DataType::Float32 => log_f32(tensor, &mut output_data)?, - DataType::Float64 => log_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "Logarithm function only supported for floating point tensors", - )); - } - } - - // Create output tensor - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - // Set up gradient function if needed - if output.requires_grad() { - let grad_fn = Arc::new(LogBackward { - input_id: tensor.id(), - input: tensor.clone().detach(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&output_with_grad, Some(grad_fn))?; - - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// log1p (log(1 + x)) function with gradient support -pub fn log1p(tensor: &Tensor) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => log1p_f32(tensor, &mut output_data)?, - DataType::Float64 => log1p_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "log1p is only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(Log1pBackward { - input_id: tensor.id(), - input: tensor.clone().detach(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// expm1 (exp(x) - 1) with gradient support -pub fn expm1(tensor: &Tensor) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => expm1_f32(tensor, &mut output_data)?, - DataType::Float64 => expm1_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "expm1 is only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(Expm1Backward { - input_id: tensor.id(), - output: output.clone().detach(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Sine function with gradient support -pub fn sin(tensor: &Tensor) -> Result { - // Create output tensor data - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - // Perform sine based on data type - match tensor.dtype() { - DataType::Float32 => sin_f32(tensor, &mut output_data)?, - DataType::Float64 => sin_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "Sine function only supported for floating point tensors", - )); - } - } - - // Create output tensor - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - // Set up gradient function if needed - if output.requires_grad() { - let grad_fn = Arc::new(SinBackward { - input_id: tensor.id(), - input: tensor.clone(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&output_with_grad, Some(grad_fn))?; - - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Cosine function with gradient support -pub fn cos(tensor: &Tensor) -> Result { - // Create output tensor data - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - // Perform cosine based on data type - match tensor.dtype() { - DataType::Float32 => cos_f32(tensor, &mut output_data)?, - DataType::Float64 => cos_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "Cosine function only supported for floating point tensors", - )); - } - } - - // Create output tensor - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - // Set up gradient function if needed - if output.requires_grad() { - let grad_fn = Arc::new(CosBackward { - input_id: tensor.id(), - input: tensor.clone(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&output_with_grad, Some(grad_fn))?; - - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Tangent function with gradient support -pub fn tan(tensor: &Tensor) -> Result { - // Create output tensor data - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - // Perform tangent based on data type - match tensor.dtype() { - DataType::Float32 => tan_f32(tensor, &mut output_data)?, - DataType::Float64 => tan_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "Tangent function only supported for floating point tensors", - )); - } - } - - // Create output tensor - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - // Set up gradient function if needed - if output.requires_grad() { - let grad_fn = Arc::new(TanBackward { - input_id: tensor.id(), - output: output.clone().detach(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&output_with_grad, Some(grad_fn))?; - - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Inverse sine function with gradient support -pub fn asin(tensor: &Tensor) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => asin_f32(tensor, &mut output_data)?, - DataType::Float64 => asin_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "Inverse sine only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(AsinBackward { - input_id: tensor.id(), - input: tensor.clone(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Inverse cosine function with gradient support -pub fn acos(tensor: &Tensor) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => acos_f32(tensor, &mut output_data)?, - DataType::Float64 => acos_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "Inverse cosine only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(AcosBackward { - input_id: tensor.id(), - input: tensor.clone(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Inverse tangent function with gradient support -pub fn atan(tensor: &Tensor) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => atan_f32(tensor, &mut output_data)?, - DataType::Float64 => atan_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "Inverse tangent only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(AtanBackward { - input_id: tensor.id(), - input: tensor.clone(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Hyperbolic sine with gradient support -pub fn sinh(tensor: &Tensor) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => sinh_f32(tensor, &mut output_data)?, - DataType::Float64 => sinh_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "sinh is only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(SinhBackward { - input_id: tensor.id(), - input: tensor.clone(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Hyperbolic cosine with gradient support -pub fn cosh(tensor: &Tensor) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => cosh_f32(tensor, &mut output_data)?, - DataType::Float64 => cosh_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "cosh is only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(CoshBackward { - input_id: tensor.id(), - input: tensor.clone(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Inverse hyperbolic sine with gradient support -pub fn asinh(tensor: &Tensor) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => asinh_f32(tensor, &mut output_data)?, - DataType::Float64 => asinh_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "asinh is only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(AsinhBackward { - input_id: tensor.id(), - input: tensor.clone(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Inverse hyperbolic cosine with gradient support -pub fn acosh(tensor: &Tensor) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => acosh_f32(tensor, &mut output_data)?, - DataType::Float64 => acosh_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "acosh is only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(AcoshBackward { - input_id: tensor.id(), - input: tensor.clone(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Inverse hyperbolic tangent with gradient support -pub fn atanh(tensor: &Tensor) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => atanh_f32(tensor, &mut output_data)?, - DataType::Float64 => atanh_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "atanh is only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(AtanhBackward { - input_id: tensor.id(), - input: tensor.clone(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Hyperbolic tangent function with gradient support -pub fn tanh(tensor: &Tensor) -> Result { - // Create output tensor data - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - // Perform tanh based on data type - match tensor.dtype() { - DataType::Float32 => tanh_f32(tensor, &mut output_data)?, - DataType::Float64 => tanh_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "Tanh function only supported for floating point tensors", - )); - } - } - - // Create output tensor - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - // Set up gradient function if needed - if output.requires_grad() { - let grad_fn = Arc::new(TanhBackward { - input_id: tensor.id(), - output: output.clone(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&output_with_grad, Some(grad_fn))?; - - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Sigmoid activation function with gradient support -pub fn sigmoid(tensor: &Tensor) -> Result { - // Create output tensor data - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - // Perform sigmoid based on data type - match tensor.dtype() { - DataType::Float32 => sigmoid_f32(tensor, &mut output_data)?, - DataType::Float64 => sigmoid_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "Sigmoid function only supported for floating point tensors", - )); - } - } - - // Create output tensor - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - // Set up gradient function if needed - if output.requires_grad() { - let grad_fn = Arc::new(SigmoidBackward { - input_id: tensor.id(), - output: output.clone(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&output_with_grad, Some(grad_fn))?; - - Ok(output_with_grad) - } else { - Ok(output) - } -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; + +use crate::{ + autograd::{ + AcosBackward, AcoshBackward, AsinBackward, AsinhBackward, AtanBackward, AtanhBackward, + CosBackward, CoshBackward, ExpBackward, Expm1Backward, Log1pBackward, LogBackward, + SigmoidBackward, SinBackward, SinhBackward, TanBackward, TanhBackward, add_to_graph, + }, + error::{MinitensorError, Result}, + tensor::{DataType, Tensor, TensorData}, +}; +use rayon::prelude::*; +use std::sync::Arc; + +pub(crate) const PAR_THRESHOLD: usize = 1 << 12; // 4096 elements + +#[inline(always)] +pub(crate) fn unary_apply(input: &[T], output: &mut [T], op: F) +where + T: Copy + Send + Sync, + F: Fn(T) -> T + Sync + Send, +{ + #[inline(always)] + fn apply_chunk(input: &[T], output: &mut [T], op: &F) + where + T: Copy, + F: Fn(T) -> T, + { + let len = input.len(); + let mut i = 0usize; + let n = len.saturating_sub(len % 8); + while i < n { + unsafe { + *output.get_unchecked_mut(i) = op(*input.get_unchecked(i)); + *output.get_unchecked_mut(i + 1) = op(*input.get_unchecked(i + 1)); + *output.get_unchecked_mut(i + 2) = op(*input.get_unchecked(i + 2)); + *output.get_unchecked_mut(i + 3) = op(*input.get_unchecked(i + 3)); + *output.get_unchecked_mut(i + 4) = op(*input.get_unchecked(i + 4)); + *output.get_unchecked_mut(i + 5) = op(*input.get_unchecked(i + 5)); + *output.get_unchecked_mut(i + 6) = op(*input.get_unchecked(i + 6)); + *output.get_unchecked_mut(i + 7) = op(*input.get_unchecked(i + 7)); + } + i += 8; + } + for j in i..len { + unsafe { + *output.get_unchecked_mut(j) = op(*input.get_unchecked(j)); + } + } + } + + let len = input.len(); + debug_assert_eq!(len, output.len()); + if len < PAR_THRESHOLD { + apply_chunk(input, output, &op); + } else { + const CHUNK: usize = 1024; + input + .par_chunks(CHUNK) + .zip(output.par_chunks_mut(CHUNK)) + .for_each(|(in_chunk, out_chunk)| apply_chunk(in_chunk, out_chunk, &op)); + } +} + +/// Exponential function with gradient support +pub fn exp(tensor: &Tensor) -> Result { + // Create output tensor data + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + // Perform exponential based on data type + match tensor.dtype() { + DataType::Float32 => exp_f32(tensor, &mut output_data)?, + DataType::Float64 => exp_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "Exponential function only supported for floating point tensors", + )); + } + } + + // Create output tensor + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + // Set up gradient function if needed + if output.requires_grad() { + let grad_fn = Arc::new(ExpBackward { + input_id: tensor.id(), + output: output.clone().detach(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&output_with_grad, Some(grad_fn))?; + + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Natural logarithm function with gradient support +pub fn log(tensor: &Tensor) -> Result { + // Create output tensor data + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + // Perform logarithm based on data type + match tensor.dtype() { + DataType::Float32 => log_f32(tensor, &mut output_data)?, + DataType::Float64 => log_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "Logarithm function only supported for floating point tensors", + )); + } + } + + // Create output tensor + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + // Set up gradient function if needed + if output.requires_grad() { + let grad_fn = Arc::new(LogBackward { + input_id: tensor.id(), + input: tensor.clone().detach(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&output_with_grad, Some(grad_fn))?; + + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// log1p (log(1 + x)) function with gradient support +pub fn log1p(tensor: &Tensor) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => log1p_f32(tensor, &mut output_data)?, + DataType::Float64 => log1p_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "log1p is only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(Log1pBackward { + input_id: tensor.id(), + input: tensor.clone().detach(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// expm1 (exp(x) - 1) with gradient support +pub fn expm1(tensor: &Tensor) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => expm1_f32(tensor, &mut output_data)?, + DataType::Float64 => expm1_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "expm1 is only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(Expm1Backward { + input_id: tensor.id(), + output: output.clone().detach(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Sine function with gradient support +pub fn sin(tensor: &Tensor) -> Result { + // Create output tensor data + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + // Perform sine based on data type + match tensor.dtype() { + DataType::Float32 => sin_f32(tensor, &mut output_data)?, + DataType::Float64 => sin_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "Sine function only supported for floating point tensors", + )); + } + } + + // Create output tensor + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + // Set up gradient function if needed + if output.requires_grad() { + let grad_fn = Arc::new(SinBackward { + input_id: tensor.id(), + input: tensor.clone(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&output_with_grad, Some(grad_fn))?; + + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Cosine function with gradient support +pub fn cos(tensor: &Tensor) -> Result { + // Create output tensor data + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + // Perform cosine based on data type + match tensor.dtype() { + DataType::Float32 => cos_f32(tensor, &mut output_data)?, + DataType::Float64 => cos_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "Cosine function only supported for floating point tensors", + )); + } + } + + // Create output tensor + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + // Set up gradient function if needed + if output.requires_grad() { + let grad_fn = Arc::new(CosBackward { + input_id: tensor.id(), + input: tensor.clone(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&output_with_grad, Some(grad_fn))?; + + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Tangent function with gradient support +pub fn tan(tensor: &Tensor) -> Result { + // Create output tensor data + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + // Perform tangent based on data type + match tensor.dtype() { + DataType::Float32 => tan_f32(tensor, &mut output_data)?, + DataType::Float64 => tan_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "Tangent function only supported for floating point tensors", + )); + } + } + + // Create output tensor + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + // Set up gradient function if needed + if output.requires_grad() { + let grad_fn = Arc::new(TanBackward { + input_id: tensor.id(), + output: output.clone().detach(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&output_with_grad, Some(grad_fn))?; + + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Inverse sine function with gradient support +pub fn asin(tensor: &Tensor) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => asin_f32(tensor, &mut output_data)?, + DataType::Float64 => asin_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "Inverse sine only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(AsinBackward { + input_id: tensor.id(), + input: tensor.clone(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Inverse cosine function with gradient support +pub fn acos(tensor: &Tensor) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => acos_f32(tensor, &mut output_data)?, + DataType::Float64 => acos_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "Inverse cosine only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(AcosBackward { + input_id: tensor.id(), + input: tensor.clone(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Inverse tangent function with gradient support +pub fn atan(tensor: &Tensor) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => atan_f32(tensor, &mut output_data)?, + DataType::Float64 => atan_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "Inverse tangent only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(AtanBackward { + input_id: tensor.id(), + input: tensor.clone(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Hyperbolic sine with gradient support +pub fn sinh(tensor: &Tensor) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => sinh_f32(tensor, &mut output_data)?, + DataType::Float64 => sinh_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "sinh is only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(SinhBackward { + input_id: tensor.id(), + input: tensor.clone(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Hyperbolic cosine with gradient support +pub fn cosh(tensor: &Tensor) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => cosh_f32(tensor, &mut output_data)?, + DataType::Float64 => cosh_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "cosh is only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(CoshBackward { + input_id: tensor.id(), + input: tensor.clone(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Inverse hyperbolic sine with gradient support +pub fn asinh(tensor: &Tensor) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => asinh_f32(tensor, &mut output_data)?, + DataType::Float64 => asinh_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "asinh is only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(AsinhBackward { + input_id: tensor.id(), + input: tensor.clone(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Inverse hyperbolic cosine with gradient support +pub fn acosh(tensor: &Tensor) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => acosh_f32(tensor, &mut output_data)?, + DataType::Float64 => acosh_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "acosh is only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(AcoshBackward { + input_id: tensor.id(), + input: tensor.clone(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Inverse hyperbolic tangent with gradient support +pub fn atanh(tensor: &Tensor) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => atanh_f32(tensor, &mut output_data)?, + DataType::Float64 => atanh_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "atanh is only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(AtanhBackward { + input_id: tensor.id(), + input: tensor.clone(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Hyperbolic tangent function with gradient support +pub fn tanh(tensor: &Tensor) -> Result { + // Create output tensor data + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + // Perform tanh based on data type + match tensor.dtype() { + DataType::Float32 => tanh_f32(tensor, &mut output_data)?, + DataType::Float64 => tanh_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "Tanh function only supported for floating point tensors", + )); + } + } + + // Create output tensor + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + // Set up gradient function if needed + if output.requires_grad() { + let grad_fn = Arc::new(TanhBackward { + input_id: tensor.id(), + output: output.clone(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&output_with_grad, Some(grad_fn))?; + + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Sigmoid activation function with gradient support +pub fn sigmoid(tensor: &Tensor) -> Result { + // Create output tensor data + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + // Perform sigmoid based on data type + match tensor.dtype() { + DataType::Float32 => sigmoid_f32(tensor, &mut output_data)?, + DataType::Float64 => sigmoid_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "Sigmoid function only supported for floating point tensors", + )); + } + } + + // Create output tensor + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + // Set up gradient function if needed + if output.requires_grad() { + let grad_fn = Arc::new(SigmoidBackward { + input_id: tensor.id(), + output: output.clone(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&output_with_grad, Some(grad_fn))?; + + Ok(output_with_grad) + } else { + Ok(output) + } +} diff --git a/engine/src/operations/activation/hyperbolic.rs b/engine/src/operations/activation/hyperbolic.rs index 83018a29..e5c5262d 100644 --- a/engine/src/operations/activation/hyperbolic.rs +++ b/engine/src/operations/activation/hyperbolic.rs @@ -1,881 +1,615 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -/// Masked softmax activation function with gradient support. -/// Masked positions are filled with zeros in the output. -pub fn masked_softmax(tensor: &Tensor, mask: &Tensor, dim: Option) -> Result { - if mask.dtype() != DataType::Bool { - return Err(MinitensorError::invalid_operation( - "masked_softmax mask must have bool dtype", - )); - } - - if tensor.device() != mask.device() { - return Err(MinitensorError::device_mismatch( - format!("{:?}", tensor.device()), - format!("{:?}", mask.device()), - )); - } - - let broadcast_shape = mask.shape().broadcast_with(tensor.shape())?; - if &broadcast_shape != tensor.shape() { - return Err(MinitensorError::shape_mismatch( - mask.shape().dims().to_vec(), - tensor.shape().dims().to_vec(), - )); - } - - if tensor.ndim() == 0 { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - let mask_value = mask - .data() - .as_bool_slice() - .ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from mask tensor") - })? - .first() - .copied() - .unwrap_or(false); - match tensor.dtype() { - DataType::Float32 => { - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from output data", - ) - })?; - output_slice[0] = if mask_value { 0.0 } else { 1.0 }; - } - DataType::Float64 => { - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from output data", - ) - })?; - output_slice[0] = if mask_value { 0.0 } else { 1.0 }; - } - _ => { - return Err(MinitensorError::invalid_operation( - "masked_softmax only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(SoftmaxBackward { - input_id: tensor.id(), - output: output.detach(), - dim: 0, - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - return Ok(output_with_grad); - } - - return Ok(output); - } - - let dim = dim.unwrap_or(tensor.ndim() - 1); - - if dim >= tensor.ndim() { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => masked_softmax_f32(tensor, mask, &mut output_data, dim)?, - DataType::Float64 => masked_softmax_f64(tensor, mask, &mut output_data, dim)?, - _ => { - return Err(MinitensorError::invalid_operation( - "masked_softmax only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(SoftmaxBackward { - input_id: tensor.id(), - output: output.detach(), - dim, - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Masked log-softmax activation function with gradient support. -/// Masked positions are filled with -inf in the output. -pub fn masked_log_softmax(tensor: &Tensor, mask: &Tensor, dim: Option) -> Result { - if mask.dtype() != DataType::Bool { - return Err(MinitensorError::invalid_operation( - "masked_log_softmax mask must have bool dtype", - )); - } - - if tensor.device() != mask.device() { - return Err(MinitensorError::device_mismatch( - format!("{:?}", tensor.device()), - format!("{:?}", mask.device()), - )); - } - - let broadcast_shape = mask.shape().broadcast_with(tensor.shape())?; - if &broadcast_shape != tensor.shape() { - return Err(MinitensorError::shape_mismatch( - mask.shape().dims().to_vec(), - tensor.shape().dims().to_vec(), - )); - } - - if tensor.ndim() == 0 { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - let mask_value = mask - .data() - .as_bool_slice() - .ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from mask tensor") - })? - .first() - .copied() - .unwrap_or(false); - match tensor.dtype() { - DataType::Float32 => { - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from output data", - ) - })?; - output_slice[0] = if mask_value { f32::NEG_INFINITY } else { 0.0 }; - } - DataType::Float64 => { - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from output data", - ) - })?; - output_slice[0] = if mask_value { f64::NEG_INFINITY } else { 0.0 }; - } - _ => { - return Err(MinitensorError::invalid_operation( - "masked_log_softmax only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(MaskedLogSoftmaxBackward { - input_id: tensor.id(), - output: output.detach(), - mask: mask.detach(), - dim: 0, - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - return Ok(output_with_grad); - } - - return Ok(output); - } - - let dim = dim.unwrap_or(tensor.ndim() - 1); - - if dim >= tensor.ndim() { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => masked_log_softmax_f32(tensor, mask, &mut output_data, dim)?, - DataType::Float64 => masked_log_softmax_f64(tensor, mask, &mut output_data, dim)?, - _ => { - return Err(MinitensorError::invalid_operation( - "masked_log_softmax only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(MaskedLogSoftmaxBackward { - input_id: tensor.id(), - output: output.detach(), - mask: mask.detach(), - dim, - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - - Ok(output_with_grad) - } else { - Ok(output) - } -} - -// Helper functions for type-specific operations - -fn exp_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - unary_apply(input_data, output_slice, f32::exp); - Ok(()) -} - -fn exp_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - unary_apply(input_data, output_slice, f64::exp); - Ok(()) -} - -fn log_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - unary_apply(input_data, output_slice, f32::ln); - Ok(()) -} - -fn log_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - unary_apply(input_data, output_slice, f64::ln); - Ok(()) -} - -fn log1p_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - unary_apply(input_data, output_slice, |val: f32| { - if val == -1.0 { - f32::NEG_INFINITY - } else if val < -1.0 { - f32::NAN - } else { - val.ln_1p() - } - }); - Ok(()) -} - -fn log1p_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - unary_apply(input_data, output_slice, |val: f64| { - if val == -1.0 { - f64::NEG_INFINITY - } else if val < -1.0 { - f64::NAN - } else { - val.ln_1p() - } - }); - Ok(()) -} - -fn expm1_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - unary_apply(input_data, output_slice, f32::exp_m1); - Ok(()) -} - -fn expm1_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - unary_apply(input_data, output_slice, f64::exp_m1); - Ok(()) -} - -fn sin_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - unary_apply(input_data, output_slice, f32::sin); - Ok(()) -} - -fn sin_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - unary_apply(input_data, output_slice, f64::sin); - Ok(()) -} - -fn cos_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - unary_apply(input_data, output_slice, f32::cos); - Ok(()) -} - -fn cos_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - unary_apply(input_data, output_slice, f64::cos); - Ok(()) -} - -fn tan_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - unary_apply(input_data, output_slice, f32::tan); - Ok(()) -} - -fn tan_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - unary_apply(input_data, output_slice, f64::tan); - Ok(()) -} - -fn asin_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - unary_apply(input_data, output_slice, f32::asin); - Ok(()) -} - -fn asin_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - unary_apply(input_data, output_slice, f64::asin); - Ok(()) -} - -fn acos_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - unary_apply(input_data, output_slice, f32::acos); - Ok(()) -} - -fn acos_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - unary_apply(input_data, output_slice, f64::acos); - Ok(()) -} - -fn atan_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - unary_apply(input_data, output_slice, f32::atan); - Ok(()) -} - -fn atan_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - unary_apply(input_data, output_slice, f64::atan); - Ok(()) -} - -fn sinh_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - unary_apply(input_data, output_slice, f32::sinh); - Ok(()) -} - -fn sinh_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - unary_apply(input_data, output_slice, f64::sinh); - Ok(()) -} - -fn cosh_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - unary_apply(input_data, output_slice, f32::cosh); - Ok(()) -} - -fn cosh_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - unary_apply(input_data, output_slice, f64::cosh); - Ok(()) -} - -fn asinh_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - unary_apply(input_data, output_slice, f32::asinh); - Ok(()) -} - -fn asinh_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - unary_apply(input_data, output_slice, f64::asinh); - Ok(()) -} - -fn acosh_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - unary_apply(input_data, output_slice, f32::acosh); - Ok(()) -} - -fn acosh_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - unary_apply(input_data, output_slice, f64::acosh); - Ok(()) -} - -fn atanh_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - unary_apply(input_data, output_slice, f32::atanh); - Ok(()) -} - -fn atanh_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - unary_apply(input_data, output_slice, f64::atanh); - Ok(()) -} - -fn softplus_f32( - tensor: &Tensor, - output_data: &mut TensorData, - beta: f32, - threshold: f32, -) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - unary_apply(input_data, output_slice, |val: f32| { - let scaled = beta * val; - if scaled > threshold { - val - } else { - scaled.exp().ln_1p() / beta - } - }); - Ok(()) -} - -fn softplus_f64( - tensor: &Tensor, - output_data: &mut TensorData, - beta: f64, - threshold: f64, -) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - unary_apply(input_data, output_slice, |val: f64| { - let scaled = beta * val; - if scaled > threshold { - val - } else { - scaled.exp().ln_1p() / beta - } - }); - Ok(()) -} - -fn gelu_f32(tensor: &Tensor, output_data: &mut TensorData, approximate: bool) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - if approximate { - let coeff = (2.0f32 / std::f32::consts::PI).sqrt(); - unary_apply(input_data, output_slice, |x: f32| { - let x3 = x * x * x; - let inner = coeff * (x + 0.044715f32 * x3); - 0.5f32 * x * (1.0f32 + inner.tanh()) - }); - } else { - let inv_sqrt_2 = std::f32::consts::FRAC_1_SQRT_2; - unary_apply(input_data, output_slice, |x: f32| { - 0.5f32 * x * (1.0f32 + erff(x * inv_sqrt_2)) - }); - } - Ok(()) -} - -fn gelu_f64(tensor: &Tensor, output_data: &mut TensorData, approximate: bool) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - if approximate { - let coeff = (2.0f64 / std::f64::consts::PI).sqrt(); - unary_apply(input_data, output_slice, |x: f64| { - let x3 = x * x * x; - let inner = coeff * (x + 0.044715f64 * x3); - 0.5f64 * x * (1.0f64 + inner.tanh()) - }); - } else { - let inv_sqrt_2 = std::f64::consts::FRAC_1_SQRT_2; - unary_apply(input_data, output_slice, |x: f64| { - 0.5f64 * x * (1.0f64 + erf(x * inv_sqrt_2)) - }); - } - Ok(()) -} - -fn elu_f32(tensor: &Tensor, output_data: &mut TensorData, alpha: f32) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - unary_apply(input_data, output_slice, |x: f32| { - if x > 0.0 { x } else { alpha * (x.exp() - 1.0) } - }); - Ok(()) -} - -fn elu_f64(tensor: &Tensor, output_data: &mut TensorData, alpha: f64) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - unary_apply(input_data, output_slice, |x: f64| { - if x > 0.0 { x } else { alpha * (x.exp() - 1.0) } - }); - Ok(()) -} - -fn selu_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - const ALPHA: f32 = 1.6732632; - const SCALE: f32 = 1.050701; - unary_apply(input_data, output_slice, |x: f32| { - if x > 0.0 { - SCALE * x - } else { - SCALE * ALPHA * (x.exp() - 1.0) - } - }); - Ok(()) -} - -fn selu_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - const ALPHA: f64 = 1.6732632423543772848170429916717; - const SCALE: f64 = 1.0507009873554804934193349852946; - unary_apply(input_data, output_slice, |x: f64| { - if x > 0.0 { - SCALE * x - } else { - SCALE * ALPHA * (x.exp() - 1.0) - } - }); - Ok(()) -} - -fn silu_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - unary_apply(input_data, output_slice, |x: f32| { - let sigmoid = 1.0 / (1.0 + (-x).exp()); - x * sigmoid - }); - Ok(()) -} - -fn silu_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - unary_apply(input_data, output_slice, |x: f64| { - let sigmoid = 1.0 / (1.0 + (-x).exp()); - x * sigmoid - }); - Ok(()) -} - -fn softsign_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - unary_apply(input_data, output_slice, |x: f32| { - let denom = 1.0 + x.abs(); - x / denom - }); - Ok(()) -} - -fn softsign_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - unary_apply(input_data, output_slice, |x: f64| { - let denom = 1.0 + x.abs(); - x / denom - }); - Ok(()) -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::autograd::MaskedLogSoftmaxBackward; +use crate::autograd::SoftmaxBackward; +use crate::{ + autograd::add_to_graph, + error::{MinitensorError, Result}, + tensor::{DataType, Tensor, TensorData}, +}; +use libm::{erf, erff}; +use std::sync::Arc; + +/// Masked softmax activation function with gradient support. +/// Masked positions are filled with zeros in the output. +pub fn masked_softmax(tensor: &Tensor, mask: &Tensor, dim: Option) -> Result { + if mask.dtype() != DataType::Bool { + return Err(MinitensorError::invalid_operation( + "masked_softmax mask must have bool dtype", + )); + } + + if tensor.device() != mask.device() { + return Err(MinitensorError::device_mismatch( + format!("{:?}", tensor.device()), + format!("{:?}", mask.device()), + )); + } + + let broadcast_shape = mask.shape().broadcast_with(tensor.shape())?; + if &broadcast_shape != tensor.shape() { + return Err(MinitensorError::shape_mismatch( + mask.shape().dims().to_vec(), + tensor.shape().dims().to_vec(), + )); + } + + if tensor.ndim() == 0 { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + let mask_value = mask + .data() + .as_bool_slice() + .ok_or_else(|| { + MinitensorError::internal_error("Failed to get bool slice from mask tensor") + })? + .first() + .copied() + .unwrap_or(false); + match tensor.dtype() { + DataType::Float32 => { + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from output data", + ) + })?; + output_slice[0] = if mask_value { 0.0 } else { 1.0 }; + } + DataType::Float64 => { + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from output data", + ) + })?; + output_slice[0] = if mask_value { 0.0 } else { 1.0 }; + } + _ => { + return Err(MinitensorError::invalid_operation( + "masked_softmax only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(SoftmaxBackward { + input_id: tensor.id(), + output: output.detach(), + dim: 0, + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + return Ok(output_with_grad); + } + + return Ok(output); + } + + let dim = dim.unwrap_or(tensor.ndim() - 1); + + if dim >= tensor.ndim() { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => masked_softmax_f32(tensor, mask, &mut output_data, dim)?, + DataType::Float64 => masked_softmax_f64(tensor, mask, &mut output_data, dim)?, + _ => { + return Err(MinitensorError::invalid_operation( + "masked_softmax only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(SoftmaxBackward { + input_id: tensor.id(), + output: output.detach(), + dim, + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Masked log-softmax activation function with gradient support. +/// Masked positions are filled with -inf in the output. +pub fn masked_log_softmax(tensor: &Tensor, mask: &Tensor, dim: Option) -> Result { + if mask.dtype() != DataType::Bool { + return Err(MinitensorError::invalid_operation( + "masked_log_softmax mask must have bool dtype", + )); + } + + if tensor.device() != mask.device() { + return Err(MinitensorError::device_mismatch( + format!("{:?}", tensor.device()), + format!("{:?}", mask.device()), + )); + } + + let broadcast_shape = mask.shape().broadcast_with(tensor.shape())?; + if &broadcast_shape != tensor.shape() { + return Err(MinitensorError::shape_mismatch( + mask.shape().dims().to_vec(), + tensor.shape().dims().to_vec(), + )); + } + + if tensor.ndim() == 0 { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + let mask_value = mask + .data() + .as_bool_slice() + .ok_or_else(|| { + MinitensorError::internal_error("Failed to get bool slice from mask tensor") + })? + .first() + .copied() + .unwrap_or(false); + match tensor.dtype() { + DataType::Float32 => { + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from output data", + ) + })?; + output_slice[0] = if mask_value { f32::NEG_INFINITY } else { 0.0 }; + } + DataType::Float64 => { + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from output data", + ) + })?; + output_slice[0] = if mask_value { f64::NEG_INFINITY } else { 0.0 }; + } + _ => { + return Err(MinitensorError::invalid_operation( + "masked_log_softmax only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(MaskedLogSoftmaxBackward { + input_id: tensor.id(), + output: output.detach(), + mask: mask.detach(), + dim: 0, + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + return Ok(output_with_grad); + } + + return Ok(output); + } + + let dim = dim.unwrap_or(tensor.ndim() - 1); + + if dim >= tensor.ndim() { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => masked_log_softmax_f32(tensor, mask, &mut output_data, dim)?, + DataType::Float64 => masked_log_softmax_f64(tensor, mask, &mut output_data, dim)?, + _ => { + return Err(MinitensorError::invalid_operation( + "masked_log_softmax only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(MaskedLogSoftmaxBackward { + input_id: tensor.id(), + output: output.detach(), + mask: mask.detach(), + dim, + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + + Ok(output_with_grad) + } else { + Ok(output) + } +} + +// Helper functions for type-specific operations + +/// Generates a float unary elementwise kernel: fetch the input and output +/// slices for the dtype and apply `$f` element-wise via `unary_apply`. Body is +/// identical to the hand-written wrappers it replaces; only the mapping +/// closure differs per op. +macro_rules! float_unary_kernel { + ($name:ident, $accessor:ident, $accessor_mut:ident, $tyname:literal, $f:expr) => { + pub(crate) fn $name(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { + let input_data = tensor.data().$accessor().ok_or_else(|| { + MinitensorError::internal_error(concat!( + "Failed to get ", + $tyname, + " slice from input tensor" + )) + })?; + let output_slice = output_data.$accessor_mut().ok_or_else(|| { + MinitensorError::internal_error(concat!( + "Failed to get mutable ", + $tyname, + " slice from output data" + )) + })?; + unary_apply(input_data, output_slice, $f); + Ok(()) + } + }; +} + +float_unary_kernel!(exp_f32, as_f32_slice, as_f32_slice_mut, "f32", f32::exp); + +float_unary_kernel!(exp_f64, as_f64_slice, as_f64_slice_mut, "f64", f64::exp); + +float_unary_kernel!(log_f32, as_f32_slice, as_f32_slice_mut, "f32", f32::ln); + +float_unary_kernel!(log_f64, as_f64_slice, as_f64_slice_mut, "f64", f64::ln); + +float_unary_kernel!( + log1p_f32, + as_f32_slice, + as_f32_slice_mut, + "f32", + |val: f32| { + if val == -1.0 { + f32::NEG_INFINITY + } else if val < -1.0 { + f32::NAN + } else { + val.ln_1p() + } + } +); + +float_unary_kernel!( + log1p_f64, + as_f64_slice, + as_f64_slice_mut, + "f64", + |val: f64| { + if val == -1.0 { + f64::NEG_INFINITY + } else if val < -1.0 { + f64::NAN + } else { + val.ln_1p() + } + } +); + +float_unary_kernel!( + expm1_f32, + as_f32_slice, + as_f32_slice_mut, + "f32", + f32::exp_m1 +); + +float_unary_kernel!( + expm1_f64, + as_f64_slice, + as_f64_slice_mut, + "f64", + f64::exp_m1 +); + +float_unary_kernel!(sin_f32, as_f32_slice, as_f32_slice_mut, "f32", f32::sin); + +float_unary_kernel!(sin_f64, as_f64_slice, as_f64_slice_mut, "f64", f64::sin); + +float_unary_kernel!(cos_f32, as_f32_slice, as_f32_slice_mut, "f32", f32::cos); + +float_unary_kernel!(cos_f64, as_f64_slice, as_f64_slice_mut, "f64", f64::cos); + +float_unary_kernel!(tan_f32, as_f32_slice, as_f32_slice_mut, "f32", f32::tan); + +float_unary_kernel!(tan_f64, as_f64_slice, as_f64_slice_mut, "f64", f64::tan); + +float_unary_kernel!(asin_f32, as_f32_slice, as_f32_slice_mut, "f32", f32::asin); + +float_unary_kernel!(asin_f64, as_f64_slice, as_f64_slice_mut, "f64", f64::asin); + +float_unary_kernel!(acos_f32, as_f32_slice, as_f32_slice_mut, "f32", f32::acos); + +float_unary_kernel!(acos_f64, as_f64_slice, as_f64_slice_mut, "f64", f64::acos); + +float_unary_kernel!(atan_f32, as_f32_slice, as_f32_slice_mut, "f32", f32::atan); + +float_unary_kernel!(atan_f64, as_f64_slice, as_f64_slice_mut, "f64", f64::atan); + +float_unary_kernel!(sinh_f32, as_f32_slice, as_f32_slice_mut, "f32", f32::sinh); + +float_unary_kernel!(sinh_f64, as_f64_slice, as_f64_slice_mut, "f64", f64::sinh); + +float_unary_kernel!(cosh_f32, as_f32_slice, as_f32_slice_mut, "f32", f32::cosh); + +float_unary_kernel!(cosh_f64, as_f64_slice, as_f64_slice_mut, "f64", f64::cosh); + +float_unary_kernel!(asinh_f32, as_f32_slice, as_f32_slice_mut, "f32", f32::asinh); + +float_unary_kernel!(asinh_f64, as_f64_slice, as_f64_slice_mut, "f64", f64::asinh); + +float_unary_kernel!(acosh_f32, as_f32_slice, as_f32_slice_mut, "f32", f32::acosh); + +float_unary_kernel!(acosh_f64, as_f64_slice, as_f64_slice_mut, "f64", f64::acosh); + +float_unary_kernel!(atanh_f32, as_f32_slice, as_f32_slice_mut, "f32", f32::atanh); + +float_unary_kernel!(atanh_f64, as_f64_slice, as_f64_slice_mut, "f64", f64::atanh); + +pub(crate) fn softplus_f32( + tensor: &Tensor, + output_data: &mut TensorData, + beta: f32, + threshold: f32, +) -> Result<()> { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + + unary_apply(input_data, output_slice, |val: f32| { + let scaled = beta * val; + if scaled > threshold { + val + } else { + scaled.exp().ln_1p() / beta + } + }); + Ok(()) +} + +pub(crate) fn softplus_f64( + tensor: &Tensor, + output_data: &mut TensorData, + beta: f64, + threshold: f64, +) -> Result<()> { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + + unary_apply(input_data, output_slice, |val: f64| { + let scaled = beta * val; + if scaled > threshold { + val + } else { + scaled.exp().ln_1p() / beta + } + }); + Ok(()) +} + +pub(crate) fn gelu_f32( + tensor: &Tensor, + output_data: &mut TensorData, + approximate: bool, +) -> Result<()> { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + + if approximate { + let coeff = (2.0f32 / std::f32::consts::PI).sqrt(); + unary_apply(input_data, output_slice, |x: f32| { + let x3 = x * x * x; + let inner = coeff * (x + 0.044715f32 * x3); + 0.5f32 * x * (1.0f32 + inner.tanh()) + }); + } else { + let inv_sqrt_2 = std::f32::consts::FRAC_1_SQRT_2; + unary_apply(input_data, output_slice, |x: f32| { + 0.5f32 * x * (1.0f32 + erff(x * inv_sqrt_2)) + }); + } + Ok(()) +} + +pub(crate) fn gelu_f64( + tensor: &Tensor, + output_data: &mut TensorData, + approximate: bool, +) -> Result<()> { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + + if approximate { + let coeff = (2.0f64 / std::f64::consts::PI).sqrt(); + unary_apply(input_data, output_slice, |x: f64| { + let x3 = x * x * x; + let inner = coeff * (x + 0.044715f64 * x3); + 0.5f64 * x * (1.0f64 + inner.tanh()) + }); + } else { + let inv_sqrt_2 = std::f64::consts::FRAC_1_SQRT_2; + unary_apply(input_data, output_slice, |x: f64| { + 0.5f64 * x * (1.0f64 + erf(x * inv_sqrt_2)) + }); + } + Ok(()) +} + +pub(crate) fn elu_f32(tensor: &Tensor, output_data: &mut TensorData, alpha: f32) -> Result<()> { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + + unary_apply(input_data, output_slice, |x: f32| { + if x > 0.0 { x } else { alpha * (x.exp() - 1.0) } + }); + Ok(()) +} + +pub(crate) fn elu_f64(tensor: &Tensor, output_data: &mut TensorData, alpha: f64) -> Result<()> { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + + unary_apply(input_data, output_slice, |x: f64| { + if x > 0.0 { x } else { alpha * (x.exp() - 1.0) } + }); + Ok(()) +} + +pub(crate) fn selu_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + + const ALPHA: f32 = 1.6732632; + const SCALE: f32 = 1.050701; + unary_apply(input_data, output_slice, |x: f32| { + if x > 0.0 { + SCALE * x + } else { + SCALE * ALPHA * (x.exp() - 1.0) + } + }); + Ok(()) +} + +pub(crate) fn selu_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + + const ALPHA: f64 = 1.6732632423543772848170429916717; + const SCALE: f64 = 1.0507009873554804934193349852946; + unary_apply(input_data, output_slice, |x: f64| { + if x > 0.0 { + SCALE * x + } else { + SCALE * ALPHA * (x.exp() - 1.0) + } + }); + Ok(()) +} + +float_unary_kernel!(silu_f32, as_f32_slice, as_f32_slice_mut, "f32", |x: f32| { + let sigmoid = 1.0 / (1.0 + (-x).exp()); + x * sigmoid +}); + +float_unary_kernel!(silu_f64, as_f64_slice, as_f64_slice_mut, "f64", |x: f64| { + let sigmoid = 1.0 / (1.0 + (-x).exp()); + x * sigmoid +}); + +float_unary_kernel!( + softsign_f32, + as_f32_slice, + as_f32_slice_mut, + "f32", + |x: f32| { + let denom = 1.0 + x.abs(); + x / denom + } +); + +float_unary_kernel!( + softsign_f64, + as_f64_slice, + as_f64_slice_mut, + "f64", + |x: f64| { + let denom = 1.0 + x.abs(); + x / denom + } +); diff --git a/engine/src/operations/activation/power.rs b/engine/src/operations/activation/power.rs index 0154a4f1..3e749c77 100644 --- a/engine/src/operations/activation/power.rs +++ b/engine/src/operations/activation/power.rs @@ -1,750 +1,765 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -/// Absolute value function -pub fn abs(tensor: &Tensor) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => abs_f32(tensor, &mut output_data)?, - DataType::Float64 => abs_f64(tensor, &mut output_data)?, - DataType::Int32 => abs_i32(tensor, &mut output_data)?, - DataType::Int64 => abs_i64(tensor, &mut output_data)?, - DataType::Bool => { - return Err(MinitensorError::invalid_operation( - "Absolute value not supported for boolean tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() && tensor.dtype().is_float() { - let grad_fn = Arc::new(AbsBackward { - input_id: tensor.id(), - input: tensor.detach(), - }); - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - return Ok(output_with_grad); - } - - Ok(output) -} - -/// Element-wise sign function (-1, 0, or 1 depending on value sign) -pub fn sign(tensor: &Tensor) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => sign_f32(tensor, &mut output_data)?, - DataType::Float64 => sign_f64(tensor, &mut output_data)?, - DataType::Int32 => sign_i32(tensor, &mut output_data)?, - DataType::Int64 => sign_i64(tensor, &mut output_data)?, - DataType::Bool => { - return Err(MinitensorError::invalid_operation( - "Sign operation not supported for boolean tensors", - )); - } - } - - Ok(Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )) -} - -/// Square root function -pub fn sqrt(tensor: &Tensor) -> Result { - // Use powf implementation for gradient support: sqrt(x) = x.powf(0.5) - powf(tensor, 0.5) -} - -/// Reciprocal square root function -pub fn rsqrt(tensor: &Tensor) -> Result { - // Use powf implementation for gradient support: rsqrt(x) = x.powf(-0.5) - powf(tensor, -0.5) -} - -/// Element-wise reciprocal (1/x) with gradient support -pub fn reciprocal(tensor: &Tensor) -> Result { - match tensor.dtype() { - DataType::Float32 | DataType::Float64 => powf(tensor, -1.0), - _ => Err(MinitensorError::invalid_operation( - "Reciprocal only supported for floating point tensors", - )), - } -} - -/// Clip tensor values to range -pub fn clip(tensor: &Tensor, min_val: Option, max_val: Option) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => clip_f32(tensor, &mut output_data, min_val, max_val)?, - DataType::Float64 => clip_f64(tensor, &mut output_data, min_val, max_val)?, - DataType::Int32 => clip_i32(tensor, &mut output_data, min_val, max_val)?, - DataType::Int64 => clip_i64(tensor, &mut output_data, min_val, max_val)?, - DataType::Bool => { - return Err(MinitensorError::invalid_operation( - "Clip not supported for boolean tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() && tensor.dtype().is_float() { - let grad_fn = Arc::new(ClampBackward { - input_id: tensor.id(), - input: tensor.detach(), - min: min_val, - max: max_val, - }); - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - return Ok(output_with_grad); - } - - Ok(output) -} - - -/// Replace NaN and infinity values in floating point tensors. -/// -/// Exact tensors cannot contain NaN or infinity, so they are returned unchanged. -pub fn nan_to_num( - tensor: &Tensor, - nan: f64, - posinf: Option, - neginf: Option, -) -> Result { - match tensor.dtype() { - DataType::Float32 | DataType::Float64 => {} - DataType::Int32 | DataType::Int64 | DataType::Bool => return Ok(tensor.clone()), - } - - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - let mut finite_mask = tensor.requires_grad().then(|| Vec::with_capacity(tensor.numel())); - - match tensor.dtype() { - DataType::Float32 => nan_to_num_f32( - tensor, - &mut output_data, - nan, - posinf, - neginf, - finite_mask.as_mut(), - )?, - DataType::Float64 => nan_to_num_f64( - tensor, - &mut output_data, - nan, - posinf, - neginf, - finite_mask.as_mut(), - )?, - DataType::Int32 | DataType::Int64 | DataType::Bool => unreachable!(), - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(NanToNumBackward { - input_id: tensor.id(), - finite_mask: finite_mask.unwrap_or_default(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Round tensor values -pub fn round(tensor: &Tensor, decimals: i32) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => round_f32(tensor, &mut output_data, decimals)?, - DataType::Float64 => round_f64(tensor, &mut output_data, decimals)?, - _ => { - return Err(MinitensorError::invalid_operation( - "Round only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - Ok(output) -} - -/// Floor tensor values -pub fn floor(tensor: &Tensor) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => floor_f32(tensor, &mut output_data)?, - DataType::Float64 => floor_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "Floor only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - Ok(output) -} - -/// Ceiling tensor values -pub fn ceil(tensor: &Tensor) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => ceil_f32(tensor, &mut output_data)?, - DataType::Float64 => ceil_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "Ceiling only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - Ok(output) -} - -// Helper functions for the new operations - -fn abs_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - unary_apply(input_data, output_slice, |v: f32| v.abs()); - Ok(()) -} - -fn abs_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - unary_apply(input_data, output_slice, |v: f64| v.abs()); - Ok(()) -} - -fn abs_i32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from input tensor") - })?; - - let output_slice = output_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i32 slice from output data") - })?; - - unary_apply(input_data, output_slice, |v: i32| v.abs()); - Ok(()) -} - -fn abs_i64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from input tensor") - })?; - - let output_slice = output_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice from output data") - })?; - - unary_apply(input_data, output_slice, |v: i64| v.abs()); - Ok(()) -} - -fn sign_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - unary_apply(input_data, output_slice, |v: f32| { - if v > 0.0 { - 1.0 - } else if v < 0.0 { - -1.0 - } else { - 0.0 - } - }); - Ok(()) -} - -fn sign_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - unary_apply(input_data, output_slice, |v: f64| { - if v > 0.0 { - 1.0 - } else if v < 0.0 { - -1.0 - } else { - 0.0 - } - }); - Ok(()) -} - -fn sign_i32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from input tensor") - })?; - - let output_slice = output_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i32 slice from output data") - })?; - - unary_apply(input_data, output_slice, |v: i32| { - if v > 0 { - 1 - } else if v < 0 { - -1 - } else { - 0 - } - }); - Ok(()) -} - -fn sign_i64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from input tensor") - })?; - - let output_slice = output_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice from output data") - })?; - - unary_apply(input_data, output_slice, |v: i64| { - if v > 0 { - 1 - } else if v < 0 { - -1 - } else { - 0 - } - }); - Ok(()) -} - -fn clip_f32( - tensor: &Tensor, - output_data: &mut TensorData, - min_val: Option, - max_val: Option, -) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - let min_f32 = min_val.map(|v| v as f32); - let max_f32 = max_val.map(|v| v as f32); - unary_apply(input_data, output_slice, |val: f32| { - let mut v = val; - if let Some(min) = min_f32 { - v = v.max(min); - } - if let Some(max) = max_f32 { - v = v.min(max); - } - v - }); - Ok(()) -} - -fn clip_f64( - tensor: &Tensor, - output_data: &mut TensorData, - min_val: Option, - max_val: Option, -) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - unary_apply(input_data, output_slice, |val: f64| { - let mut v = val; - if let Some(min) = min_val { - v = v.max(min); - } - if let Some(max) = max_val { - v = v.min(max); - } - v - }); - Ok(()) -} - -fn clip_i32( - tensor: &Tensor, - output_data: &mut TensorData, - min_val: Option, - max_val: Option, -) -> Result<()> { - let input_data = tensor.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from input tensor") - })?; - - let output_slice = output_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i32 slice from output data") - })?; - - let min_i32 = min_val.map(|v| v as i32); - let max_i32 = max_val.map(|v| v as i32); - unary_apply(input_data, output_slice, |val: i32| { - let mut v = val; - if let Some(min) = min_i32 { - v = v.max(min); - } - if let Some(max) = max_i32 { - v = v.min(max); - } - v - }); - Ok(()) -} - -fn clip_i64( - tensor: &Tensor, - output_data: &mut TensorData, - min_val: Option, - max_val: Option, -) -> Result<()> { - let input_data = tensor.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from input tensor") - })?; - - let output_slice = output_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice from output data") - })?; - - let min_i64 = min_val.map(|v| v as i64); - let max_i64 = max_val.map(|v| v as i64); - unary_apply(input_data, output_slice, |val: i64| { - let mut v = val; - if let Some(min) = min_i64 { - v = v.max(min); - } - if let Some(max) = max_i64 { - v = v.min(max); - } - v - }); - Ok(()) -} - - -fn nan_to_num_f32( - tensor: &Tensor, - output_data: &mut TensorData, - nan: f64, - posinf: Option, - neginf: Option, - finite_mask: Option<&mut Vec>, -) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - replace_non_finite( - input_data, - output_slice, - nan as f32, - posinf.map_or(f32::MAX, |v| v as f32), - neginf.map_or(f32::MIN, |v| v as f32), - finite_mask, - ); - Ok(()) -} - -fn nan_to_num_f64( - tensor: &Tensor, - output_data: &mut TensorData, - nan: f64, - posinf: Option, - neginf: Option, - finite_mask: Option<&mut Vec>, -) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - replace_non_finite( - input_data, - output_slice, - nan, - posinf.unwrap_or(f64::MAX), - neginf.unwrap_or(f64::MIN), - finite_mask, - ); - Ok(()) -} - -fn replace_non_finite( - input: &[T], - output: &mut [T], - nan: T, - posinf: T, - neginf: T, - finite_mask: Option<&mut Vec>, -) where - T: Copy + PartialEq + Send + Sync, - T: FloatClassify, -{ - match finite_mask { - Some(mask) => { - mask.resize(input.len(), false); - if input.len() < PAR_THRESHOLD { - for ((&val, out), is_finite) in input - .iter() - .zip(output.iter_mut()) - .zip(mask.iter_mut()) - { - *is_finite = val.is_finite_value(); - *out = classify_nan_to_num(val, nan, posinf, neginf); - } - } else { - input - .par_iter() - .zip(output.par_iter_mut()) - .zip(mask.par_iter_mut()) - .for_each(|((val, out), is_finite)| { - *is_finite = val.is_finite_value(); - *out = classify_nan_to_num(*val, nan, posinf, neginf); - }); - } - } - None => unary_apply(input, output, |val| classify_nan_to_num(val, nan, posinf, neginf)), - } -} - -#[inline(always)] -fn classify_nan_to_num(val: T, nan: T, posinf: T, neginf: T) -> T -where - T: Copy + PartialEq + FloatClassify, -{ - if val.is_nan_value() { - nan - } else if val.is_positive_infinity() { - posinf - } else if val.is_negative_infinity() { - neginf - } else { - val - } -} - -trait FloatClassify { - fn is_nan_value(self) -> bool; - fn is_finite_value(self) -> bool; - fn is_positive_infinity(self) -> bool; - fn is_negative_infinity(self) -> bool; -} - -impl FloatClassify for f32 { - #[inline(always)] - fn is_nan_value(self) -> bool { - self.is_nan() - } - - #[inline(always)] - fn is_finite_value(self) -> bool { - self.is_finite() - } - - #[inline(always)] - fn is_positive_infinity(self) -> bool { - self == f32::INFINITY - } - - #[inline(always)] - fn is_negative_infinity(self) -> bool { - self == f32::NEG_INFINITY - } -} - -impl FloatClassify for f64 { - #[inline(always)] - fn is_nan_value(self) -> bool { - self.is_nan() - } - - #[inline(always)] - fn is_finite_value(self) -> bool { - self.is_finite() - } - - #[inline(always)] - fn is_positive_infinity(self) -> bool { - self == f64::INFINITY - } - - #[inline(always)] - fn is_negative_infinity(self) -> bool { - self == f64::NEG_INFINITY - } -} - -fn round_f32(tensor: &Tensor, output_data: &mut TensorData, decimals: i32) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - let multiplier = 10.0_f32.powi(decimals); - unary_apply(input_data, output_slice, |val: f32| { - (val * multiplier).round() / multiplier - }); - Ok(()) -} - -fn round_f64(tensor: &Tensor, output_data: &mut TensorData, decimals: i32) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - let multiplier = 10.0_f64.powi(decimals); - unary_apply(input_data, output_slice, |val: f64| { - (val * multiplier).round() / multiplier - }); - Ok(()) -} - -fn floor_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - unary_apply(input_data, output_slice, f32::floor); - Ok(()) -} - -fn floor_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - unary_apply(input_data, output_slice, f64::floor); - Ok(()) -} - -fn ceil_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - unary_apply(input_data, output_slice, f32::ceil); - Ok(()) -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::autograd::AbsBackward; +use crate::autograd::ClampBackward; +use crate::autograd::NanToNumBackward; +use crate::{ + autograd::add_to_graph, + error::{MinitensorError, Result}, + tensor::{DataType, Tensor, TensorData}, +}; +use rayon::prelude::*; +use std::sync::Arc; + +/// Absolute value function +pub fn abs(tensor: &Tensor) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => abs_f32(tensor, &mut output_data)?, + DataType::Float64 => abs_f64(tensor, &mut output_data)?, + DataType::Int32 => abs_i32(tensor, &mut output_data)?, + DataType::Int64 => abs_i64(tensor, &mut output_data)?, + DataType::Bool => { + return Err(MinitensorError::invalid_operation( + "Absolute value not supported for boolean tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() && tensor.dtype().is_float() { + let grad_fn = Arc::new(AbsBackward { + input_id: tensor.id(), + input: tensor.detach(), + }); + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + return Ok(output_with_grad); + } + + Ok(output) +} + +/// Element-wise sign function (-1, 0, or 1 depending on value sign) +pub fn sign(tensor: &Tensor) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => sign_f32(tensor, &mut output_data)?, + DataType::Float64 => sign_f64(tensor, &mut output_data)?, + DataType::Int32 => sign_i32(tensor, &mut output_data)?, + DataType::Int64 => sign_i64(tensor, &mut output_data)?, + DataType::Bool => { + return Err(MinitensorError::invalid_operation( + "Sign operation not supported for boolean tensors", + )); + } + } + + Ok(Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )) +} + +/// Square root function +pub fn sqrt(tensor: &Tensor) -> Result { + // Use powf implementation for gradient support: sqrt(x) = x.powf(0.5) + powf(tensor, 0.5) +} + +/// Reciprocal square root function +pub fn rsqrt(tensor: &Tensor) -> Result { + // Use powf implementation for gradient support: rsqrt(x) = x.powf(-0.5) + powf(tensor, -0.5) +} + +/// Element-wise reciprocal (1/x) with gradient support +pub fn reciprocal(tensor: &Tensor) -> Result { + match tensor.dtype() { + DataType::Float32 | DataType::Float64 => powf(tensor, -1.0), + _ => Err(MinitensorError::invalid_operation( + "Reciprocal only supported for floating point tensors", + )), + } +} + +/// Clip tensor values to range +pub fn clip(tensor: &Tensor, min_val: Option, max_val: Option) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => clip_f32(tensor, &mut output_data, min_val, max_val)?, + DataType::Float64 => clip_f64(tensor, &mut output_data, min_val, max_val)?, + DataType::Int32 => clip_i32(tensor, &mut output_data, min_val, max_val)?, + DataType::Int64 => clip_i64(tensor, &mut output_data, min_val, max_val)?, + DataType::Bool => { + return Err(MinitensorError::invalid_operation( + "Clip not supported for boolean tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() && tensor.dtype().is_float() { + let grad_fn = Arc::new(ClampBackward { + input_id: tensor.id(), + input: tensor.detach(), + min: min_val, + max: max_val, + }); + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + return Ok(output_with_grad); + } + + Ok(output) +} + +/// Replace NaN and infinity values in floating point tensors. +/// +/// Exact tensors cannot contain NaN or infinity, so they are returned unchanged. +pub fn nan_to_num( + tensor: &Tensor, + nan: f64, + posinf: Option, + neginf: Option, +) -> Result { + match tensor.dtype() { + DataType::Float32 | DataType::Float64 => {} + DataType::Int32 | DataType::Int64 | DataType::Bool => return Ok(tensor.clone()), + } + + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + let mut finite_mask = tensor + .requires_grad() + .then(|| Vec::with_capacity(tensor.numel())); + + match tensor.dtype() { + DataType::Float32 => nan_to_num_f32( + tensor, + &mut output_data, + nan, + posinf, + neginf, + finite_mask.as_mut(), + )?, + DataType::Float64 => nan_to_num_f64( + tensor, + &mut output_data, + nan, + posinf, + neginf, + finite_mask.as_mut(), + )?, + DataType::Int32 | DataType::Int64 | DataType::Bool => unreachable!(), + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(NanToNumBackward { + input_id: tensor.id(), + finite_mask: finite_mask.unwrap_or_default(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Round tensor values +pub fn round(tensor: &Tensor, decimals: i32) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => round_f32(tensor, &mut output_data, decimals)?, + DataType::Float64 => round_f64(tensor, &mut output_data, decimals)?, + _ => { + return Err(MinitensorError::invalid_operation( + "Round only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + Ok(output) +} + +/// Floor tensor values +pub fn floor(tensor: &Tensor) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => floor_f32(tensor, &mut output_data)?, + DataType::Float64 => floor_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "Floor only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + Ok(output) +} + +/// Ceiling tensor values +pub fn ceil(tensor: &Tensor) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => ceil_f32(tensor, &mut output_data)?, + DataType::Float64 => ceil_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "Ceiling only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + Ok(output) +} + +// Helper functions for the new operations + +fn abs_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + + unary_apply(input_data, output_slice, |v: f32| v.abs()); + Ok(()) +} + +fn abs_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + + unary_apply(input_data, output_slice, |v: f64| v.abs()); + Ok(()) +} + +fn abs_i32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { + let input_data = tensor.data().as_i32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i32 slice from input tensor") + })?; + + let output_slice = output_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i32 slice from output data") + })?; + + unary_apply(input_data, output_slice, |v: i32| v.abs()); + Ok(()) +} + +fn abs_i64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { + let input_data = tensor.data().as_i64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i64 slice from input tensor") + })?; + + let output_slice = output_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice from output data") + })?; + + unary_apply(input_data, output_slice, |v: i64| v.abs()); + Ok(()) +} + +fn sign_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + + unary_apply(input_data, output_slice, |v: f32| { + if v > 0.0 { + 1.0 + } else if v < 0.0 { + -1.0 + } else { + 0.0 + } + }); + Ok(()) +} + +fn sign_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + + unary_apply(input_data, output_slice, |v: f64| { + if v > 0.0 { + 1.0 + } else if v < 0.0 { + -1.0 + } else { + 0.0 + } + }); + Ok(()) +} + +fn sign_i32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { + let input_data = tensor.data().as_i32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i32 slice from input tensor") + })?; + + let output_slice = output_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i32 slice from output data") + })?; + + unary_apply(input_data, output_slice, |v: i32| { + if v > 0 { + 1 + } else if v < 0 { + -1 + } else { + 0 + } + }); + Ok(()) +} + +fn sign_i64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { + let input_data = tensor.data().as_i64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i64 slice from input tensor") + })?; + + let output_slice = output_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice from output data") + })?; + + unary_apply(input_data, output_slice, |v: i64| { + if v > 0 { + 1 + } else if v < 0 { + -1 + } else { + 0 + } + }); + Ok(()) +} + +fn clip_f32( + tensor: &Tensor, + output_data: &mut TensorData, + min_val: Option, + max_val: Option, +) -> Result<()> { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + + let min_f32 = min_val.map(|v| v as f32); + let max_f32 = max_val.map(|v| v as f32); + unary_apply(input_data, output_slice, |val: f32| { + let mut v = val; + if let Some(min) = min_f32 { + v = v.max(min); + } + if let Some(max) = max_f32 { + v = v.min(max); + } + v + }); + Ok(()) +} + +fn clip_f64( + tensor: &Tensor, + output_data: &mut TensorData, + min_val: Option, + max_val: Option, +) -> Result<()> { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + unary_apply(input_data, output_slice, |val: f64| { + let mut v = val; + if let Some(min) = min_val { + v = v.max(min); + } + if let Some(max) = max_val { + v = v.min(max); + } + v + }); + Ok(()) +} + +fn clip_i32( + tensor: &Tensor, + output_data: &mut TensorData, + min_val: Option, + max_val: Option, +) -> Result<()> { + let input_data = tensor.data().as_i32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i32 slice from input tensor") + })?; + + let output_slice = output_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i32 slice from output data") + })?; + + let min_i32 = min_val.map(|v| v as i32); + let max_i32 = max_val.map(|v| v as i32); + unary_apply(input_data, output_slice, |val: i32| { + let mut v = val; + if let Some(min) = min_i32 { + v = v.max(min); + } + if let Some(max) = max_i32 { + v = v.min(max); + } + v + }); + Ok(()) +} + +fn clip_i64( + tensor: &Tensor, + output_data: &mut TensorData, + min_val: Option, + max_val: Option, +) -> Result<()> { + let input_data = tensor.data().as_i64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i64 slice from input tensor") + })?; + + let output_slice = output_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice from output data") + })?; + + let min_i64 = min_val.map(|v| v as i64); + let max_i64 = max_val.map(|v| v as i64); + unary_apply(input_data, output_slice, |val: i64| { + let mut v = val; + if let Some(min) = min_i64 { + v = v.max(min); + } + if let Some(max) = max_i64 { + v = v.min(max); + } + v + }); + Ok(()) +} + +fn nan_to_num_f32( + tensor: &Tensor, + output_data: &mut TensorData, + nan: f64, + posinf: Option, + neginf: Option, + finite_mask: Option<&mut Vec>, +) -> Result<()> { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + + replace_non_finite( + input_data, + output_slice, + nan as f32, + posinf.map_or(f32::MAX, |v| v as f32), + neginf.map_or(f32::MIN, |v| v as f32), + finite_mask, + ); + Ok(()) +} + +fn nan_to_num_f64( + tensor: &Tensor, + output_data: &mut TensorData, + nan: f64, + posinf: Option, + neginf: Option, + finite_mask: Option<&mut Vec>, +) -> Result<()> { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + + replace_non_finite( + input_data, + output_slice, + nan, + posinf.unwrap_or(f64::MAX), + neginf.unwrap_or(f64::MIN), + finite_mask, + ); + Ok(()) +} + +fn replace_non_finite( + input: &[T], + output: &mut [T], + nan: T, + posinf: T, + neginf: T, + finite_mask: Option<&mut Vec>, +) where + T: Copy + PartialEq + Send + Sync, + T: FloatClassify, +{ + match finite_mask { + Some(mask) => { + mask.resize(input.len(), false); + if input.len() < PAR_THRESHOLD { + for ((&val, out), is_finite) in + input.iter().zip(output.iter_mut()).zip(mask.iter_mut()) + { + *is_finite = val.is_finite_value(); + *out = classify_nan_to_num(val, nan, posinf, neginf); + } + } else { + input + .par_iter() + .zip(output.par_iter_mut()) + .zip(mask.par_iter_mut()) + .for_each(|((val, out), is_finite)| { + *is_finite = val.is_finite_value(); + *out = classify_nan_to_num(*val, nan, posinf, neginf); + }); + } + } + None => unary_apply(input, output, |val| { + classify_nan_to_num(val, nan, posinf, neginf) + }), + } +} + +#[inline(always)] +fn classify_nan_to_num(val: T, nan: T, posinf: T, neginf: T) -> T +where + T: Copy + PartialEq + FloatClassify, +{ + if val.is_nan_value() { + nan + } else if val.is_positive_infinity() { + posinf + } else if val.is_negative_infinity() { + neginf + } else { + val + } +} + +// Copy-scalar helpers mirroring `f32::is_nan` and friends, which also take +// `self` by value. +#[allow(clippy::wrong_self_convention)] +trait FloatClassify { + fn is_nan_value(self) -> bool; + fn is_finite_value(self) -> bool; + fn is_positive_infinity(self) -> bool; + fn is_negative_infinity(self) -> bool; +} + +impl FloatClassify for f32 { + #[inline(always)] + fn is_nan_value(self) -> bool { + self.is_nan() + } + + #[inline(always)] + fn is_finite_value(self) -> bool { + self.is_finite() + } + + #[inline(always)] + fn is_positive_infinity(self) -> bool { + self == f32::INFINITY + } + + #[inline(always)] + fn is_negative_infinity(self) -> bool { + self == f32::NEG_INFINITY + } +} + +impl FloatClassify for f64 { + #[inline(always)] + fn is_nan_value(self) -> bool { + self.is_nan() + } + + #[inline(always)] + fn is_finite_value(self) -> bool { + self.is_finite() + } + + #[inline(always)] + fn is_positive_infinity(self) -> bool { + self == f64::INFINITY + } + + #[inline(always)] + fn is_negative_infinity(self) -> bool { + self == f64::NEG_INFINITY + } +} + +fn round_f32(tensor: &Tensor, output_data: &mut TensorData, decimals: i32) -> Result<()> { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + + let multiplier = 10.0_f32.powi(decimals); + unary_apply(input_data, output_slice, |val: f32| { + (val * multiplier).round() / multiplier + }); + Ok(()) +} + +fn round_f64(tensor: &Tensor, output_data: &mut TensorData, decimals: i32) -> Result<()> { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + + let multiplier = 10.0_f64.powi(decimals); + unary_apply(input_data, output_slice, |val: f64| { + (val * multiplier).round() / multiplier + }); + Ok(()) +} + +fn floor_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + unary_apply(input_data, output_slice, f32::floor); + Ok(()) +} + +fn floor_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + unary_apply(input_data, output_slice, f64::floor); + Ok(()) +} + +fn ceil_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + unary_apply(input_data, output_slice, f32::ceil); + Ok(()) +} diff --git a/engine/src/operations/activation/softmax.rs b/engine/src/operations/activation/softmax.rs index f73f3b32..5ec41a4e 100644 --- a/engine/src/operations/activation/softmax.rs +++ b/engine/src/operations/activation/softmax.rs @@ -1,1170 +1,1187 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -use num_traits::Float; - -fn logaddexp_f32( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from rhs tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - crate::operations::arithmetic::broadcast_binary_op( - lhs_data, - rhs_data, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - |a, b| { - if a.is_nan() || b.is_nan() { - f32::NAN - } else { - let max = a.max(b); - if max.is_infinite() { - max - } else { - let exp_a = (a - max).exp(); - let exp_b = (b - max).exp(); - max + (exp_a + exp_b).ln() - } - } - }, - ) -} - -fn logaddexp_f64( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from rhs tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - crate::operations::arithmetic::broadcast_binary_op( - lhs_data, - rhs_data, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - |a, b| { - if a.is_nan() || b.is_nan() { - f64::NAN - } else { - let max = a.max(b); - if max.is_infinite() { - max - } else { - let exp_a = (a - max).exp(); - let exp_b = (b - max).exp(); - max + (exp_a + exp_b).ln() - } - } - }, - ) -} - -fn tanh_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - unary_apply(input_data, output_slice, f32::tanh); - Ok(()) -} - -fn tanh_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - unary_apply(input_data, output_slice, f64::tanh); - Ok(()) -} - -fn sigmoid_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - unary_apply(input_data, output_slice, stable_sigmoid_f32); - Ok(()) -} - -fn sigmoid_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - unary_apply(input_data, output_slice, stable_sigmoid_f64); - Ok(()) -} - -#[inline] -fn stable_sigmoid_f32(val: f32) -> f32 { - if val >= 0.0 { - let exp_neg = (-val).exp(); - 1.0 / (1.0 + exp_neg) - } else { - let exp_pos = val.exp(); - exp_pos / (1.0 + exp_pos) - } -} - -#[inline] -fn stable_sigmoid_f64(val: f64) -> f64 { - if val >= 0.0 { - let exp_neg = (-val).exp(); - 1.0 / (1.0 + exp_neg) - } else { - let exp_pos = val.exp(); - exp_pos / (1.0 + exp_pos) - } -} - -fn relu_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - let len = input_data.len(); - let mut mask = vec![false; len]; - if len >= PAR_THRESHOLD { - output_slice - .par_iter_mut() - .zip(input_data.par_iter()) - .zip(mask.par_iter_mut()) - .for_each(|((o, &v), m)| { - if v.is_nan() { - *o = v; - } else if v > 0.0 { - *o = v; - *m = true; - } else { - *o = 0.0; - } - }); - } else { - for ((o, &v), m) in output_slice - .iter_mut() - .zip(input_data.iter()) - .zip(mask.iter_mut()) - { - if v.is_nan() { - *o = v; - } else if v > 0.0 { - *o = v; - *m = true; - } else { - *o = 0.0; - } - } - } - Ok(mask) -} - -fn relu_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - let len = input_data.len(); - let mut mask = vec![false; len]; - if len >= PAR_THRESHOLD { - output_slice - .par_iter_mut() - .zip(input_data.par_iter()) - .zip(mask.par_iter_mut()) - .for_each(|((o, &v), m)| { - if v.is_nan() { - *o = v; - } else if v > 0.0 { - *o = v; - *m = true; - } else { - *o = 0.0; - } - }); - } else { - for ((o, &v), m) in output_slice - .iter_mut() - .zip(input_data.iter()) - .zip(mask.iter_mut()) - { - if v.is_nan() { - *o = v; - } else if v > 0.0 { - *o = v; - *m = true; - } else { - *o = 0.0; - } - } - } - Ok(mask) -} - -fn relu_i32(tensor: &Tensor, output_data: &mut TensorData) -> Result> { - let input_data = tensor.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from input tensor") - })?; - - let output_slice = output_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i32 slice from output data") - })?; - let len = input_data.len(); - let mut mask = vec![false; len]; - if len >= PAR_THRESHOLD { - output_slice - .par_iter_mut() - .zip(input_data.par_iter()) - .zip(mask.par_iter_mut()) - .for_each(|((o, &v), m)| { - if v > 0 { - *o = v; - *m = true; - } else { - *o = 0; - } - }); - } else { - for ((o, &v), m) in output_slice - .iter_mut() - .zip(input_data.iter()) - .zip(mask.iter_mut()) - { - if v > 0 { - *o = v; - *m = true; - } else { - *o = 0; - } - } - } - Ok(mask) -} - -fn relu_i64(tensor: &Tensor, output_data: &mut TensorData) -> Result> { - let input_data = tensor.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from input tensor") - })?; - - let output_slice = output_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice from output data") - })?; - let len = input_data.len(); - let mut mask = vec![false; len]; - if len >= PAR_THRESHOLD { - output_slice - .par_iter_mut() - .zip(input_data.par_iter()) - .zip(mask.par_iter_mut()) - .for_each(|((o, &v), m)| { - if v > 0 { - *o = v; - *m = true; - } else { - *o = 0; - } - }); - } else { - for ((o, &v), m) in output_slice - .iter_mut() - .zip(input_data.iter()) - .zip(mask.iter_mut()) - { - if v > 0 { - *o = v; - *m = true; - } else { - *o = 0; - } - } - } - Ok(mask) -} - -fn hardshrink_f32( - tensor: &Tensor, - output_data: &mut TensorData, - lambd: f32, - store_mask: bool, -) -> Result>> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - let mut mask = if store_mask { - Some(Vec::with_capacity(input_data.len())) - } else { - None - }; - - for (&value, out_slot) in input_data.iter().zip(output_slice.iter_mut()) { - let keep = value > lambd || value < -lambd; - *out_slot = if keep { value } else { 0.0 }; - if let Some(ref mut mask_vec) = mask { - mask_vec.push(keep); - } - } - - Ok(mask) -} - -fn hardshrink_f64( - tensor: &Tensor, - output_data: &mut TensorData, - lambd: f64, - store_mask: bool, -) -> Result>> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - let mut mask = if store_mask { - Some(Vec::with_capacity(input_data.len())) - } else { - None - }; - - for (&value, out_slot) in input_data.iter().zip(output_slice.iter_mut()) { - let keep = value > lambd || value < -lambd; - *out_slot = if keep { value } else { 0.0 }; - if let Some(ref mut mask_vec) = mask { - mask_vec.push(keep); - } - } - - Ok(mask) -} - -fn leaky_relu_f32( - tensor: &Tensor, - output_data: &mut TensorData, - negative_slope: f32, -) -> Result> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - let len = input_data.len(); - let mut mask = vec![false; len]; - let mask_ptr = mask.as_mut_ptr() as usize; - let in_ptr = input_data.as_ptr() as usize; - let out_ptr = output_slice.as_mut_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let in_ptr = in_ptr as *const f32; - let out_ptr = out_ptr as *mut f32; - let mask_ptr = mask_ptr as *mut bool; - let val = *in_ptr.add(i); - if val >= 0.0 { - *out_ptr.add(i) = val; - *mask_ptr.add(i) = true; - } else { - *out_ptr.add(i) = negative_slope * val; - } - }); - Ok(mask) -} - -fn leaky_relu_f64( - tensor: &Tensor, - output_data: &mut TensorData, - negative_slope: f64, -) -> Result> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - let len = input_data.len(); - let mut mask = vec![false; len]; - let mask_ptr = mask.as_mut_ptr() as usize; - let in_ptr = input_data.as_ptr() as usize; - let out_ptr = output_slice.as_mut_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let in_ptr = in_ptr as *const f64; - let out_ptr = out_ptr as *mut f64; - let mask_ptr = mask_ptr as *mut bool; - let val = *in_ptr.add(i); - if val >= 0.0 { - *out_ptr.add(i) = val; - *mask_ptr.add(i) = true; - } else { - *out_ptr.add(i) = negative_slope * val; - } - }); - Ok(mask) -} - -fn softmax_f32(tensor: &Tensor, output_data: &mut TensorData, dim: usize) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - let dims = tensor.shape().dims(); - let dim_size = dims[dim]; - - if dim_size == 0 { - return Ok(()); - } - - // Compute the number of groups before and after the softmax dimension. This - // allows us to iterate over all slices along `dim` for tensors of arbitrary - // rank using row-major indexing. - let after: usize = if dim + 1 >= dims.len() { - 1 - } else { - dims[dim + 1..].iter().product() - }; - let group = dim_size * after; - input_data - .par_chunks(group) - .zip(output_slice.par_chunks_mut(group)) - .for_each(|(in_block, out_block)| { - for a in 0..after { - let base = a; - let mut max_val = f32::NEG_INFINITY; - for k in 0..dim_size { - let idx = base + k * after; - max_val = max_val.max(in_block[idx]); - } - if max_val.is_infinite() && max_val.is_sign_negative() { - for k in 0..dim_size { - let idx = base + k * after; - out_block[idx] = 0.0; - } - continue; - } - let mut sum = 0.0f32; - for k in 0..dim_size { - let idx = base + k * after; - let val = (in_block[idx] - max_val).exp(); - out_block[idx] = val; - sum += val; - } - for k in 0..dim_size { - let idx = base + k * after; - out_block[idx] /= sum; - } - } - }); - - Ok(()) -} - -fn softmax_f64(tensor: &Tensor, output_data: &mut TensorData, dim: usize) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - let dims = tensor.shape().dims(); - let dim_size = dims[dim]; - - if dim_size == 0 { - return Ok(()); - } - - let after: usize = if dim + 1 >= dims.len() { - 1 - } else { - dims[dim + 1..].iter().product() - }; - let group = dim_size * after; - input_data - .par_chunks(group) - .zip(output_slice.par_chunks_mut(group)) - .for_each(|(in_block, out_block)| { - for a in 0..after { - let base = a; - let mut max_val = f64::NEG_INFINITY; - for k in 0..dim_size { - let idx = base + k * after; - max_val = max_val.max(in_block[idx]); - } - if max_val.is_infinite() && max_val.is_sign_negative() { - for k in 0..dim_size { - let idx = base + k * after; - out_block[idx] = 0.0; - } - continue; - } - let mut sum = 0.0f64; - for k in 0..dim_size { - let idx = base + k * after; - let val = (in_block[idx] - max_val).exp(); - out_block[idx] = val; - sum += val; - } - for k in 0..dim_size { - let idx = base + k * after; - out_block[idx] /= sum; - } - } - }); - - Ok(()) -} - -fn broadcast_mask_index( - linear_idx: usize, - output_dims: &[usize], - output_strides: &[usize], - mask_dims: &[usize], - mask_strides: &[usize], -) -> usize { - if mask_dims.is_empty() { - return 0; - } - - let output_ndim = output_dims.len(); - let mask_ndim = mask_dims.len(); - let mut mask_index = 0usize; - - for i in 0..mask_ndim { - let output_dim_idx = output_ndim - 1 - i; - let mask_dim_idx = mask_ndim - 1 - i; - let stride = output_strides[output_dim_idx]; - let coord = if stride == 0 { - 0 - } else { - (linear_idx / stride) % output_dims[output_dim_idx] - }; - let mask_dim = mask_dims[mask_dim_idx]; - let mask_coord = if mask_dim == 1 { 0 } else { coord }; - mask_index += mask_coord * mask_strides[mask_dim_idx]; - } - - mask_index -} - -fn masked_softmax_f32( - tensor: &Tensor, - mask: &Tensor, - output_data: &mut TensorData, - dim: usize, -) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - let mask_data = mask.data().as_bool_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from mask tensor") - })?; - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - let dims = tensor.shape().dims(); - let mask_dims = mask.shape().dims(); - let dim_size = dims[dim]; - if dim_size == 0 { - return Ok(()); - } - - let after: usize = if dim + 1 >= dims.len() { - 1 - } else { - dims[dim + 1..].iter().product() - }; - let group = dim_size * after; - let same_shape = mask_dims == dims; - let output_strides = if same_shape { - None - } else { - Some(Strides::from_shape(tensor.shape())) - }; - let mask_strides = if same_shape { - None - } else { - Some(Strides::from_shape(mask.shape())) - }; - - input_data - .par_chunks(group) - .zip(output_slice.par_chunks_mut(group)) - .enumerate() - .for_each(|(block_idx, (in_block, out_block))| { - let block_offset = block_idx * group; - for a in 0..after { - let base = a; - let mut max_val = f32::NEG_INFINITY; - let mut has_unmasked = false; - for k in 0..dim_size { - let idx = base + k * after; - let linear_idx = block_offset + idx; - let masked = if same_shape { - mask_data[linear_idx] - } else { - let mask_index = broadcast_mask_index( - linear_idx, - dims, - output_strides.as_ref().unwrap().as_slice(), - mask_dims, - mask_strides.as_ref().unwrap().as_slice(), - ); - mask_data[mask_index] - }; - if !masked { - has_unmasked = true; - max_val = max_val.max(in_block[idx]); - } - } - if !has_unmasked { - for k in 0..dim_size { - let idx = base + k * after; - out_block[idx] = 0.0; - } - continue; - } - if max_val.is_infinite() && max_val.is_sign_negative() { - for k in 0..dim_size { - let idx = base + k * after; - out_block[idx] = 0.0; - } - continue; - } - let mut sum = 0.0f32; - for k in 0..dim_size { - let idx = base + k * after; - let linear_idx = block_offset + idx; - let masked = if same_shape { - mask_data[linear_idx] - } else { - let mask_index = broadcast_mask_index( - linear_idx, - dims, - output_strides.as_ref().unwrap().as_slice(), - mask_dims, - mask_strides.as_ref().unwrap().as_slice(), - ); - mask_data[mask_index] - }; - if masked { - out_block[idx] = 0.0; - } else { - let val = (in_block[idx] - max_val).exp(); - out_block[idx] = val; - sum += val; - } - } - if sum != 0.0 { - for k in 0..dim_size { - let idx = base + k * after; - out_block[idx] /= sum; - } - } - } - }); - - Ok(()) -} - -fn masked_softmax_f64( - tensor: &Tensor, - mask: &Tensor, - output_data: &mut TensorData, - dim: usize, -) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - let mask_data = mask.data().as_bool_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from mask tensor") - })?; - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - let dims = tensor.shape().dims(); - let mask_dims = mask.shape().dims(); - let dim_size = dims[dim]; - if dim_size == 0 { - return Ok(()); - } - - let after: usize = if dim + 1 >= dims.len() { - 1 - } else { - dims[dim + 1..].iter().product() - }; - let group = dim_size * after; - let same_shape = mask_dims == dims; - let output_strides = if same_shape { - None - } else { - Some(Strides::from_shape(tensor.shape())) - }; - let mask_strides = if same_shape { - None - } else { - Some(Strides::from_shape(mask.shape())) - }; - - input_data - .par_chunks(group) - .zip(output_slice.par_chunks_mut(group)) - .enumerate() - .for_each(|(block_idx, (in_block, out_block))| { - let block_offset = block_idx * group; - for a in 0..after { - let base = a; - let mut max_val = f64::NEG_INFINITY; - let mut has_unmasked = false; - for k in 0..dim_size { - let idx = base + k * after; - let linear_idx = block_offset + idx; - let masked = if same_shape { - mask_data[linear_idx] - } else { - let mask_index = broadcast_mask_index( - linear_idx, - dims, - output_strides.as_ref().unwrap().as_slice(), - mask_dims, - mask_strides.as_ref().unwrap().as_slice(), - ); - mask_data[mask_index] - }; - if !masked { - has_unmasked = true; - max_val = max_val.max(in_block[idx]); - } - } - if !has_unmasked { - for k in 0..dim_size { - let idx = base + k * after; - out_block[idx] = 0.0; - } - continue; - } - if max_val.is_infinite() && max_val.is_sign_negative() { - for k in 0..dim_size { - let idx = base + k * after; - out_block[idx] = 0.0; - } - continue; - } - let mut sum = 0.0f64; - for k in 0..dim_size { - let idx = base + k * after; - let linear_idx = block_offset + idx; - let masked = if same_shape { - mask_data[linear_idx] - } else { - let mask_index = broadcast_mask_index( - linear_idx, - dims, - output_strides.as_ref().unwrap().as_slice(), - mask_dims, - mask_strides.as_ref().unwrap().as_slice(), - ); - mask_data[mask_index] - }; - if masked { - out_block[idx] = 0.0; - } else { - let val = (in_block[idx] - max_val).exp(); - out_block[idx] = val; - sum += val; - } - } - if sum != 0.0 { - for k in 0..dim_size { - let idx = base + k * after; - out_block[idx] /= sum; - } - } - } - }); - - Ok(()) -} - -fn log_softmax_core( - input_data: &[T], - output_slice: &mut [T], - dims: &[usize], - dim: usize, - neg_inf: T, -) -> Result<()> { - let dim_size = dims[dim]; - if dim_size == 0 { - return Ok(()); - } - - let after: usize = if dim + 1 >= dims.len() { - 1 - } else { - dims[dim + 1..].iter().product() - }; - let group = dim_size * after; - input_data - .par_chunks(group) - .zip(output_slice.par_chunks_mut(group)) - .for_each(|(in_block, out_block)| { - for a in 0..after { - let base = a; - let mut max_val = neg_inf; - for k in 0..dim_size { - let idx = base + k * after; - let val = in_block[idx]; - if val > max_val { - max_val = val; - } - } - if max_val.is_infinite() && max_val.is_sign_negative() { - for k in 0..dim_size { - let idx = base + k * after; - out_block[idx] = neg_inf; - } - continue; - } - let mut sum = T::zero(); - for k in 0..dim_size { - let idx = base + k * after; - sum = sum + (in_block[idx] - max_val).exp(); - } - let logsum = sum.ln() + max_val; - for k in 0..dim_size { - let idx = base + k * after; - out_block[idx] = in_block[idx] - logsum; - } - } - }); - - Ok(()) -} - -fn masked_log_softmax_core( - input_data: &[T], - output_slice: &mut [T], - mask_data: &[bool], - tensor_shape: &Shape, - mask_shape: &Shape, - dim: usize, - neg_inf: T, -) -> Result<()> { - let dims = tensor_shape.dims(); - let mask_dims = mask_shape.dims(); - let dim_size = dims[dim]; - if dim_size == 0 { - return Ok(()); - } - - let after: usize = if dim + 1 >= dims.len() { - 1 - } else { - dims[dim + 1..].iter().product() - }; - let group = dim_size * after; - let same_shape = mask_dims == dims; - - if same_shape { - input_data - .par_chunks(group) - .zip(output_slice.par_chunks_mut(group)) - .enumerate() - .for_each(|(block_idx, (in_block, out_block))| { - let block_offset = block_idx * group; - for a in 0..after { - let base = a; - let mut max_val = neg_inf; - let mut has_unmasked = false; - for k in 0..dim_size { - let idx = base + k * after; - let linear_idx = block_offset + idx; - if !mask_data[linear_idx] { - has_unmasked = true; - let val = in_block[idx]; - if val > max_val { - max_val = val; - } - } - } - if !has_unmasked || (max_val.is_infinite() && max_val.is_sign_negative()) { - for k in 0..dim_size { - let idx = base + k * after; - out_block[idx] = neg_inf; - } - continue; - } - let mut sum = T::zero(); - for k in 0..dim_size { - let idx = base + k * after; - let linear_idx = block_offset + idx; - if !mask_data[linear_idx] { - sum = sum + (in_block[idx] - max_val).exp(); - } - } - let logsum = sum.ln() + max_val; - for k in 0..dim_size { - let idx = base + k * after; - let linear_idx = block_offset + idx; - if mask_data[linear_idx] { - out_block[idx] = neg_inf; - } else { - out_block[idx] = in_block[idx] - logsum; - } - } - } - }); - return Ok(()); - } - - let output_strides = Strides::from_shape(tensor_shape); - let mask_strides = Strides::from_shape(mask_shape); - let output_stride_slice = output_strides.as_slice(); - let mask_stride_slice = mask_strides.as_slice(); - - input_data - .par_chunks(group) - .zip(output_slice.par_chunks_mut(group)) - .enumerate() - .for_each(|(block_idx, (in_block, out_block))| { - let block_offset = block_idx * group; - for a in 0..after { - let base = a; - let mut max_val = neg_inf; - let mut has_unmasked = false; - for k in 0..dim_size { - let idx = base + k * after; - let linear_idx = block_offset + idx; - let mask_index = broadcast_mask_index( - linear_idx, - dims, - output_stride_slice, - mask_dims, - mask_stride_slice, - ); - if !mask_data[mask_index] { - has_unmasked = true; - let val = in_block[idx]; - if val > max_val { - max_val = val; - } - } - } - if !has_unmasked || (max_val.is_infinite() && max_val.is_sign_negative()) { - for k in 0..dim_size { - let idx = base + k * after; - out_block[idx] = neg_inf; - } - continue; - } - let mut sum = T::zero(); - for k in 0..dim_size { - let idx = base + k * after; - let linear_idx = block_offset + idx; - let mask_index = broadcast_mask_index( - linear_idx, - dims, - output_stride_slice, - mask_dims, - mask_stride_slice, - ); - if !mask_data[mask_index] { - sum = sum + (in_block[idx] - max_val).exp(); - } - } - let logsum = sum.ln() + max_val; - for k in 0..dim_size { - let idx = base + k * after; - let linear_idx = block_offset + idx; - let mask_index = broadcast_mask_index( - linear_idx, - dims, - output_stride_slice, - mask_dims, - mask_stride_slice, - ); - if mask_data[mask_index] { - out_block[idx] = neg_inf; - } else { - out_block[idx] = in_block[idx] - logsum; - } - } - } - }); - - Ok(()) -} - -macro_rules! masked_log_softmax_impl { - ( - $tensor:expr, - $mask:expr, - $output_data:expr, - $dim:expr, - $input_ty:ty, - $as_input:ident, - $as_output:ident, - $neg_inf:expr - ) => {{ - let input_data = $tensor.data().$as_input().ok_or_else(|| { - MinitensorError::internal_error("Failed to get input slice from tensor") - })?; - let mask_data = $mask.data().as_bool_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from mask tensor") - })?; - let output_slice = $output_data.$as_output().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable output slice from data") - })?; - - masked_log_softmax_core( - input_data, - output_slice, - mask_data, - $tensor.shape(), - $mask.shape(), - $dim, - $neg_inf, - ) - }}; -} - -macro_rules! log_softmax_impl { - ( - $tensor:expr, - $output_data:expr, - $dim:expr, - $input_ty:ty, - $as_input:ident, - $as_output:ident, - $neg_inf:expr - ) => {{ - let input_data = $tensor.data().$as_input().ok_or_else(|| { - MinitensorError::internal_error("Failed to get input slice from tensor") - })?; - let output_slice = $output_data.$as_output().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable output slice from data") - })?; - - let dims = $tensor.shape().dims(); - log_softmax_core(input_data, output_slice, dims, $dim, $neg_inf) - }}; -} - -fn masked_log_softmax_f32( - tensor: &Tensor, - mask: &Tensor, - output_data: &mut TensorData, - dim: usize, -) -> Result<()> { - masked_log_softmax_impl!( - tensor, - mask, - output_data, - dim, - f32, - as_f32_slice, - as_f32_slice_mut, - f32::NEG_INFINITY - ) -} - -fn masked_log_softmax_f64( - tensor: &Tensor, - mask: &Tensor, - output_data: &mut TensorData, - dim: usize, -) -> Result<()> { - masked_log_softmax_impl!( - tensor, - mask, - output_data, - dim, - f64, - as_f64_slice, - as_f64_slice_mut, - f64::NEG_INFINITY - ) -} - -fn log_softmax_f32(tensor: &Tensor, output_data: &mut TensorData, dim: usize) -> Result<()> { - log_softmax_impl!( - tensor, - output_data, - dim, - f32, - as_f32_slice, - as_f32_slice_mut, - f32::NEG_INFINITY - ) -} - -fn log_softmax_f64(tensor: &Tensor, output_data: &mut TensorData, dim: usize) -> Result<()> { - log_softmax_impl!( - tensor, - output_data, - dim, - f64, - as_f64_slice, - as_f64_slice_mut, - f64::NEG_INFINITY - ) -} +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::error::MinitensorError; +use crate::error::Result; +use crate::tensor::Shape; +use crate::tensor::Strides; +use crate::tensor::Tensor; +use crate::tensor::TensorData; +use rayon::prelude::*; + +use num_traits::Float; + +pub(crate) fn logaddexp_f32( + lhs: &Tensor, + rhs: &Tensor, + output_data: &mut TensorData, + output_shape: &Shape, +) -> Result<()> { + let lhs_data = lhs.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from lhs tensor") + })?; + let rhs_data = rhs.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from rhs tensor") + })?; + + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + + crate::operations::arithmetic::broadcast_binary_op( + lhs_data, + rhs_data, + output_slice, + lhs.shape(), + rhs.shape(), + output_shape, + |a, b| { + if a.is_nan() || b.is_nan() { + f32::NAN + } else { + let max = a.max(b); + if max.is_infinite() { + max + } else { + let exp_a = (a - max).exp(); + let exp_b = (b - max).exp(); + max + (exp_a + exp_b).ln() + } + } + }, + ) +} + +pub(crate) fn logaddexp_f64( + lhs: &Tensor, + rhs: &Tensor, + output_data: &mut TensorData, + output_shape: &Shape, +) -> Result<()> { + let lhs_data = lhs.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from lhs tensor") + })?; + let rhs_data = rhs.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from rhs tensor") + })?; + + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + + crate::operations::arithmetic::broadcast_binary_op( + lhs_data, + rhs_data, + output_slice, + lhs.shape(), + rhs.shape(), + output_shape, + |a, b| { + if a.is_nan() || b.is_nan() { + f64::NAN + } else { + let max = a.max(b); + if max.is_infinite() { + max + } else { + let exp_a = (a - max).exp(); + let exp_b = (b - max).exp(); + max + (exp_a + exp_b).ln() + } + } + }, + ) +} + +pub(crate) fn tanh_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + unary_apply(input_data, output_slice, f32::tanh); + Ok(()) +} + +pub(crate) fn tanh_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + unary_apply(input_data, output_slice, f64::tanh); + Ok(()) +} + +pub(crate) fn sigmoid_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + unary_apply(input_data, output_slice, stable_sigmoid_f32); + Ok(()) +} + +pub(crate) fn sigmoid_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result<()> { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + unary_apply(input_data, output_slice, stable_sigmoid_f64); + Ok(()) +} + +#[inline] +fn stable_sigmoid_f32(val: f32) -> f32 { + if val >= 0.0 { + let exp_neg = (-val).exp(); + 1.0 / (1.0 + exp_neg) + } else { + let exp_pos = val.exp(); + exp_pos / (1.0 + exp_pos) + } +} + +#[inline] +fn stable_sigmoid_f64(val: f64) -> f64 { + if val >= 0.0 { + let exp_neg = (-val).exp(); + 1.0 / (1.0 + exp_neg) + } else { + let exp_pos = val.exp(); + exp_pos / (1.0 + exp_pos) + } +} + +pub(crate) fn relu_f32(tensor: &Tensor, output_data: &mut TensorData) -> Result> { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + let len = input_data.len(); + let mut mask = vec![false; len]; + if len >= PAR_THRESHOLD { + output_slice + .par_iter_mut() + .zip(input_data.par_iter()) + .zip(mask.par_iter_mut()) + .for_each(|((o, &v), m)| { + if v.is_nan() { + *o = v; + } else if v > 0.0 { + *o = v; + *m = true; + } else { + *o = 0.0; + } + }); + } else { + for ((o, &v), m) in output_slice + .iter_mut() + .zip(input_data.iter()) + .zip(mask.iter_mut()) + { + if v.is_nan() { + *o = v; + } else if v > 0.0 { + *o = v; + *m = true; + } else { + *o = 0.0; + } + } + } + Ok(mask) +} + +pub(crate) fn relu_f64(tensor: &Tensor, output_data: &mut TensorData) -> Result> { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + let len = input_data.len(); + let mut mask = vec![false; len]; + if len >= PAR_THRESHOLD { + output_slice + .par_iter_mut() + .zip(input_data.par_iter()) + .zip(mask.par_iter_mut()) + .for_each(|((o, &v), m)| { + if v.is_nan() { + *o = v; + } else if v > 0.0 { + *o = v; + *m = true; + } else { + *o = 0.0; + } + }); + } else { + for ((o, &v), m) in output_slice + .iter_mut() + .zip(input_data.iter()) + .zip(mask.iter_mut()) + { + if v.is_nan() { + *o = v; + } else if v > 0.0 { + *o = v; + *m = true; + } else { + *o = 0.0; + } + } + } + Ok(mask) +} + +pub(crate) fn relu_i32(tensor: &Tensor, output_data: &mut TensorData) -> Result> { + let input_data = tensor.data().as_i32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i32 slice from input tensor") + })?; + + let output_slice = output_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i32 slice from output data") + })?; + let len = input_data.len(); + let mut mask = vec![false; len]; + if len >= PAR_THRESHOLD { + output_slice + .par_iter_mut() + .zip(input_data.par_iter()) + .zip(mask.par_iter_mut()) + .for_each(|((o, &v), m)| { + if v > 0 { + *o = v; + *m = true; + } else { + *o = 0; + } + }); + } else { + for ((o, &v), m) in output_slice + .iter_mut() + .zip(input_data.iter()) + .zip(mask.iter_mut()) + { + if v > 0 { + *o = v; + *m = true; + } else { + *o = 0; + } + } + } + Ok(mask) +} + +pub(crate) fn relu_i64(tensor: &Tensor, output_data: &mut TensorData) -> Result> { + let input_data = tensor.data().as_i64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i64 slice from input tensor") + })?; + + let output_slice = output_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice from output data") + })?; + let len = input_data.len(); + let mut mask = vec![false; len]; + if len >= PAR_THRESHOLD { + output_slice + .par_iter_mut() + .zip(input_data.par_iter()) + .zip(mask.par_iter_mut()) + .for_each(|((o, &v), m)| { + if v > 0 { + *o = v; + *m = true; + } else { + *o = 0; + } + }); + } else { + for ((o, &v), m) in output_slice + .iter_mut() + .zip(input_data.iter()) + .zip(mask.iter_mut()) + { + if v > 0 { + *o = v; + *m = true; + } else { + *o = 0; + } + } + } + Ok(mask) +} + +pub(crate) fn hardshrink_f32( + tensor: &Tensor, + output_data: &mut TensorData, + lambd: f32, + store_mask: bool, +) -> Result>> { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + + let mut mask = if store_mask { + Some(Vec::with_capacity(input_data.len())) + } else { + None + }; + + for (&value, out_slot) in input_data.iter().zip(output_slice.iter_mut()) { + let keep = value > lambd || value < -lambd; + *out_slot = if keep { value } else { 0.0 }; + if let Some(ref mut mask_vec) = mask { + mask_vec.push(keep); + } + } + + Ok(mask) +} + +pub(crate) fn hardshrink_f64( + tensor: &Tensor, + output_data: &mut TensorData, + lambd: f64, + store_mask: bool, +) -> Result>> { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + + let mut mask = if store_mask { + Some(Vec::with_capacity(input_data.len())) + } else { + None + }; + + for (&value, out_slot) in input_data.iter().zip(output_slice.iter_mut()) { + let keep = value > lambd || value < -lambd; + *out_slot = if keep { value } else { 0.0 }; + if let Some(ref mut mask_vec) = mask { + mask_vec.push(keep); + } + } + + Ok(mask) +} + +pub(crate) fn leaky_relu_f32( + tensor: &Tensor, + output_data: &mut TensorData, + negative_slope: f32, +) -> Result> { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + + let len = input_data.len(); + let mut mask = vec![false; len]; + let mask_ptr = mask.as_mut_ptr() as usize; + let in_ptr = input_data.as_ptr() as usize; + let out_ptr = output_slice.as_mut_ptr() as usize; + (0..len).into_par_iter().for_each(|i| unsafe { + let in_ptr = in_ptr as *const f32; + let out_ptr = out_ptr as *mut f32; + let mask_ptr = mask_ptr as *mut bool; + let val = *in_ptr.add(i); + if val >= 0.0 { + *out_ptr.add(i) = val; + *mask_ptr.add(i) = true; + } else { + *out_ptr.add(i) = negative_slope * val; + } + }); + Ok(mask) +} + +pub(crate) fn leaky_relu_f64( + tensor: &Tensor, + output_data: &mut TensorData, + negative_slope: f64, +) -> Result> { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + + let len = input_data.len(); + let mut mask = vec![false; len]; + let mask_ptr = mask.as_mut_ptr() as usize; + let in_ptr = input_data.as_ptr() as usize; + let out_ptr = output_slice.as_mut_ptr() as usize; + (0..len).into_par_iter().for_each(|i| unsafe { + let in_ptr = in_ptr as *const f64; + let out_ptr = out_ptr as *mut f64; + let mask_ptr = mask_ptr as *mut bool; + let val = *in_ptr.add(i); + if val >= 0.0 { + *out_ptr.add(i) = val; + *mask_ptr.add(i) = true; + } else { + *out_ptr.add(i) = negative_slope * val; + } + }); + Ok(mask) +} + +pub(crate) fn softmax_f32(tensor: &Tensor, output_data: &mut TensorData, dim: usize) -> Result<()> { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + + let dims = tensor.shape().dims(); + let dim_size = dims[dim]; + + if dim_size == 0 { + return Ok(()); + } + + // Compute the number of groups before and after the softmax dimension. This + // allows us to iterate over all slices along `dim` for tensors of arbitrary + // rank using row-major indexing. + let after: usize = if dim + 1 >= dims.len() { + 1 + } else { + dims[dim + 1..].iter().product() + }; + let group = dim_size * after; + input_data + .par_chunks(group) + .zip(output_slice.par_chunks_mut(group)) + .for_each(|(in_block, out_block)| { + for a in 0..after { + let base = a; + let mut max_val = f32::NEG_INFINITY; + for k in 0..dim_size { + let idx = base + k * after; + max_val = max_val.max(in_block[idx]); + } + if max_val.is_infinite() && max_val.is_sign_negative() { + for k in 0..dim_size { + let idx = base + k * after; + out_block[idx] = 0.0; + } + continue; + } + let mut sum = 0.0f32; + for k in 0..dim_size { + let idx = base + k * after; + let val = (in_block[idx] - max_val).exp(); + out_block[idx] = val; + sum += val; + } + for k in 0..dim_size { + let idx = base + k * after; + out_block[idx] /= sum; + } + } + }); + + Ok(()) +} + +pub(crate) fn softmax_f64(tensor: &Tensor, output_data: &mut TensorData, dim: usize) -> Result<()> { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + + let dims = tensor.shape().dims(); + let dim_size = dims[dim]; + + if dim_size == 0 { + return Ok(()); + } + + let after: usize = if dim + 1 >= dims.len() { + 1 + } else { + dims[dim + 1..].iter().product() + }; + let group = dim_size * after; + input_data + .par_chunks(group) + .zip(output_slice.par_chunks_mut(group)) + .for_each(|(in_block, out_block)| { + for a in 0..after { + let base = a; + let mut max_val = f64::NEG_INFINITY; + for k in 0..dim_size { + let idx = base + k * after; + max_val = max_val.max(in_block[idx]); + } + if max_val.is_infinite() && max_val.is_sign_negative() { + for k in 0..dim_size { + let idx = base + k * after; + out_block[idx] = 0.0; + } + continue; + } + let mut sum = 0.0f64; + for k in 0..dim_size { + let idx = base + k * after; + let val = (in_block[idx] - max_val).exp(); + out_block[idx] = val; + sum += val; + } + for k in 0..dim_size { + let idx = base + k * after; + out_block[idx] /= sum; + } + } + }); + + Ok(()) +} + +fn broadcast_mask_index( + linear_idx: usize, + output_dims: &[usize], + output_strides: &[usize], + mask_dims: &[usize], + mask_strides: &[usize], +) -> usize { + if mask_dims.is_empty() { + return 0; + } + + let output_ndim = output_dims.len(); + let mask_ndim = mask_dims.len(); + let mut mask_index = 0usize; + + for i in 0..mask_ndim { + let output_dim_idx = output_ndim - 1 - i; + let mask_dim_idx = mask_ndim - 1 - i; + let stride = output_strides[output_dim_idx]; + let coord = if stride == 0 { + 0 + } else { + (linear_idx / stride) % output_dims[output_dim_idx] + }; + let mask_dim = mask_dims[mask_dim_idx]; + let mask_coord = if mask_dim == 1 { 0 } else { coord }; + mask_index += mask_coord * mask_strides[mask_dim_idx]; + } + + mask_index +} + +pub(crate) fn masked_softmax_f32( + tensor: &Tensor, + mask: &Tensor, + output_data: &mut TensorData, + dim: usize, +) -> Result<()> { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + let mask_data = mask.data().as_bool_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get bool slice from mask tensor") + })?; + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + + let dims = tensor.shape().dims(); + let mask_dims = mask.shape().dims(); + let dim_size = dims[dim]; + if dim_size == 0 { + return Ok(()); + } + + let after: usize = if dim + 1 >= dims.len() { + 1 + } else { + dims[dim + 1..].iter().product() + }; + let group = dim_size * after; + let same_shape = mask_dims == dims; + let output_strides = if same_shape { + None + } else { + Some(Strides::from_shape(tensor.shape())) + }; + let mask_strides = if same_shape { + None + } else { + Some(Strides::from_shape(mask.shape())) + }; + + input_data + .par_chunks(group) + .zip(output_slice.par_chunks_mut(group)) + .enumerate() + .for_each(|(block_idx, (in_block, out_block))| { + let block_offset = block_idx * group; + for a in 0..after { + let base = a; + let mut max_val = f32::NEG_INFINITY; + let mut has_unmasked = false; + for k in 0..dim_size { + let idx = base + k * after; + let linear_idx = block_offset + idx; + let masked = if same_shape { + mask_data[linear_idx] + } else { + let mask_index = broadcast_mask_index( + linear_idx, + dims, + output_strides.as_ref().unwrap().as_slice(), + mask_dims, + mask_strides.as_ref().unwrap().as_slice(), + ); + mask_data[mask_index] + }; + if !masked { + has_unmasked = true; + max_val = max_val.max(in_block[idx]); + } + } + if !has_unmasked { + for k in 0..dim_size { + let idx = base + k * after; + out_block[idx] = 0.0; + } + continue; + } + if max_val.is_infinite() && max_val.is_sign_negative() { + for k in 0..dim_size { + let idx = base + k * after; + out_block[idx] = 0.0; + } + continue; + } + let mut sum = 0.0f32; + for k in 0..dim_size { + let idx = base + k * after; + let linear_idx = block_offset + idx; + let masked = if same_shape { + mask_data[linear_idx] + } else { + let mask_index = broadcast_mask_index( + linear_idx, + dims, + output_strides.as_ref().unwrap().as_slice(), + mask_dims, + mask_strides.as_ref().unwrap().as_slice(), + ); + mask_data[mask_index] + }; + if masked { + out_block[idx] = 0.0; + } else { + let val = (in_block[idx] - max_val).exp(); + out_block[idx] = val; + sum += val; + } + } + if sum != 0.0 { + for k in 0..dim_size { + let idx = base + k * after; + out_block[idx] /= sum; + } + } + } + }); + + Ok(()) +} + +pub(crate) fn masked_softmax_f64( + tensor: &Tensor, + mask: &Tensor, + output_data: &mut TensorData, + dim: usize, +) -> Result<()> { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + let mask_data = mask.data().as_bool_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get bool slice from mask tensor") + })?; + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + + let dims = tensor.shape().dims(); + let mask_dims = mask.shape().dims(); + let dim_size = dims[dim]; + if dim_size == 0 { + return Ok(()); + } + + let after: usize = if dim + 1 >= dims.len() { + 1 + } else { + dims[dim + 1..].iter().product() + }; + let group = dim_size * after; + let same_shape = mask_dims == dims; + let output_strides = if same_shape { + None + } else { + Some(Strides::from_shape(tensor.shape())) + }; + let mask_strides = if same_shape { + None + } else { + Some(Strides::from_shape(mask.shape())) + }; + + input_data + .par_chunks(group) + .zip(output_slice.par_chunks_mut(group)) + .enumerate() + .for_each(|(block_idx, (in_block, out_block))| { + let block_offset = block_idx * group; + for a in 0..after { + let base = a; + let mut max_val = f64::NEG_INFINITY; + let mut has_unmasked = false; + for k in 0..dim_size { + let idx = base + k * after; + let linear_idx = block_offset + idx; + let masked = if same_shape { + mask_data[linear_idx] + } else { + let mask_index = broadcast_mask_index( + linear_idx, + dims, + output_strides.as_ref().unwrap().as_slice(), + mask_dims, + mask_strides.as_ref().unwrap().as_slice(), + ); + mask_data[mask_index] + }; + if !masked { + has_unmasked = true; + max_val = max_val.max(in_block[idx]); + } + } + if !has_unmasked { + for k in 0..dim_size { + let idx = base + k * after; + out_block[idx] = 0.0; + } + continue; + } + if max_val.is_infinite() && max_val.is_sign_negative() { + for k in 0..dim_size { + let idx = base + k * after; + out_block[idx] = 0.0; + } + continue; + } + let mut sum = 0.0f64; + for k in 0..dim_size { + let idx = base + k * after; + let linear_idx = block_offset + idx; + let masked = if same_shape { + mask_data[linear_idx] + } else { + let mask_index = broadcast_mask_index( + linear_idx, + dims, + output_strides.as_ref().unwrap().as_slice(), + mask_dims, + mask_strides.as_ref().unwrap().as_slice(), + ); + mask_data[mask_index] + }; + if masked { + out_block[idx] = 0.0; + } else { + let val = (in_block[idx] - max_val).exp(); + out_block[idx] = val; + sum += val; + } + } + if sum != 0.0 { + for k in 0..dim_size { + let idx = base + k * after; + out_block[idx] /= sum; + } + } + } + }); + + Ok(()) +} + +fn log_softmax_core( + input_data: &[T], + output_slice: &mut [T], + dims: &[usize], + dim: usize, + neg_inf: T, +) -> Result<()> { + let dim_size = dims[dim]; + if dim_size == 0 { + return Ok(()); + } + + let after: usize = if dim + 1 >= dims.len() { + 1 + } else { + dims[dim + 1..].iter().product() + }; + let group = dim_size * after; + input_data + .par_chunks(group) + .zip(output_slice.par_chunks_mut(group)) + .for_each(|(in_block, out_block)| { + for a in 0..after { + let base = a; + let mut max_val = neg_inf; + for k in 0..dim_size { + let idx = base + k * after; + let val = in_block[idx]; + if val > max_val { + max_val = val; + } + } + if max_val.is_infinite() && max_val.is_sign_negative() { + for k in 0..dim_size { + let idx = base + k * after; + out_block[idx] = neg_inf; + } + continue; + } + let mut sum = T::zero(); + for k in 0..dim_size { + let idx = base + k * after; + sum = sum + (in_block[idx] - max_val).exp(); + } + let logsum = sum.ln() + max_val; + for k in 0..dim_size { + let idx = base + k * after; + out_block[idx] = in_block[idx] - logsum; + } + } + }); + + Ok(()) +} + +fn masked_log_softmax_core( + input_data: &[T], + output_slice: &mut [T], + mask_data: &[bool], + tensor_shape: &Shape, + mask_shape: &Shape, + dim: usize, + neg_inf: T, +) -> Result<()> { + let dims = tensor_shape.dims(); + let mask_dims = mask_shape.dims(); + let dim_size = dims[dim]; + if dim_size == 0 { + return Ok(()); + } + + let after: usize = if dim + 1 >= dims.len() { + 1 + } else { + dims[dim + 1..].iter().product() + }; + let group = dim_size * after; + let same_shape = mask_dims == dims; + + if same_shape { + input_data + .par_chunks(group) + .zip(output_slice.par_chunks_mut(group)) + .enumerate() + .for_each(|(block_idx, (in_block, out_block))| { + let block_offset = block_idx * group; + for a in 0..after { + let base = a; + let mut max_val = neg_inf; + let mut has_unmasked = false; + for k in 0..dim_size { + let idx = base + k * after; + let linear_idx = block_offset + idx; + if !mask_data[linear_idx] { + has_unmasked = true; + let val = in_block[idx]; + if val > max_val { + max_val = val; + } + } + } + if !has_unmasked || (max_val.is_infinite() && max_val.is_sign_negative()) { + for k in 0..dim_size { + let idx = base + k * after; + out_block[idx] = neg_inf; + } + continue; + } + let mut sum = T::zero(); + for k in 0..dim_size { + let idx = base + k * after; + let linear_idx = block_offset + idx; + if !mask_data[linear_idx] { + sum = sum + (in_block[idx] - max_val).exp(); + } + } + let logsum = sum.ln() + max_val; + for k in 0..dim_size { + let idx = base + k * after; + let linear_idx = block_offset + idx; + if mask_data[linear_idx] { + out_block[idx] = neg_inf; + } else { + out_block[idx] = in_block[idx] - logsum; + } + } + } + }); + return Ok(()); + } + + let output_strides = Strides::from_shape(tensor_shape); + let mask_strides = Strides::from_shape(mask_shape); + let output_stride_slice = output_strides.as_slice(); + let mask_stride_slice = mask_strides.as_slice(); + + input_data + .par_chunks(group) + .zip(output_slice.par_chunks_mut(group)) + .enumerate() + .for_each(|(block_idx, (in_block, out_block))| { + let block_offset = block_idx * group; + for a in 0..after { + let base = a; + let mut max_val = neg_inf; + let mut has_unmasked = false; + for k in 0..dim_size { + let idx = base + k * after; + let linear_idx = block_offset + idx; + let mask_index = broadcast_mask_index( + linear_idx, + dims, + output_stride_slice, + mask_dims, + mask_stride_slice, + ); + if !mask_data[mask_index] { + has_unmasked = true; + let val = in_block[idx]; + if val > max_val { + max_val = val; + } + } + } + if !has_unmasked || (max_val.is_infinite() && max_val.is_sign_negative()) { + for k in 0..dim_size { + let idx = base + k * after; + out_block[idx] = neg_inf; + } + continue; + } + let mut sum = T::zero(); + for k in 0..dim_size { + let idx = base + k * after; + let linear_idx = block_offset + idx; + let mask_index = broadcast_mask_index( + linear_idx, + dims, + output_stride_slice, + mask_dims, + mask_stride_slice, + ); + if !mask_data[mask_index] { + sum = sum + (in_block[idx] - max_val).exp(); + } + } + let logsum = sum.ln() + max_val; + for k in 0..dim_size { + let idx = base + k * after; + let linear_idx = block_offset + idx; + let mask_index = broadcast_mask_index( + linear_idx, + dims, + output_stride_slice, + mask_dims, + mask_stride_slice, + ); + if mask_data[mask_index] { + out_block[idx] = neg_inf; + } else { + out_block[idx] = in_block[idx] - logsum; + } + } + } + }); + + Ok(()) +} + +macro_rules! masked_log_softmax_impl { + ( + $tensor:expr, + $mask:expr, + $output_data:expr, + $dim:expr, + $input_ty:ty, + $as_input:ident, + $as_output:ident, + $neg_inf:expr + ) => {{ + let input_data = $tensor.data().$as_input().ok_or_else(|| { + MinitensorError::internal_error("Failed to get input slice from tensor") + })?; + let mask_data = $mask.data().as_bool_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get bool slice from mask tensor") + })?; + let output_slice = $output_data.$as_output().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable output slice from data") + })?; + + masked_log_softmax_core( + input_data, + output_slice, + mask_data, + $tensor.shape(), + $mask.shape(), + $dim, + $neg_inf, + ) + }}; +} + +macro_rules! log_softmax_impl { + ( + $tensor:expr, + $output_data:expr, + $dim:expr, + $input_ty:ty, + $as_input:ident, + $as_output:ident, + $neg_inf:expr + ) => {{ + let input_data = $tensor.data().$as_input().ok_or_else(|| { + MinitensorError::internal_error("Failed to get input slice from tensor") + })?; + let output_slice = $output_data.$as_output().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable output slice from data") + })?; + + let dims = $tensor.shape().dims(); + log_softmax_core(input_data, output_slice, dims, $dim, $neg_inf) + }}; +} + +pub(crate) fn masked_log_softmax_f32( + tensor: &Tensor, + mask: &Tensor, + output_data: &mut TensorData, + dim: usize, +) -> Result<()> { + masked_log_softmax_impl!( + tensor, + mask, + output_data, + dim, + f32, + as_f32_slice, + as_f32_slice_mut, + f32::NEG_INFINITY + ) +} + +pub(crate) fn masked_log_softmax_f64( + tensor: &Tensor, + mask: &Tensor, + output_data: &mut TensorData, + dim: usize, +) -> Result<()> { + masked_log_softmax_impl!( + tensor, + mask, + output_data, + dim, + f64, + as_f64_slice, + as_f64_slice_mut, + f64::NEG_INFINITY + ) +} + +pub(crate) fn log_softmax_f32( + tensor: &Tensor, + output_data: &mut TensorData, + dim: usize, +) -> Result<()> { + log_softmax_impl!( + tensor, + output_data, + dim, + f32, + as_f32_slice, + as_f32_slice_mut, + f32::NEG_INFINITY + ) +} + +pub(crate) fn log_softmax_f64( + tensor: &Tensor, + output_data: &mut TensorData, + dim: usize, +) -> Result<()> { + log_softmax_impl!( + tensor, + output_data, + dim, + f64, + as_f64_slice, + as_f64_slice_mut, + f64::NEG_INFINITY + ) +} diff --git a/engine/src/operations/activation/trigonometry.rs b/engine/src/operations/activation/trigonometry.rs index 2cf59fd1..4bc376db 100644 --- a/engine/src/operations/activation/trigonometry.rs +++ b/engine/src/operations/activation/trigonometry.rs @@ -1,841 +1,864 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -/// Element-wise power with tensor exponent and gradient support -pub fn pow(base: &Tensor, exponent: &Tensor) -> Result { - // Check device and dtype compatibility - if base.device() != exponent.device() { - return Err(MinitensorError::device_mismatch( - format!("{:?}", base.device()), - format!("{:?}", exponent.device()), - )); - } - - if base.dtype() != exponent.dtype() { - return Err(MinitensorError::type_mismatch( - format!("{:?}", base.dtype()), - format!("{:?}", exponent.dtype()), - )); - } - - let base_shape = base.shape().clone(); - let exponent_shape = exponent.shape().clone(); - let base_numel = base_shape.numel(); - let exp_numel = exponent_shape.numel(); - - let broadcast = if base_shape == exponent_shape { - PowBroadcast::None - } else if base_numel == 1 { - PowBroadcast::BaseScalar - } else if exp_numel == 1 { - PowBroadcast::ExponentScalar - } else { - // General broadcasting: materialize both operands at the broadcast - // shape (expand and contiguous are grad-aware) and recurse into the - // same-shape fast path. - let out_shape = base_shape.broadcast_with(&exponent_shape)?; - let dims: Vec = out_shape.dims().iter().map(|&d| d as isize).collect(); - let base_b = base.expand(dims.clone())?.contiguous()?; - let exp_b = exponent.expand(dims)?.contiguous()?; - return pow(&base_b, &exp_b); - }; - - let output_shape = match broadcast { - PowBroadcast::None | PowBroadcast::ExponentScalar => base_shape.clone(), - PowBroadcast::BaseScalar => exponent_shape.clone(), - }; - - let mut output_data = - TensorData::uninitialized_on_device(output_shape.numel(), base.dtype(), base.device()); - - match base.dtype() { - DataType::Float32 => { - let b = base.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from base tensor") - })?; - let e = exponent.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from exponent tensor") - })?; - let out = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - match broadcast { - PowBroadcast::None => { - for i in 0..b.len() { - out[i] = b[i].powf(e[i]); - } - } - PowBroadcast::BaseScalar => { - let base_val = b[0]; - for i in 0..e.len() { - out[i] = base_val.powf(e[i]); - } - } - PowBroadcast::ExponentScalar => { - let exp_val = e[0]; - for i in 0..b.len() { - out[i] = b[i].powf(exp_val); - } - } - } - } - DataType::Float64 => { - let b = base.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from base tensor") - })?; - let e = exponent.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from exponent tensor") - })?; - let out = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - match broadcast { - PowBroadcast::None => { - for i in 0..b.len() { - out[i] = b[i].powf(e[i]); - } - } - PowBroadcast::BaseScalar => { - let base_val = b[0]; - for i in 0..e.len() { - out[i] = base_val.powf(e[i]); - } - } - PowBroadcast::ExponentScalar => { - let exp_val = e[0]; - for i in 0..b.len() { - out[i] = b[i].powf(exp_val); - } - } - } - } - _ => { - return Err(MinitensorError::invalid_operation( - "Power operation only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - output_shape, - base.dtype(), - base.device(), - base.requires_grad() || exponent.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(PowBackward { - base: base.detach(), - exponent: exponent.detach(), - output: output.clone().detach(), - input_ids: [base.id(), exponent.id()], - base_requires_grad: base.requires_grad(), - exp_requires_grad: exponent.requires_grad(), - broadcast, - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Element-wise power with scalar exponent and gradient support -pub fn powf(tensor: &Tensor, exponent: f64) -> Result { - // Create exponent tensor filled with scalar value - let mut exp_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - match tensor.dtype() { - DataType::Float32 => { - let slice = exp_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from exponent data", - ) - })?; - for val in slice.iter_mut() { - *val = exponent as f32; - } - } - DataType::Float64 => { - let slice = exp_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from exponent data", - ) - })?; - for val in slice.iter_mut() { - *val = exponent; - } - } - _ => { - return Err(MinitensorError::invalid_operation( - "Power operation only supported for floating point tensors", - )); - } - } - let exp_tensor = Tensor::new( - Arc::new(exp_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - false, - ); - pow(tensor, &exp_tensor) -} - -/// Numerically stable logaddexp with gradient support -pub fn logaddexp(lhs: &Tensor, rhs: &Tensor) -> Result { - if lhs.device() != rhs.device() { - return Err(MinitensorError::device_mismatch( - format!("{:?}", lhs.device()), - format!("{:?}", rhs.device()), - )); - } - - let requires_grad = lhs.requires_grad() || rhs.requires_grad(); - use crate::operations::binary::{BinaryOpKind, coerce_binary_operands}; - let (lhs_cast, rhs_cast, result_dtype) = coerce_binary_operands(lhs, rhs, BinaryOpKind::Add)?; - - let lhs_tensor = match lhs_cast { - std::borrow::Cow::Borrowed(t) => t.clone(), - std::borrow::Cow::Owned(t) => t, - }; - let rhs_tensor = match rhs_cast { - std::borrow::Cow::Borrowed(t) => t.clone(), - std::borrow::Cow::Owned(t) => t, - }; - - if result_dtype != DataType::Float32 && result_dtype != DataType::Float64 { - return Err(MinitensorError::invalid_operation( - "logaddexp is only supported for floating point tensors", - )); - } - - let output_shape = lhs_tensor.shape().broadcast_with(rhs_tensor.shape())?; - let mut output_data = - TensorData::uninitialized_on_device(output_shape.numel(), result_dtype, lhs.device()); - - match result_dtype { - DataType::Float32 => { - logaddexp_f32(&lhs_tensor, &rhs_tensor, &mut output_data, &output_shape)? - } - DataType::Float64 => { - logaddexp_f64(&lhs_tensor, &rhs_tensor, &mut output_data, &output_shape)? - } - _ => unreachable!(), - } - - let output = Tensor::new( - Arc::new(output_data), - output_shape.clone(), - result_dtype, - lhs.device(), - requires_grad, - ); - - if requires_grad { - let grad_fn = Arc::new(LogAddExpBackward { - lhs: lhs_tensor.detach(), - rhs: rhs_tensor.detach(), - output: output.clone().detach(), - input_ids: [lhs.id(), rhs.id()], - input_shapes: [lhs.shape().dims().to_vec(), rhs.shape().dims().to_vec()], - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Softplus activation function with gradient support -pub fn softplus(tensor: &Tensor, beta: f64, threshold: f64) -> Result { - if beta <= 0.0 { - return Err(MinitensorError::invalid_argument( - "softplus beta must be positive", - )); - } - - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => softplus_f32(tensor, &mut output_data, beta as f32, threshold as f32)?, - DataType::Float64 => softplus_f64(tensor, &mut output_data, beta, threshold)?, - _ => { - return Err(MinitensorError::invalid_operation( - "Softplus is only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(SoftplusBackward { - input_id: tensor.id(), - input: tensor.clone().detach(), - beta, - threshold, - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// GELU activation function with optional tanh approximation -pub fn gelu(tensor: &Tensor, approximate: bool) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => gelu_f32(tensor, &mut output_data, approximate)?, - DataType::Float64 => gelu_f64(tensor, &mut output_data, approximate)?, - _ => { - return Err(MinitensorError::invalid_operation( - "GELU is only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(GeluBackward { - input_id: tensor.id(), - input: tensor.clone().detach(), - approximate, - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// ELU activation function with configurable alpha -pub fn elu(tensor: &Tensor, alpha: f64) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => elu_f32(tensor, &mut output_data, alpha as f32)?, - DataType::Float64 => elu_f64(tensor, &mut output_data, alpha)?, - _ => { - return Err(MinitensorError::invalid_operation( - "ELU is only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(EluBackward { - input_id: tensor.id(), - output: output.clone().detach(), - alpha, - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// SELU activation function. -pub fn selu(tensor: &Tensor) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => selu_f32(tensor, &mut output_data)?, - DataType::Float64 => selu_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "SELU is only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(SeluBackward { - input_id: tensor.id(), - output: output.clone().detach(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// SiLU (Swish) activation function with gradient support -pub fn silu(tensor: &Tensor) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => silu_f32(tensor, &mut output_data)?, - DataType::Float64 => silu_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "SiLU is only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(SiluBackward { - input_id: tensor.id(), - input: tensor.clone().detach(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Softsign activation function with gradient support -pub fn softsign(tensor: &Tensor) -> Result { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => softsign_f32(tensor, &mut output_data)?, - DataType::Float64 => softsign_f64(tensor, &mut output_data)?, - _ => { - return Err(MinitensorError::invalid_operation( - "Softsign is only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(SoftsignBackward { - input_id: tensor.id(), - input: tensor.clone().detach(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// ReLU activation function with gradient support -pub fn relu(tensor: &Tensor) -> Result { - // Create output tensor data - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - // Perform ReLU based on data type while capturing mask of positive inputs - let mask = match tensor.dtype() { - DataType::Float32 => relu_f32(tensor, &mut output_data)?, - DataType::Float64 => relu_f64(tensor, &mut output_data)?, - DataType::Int32 => relu_i32(tensor, &mut output_data)?, - DataType::Int64 => relu_i64(tensor, &mut output_data)?, - DataType::Bool => { - return Err(MinitensorError::invalid_operation( - "ReLU function not supported for boolean tensors", - )); - } - }; - - // Create output tensor - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - // Set up gradient function if needed - if output.requires_grad() { - let grad_fn = Arc::new(ReluBackward { - input_id: tensor.id(), - mask, - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&output_with_grad, Some(grad_fn))?; - - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Hardshrink activation that thresholds values to zero within ``[-lambd, lambd]`` -pub fn hardshrink(tensor: &Tensor, lambd: f64) -> Result { - if lambd < 0.0 { - return Err(MinitensorError::invalid_operation( - "hardshrink requires lambd to be non-negative", - )); - } - - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - let store_mask = tensor.requires_grad(); - let mask = match tensor.dtype() { - DataType::Float32 => hardshrink_f32(tensor, &mut output_data, lambd as f32, store_mask)?, - DataType::Float64 => hardshrink_f64(tensor, &mut output_data, lambd, store_mask)?, - _ => { - return Err(MinitensorError::invalid_operation( - "hardshrink is only supported for floating point tensors", - )); - } - }; - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(HardshrinkBackward { - input_id: tensor.id(), - mask: mask.ok_or_else(|| { - MinitensorError::internal_error( - "hardshrink mask missing despite gradients being required", - ) - })?, - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - debug_assert!(mask.is_none()); - Ok(output) - } -} - -/// LeakyReLU activation function with gradient support -pub fn leaky_relu(tensor: &Tensor, negative_slope: f64) -> Result { - // Create output tensor data - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - // Perform LeakyReLU based on data type and capture mask of positive inputs - let mask = match tensor.dtype() { - DataType::Float32 => leaky_relu_f32(tensor, &mut output_data, negative_slope as f32)?, - DataType::Float64 => leaky_relu_f64(tensor, &mut output_data, negative_slope)?, - _ => { - return Err(MinitensorError::invalid_operation( - "LeakyReLU function only supported for floating point tensors", - )); - } - }; - - // Create output tensor - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - // Set up gradient function if needed - if output.requires_grad() { - let grad_fn = Arc::new(LeakyReluBackward { - input_id: tensor.id(), - negative_slope, - mask, - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&output_with_grad, Some(grad_fn))?; - - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Softmax activation function with gradient support -pub fn softmax(tensor: &Tensor, dim: Option) -> Result { - if tensor.ndim() == 0 { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - match tensor.dtype() { - DataType::Float32 => { - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from output data", - ) - })?; - output_slice[0] = 1.0; - } - DataType::Float64 => { - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from output data", - ) - })?; - output_slice[0] = 1.0; - } - _ => { - return Err(MinitensorError::invalid_operation( - "Softmax function only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(SoftmaxBackward { - input_id: tensor.id(), - output: output.detach(), - dim: 0, - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - return Ok(output_with_grad); - } - - return Ok(output); - } - - let dim = dim.unwrap_or(tensor.ndim() - 1); - - if dim >= tensor.ndim() { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - - // Create output tensor data - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - // Perform softmax based on data type - match tensor.dtype() { - DataType::Float32 => softmax_f32(tensor, &mut output_data, dim)?, - DataType::Float64 => softmax_f64(tensor, &mut output_data, dim)?, - _ => { - return Err(MinitensorError::invalid_operation( - "Softmax function only supported for floating point tensors", - )); - } - } - - // Create output tensor - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - // Set up gradient function if needed - if output.requires_grad() { - let grad_fn = Arc::new(SoftmaxBackward { - input_id: tensor.id(), - output: output.detach(), - dim, - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&output_with_grad, Some(grad_fn))?; - - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Log-Softmax activation function with gradient support -pub fn log_softmax(tensor: &Tensor, dim: Option) -> Result { - if tensor.ndim() == 0 { - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - match tensor.dtype() { - DataType::Float32 => { - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from output data", - ) - })?; - output_slice[0] = 0.0; - } - DataType::Float64 => { - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from output data", - ) - })?; - output_slice[0] = 0.0; - } - _ => { - return Err(MinitensorError::invalid_operation( - "LogSoftmax function only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(LogSoftmaxBackward { - input_id: tensor.id(), - output: output.detach(), - dim: 0, - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - return Ok(output_with_grad); - } - - return Ok(output); - } - - let dim = dim.unwrap_or(tensor.ndim() - 1); - - if dim >= tensor.ndim() { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => log_softmax_f32(tensor, &mut output_data, dim)?, - DataType::Float64 => log_softmax_f64(tensor, &mut output_data, dim)?, - _ => { - return Err(MinitensorError::invalid_operation( - "LogSoftmax function only supported for floating point tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(LogSoftmaxBackward { - input_id: tensor.id(), - output: output.detach(), - dim, - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - - add_to_graph(&output_with_grad, Some(grad_fn))?; - - Ok(output_with_grad) - } else { - Ok(output) - } -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::autograd::EluBackward; +use crate::autograd::GeluBackward; +use crate::autograd::HardshrinkBackward; +use crate::autograd::LeakyReluBackward; +use crate::autograd::LogAddExpBackward; +use crate::autograd::LogSoftmaxBackward; +use crate::autograd::PowBackward; +use crate::autograd::PowBroadcast; +use crate::autograd::ReluBackward; +use crate::autograd::SeluBackward; +use crate::autograd::SiluBackward; +use crate::autograd::SoftmaxBackward; +use crate::autograd::SoftplusBackward; +use crate::autograd::SoftsignBackward; +use crate::{ + autograd::add_to_graph, + error::{MinitensorError, Result}, + tensor::{DataType, Tensor, TensorData}, +}; +use std::sync::Arc; + +/// Element-wise power with tensor exponent and gradient support +pub fn pow(base: &Tensor, exponent: &Tensor) -> Result { + // Check device and dtype compatibility + if base.device() != exponent.device() { + return Err(MinitensorError::device_mismatch( + format!("{:?}", base.device()), + format!("{:?}", exponent.device()), + )); + } + + if base.dtype() != exponent.dtype() { + return Err(MinitensorError::type_mismatch( + format!("{:?}", base.dtype()), + format!("{:?}", exponent.dtype()), + )); + } + + let base_shape = base.shape().clone(); + let exponent_shape = exponent.shape().clone(); + let base_numel = base_shape.numel(); + let exp_numel = exponent_shape.numel(); + + let broadcast = if base_shape == exponent_shape { + PowBroadcast::None + } else if base_numel == 1 { + PowBroadcast::BaseScalar + } else if exp_numel == 1 { + PowBroadcast::ExponentScalar + } else { + // General broadcasting: materialize both operands at the broadcast + // shape (expand and contiguous are grad-aware) and recurse into the + // same-shape fast path. + let out_shape = base_shape.broadcast_with(&exponent_shape)?; + let dims: Vec = out_shape.dims().iter().map(|&d| d as isize).collect(); + let base_b = base.expand(dims.clone())?.contiguous()?; + let exp_b = exponent.expand(dims)?.contiguous()?; + return pow(&base_b, &exp_b); + }; + + let output_shape = match broadcast { + PowBroadcast::None | PowBroadcast::ExponentScalar => base_shape.clone(), + PowBroadcast::BaseScalar => exponent_shape.clone(), + }; + + let mut output_data = + TensorData::uninitialized_on_device(output_shape.numel(), base.dtype(), base.device()); + + match base.dtype() { + DataType::Float32 => { + let b = base.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from base tensor") + })?; + let e = exponent.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from exponent tensor") + })?; + let out = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + match broadcast { + PowBroadcast::None => { + for i in 0..b.len() { + out[i] = b[i].powf(e[i]); + } + } + PowBroadcast::BaseScalar => { + let base_val = b[0]; + for i in 0..e.len() { + out[i] = base_val.powf(e[i]); + } + } + PowBroadcast::ExponentScalar => { + let exp_val = e[0]; + for i in 0..b.len() { + out[i] = b[i].powf(exp_val); + } + } + } + } + DataType::Float64 => { + let b = base.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from base tensor") + })?; + let e = exponent.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from exponent tensor") + })?; + let out = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + match broadcast { + PowBroadcast::None => { + for i in 0..b.len() { + out[i] = b[i].powf(e[i]); + } + } + PowBroadcast::BaseScalar => { + let base_val = b[0]; + for i in 0..e.len() { + out[i] = base_val.powf(e[i]); + } + } + PowBroadcast::ExponentScalar => { + let exp_val = e[0]; + for i in 0..b.len() { + out[i] = b[i].powf(exp_val); + } + } + } + } + _ => { + return Err(MinitensorError::invalid_operation( + "Power operation only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + output_shape, + base.dtype(), + base.device(), + base.requires_grad() || exponent.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(PowBackward { + base: base.detach(), + exponent: exponent.detach(), + output: output.clone().detach(), + input_ids: [base.id(), exponent.id()], + base_requires_grad: base.requires_grad(), + exp_requires_grad: exponent.requires_grad(), + broadcast, + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Element-wise power with scalar exponent and gradient support +pub fn powf(tensor: &Tensor, exponent: f64) -> Result { + // Create exponent tensor filled with scalar value + let mut exp_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + match tensor.dtype() { + DataType::Float32 => { + let slice = exp_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from exponent data", + ) + })?; + for val in slice.iter_mut() { + *val = exponent as f32; + } + } + DataType::Float64 => { + let slice = exp_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from exponent data", + ) + })?; + for val in slice.iter_mut() { + *val = exponent; + } + } + _ => { + return Err(MinitensorError::invalid_operation( + "Power operation only supported for floating point tensors", + )); + } + } + let exp_tensor = Tensor::new( + Arc::new(exp_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + false, + ); + pow(tensor, &exp_tensor) +} + +/// Numerically stable logaddexp with gradient support +pub fn logaddexp(lhs: &Tensor, rhs: &Tensor) -> Result { + if lhs.device() != rhs.device() { + return Err(MinitensorError::device_mismatch( + format!("{:?}", lhs.device()), + format!("{:?}", rhs.device()), + )); + } + + let requires_grad = lhs.requires_grad() || rhs.requires_grad(); + use crate::operations::binary::{BinaryOpKind, coerce_binary_operands}; + let (lhs_cast, rhs_cast, result_dtype) = coerce_binary_operands(lhs, rhs, BinaryOpKind::Add)?; + + let lhs_tensor = match lhs_cast { + std::borrow::Cow::Borrowed(t) => t.clone(), + std::borrow::Cow::Owned(t) => t, + }; + let rhs_tensor = match rhs_cast { + std::borrow::Cow::Borrowed(t) => t.clone(), + std::borrow::Cow::Owned(t) => t, + }; + + if result_dtype != DataType::Float32 && result_dtype != DataType::Float64 { + return Err(MinitensorError::invalid_operation( + "logaddexp is only supported for floating point tensors", + )); + } + + let output_shape = lhs_tensor.shape().broadcast_with(rhs_tensor.shape())?; + let mut output_data = + TensorData::uninitialized_on_device(output_shape.numel(), result_dtype, lhs.device()); + + match result_dtype { + DataType::Float32 => { + logaddexp_f32(&lhs_tensor, &rhs_tensor, &mut output_data, &output_shape)? + } + DataType::Float64 => { + logaddexp_f64(&lhs_tensor, &rhs_tensor, &mut output_data, &output_shape)? + } + _ => unreachable!(), + } + + let output = Tensor::new( + Arc::new(output_data), + output_shape.clone(), + result_dtype, + lhs.device(), + requires_grad, + ); + + if requires_grad { + let grad_fn = Arc::new(LogAddExpBackward { + lhs: lhs_tensor.detach(), + rhs: rhs_tensor.detach(), + output: output.clone().detach(), + input_ids: [lhs.id(), rhs.id()], + input_shapes: [lhs.shape().dims().to_vec(), rhs.shape().dims().to_vec()], + input_requires_grad: [lhs.requires_grad(), rhs.requires_grad()], + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Softplus activation function with gradient support +pub fn softplus(tensor: &Tensor, beta: f64, threshold: f64) -> Result { + if beta <= 0.0 { + return Err(MinitensorError::invalid_argument( + "softplus beta must be positive", + )); + } + + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => softplus_f32(tensor, &mut output_data, beta as f32, threshold as f32)?, + DataType::Float64 => softplus_f64(tensor, &mut output_data, beta, threshold)?, + _ => { + return Err(MinitensorError::invalid_operation( + "Softplus is only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(SoftplusBackward { + input_id: tensor.id(), + input: tensor.clone().detach(), + beta, + threshold, + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// GELU activation function with optional tanh approximation +pub fn gelu(tensor: &Tensor, approximate: bool) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => gelu_f32(tensor, &mut output_data, approximate)?, + DataType::Float64 => gelu_f64(tensor, &mut output_data, approximate)?, + _ => { + return Err(MinitensorError::invalid_operation( + "GELU is only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(GeluBackward { + input_id: tensor.id(), + input: tensor.clone().detach(), + approximate, + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// ELU activation function with configurable alpha +pub fn elu(tensor: &Tensor, alpha: f64) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => elu_f32(tensor, &mut output_data, alpha as f32)?, + DataType::Float64 => elu_f64(tensor, &mut output_data, alpha)?, + _ => { + return Err(MinitensorError::invalid_operation( + "ELU is only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(EluBackward { + input_id: tensor.id(), + output: output.clone().detach(), + alpha, + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// SELU activation function. +pub fn selu(tensor: &Tensor) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => selu_f32(tensor, &mut output_data)?, + DataType::Float64 => selu_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "SELU is only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(SeluBackward { + input_id: tensor.id(), + output: output.clone().detach(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// SiLU (Swish) activation function with gradient support +pub fn silu(tensor: &Tensor) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => silu_f32(tensor, &mut output_data)?, + DataType::Float64 => silu_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "SiLU is only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(SiluBackward { + input_id: tensor.id(), + input: tensor.clone().detach(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Softsign activation function with gradient support +pub fn softsign(tensor: &Tensor) -> Result { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => softsign_f32(tensor, &mut output_data)?, + DataType::Float64 => softsign_f64(tensor, &mut output_data)?, + _ => { + return Err(MinitensorError::invalid_operation( + "Softsign is only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(SoftsignBackward { + input_id: tensor.id(), + input: tensor.clone().detach(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// ReLU activation function with gradient support +pub fn relu(tensor: &Tensor) -> Result { + // Create output tensor data + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + // Perform ReLU based on data type while capturing mask of positive inputs + let mask = match tensor.dtype() { + DataType::Float32 => relu_f32(tensor, &mut output_data)?, + DataType::Float64 => relu_f64(tensor, &mut output_data)?, + DataType::Int32 => relu_i32(tensor, &mut output_data)?, + DataType::Int64 => relu_i64(tensor, &mut output_data)?, + DataType::Bool => { + return Err(MinitensorError::invalid_operation( + "ReLU function not supported for boolean tensors", + )); + } + }; + + // Create output tensor + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + // Set up gradient function if needed + if output.requires_grad() { + let grad_fn = Arc::new(ReluBackward { + input_id: tensor.id(), + mask, + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&output_with_grad, Some(grad_fn))?; + + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Hardshrink activation that thresholds values to zero within ``[-lambd, lambd]`` +pub fn hardshrink(tensor: &Tensor, lambd: f64) -> Result { + if lambd < 0.0 { + return Err(MinitensorError::invalid_operation( + "hardshrink requires lambd to be non-negative", + )); + } + + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + let store_mask = tensor.requires_grad(); + let mask = match tensor.dtype() { + DataType::Float32 => hardshrink_f32(tensor, &mut output_data, lambd as f32, store_mask)?, + DataType::Float64 => hardshrink_f64(tensor, &mut output_data, lambd, store_mask)?, + _ => { + return Err(MinitensorError::invalid_operation( + "hardshrink is only supported for floating point tensors", + )); + } + }; + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(HardshrinkBackward { + input_id: tensor.id(), + mask: mask.ok_or_else(|| { + MinitensorError::internal_error( + "hardshrink mask missing despite gradients being required", + ) + })?, + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + debug_assert!(mask.is_none()); + Ok(output) + } +} + +/// LeakyReLU activation function with gradient support +pub fn leaky_relu(tensor: &Tensor, negative_slope: f64) -> Result { + // Create output tensor data + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + // Perform LeakyReLU based on data type and capture mask of positive inputs + let mask = match tensor.dtype() { + DataType::Float32 => leaky_relu_f32(tensor, &mut output_data, negative_slope as f32)?, + DataType::Float64 => leaky_relu_f64(tensor, &mut output_data, negative_slope)?, + _ => { + return Err(MinitensorError::invalid_operation( + "LeakyReLU function only supported for floating point tensors", + )); + } + }; + + // Create output tensor + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + // Set up gradient function if needed + if output.requires_grad() { + let grad_fn = Arc::new(LeakyReluBackward { + input_id: tensor.id(), + negative_slope, + mask, + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&output_with_grad, Some(grad_fn))?; + + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Softmax activation function with gradient support +pub fn softmax(tensor: &Tensor, dim: Option) -> Result { + if tensor.ndim() == 0 { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + match tensor.dtype() { + DataType::Float32 => { + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from output data", + ) + })?; + output_slice[0] = 1.0; + } + DataType::Float64 => { + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from output data", + ) + })?; + output_slice[0] = 1.0; + } + _ => { + return Err(MinitensorError::invalid_operation( + "Softmax function only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(SoftmaxBackward { + input_id: tensor.id(), + output: output.detach(), + dim: 0, + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + return Ok(output_with_grad); + } + + return Ok(output); + } + + let dim = dim.unwrap_or(tensor.ndim() - 1); + + if dim >= tensor.ndim() { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + + // Create output tensor data + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + // Perform softmax based on data type + match tensor.dtype() { + DataType::Float32 => softmax_f32(tensor, &mut output_data, dim)?, + DataType::Float64 => softmax_f64(tensor, &mut output_data, dim)?, + _ => { + return Err(MinitensorError::invalid_operation( + "Softmax function only supported for floating point tensors", + )); + } + } + + // Create output tensor + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + // Set up gradient function if needed + if output.requires_grad() { + let grad_fn = Arc::new(SoftmaxBackward { + input_id: tensor.id(), + output: output.detach(), + dim, + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&output_with_grad, Some(grad_fn))?; + + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Log-Softmax activation function with gradient support +pub fn log_softmax(tensor: &Tensor, dim: Option) -> Result { + if tensor.ndim() == 0 { + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + match tensor.dtype() { + DataType::Float32 => { + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from output data", + ) + })?; + output_slice[0] = 0.0; + } + DataType::Float64 => { + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from output data", + ) + })?; + output_slice[0] = 0.0; + } + _ => { + return Err(MinitensorError::invalid_operation( + "LogSoftmax function only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(LogSoftmaxBackward { + input_id: tensor.id(), + output: output.detach(), + dim: 0, + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + return Ok(output_with_grad); + } + + return Ok(output); + } + + let dim = dim.unwrap_or(tensor.ndim() - 1); + + if dim >= tensor.ndim() { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => log_softmax_f32(tensor, &mut output_data, dim)?, + DataType::Float64 => log_softmax_f64(tensor, &mut output_data, dim)?, + _ => { + return Err(MinitensorError::invalid_operation( + "LogSoftmax function only supported for floating point tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(LogSoftmaxBackward { + input_id: tensor.id(), + output: output.detach(), + dim, + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + + add_to_graph(&output_with_grad, Some(grad_fn))?; + + Ok(output_with_grad) + } else { + Ok(output) + } +} diff --git a/engine/src/operations/arithmetic.rs b/engine/src/operations/arithmetic.rs index b0552df4..cb8d3ab7 100644 --- a/engine/src/operations/arithmetic.rs +++ b/engine/src/operations/arithmetic.rs @@ -4,5 +4,10 @@ // This source code is licensed under the Apache-style license found in the // LICENSE file in the root directory of this source tree. -include!("arithmetic/elementwise.rs"); -include!("arithmetic/kernels.rs"); +#[path = "arithmetic/elementwise.rs"] +mod elementwise_impl; +#[path = "arithmetic/kernels.rs"] +mod kernels_impl; + +pub use self::elementwise_impl::*; +pub(crate) use self::kernels_impl::*; diff --git a/engine/src/operations/arithmetic/elementwise.rs b/engine/src/operations/arithmetic/elementwise.rs index 1e3ab589..9d1777cd 100644 --- a/engine/src/operations/arithmetic/elementwise.rs +++ b/engine/src/operations/arithmetic/elementwise.rs @@ -1,883 +1,507 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -use crate::{ - autograd::{AddBackward, DivBackward, MulBackward, NegBackward, SubBackward, add_to_graph}, - error::{MinitensorError, Result}, - operations::{ - binary::{BinaryOpKind, coerce_binary_operands}, - simd::{ - can_use_simd_fast_path, simd_add_f32, simd_add_f64, simd_div_f32, simd_div_f64, - simd_mul_f32, simd_mul_f64, simd_sub_f32, simd_sub_f64, - }, - }, - tensor::{DataType, Shape, Strides, Tensor, TensorData}, -}; -use rayon::prelude::*; -use smallvec::{SmallVec, smallvec}; -use std::sync::Arc; - -const PAR_THRESHOLD: usize = 1 << 12; // 4096 elements - -/// Element-wise addition with broadcasting support -pub fn add(lhs: &Tensor, rhs: &Tensor) -> Result { - // Check device compatibility - if lhs.device() != rhs.device() { - return Err(MinitensorError::device_mismatch( - format!("{:?}", lhs.device()), - format!("{:?}", rhs.device()), - )); - } - - let requires_grad = lhs.requires_grad() || rhs.requires_grad(); - let (lhs_cast, rhs_cast, result_dtype) = coerce_binary_operands(lhs, rhs, BinaryOpKind::Add)?; - let lhs_ref = lhs_cast.as_ref(); - let rhs_ref = rhs_cast.as_ref(); - - // Compute broadcasted shape - let output_shape = lhs_ref.shape().broadcast_with(rhs_ref.shape())?; - - if output_shape.numel() == 0 { - let mut output = Tensor::empty( - output_shape.clone(), - result_dtype, - lhs.device(), - requires_grad, - ); - - if requires_grad { - let grad_fn = Arc::new(AddBackward { - input_shapes: [lhs.shape().dims().to_vec(), rhs.shape().dims().to_vec()], - input_ids: [lhs.id(), rhs.id()], - }); - output.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output, Some(grad_fn))?; - } - - return Ok(output); - } - - // Create output tensor data - let mut output_data = - TensorData::uninitialized_on_device(output_shape.numel(), result_dtype, lhs.device()); - - // Perform element-wise addition based on data type - match result_dtype { - DataType::Float32 => add_f32_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, - DataType::Float64 => add_f64_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, - DataType::Int32 => add_i32_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, - DataType::Int64 => add_i64_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, - DataType::Bool => add_bool_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, - } - - // Create output tensor - let mut output = Tensor::new( - Arc::new(output_data), - output_shape.clone(), - result_dtype, - lhs.device(), - requires_grad, - ); - - // Set up gradient function if needed - if requires_grad { - let grad_fn = Arc::new(AddBackward { - input_shapes: [lhs.shape().dims().to_vec(), rhs.shape().dims().to_vec()], - input_ids: [lhs.id(), rhs.id()], - }); - - output.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&output, Some(grad_fn))?; - } - - Ok(output) -} - -/// In-place element-wise addition used for gradient accumulation -pub fn add_inplace(lhs: &mut Tensor, rhs: &Tensor) -> Result<()> { - if lhs.shape() != rhs.shape() { - return Err(MinitensorError::shape_mismatch( - lhs.shape().dims().to_vec(), - rhs.shape().dims().to_vec(), - )); - } - if lhs.dtype() != rhs.dtype() { - return Err(MinitensorError::type_mismatch( - format!("{:?}", lhs.dtype()), - format!("{:?}", rhs.dtype()), - )); - } - if lhs.device() != rhs.device() { - return Err(MinitensorError::device_mismatch( - format!("{:?}", lhs.device()), - format!("{:?}", rhs.device()), - )); - } - if std::sync::Arc::strong_count(lhs.data()) > 1 { - // Fallback to out-of-place addition if data is shared - let tmp = add(lhs, rhs)?; - *lhs = tmp; - return Ok(()); - } - - match lhs.dtype() { - DataType::Float32 => { - let lhs_slice = lhs.data_mut().as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from lhs tensor") - })?; - let rhs_slice = rhs.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from rhs tensor") - })?; - let len = lhs_slice.len(); - if len < PAR_THRESHOLD { - for i in 0..len { - lhs_slice[i] += rhs_slice[i]; - } - } else { - let lhs_ptr = lhs_slice.as_mut_ptr() as usize; - let rhs_ptr = rhs_slice.as_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let lhs_ptr = lhs_ptr as *mut f32; - let rhs_ptr = rhs_ptr as *const f32; - *lhs_ptr.add(i) += *rhs_ptr.add(i); - }); - } - } - DataType::Float64 => { - let lhs_slice = lhs.data_mut().as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from lhs tensor") - })?; - let rhs_slice = rhs.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from rhs tensor") - })?; - let len = lhs_slice.len(); - if len < PAR_THRESHOLD { - for i in 0..len { - lhs_slice[i] += rhs_slice[i]; - } - } else { - let lhs_ptr = lhs_slice.as_mut_ptr() as usize; - let rhs_ptr = rhs_slice.as_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let lhs_ptr = lhs_ptr as *mut f64; - let rhs_ptr = rhs_ptr as *const f64; - *lhs_ptr.add(i) += *rhs_ptr.add(i); - }); - } - } - DataType::Int32 => { - let lhs_slice = lhs.data_mut().as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i32 slice from lhs tensor") - })?; - let rhs_slice = rhs.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from rhs tensor") - })?; - let len = lhs_slice.len(); - if len < PAR_THRESHOLD { - for i in 0..len { - lhs_slice[i] += rhs_slice[i]; - } - } else { - let lhs_ptr = lhs_slice.as_mut_ptr() as usize; - let rhs_ptr = rhs_slice.as_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let lhs_ptr = lhs_ptr as *mut i32; - let rhs_ptr = rhs_ptr as *const i32; - *lhs_ptr.add(i) += *rhs_ptr.add(i); - }); - } - } - DataType::Int64 => { - let lhs_slice = lhs.data_mut().as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice from lhs tensor") - })?; - let rhs_slice = rhs.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from rhs tensor") - })?; - let len = lhs_slice.len(); - if len < PAR_THRESHOLD { - for i in 0..len { - lhs_slice[i] += rhs_slice[i]; - } - } else { - let lhs_ptr = lhs_slice.as_mut_ptr() as usize; - let rhs_ptr = rhs_slice.as_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let lhs_ptr = lhs_ptr as *mut i64; - let rhs_ptr = rhs_ptr as *const i64; - *lhs_ptr.add(i) += *rhs_ptr.add(i); - }); - } - } - DataType::Bool => { - let lhs_slice = lhs.data_mut().as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice from lhs tensor") - })?; - let rhs_slice = rhs.data().as_bool_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from rhs tensor") - })?; - let len = lhs_slice.len(); - if len < PAR_THRESHOLD { - for i in 0..len { - lhs_slice[i] = lhs_slice[i] || rhs_slice[i]; - } - } else { - let lhs_ptr = lhs_slice.as_mut_ptr() as usize; - let rhs_ptr = rhs_slice.as_ptr() as usize; - (0..len).into_par_iter().for_each(|i| unsafe { - let lhs_ptr = lhs_ptr as *mut bool; - let rhs_ptr = rhs_ptr as *const bool; - *lhs_ptr.add(i) = *lhs_ptr.add(i) || *rhs_ptr.add(i); - }); - } - } - } - Ok(()) -} - -/// Element-wise subtraction with broadcasting support -pub fn sub(lhs: &Tensor, rhs: &Tensor) -> Result { - // Check device compatibility - if lhs.device() != rhs.device() { - return Err(MinitensorError::device_mismatch( - format!("{:?}", lhs.device()), - format!("{:?}", rhs.device()), - )); - } - - let requires_grad = lhs.requires_grad() || rhs.requires_grad(); - let (lhs_cast, rhs_cast, result_dtype) = coerce_binary_operands(lhs, rhs, BinaryOpKind::Sub)?; - let lhs_ref = lhs_cast.as_ref(); - let rhs_ref = rhs_cast.as_ref(); - - // Compute broadcasted shape - let output_shape = lhs_ref.shape().broadcast_with(rhs_ref.shape())?; - - if output_shape.numel() == 0 { - let mut output = Tensor::empty( - output_shape.clone(), - result_dtype, - lhs.device(), - requires_grad, - ); - - if requires_grad { - let grad_fn = Arc::new(SubBackward { - input_shapes: [lhs.shape().dims().to_vec(), rhs.shape().dims().to_vec()], - input_ids: [lhs.id(), rhs.id()], - }); - output.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output, Some(grad_fn))?; - } - - return Ok(output); - } - - // Create output tensor data - let mut output_data = - TensorData::uninitialized_on_device(output_shape.numel(), result_dtype, lhs.device()); - - // Perform element-wise subtraction based on data type - match result_dtype { - DataType::Float32 => sub_f32_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, - DataType::Float64 => sub_f64_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, - DataType::Int32 => sub_i32_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, - DataType::Int64 => sub_i64_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, - DataType::Bool => unreachable!("boolean subtraction should be rejected during coercion"), - } - - // Create output tensor - let mut output = Tensor::new( - Arc::new(output_data), - output_shape.clone(), - result_dtype, - lhs.device(), - requires_grad, - ); - - // Set up gradient function if needed - if requires_grad { - let grad_fn = Arc::new(SubBackward { - input_shapes: [lhs.shape().dims().to_vec(), rhs.shape().dims().to_vec()], - input_ids: [lhs.id(), rhs.id()], - }); - - output.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output, Some(grad_fn))?; - } - - Ok(output) -} - -/// Element-wise multiplication with broadcasting support -pub fn mul(lhs: &Tensor, rhs: &Tensor) -> Result { - // Check device compatibility - if lhs.device() != rhs.device() { - return Err(MinitensorError::device_mismatch( - format!("{:?}", lhs.device()), - format!("{:?}", rhs.device()), - )); - } - - let requires_grad = lhs.requires_grad() || rhs.requires_grad(); - let (lhs_cast, rhs_cast, result_dtype) = coerce_binary_operands(lhs, rhs, BinaryOpKind::Mul)?; - let lhs_ref = lhs_cast.as_ref(); - let rhs_ref = rhs_cast.as_ref(); - - // Compute broadcasted shape - let output_shape = lhs_ref.shape().broadcast_with(rhs_ref.shape())?; - - if output_shape.numel() == 0 { - let mut output = Tensor::empty( - output_shape.clone(), - result_dtype, - lhs.device(), - requires_grad, - ); - - if requires_grad { - let grad_fn = Arc::new(MulBackward { - lhs: lhs.clone(), - rhs: rhs.clone(), - input_ids: [lhs.id(), rhs.id()], - }); - output.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output, Some(grad_fn))?; - } - - return Ok(output); - } - - // Create output tensor data - let mut output_data = - TensorData::uninitialized_on_device(output_shape.numel(), result_dtype, lhs.device()); - - // Perform element-wise multiplication based on data type - match result_dtype { - DataType::Float32 => mul_f32_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, - DataType::Float64 => mul_f64_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, - DataType::Int32 => mul_i32_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, - DataType::Int64 => mul_i64_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, - DataType::Bool => mul_bool_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, - } - - // Create output tensor - let mut output = Tensor::new( - Arc::new(output_data), - output_shape.clone(), - result_dtype, - lhs.device(), - requires_grad, - ); - - // Set up gradient function if needed - if requires_grad { - let grad_fn = Arc::new(MulBackward { - lhs: lhs.clone(), - rhs: rhs.clone(), - input_ids: [lhs.id(), rhs.id()], - }); - - output.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output, Some(grad_fn))?; - } - - Ok(output) -} - -/// Element-wise division with broadcasting support -pub fn div(lhs: &Tensor, rhs: &Tensor) -> Result { - // Check device compatibility - if lhs.device() != rhs.device() { - return Err(MinitensorError::device_mismatch( - format!("{:?}", lhs.device()), - format!("{:?}", rhs.device()), - )); - } - - let requires_grad = lhs.requires_grad() || rhs.requires_grad(); - let (lhs_cast, rhs_cast, result_dtype) = coerce_binary_operands(lhs, rhs, BinaryOpKind::Div)?; - let lhs_ref = lhs_cast.as_ref(); - let rhs_ref = rhs_cast.as_ref(); - - // Compute broadcasted shape - let output_shape = lhs_ref.shape().broadcast_with(rhs_ref.shape())?; - - if output_shape.numel() == 0 { - let mut output = Tensor::empty( - output_shape.clone(), - result_dtype, - lhs.device(), - requires_grad, - ); - - if requires_grad { - let grad_fn = Arc::new(DivBackward { - lhs: lhs.clone(), - rhs: rhs.clone(), - input_ids: [lhs.id(), rhs.id()], - }); - output.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output, Some(grad_fn))?; - } - - return Ok(output); - } - - // Create output tensor data - let mut output_data = - TensorData::uninitialized_on_device(output_shape.numel(), result_dtype, lhs.device()); - - // Perform element-wise division based on data type - match result_dtype { - DataType::Float32 => div_f32_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, - DataType::Float64 => div_f64_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, - DataType::Int32 | DataType::Int64 | DataType::Bool => { - unreachable!("integer and boolean division should coerce to floating point") - } - } - - // Create output tensor - let mut output = Tensor::new( - Arc::new(output_data), - output_shape.clone(), - result_dtype, - lhs.device(), - requires_grad, - ); - - // Set up gradient function if needed - if requires_grad { - let grad_fn = Arc::new(DivBackward { - lhs: lhs.clone(), - rhs: rhs.clone(), - input_ids: [lhs.id(), rhs.id()], - }); - - output.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output, Some(grad_fn))?; - } - - Ok(output) -} - -/// Element-wise negation -pub fn neg(tensor: &Tensor) -> Result { - let mut output_data = TensorData::uninitialized_on_device( - tensor.shape().numel(), - tensor.dtype(), - tensor.device(), - ); - - match tensor.dtype() { - DataType::Float32 => { - let input = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from tensor") - })?; - let output = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output") - })?; - if input.len() >= PAR_THRESHOLD { - output - .par_iter_mut() - .zip(input.par_iter()) - .for_each(|(o, &i)| *o = -i); - } else { - for (o, &i) in output.iter_mut().zip(input.iter()) { - *o = -i; - } - } - } - DataType::Float64 => { - let input = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from tensor") - })?; - let output = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output") - })?; - if input.len() >= PAR_THRESHOLD { - output - .par_iter_mut() - .zip(input.par_iter()) - .for_each(|(o, &i)| *o = -i); - } else { - for (o, &i) in output.iter_mut().zip(input.iter()) { - *o = -i; - } - } - } - DataType::Int32 => { - let input = tensor.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from tensor") - })?; - let output = output_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i32 slice from output") - })?; - if input.len() >= PAR_THRESHOLD { - output - .par_iter_mut() - .zip(input.par_iter()) - .for_each(|(o, &i)| *o = -i); - } else { - for (o, &i) in output.iter_mut().zip(input.iter()) { - *o = -i; - } - } - } - DataType::Int64 => { - let input = tensor.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from tensor") - })?; - let output = output_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice from output") - })?; - if input.len() >= PAR_THRESHOLD { - output - .par_iter_mut() - .zip(input.par_iter()) - .for_each(|(o, &i)| *o = -i); - } else { - for (o, &i) in output.iter_mut().zip(input.iter()) { - *o = -i; - } - } - } - DataType::Bool => { - return Err(MinitensorError::invalid_operation( - "Negation not supported for boolean tensors", - )); - } - } - - let output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if output.requires_grad() { - let grad_fn = Arc::new(NegBackward { - input_id: tensor.id(), - }); - let mut out_with_grad = output; - out_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&out_with_grad, Some(grad_fn))?; - Ok(out_with_grad) - } else { - Ok(output) - } -} - -// Helper functions for type-specific operations - -fn add_f32_direct( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from rhs tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - // Use SIMD fast path if possible (no broadcasting, same shapes) - if can_use_simd_fast_path(lhs.shape(), rhs.shape(), output_shape) { - simd_add_f32(lhs_data, rhs_data, output_slice) - } else { - broadcast_binary_op( - lhs_data, - rhs_data, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - |a, b| a + b, - ) - } -} - -fn add_f64_direct( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from rhs tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - // Use SIMD fast path if possible (no broadcasting, same shapes) - if can_use_simd_fast_path(lhs.shape(), rhs.shape(), output_shape) { - simd_add_f64(lhs_data, rhs_data, output_slice) - } else { - broadcast_binary_op( - lhs_data, - rhs_data, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - |a, b| a + b, - ) - } -} - -fn add_i32_direct( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from rhs tensor") - })?; - - let output_slice = output_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i32 slice from output data") - })?; - - broadcast_binary_op( - lhs_data, - rhs_data, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - |a, b| a + b, - ) -} - -fn add_i64_direct( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from rhs tensor") - })?; - - let output_slice = output_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice from output data") - })?; - - broadcast_binary_op( - lhs_data, - rhs_data, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - |a, b| a + b, - ) -} - -fn add_bool_direct( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_bool_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_bool_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from rhs tensor") - })?; - - let output_slice = output_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice from output data") - })?; - - broadcast_binary_op( - lhs_data, - rhs_data, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - |a, b| a || b, - ) -} - -fn sub_f32_direct( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from rhs tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - // Use SIMD fast path if possible (no broadcasting, same shapes) - if can_use_simd_fast_path(lhs.shape(), rhs.shape(), output_shape) { - simd_sub_f32(lhs_data, rhs_data, output_slice) - } else { - broadcast_binary_op( - lhs_data, - rhs_data, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - |a, b| a - b, - ) - } -} - -fn sub_f64_direct( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from rhs tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - // Use SIMD fast path if possible (no broadcasting, same shapes) - if can_use_simd_fast_path(lhs.shape(), rhs.shape(), output_shape) { - simd_sub_f64(lhs_data, rhs_data, output_slice) - } else { - broadcast_binary_op( - lhs_data, - rhs_data, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - |a, b| a - b, - ) - } -} - -fn sub_i32_direct( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from rhs tensor") - })?; - - let output_slice = output_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i32 slice from output data") - })?; - - broadcast_binary_op( - lhs_data, - rhs_data, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - |a, b| a - b, - ) -} - -fn sub_i64_direct( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from rhs tensor") - })?; - - let output_slice = output_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice from output data") - })?; - - broadcast_binary_op( - lhs_data, - rhs_data, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - |a, b| a - b, - ) -} - -fn mul_f32_direct( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from rhs tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - // Use SIMD fast path if possible (no broadcasting, same shapes) - if can_use_simd_fast_path(lhs.shape(), rhs.shape(), output_shape) { - simd_mul_f32(lhs_data, rhs_data, output_slice) - } else { - broadcast_binary_op( - lhs_data, - rhs_data, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - |a, b| a * b, - ) - } -} +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; + +use crate::{ + autograd::{AddBackward, DivBackward, MulBackward, NegBackward, SubBackward, add_to_graph}, + error::{MinitensorError, Result}, + operations::binary::{BinaryOpKind, coerce_binary_operands}, + tensor::{DataType, Tensor, TensorData}, +}; +use rayon::prelude::*; +use std::sync::Arc; + +pub(crate) const PAR_THRESHOLD: usize = 1 << 12; // 4096 elements + +/// Element-wise addition with broadcasting support +pub fn add(lhs: &Tensor, rhs: &Tensor) -> Result { + // Check device compatibility + if lhs.device() != rhs.device() { + return Err(MinitensorError::device_mismatch( + format!("{:?}", lhs.device()), + format!("{:?}", rhs.device()), + )); + } + + let requires_grad = lhs.requires_grad() || rhs.requires_grad(); + let (lhs_cast, rhs_cast, result_dtype) = coerce_binary_operands(lhs, rhs, BinaryOpKind::Add)?; + let lhs_ref = lhs_cast.as_ref(); + let rhs_ref = rhs_cast.as_ref(); + + // Compute broadcasted shape + let output_shape = lhs_ref.shape().broadcast_with(rhs_ref.shape())?; + + if output_shape.numel() == 0 { + let mut output = Tensor::empty( + output_shape.clone(), + result_dtype, + lhs.device(), + requires_grad, + ); + + if requires_grad { + let grad_fn = Arc::new(AddBackward { + input_shapes: [lhs.shape().dims().to_vec(), rhs.shape().dims().to_vec()], + input_ids: [lhs.id(), rhs.id()], + input_requires_grad: [lhs.requires_grad(), rhs.requires_grad()], + }); + output.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output, Some(grad_fn))?; + } + + return Ok(output); + } + + // Create output tensor data + let mut output_data = + TensorData::uninitialized_on_device(output_shape.numel(), result_dtype, lhs.device()); + + // Perform element-wise addition based on data type + match result_dtype { + DataType::Float32 => add_f32_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, + DataType::Float64 => add_f64_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, + DataType::Int32 => add_i32_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, + DataType::Int64 => add_i64_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, + DataType::Bool => add_bool_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, + } + + // Create output tensor + let mut output = Tensor::new( + Arc::new(output_data), + output_shape.clone(), + result_dtype, + lhs.device(), + requires_grad, + ); + + // Set up gradient function if needed + if requires_grad { + let grad_fn = Arc::new(AddBackward { + input_shapes: [lhs.shape().dims().to_vec(), rhs.shape().dims().to_vec()], + input_ids: [lhs.id(), rhs.id()], + input_requires_grad: [lhs.requires_grad(), rhs.requires_grad()], + }); + + output.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&output, Some(grad_fn))?; + } + + Ok(output) +} + +/// In-place element-wise addition used for gradient accumulation +pub fn add_inplace(lhs: &mut Tensor, rhs: &Tensor) -> Result<()> { + if lhs.shape() != rhs.shape() { + return Err(MinitensorError::shape_mismatch( + lhs.shape().dims().to_vec(), + rhs.shape().dims().to_vec(), + )); + } + if lhs.dtype() != rhs.dtype() { + return Err(MinitensorError::type_mismatch( + format!("{:?}", lhs.dtype()), + format!("{:?}", rhs.dtype()), + )); + } + if lhs.device() != rhs.device() { + return Err(MinitensorError::device_mismatch( + format!("{:?}", lhs.device()), + format!("{:?}", rhs.device()), + )); + } + if std::sync::Arc::strong_count(lhs.data()) > 1 { + // Fallback to out-of-place addition if data is shared + let tmp = add(lhs, rhs)?; + *lhs = tmp; + return Ok(()); + } + + match lhs.dtype() { + DataType::Float32 => { + let lhs_slice = lhs.data_mut().as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from lhs tensor") + })?; + let rhs_slice = rhs.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from rhs tensor") + })?; + binary_assign_slices(lhs_slice, rhs_slice, |l, r| l + r); + } + DataType::Float64 => { + let lhs_slice = lhs.data_mut().as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from lhs tensor") + })?; + let rhs_slice = rhs.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from rhs tensor") + })?; + binary_assign_slices(lhs_slice, rhs_slice, |l, r| l + r); + } + DataType::Int32 => { + let lhs_slice = lhs.data_mut().as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i32 slice from lhs tensor") + })?; + let rhs_slice = rhs.data().as_i32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i32 slice from rhs tensor") + })?; + binary_assign_slices(lhs_slice, rhs_slice, |l, r| l + r); + } + DataType::Int64 => { + let lhs_slice = lhs.data_mut().as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice from lhs tensor") + })?; + let rhs_slice = rhs.data().as_i64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i64 slice from rhs tensor") + })?; + binary_assign_slices(lhs_slice, rhs_slice, |l, r| l + r); + } + DataType::Bool => { + let lhs_slice = lhs.data_mut().as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice from lhs tensor") + })?; + let rhs_slice = rhs.data().as_bool_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get bool slice from rhs tensor") + })?; + binary_assign_slices(lhs_slice, rhs_slice, |l, r| l || r); + } + } + Ok(()) +} + +/// Apply `op` element-wise, writing the result into `lhs`. +/// +/// Safe replacement for the previous raw-pointer parallel loops: chunked +/// `rayon` iteration keeps bounds information visible to the compiler (so the +/// inner loops still vectorise) without any `unsafe`. +#[inline] +fn binary_assign_slices( + lhs: &mut [T], + rhs: &[T], + op: impl Fn(T, T) -> T + Send + Sync, +) { + debug_assert_eq!(lhs.len(), rhs.len()); + const CHUNK: usize = 4096; + if lhs.len() < PAR_THRESHOLD { + for (l, &r) in lhs.iter_mut().zip(rhs.iter()) { + *l = op(*l, r); + } + } else { + lhs.par_chunks_mut(CHUNK) + .zip(rhs.par_chunks(CHUNK)) + .for_each(|(lhs_chunk, rhs_chunk)| { + for (l, &r) in lhs_chunk.iter_mut().zip(rhs_chunk.iter()) { + *l = op(*l, r); + } + }); + } +} + +/// Element-wise subtraction with broadcasting support +pub fn sub(lhs: &Tensor, rhs: &Tensor) -> Result { + // Check device compatibility + if lhs.device() != rhs.device() { + return Err(MinitensorError::device_mismatch( + format!("{:?}", lhs.device()), + format!("{:?}", rhs.device()), + )); + } + + let requires_grad = lhs.requires_grad() || rhs.requires_grad(); + let (lhs_cast, rhs_cast, result_dtype) = coerce_binary_operands(lhs, rhs, BinaryOpKind::Sub)?; + let lhs_ref = lhs_cast.as_ref(); + let rhs_ref = rhs_cast.as_ref(); + + // Compute broadcasted shape + let output_shape = lhs_ref.shape().broadcast_with(rhs_ref.shape())?; + + if output_shape.numel() == 0 { + let mut output = Tensor::empty( + output_shape.clone(), + result_dtype, + lhs.device(), + requires_grad, + ); + + if requires_grad { + let grad_fn = Arc::new(SubBackward { + input_shapes: [lhs.shape().dims().to_vec(), rhs.shape().dims().to_vec()], + input_ids: [lhs.id(), rhs.id()], + input_requires_grad: [lhs.requires_grad(), rhs.requires_grad()], + }); + output.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output, Some(grad_fn))?; + } + + return Ok(output); + } + + // Create output tensor data + let mut output_data = + TensorData::uninitialized_on_device(output_shape.numel(), result_dtype, lhs.device()); + + // Perform element-wise subtraction based on data type + match result_dtype { + DataType::Float32 => sub_f32_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, + DataType::Float64 => sub_f64_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, + DataType::Int32 => sub_i32_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, + DataType::Int64 => sub_i64_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, + DataType::Bool => unreachable!("boolean subtraction should be rejected during coercion"), + } + + // Create output tensor + let mut output = Tensor::new( + Arc::new(output_data), + output_shape.clone(), + result_dtype, + lhs.device(), + requires_grad, + ); + + // Set up gradient function if needed + if requires_grad { + let grad_fn = Arc::new(SubBackward { + input_shapes: [lhs.shape().dims().to_vec(), rhs.shape().dims().to_vec()], + input_ids: [lhs.id(), rhs.id()], + input_requires_grad: [lhs.requires_grad(), rhs.requires_grad()], + }); + + output.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output, Some(grad_fn))?; + } + + Ok(output) +} + +/// Element-wise multiplication with broadcasting support +pub fn mul(lhs: &Tensor, rhs: &Tensor) -> Result { + // Check device compatibility + if lhs.device() != rhs.device() { + return Err(MinitensorError::device_mismatch( + format!("{:?}", lhs.device()), + format!("{:?}", rhs.device()), + )); + } + + let requires_grad = lhs.requires_grad() || rhs.requires_grad(); + let (lhs_cast, rhs_cast, result_dtype) = coerce_binary_operands(lhs, rhs, BinaryOpKind::Mul)?; + let lhs_ref = lhs_cast.as_ref(); + let rhs_ref = rhs_cast.as_ref(); + + // Compute broadcasted shape + let output_shape = lhs_ref.shape().broadcast_with(rhs_ref.shape())?; + + if output_shape.numel() == 0 { + let mut output = Tensor::empty( + output_shape.clone(), + result_dtype, + lhs.device(), + requires_grad, + ); + + if requires_grad { + let grad_fn = Arc::new(MulBackward { + lhs: lhs.clone(), + rhs: rhs.clone(), + input_ids: [lhs.id(), rhs.id()], + input_requires_grad: [lhs.requires_grad(), rhs.requires_grad()], + }); + output.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output, Some(grad_fn))?; + } + + return Ok(output); + } + + // Create output tensor data + let mut output_data = + TensorData::uninitialized_on_device(output_shape.numel(), result_dtype, lhs.device()); + + // Perform element-wise multiplication based on data type + match result_dtype { + DataType::Float32 => mul_f32_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, + DataType::Float64 => mul_f64_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, + DataType::Int32 => mul_i32_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, + DataType::Int64 => mul_i64_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, + DataType::Bool => mul_bool_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, + } + + // Create output tensor + let mut output = Tensor::new( + Arc::new(output_data), + output_shape.clone(), + result_dtype, + lhs.device(), + requires_grad, + ); + + // Set up gradient function if needed + if requires_grad { + let grad_fn = Arc::new(MulBackward { + lhs: lhs.clone(), + rhs: rhs.clone(), + input_ids: [lhs.id(), rhs.id()], + input_requires_grad: [lhs.requires_grad(), rhs.requires_grad()], + }); + + output.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output, Some(grad_fn))?; + } + + Ok(output) +} + +/// Element-wise division with broadcasting support +pub fn div(lhs: &Tensor, rhs: &Tensor) -> Result { + // Check device compatibility + if lhs.device() != rhs.device() { + return Err(MinitensorError::device_mismatch( + format!("{:?}", lhs.device()), + format!("{:?}", rhs.device()), + )); + } + + let requires_grad = lhs.requires_grad() || rhs.requires_grad(); + let (lhs_cast, rhs_cast, result_dtype) = coerce_binary_operands(lhs, rhs, BinaryOpKind::Div)?; + let lhs_ref = lhs_cast.as_ref(); + let rhs_ref = rhs_cast.as_ref(); + + // Compute broadcasted shape + let output_shape = lhs_ref.shape().broadcast_with(rhs_ref.shape())?; + + if output_shape.numel() == 0 { + let mut output = Tensor::empty( + output_shape.clone(), + result_dtype, + lhs.device(), + requires_grad, + ); + + if requires_grad { + let grad_fn = Arc::new(DivBackward { + lhs: lhs.clone(), + rhs: rhs.clone(), + input_ids: [lhs.id(), rhs.id()], + input_requires_grad: [lhs.requires_grad(), rhs.requires_grad()], + }); + output.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output, Some(grad_fn))?; + } + + return Ok(output); + } + + // Create output tensor data + let mut output_data = + TensorData::uninitialized_on_device(output_shape.numel(), result_dtype, lhs.device()); + + // Perform element-wise division based on data type + match result_dtype { + DataType::Float32 => div_f32_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, + DataType::Float64 => div_f64_direct(lhs_ref, rhs_ref, &mut output_data, &output_shape)?, + DataType::Int32 | DataType::Int64 | DataType::Bool => { + unreachable!("integer and boolean division should coerce to floating point") + } + } + + // Create output tensor + let mut output = Tensor::new( + Arc::new(output_data), + output_shape.clone(), + result_dtype, + lhs.device(), + requires_grad, + ); + + // Set up gradient function if needed + if requires_grad { + let grad_fn = Arc::new(DivBackward { + lhs: lhs.clone(), + rhs: rhs.clone(), + input_ids: [lhs.id(), rhs.id()], + input_requires_grad: [lhs.requires_grad(), rhs.requires_grad()], + }); + + output.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output, Some(grad_fn))?; + } + + Ok(output) +} + +/// Element-wise negation +pub fn neg(tensor: &Tensor) -> Result { + let mut output_data = TensorData::uninitialized_on_device( + tensor.shape().numel(), + tensor.dtype(), + tensor.device(), + ); + + /// Applies negation for one dtype: fetch the input/output slices and map + /// element-wise (parallel above `PAR_THRESHOLD`). + macro_rules! neg_arm { + ($accessor:ident, $accessor_mut:ident, $tyname:literal) => {{ + let input = tensor.data().$accessor().ok_or_else(|| { + MinitensorError::internal_error(concat!( + "Failed to get ", + $tyname, + " slice from tensor" + )) + })?; + let output = output_data.$accessor_mut().ok_or_else(|| { + MinitensorError::internal_error(concat!( + "Failed to get mutable ", + $tyname, + " slice from output" + )) + })?; + if input.len() >= PAR_THRESHOLD { + output + .par_iter_mut() + .zip(input.par_iter()) + .for_each(|(o, &i)| *o = -i); + } else { + for (o, &i) in output.iter_mut().zip(input.iter()) { + *o = -i; + } + } + }}; + } + + match tensor.dtype() { + DataType::Float32 => neg_arm!(as_f32_slice, as_f32_slice_mut, "f32"), + DataType::Float64 => neg_arm!(as_f64_slice, as_f64_slice_mut, "f64"), + DataType::Int32 => neg_arm!(as_i32_slice, as_i32_slice_mut, "i32"), + DataType::Int64 => neg_arm!(as_i64_slice, as_i64_slice_mut, "i64"), + DataType::Bool => { + return Err(MinitensorError::invalid_operation( + "Negation not supported for boolean tensors", + )); + } + } + + let output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if output.requires_grad() { + let grad_fn = Arc::new(NegBackward { + input_id: tensor.id(), + }); + let mut out_with_grad = output; + out_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&out_with_grad, Some(grad_fn))?; + Ok(out_with_grad) + } else { + Ok(output) + } +} + +// Helper functions for type-specific operations diff --git a/engine/src/operations/arithmetic/kernels.rs b/engine/src/operations/arithmetic/kernels.rs index af953e05..b35e9f7b 100644 --- a/engine/src/operations/arithmetic/kernels.rs +++ b/engine/src/operations/arithmetic/kernels.rs @@ -1,633 +1,684 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -fn mul_f64_direct( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from rhs tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - // Use SIMD fast path if possible (no broadcasting, same shapes) - if can_use_simd_fast_path(lhs.shape(), rhs.shape(), output_shape) { - simd_mul_f64(lhs_data, rhs_data, output_slice) - } else { - broadcast_binary_op( - lhs_data, - rhs_data, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - |a, b| a * b, - ) - } -} - -fn mul_i32_direct( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from rhs tensor") - })?; - - let output_slice = output_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i32 slice from output data") - })?; - - broadcast_binary_op( - lhs_data, - rhs_data, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - |a, b| a * b, - ) -} - -fn mul_i64_direct( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from rhs tensor") - })?; - - let output_slice = output_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice from output data") - })?; - - broadcast_binary_op( - lhs_data, - rhs_data, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - |a, b| a * b, - ) -} - -fn mul_bool_direct( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_bool_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_bool_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from rhs tensor") - })?; - - let output_slice = output_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice from output data") - })?; - - broadcast_binary_op( - lhs_data, - rhs_data, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - |a, b| a && b, - ) -} - -fn div_f32_direct( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from rhs tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - // Use SIMD fast path if possible (no broadcasting, same shapes) - if can_use_simd_fast_path(lhs.shape(), rhs.shape(), output_shape) { - simd_div_f32(lhs_data, rhs_data, output_slice) - } else { - broadcast_binary_op( - lhs_data, - rhs_data, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - |a, b| a / b, - ) - } -} - -fn div_f64_direct( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from rhs tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - // Use SIMD fast path if possible (no broadcasting, same shapes) - if can_use_simd_fast_path(lhs.shape(), rhs.shape(), output_shape) { - simd_div_f64(lhs_data, rhs_data, output_slice) - } else { - broadcast_binary_op( - lhs_data, - rhs_data, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - |a, b| a / b, - ) - } -} - -/// Generic broadcasting binary operation -pub(crate) fn broadcast_binary_op( - lhs_data: &[T], - rhs_data: &[T], - output_data: &mut [T], - lhs_shape: &Shape, - rhs_shape: &Shape, - output_shape: &Shape, - op: F, -) -> Result<()> -where - T: Copy + Send + Sync, - F: Fn(T, T) -> T + Send + Sync, -{ - let output_dims = output_shape.dims(); - let lhs_dims = lhs_shape.dims(); - let rhs_dims = rhs_shape.dims(); - let rank = output_dims.len(); - - if output_shape.numel() == 0 || output_dims.iter().any(|&dim| dim == 0) { - return Ok(()); - } - - // Fast path when no broadcasting is required. This avoids the - // relatively expensive index mapping logic below and simply applies the - // operation element-wise. We use parallel iteration for large tensors and - // fall back to a simple loop for smaller ones to reduce rayon overhead. - if lhs_dims == output_dims && rhs_dims == output_dims { - if output_data.len() >= 1024 { - output_data - .par_iter_mut() - .zip(lhs_data.par_iter().zip(rhs_data.par_iter())) - .for_each(|(out, (l, r))| { - *out = op(*l, *r); - }); - } else { - for ((out, &l), &r) in output_data - .iter_mut() - .zip(lhs_data.iter()) - .zip(rhs_data.iter()) - { - *out = op(l, r); - } - } - return Ok(()); - } - - // Fast path when one side is a scalar and the other already matches the - // output shape. This avoids the more expensive coordinate calculation - // used for general broadcasting. We again switch between parallel and - // sequential execution based on tensor size to minimize overhead. - if lhs_data.len() == 1 && rhs_dims == output_dims { - let lhs_val = lhs_data[0]; - if output_data.len() >= 1024 { - output_data - .par_iter_mut() - .zip(rhs_data.par_iter()) - .for_each(|(out, &r)| { - *out = op(lhs_val, r); - }); - } else { - for (out, &r) in output_data.iter_mut().zip(rhs_data.iter()) { - *out = op(lhs_val, r); - } - } - return Ok(()); - } - - if rhs_data.len() == 1 && lhs_dims == output_dims { - let rhs_val = rhs_data[0]; - if output_data.len() >= 1024 { - output_data - .par_iter_mut() - .zip(lhs_data.par_iter()) - .for_each(|(out, &l)| { - *out = op(l, rhs_val); - }); - } else { - for (out, &l) in output_data.iter_mut().zip(lhs_data.iter()) { - *out = op(l, rhs_val); - } - } - return Ok(()); - } - - let lhs_contiguous = Strides::from_shape(lhs_shape); - let rhs_contiguous = Strides::from_shape(rhs_shape); - let lhs_strides = lhs_contiguous.as_slice(); - let rhs_strides = rhs_contiguous.as_slice(); - - let mut lhs_aligned: SmallVec<[usize; 8]> = smallvec![0; rank]; - let mut rhs_aligned: SmallVec<[usize; 8]> = smallvec![0; rank]; - - let lhs_offset = rank.saturating_sub(lhs_dims.len()); - for (i, &dim) in lhs_dims.iter().enumerate() { - lhs_aligned[lhs_offset + i] = if dim == 1 { 0 } else { lhs_strides[i] }; - } - - let rhs_offset = rank.saturating_sub(rhs_dims.len()); - for (i, &dim) in rhs_dims.iter().enumerate() { - rhs_aligned[rhs_offset + i] = if dim == 1 { 0 } else { rhs_strides[i] }; - } - - // For small tensors, a simple sequential loop is faster than spawning - // rayon tasks. We use the same index mapping logic but without parallel - // chunking to minimize overhead. - if output_data.len() < 1024 { - let lhs_ptr = lhs_data.as_ptr(); - let rhs_ptr = rhs_data.as_ptr(); - for (idx, out) in output_data.iter_mut().enumerate() { - let mut lhs_idx = 0usize; - let mut rhs_idx = 0usize; - let mut tmp = idx; - for i in (0..rank).rev() { - let coord = tmp % output_dims[i]; - tmp /= output_dims[i]; - lhs_idx += coord * lhs_aligned[i]; - rhs_idx += coord * rhs_aligned[i]; - } - unsafe { - *out = op(*lhs_ptr.add(lhs_idx), *rhs_ptr.add(rhs_idx)); - } - } - return Ok(()); - } - - const CHUNK: usize = 1024; - output_data - .par_chunks_mut(CHUNK) - .enumerate() - .for_each(|(chunk_idx, out_chunk)| { - let start = chunk_idx * CHUNK; - let mut coord: SmallVec<[usize; 8]> = smallvec![0; rank]; - let mut tmp = start; - for i in (0..rank).rev() { - coord[i] = tmp % output_dims[i]; - tmp /= output_dims[i]; - } - - let mut lhs_idx = 0usize; - let mut rhs_idx = 0usize; - for i in 0..rank { - lhs_idx += coord[i] * lhs_aligned[i]; - rhs_idx += coord[i] * rhs_aligned[i]; - } - - let lhs_ptr = lhs_data.as_ptr(); - let rhs_ptr = rhs_data.as_ptr(); - for out in out_chunk.iter_mut() { - unsafe { - *out = op(*lhs_ptr.add(lhs_idx), *rhs_ptr.add(rhs_idx)); - } - for i in (0..rank).rev() { - coord[i] += 1; - lhs_idx += lhs_aligned[i]; - rhs_idx += rhs_aligned[i]; - if coord[i] < output_dims[i] { - break; - } - coord[i] = 0; - lhs_idx -= lhs_aligned[i] * output_dims[i]; - rhs_idx -= rhs_aligned[i] * output_dims[i]; - } - } - }); - - Ok(()) -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::device::Device; - - fn create_test_tensor_f32(data: Vec, shape: Vec, requires_grad: bool) -> Tensor { - let shape_obj = Shape::new(shape); - let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Float32); - - if let Some(slice) = tensor_data.as_f32_slice_mut() { - slice.copy_from_slice(&data); - } - - Tensor::new( - Arc::new(tensor_data), - shape_obj, - DataType::Float32, - Device::cpu(), - requires_grad, - ) - } - - #[test] - fn test_add_basic() { - let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], false); - let b = create_test_tensor_f32(vec![4.0, 5.0, 6.0], vec![3], false); - - let result = add(&a, &b).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - - assert_eq!(result_data, &[5.0, 7.0, 9.0]); - assert_eq!(result.shape().dims(), &[3]); - } - - #[test] - fn test_add_broadcasting() { - let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], false); - let b = create_test_tensor_f32(vec![10.0], vec![1], false); - - let result = add(&a, &b).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - - assert_eq!(result_data, &[11.0, 12.0, 13.0]); - assert_eq!(result.shape().dims(), &[3]); - } - - #[test] - fn test_sub_basic() { - let a = create_test_tensor_f32(vec![5.0, 7.0, 9.0], vec![3], false); - let b = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], false); - - let result = sub(&a, &b).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - - assert_eq!(result_data, &[4.0, 5.0, 6.0]); - } - - #[test] - fn test_mul_basic() { - let a = create_test_tensor_f32(vec![2.0, 3.0, 4.0], vec![3], false); - let b = create_test_tensor_f32(vec![5.0, 6.0, 7.0], vec![3], false); - - let result = mul(&a, &b).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - - assert_eq!(result_data, &[10.0, 18.0, 28.0]); - } - - #[test] - fn test_div_basic() { - let a = create_test_tensor_f32(vec![10.0, 15.0, 20.0], vec![3], false); - let b = create_test_tensor_f32(vec![2.0, 3.0, 4.0], vec![3], false); - - let result = div(&a, &b).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - - assert_eq!(result_data, &[5.0, 5.0, 5.0]); - } - - #[test] - fn test_neg_basic() { - let a = create_test_tensor_f32(vec![1.0, -2.0, 3.5], vec![3], false); - let result = neg(&a).unwrap(); - let data = result.data().as_f32_slice().unwrap(); - assert_eq!(data, &[-1.0, 2.0, -3.5]); - } - - #[test] - fn test_gradient_tracking() { - let a = create_test_tensor_f32(vec![1.0, 2.0], vec![2], true); - let b = create_test_tensor_f32(vec![3.0, 4.0], vec![2], true); - - let result = add(&a, &b).unwrap(); - - assert!(result.requires_grad()); - assert!(result.grad_fn().is_some()); - } - - #[test] - fn test_device_mismatch_error() { - let a = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - let b = create_test_tensor_f32(vec![3.0, 4.0], vec![2], false); - - // This would normally fail, but we can't easily create different device tensors in tests - // So we'll just test that same device works - let result = add(&a, &b); - assert!(result.is_ok()); - } - - #[test] - fn test_mixed_dtype_promotion() { - let a = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - - // Create an i32 tensor - let shape_obj = Shape::new(vec![2]); - let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Int32); - if let Some(slice) = tensor_data.as_i32_slice_mut() { - slice.copy_from_slice(&[3, 4]); - } - let b = Tensor::new( - Arc::new(tensor_data), - shape_obj, - DataType::Int32, - Device::cpu(), - false, - ); - - let result = add(&a, &b).unwrap(); - assert_eq!(result.dtype(), DataType::Float32); - assert_eq!(result.data().as_f32_slice().unwrap(), &[4.0, 6.0]); - } - - #[test] - fn test_sub_broadcasting_2d() { - let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); - let b = create_test_tensor_f32(vec![1.0, 2.0], vec![1, 2], false); - let result = sub(&a, &b).unwrap(); - let expected = vec![0.0, 0.0, 2.0, 2.0]; - assert_eq!(result.data().as_f32_slice().unwrap(), expected.as_slice()); - assert_eq!(result.shape().dims(), &[2, 2]); - } - - #[test] - fn test_mul_broadcasting_2d() { - let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); - let b = create_test_tensor_f32(vec![2.0], vec![1, 1], false); - let result = mul(&a, &b).unwrap(); - assert_eq!(result.data().as_f32_slice().unwrap(), &[2.0, 4.0, 6.0, 8.0]); - assert_eq!(result.shape().dims(), &[2, 2]); - } - - #[test] - fn test_div_broadcasting_2d() { - let a = create_test_tensor_f32(vec![2.0, 4.0, 6.0, 8.0], vec![2, 2], false); - let b = create_test_tensor_f32(vec![2.0], vec![1, 1], false); - let result = div(&a, &b).unwrap(); - assert_eq!(result.data().as_f32_slice().unwrap(), &[1.0, 2.0, 3.0, 4.0]); - assert_eq!(result.shape().dims(), &[2, 2]); - } - - #[test] - fn test_bool_arithmetic_behaviour() { - // Create boolean tensors - let shape_obj = Shape::new(vec![2]); - let mut data_a = TensorData::zeros(shape_obj.numel(), DataType::Bool); - if let Some(slice) = data_a.as_bool_slice_mut() { - slice.copy_from_slice(&[true, false]); - } - let a = Tensor::new( - Arc::new(data_a), - shape_obj.clone(), - DataType::Bool, - Device::cpu(), - false, - ); - - let mut data_b = TensorData::zeros(shape_obj.numel(), DataType::Bool); - if let Some(slice) = data_b.as_bool_slice_mut() { - slice.copy_from_slice(&[false, true]); - } - let b = Tensor::new( - Arc::new(data_b), - shape_obj, - DataType::Bool, - Device::cpu(), - false, - ); - - let add_result = add(&a, &b).unwrap(); - assert_eq!(add_result.dtype(), DataType::Bool); - assert_eq!(add_result.data().as_bool_slice().unwrap(), &[true, true]); - assert!(sub(&a, &b).is_err()); - let mul_result = mul(&a, &b).unwrap(); - assert_eq!(mul_result.dtype(), DataType::Bool); - assert_eq!(mul_result.data().as_bool_slice().unwrap(), &[false, false]); - let div_result = div(&a, &b).unwrap(); - assert_eq!(div_result.dtype(), DataType::Float32); - assert_eq!( - div_result.data().as_f32_slice().unwrap(), - &[f32::INFINITY, 0.0] - ); - assert!(neg(&a).is_err()); - } - - #[test] - fn test_incompatible_shapes_error() { - let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], false); - let b = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - assert!(sub(&a, &b).is_err()); - assert!(mul(&a, &b).is_err()); - assert!(div(&a, &b).is_err()); - } - - #[test] - fn test_division_by_zero_returns_inf() { - let a = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - let b = create_test_tensor_f32(vec![0.0, 1.0], vec![2], false); - let result = div(&a, &b).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - assert!(result_data[0].is_infinite()); - assert_eq!(result_data[1], 2.0); - } - - #[test] - fn test_division_by_zero_ieee_semantics_broadcast_path() { - // Broadcasting a scalar zero divisor exercises the non-SIMD path, - // which must match IEEE 754 (and the SIMD path): -1/0 = -inf, - // 0/0 = NaN, 1/0 = inf. - let a = create_test_tensor_f32(vec![-1.0, 0.0, 1.0], vec![3], false); - let b = create_test_tensor_f32(vec![0.0], vec![1], false); - let result = div(&a, &b).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - assert_eq!(result_data[0], f32::NEG_INFINITY); - assert!(result_data[1].is_nan()); - assert_eq!(result_data[2], f32::INFINITY); - } - - #[test] - fn test_add_handles_zero_sized_broadcast() { - let a = create_test_tensor_f32(vec![], vec![0, 3], false); - let b = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], false); - - let result = add(&a, &b).unwrap(); - assert_eq!(result.shape().dims(), &[0, 3]); - assert_eq!(result.data().as_f32_slice().unwrap().len(), 0); - } - - #[test] - fn test_add_handles_zero_sized_broadcast_from_vec() { - use crate::tensor::TensorData; - - let a_data = TensorData::from_vec::(vec![], DataType::Float32, Device::cpu()); - let a = Tensor::new( - Arc::new(a_data), - Shape::new(vec![0, 3]), - DataType::Float32, - Device::cpu(), - false, - ); - - let b_data = - TensorData::from_vec::(vec![1.0_f32, 2.0, 3.0], DataType::Float32, Device::cpu()); - let b = Tensor::new( - Arc::new(b_data), - Shape::new(vec![3]), - DataType::Float32, - Device::cpu(), - false, - ); - - let result = add(&a, &b).unwrap(); - assert_eq!(result.shape().dims(), &[0, 3]); - assert_eq!(result.data().as_f32_slice().unwrap().len(), 0); - } -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use crate::operations::simd::*; +use crate::{ + error::{MinitensorError, Result}, + tensor::{Shape, Strides, Tensor, TensorData}, +}; +use rayon::prelude::*; +use smallvec::SmallVec; +use smallvec::smallvec; + +/// Generates a dtype-specialized broadcasting binary kernel. +/// +/// Every kernel has the same shape: fetch both input slices and the output +/// slice for the dtype, then apply `$op` element-wise with broadcasting. +macro_rules! binary_kernel { + ($name:ident, $accessor:ident, $accessor_mut:ident, $tyname:literal, $op:expr) => { + pub(crate) fn $name( + lhs: &Tensor, + rhs: &Tensor, + output_data: &mut TensorData, + output_shape: &Shape, + ) -> Result<()> { + let lhs_data = lhs.data().$accessor().ok_or_else(|| { + MinitensorError::internal_error(concat!( + "Failed to get ", + $tyname, + " slice from lhs tensor" + )) + })?; + let rhs_data = rhs.data().$accessor().ok_or_else(|| { + MinitensorError::internal_error(concat!( + "Failed to get ", + $tyname, + " slice from rhs tensor" + )) + })?; + let output_slice = output_data.$accessor_mut().ok_or_else(|| { + MinitensorError::internal_error(concat!( + "Failed to get mutable ", + $tyname, + " slice from output data" + )) + })?; + broadcast_binary_op( + lhs_data, + rhs_data, + output_slice, + lhs.shape(), + rhs.shape(), + output_shape, + $op, + ) + } + }; +} + +/// Same as [`binary_kernel!`], with a SIMD fast path for same-shape inputs +/// (f32/f64 only). +macro_rules! binary_kernel_simd { + ($name:ident, $accessor:ident, $accessor_mut:ident, $tyname:literal, $simd:ident, $op:expr) => { + pub(crate) fn $name( + lhs: &Tensor, + rhs: &Tensor, + output_data: &mut TensorData, + output_shape: &Shape, + ) -> Result<()> { + let lhs_data = lhs.data().$accessor().ok_or_else(|| { + MinitensorError::internal_error(concat!( + "Failed to get ", + $tyname, + " slice from lhs tensor" + )) + })?; + let rhs_data = rhs.data().$accessor().ok_or_else(|| { + MinitensorError::internal_error(concat!( + "Failed to get ", + $tyname, + " slice from rhs tensor" + )) + })?; + let output_slice = output_data.$accessor_mut().ok_or_else(|| { + MinitensorError::internal_error(concat!( + "Failed to get mutable ", + $tyname, + " slice from output data" + )) + })?; + if can_use_simd_fast_path(lhs.shape(), rhs.shape(), output_shape) { + $simd(lhs_data, rhs_data, output_slice) + } else { + broadcast_binary_op( + lhs_data, + rhs_data, + output_slice, + lhs.shape(), + rhs.shape(), + output_shape, + $op, + ) + } + } + }; +} + +// Addition: `+` for numeric dtypes, logical OR for bool. +binary_kernel_simd!( + add_f32_direct, + as_f32_slice, + as_f32_slice_mut, + "f32", + simd_add_f32, + |a, b| a + b +); +binary_kernel_simd!( + add_f64_direct, + as_f64_slice, + as_f64_slice_mut, + "f64", + simd_add_f64, + |a, b| a + b +); +binary_kernel!( + add_i32_direct, + as_i32_slice, + as_i32_slice_mut, + "i32", + |a, b| a + b +); +binary_kernel!( + add_i64_direct, + as_i64_slice, + as_i64_slice_mut, + "i64", + |a, b| a + b +); +binary_kernel!( + add_bool_direct, + as_bool_slice, + as_bool_slice_mut, + "bool", + |a, b| a || b +); + +// Subtraction: bool is rejected during operand coercion. +binary_kernel_simd!( + sub_f32_direct, + as_f32_slice, + as_f32_slice_mut, + "f32", + simd_sub_f32, + |a, b| a - b +); +binary_kernel_simd!( + sub_f64_direct, + as_f64_slice, + as_f64_slice_mut, + "f64", + simd_sub_f64, + |a, b| a - b +); +binary_kernel!( + sub_i32_direct, + as_i32_slice, + as_i32_slice_mut, + "i32", + |a, b| a - b +); +binary_kernel!( + sub_i64_direct, + as_i64_slice, + as_i64_slice_mut, + "i64", + |a, b| a - b +); + +// Multiplication: `*` for numeric dtypes, logical AND for bool. +binary_kernel_simd!( + mul_f32_direct, + as_f32_slice, + as_f32_slice_mut, + "f32", + simd_mul_f32, + |a, b| a * b +); +binary_kernel_simd!( + mul_f64_direct, + as_f64_slice, + as_f64_slice_mut, + "f64", + simd_mul_f64, + |a, b| a * b +); +binary_kernel!( + mul_i32_direct, + as_i32_slice, + as_i32_slice_mut, + "i32", + |a, b| a * b +); +binary_kernel!( + mul_i64_direct, + as_i64_slice, + as_i64_slice_mut, + "i64", + |a, b| a * b +); +binary_kernel!( + mul_bool_direct, + as_bool_slice, + as_bool_slice_mut, + "bool", + |a, b| a && b +); + +// Division: integer and bool operands coerce to floating point beforehand. +binary_kernel_simd!( + div_f32_direct, + as_f32_slice, + as_f32_slice_mut, + "f32", + simd_div_f32, + |a, b| a / b +); +binary_kernel_simd!( + div_f64_direct, + as_f64_slice, + as_f64_slice_mut, + "f64", + simd_div_f64, + |a, b| a / b +); + +/// Generic broadcasting binary operation +pub(crate) fn broadcast_binary_op( + lhs_data: &[T], + rhs_data: &[T], + output_data: &mut [T], + lhs_shape: &Shape, + rhs_shape: &Shape, + output_shape: &Shape, + op: F, +) -> Result<()> +where + T: Copy + Send + Sync, + F: Fn(T, T) -> T + Send + Sync, +{ + let output_dims = output_shape.dims(); + let lhs_dims = lhs_shape.dims(); + let rhs_dims = rhs_shape.dims(); + let rank = output_dims.len(); + + if output_shape.numel() == 0 || output_dims.contains(&0) { + return Ok(()); + } + + // Fast path when no broadcasting is required. This avoids the + // relatively expensive index mapping logic below and simply applies the + // operation element-wise. We use parallel iteration for large tensors and + // fall back to a simple loop for smaller ones to reduce rayon overhead. + if lhs_dims == output_dims && rhs_dims == output_dims { + if output_data.len() >= 1024 { + output_data + .par_iter_mut() + .zip(lhs_data.par_iter().zip(rhs_data.par_iter())) + .for_each(|(out, (l, r))| { + *out = op(*l, *r); + }); + } else { + for ((out, &l), &r) in output_data + .iter_mut() + .zip(lhs_data.iter()) + .zip(rhs_data.iter()) + { + *out = op(l, r); + } + } + return Ok(()); + } + + // Fast path when one side is a scalar and the other already matches the + // output shape. This avoids the more expensive coordinate calculation + // used for general broadcasting. We again switch between parallel and + // sequential execution based on tensor size to minimize overhead. + if lhs_data.len() == 1 && rhs_dims == output_dims { + let lhs_val = lhs_data[0]; + if output_data.len() >= 1024 { + output_data + .par_iter_mut() + .zip(rhs_data.par_iter()) + .for_each(|(out, &r)| { + *out = op(lhs_val, r); + }); + } else { + for (out, &r) in output_data.iter_mut().zip(rhs_data.iter()) { + *out = op(lhs_val, r); + } + } + return Ok(()); + } + + if rhs_data.len() == 1 && lhs_dims == output_dims { + let rhs_val = rhs_data[0]; + if output_data.len() >= 1024 { + output_data + .par_iter_mut() + .zip(lhs_data.par_iter()) + .for_each(|(out, &l)| { + *out = op(l, rhs_val); + }); + } else { + for (out, &l) in output_data.iter_mut().zip(lhs_data.iter()) { + *out = op(l, rhs_val); + } + } + return Ok(()); + } + + let lhs_contiguous = Strides::from_shape(lhs_shape); + let rhs_contiguous = Strides::from_shape(rhs_shape); + let lhs_strides = lhs_contiguous.as_slice(); + let rhs_strides = rhs_contiguous.as_slice(); + + let mut lhs_aligned: SmallVec<[usize; 8]> = smallvec![0; rank]; + let mut rhs_aligned: SmallVec<[usize; 8]> = smallvec![0; rank]; + + let lhs_offset = rank.saturating_sub(lhs_dims.len()); + for (i, &dim) in lhs_dims.iter().enumerate() { + lhs_aligned[lhs_offset + i] = if dim == 1 { 0 } else { lhs_strides[i] }; + } + + let rhs_offset = rank.saturating_sub(rhs_dims.len()); + for (i, &dim) in rhs_dims.iter().enumerate() { + rhs_aligned[rhs_offset + i] = if dim == 1 { 0 } else { rhs_strides[i] }; + } + + // For small tensors, a simple sequential loop is faster than spawning + // rayon tasks. We use the same index mapping logic but without parallel + // chunking to minimize overhead. + if output_data.len() < 1024 { + let lhs_ptr = lhs_data.as_ptr(); + let rhs_ptr = rhs_data.as_ptr(); + for (idx, out) in output_data.iter_mut().enumerate() { + let mut lhs_idx = 0usize; + let mut rhs_idx = 0usize; + let mut tmp = idx; + for i in (0..rank).rev() { + let coord = tmp % output_dims[i]; + tmp /= output_dims[i]; + lhs_idx += coord * lhs_aligned[i]; + rhs_idx += coord * rhs_aligned[i]; + } + unsafe { + *out = op(*lhs_ptr.add(lhs_idx), *rhs_ptr.add(rhs_idx)); + } + } + return Ok(()); + } + + const CHUNK: usize = 1024; + output_data + .par_chunks_mut(CHUNK) + .enumerate() + .for_each(|(chunk_idx, out_chunk)| { + let start = chunk_idx * CHUNK; + let mut coord: SmallVec<[usize; 8]> = smallvec![0; rank]; + let mut tmp = start; + for i in (0..rank).rev() { + coord[i] = tmp % output_dims[i]; + tmp /= output_dims[i]; + } + + let mut lhs_idx = 0usize; + let mut rhs_idx = 0usize; + for i in 0..rank { + lhs_idx += coord[i] * lhs_aligned[i]; + rhs_idx += coord[i] * rhs_aligned[i]; + } + + let lhs_ptr = lhs_data.as_ptr(); + let rhs_ptr = rhs_data.as_ptr(); + for out in out_chunk.iter_mut() { + unsafe { + *out = op(*lhs_ptr.add(lhs_idx), *rhs_ptr.add(rhs_idx)); + } + for i in (0..rank).rev() { + coord[i] += 1; + lhs_idx += lhs_aligned[i]; + rhs_idx += rhs_aligned[i]; + if coord[i] < output_dims[i] { + break; + } + coord[i] = 0; + lhs_idx -= lhs_aligned[i] * output_dims[i]; + rhs_idx -= rhs_aligned[i] * output_dims[i]; + } + } + }); + + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::device::Device; + use crate::operations::arithmetic::{add, div, mul, neg, sub}; + use crate::tensor::DataType; + use std::sync::Arc; + + fn create_test_tensor_f32(data: Vec, shape: Vec, requires_grad: bool) -> Tensor { + let shape_obj = Shape::new(shape); + let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Float32); + + if let Some(slice) = tensor_data.as_f32_slice_mut() { + slice.copy_from_slice(&data); + } + + Tensor::new( + Arc::new(tensor_data), + shape_obj, + DataType::Float32, + Device::cpu(), + requires_grad, + ) + } + + #[test] + fn test_add_basic() { + let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], false); + let b = create_test_tensor_f32(vec![4.0, 5.0, 6.0], vec![3], false); + + let result = add(&a, &b).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + + assert_eq!(result_data, &[5.0, 7.0, 9.0]); + assert_eq!(result.shape().dims(), &[3]); + } + + #[test] + fn test_add_broadcasting() { + let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], false); + let b = create_test_tensor_f32(vec![10.0], vec![1], false); + + let result = add(&a, &b).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + + assert_eq!(result_data, &[11.0, 12.0, 13.0]); + assert_eq!(result.shape().dims(), &[3]); + } + + #[test] + fn test_sub_basic() { + let a = create_test_tensor_f32(vec![5.0, 7.0, 9.0], vec![3], false); + let b = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], false); + + let result = sub(&a, &b).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + + assert_eq!(result_data, &[4.0, 5.0, 6.0]); + } + + #[test] + fn test_mul_basic() { + let a = create_test_tensor_f32(vec![2.0, 3.0, 4.0], vec![3], false); + let b = create_test_tensor_f32(vec![5.0, 6.0, 7.0], vec![3], false); + + let result = mul(&a, &b).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + + assert_eq!(result_data, &[10.0, 18.0, 28.0]); + } + + #[test] + fn test_div_basic() { + let a = create_test_tensor_f32(vec![10.0, 15.0, 20.0], vec![3], false); + let b = create_test_tensor_f32(vec![2.0, 3.0, 4.0], vec![3], false); + + let result = div(&a, &b).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + + assert_eq!(result_data, &[5.0, 5.0, 5.0]); + } + + #[test] + fn test_neg_basic() { + let a = create_test_tensor_f32(vec![1.0, -2.0, 3.5], vec![3], false); + let result = neg(&a).unwrap(); + let data = result.data().as_f32_slice().unwrap(); + assert_eq!(data, &[-1.0, 2.0, -3.5]); + } + + #[test] + fn test_gradient_tracking() { + let a = create_test_tensor_f32(vec![1.0, 2.0], vec![2], true); + let b = create_test_tensor_f32(vec![3.0, 4.0], vec![2], true); + + let result = add(&a, &b).unwrap(); + + assert!(result.requires_grad()); + assert!(result.grad_fn().is_some()); + } + + #[test] + fn test_device_mismatch_error() { + let a = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + let b = create_test_tensor_f32(vec![3.0, 4.0], vec![2], false); + + // This would normally fail, but we can't easily create different device tensors in tests + // So we'll just test that same device works + let result = add(&a, &b); + assert!(result.is_ok()); + } + + #[test] + fn test_mixed_dtype_promotion() { + let a = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + + // Create an i32 tensor + let shape_obj = Shape::new(vec![2]); + let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Int32); + if let Some(slice) = tensor_data.as_i32_slice_mut() { + slice.copy_from_slice(&[3, 4]); + } + let b = Tensor::new( + Arc::new(tensor_data), + shape_obj, + DataType::Int32, + Device::cpu(), + false, + ); + + let result = add(&a, &b).unwrap(); + assert_eq!(result.dtype(), DataType::Float32); + assert_eq!(result.data().as_f32_slice().unwrap(), &[4.0, 6.0]); + } + + #[test] + fn test_sub_broadcasting_2d() { + let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); + let b = create_test_tensor_f32(vec![1.0, 2.0], vec![1, 2], false); + let result = sub(&a, &b).unwrap(); + let expected = vec![0.0, 0.0, 2.0, 2.0]; + assert_eq!(result.data().as_f32_slice().unwrap(), expected.as_slice()); + assert_eq!(result.shape().dims(), &[2, 2]); + } + + #[test] + fn test_mul_broadcasting_2d() { + let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); + let b = create_test_tensor_f32(vec![2.0], vec![1, 1], false); + let result = mul(&a, &b).unwrap(); + assert_eq!(result.data().as_f32_slice().unwrap(), &[2.0, 4.0, 6.0, 8.0]); + assert_eq!(result.shape().dims(), &[2, 2]); + } + + #[test] + fn test_div_broadcasting_2d() { + let a = create_test_tensor_f32(vec![2.0, 4.0, 6.0, 8.0], vec![2, 2], false); + let b = create_test_tensor_f32(vec![2.0], vec![1, 1], false); + let result = div(&a, &b).unwrap(); + assert_eq!(result.data().as_f32_slice().unwrap(), &[1.0, 2.0, 3.0, 4.0]); + assert_eq!(result.shape().dims(), &[2, 2]); + } + + #[test] + fn test_bool_arithmetic_behaviour() { + // Create boolean tensors + let shape_obj = Shape::new(vec![2]); + let mut data_a = TensorData::zeros(shape_obj.numel(), DataType::Bool); + if let Some(slice) = data_a.as_bool_slice_mut() { + slice.copy_from_slice(&[true, false]); + } + let a = Tensor::new( + Arc::new(data_a), + shape_obj.clone(), + DataType::Bool, + Device::cpu(), + false, + ); + + let mut data_b = TensorData::zeros(shape_obj.numel(), DataType::Bool); + if let Some(slice) = data_b.as_bool_slice_mut() { + slice.copy_from_slice(&[false, true]); + } + let b = Tensor::new( + Arc::new(data_b), + shape_obj, + DataType::Bool, + Device::cpu(), + false, + ); + + let add_result = add(&a, &b).unwrap(); + assert_eq!(add_result.dtype(), DataType::Bool); + assert_eq!(add_result.data().as_bool_slice().unwrap(), &[true, true]); + assert!(sub(&a, &b).is_err()); + let mul_result = mul(&a, &b).unwrap(); + assert_eq!(mul_result.dtype(), DataType::Bool); + assert_eq!(mul_result.data().as_bool_slice().unwrap(), &[false, false]); + let div_result = div(&a, &b).unwrap(); + assert_eq!(div_result.dtype(), DataType::Float32); + assert_eq!( + div_result.data().as_f32_slice().unwrap(), + &[f32::INFINITY, 0.0] + ); + assert!(neg(&a).is_err()); + } + + #[test] + fn test_incompatible_shapes_error() { + let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], false); + let b = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + assert!(sub(&a, &b).is_err()); + assert!(mul(&a, &b).is_err()); + assert!(div(&a, &b).is_err()); + } + + #[test] + fn test_division_by_zero_returns_inf() { + let a = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + let b = create_test_tensor_f32(vec![0.0, 1.0], vec![2], false); + let result = div(&a, &b).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + assert!(result_data[0].is_infinite()); + assert_eq!(result_data[1], 2.0); + } + + #[test] + fn test_division_by_zero_ieee_semantics_broadcast_path() { + // Broadcasting a scalar zero divisor exercises the non-SIMD path, + // which must match IEEE 754 (and the SIMD path): -1/0 = -inf, + // 0/0 = NaN, 1/0 = inf. + let a = create_test_tensor_f32(vec![-1.0, 0.0, 1.0], vec![3], false); + let b = create_test_tensor_f32(vec![0.0], vec![1], false); + let result = div(&a, &b).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + assert_eq!(result_data[0], f32::NEG_INFINITY); + assert!(result_data[1].is_nan()); + assert_eq!(result_data[2], f32::INFINITY); + } + + #[test] + fn test_add_handles_zero_sized_broadcast() { + let a = create_test_tensor_f32(vec![], vec![0, 3], false); + let b = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], false); + + let result = add(&a, &b).unwrap(); + assert_eq!(result.shape().dims(), &[0, 3]); + assert_eq!(result.data().as_f32_slice().unwrap().len(), 0); + } + + #[test] + fn test_add_handles_zero_sized_broadcast_from_vec() { + use crate::tensor::TensorData; + + let a_data = TensorData::from_vec::(vec![], DataType::Float32, Device::cpu()); + let a = Tensor::new( + Arc::new(a_data), + Shape::new(vec![0, 3]), + DataType::Float32, + Device::cpu(), + false, + ); + + let b_data = + TensorData::from_vec::(vec![1.0_f32, 2.0, 3.0], DataType::Float32, Device::cpu()); + let b = Tensor::new( + Arc::new(b_data), + Shape::new(vec![3]), + DataType::Float32, + Device::cpu(), + false, + ); + + let result = add(&a, &b).unwrap(); + assert_eq!(result.shape().dims(), &[0, 3]); + assert_eq!(result.data().as_f32_slice().unwrap().len(), 0); + } +} diff --git a/engine/src/operations/comparison.rs b/engine/src/operations/comparison.rs index 98550281..dba3bf46 100644 --- a/engine/src/operations/comparison.rs +++ b/engine/src/operations/comparison.rs @@ -141,140 +141,53 @@ where Ok(()) } -fn cmp_f32( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, - op: impl Fn(f32, f32) -> bool + Sync + Send, -) -> Result<()> { - let lhs_slice = lhs.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from lhs tensor") - })?; - let rhs_slice = rhs.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from rhs tensor") - })?; - let output_slice = output_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from output data") - })?; - broadcast_compare_op( - lhs_slice, - rhs_slice, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - op, - ) -} - -fn cmp_f64( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, - op: impl Fn(f64, f64) -> bool + Sync + Send, -) -> Result<()> { - let lhs_slice = lhs.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from lhs tensor") - })?; - let rhs_slice = rhs.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from rhs tensor") - })?; - let output_slice = output_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from output data") - })?; - broadcast_compare_op( - lhs_slice, - rhs_slice, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - op, - ) -} - -fn cmp_i32( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, - op: impl Fn(i32, i32) -> bool + Sync + Send, -) -> Result<()> { - let lhs_slice = lhs.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from lhs tensor") - })?; - let rhs_slice = rhs.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from rhs tensor") - })?; - let output_slice = output_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from output data") - })?; - broadcast_compare_op( - lhs_slice, - rhs_slice, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - op, - ) -} - -fn cmp_i64( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, - op: impl Fn(i64, i64) -> bool + Sync + Send, -) -> Result<()> { - let lhs_slice = lhs.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from lhs tensor") - })?; - let rhs_slice = rhs.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from rhs tensor") - })?; - let output_slice = output_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from output data") - })?; - broadcast_compare_op( - lhs_slice, - rhs_slice, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - op, - ) +/// Generates a dtype-specialized comparison kernel: fetch both input slices +/// for the dtype, fetch the bool output slice, and apply `op` element-wise +/// with broadcasting. +macro_rules! cmp_kernel { + ($name:ident, $ty:ty, $accessor:ident, $tyname:literal) => { + fn $name( + lhs: &Tensor, + rhs: &Tensor, + output_data: &mut TensorData, + output_shape: &Shape, + op: impl Fn($ty, $ty) -> bool + Sync + Send, + ) -> Result<()> { + let lhs_slice = lhs.data().$accessor().ok_or_else(|| { + MinitensorError::internal_error(concat!( + "Failed to get ", + $tyname, + " slice from lhs tensor" + )) + })?; + let rhs_slice = rhs.data().$accessor().ok_or_else(|| { + MinitensorError::internal_error(concat!( + "Failed to get ", + $tyname, + " slice from rhs tensor" + )) + })?; + let output_slice = output_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get bool slice from output data") + })?; + broadcast_compare_op( + lhs_slice, + rhs_slice, + output_slice, + lhs.shape(), + rhs.shape(), + output_shape, + op, + ) + } + }; } -fn cmp_bool( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, - op: impl Fn(bool, bool) -> bool + Sync + Send, -) -> Result<()> { - let lhs_slice = lhs.data().as_bool_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from lhs tensor") - })?; - let rhs_slice = rhs.data().as_bool_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from rhs tensor") - })?; - let output_slice = output_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from output data") - })?; - broadcast_compare_op( - lhs_slice, - rhs_slice, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - op, - ) -} +cmp_kernel!(cmp_f32, f32, as_f32_slice, "f32"); +cmp_kernel!(cmp_f64, f64, as_f64_slice, "f64"); +cmp_kernel!(cmp_i32, i32, as_i32_slice, "i32"); +cmp_kernel!(cmp_i64, i64, as_i64_slice, "i64"); +cmp_kernel!(cmp_bool, bool, as_bool_slice, "bool"); macro_rules! cmp_op { ($fn_name:ident, $op:tt, $bool_ok:expr) => { diff --git a/engine/src/operations/conv.rs b/engine/src/operations/conv.rs index 768656a5..ab8617b5 100644 --- a/engine/src/operations/conv.rs +++ b/engine/src/operations/conv.rs @@ -57,13 +57,13 @@ pub fn conv2d( )); } - if let Some(b) = bias { - if b.ndim() != 1 || b.size(0)? != out_channels { - return Err(MinitensorError::shape_mismatch( - vec![out_channels], - vec![b.size(0)?], - )); - } + if let Some(b) = bias + && (b.ndim() != 1 || b.size(0)? != out_channels) + { + return Err(MinitensorError::shape_mismatch( + vec![out_channels], + vec![b.size(0)?], + )); } if stride.0 == 0 || stride.1 == 0 { @@ -158,7 +158,7 @@ pub fn conv2d( let requires_grad = input.requires_grad() || weight.requires_grad() - || bias.map_or(false, |b| b.requires_grad()); + || bias.is_some_and(|b| b.requires_grad()); let output_data = TensorData::from_vec_f32(output_vec, input.device()); let mut output = Tensor::new( Arc::new(output_data), @@ -176,7 +176,7 @@ pub fn conv2d( if weight.requires_grad() { deps.push(weight.id()); } - let bias_requires_grad = bias.map_or(false, |b| b.requires_grad()); + let bias_requires_grad = bias.is_some_and(|b| b.requires_grad()); if bias_requires_grad { deps.push(bias.unwrap().id()); } diff --git a/engine/src/operations/fusion.rs b/engine/src/operations/fusion.rs deleted file mode 100644 index e559ab71..00000000 --- a/engine/src/operations/fusion.rs +++ /dev/null @@ -1,623 +0,0 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -use crate::{ - device::Device, - error::{MinitensorError, Result}, - tensor::{DataType, Shape, Tensor, TensorData}, -}; -use rayon::prelude::*; -use std::{ - collections::{HashMap, VecDeque}, - sync::{Arc, Mutex}, -}; - -const CHUNK: usize = 1024; - -#[inline(always)] -fn unary_apply_f32(input: &[f32], output: &mut [f32], f: F) -where - F: Fn(f32) -> f32 + Sync, -{ - output - .par_chunks_mut(CHUNK) - .zip(input.par_chunks(CHUNK)) - .for_each(|(out, inp)| unsafe { - let in_ptr = inp.as_ptr(); - let out_ptr = out.as_mut_ptr(); - for i in 0..out.len() { - *out_ptr.add(i) = f(*in_ptr.add(i)); - } - }); -} - -#[inline(always)] -fn binary_apply_f32(lhs: &[f32], rhs: &[f32], output: &mut [f32], f: F) -where - F: Fn(f32, f32) -> f32 + Sync, -{ - output - .par_chunks_mut(CHUNK) - .zip(lhs.par_chunks(CHUNK).zip(rhs.par_chunks(CHUNK))) - .for_each(|(out, (a, b))| unsafe { - let a_ptr = a.as_ptr(); - let b_ptr = b.as_ptr(); - let out_ptr = out.as_mut_ptr(); - for i in 0..out.len() { - *out_ptr.add(i) = f(*a_ptr.add(i), *b_ptr.add(i)); - } - }); -} - -#[inline(always)] -fn unary_apply_f64(input: &[f64], output: &mut [f64], f: F) -where - F: Fn(f64) -> f64 + Sync, -{ - output - .par_chunks_mut(CHUNK) - .zip(input.par_chunks(CHUNK)) - .for_each(|(out, inp)| unsafe { - let in_ptr = inp.as_ptr(); - let out_ptr = out.as_mut_ptr(); - for i in 0..out.len() { - *out_ptr.add(i) = f(*in_ptr.add(i)); - } - }); -} - -#[inline(always)] -fn binary_apply_f64(lhs: &[f64], rhs: &[f64], output: &mut [f64], f: F) -where - F: Fn(f64, f64) -> f64 + Sync, -{ - output - .par_chunks_mut(CHUNK) - .zip(lhs.par_chunks(CHUNK).zip(rhs.par_chunks(CHUNK))) - .for_each(|(out, (a, b))| unsafe { - let a_ptr = a.as_ptr(); - let b_ptr = b.as_ptr(); - let out_ptr = out.as_mut_ptr(); - for i in 0..out.len() { - *out_ptr.add(i) = f(*a_ptr.add(i), *b_ptr.add(i)); - } - }); -} - -/// Represents a fused operation that can combine multiple tensor operations -#[derive(Debug, Clone)] -pub enum FusedOp { - /// Element-wise addition - Add, - /// Element-wise subtraction - Sub, - /// Element-wise multiplication - Mul, - /// Element-wise division - Div, - /// ReLU activation - ReLU, - /// Sigmoid activation - Sigmoid, - /// Tanh activation - Tanh, - /// Exponential - Exp, - /// Natural logarithm - Log, -} - -/// A sequence of operations that can be fused together -#[derive(Debug, Clone)] -pub struct FusionSequence { - pub operations: Vec, - pub input_shapes: Vec, - pub output_shape: Shape, - pub dtype: DataType, -} - -impl FusionSequence { - /// Create a new fusion sequence - pub fn new(dtype: DataType) -> Self { - Self { - operations: Vec::new(), - input_shapes: Vec::new(), - output_shape: Shape::new(vec![]), - dtype, - } - } - - /// Add an operation to the fusion sequence - pub fn add_operation(&mut self, op: FusedOp, input_shape: Shape) -> Result<()> { - self.operations.push(op); - self.input_shapes.push(input_shape.clone()); - - // Update output shape (for now, assume same as input for element-wise ops) - self.output_shape = input_shape; - - Ok(()) - } - - /// Check if this sequence can be fused with another operation - pub fn can_fuse_with(&self, op: &FusedOp, shape: &Shape) -> bool { - // For now, only fuse element-wise operations with compatible shapes - match op { - FusedOp::Add - | FusedOp::Sub - | FusedOp::Mul - | FusedOp::Div - | FusedOp::ReLU - | FusedOp::Sigmoid - | FusedOp::Tanh - | FusedOp::Exp - | FusedOp::Log => { - // Check if shapes are compatible - self.output_shape.dims() == shape.dims() && self.operations.len() < 8 - // Limit fusion depth - } - } - } - - /// Execute the fused operation sequence - pub fn execute(&self, inputs: &[&Tensor]) -> Result { - if inputs.is_empty() { - return Err(MinitensorError::invalid_operation( - "No input tensors provided", - )); - } - - let device = inputs[0].device(); - let mut current_data = inputs[0].data().clone(); - - // Execute operations in sequence - for (i, op) in self.operations.iter().enumerate() { - let second_input = inputs.get(i + 1).map(|t| *t); - current_data = Arc::new(self.execute_single_op(op, ¤t_data, second_input)?); - } - - Ok(Tensor::new( - current_data, - self.output_shape.clone(), - self.dtype, - device, - false, // For now, don't track gradients in fused ops - )) - } - - /// Execute a single operation in the fusion sequence - fn execute_single_op( - &self, - op: &FusedOp, - input: &TensorData, - second_input: Option<&Tensor>, - ) -> Result { - match self.dtype { - DataType::Float32 => self.execute_f32_op(op, input, second_input), - DataType::Float64 => self.execute_f64_op(op, input, second_input), - _ => Err(MinitensorError::invalid_operation( - "Unsupported data type for fusion", - )), - } - } - - /// Execute f32 operations - fn execute_f32_op( - &self, - op: &FusedOp, - input: &TensorData, - second_input: Option<&Tensor>, - ) -> Result { - let input_slice = input - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice from input"))?; - - let mut output_data = TensorData::zeros(input_slice.len(), DataType::Float32); - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output") - })?; - - match op { - FusedOp::Add => { - if let Some(second) = second_input { - let second_slice = second.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from second input") - })?; - binary_apply_f32(input_slice, second_slice, output_slice, |a, b| a + b); - } else { - return Err(MinitensorError::invalid_operation( - "Add operation requires two inputs", - )); - } - } - FusedOp::Sub => { - if let Some(second) = second_input { - let second_slice = second.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from second input") - })?; - binary_apply_f32(input_slice, second_slice, output_slice, |a, b| a - b); - } else { - return Err(MinitensorError::invalid_operation( - "Sub operation requires two inputs", - )); - } - } - FusedOp::Mul => { - if let Some(second) = second_input { - let second_slice = second.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from second input") - })?; - binary_apply_f32(input_slice, second_slice, output_slice, |a, b| a * b); - } else { - return Err(MinitensorError::invalid_operation( - "Mul operation requires two inputs", - )); - } - } - FusedOp::Div => { - if let Some(second) = second_input { - let second_slice = second.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from second input") - })?; - binary_apply_f32(input_slice, second_slice, output_slice, |a, b| a / b); - } else { - return Err(MinitensorError::invalid_operation( - "Div operation requires two inputs", - )); - } - } - FusedOp::ReLU => { - unary_apply_f32(input_slice, output_slice, |x| x.max(0.0)); - } - FusedOp::Sigmoid => { - unary_apply_f32(input_slice, output_slice, |x| 1.0 / (1.0 + (-x).exp())); - } - FusedOp::Tanh => { - unary_apply_f32(input_slice, output_slice, |x| x.tanh()); - } - FusedOp::Exp => { - unary_apply_f32(input_slice, output_slice, |x| x.exp()); - } - FusedOp::Log => { - unary_apply_f32(input_slice, output_slice, |x| x.ln()); - } - } - - Ok(output_data) - } - - /// Execute f64 operations - fn execute_f64_op( - &self, - op: &FusedOp, - input: &TensorData, - second_input: Option<&Tensor>, - ) -> Result { - let input_slice = input - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice from input"))?; - - let mut output_data = TensorData::zeros(input_slice.len(), DataType::Float64); - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output") - })?; - - match op { - FusedOp::Add => { - if let Some(second) = second_input { - let second_slice = second.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from second input") - })?; - binary_apply_f64(input_slice, second_slice, output_slice, |a, b| a + b); - } else { - return Err(MinitensorError::invalid_operation( - "Add operation requires two inputs", - )); - } - } - FusedOp::Sub => { - if let Some(second) = second_input { - let second_slice = second.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from second input") - })?; - binary_apply_f64(input_slice, second_slice, output_slice, |a, b| a - b); - } else { - return Err(MinitensorError::invalid_operation( - "Sub operation requires two inputs", - )); - } - } - FusedOp::Mul => { - if let Some(second) = second_input { - let second_slice = second.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from second input") - })?; - binary_apply_f64(input_slice, second_slice, output_slice, |a, b| a * b); - } else { - return Err(MinitensorError::invalid_operation( - "Mul operation requires two inputs", - )); - } - } - FusedOp::Div => { - if let Some(second) = second_input { - let second_slice = second.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from second input") - })?; - binary_apply_f64(input_slice, second_slice, output_slice, |a, b| a / b); - } else { - return Err(MinitensorError::invalid_operation( - "Div operation requires two inputs", - )); - } - } - FusedOp::ReLU => { - unary_apply_f64(input_slice, output_slice, |x| x.max(0.0)); - } - FusedOp::Sigmoid => { - unary_apply_f64(input_slice, output_slice, |x| 1.0 / (1.0 + (-x).exp())); - } - FusedOp::Tanh => { - unary_apply_f64(input_slice, output_slice, |x| x.tanh()); - } - FusedOp::Exp => { - unary_apply_f64(input_slice, output_slice, |x| x.exp()); - } - FusedOp::Log => { - unary_apply_f64(input_slice, output_slice, |x| x.ln()); - } - } - - Ok(output_data) - } -} - -/// Memory pool for reusing tensor allocations -pub struct MemoryPool { - pools: HashMap<(DataType, Device), VecDeque>, - max_pool_size: usize, -} - -impl MemoryPool { - /// Create a new memory pool - pub fn new(max_pool_size: usize) -> Self { - Self { - pools: HashMap::new(), - max_pool_size, - } - } - - /// Get a tensor data from the pool or allocate a new one - pub fn get_or_allocate(&mut self, size: usize, dtype: DataType, device: Device) -> TensorData { - let key = (dtype, device); - - if let Some(pool) = self.pools.get_mut(&key) { - // Try to find a suitable tensor in the pool - for i in 0..pool.len() { - if pool[i].len() >= size { - let mut data = pool.remove(i).unwrap(); - // Resize if necessary (this is a no-op if size matches) - if data.len() != size { - data = TensorData::zeros(size, dtype); - } - return data; - } - } - } - - // No suitable tensor found, allocate a new one - TensorData::zeros_on_device(size, dtype, device) - } - - /// Return a tensor data to the pool for reuse - pub fn return_to_pool(&mut self, data: TensorData, dtype: DataType, device: Device) { - let key = (dtype, device); - - let pool = self.pools.entry(key).or_insert_with(VecDeque::new); - - // Only keep the tensor if the pool isn't full - if pool.len() < self.max_pool_size { - pool.push_back(data); - } - // Otherwise, let the tensor be dropped and deallocated - } - - /// Clear all pools - pub fn clear(&mut self) { - self.pools.clear(); - } - - /// Get statistics about the memory pool - pub fn stats(&self) -> MemoryPoolStats { - let mut total_tensors = 0; - let mut total_memory = 0; - - for pool in self.pools.values() { - total_tensors += pool.len(); - for data in pool { - total_memory += data.len() * data.dtype().size_in_bytes(); - } - } - - MemoryPoolStats { - total_tensors, - total_memory_bytes: total_memory, - pool_count: self.pools.len(), - } - } -} - -/// Statistics about memory pool usage -#[derive(Debug, Clone)] -pub struct MemoryPoolStats { - pub total_tensors: usize, - pub total_memory_bytes: usize, - pub pool_count: usize, -} - -/// Global memory pool instance -static GLOBAL_MEMORY_POOL: Mutex> = Mutex::new(None); - -/// Initialize the global memory pool -pub fn init_memory_pool(max_pool_size: usize) { - let mut pool = GLOBAL_MEMORY_POOL.lock().unwrap(); - *pool = Some(MemoryPool::new(max_pool_size)); -} - -/// Get a tensor from the global memory pool -pub fn get_pooled_tensor(size: usize, dtype: DataType, device: Device) -> TensorData { - let mut pool_guard = GLOBAL_MEMORY_POOL.lock().unwrap(); - if let Some(ref mut pool) = *pool_guard { - pool.get_or_allocate(size, dtype, device) - } else { - // Pool not initialized, allocate directly - TensorData::zeros_on_device(size, dtype, device) - } -} - -/// Return a tensor to the global memory pool -pub fn return_pooled_tensor(data: TensorData, dtype: DataType, device: Device) { - let mut pool_guard = GLOBAL_MEMORY_POOL.lock().unwrap(); - if let Some(ref mut pool) = *pool_guard { - pool.return_to_pool(data, dtype, device); - } - // If pool not initialized, just let the tensor be dropped -} - -/// Get memory pool statistics -pub fn memory_pool_stats() -> Option { - let pool_guard = GLOBAL_MEMORY_POOL.lock().unwrap(); - pool_guard.as_ref().map(|pool| pool.stats()) -} - -/// Lazy evaluation system for deferred computation -pub struct LazyTensor { - pub operation: FusedOp, - pub inputs: Vec>, - pub shape: Shape, - pub dtype: DataType, - pub device: Device, - pub computed_value: Option>, -} - -impl LazyTensor { - /// Create a new lazy tensor from a concrete value - pub fn from_tensor(tensor: &Tensor) -> Self { - Self { - operation: FusedOp::Add, // Placeholder for leaf nodes - inputs: Vec::new(), - shape: tensor.shape().clone(), - dtype: tensor.dtype(), - device: tensor.device(), - computed_value: Some(tensor.data().clone()), - } - } - - /// Create a new lazy tensor from an operation - pub fn from_operation( - operation: FusedOp, - inputs: Vec>, - shape: Shape, - dtype: DataType, - device: Device, - ) -> Self { - Self { - operation, - inputs, - shape, - dtype, - device, - computed_value: None, - } - } - - /// Compute the value of this lazy tensor - pub fn compute(&mut self) -> Result> { - if let Some(ref value) = self.computed_value { - return Ok(value.clone()); - } - - // Compute input values first - let mut input_values = Vec::new(); - for input in &mut self.inputs { - // This is a simplified approach - in practice, you'd need interior mutability - // or a different design to handle mutable references properly - input_values.push( - input - .computed_value - .clone() - .ok_or_else(|| MinitensorError::internal_error("Input value not computed"))?, - ); - } - - // Create a fusion sequence and execute it - let mut sequence = FusionSequence::new(self.dtype); - sequence.add_operation(self.operation.clone(), self.shape.clone())?; - - // For now, this is a simplified implementation - // In practice, you'd need to handle the conversion from TensorData to Tensor properly - let result_data = get_pooled_tensor(self.shape.numel(), self.dtype, self.device); - let result = Arc::new(result_data); - - self.computed_value = Some(result.clone()); - Ok(result) - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::device::Device; - - #[test] - fn test_fusion_sequence_creation() { - let mut sequence = FusionSequence::new(DataType::Float32); - let shape = Shape::new(vec![4]); - - sequence.add_operation(FusedOp::Add, shape.clone()).unwrap(); - sequence.add_operation(FusedOp::ReLU, shape).unwrap(); - - assert_eq!(sequence.operations.len(), 2); - } - - #[test] - fn test_memory_pool() { - let mut pool = MemoryPool::new(5); - - // Allocate some tensors - let data1 = pool.get_or_allocate(100, DataType::Float32, Device::cpu()); - let data2 = pool.get_or_allocate(200, DataType::Float32, Device::cpu()); - - // Return them to pool - pool.return_to_pool(data1, DataType::Float32, Device::cpu()); - pool.return_to_pool(data2, DataType::Float32, Device::cpu()); - - let stats = pool.stats(); - assert_eq!(stats.total_tensors, 2); - } - - #[test] - fn test_can_fuse_operations() { - let mut sequence = FusionSequence::new(DataType::Float32); - let shape = Shape::new(vec![4]); - - sequence.add_operation(FusedOp::Add, shape.clone()).unwrap(); - - assert!(sequence.can_fuse_with(&FusedOp::ReLU, &shape)); - assert!(sequence.can_fuse_with(&FusedOp::Sigmoid, &shape)); - } - - #[test] - fn test_global_memory_pool() { - init_memory_pool(10); - - let data = get_pooled_tensor(50, DataType::Float32, Device::cpu()); - return_pooled_tensor(data, DataType::Float32, Device::cpu()); - - if let Some(stats) = memory_pool_stats() { - assert!(stats.total_tensors > 0); - } - } -} diff --git a/engine/src/operations/linalg.rs b/engine/src/operations/linalg.rs index 9646912d..0d314dd6 100644 --- a/engine/src/operations/linalg.rs +++ b/engine/src/operations/linalg.rs @@ -4,6 +4,13 @@ // This source code is licensed under the Apache-style license found in the // LICENSE file in the root directory of this source tree. -include!("linalg/matmul.rs"); -include!("linalg/diagonal.rs"); -include!("linalg/triangular.rs"); +#[path = "linalg/diagonal.rs"] +mod diagonal_impl; +#[path = "linalg/matmul.rs"] +mod matmul_impl; +#[path = "linalg/triangular.rs"] +mod triangular_impl; + +pub use self::diagonal_impl::*; +pub use self::matmul_impl::*; +pub(crate) use self::triangular_impl::*; diff --git a/engine/src/operations/linalg/diagonal.rs b/engine/src/operations/linalg/diagonal.rs index a5358595..d1ed6bcb 100644 --- a/engine/src/operations/linalg/diagonal.rs +++ b/engine/src/operations/linalg/diagonal.rs @@ -1,767 +1,778 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -/// Transpose operation with gradient support -pub fn transpose(tensor: &Tensor, dim0: isize, dim1: isize) -> Result { - let ndim = tensor.ndim() as isize; - let dim0 = if dim0 < 0 { dim0 + ndim } else { dim0 }; - let dim1 = if dim1 < 0 { dim1 + ndim } else { dim1 }; - - if dim0 < 0 || dim0 >= ndim || dim1 < 0 || dim1 >= ndim { - return Err(MinitensorError::index_error( - dim0.max(dim1), - 0, - ndim as usize, - )); - } - - if dim0 == dim1 { - // No-op transpose - return Ok(tensor.clone()); - } - - let dim0_usize = dim0 as usize; - let dim1_usize = dim1 as usize; - - // Create new shape with swapped dimensions - let mut new_shape = tensor.shape().dims().to_vec(); - new_shape.swap(dim0_usize, dim1_usize); - let new_shape_obj = Shape::new(new_shape); - - // Create new strides with swapped dimensions - let old_strides = tensor.strides().as_slice(); - let mut new_strides = old_strides.to_vec(); - new_strides.swap(dim0_usize, dim1_usize); - - // Create output tensor data by copying and rearranging - let mut output_data = - TensorData::zeros_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - // Perform transpose based on data type - match tensor.dtype() { - DataType::Float32 => transpose_f32( - tensor, - &mut output_data, - &new_shape_obj, - dim0_usize, - dim1_usize, - )?, - DataType::Float64 => transpose_f64( - tensor, - &mut output_data, - &new_shape_obj, - dim0_usize, - dim1_usize, - )?, - DataType::Int32 => transpose_i32( - tensor, - &mut output_data, - &new_shape_obj, - dim0_usize, - dim1_usize, - )?, - DataType::Int64 => transpose_i64( - tensor, - &mut output_data, - &new_shape_obj, - dim0_usize, - dim1_usize, - )?, - DataType::Bool => transpose_bool( - tensor, - &mut output_data, - &new_shape_obj, - dim0_usize, - dim1_usize, - )?, - } - - // Create output tensor - let output = Tensor::new( - Arc::new(output_data), - new_shape_obj, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - // Set up gradient function if needed - if output.requires_grad() { - let grad_fn = Arc::new(TransposeBackward { - dims: vec![dim0_usize, dim1_usize], - input_id: tensor.id(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&output_with_grad, Some(grad_fn))?; - - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Extract a diagonal from the tensor, reducing two dimensions into one. -pub fn diagonal(tensor: &Tensor, offset: isize, dim1: isize, dim2: isize) -> Result { - if tensor.ndim() < 2 { - return Err(MinitensorError::invalid_operation( - "diagonal requires tensors with at least 2 dimensions", - )); - } - - if !tensor.device().is_cpu() { - return Err(MinitensorError::invalid_operation( - "diagonal currently supports only CPU tensors", - )); - } - - let ndim = tensor.ndim(); - let dim1 = normalize_dim(dim1, ndim)?; - let dim2 = normalize_dim(dim2, ndim)?; - if dim1 == dim2 { - return Err(MinitensorError::invalid_operation( - "diagonal dimensions must be distinct", - )); - } - - let dims = tensor.shape().dims(); - let strides = tensor.strides().as_slice(); - let spec = compute_diagonal_spec(dims, strides, dim1, dim2, offset)?; - let out_shape = Shape::new(spec.output_dims.clone()); - let dtype = tensor.dtype(); - let device = tensor.device(); - let mut output_data = TensorData::zeros_on_device(out_shape.numel(), dtype, device); - - if out_shape.numel() > 0 { - match dtype { - DataType::Float32 => { - let input = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from tensor") - })?; - let output = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice for diagonal output", - ) - })?; - diagonal_copy(input, output, dims, strides, &spec); - } - DataType::Float64 => { - let input = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from tensor") - })?; - let output = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice for diagonal output", - ) - })?; - diagonal_copy(input, output, dims, strides, &spec); - } - DataType::Int32 => { - let input = tensor.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from tensor") - })?; - let output = output_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable i32 slice for diagonal output", - ) - })?; - diagonal_copy(input, output, dims, strides, &spec); - } - DataType::Int64 => { - let input = tensor.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from tensor") - })?; - let output = output_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable i64 slice for diagonal output", - ) - })?; - diagonal_copy(input, output, dims, strides, &spec); - } - DataType::Bool => { - let input = tensor.data().as_bool_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from tensor") - })?; - let output = output_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable bool slice for diagonal output", - ) - })?; - diagonal_copy(input, output, dims, strides, &spec); - } - } - } - - let mut output = Tensor::new( - Arc::new(output_data), - out_shape, - dtype, - device, - tensor.requires_grad(), - ); - - if tensor.requires_grad() { - let grad_fn = Arc::new(crate::autograd::DiagonalBackward { - input_shape: dims.to_vec(), - input_strides: strides.to_vec(), - input_dtype: dtype, - dim1, - dim2, - offset, - input_requires_grad: tensor.requires_grad(), - input_id: tensor.id(), - }); - - output.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output, Some(grad_fn))?; - } - - Ok(output) -} - -/// Sum of the diagonal elements along two dimensions. -pub fn trace(tensor: &Tensor, offset: isize, dim1: isize, dim2: isize) -> Result { - let diag = diagonal(tensor, offset, dim1, dim2)?; - if diag.ndim() == 0 { - return Ok(diag); - } - - reduction::sum(&diag, Some(vec![-1]), false) -} - -/// Return the upper triangular part of a matrix (or batch of matrices). -pub fn triu(tensor: &Tensor, diagonal: i64) -> Result { - triangular_op(tensor, diagonal, true) -} - -/// Return the lower triangular part of a matrix (or batch of matrices). -pub fn tril(tensor: &Tensor, diagonal: i64) -> Result { - triangular_op(tensor, diagonal, false) -} - -fn triangular_op(tensor: &Tensor, diagonal: i64, upper: bool) -> Result { - if tensor.ndim() < 2 { - return Err(MinitensorError::invalid_operation( - "triangular operations require tensors with at least 2 dimensions", - )); - } - - let clamped_diagonal = diagonal.clamp(isize::MIN as i64, isize::MAX as i64) as isize; - - let mut output_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - apply_triangular_mask(tensor, &mut output_data, clamped_diagonal, upper)?; - - let mut output = Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if tensor.requires_grad() { - let grad_fn = Arc::new(crate::autograd::TriangularBackward { - input_shape: tensor.shape().dims().to_vec(), - diagonal: clamped_diagonal, - upper, - input_requires_grad: tensor.requires_grad(), - input_id: tensor.id(), - }); - - output.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output, Some(grad_fn))?; - } - - Ok(output) -} - -pub(crate) fn apply_triangular_mask( - tensor: &Tensor, - output_data: &mut TensorData, - diagonal: isize, - upper: bool, -) -> Result<()> { - match tensor.dtype() { - DataType::Float32 => { - let input = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from tensor") - })?; - let output = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice for triangular output", - ) - })?; - copy_and_mask(input, output, tensor.shape(), diagonal, upper); - } - DataType::Float64 => { - let input = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from tensor") - })?; - let output = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice for triangular output", - ) - })?; - copy_and_mask(input, output, tensor.shape(), diagonal, upper); - } - DataType::Int32 => { - let input = tensor.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from tensor") - })?; - let output = output_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable i32 slice for triangular output", - ) - })?; - copy_and_mask(input, output, tensor.shape(), diagonal, upper); - } - DataType::Int64 => { - let input = tensor.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from tensor") - })?; - let output = output_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable i64 slice for triangular output", - ) - })?; - copy_and_mask(input, output, tensor.shape(), diagonal, upper); - } - DataType::Bool => { - let input = tensor.data().as_bool_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from tensor") - })?; - let output = output_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable bool slice for triangular output", - ) - })?; - copy_and_mask(input, output, tensor.shape(), diagonal, upper); - } - } - - Ok(()) -} - -fn copy_and_mask( - input: &[T], - output: &mut [T], - shape: &Shape, - diagonal: isize, - upper: bool, -) { - if input.is_empty() { - return; - } - - output.copy_from_slice(input); - - let dims = shape.dims(); - debug_assert!(dims.len() >= 2); - let rows = dims[dims.len() - 2]; - let cols = dims[dims.len() - 1]; - - if rows == 0 || cols == 0 { - return; - } - - let batch = shape.numel() / (rows * cols); - let zero = T::default(); - - for b in 0..batch { - let base = b * rows * cols; - for r in 0..rows { - let row_offset = base + r * cols; - let row_idx = r as isize; - for c in 0..cols { - let col_idx = c as isize; - let keep = if upper { - col_idx - row_idx >= diagonal - } else { - col_idx - row_idx <= diagonal - }; - if !keep { - output[row_offset + c] = zero; - } - } - } - } -} - -// Helper functions for matrix multiplication - -fn matmul_f32( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - _output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from rhs tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - optimized_matmul_f32(lhs_data, rhs_data, output_slice, lhs.shape(), rhs.shape()) -} - -fn matmul_f64( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - _output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from rhs tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - optimized_matmul_f64(lhs_data, rhs_data, output_slice, lhs.shape(), rhs.shape()) -} - -fn matmul_i32( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from rhs tensor") - })?; - - let output_slice = output_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i32 slice from output data") - })?; - - naive_matmul( - lhs_data, - rhs_data, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - ) -} - -fn matmul_i64( - lhs: &Tensor, - rhs: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, -) -> Result<()> { - let lhs_data = lhs.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from lhs tensor") - })?; - let rhs_data = rhs.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from rhs tensor") - })?; - - let output_slice = output_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice from output data") - })?; - - naive_matmul( - lhs_data, - rhs_data, - output_slice, - lhs.shape(), - rhs.shape(), - output_shape, - ) -} - -/// Naive matrix multiplication implementation (O(n^3)) with batch support -fn naive_matmul( - lhs_data: &[T], - rhs_data: &[T], - output_data: &mut [T], - lhs_shape: &Shape, - rhs_shape: &Shape, - _output_shape: &Shape, -) -> Result<()> -where - T: Copy + std::ops::Add + std::ops::Mul + Default + Send + Sync, -{ - let lhs_dims = lhs_shape.dims(); - let rhs_dims = rhs_shape.dims(); - - let m = lhs_dims[lhs_dims.len() - 2]; - let k = lhs_dims[lhs_dims.len() - 1]; - let n = rhs_dims[rhs_dims.len() - 1]; - let batch = lhs_data.len() / (m * k); - if batch == 1 && m * n * k < PAR_THRESHOLD { - // For small single-batch matrices, avoid parallel overhead - for i in 0..m { - for j in 0..n { - let mut sum = T::default(); - for l in 0..k { - let lhs_idx = i * k + l; - let rhs_idx = l * n + j; - sum = sum + lhs_data[lhs_idx] * rhs_data[rhs_idx]; - } - output_data[i * n + j] = sum; - } - } - } else { - output_data - .par_chunks_mut(m * n) - .enumerate() - .for_each(|(b, chunk)| { - let lhs_batch = &lhs_data[b * m * k..(b + 1) * m * k]; - let rhs_batch = &rhs_data[b * k * n..(b + 1) * k * n]; - chunk.par_chunks_mut(n).enumerate().for_each(|(i, row)| { - for j in 0..n { - let mut sum = T::default(); - for l in 0..k { - let lhs_idx = i * k + l; - let rhs_idx = l * n + j; - sum = sum + lhs_batch[lhs_idx] * rhs_batch[rhs_idx]; - } - row[j] = sum; - } - }); - }); - } - - Ok(()) -} - -fn optimized_matmul_f32( - lhs_data: &[f32], - rhs_data: &[f32], - output_data: &mut [f32], - lhs_shape: &Shape, - rhs_shape: &Shape, -) -> Result<()> { - let lhs_dims = lhs_shape.dims(); - let rhs_dims = rhs_shape.dims(); - let m = lhs_dims[lhs_dims.len() - 2]; - let k = lhs_dims[lhs_dims.len() - 1]; - let n = rhs_dims[rhs_dims.len() - 1]; - - if m == 0 || k == 0 || n == 0 { - // Nothing to compute for zero-sized dimensions - return Ok(()); - } - - let batch = lhs_data.len() / (m * k); - if batch == 1 { - // Avoid parallel overhead for single matrix multiplication - unsafe { - gemm_f32( - m, - k, - n, - lhs_data.as_ptr(), - rhs_data.as_ptr(), - output_data.as_mut_ptr(), - ) - }; - } else { - output_data - .par_chunks_mut(m * n) - .enumerate() - .for_each(|(b, chunk)| { - let a = &lhs_data[b * m * k..(b + 1) * m * k]; - let r = &rhs_data[b * k * n..(b + 1) * k * n]; - unsafe { - gemm_f32(m, k, n, a.as_ptr(), r.as_ptr(), chunk.as_mut_ptr()); - } - }); - } - - Ok(()) -} - -fn optimized_matmul_f64( - lhs_data: &[f64], - rhs_data: &[f64], - output_data: &mut [f64], - lhs_shape: &Shape, - rhs_shape: &Shape, -) -> Result<()> { - let lhs_dims = lhs_shape.dims(); - let rhs_dims = rhs_shape.dims(); - let m = lhs_dims[lhs_dims.len() - 2]; - let k = lhs_dims[lhs_dims.len() - 1]; - let n = rhs_dims[rhs_dims.len() - 1]; - - if m == 0 || k == 0 || n == 0 { - return Ok(()); - } - - let batch = lhs_data.len() / (m * k); - if batch == 1 { - unsafe { - gemm_f64( - m, - k, - n, - lhs_data.as_ptr(), - rhs_data.as_ptr(), - output_data.as_mut_ptr(), - ) - }; - } else { - output_data - .par_chunks_mut(m * n) - .enumerate() - .for_each(|(b, chunk)| { - let a = &lhs_data[b * m * k..(b + 1) * m * k]; - let r = &rhs_data[b * k * n..(b + 1) * k * n]; - unsafe { - gemm_f64(m, k, n, a.as_ptr(), r.as_ptr(), chunk.as_mut_ptr()); - } - }); - } - - Ok(()) -} - -// Helper functions for transpose operations - -fn transpose_f32( - tensor: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, - dim0: usize, - dim1: usize, -) -> Result<()> { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from input tensor") - })?; - - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output data") - })?; - - transpose_generic( - input_data, - output_slice, - tensor.shape(), - output_shape, - dim0, - dim1, - ) -} - -fn transpose_f64( - tensor: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, - dim0: usize, - dim1: usize, -) -> Result<()> { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from input tensor") - })?; - - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output data") - })?; - - transpose_generic( - input_data, - output_slice, - tensor.shape(), - output_shape, - dim0, - dim1, - ) -} - -fn transpose_i32( - tensor: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, - dim0: usize, - dim1: usize, -) -> Result<()> { - let input_data = tensor.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from input tensor") - })?; - - let output_slice = output_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i32 slice from output data") - })?; - - transpose_generic( - input_data, - output_slice, - tensor.shape(), - output_shape, - dim0, - dim1, - ) -} - -fn transpose_i64( - tensor: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, - dim0: usize, - dim1: usize, -) -> Result<()> { - let input_data = tensor.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from input tensor") - })?; - - let output_slice = output_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice from output data") - })?; - - transpose_generic( - input_data, - output_slice, - tensor.shape(), - output_shape, - dim0, - dim1, - ) -} - -fn transpose_bool( - tensor: &Tensor, - output_data: &mut TensorData, - output_shape: &Shape, - dim0: usize, - dim1: usize, -) -> Result<()> { - let input_data = tensor.data().as_bool_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from input tensor") - })?; - - let output_slice = output_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice from output data") - })?; - - transpose_generic( - input_data, - output_slice, - tensor.shape(), - output_shape, - dim0, - dim1, - ) -} +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::autograd::TransposeBackward; +use crate::operations::reduction; +use crate::{ + autograd::add_to_graph, + error::{MinitensorError, Result}, + tensor::{DataType, Shape, Tensor, TensorData}, +}; +use rayon::prelude::*; +use std::sync::Arc; + +/// Transpose operation with gradient support +pub fn transpose(tensor: &Tensor, dim0: isize, dim1: isize) -> Result { + let ndim = tensor.ndim() as isize; + let dim0 = if dim0 < 0 { dim0 + ndim } else { dim0 }; + let dim1 = if dim1 < 0 { dim1 + ndim } else { dim1 }; + + if dim0 < 0 || dim0 >= ndim || dim1 < 0 || dim1 >= ndim { + return Err(MinitensorError::index_error( + dim0.max(dim1), + 0, + ndim as usize, + )); + } + + if dim0 == dim1 { + // No-op transpose + return Ok(tensor.clone()); + } + + let dim0_usize = dim0 as usize; + let dim1_usize = dim1 as usize; + + // Create new shape with swapped dimensions + let mut new_shape = tensor.shape().dims().to_vec(); + new_shape.swap(dim0_usize, dim1_usize); + let new_shape_obj = Shape::new(new_shape); + + // Create new strides with swapped dimensions + let old_strides = tensor.strides().as_slice(); + let mut new_strides = old_strides.to_vec(); + new_strides.swap(dim0_usize, dim1_usize); + + // Create output tensor data by copying and rearranging + let mut output_data = + TensorData::zeros_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + // Perform transpose based on data type + match tensor.dtype() { + DataType::Float32 => transpose_f32( + tensor, + &mut output_data, + &new_shape_obj, + dim0_usize, + dim1_usize, + )?, + DataType::Float64 => transpose_f64( + tensor, + &mut output_data, + &new_shape_obj, + dim0_usize, + dim1_usize, + )?, + DataType::Int32 => transpose_i32( + tensor, + &mut output_data, + &new_shape_obj, + dim0_usize, + dim1_usize, + )?, + DataType::Int64 => transpose_i64( + tensor, + &mut output_data, + &new_shape_obj, + dim0_usize, + dim1_usize, + )?, + DataType::Bool => transpose_bool( + tensor, + &mut output_data, + &new_shape_obj, + dim0_usize, + dim1_usize, + )?, + } + + // Create output tensor + let output = Tensor::new( + Arc::new(output_data), + new_shape_obj, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + // Set up gradient function if needed + if output.requires_grad() { + let grad_fn = Arc::new(TransposeBackward { + dims: vec![dim0_usize, dim1_usize], + input_id: tensor.id(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&output_with_grad, Some(grad_fn))?; + + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Extract a diagonal from the tensor, reducing two dimensions into one. +pub fn diagonal(tensor: &Tensor, offset: isize, dim1: isize, dim2: isize) -> Result { + if tensor.ndim() < 2 { + return Err(MinitensorError::invalid_operation( + "diagonal requires tensors with at least 2 dimensions", + )); + } + + if !tensor.device().is_cpu() { + return Err(MinitensorError::invalid_operation( + "diagonal currently supports only CPU tensors", + )); + } + + let ndim = tensor.ndim(); + let dim1 = normalize_dim(dim1, ndim)?; + let dim2 = normalize_dim(dim2, ndim)?; + if dim1 == dim2 { + return Err(MinitensorError::invalid_operation( + "diagonal dimensions must be distinct", + )); + } + + let dims = tensor.shape().dims(); + let strides = tensor.strides().as_slice(); + let spec = compute_diagonal_spec(dims, strides, dim1, dim2, offset)?; + let out_shape = Shape::new(spec.output_dims.clone()); + let dtype = tensor.dtype(); + let device = tensor.device(); + let mut output_data = TensorData::zeros_on_device(out_shape.numel(), dtype, device); + + if out_shape.numel() > 0 { + match dtype { + DataType::Float32 => { + let input = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from tensor") + })?; + let output = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice for diagonal output", + ) + })?; + diagonal_copy(input, output, dims, strides, &spec); + } + DataType::Float64 => { + let input = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from tensor") + })?; + let output = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice for diagonal output", + ) + })?; + diagonal_copy(input, output, dims, strides, &spec); + } + DataType::Int32 => { + let input = tensor.data().as_i32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i32 slice from tensor") + })?; + let output = output_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable i32 slice for diagonal output", + ) + })?; + diagonal_copy(input, output, dims, strides, &spec); + } + DataType::Int64 => { + let input = tensor.data().as_i64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i64 slice from tensor") + })?; + let output = output_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable i64 slice for diagonal output", + ) + })?; + diagonal_copy(input, output, dims, strides, &spec); + } + DataType::Bool => { + let input = tensor.data().as_bool_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get bool slice from tensor") + })?; + let output = output_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable bool slice for diagonal output", + ) + })?; + diagonal_copy(input, output, dims, strides, &spec); + } + } + } + + let mut output = Tensor::new( + Arc::new(output_data), + out_shape, + dtype, + device, + tensor.requires_grad(), + ); + + if tensor.requires_grad() { + let grad_fn = Arc::new(crate::autograd::DiagonalBackward { + input_shape: dims.to_vec(), + input_strides: strides.to_vec(), + input_dtype: dtype, + dim1, + dim2, + offset, + input_requires_grad: tensor.requires_grad(), + input_id: tensor.id(), + }); + + output.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output, Some(grad_fn))?; + } + + Ok(output) +} + +/// Sum of the diagonal elements along two dimensions. +pub fn trace(tensor: &Tensor, offset: isize, dim1: isize, dim2: isize) -> Result { + let diag = diagonal(tensor, offset, dim1, dim2)?; + if diag.ndim() == 0 { + return Ok(diag); + } + + reduction::sum(&diag, Some(vec![-1]), false) +} + +/// Return the upper triangular part of a matrix (or batch of matrices). +pub fn triu(tensor: &Tensor, diagonal: i64) -> Result { + triangular_op(tensor, diagonal, true) +} + +/// Return the lower triangular part of a matrix (or batch of matrices). +pub fn tril(tensor: &Tensor, diagonal: i64) -> Result { + triangular_op(tensor, diagonal, false) +} + +fn triangular_op(tensor: &Tensor, diagonal: i64, upper: bool) -> Result { + if tensor.ndim() < 2 { + return Err(MinitensorError::invalid_operation( + "triangular operations require tensors with at least 2 dimensions", + )); + } + + let clamped_diagonal = diagonal.clamp(isize::MIN as i64, isize::MAX as i64) as isize; + + let mut output_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + apply_triangular_mask(tensor, &mut output_data, clamped_diagonal, upper)?; + + let mut output = Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if tensor.requires_grad() { + let grad_fn = Arc::new(crate::autograd::TriangularBackward { + input_shape: tensor.shape().dims().to_vec(), + diagonal: clamped_diagonal, + upper, + input_requires_grad: tensor.requires_grad(), + input_id: tensor.id(), + }); + + output.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output, Some(grad_fn))?; + } + + Ok(output) +} + +pub(crate) fn apply_triangular_mask( + tensor: &Tensor, + output_data: &mut TensorData, + diagonal: isize, + upper: bool, +) -> Result<()> { + match tensor.dtype() { + DataType::Float32 => { + let input = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from tensor") + })?; + let output = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice for triangular output", + ) + })?; + copy_and_mask(input, output, tensor.shape(), diagonal, upper); + } + DataType::Float64 => { + let input = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from tensor") + })?; + let output = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice for triangular output", + ) + })?; + copy_and_mask(input, output, tensor.shape(), diagonal, upper); + } + DataType::Int32 => { + let input = tensor.data().as_i32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i32 slice from tensor") + })?; + let output = output_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable i32 slice for triangular output", + ) + })?; + copy_and_mask(input, output, tensor.shape(), diagonal, upper); + } + DataType::Int64 => { + let input = tensor.data().as_i64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i64 slice from tensor") + })?; + let output = output_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable i64 slice for triangular output", + ) + })?; + copy_and_mask(input, output, tensor.shape(), diagonal, upper); + } + DataType::Bool => { + let input = tensor.data().as_bool_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get bool slice from tensor") + })?; + let output = output_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable bool slice for triangular output", + ) + })?; + copy_and_mask(input, output, tensor.shape(), diagonal, upper); + } + } + + Ok(()) +} + +fn copy_and_mask( + input: &[T], + output: &mut [T], + shape: &Shape, + diagonal: isize, + upper: bool, +) { + if input.is_empty() { + return; + } + + output.copy_from_slice(input); + + let dims = shape.dims(); + debug_assert!(dims.len() >= 2); + let rows = dims[dims.len() - 2]; + let cols = dims[dims.len() - 1]; + + if rows == 0 || cols == 0 { + return; + } + + let batch = shape.numel() / (rows * cols); + let zero = T::default(); + + for b in 0..batch { + let base = b * rows * cols; + for r in 0..rows { + let row_offset = base + r * cols; + let row_idx = r as isize; + for c in 0..cols { + let col_idx = c as isize; + let keep = if upper { + col_idx - row_idx >= diagonal + } else { + col_idx - row_idx <= diagonal + }; + if !keep { + output[row_offset + c] = zero; + } + } + } + } +} + +// Helper functions for matrix multiplication + +pub(crate) fn matmul_f32( + lhs: &Tensor, + rhs: &Tensor, + output_data: &mut TensorData, + _output_shape: &Shape, +) -> Result<()> { + let lhs_data = lhs.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from lhs tensor") + })?; + let rhs_data = rhs.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from rhs tensor") + })?; + + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + + optimized_matmul_f32(lhs_data, rhs_data, output_slice, lhs.shape(), rhs.shape()) +} + +pub(crate) fn matmul_f64( + lhs: &Tensor, + rhs: &Tensor, + output_data: &mut TensorData, + _output_shape: &Shape, +) -> Result<()> { + let lhs_data = lhs.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from lhs tensor") + })?; + let rhs_data = rhs.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from rhs tensor") + })?; + + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + + optimized_matmul_f64(lhs_data, rhs_data, output_slice, lhs.shape(), rhs.shape()) +} + +pub(crate) fn matmul_i32( + lhs: &Tensor, + rhs: &Tensor, + output_data: &mut TensorData, + output_shape: &Shape, +) -> Result<()> { + let lhs_data = lhs.data().as_i32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i32 slice from lhs tensor") + })?; + let rhs_data = rhs.data().as_i32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i32 slice from rhs tensor") + })?; + + let output_slice = output_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i32 slice from output data") + })?; + + naive_matmul( + lhs_data, + rhs_data, + output_slice, + lhs.shape(), + rhs.shape(), + output_shape, + ) +} + +pub(crate) fn matmul_i64( + lhs: &Tensor, + rhs: &Tensor, + output_data: &mut TensorData, + output_shape: &Shape, +) -> Result<()> { + let lhs_data = lhs.data().as_i64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i64 slice from lhs tensor") + })?; + let rhs_data = rhs.data().as_i64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i64 slice from rhs tensor") + })?; + + let output_slice = output_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice from output data") + })?; + + naive_matmul( + lhs_data, + rhs_data, + output_slice, + lhs.shape(), + rhs.shape(), + output_shape, + ) +} + +/// Naive matrix multiplication implementation (O(n^3)) with batch support +fn naive_matmul( + lhs_data: &[T], + rhs_data: &[T], + output_data: &mut [T], + lhs_shape: &Shape, + rhs_shape: &Shape, + _output_shape: &Shape, +) -> Result<()> +where + T: Copy + std::ops::Add + std::ops::Mul + Default + Send + Sync, +{ + let lhs_dims = lhs_shape.dims(); + let rhs_dims = rhs_shape.dims(); + + let m = lhs_dims[lhs_dims.len() - 2]; + let k = lhs_dims[lhs_dims.len() - 1]; + let n = rhs_dims[rhs_dims.len() - 1]; + let batch = lhs_data.len() / (m * k); + if batch == 1 && m * n * k < PAR_THRESHOLD { + // For small single-batch matrices, avoid parallel overhead + for i in 0..m { + for j in 0..n { + let mut sum = T::default(); + for l in 0..k { + let lhs_idx = i * k + l; + let rhs_idx = l * n + j; + sum = sum + lhs_data[lhs_idx] * rhs_data[rhs_idx]; + } + output_data[i * n + j] = sum; + } + } + } else { + output_data + .par_chunks_mut(m * n) + .enumerate() + .for_each(|(b, chunk)| { + let lhs_batch = &lhs_data[b * m * k..(b + 1) * m * k]; + let rhs_batch = &rhs_data[b * k * n..(b + 1) * k * n]; + chunk.par_chunks_mut(n).enumerate().for_each(|(i, row)| { + for j in 0..n { + let mut sum = T::default(); + for l in 0..k { + let lhs_idx = i * k + l; + let rhs_idx = l * n + j; + sum = sum + lhs_batch[lhs_idx] * rhs_batch[rhs_idx]; + } + row[j] = sum; + } + }); + }); + } + + Ok(()) +} + +fn optimized_matmul_f32( + lhs_data: &[f32], + rhs_data: &[f32], + output_data: &mut [f32], + lhs_shape: &Shape, + rhs_shape: &Shape, +) -> Result<()> { + let lhs_dims = lhs_shape.dims(); + let rhs_dims = rhs_shape.dims(); + let m = lhs_dims[lhs_dims.len() - 2]; + let k = lhs_dims[lhs_dims.len() - 1]; + let n = rhs_dims[rhs_dims.len() - 1]; + + if m == 0 || k == 0 || n == 0 { + // Nothing to compute for zero-sized dimensions + return Ok(()); + } + + let batch = lhs_data.len() / (m * k); + if batch == 1 { + // Avoid parallel overhead for single matrix multiplication + unsafe { + gemm_f32( + m, + k, + n, + lhs_data.as_ptr(), + rhs_data.as_ptr(), + output_data.as_mut_ptr(), + ) + }; + } else { + output_data + .par_chunks_mut(m * n) + .enumerate() + .for_each(|(b, chunk)| { + let a = &lhs_data[b * m * k..(b + 1) * m * k]; + let r = &rhs_data[b * k * n..(b + 1) * k * n]; + unsafe { + gemm_f32(m, k, n, a.as_ptr(), r.as_ptr(), chunk.as_mut_ptr()); + } + }); + } + + Ok(()) +} + +fn optimized_matmul_f64( + lhs_data: &[f64], + rhs_data: &[f64], + output_data: &mut [f64], + lhs_shape: &Shape, + rhs_shape: &Shape, +) -> Result<()> { + let lhs_dims = lhs_shape.dims(); + let rhs_dims = rhs_shape.dims(); + let m = lhs_dims[lhs_dims.len() - 2]; + let k = lhs_dims[lhs_dims.len() - 1]; + let n = rhs_dims[rhs_dims.len() - 1]; + + if m == 0 || k == 0 || n == 0 { + return Ok(()); + } + + let batch = lhs_data.len() / (m * k); + if batch == 1 { + unsafe { + gemm_f64( + m, + k, + n, + lhs_data.as_ptr(), + rhs_data.as_ptr(), + output_data.as_mut_ptr(), + ) + }; + } else { + output_data + .par_chunks_mut(m * n) + .enumerate() + .for_each(|(b, chunk)| { + let a = &lhs_data[b * m * k..(b + 1) * m * k]; + let r = &rhs_data[b * k * n..(b + 1) * k * n]; + unsafe { + gemm_f64(m, k, n, a.as_ptr(), r.as_ptr(), chunk.as_mut_ptr()); + } + }); + } + + Ok(()) +} + +// Helper functions for transpose operations + +fn transpose_f32( + tensor: &Tensor, + output_data: &mut TensorData, + output_shape: &Shape, + dim0: usize, + dim1: usize, +) -> Result<()> { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from input tensor") + })?; + + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output data") + })?; + + transpose_generic( + input_data, + output_slice, + tensor.shape(), + output_shape, + dim0, + dim1, + ) +} + +fn transpose_f64( + tensor: &Tensor, + output_data: &mut TensorData, + output_shape: &Shape, + dim0: usize, + dim1: usize, +) -> Result<()> { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from input tensor") + })?; + + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output data") + })?; + + transpose_generic( + input_data, + output_slice, + tensor.shape(), + output_shape, + dim0, + dim1, + ) +} + +fn transpose_i32( + tensor: &Tensor, + output_data: &mut TensorData, + output_shape: &Shape, + dim0: usize, + dim1: usize, +) -> Result<()> { + let input_data = tensor.data().as_i32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i32 slice from input tensor") + })?; + + let output_slice = output_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i32 slice from output data") + })?; + + transpose_generic( + input_data, + output_slice, + tensor.shape(), + output_shape, + dim0, + dim1, + ) +} + +fn transpose_i64( + tensor: &Tensor, + output_data: &mut TensorData, + output_shape: &Shape, + dim0: usize, + dim1: usize, +) -> Result<()> { + let input_data = tensor.data().as_i64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i64 slice from input tensor") + })?; + + let output_slice = output_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice from output data") + })?; + + transpose_generic( + input_data, + output_slice, + tensor.shape(), + output_shape, + dim0, + dim1, + ) +} + +fn transpose_bool( + tensor: &Tensor, + output_data: &mut TensorData, + output_shape: &Shape, + dim0: usize, + dim1: usize, +) -> Result<()> { + let input_data = tensor.data().as_bool_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get bool slice from input tensor") + })?; + + let output_slice = output_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice from output data") + })?; + + transpose_generic( + input_data, + output_slice, + tensor.shape(), + output_shape, + dim0, + dim1, + ) +} diff --git a/engine/src/operations/linalg/matmul.rs b/engine/src/operations/linalg/matmul.rs index 659b8411..855ddaf7 100644 --- a/engine/src/operations/linalg/matmul.rs +++ b/engine/src/operations/linalg/matmul.rs @@ -1,914 +1,941 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -use crate::{ - autograd::{DotBackward, MatMulBackward, SolveBackward, TransposeBackward, add_to_graph}, - error::{MinitensorError, Result}, - operations::{ - binary::{BinaryOpKind, coerce_binary_operands}, - reduction, - }, - tensor::{DataType, Shape, Strides, Tensor, TensorData}, -}; -use rayon::prelude::*; -use std::sync::Arc; - -#[cfg(feature = "blas")] -use cblas::{Layout, Transpose}; - -const PAR_THRESHOLD: usize = 1 << 12; - -#[derive(Debug, Clone)] -pub(crate) struct DiagonalSpec { - pub diag_len: usize, - pub base_offset: usize, - pub diag_stride: usize, - pub kept_dims: Vec, - pub output_dims: Vec, -} - -fn normalize_dim(dim: isize, ndim: usize) -> Result { - let dim = if dim < 0 { dim + ndim as isize } else { dim }; - if dim < 0 || dim >= ndim as isize { - Err(MinitensorError::index_error(dim, 0, ndim)) - } else { - Ok(dim as usize) - } -} - -pub(crate) fn compute_diagonal_spec( - dims: &[usize], - strides: &[usize], - dim1: usize, - dim2: usize, - offset: isize, -) -> Result { - debug_assert!(dim1 != dim2); - - let dim1_size = dims - .get(dim1) - .ok_or_else(|| MinitensorError::index_error(dim1 as isize, 0, dims.len()))?; - let dim2_size = dims - .get(dim2) - .ok_or_else(|| MinitensorError::index_error(dim2 as isize, 0, dims.len()))?; - let stride1 = strides - .get(dim1) - .ok_or_else(|| MinitensorError::index_error(dim1 as isize, 0, strides.len()))?; - let stride2 = strides - .get(dim2) - .ok_or_else(|| MinitensorError::index_error(dim2 as isize, 0, strides.len()))?; - - let diag_stride = stride1.saturating_add(*stride2); - - let (diag_len, base_offset) = if offset >= 0 { - let offset = offset as usize; - if offset >= *dim2_size { - (0, 0) - } else { - ( - (*dim1_size).min(dim2_size - offset), - offset.saturating_mul(*stride2), - ) - } - } else { - let neg = (-offset) as usize; - if neg >= *dim1_size { - (0, 0) - } else { - ( - (dim1_size - neg).min(*dim2_size), - neg.saturating_mul(*stride1), - ) - } - }; - - let mut kept_dims = Vec::with_capacity(dims.len().saturating_sub(2)); - let mut output_dims = Vec::with_capacity(kept_dims.capacity() + 1); - for (idx, &size) in dims.iter().enumerate() { - if idx == dim1 || idx == dim2 { - continue; - } - kept_dims.push(idx); - output_dims.push(size); - } - output_dims.push(diag_len); - - Ok(DiagonalSpec { - diag_len, - base_offset, - diag_stride, - kept_dims, - output_dims, - }) -} - -fn diagonal_copy( - input: &[T], - output: &mut [T], - dims: &[usize], - strides: &[usize], - spec: &DiagonalSpec, -) { - if output.is_empty() { - return; - } - - let mut axis_sizes: Vec = spec.kept_dims.iter().map(|&dim| dims[dim]).collect(); - axis_sizes.push(spec.diag_len); - - let mut axis_strides: Vec = spec.kept_dims.iter().map(|&dim| strides[dim]).collect(); - axis_strides.push(spec.diag_stride); - - let axes = axis_sizes.len(); - let mut indices = vec![0usize; axes]; - let mut out_idx = 0usize; - - loop { - let mut input_offset = spec.base_offset; - for axis in 0..axes { - input_offset += indices[axis] * axis_strides[axis]; - } - output[out_idx] = input[input_offset]; - out_idx += 1; - - let mut done = true; - for axis in (0..axes).rev() { - indices[axis] += 1; - if indices[axis] < axis_sizes[axis] { - done = false; - break; - } - indices[axis] = 0; - } - if done { - break; - } - } -} - -pub(crate) fn diagonal_scatter( - grad_output: &[T], - grad_input: &mut [T], - dims: &[usize], - strides: &[usize], - spec: &DiagonalSpec, -) where - T: Copy + Send + Sync + std::ops::AddAssign, -{ - if grad_output.is_empty() { - return; - } - - let mut axis_sizes: Vec = spec.kept_dims.iter().map(|&dim| dims[dim]).collect(); - axis_sizes.push(spec.diag_len); - - let mut axis_strides: Vec = spec.kept_dims.iter().map(|&dim| strides[dim]).collect(); - axis_strides.push(spec.diag_stride); - - let axes = axis_sizes.len(); - let mut indices = vec![0usize; axes]; - let mut out_idx = 0usize; - - loop { - let mut input_offset = spec.base_offset; - for axis in 0..axes { - input_offset += indices[axis] * axis_strides[axis]; - } - grad_input[input_offset] += grad_output[out_idx]; - out_idx += 1; - - let mut done = true; - for axis in (0..axes).rev() { - indices[axis] += 1; - if indices[axis] < axis_sizes[axis] { - done = false; - break; - } - indices[axis] = 0; - } - if done { - break; - } - } -} - -#[cfg(feature = "blas")] -#[inline] -unsafe fn gemm_f32(m: usize, k: usize, n: usize, a: *const f32, b: *const f32, c: *mut f32) { - cblas::sgemm( - Layout::RowMajor, - Transpose::None, - Transpose::None, - m as i32, - n as i32, - k as i32, - 1.0, - a, - k as i32, - b, - n as i32, - 0.0, - c, - n as i32, - ); -} - -#[cfg(feature = "blas")] -#[inline] -unsafe fn gemm_f64(m: usize, k: usize, n: usize, a: *const f64, b: *const f64, c: *mut f64) { - cblas::dgemm( - Layout::RowMajor, - Transpose::None, - Transpose::None, - m as i32, - n as i32, - k as i32, - 1.0, - a, - k as i32, - b, - n as i32, - 0.0, - c, - n as i32, - ); -} - -#[cfg(not(feature = "blas"))] -#[inline] -unsafe fn gemm_f32(m: usize, k: usize, n: usize, a: *const f32, b: *const f32, c: *mut f32) { - unsafe { - matrixmultiply::sgemm( - m, k, n, 1.0, a, k as isize, 1, b, n as isize, 1, 0.0, c, n as isize, 1, - ) - }; -} - -#[cfg(not(feature = "blas"))] -#[inline] -unsafe fn gemm_f64(m: usize, k: usize, n: usize, a: *const f64, b: *const f64, c: *mut f64) { - unsafe { - matrixmultiply::dgemm( - m, k, n, 1.0, a, k as isize, 1, b, n as isize, 1, 0.0, c, n as isize, 1, - ) - }; -} - -/// Matrix multiplication with gradient support -pub fn matmul(lhs: &Tensor, rhs: &Tensor) -> Result { - // Check device compatibility - if lhs.device() != rhs.device() { - return Err(MinitensorError::device_mismatch( - format!("{:?}", lhs.device()), - format!("{:?}", rhs.device()), - )); - } - - // Check data type compatibility - if lhs.dtype() != rhs.dtype() { - return Err(MinitensorError::type_mismatch( - format!("{:?}", lhs.dtype()), - format!("{:?}", rhs.dtype()), - )); - } - - // For 1-D vectors, `lhs` is promoted by prepending a 1 and `rhs` by appending - // a 1; the added axes are removed from the result - // (so mat@vec -> vec, vec@mat -> vec, vec@vec -> scalar). Reshapes are - // grad-aware, so the gradient flows through the promotion. - let lhs_1d = lhs.ndim() == 1; - let rhs_1d = rhs.ndim() == 1; - if lhs_1d || rhs_1d { - use crate::operations::shape_ops::reshape; - let lhs2 = if lhs_1d { - reshape(lhs, Shape::new(vec![1, lhs.shape().dims()[0]]))? - } else { - lhs.clone() - }; - let rhs2 = if rhs_1d { - reshape(rhs, Shape::new(vec![rhs.shape().dims()[0], 1]))? - } else { - rhs.clone() - }; - let promoted = matmul(&lhs2, &rhs2)?; - // Drop the promoted axes (remove the trailing column before the leading - // row so the earlier index stays valid). - let mut dims = promoted.shape().dims().to_vec(); - let len = dims.len(); - if rhs_1d { - dims.remove(len - 1); - } - if lhs_1d { - dims.remove(len - 2); - } - return reshape(&promoted, Shape::new(dims)); - } - - // Validate matrix multiplication dimensions - if lhs.ndim() < 2 || rhs.ndim() < 2 { - return Err(MinitensorError::invalid_operation( - "Matrix multiplication requires tensors with at least 1 dimension (scalars are not valid operands)", - )); - } - - let lhs_shape = lhs.shape().dims(); - let rhs_shape = rhs.shape().dims(); - - // Broadcast batch dimensions when they differ, e.g. - // [2, 3, 4] @ [4, 5] or [1, 3, 4] @ [7, 4, 5]. The expanded operands are - // materialized contiguously; expand/contiguous are grad-aware so the - // gradient reduces back over the broadcast batch dimensions. - if lhs_shape[..lhs_shape.len() - 2] != rhs_shape[..rhs_shape.len() - 2] { - let lhs_batch = Shape::new(lhs_shape[..lhs_shape.len() - 2].to_vec()); - let rhs_batch = Shape::new(rhs_shape[..rhs_shape.len() - 2].to_vec()); - let batch = lhs_batch.broadcast_with(&rhs_batch).map_err(|_| { - MinitensorError::shape_mismatch(lhs_shape.to_vec(), rhs_shape.to_vec()) - })?; - - let mut lhs_target: Vec = batch.dims().iter().map(|&d| d as isize).collect(); - lhs_target.extend_from_slice(&[ - lhs_shape[lhs_shape.len() - 2] as isize, - lhs_shape[lhs_shape.len() - 1] as isize, - ]); - let mut rhs_target: Vec = batch.dims().iter().map(|&d| d as isize).collect(); - rhs_target.extend_from_slice(&[ - rhs_shape[rhs_shape.len() - 2] as isize, - rhs_shape[rhs_shape.len() - 1] as isize, - ]); - - let lhs_b = lhs.expand(lhs_target)?.contiguous()?; - let rhs_b = rhs.expand(rhs_target)?.contiguous()?; - return matmul(&lhs_b, &rhs_b); - } - - // Get the last two dimensions for matrix multiplication - let lhs_rows = lhs_shape[lhs_shape.len() - 2]; - let lhs_cols = lhs_shape[lhs_shape.len() - 1]; - let rhs_rows = rhs_shape[rhs_shape.len() - 2]; - let rhs_cols = rhs_shape[rhs_shape.len() - 1]; - - if lhs_cols != rhs_rows { - return Err(MinitensorError::shape_mismatch( - vec![lhs_rows, lhs_cols], - vec![rhs_rows, rhs_cols], - )); - } - - // Compute output shape - let mut output_shape = lhs_shape[..lhs_shape.len() - 2].to_vec(); - output_shape.push(lhs_rows); - output_shape.push(rhs_cols); - let output_shape_obj = Shape::new(output_shape); - - if lhs.dtype() == DataType::Bool { - return Err(MinitensorError::invalid_operation( - "Matrix multiplication not supported for boolean tensors", - )); - } - - // Create output tensor data - let mut output_data = - TensorData::zeros_on_device(output_shape_obj.numel(), lhs.dtype(), lhs.device()); - - if output_shape_obj.numel() != 0 && lhs_cols != 0 { - // Perform matrix multiplication based on data type - match lhs.dtype() { - DataType::Float32 => matmul_f32(lhs, rhs, &mut output_data, &output_shape_obj)?, - DataType::Float64 => matmul_f64(lhs, rhs, &mut output_data, &output_shape_obj)?, - DataType::Int32 => matmul_i32(lhs, rhs, &mut output_data, &output_shape_obj)?, - DataType::Int64 => matmul_i64(lhs, rhs, &mut output_data, &output_shape_obj)?, - DataType::Bool => unreachable!("bool dtype checked above"), - } - } - - // Create output tensor - let output = Tensor::new( - Arc::new(output_data), - output_shape_obj, - lhs.dtype(), - lhs.device(), - lhs.requires_grad() || rhs.requires_grad(), - ); - - // Set up gradient function if needed - if output.requires_grad() { - let grad_fn = Arc::new(MatMulBackward { - lhs: lhs.detach(), - rhs: rhs.detach(), - input_ids: [lhs.id(), rhs.id()], - lhs_requires_grad: lhs.requires_grad(), - rhs_requires_grad: rhs.requires_grad(), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&output_with_grad, Some(grad_fn))?; - - Ok(output_with_grad) - } else { - Ok(output) - } -} - -/// Solve a linear system of equations `AX = B` for `X`. -/// -/// Both `lhs` (`A`) and `rhs` (`B`) must be float tensors that live on the CPU. -/// `lhs` must have shape `[..., n, n]` (square matrices) and `rhs` can either have -/// shape `[..., n]` (a collection of vectors) or `[..., n, k]` (multiple right -/// hand sides). Batch dimensions need to match exactly across the operands. -pub fn solve(lhs: &Tensor, rhs: &Tensor) -> Result { - if lhs.device() != rhs.device() { - return Err(MinitensorError::device_mismatch( - format!("{:?}", lhs.device()), - format!("{:?}", rhs.device()), - )); - } - - if lhs.dtype() != rhs.dtype() { - return Err(MinitensorError::type_mismatch( - format!("{:?}", lhs.dtype()), - format!("{:?}", rhs.dtype()), - )); - } - - let lhs_ndim = lhs.ndim(); - if lhs_ndim < 2 { - return Err(MinitensorError::invalid_operation( - "solve expects lhs to have at least 2 dimensions", - )); - } - - let lhs_shape = lhs.shape().dims(); - let n = lhs_shape[lhs_ndim - 1]; - let m = lhs_shape[lhs_ndim - 2]; - if n != m { - return Err(MinitensorError::invalid_operation( - "solve expects lhs matrices to be square", - )); - } - - let rhs_ndim = rhs.ndim(); - if rhs_ndim < 1 { - return Err(MinitensorError::invalid_operation( - "solve expects rhs to have at least 1 dimension", - )); - } - - let rhs_shape = rhs.shape().dims(); - let (rhs_cols, rhs_batch_dims) = if rhs_ndim == lhs_ndim { - if rhs_shape[rhs_ndim - 2] != n { - return Err(MinitensorError::shape_mismatch( - vec![n], - vec![rhs_shape[rhs_ndim - 2]], - )); - } - (rhs_shape[rhs_ndim - 1], &rhs_shape[..rhs_ndim - 2]) - } else if rhs_ndim + 1 == lhs_ndim { - if rhs_shape[rhs_ndim - 1] != n { - return Err(MinitensorError::shape_mismatch( - vec![n], - vec![rhs_shape[rhs_ndim - 1]], - )); - } - (1usize, &rhs_shape[..rhs_ndim - 1]) - } else { - return Err(MinitensorError::invalid_operation( - "solve expects rhs to have either the same rank as lhs or one less", - )); - }; - - if &lhs_shape[..lhs_ndim - 2] != rhs_batch_dims { - return Err(MinitensorError::shape_mismatch( - lhs_shape[..lhs_ndim - 2].to_vec(), - rhs_batch_dims.to_vec(), - )); - } - - let requires_grad = lhs.requires_grad() || rhs.requires_grad(); - - let output_shape = rhs_shape.to_vec(); - let output_shape = Shape::new(output_shape); - - let mut output_data = - TensorData::zeros_on_device(output_shape.numel(), lhs.dtype(), lhs.device()); - - match lhs.dtype() { - DataType::Float32 => solve_f32(lhs, rhs, &mut output_data, rhs_cols)?, - DataType::Float64 => solve_f64(lhs, rhs, &mut output_data, rhs_cols)?, - _ => { - return Err(MinitensorError::invalid_operation( - "solve currently supports only Float32 and Float64 tensors", - )); - } - } - - let mut output = Tensor::new( - Arc::new(output_data), - output_shape, - lhs.dtype(), - lhs.device(), - requires_grad, - ); - - if output.requires_grad() { - let grad_fn = Arc::new(SolveBackward { - lhs: lhs.detach(), - solution: output.detach(), - input_ids: [lhs.id(), rhs.id()], - lhs_requires_grad: lhs.requires_grad(), - rhs_requires_grad: rhs.requires_grad(), - }); - output.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output, Some(grad_fn))?; - } - - Ok(output) -} - -fn solve_f32(lhs: &Tensor, rhs: &Tensor, output: &mut TensorData, rhs_cols: usize) -> Result<()> { - use std::borrow::Cow; - - let lhs_view = if lhs.is_contiguous() && lhs.data().is_contiguous() { - Cow::Borrowed(lhs) - } else { - Cow::Owned(lhs.contiguous()?) - }; - let rhs_view = if rhs.is_contiguous() && rhs.data().is_contiguous() { - Cow::Borrowed(rhs) - } else { - Cow::Owned(rhs.contiguous()?) - }; - - let lhs_slice = lhs_view - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to access f32 data for lhs"))?; - let rhs_slice = rhs_view - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to access f32 data for rhs"))?; - let out_slice = output - .as_f32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to access f32 output slice"))?; - - solve_batched( - lhs.shape().dims(), - rhs_cols, - lhs_slice, - rhs_slice, - out_slice, - ) -} - -fn solve_f64(lhs: &Tensor, rhs: &Tensor, output: &mut TensorData, rhs_cols: usize) -> Result<()> { - use std::borrow::Cow; - - let lhs_view = if lhs.is_contiguous() && lhs.data().is_contiguous() { - Cow::Borrowed(lhs) - } else { - Cow::Owned(lhs.contiguous()?) - }; - let rhs_view = if rhs.is_contiguous() && rhs.data().is_contiguous() { - Cow::Borrowed(rhs) - } else { - Cow::Owned(rhs.contiguous()?) - }; - - let lhs_slice = lhs_view - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to access f64 data for lhs"))?; - let rhs_slice = rhs_view - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to access f64 data for rhs"))?; - let out_slice = output - .as_f64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to access f64 output slice"))?; - - solve_batched( - lhs.shape().dims(), - rhs_cols, - lhs_slice, - rhs_slice, - out_slice, - ) -} - -fn solve_batched( - lhs_shape: &[usize], - rhs_cols: usize, - lhs_slice: &[T], - rhs_slice: &[T], - out_slice: &mut [T], -) -> Result<()> -where - T: Copy - + Send - + Sync - + std::ops::SubAssign - + std::ops::Mul - + std::ops::Div - + std::ops::Neg - + PartialOrd - + Default - + PartialEq, -{ - let n = *lhs_shape.last().expect("lhs has at least 2 dims"); - let batch = lhs_shape[..lhs_shape.len() - 2] - .iter() - .copied() - .product::() - .max(1); - let rhs_stride = n * rhs_cols; - - let matrix_stride = n * n; - let mut matrix = vec![T::default(); matrix_stride]; - let mut rhs_buf = vec![T::default(); rhs_stride]; - - for batch_idx in 0..batch { - let lhs_offset = batch_idx * matrix_stride; - let rhs_offset = batch_idx * rhs_stride; - - matrix.copy_from_slice(&lhs_slice[lhs_offset..lhs_offset + matrix_stride]); - rhs_buf[..rhs_stride].copy_from_slice(&rhs_slice[rhs_offset..rhs_offset + rhs_stride]); - - gaussian_elimination(&mut matrix, &mut rhs_buf, n, rhs_cols)?; - - out_slice[rhs_offset..rhs_offset + rhs_stride].copy_from_slice(&rhs_buf[..rhs_stride]); - } - - Ok(()) -} - -fn gaussian_elimination(matrix: &mut [T], rhs: &mut [T], n: usize, rhs_cols: usize) -> Result<()> -where - T: Copy - + Send - + Sync - + std::ops::SubAssign - + std::ops::Mul - + std::ops::Div - + std::ops::Neg - + PartialOrd - + Default - + PartialEq, -{ - for k in 0..n { - // Pivot selection - let mut pivot_row = k; - let mut pivot_val = abs(matrix[k * n + k]); - for i in (k + 1)..n { - let candidate = abs(matrix[i * n + k]); - if candidate > pivot_val { - pivot_val = candidate; - pivot_row = i; - } - } - - if pivot_val == T::default() { - return Err(MinitensorError::invalid_operation( - "solve received a singular matrix", - )); - } - - if pivot_row != k { - for col in 0..n { - matrix.swap(k * n + col, pivot_row * n + col); - } - for col in 0..rhs_cols { - rhs.swap(k * rhs_cols + col, pivot_row * rhs_cols + col); - } - } - - let pivot = matrix[k * n + k]; - - for i in (k + 1)..n { - let factor = matrix[i * n + k] / pivot; - matrix[i * n + k] = T::default(); - for j in (k + 1)..n { - let idx = i * n + j; - matrix[idx] -= factor * matrix[k * n + j]; - } - for col in 0..rhs_cols { - let idx = i * rhs_cols + col; - rhs[idx] -= factor * rhs[k * rhs_cols + col]; - } - } - } - - for i in (0..n).rev() { - let pivot = matrix[i * n + i]; - if abs(pivot) == T::default() { - return Err(MinitensorError::invalid_operation( - "solve received a singular matrix", - )); - } - for col in 0..rhs_cols { - let mut value = rhs[i * rhs_cols + col]; - for j in (i + 1)..n { - value -= matrix[i * n + j] * rhs[j * rhs_cols + col]; - } - rhs[i * rhs_cols + col] = value / pivot; - } - } - - Ok(()) -} - -fn abs(value: T) -> T -where - T: Copy + PartialOrd + std::ops::Neg + Default, -{ - if value < T::default() { -value } else { value } -} - -/// Batched matrix multiplication specialized for 3D tensors. -/// -/// This is a thin convenience wrapper around [`matmul`] that enforces the -/// traditional batch matrix multiply constraints: both operands must be -/// rank-3 tensors with matching batch dimensions. The actual computation is -/// still delegated to the highly optimised [`matmul`] implementation so all -/// execution happens inside the Rust backend. -pub fn bmm(lhs: &Tensor, rhs: &Tensor) -> Result { - if lhs.ndim() != 3 || rhs.ndim() != 3 { - return Err(MinitensorError::invalid_operation( - "bmm expects both inputs to be 3D tensors".to_string(), - )); - } - - let lhs_shape = lhs.shape().dims(); - let rhs_shape = rhs.shape().dims(); - - if lhs_shape[0] != rhs_shape[0] { - return Err(MinitensorError::shape_mismatch( - lhs_shape.to_vec(), - rhs_shape.to_vec(), - )); - } - - if lhs_shape[2] != rhs_shape[1] { - return Err(MinitensorError::shape_mismatch( - vec![lhs_shape[2]], - vec![rhs_shape[1]], - )); - } - - matmul(lhs, rhs) -} - -/// Dot product of two 1D tensors with gradient support -pub fn dot(lhs: &Tensor, rhs: &Tensor) -> Result { - if lhs.device() != rhs.device() { - return Err(MinitensorError::device_mismatch( - format!("{:?}", lhs.device()), - format!("{:?}", rhs.device()), - )); - } - - let lhs_dims = lhs.ndim(); - let rhs_dims = rhs.ndim(); - if lhs_dims != 1 || rhs_dims != 1 { - return Err(MinitensorError::invalid_operation(format!( - "dot: expected 1D tensors but got {}D and {}D tensors", - lhs_dims, rhs_dims - ))); - } - - if lhs.numel() != rhs.numel() { - return Err(MinitensorError::shape_mismatch( - lhs.shape().dims().to_vec(), - rhs.shape().dims().to_vec(), - )); - } - - let (lhs_cast, rhs_cast, result_dtype) = coerce_binary_operands(lhs, rhs, BinaryOpKind::Mul)?; - - if result_dtype == DataType::Bool { - return Err(MinitensorError::invalid_operation( - "dot does not support bool tensors", - )); - } - - let lhs_view = lhs_cast.as_ref(); - let rhs_view = rhs_cast.as_ref(); - - let numel = lhs_view.numel(); - let device = lhs.device(); - let requires_grad = lhs.requires_grad() || rhs.requires_grad(); - - let output_data = match result_dtype { - DataType::Float32 => { - let lhs_slice = lhs_view.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice for dot input") - })?; - let rhs_slice = rhs_view.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice for dot input") - })?; - - let dot = if numel >= PAR_THRESHOLD { - lhs_slice - .par_iter() - .zip(rhs_slice.par_iter()) - .map(|(&a, &b)| a * b) - .sum::() - } else { - lhs_slice - .iter() - .zip(rhs_slice.iter()) - .map(|(&a, &b)| a * b) - .sum::() - }; - - TensorData::from_vec_f32(vec![dot], device) - } - DataType::Float64 => { - let lhs_slice = lhs_view.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice for dot input") - })?; - let rhs_slice = rhs_view.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice for dot input") - })?; - - let dot = if numel >= PAR_THRESHOLD { - lhs_slice - .par_iter() - .zip(rhs_slice.par_iter()) - .map(|(&a, &b)| a * b) - .sum::() - } else { - lhs_slice - .iter() - .zip(rhs_slice.iter()) - .map(|(&a, &b)| a * b) - .sum::() - }; - - TensorData::from_vec_f64(vec![dot], device) - } - DataType::Int32 => { - let lhs_slice = lhs_view.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice for dot input") - })?; - let rhs_slice = rhs_view.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice for dot input") - })?; - - let mut dot: i32 = 0; - for (&a, &b) in lhs_slice.iter().zip(rhs_slice.iter()) { - dot = dot.wrapping_add(a.wrapping_mul(b)); - } - - TensorData::from_vec_i32(vec![dot], device) - } - DataType::Int64 => { - let lhs_slice = lhs_view.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice for dot input") - })?; - let rhs_slice = rhs_view.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice for dot input") - })?; - - let mut dot: i64 = 0; - for (&a, &b) in lhs_slice.iter().zip(rhs_slice.iter()) { - dot = dot.wrapping_add(a.wrapping_mul(b)); - } - - TensorData::from_vec_i64(vec![dot], device) - } - DataType::Bool => unreachable!("Bool dtype handled earlier"), - }; - - let output_shape = Shape::new(Vec::new()); - let output = Tensor::new( - Arc::new(output_data), - output_shape, - result_dtype, - device, - requires_grad, - ); - - if output.requires_grad() { - let lhs_requires_grad = lhs.requires_grad(); - let rhs_requires_grad = rhs.requires_grad(); - let grad_fn = Arc::new(DotBackward { - lhs: lhs_cast.into_owned().detach(), - rhs: rhs_cast.into_owned().detach(), - input_ids: [lhs.id(), rhs.id()], - lhs_requires_grad, - rhs_requires_grad, - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) - } else { - Ok(output) - } -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; + +use crate::{ + autograd::{DotBackward, MatMulBackward, SolveBackward, add_to_graph}, + error::{MinitensorError, Result}, + operations::binary::{BinaryOpKind, coerce_binary_operands}, + tensor::{DataType, Shape, Tensor, TensorData}, +}; +use rayon::prelude::*; +use std::sync::Arc; + +#[cfg(feature = "blas")] +use cblas::{Layout, Transpose}; + +pub(crate) const PAR_THRESHOLD: usize = 1 << 12; + +#[derive(Debug, Clone)] +pub(crate) struct DiagonalSpec { + pub diag_len: usize, + pub base_offset: usize, + pub diag_stride: usize, + pub kept_dims: Vec, + pub output_dims: Vec, +} + +pub(crate) fn normalize_dim(dim: isize, ndim: usize) -> Result { + let dim = if dim < 0 { dim + ndim as isize } else { dim }; + if dim < 0 || dim >= ndim as isize { + Err(MinitensorError::index_error(dim, 0, ndim)) + } else { + Ok(dim as usize) + } +} + +pub(crate) fn compute_diagonal_spec( + dims: &[usize], + strides: &[usize], + dim1: usize, + dim2: usize, + offset: isize, +) -> Result { + debug_assert!(dim1 != dim2); + + let dim1_size = dims + .get(dim1) + .ok_or_else(|| MinitensorError::index_error(dim1 as isize, 0, dims.len()))?; + let dim2_size = dims + .get(dim2) + .ok_or_else(|| MinitensorError::index_error(dim2 as isize, 0, dims.len()))?; + let stride1 = strides + .get(dim1) + .ok_or_else(|| MinitensorError::index_error(dim1 as isize, 0, strides.len()))?; + let stride2 = strides + .get(dim2) + .ok_or_else(|| MinitensorError::index_error(dim2 as isize, 0, strides.len()))?; + + let diag_stride = stride1.saturating_add(*stride2); + + let (diag_len, base_offset) = if offset >= 0 { + let offset = offset as usize; + if offset >= *dim2_size { + (0, 0) + } else { + ( + (*dim1_size).min(dim2_size - offset), + offset.saturating_mul(*stride2), + ) + } + } else { + let neg = (-offset) as usize; + if neg >= *dim1_size { + (0, 0) + } else { + ( + (dim1_size - neg).min(*dim2_size), + neg.saturating_mul(*stride1), + ) + } + }; + + let mut kept_dims = Vec::with_capacity(dims.len().saturating_sub(2)); + let mut output_dims = Vec::with_capacity(kept_dims.capacity() + 1); + for (idx, &size) in dims.iter().enumerate() { + if idx == dim1 || idx == dim2 { + continue; + } + kept_dims.push(idx); + output_dims.push(size); + } + output_dims.push(diag_len); + + Ok(DiagonalSpec { + diag_len, + base_offset, + diag_stride, + kept_dims, + output_dims, + }) +} + +pub(crate) fn diagonal_copy( + input: &[T], + output: &mut [T], + dims: &[usize], + strides: &[usize], + spec: &DiagonalSpec, +) { + if output.is_empty() { + return; + } + + let mut axis_sizes: Vec = spec.kept_dims.iter().map(|&dim| dims[dim]).collect(); + axis_sizes.push(spec.diag_len); + + let mut axis_strides: Vec = spec.kept_dims.iter().map(|&dim| strides[dim]).collect(); + axis_strides.push(spec.diag_stride); + + let axes = axis_sizes.len(); + let mut indices = vec![0usize; axes]; + let mut out_idx = 0usize; + + loop { + let mut input_offset = spec.base_offset; + for axis in 0..axes { + input_offset += indices[axis] * axis_strides[axis]; + } + output[out_idx] = input[input_offset]; + out_idx += 1; + + let mut done = true; + for axis in (0..axes).rev() { + indices[axis] += 1; + if indices[axis] < axis_sizes[axis] { + done = false; + break; + } + indices[axis] = 0; + } + if done { + break; + } + } +} + +pub(crate) fn diagonal_scatter( + grad_output: &[T], + grad_input: &mut [T], + dims: &[usize], + strides: &[usize], + spec: &DiagonalSpec, +) where + T: Copy + Send + Sync + std::ops::AddAssign, +{ + if grad_output.is_empty() { + return; + } + + let mut axis_sizes: Vec = spec.kept_dims.iter().map(|&dim| dims[dim]).collect(); + axis_sizes.push(spec.diag_len); + + let mut axis_strides: Vec = spec.kept_dims.iter().map(|&dim| strides[dim]).collect(); + axis_strides.push(spec.diag_stride); + + let axes = axis_sizes.len(); + let mut indices = vec![0usize; axes]; + let mut out_idx = 0usize; + + loop { + let mut input_offset = spec.base_offset; + for axis in 0..axes { + input_offset += indices[axis] * axis_strides[axis]; + } + grad_input[input_offset] += grad_output[out_idx]; + out_idx += 1; + + let mut done = true; + for axis in (0..axes).rev() { + indices[axis] += 1; + if indices[axis] < axis_sizes[axis] { + done = false; + break; + } + indices[axis] = 0; + } + if done { + break; + } + } +} + +#[cfg(feature = "blas")] +#[inline] +pub(crate) unsafe fn gemm_f32( + m: usize, + k: usize, + n: usize, + a: *const f32, + b: *const f32, + c: *mut f32, +) { + cblas::sgemm( + Layout::RowMajor, + Transpose::None, + Transpose::None, + m as i32, + n as i32, + k as i32, + 1.0, + a, + k as i32, + b, + n as i32, + 0.0, + c, + n as i32, + ); +} + +#[cfg(feature = "blas")] +#[inline] +pub(crate) unsafe fn gemm_f64( + m: usize, + k: usize, + n: usize, + a: *const f64, + b: *const f64, + c: *mut f64, +) { + cblas::dgemm( + Layout::RowMajor, + Transpose::None, + Transpose::None, + m as i32, + n as i32, + k as i32, + 1.0, + a, + k as i32, + b, + n as i32, + 0.0, + c, + n as i32, + ); +} + +#[cfg(not(feature = "blas"))] +#[inline] +pub(crate) unsafe fn gemm_f32( + m: usize, + k: usize, + n: usize, + a: *const f32, + b: *const f32, + c: *mut f32, +) { + unsafe { + matrixmultiply::sgemm( + m, k, n, 1.0, a, k as isize, 1, b, n as isize, 1, 0.0, c, n as isize, 1, + ) + }; +} + +#[cfg(not(feature = "blas"))] +#[inline] +pub(crate) unsafe fn gemm_f64( + m: usize, + k: usize, + n: usize, + a: *const f64, + b: *const f64, + c: *mut f64, +) { + unsafe { + matrixmultiply::dgemm( + m, k, n, 1.0, a, k as isize, 1, b, n as isize, 1, 0.0, c, n as isize, 1, + ) + }; +} + +/// Matrix multiplication with gradient support +pub fn matmul(lhs: &Tensor, rhs: &Tensor) -> Result { + // Check device compatibility + if lhs.device() != rhs.device() { + return Err(MinitensorError::device_mismatch( + format!("{:?}", lhs.device()), + format!("{:?}", rhs.device()), + )); + } + + // Check data type compatibility + if lhs.dtype() != rhs.dtype() { + return Err(MinitensorError::type_mismatch( + format!("{:?}", lhs.dtype()), + format!("{:?}", rhs.dtype()), + )); + } + + // For 1-D vectors, `lhs` is promoted by prepending a 1 and `rhs` by appending + // a 1; the added axes are removed from the result + // (so mat@vec -> vec, vec@mat -> vec, vec@vec -> scalar). Reshapes are + // grad-aware, so the gradient flows through the promotion. + let lhs_1d = lhs.ndim() == 1; + let rhs_1d = rhs.ndim() == 1; + if lhs_1d || rhs_1d { + use crate::operations::shape_ops::reshape; + let lhs2 = if lhs_1d { + reshape(lhs, Shape::new(vec![1, lhs.shape().dims()[0]]))? + } else { + lhs.clone() + }; + let rhs2 = if rhs_1d { + reshape(rhs, Shape::new(vec![rhs.shape().dims()[0], 1]))? + } else { + rhs.clone() + }; + let promoted = matmul(&lhs2, &rhs2)?; + // Drop the promoted axes (remove the trailing column before the leading + // row so the earlier index stays valid). + let mut dims = promoted.shape().dims().to_vec(); + let len = dims.len(); + if rhs_1d { + dims.remove(len - 1); + } + if lhs_1d { + dims.remove(len - 2); + } + return reshape(&promoted, Shape::new(dims)); + } + + // Validate matrix multiplication dimensions + if lhs.ndim() < 2 || rhs.ndim() < 2 { + return Err(MinitensorError::invalid_operation( + "Matrix multiplication requires tensors with at least 1 dimension (scalars are not valid operands)", + )); + } + + let lhs_shape = lhs.shape().dims(); + let rhs_shape = rhs.shape().dims(); + + // Broadcast batch dimensions when they differ, e.g. + // [2, 3, 4] @ [4, 5] or [1, 3, 4] @ [7, 4, 5]. The expanded operands are + // materialized contiguously; expand/contiguous are grad-aware so the + // gradient reduces back over the broadcast batch dimensions. + if lhs_shape[..lhs_shape.len() - 2] != rhs_shape[..rhs_shape.len() - 2] { + let lhs_batch = Shape::new(lhs_shape[..lhs_shape.len() - 2].to_vec()); + let rhs_batch = Shape::new(rhs_shape[..rhs_shape.len() - 2].to_vec()); + let batch = lhs_batch + .broadcast_with(&rhs_batch) + .map_err(|_| MinitensorError::shape_mismatch(lhs_shape.to_vec(), rhs_shape.to_vec()))?; + + let mut lhs_target: Vec = batch.dims().iter().map(|&d| d as isize).collect(); + lhs_target.extend_from_slice(&[ + lhs_shape[lhs_shape.len() - 2] as isize, + lhs_shape[lhs_shape.len() - 1] as isize, + ]); + let mut rhs_target: Vec = batch.dims().iter().map(|&d| d as isize).collect(); + rhs_target.extend_from_slice(&[ + rhs_shape[rhs_shape.len() - 2] as isize, + rhs_shape[rhs_shape.len() - 1] as isize, + ]); + + let lhs_b = lhs.expand(lhs_target)?.contiguous()?; + let rhs_b = rhs.expand(rhs_target)?.contiguous()?; + return matmul(&lhs_b, &rhs_b); + } + + // Get the last two dimensions for matrix multiplication + let lhs_rows = lhs_shape[lhs_shape.len() - 2]; + let lhs_cols = lhs_shape[lhs_shape.len() - 1]; + let rhs_rows = rhs_shape[rhs_shape.len() - 2]; + let rhs_cols = rhs_shape[rhs_shape.len() - 1]; + + if lhs_cols != rhs_rows { + return Err(MinitensorError::shape_mismatch( + vec![lhs_rows, lhs_cols], + vec![rhs_rows, rhs_cols], + )); + } + + // Compute output shape + let mut output_shape = lhs_shape[..lhs_shape.len() - 2].to_vec(); + output_shape.push(lhs_rows); + output_shape.push(rhs_cols); + let output_shape_obj = Shape::new(output_shape); + + if lhs.dtype() == DataType::Bool { + return Err(MinitensorError::invalid_operation( + "Matrix multiplication not supported for boolean tensors", + )); + } + + // Create output tensor data + let mut output_data = + TensorData::zeros_on_device(output_shape_obj.numel(), lhs.dtype(), lhs.device()); + + if output_shape_obj.numel() != 0 && lhs_cols != 0 { + // Perform matrix multiplication based on data type + match lhs.dtype() { + DataType::Float32 => matmul_f32(lhs, rhs, &mut output_data, &output_shape_obj)?, + DataType::Float64 => matmul_f64(lhs, rhs, &mut output_data, &output_shape_obj)?, + DataType::Int32 => matmul_i32(lhs, rhs, &mut output_data, &output_shape_obj)?, + DataType::Int64 => matmul_i64(lhs, rhs, &mut output_data, &output_shape_obj)?, + DataType::Bool => unreachable!("bool dtype checked above"), + } + } + + // Create output tensor + let output = Tensor::new( + Arc::new(output_data), + output_shape_obj, + lhs.dtype(), + lhs.device(), + lhs.requires_grad() || rhs.requires_grad(), + ); + + // Set up gradient function if needed + if output.requires_grad() { + let grad_fn = Arc::new(MatMulBackward { + lhs: lhs.detach(), + rhs: rhs.detach(), + input_ids: [lhs.id(), rhs.id()], + lhs_requires_grad: lhs.requires_grad(), + rhs_requires_grad: rhs.requires_grad(), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&output_with_grad, Some(grad_fn))?; + + Ok(output_with_grad) + } else { + Ok(output) + } +} + +/// Solve a linear system of equations `AX = B` for `X`. +/// +/// Both `lhs` (`A`) and `rhs` (`B`) must be float tensors that live on the CPU. +/// `lhs` must have shape `[..., n, n]` (square matrices) and `rhs` can either have +/// shape `[..., n]` (a collection of vectors) or `[..., n, k]` (multiple right +/// hand sides). Batch dimensions need to match exactly across the operands. +pub fn solve(lhs: &Tensor, rhs: &Tensor) -> Result { + if lhs.device() != rhs.device() { + return Err(MinitensorError::device_mismatch( + format!("{:?}", lhs.device()), + format!("{:?}", rhs.device()), + )); + } + + if lhs.dtype() != rhs.dtype() { + return Err(MinitensorError::type_mismatch( + format!("{:?}", lhs.dtype()), + format!("{:?}", rhs.dtype()), + )); + } + + let lhs_ndim = lhs.ndim(); + if lhs_ndim < 2 { + return Err(MinitensorError::invalid_operation( + "solve expects lhs to have at least 2 dimensions", + )); + } + + let lhs_shape = lhs.shape().dims(); + let n = lhs_shape[lhs_ndim - 1]; + let m = lhs_shape[lhs_ndim - 2]; + if n != m { + return Err(MinitensorError::invalid_operation( + "solve expects lhs matrices to be square", + )); + } + + let rhs_ndim = rhs.ndim(); + if rhs_ndim < 1 { + return Err(MinitensorError::invalid_operation( + "solve expects rhs to have at least 1 dimension", + )); + } + + let rhs_shape = rhs.shape().dims(); + let (rhs_cols, rhs_batch_dims) = if rhs_ndim == lhs_ndim { + if rhs_shape[rhs_ndim - 2] != n { + return Err(MinitensorError::shape_mismatch( + vec![n], + vec![rhs_shape[rhs_ndim - 2]], + )); + } + (rhs_shape[rhs_ndim - 1], &rhs_shape[..rhs_ndim - 2]) + } else if rhs_ndim + 1 == lhs_ndim { + if rhs_shape[rhs_ndim - 1] != n { + return Err(MinitensorError::shape_mismatch( + vec![n], + vec![rhs_shape[rhs_ndim - 1]], + )); + } + (1usize, &rhs_shape[..rhs_ndim - 1]) + } else { + return Err(MinitensorError::invalid_operation( + "solve expects rhs to have either the same rank as lhs or one less", + )); + }; + + if &lhs_shape[..lhs_ndim - 2] != rhs_batch_dims { + return Err(MinitensorError::shape_mismatch( + lhs_shape[..lhs_ndim - 2].to_vec(), + rhs_batch_dims.to_vec(), + )); + } + + let requires_grad = lhs.requires_grad() || rhs.requires_grad(); + + let output_shape = rhs_shape.to_vec(); + let output_shape = Shape::new(output_shape); + + let mut output_data = + TensorData::zeros_on_device(output_shape.numel(), lhs.dtype(), lhs.device()); + + match lhs.dtype() { + DataType::Float32 => solve_f32(lhs, rhs, &mut output_data, rhs_cols)?, + DataType::Float64 => solve_f64(lhs, rhs, &mut output_data, rhs_cols)?, + _ => { + return Err(MinitensorError::invalid_operation( + "solve currently supports only Float32 and Float64 tensors", + )); + } + } + + let mut output = Tensor::new( + Arc::new(output_data), + output_shape, + lhs.dtype(), + lhs.device(), + requires_grad, + ); + + if output.requires_grad() { + let grad_fn = Arc::new(SolveBackward { + lhs: lhs.detach(), + solution: output.detach(), + input_ids: [lhs.id(), rhs.id()], + lhs_requires_grad: lhs.requires_grad(), + rhs_requires_grad: rhs.requires_grad(), + }); + output.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output, Some(grad_fn))?; + } + + Ok(output) +} + +fn solve_f32(lhs: &Tensor, rhs: &Tensor, output: &mut TensorData, rhs_cols: usize) -> Result<()> { + use std::borrow::Cow; + + let lhs_view = if lhs.is_contiguous() && lhs.data().is_contiguous() { + Cow::Borrowed(lhs) + } else { + Cow::Owned(lhs.contiguous()?) + }; + let rhs_view = if rhs.is_contiguous() && rhs.data().is_contiguous() { + Cow::Borrowed(rhs) + } else { + Cow::Owned(rhs.contiguous()?) + }; + + let lhs_slice = lhs_view + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to access f32 data for lhs"))?; + let rhs_slice = rhs_view + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to access f32 data for rhs"))?; + let out_slice = output + .as_f32_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to access f32 output slice"))?; + + solve_batched( + lhs.shape().dims(), + rhs_cols, + lhs_slice, + rhs_slice, + out_slice, + ) +} + +fn solve_f64(lhs: &Tensor, rhs: &Tensor, output: &mut TensorData, rhs_cols: usize) -> Result<()> { + use std::borrow::Cow; + + let lhs_view = if lhs.is_contiguous() && lhs.data().is_contiguous() { + Cow::Borrowed(lhs) + } else { + Cow::Owned(lhs.contiguous()?) + }; + let rhs_view = if rhs.is_contiguous() && rhs.data().is_contiguous() { + Cow::Borrowed(rhs) + } else { + Cow::Owned(rhs.contiguous()?) + }; + + let lhs_slice = lhs_view + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to access f64 data for lhs"))?; + let rhs_slice = rhs_view + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to access f64 data for rhs"))?; + let out_slice = output + .as_f64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to access f64 output slice"))?; + + solve_batched( + lhs.shape().dims(), + rhs_cols, + lhs_slice, + rhs_slice, + out_slice, + ) +} + +fn solve_batched( + lhs_shape: &[usize], + rhs_cols: usize, + lhs_slice: &[T], + rhs_slice: &[T], + out_slice: &mut [T], +) -> Result<()> +where + T: Copy + + Send + + Sync + + std::ops::SubAssign + + std::ops::Mul + + std::ops::Div + + std::ops::Neg + + PartialOrd + + Default + + PartialEq, +{ + let n = *lhs_shape.last().expect("lhs has at least 2 dims"); + let batch = lhs_shape[..lhs_shape.len() - 2] + .iter() + .copied() + .product::() + .max(1); + let rhs_stride = n * rhs_cols; + + let matrix_stride = n * n; + let mut matrix = vec![T::default(); matrix_stride]; + let mut rhs_buf = vec![T::default(); rhs_stride]; + + for batch_idx in 0..batch { + let lhs_offset = batch_idx * matrix_stride; + let rhs_offset = batch_idx * rhs_stride; + + matrix.copy_from_slice(&lhs_slice[lhs_offset..lhs_offset + matrix_stride]); + rhs_buf[..rhs_stride].copy_from_slice(&rhs_slice[rhs_offset..rhs_offset + rhs_stride]); + + gaussian_elimination(&mut matrix, &mut rhs_buf, n, rhs_cols)?; + + out_slice[rhs_offset..rhs_offset + rhs_stride].copy_from_slice(&rhs_buf[..rhs_stride]); + } + + Ok(()) +} + +fn gaussian_elimination(matrix: &mut [T], rhs: &mut [T], n: usize, rhs_cols: usize) -> Result<()> +where + T: Copy + + Send + + Sync + + std::ops::SubAssign + + std::ops::Mul + + std::ops::Div + + std::ops::Neg + + PartialOrd + + Default + + PartialEq, +{ + for k in 0..n { + // Pivot selection + let mut pivot_row = k; + let mut pivot_val = abs(matrix[k * n + k]); + for i in (k + 1)..n { + let candidate = abs(matrix[i * n + k]); + if candidate > pivot_val { + pivot_val = candidate; + pivot_row = i; + } + } + + if pivot_val == T::default() { + return Err(MinitensorError::invalid_operation( + "solve received a singular matrix", + )); + } + + if pivot_row != k { + for col in 0..n { + matrix.swap(k * n + col, pivot_row * n + col); + } + for col in 0..rhs_cols { + rhs.swap(k * rhs_cols + col, pivot_row * rhs_cols + col); + } + } + + let pivot = matrix[k * n + k]; + + for i in (k + 1)..n { + let factor = matrix[i * n + k] / pivot; + matrix[i * n + k] = T::default(); + for j in (k + 1)..n { + let idx = i * n + j; + matrix[idx] -= factor * matrix[k * n + j]; + } + for col in 0..rhs_cols { + let idx = i * rhs_cols + col; + rhs[idx] -= factor * rhs[k * rhs_cols + col]; + } + } + } + + for i in (0..n).rev() { + let pivot = matrix[i * n + i]; + if abs(pivot) == T::default() { + return Err(MinitensorError::invalid_operation( + "solve received a singular matrix", + )); + } + for col in 0..rhs_cols { + let mut value = rhs[i * rhs_cols + col]; + for j in (i + 1)..n { + value -= matrix[i * n + j] * rhs[j * rhs_cols + col]; + } + rhs[i * rhs_cols + col] = value / pivot; + } + } + + Ok(()) +} + +fn abs(value: T) -> T +where + T: Copy + PartialOrd + std::ops::Neg + Default, +{ + if value < T::default() { -value } else { value } +} + +/// Batched matrix multiplication specialized for 3D tensors. +/// +/// This is a thin convenience wrapper around [`matmul`] that enforces the +/// traditional batch matrix multiply constraints: both operands must be +/// rank-3 tensors with matching batch dimensions. The actual computation is +/// still delegated to the highly optimised [`matmul`] implementation so all +/// execution happens inside the Rust backend. +pub fn bmm(lhs: &Tensor, rhs: &Tensor) -> Result { + if lhs.ndim() != 3 || rhs.ndim() != 3 { + return Err(MinitensorError::invalid_operation( + "bmm expects both inputs to be 3D tensors".to_string(), + )); + } + + let lhs_shape = lhs.shape().dims(); + let rhs_shape = rhs.shape().dims(); + + if lhs_shape[0] != rhs_shape[0] { + return Err(MinitensorError::shape_mismatch( + lhs_shape.to_vec(), + rhs_shape.to_vec(), + )); + } + + if lhs_shape[2] != rhs_shape[1] { + return Err(MinitensorError::shape_mismatch( + vec![lhs_shape[2]], + vec![rhs_shape[1]], + )); + } + + matmul(lhs, rhs) +} + +/// Dot product of two 1D tensors with gradient support +pub fn dot(lhs: &Tensor, rhs: &Tensor) -> Result { + if lhs.device() != rhs.device() { + return Err(MinitensorError::device_mismatch( + format!("{:?}", lhs.device()), + format!("{:?}", rhs.device()), + )); + } + + let lhs_dims = lhs.ndim(); + let rhs_dims = rhs.ndim(); + if lhs_dims != 1 || rhs_dims != 1 { + return Err(MinitensorError::invalid_operation(format!( + "dot: expected 1D tensors but got {}D and {}D tensors", + lhs_dims, rhs_dims + ))); + } + + if lhs.numel() != rhs.numel() { + return Err(MinitensorError::shape_mismatch( + lhs.shape().dims().to_vec(), + rhs.shape().dims().to_vec(), + )); + } + + let (lhs_cast, rhs_cast, result_dtype) = coerce_binary_operands(lhs, rhs, BinaryOpKind::Mul)?; + + if result_dtype == DataType::Bool { + return Err(MinitensorError::invalid_operation( + "dot does not support bool tensors", + )); + } + + let lhs_view = lhs_cast.as_ref(); + let rhs_view = rhs_cast.as_ref(); + + let numel = lhs_view.numel(); + let device = lhs.device(); + let requires_grad = lhs.requires_grad() || rhs.requires_grad(); + + let output_data = match result_dtype { + DataType::Float32 => { + let lhs_slice = lhs_view.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice for dot input") + })?; + let rhs_slice = rhs_view.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice for dot input") + })?; + + let dot = if numel >= PAR_THRESHOLD { + lhs_slice + .par_iter() + .zip(rhs_slice.par_iter()) + .map(|(&a, &b)| a * b) + .sum::() + } else { + lhs_slice + .iter() + .zip(rhs_slice.iter()) + .map(|(&a, &b)| a * b) + .sum::() + }; + + TensorData::from_vec_f32(vec![dot], device) + } + DataType::Float64 => { + let lhs_slice = lhs_view.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice for dot input") + })?; + let rhs_slice = rhs_view.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice for dot input") + })?; + + let dot = if numel >= PAR_THRESHOLD { + lhs_slice + .par_iter() + .zip(rhs_slice.par_iter()) + .map(|(&a, &b)| a * b) + .sum::() + } else { + lhs_slice + .iter() + .zip(rhs_slice.iter()) + .map(|(&a, &b)| a * b) + .sum::() + }; + + TensorData::from_vec_f64(vec![dot], device) + } + DataType::Int32 => { + let lhs_slice = lhs_view.data().as_i32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i32 slice for dot input") + })?; + let rhs_slice = rhs_view.data().as_i32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i32 slice for dot input") + })?; + + let mut dot: i32 = 0; + for (&a, &b) in lhs_slice.iter().zip(rhs_slice.iter()) { + dot = dot.wrapping_add(a.wrapping_mul(b)); + } + + TensorData::from_vec_i32(vec![dot], device) + } + DataType::Int64 => { + let lhs_slice = lhs_view.data().as_i64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i64 slice for dot input") + })?; + let rhs_slice = rhs_view.data().as_i64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i64 slice for dot input") + })?; + + let mut dot: i64 = 0; + for (&a, &b) in lhs_slice.iter().zip(rhs_slice.iter()) { + dot = dot.wrapping_add(a.wrapping_mul(b)); + } + + TensorData::from_vec_i64(vec![dot], device) + } + DataType::Bool => unreachable!("Bool dtype handled earlier"), + }; + + let output_shape = Shape::new(Vec::new()); + let output = Tensor::new( + Arc::new(output_data), + output_shape, + result_dtype, + device, + requires_grad, + ); + + if output.requires_grad() { + let lhs_requires_grad = lhs.requires_grad(); + let rhs_requires_grad = rhs.requires_grad(); + let grad_fn = Arc::new(DotBackward { + lhs: lhs_cast.into_owned().detach(), + rhs: rhs_cast.into_owned().detach(), + input_ids: [lhs.id(), rhs.id()], + lhs_requires_grad, + rhs_requires_grad, + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) + } else { + Ok(output) + } +} diff --git a/engine/src/operations/linalg/triangular.rs b/engine/src/operations/linalg/triangular.rs index 133c7220..85639c66 100644 --- a/engine/src/operations/linalg/triangular.rs +++ b/engine/src/operations/linalg/triangular.rs @@ -1,426 +1,436 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -/// Generic transpose implementation -fn transpose_generic( - input_data: &[T], - output_data: &mut [T], - input_shape: &Shape, - output_shape: &Shape, - dim0: usize, - dim1: usize, -) -> Result<()> { - // Fast path for 2D matrix transpose - if input_shape.ndim() == 2 && dim0 == 0 && dim1 == 1 { - let rows = input_shape.dims()[0]; - let cols = input_shape.dims()[1]; - if rows * cols < PAR_THRESHOLD { - for i in 0..rows { - for j in 0..cols { - unsafe { - *output_data.get_unchecked_mut(j * rows + i) = - *input_data.get_unchecked(i * cols + j); - } - } - } - } else { - output_data - .par_chunks_mut(rows) - .enumerate() - .for_each(|(j, col)| { - for i in 0..rows { - unsafe { - col[i] = *input_data.get_unchecked(i * cols + j); - } - } - }); - } - return Ok(()); - } - - let input_strides = Strides::from_shape(input_shape); - let output_strides = Strides::from_shape(output_shape); - let in_strides = input_strides.as_slice().to_vec(); - let out_strides = output_strides.as_slice().to_vec(); - let out_dims = output_shape.dims().to_vec(); - - output_data - .par_iter_mut() - .enumerate() - .for_each(|(idx, out)| { - let mut remaining = idx; - let mut input_linear = 0; - for dim in 0..out_dims.len() { - let stride = out_strides[dim]; - let coord = remaining / stride; - remaining %= stride; - let in_dim = if dim == dim0 { - dim1 - } else if dim == dim1 { - dim0 - } else { - dim - }; - input_linear += coord * in_strides[in_dim]; - } - *out = input_data[input_linear]; - }); - - Ok(()) -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::{autograd::GradientFunction, device::Device, tensor::TensorData}; - - fn create_test_tensor_f32(data: Vec, shape: Vec, requires_grad: bool) -> Tensor { - let shape_obj = Shape::new(shape); - let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Float32); - - if let Some(slice) = tensor_data.as_f32_slice_mut() { - slice.copy_from_slice(&data); - } - - Tensor::new( - Arc::new(tensor_data), - shape_obj, - DataType::Float32, - Device::cpu(), - requires_grad, - ) - } - - fn create_test_tensor_f64(data: Vec, shape: Vec, requires_grad: bool) -> Tensor { - let shape_obj = Shape::new(shape); - let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Float64); - - if let Some(slice) = tensor_data.as_f64_slice_mut() { - slice.copy_from_slice(&data); - } - - Tensor::new( - Arc::new(tensor_data), - shape_obj, - DataType::Float64, - Device::cpu(), - requires_grad, - ) - } - - fn create_test_tensor_i32(data: Vec, shape: Vec) -> Tensor { - let shape_obj = Shape::new(shape); - let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Int32); - - if let Some(slice) = tensor_data.as_i32_slice_mut() { - slice.copy_from_slice(&data); - } - - Tensor::new( - Arc::new(tensor_data), - shape_obj, - DataType::Int32, - Device::cpu(), - false, - ) - } - - fn create_test_tensor_bool(data: Vec, shape: Vec) -> Tensor { - let shape_obj = Shape::new(shape); - let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Bool); - - if let Some(slice) = tensor_data.as_bool_slice_mut() { - slice.copy_from_slice(&data); - } - - Tensor::new( - Arc::new(tensor_data), - shape_obj, - DataType::Bool, - Device::cpu(), - false, - ) - } - - fn create_test_tensor_f32_on_device( - data: Vec, - shape: Vec, - device: Device, - ) -> Tensor { - let shape_obj = Shape::new(shape); - let mut tensor_data = - TensorData::zeros_on_device(shape_obj.numel(), DataType::Float32, device); - - if let Some(slice) = tensor_data.as_f32_slice_mut() { - slice.copy_from_slice(&data); - } - - Tensor::new( - Arc::new(tensor_data), - shape_obj, - DataType::Float32, - device, - false, - ) - } - - #[test] - fn test_matmul_basic() { - // 2x3 * 3x2 = 2x2 - let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3], false); - let b = create_test_tensor_f32(vec![7.0, 8.0, 9.0, 10.0, 11.0, 12.0], vec![3, 2], false); - - let result = matmul(&a, &b).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - - // Expected: [1*7+2*9+3*11, 1*8+2*10+3*12; 4*7+5*9+6*11, 4*8+5*10+6*12] - // = [58, 64; 139, 154] - assert_eq!(result_data, &[58.0, 64.0, 139.0, 154.0]); - assert_eq!(result.shape().dims(), &[2, 2]); - } - - #[test] - fn test_matmul_i32_zero_k_dimension() { - let a = create_test_tensor_i32(vec![], vec![2, 0]); - let b = create_test_tensor_i32(vec![], vec![0, 3]); - - let result = matmul(&a, &b).unwrap(); - assert_eq!(result.shape().dims(), &[2, 3]); - assert_eq!(result.data().as_i32_slice().unwrap(), &[0, 0, 0, 0, 0, 0]); - } - - #[test] - fn test_transpose_2d() { - let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3], false); - - let result = transpose(&a, 0, 1).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - - // Original: [[1, 2, 3], [4, 5, 6]] - // Transposed: [[1, 4], [2, 5], [3, 6]] - assert_eq!(result_data, &[1.0, 4.0, 2.0, 5.0, 3.0, 6.0]); - assert_eq!(result.shape().dims(), &[3, 2]); - } - - #[test] - fn test_matmul_dimension_mismatch() { - let a = create_test_tensor_f32(vec![1.0, 2.0], vec![1, 2], false); - let b = create_test_tensor_f32(vec![3.0, 4.0, 5.0], vec![3, 1], false); - - let result = matmul(&a, &b); - assert!(result.is_err()); - } - - #[test] - fn test_transpose_same_dim() { - let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); - - let result = transpose(&a, 0, 0).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - - // Should be unchanged - assert_eq!(result_data, &[1.0, 2.0, 3.0, 4.0]); - assert_eq!(result.shape().dims(), &[2, 2]); - } - - #[test] - fn test_gradient_tracking() { - let a = create_test_tensor_f32(vec![1.0, 2.0], vec![1, 2], true); - let b = create_test_tensor_f32(vec![3.0, 4.0], vec![2, 1], true); - - let result = matmul(&a, &b).unwrap(); - - assert!(result.requires_grad()); - assert!(result.grad_fn().is_some()); - } - - #[test] - fn test_matmul_dtype_mismatch() { - let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); - let b = create_test_tensor_f64(vec![5.0, 6.0, 7.0, 8.0], vec![2, 2], false); - - let result = matmul(&a, &b); - assert!(result.is_err()); - } - - #[test] - fn test_matmul_device_mismatch() { - let a = - create_test_tensor_f32_on_device(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], Device::cpu()); - let b = create_test_tensor_f32_on_device( - vec![5.0, 6.0, 7.0, 8.0], - vec![2, 2], - Device::cuda(None), - ); - - let result = matmul(&a, &b); - assert!(result.is_err()); - } - - #[test] - fn test_matmul_bool_error() { - let a = create_test_tensor_bool(vec![true, false, true, false], vec![2, 2]); - let b = create_test_tensor_bool(vec![true, true, false, false], vec![2, 2]); - - let result = matmul(&a, &b); - assert!(result.is_err()); - } - - #[test] - fn test_matmul_vector_operands() { +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::{ + error::Result, + tensor::{Shape, Strides}, +}; +use rayon::prelude::*; + +/// Generic transpose implementation +pub(crate) fn transpose_generic( + input_data: &[T], + output_data: &mut [T], + input_shape: &Shape, + output_shape: &Shape, + dim0: usize, + dim1: usize, +) -> Result<()> { + // Fast path for 2D matrix transpose + if input_shape.ndim() == 2 && dim0 == 0 && dim1 == 1 { + let rows = input_shape.dims()[0]; + let cols = input_shape.dims()[1]; + if rows * cols < PAR_THRESHOLD { + for i in 0..rows { + for j in 0..cols { + unsafe { + *output_data.get_unchecked_mut(j * rows + i) = + *input_data.get_unchecked(i * cols + j); + } + } + } + } else { + output_data + .par_chunks_mut(rows) + .enumerate() + .for_each(|(j, col)| { + for i in 0..rows { + unsafe { + col[i] = *input_data.get_unchecked(i * cols + j); + } + } + }); + } + return Ok(()); + } + + let input_strides = Strides::from_shape(input_shape); + let output_strides = Strides::from_shape(output_shape); + let in_strides = input_strides.as_slice().to_vec(); + let out_strides = output_strides.as_slice().to_vec(); + let out_dims = output_shape.dims().to_vec(); + + output_data + .par_iter_mut() + .enumerate() + .for_each(|(idx, out)| { + let mut remaining = idx; + let mut input_linear = 0; + for dim in 0..out_dims.len() { + let stride = out_strides[dim]; + let coord = remaining / stride; + remaining %= stride; + let in_dim = if dim == dim0 { + dim1 + } else if dim == dim1 { + dim0 + } else { + dim + }; + input_linear += coord * in_strides[in_dim]; + } + *out = input_data[input_linear]; + }); + + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::tensor::DataType; + use crate::tensor::Tensor; + use crate::{autograd::GradientFunction, device::Device, tensor::TensorData}; + use std::sync::Arc; + + fn create_test_tensor_f32(data: Vec, shape: Vec, requires_grad: bool) -> Tensor { + let shape_obj = Shape::new(shape); + let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Float32); + + if let Some(slice) = tensor_data.as_f32_slice_mut() { + slice.copy_from_slice(&data); + } + + Tensor::new( + Arc::new(tensor_data), + shape_obj, + DataType::Float32, + Device::cpu(), + requires_grad, + ) + } + + fn create_test_tensor_f64(data: Vec, shape: Vec, requires_grad: bool) -> Tensor { + let shape_obj = Shape::new(shape); + let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Float64); + + if let Some(slice) = tensor_data.as_f64_slice_mut() { + slice.copy_from_slice(&data); + } + + Tensor::new( + Arc::new(tensor_data), + shape_obj, + DataType::Float64, + Device::cpu(), + requires_grad, + ) + } + + fn create_test_tensor_i32(data: Vec, shape: Vec) -> Tensor { + let shape_obj = Shape::new(shape); + let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Int32); + + if let Some(slice) = tensor_data.as_i32_slice_mut() { + slice.copy_from_slice(&data); + } + + Tensor::new( + Arc::new(tensor_data), + shape_obj, + DataType::Int32, + Device::cpu(), + false, + ) + } + + fn create_test_tensor_bool(data: Vec, shape: Vec) -> Tensor { + let shape_obj = Shape::new(shape); + let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Bool); + + if let Some(slice) = tensor_data.as_bool_slice_mut() { + slice.copy_from_slice(&data); + } + + Tensor::new( + Arc::new(tensor_data), + shape_obj, + DataType::Bool, + Device::cpu(), + false, + ) + } + + fn create_test_tensor_f32_on_device( + data: Vec, + shape: Vec, + device: Device, + ) -> Tensor { + let shape_obj = Shape::new(shape); + let mut tensor_data = + TensorData::zeros_on_device(shape_obj.numel(), DataType::Float32, device); + + if let Some(slice) = tensor_data.as_f32_slice_mut() { + slice.copy_from_slice(&data); + } + + Tensor::new( + Arc::new(tensor_data), + shape_obj, + DataType::Float32, + device, + false, + ) + } + + #[test] + fn test_matmul_basic() { + // 2x3 * 3x2 = 2x2 + let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3], false); + let b = create_test_tensor_f32(vec![7.0, 8.0, 9.0, 10.0, 11.0, 12.0], vec![3, 2], false); + + let result = matmul(&a, &b).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + + // Expected: [1*7+2*9+3*11, 1*8+2*10+3*12; 4*7+5*9+6*11, 4*8+5*10+6*12] + // = [58, 64; 139, 154] + assert_eq!(result_data, &[58.0, 64.0, 139.0, 154.0]); + assert_eq!(result.shape().dims(), &[2, 2]); + } + + #[test] + fn test_matmul_i32_zero_k_dimension() { + let a = create_test_tensor_i32(vec![], vec![2, 0]); + let b = create_test_tensor_i32(vec![], vec![0, 3]); + + let result = matmul(&a, &b).unwrap(); + assert_eq!(result.shape().dims(), &[2, 3]); + assert_eq!(result.data().as_i32_slice().unwrap(), &[0, 0, 0, 0, 0, 0]); + } + + #[test] + fn test_transpose_2d() { + let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3], false); + + let result = transpose(&a, 0, 1).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + + // Original: [[1, 2, 3], [4, 5, 6]] + // Transposed: [[1, 4], [2, 5], [3, 6]] + assert_eq!(result_data, &[1.0, 4.0, 2.0, 5.0, 3.0, 6.0]); + assert_eq!(result.shape().dims(), &[3, 2]); + } + + #[test] + fn test_matmul_dimension_mismatch() { + let a = create_test_tensor_f32(vec![1.0, 2.0], vec![1, 2], false); + let b = create_test_tensor_f32(vec![3.0, 4.0, 5.0], vec![3, 1], false); + + let result = matmul(&a, &b); + assert!(result.is_err()); + } + + #[test] + fn test_transpose_same_dim() { + let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); + + let result = transpose(&a, 0, 0).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + + // Should be unchanged + assert_eq!(result_data, &[1.0, 2.0, 3.0, 4.0]); + assert_eq!(result.shape().dims(), &[2, 2]); + } + + #[test] + fn test_gradient_tracking() { + let a = create_test_tensor_f32(vec![1.0, 2.0], vec![1, 2], true); + let b = create_test_tensor_f32(vec![3.0, 4.0], vec![2, 1], true); + + let result = matmul(&a, &b).unwrap(); + + assert!(result.requires_grad()); + assert!(result.grad_fn().is_some()); + } + + #[test] + fn test_matmul_dtype_mismatch() { + let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); + let b = create_test_tensor_f64(vec![5.0, 6.0, 7.0, 8.0], vec![2, 2], false); + + let result = matmul(&a, &b); + assert!(result.is_err()); + } + + #[test] + fn test_matmul_device_mismatch() { + let a = + create_test_tensor_f32_on_device(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], Device::cpu()); + let b = create_test_tensor_f32_on_device( + vec![5.0, 6.0, 7.0, 8.0], + vec![2, 2], + Device::cuda(None), + ); + + let result = matmul(&a, &b); + assert!(result.is_err()); + } + + #[test] + fn test_matmul_bool_error() { + let a = create_test_tensor_bool(vec![true, false, true, false], vec![2, 2]); + let b = create_test_tensor_bool(vec![true, true, false, false], vec![2, 2]); + + let result = matmul(&a, &b); + assert!(result.is_err()); + } + + #[test] + fn test_matmul_vector_operands() { // 1-D @ 1-D is a dot product returning a scalar. - let a = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - let b = create_test_tensor_f32(vec![3.0, 4.0], vec![2], false); - let dot = matmul(&a, &b).unwrap(); - assert_eq!(dot.shape().dims(), &[] as &[usize]); - assert_eq!(dot.data().as_f32_slice().unwrap(), &[11.0]); - - // matrix @ vector -> vector. - let m = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); - let mv = matmul(&m, &a).unwrap(); - assert_eq!(mv.shape().dims(), &[2]); - assert_eq!(mv.data().as_f32_slice().unwrap(), &[5.0, 11.0]); - - // 0-D scalars remain invalid operands. - let s = create_test_tensor_f32(vec![1.0], vec![], false); - assert!(matmul(&s, &s).is_err()); - } - - #[test] - fn test_bmm_basic() { - let a = create_test_tensor_f32( - vec![ - 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, // batch 0 - 7.0, 8.0, 9.0, 10.0, 11.0, 12.0, // batch 1 - ], - vec![2, 2, 3], - false, - ); - let b = create_test_tensor_f32( - vec![ - 0.5, 1.0, 1.5, 2.0, 2.5, 3.0, // batch 0 - 3.5, 4.0, 4.5, 5.0, 5.5, 6.0, // batch 1 - ], - vec![2, 3, 2], - false, - ); - - let result = bmm(&a, &b).unwrap(); - let result_data = result.data().as_f32_slice().unwrap(); - assert_eq!(result.shape().dims(), &[2, 2, 2]); - assert_eq!( - result_data, - &[11.0, 14.0, 24.5, 32.0, 110.0, 122.0, 150.5, 167.0] - ); - } - - #[test] - fn test_bmm_batch_mismatch() { - let a = create_test_tensor_f32(vec![1.0; 12], vec![2, 2, 3], false); - let b = create_test_tensor_f32(vec![2.0; 18], vec![3, 3, 2], false); - - let result = bmm(&a, &b); - assert!(result.is_err()); - } - - #[test] - fn test_bmm_rank_error() { - let a = create_test_tensor_f32(vec![1.0; 6], vec![2, 3], false); - let b = create_test_tensor_f32(vec![2.0; 6], vec![1, 3, 2], false); - - let result = bmm(&a, &b); - assert!(result.is_err()); - } - - #[test] - fn test_diagonal_main() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); - let result = diagonal(&tensor, 0, 0, 1).unwrap(); - let data = result.data().as_f32_slice().unwrap(); - assert_eq!(data, &[1.0, 4.0]); - assert_eq!(result.shape().dims(), &[2]); - } - - #[test] - fn test_diagonal_with_offset() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3], false); - let upper = diagonal(&tensor, 1, 0, 1).unwrap(); - assert_eq!(upper.data().as_f32_slice().unwrap(), &[2.0, 6.0]); - - let lower = diagonal(&tensor, -1, 0, 1).unwrap(); - assert_eq!(lower.data().as_f32_slice().unwrap(), &[4.0]); - } - - #[test] - fn test_diagonal_high_dim_shape() { - let tensor = - create_test_tensor_f32((0..24).map(|v| v as f32).collect(), vec![2, 3, 4], false); - let result = diagonal(&tensor, 0, 1, 2).unwrap(); - assert_eq!(result.shape().dims(), &[2, 3]); - } - - #[test] - fn test_diagonal_backward_gradients() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], true); - let grad_output = create_test_tensor_f32(vec![1.0, 1.0], vec![2], false); - - let backward_fn = crate::autograd::DiagonalBackward { - input_shape: tensor.shape().dims().to_vec(), - input_strides: tensor.strides().as_slice().to_vec(), - input_dtype: DataType::Float32, - dim1: 0, - dim2: 1, - offset: 0, - input_requires_grad: true, - input_id: tensor.id(), - }; - - let gradients = backward_fn.backward(&grad_output).unwrap(); - let grad_tensor = gradients.get(&tensor.id()).unwrap(); - let grad = grad_tensor.data().as_f32_slice().unwrap(); - assert_eq!(grad, &[1.0, 0.0, 0.0, 1.0]); - } - - #[test] - fn test_trace_matches_manual_sum() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); - let traced = trace(&tensor, 0, 0, 1).unwrap(); - let value = traced.data().as_f32_slice().unwrap(); - assert_eq!(value, &[5.0]); - } - - #[test] - fn test_triu_basic() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); - let result = triu(&tensor, 0).unwrap(); - let data = result.data().as_f32_slice().unwrap(); - assert_eq!(data, &[1.0, 2.0, 0.0, 4.0]); - } - - #[test] - fn test_triu_with_positive_diagonal() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); - let result = triu(&tensor, 1).unwrap(); - let data = result.data().as_f32_slice().unwrap(); - assert_eq!(data, &[0.0, 2.0, 0.0, 0.0]); - } - - #[test] - fn test_tril_basic() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); - let result = tril(&tensor, 0).unwrap(); - let data = result.data().as_f32_slice().unwrap(); - assert_eq!(data, &[1.0, 0.0, 3.0, 4.0]); - } - - #[test] - fn test_tril_with_negative_diagonal() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); - let result = tril(&tensor, -1).unwrap(); - let data = result.data().as_f32_slice().unwrap(); - assert_eq!(data, &[0.0, 0.0, 3.0, 0.0]); - } -} + let a = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + let b = create_test_tensor_f32(vec![3.0, 4.0], vec![2], false); + let dot = matmul(&a, &b).unwrap(); + assert_eq!(dot.shape().dims(), &[] as &[usize]); + assert_eq!(dot.data().as_f32_slice().unwrap(), &[11.0]); + + // matrix @ vector -> vector. + let m = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); + let mv = matmul(&m, &a).unwrap(); + assert_eq!(mv.shape().dims(), &[2]); + assert_eq!(mv.data().as_f32_slice().unwrap(), &[5.0, 11.0]); + + // 0-D scalars remain invalid operands. + let s = create_test_tensor_f32(vec![1.0], vec![], false); + assert!(matmul(&s, &s).is_err()); + } + + #[test] + fn test_bmm_basic() { + let a = create_test_tensor_f32( + vec![ + 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, // batch 0 + 7.0, 8.0, 9.0, 10.0, 11.0, 12.0, // batch 1 + ], + vec![2, 2, 3], + false, + ); + let b = create_test_tensor_f32( + vec![ + 0.5, 1.0, 1.5, 2.0, 2.5, 3.0, // batch 0 + 3.5, 4.0, 4.5, 5.0, 5.5, 6.0, // batch 1 + ], + vec![2, 3, 2], + false, + ); + + let result = bmm(&a, &b).unwrap(); + let result_data = result.data().as_f32_slice().unwrap(); + assert_eq!(result.shape().dims(), &[2, 2, 2]); + assert_eq!( + result_data, + &[11.0, 14.0, 24.5, 32.0, 110.0, 122.0, 150.5, 167.0] + ); + } + + #[test] + fn test_bmm_batch_mismatch() { + let a = create_test_tensor_f32(vec![1.0; 12], vec![2, 2, 3], false); + let b = create_test_tensor_f32(vec![2.0; 18], vec![3, 3, 2], false); + + let result = bmm(&a, &b); + assert!(result.is_err()); + } + + #[test] + fn test_bmm_rank_error() { + let a = create_test_tensor_f32(vec![1.0; 6], vec![2, 3], false); + let b = create_test_tensor_f32(vec![2.0; 6], vec![1, 3, 2], false); + + let result = bmm(&a, &b); + assert!(result.is_err()); + } + + #[test] + fn test_diagonal_main() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); + let result = diagonal(&tensor, 0, 0, 1).unwrap(); + let data = result.data().as_f32_slice().unwrap(); + assert_eq!(data, &[1.0, 4.0]); + assert_eq!(result.shape().dims(), &[2]); + } + + #[test] + fn test_diagonal_with_offset() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3], false); + let upper = diagonal(&tensor, 1, 0, 1).unwrap(); + assert_eq!(upper.data().as_f32_slice().unwrap(), &[2.0, 6.0]); + + let lower = diagonal(&tensor, -1, 0, 1).unwrap(); + assert_eq!(lower.data().as_f32_slice().unwrap(), &[4.0]); + } + + #[test] + fn test_diagonal_high_dim_shape() { + let tensor = + create_test_tensor_f32((0..24).map(|v| v as f32).collect(), vec![2, 3, 4], false); + let result = diagonal(&tensor, 0, 1, 2).unwrap(); + assert_eq!(result.shape().dims(), &[2, 3]); + } + + #[test] + fn test_diagonal_backward_gradients() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], true); + let grad_output = create_test_tensor_f32(vec![1.0, 1.0], vec![2], false); + + let backward_fn = crate::autograd::DiagonalBackward { + input_shape: tensor.shape().dims().to_vec(), + input_strides: tensor.strides().as_slice().to_vec(), + input_dtype: DataType::Float32, + dim1: 0, + dim2: 1, + offset: 0, + input_requires_grad: true, + input_id: tensor.id(), + }; + + let gradients = backward_fn.backward(&grad_output).unwrap(); + let grad_tensor = gradients.get(&tensor.id()).unwrap(); + let grad = grad_tensor.data().as_f32_slice().unwrap(); + assert_eq!(grad, &[1.0, 0.0, 0.0, 1.0]); + } + + #[test] + fn test_trace_matches_manual_sum() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); + let traced = trace(&tensor, 0, 0, 1).unwrap(); + let value = traced.data().as_f32_slice().unwrap(); + assert_eq!(value, &[5.0]); + } + + #[test] + fn test_triu_basic() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); + let result = triu(&tensor, 0).unwrap(); + let data = result.data().as_f32_slice().unwrap(); + assert_eq!(data, &[1.0, 2.0, 0.0, 4.0]); + } + + #[test] + fn test_triu_with_positive_diagonal() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); + let result = triu(&tensor, 1).unwrap(); + let data = result.data().as_f32_slice().unwrap(); + assert_eq!(data, &[0.0, 2.0, 0.0, 0.0]); + } + + #[test] + fn test_tril_basic() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); + let result = tril(&tensor, 0).unwrap(); + let data = result.data().as_f32_slice().unwrap(); + assert_eq!(data, &[1.0, 0.0, 3.0, 4.0]); + } + + #[test] + fn test_tril_with_negative_diagonal() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); + let result = tril(&tensor, -1).unwrap(); + let data = result.data().as_f32_slice().unwrap(); + assert_eq!(data, &[0.0, 0.0, 3.0, 0.0]); + } +} diff --git a/engine/src/operations/loss.rs b/engine/src/operations/loss.rs index 7b39ea54..fa8a2262 100644 --- a/engine/src/operations/loss.rs +++ b/engine/src/operations/loss.rs @@ -4,5 +4,10 @@ // This source code is licensed under the Apache-style license found in the // LICENSE file in the root directory of this source tree. -include!("loss/regression.rs"); -include!("loss/classification.rs"); +#[path = "loss/classification.rs"] +mod classification_impl; +#[path = "loss/regression.rs"] +mod regression_impl; + +pub(crate) use self::classification_impl::*; +pub use self::regression_impl::*; diff --git a/engine/src/operations/loss/classification.rs b/engine/src/operations/loss/classification.rs index 101c5f09..8e9b7c2e 100644 --- a/engine/src/operations/loss/classification.rs +++ b/engine/src/operations/loss/classification.rs @@ -1,710 +1,722 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -fn fill_one_hot_f64( - indices: &[T], - out: &mut [f64], - num_classes: usize, - to_index: F, -) -> Result<()> -where - F: Fn(&T) -> Result, -{ - for (i, value) in indices.iter().enumerate() { - let class = to_index(value)?; - out[i * num_classes + class] = 1.0; - } - Ok(()) -} - -/// Compute the sign of each tensor element (-1.0, 0.0, or 1.0) -fn sign(tensor: &Tensor) -> Result { - let mut output_data = - TensorData::zeros_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from tensor") - })?; - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output") - })?; - - output_slice - .par_chunks_mut(CHUNK) - .zip(input_data.par_chunks(CHUNK)) - .for_each(|(out, inp)| unsafe { - let in_ptr = inp.as_ptr(); - let out_ptr = out.as_mut_ptr(); - for i in 0..out.len() { - let v = *in_ptr.add(i); - *out_ptr.add(i) = if v > 0.0 { - 1.0 - } else if v < 0.0 { - -1.0 - } else { - 0.0 - }; - } - }); - } - DataType::Float64 => { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from tensor") - })?; - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output") - })?; - - output_slice - .par_chunks_mut(CHUNK) - .zip(input_data.par_chunks(CHUNK)) - .for_each(|(out, inp)| unsafe { - let in_ptr = inp.as_ptr(); - let out_ptr = out.as_mut_ptr(); - for i in 0..out.len() { - let v = *in_ptr.add(i); - *out_ptr.add(i) = if v > 0.0 { - 1.0 - } else if v < 0.0 { - -1.0 - } else { - 0.0 - }; - } - }); - } - _ => { - return Err(MinitensorError::invalid_operation( - "Sign operation only supported for floating point tensors", - )); - } - } - - Ok(Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - false, - )) -} - -/// Sum all elements in a tensor to produce a scalar -fn sum_all_elements(tensor: &Tensor) -> Result { +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::operations::arithmetic::mul; +use crate::{ + error::{MinitensorError, Result}, + tensor::{DataType, Shape, Tensor, TensorData}, +}; +use rayon::prelude::*; +use std::sync::Arc; + +pub(crate) fn fill_one_hot_f64( + indices: &[T], + out: &mut [f64], + num_classes: usize, + to_index: F, +) -> Result<()> +where + F: Fn(&T) -> Result, +{ + for (i, value) in indices.iter().enumerate() { + let class = to_index(value)?; + out[i * num_classes + class] = 1.0; + } + Ok(()) +} + +/// Compute the sign of each tensor element (-1.0, 0.0, or 1.0) +pub(crate) fn sign_tensor(tensor: &Tensor) -> Result { + let mut output_data = + TensorData::zeros_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from tensor") + })?; + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output") + })?; + + output_slice + .par_chunks_mut(CHUNK) + .zip(input_data.par_chunks(CHUNK)) + .for_each(|(out, inp)| unsafe { + let in_ptr = inp.as_ptr(); + let out_ptr = out.as_mut_ptr(); + for i in 0..out.len() { + let v = *in_ptr.add(i); + *out_ptr.add(i) = if v > 0.0 { + 1.0 + } else if v < 0.0 { + -1.0 + } else { + 0.0 + }; + } + }); + } + DataType::Float64 => { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from tensor") + })?; + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output") + })?; + + output_slice + .par_chunks_mut(CHUNK) + .zip(input_data.par_chunks(CHUNK)) + .for_each(|(out, inp)| unsafe { + let in_ptr = inp.as_ptr(); + let out_ptr = out.as_mut_ptr(); + for i in 0..out.len() { + let v = *in_ptr.add(i); + *out_ptr.add(i) = if v > 0.0 { + 1.0 + } else if v < 0.0 { + -1.0 + } else { + 0.0 + }; + } + }); + } + _ => { + return Err(MinitensorError::invalid_operation( + "Sign operation only supported for floating point tensors", + )); + } + } + + Ok(Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + false, + )) +} + +/// Sum all elements in a tensor to produce a scalar +pub(crate) fn sum_all_elements(tensor: &Tensor) -> Result { // Reduced losses are 0-dim scalars; a shape-[1] result breaks float(loss). - let scalar_shape = Shape::scalar(); - let mut output_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from tensor") - })?; - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output") - })?; - - let sum: f32 = input_data - .par_chunks(CHUNK) - .map(|chunk| unsafe { - let mut acc = 0f32; - let ptr = chunk.as_ptr(); - for i in 0..chunk.len() { - acc += *ptr.add(i); - } - acc - }) - .sum(); - output_slice[0] = sum; - } - DataType::Float64 => { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from tensor") - })?; - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output") - })?; - - let sum: f64 = input_data - .par_chunks(CHUNK) - .map(|chunk| unsafe { - let mut acc = 0f64; - let ptr = chunk.as_ptr(); - for i in 0..chunk.len() { - acc += *ptr.add(i); - } - acc - }) - .sum(); - output_slice[0] = sum; - } - _ => { - return Err(MinitensorError::invalid_operation( - "Sum only supported for floating point tensors", - )); - } - } - - Ok(Tensor::new( - Arc::new(output_data), - scalar_shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )) -} - -/// Divide tensor by a scalar value -fn divide_by_scalar(tensor: &Tensor, scalar: f64) -> Result { - let mut output_data = - TensorData::zeros_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from tensor") - })?; - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output") - })?; - - let scalar_f32 = scalar as f32; - output_slice - .par_chunks_mut(CHUNK) - .zip(input_data.par_chunks(CHUNK)) - .for_each(|(out, inp)| unsafe { - let in_ptr = inp.as_ptr(); - let out_ptr = out.as_mut_ptr(); - for i in 0..out.len() { - *out_ptr.add(i) = *in_ptr.add(i) / scalar_f32; - } - }); - } - DataType::Float64 => { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from tensor") - })?; - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output") - })?; - - output_slice - .par_chunks_mut(CHUNK) - .zip(input_data.par_chunks(CHUNK)) - .for_each(|(out, inp)| unsafe { - let in_ptr = inp.as_ptr(); - let out_ptr = out.as_mut_ptr(); - for i in 0..out.len() { - *out_ptr.add(i) = *in_ptr.add(i) / scalar; - } - }); - } - _ => { - return Err(MinitensorError::invalid_operation( - "Division only supported for floating point tensors", - )); - } - } - - Ok(Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )) -} - -/// Create a scalar tensor with the given value -fn create_scalar_tensor( - value: f64, - dtype: DataType, - device: crate::device::Device, -) -> Result { - let scalar_shape = Shape::new(vec![1]); - let mut tensor_data = TensorData::zeros_on_device(1, dtype, device); - - match dtype { - DataType::Float32 => { - let slice = tensor_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?; - slice[0] = value as f32; - } - DataType::Float64 => { - let slice = tensor_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?; - slice[0] = value; - } - _ => { - return Err(MinitensorError::invalid_operation( - "Scalar tensor creation only supported for floating point types", - )); - } - } - - Ok(Tensor::new( - Arc::new(tensor_data), - scalar_shape, - dtype, - device, - false, - )) -} - -/// Compute Huber loss element-wise -fn compute_huber_elementwise( - abs_diff: &Tensor, - diff: &Tensor, - _delta_tensor: &Tensor, - delta: f64, -) -> Result { - let mut output_data = - TensorData::zeros_on_device(abs_diff.numel(), abs_diff.dtype(), abs_diff.device()); - - match abs_diff.dtype() { - DataType::Float32 => { - let abs_data = abs_diff.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from abs_diff") - })?; - let diff_data = diff.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from diff") - })?; - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output") - })?; - - let delta_f32 = delta as f32; - output_slice - .par_chunks_mut(CHUNK) - .zip(abs_data.par_chunks(CHUNK).zip(diff_data.par_chunks(CHUNK))) - .for_each(|(out, (abs_chunk, diff_chunk))| unsafe { - let abs_ptr = abs_chunk.as_ptr(); - let diff_ptr = diff_chunk.as_ptr(); - let out_ptr = out.as_mut_ptr(); - for i in 0..out.len() { - let abs_val = *abs_ptr.add(i); - *out_ptr.add(i) = if abs_val <= delta_f32 { - 0.5 * *diff_ptr.add(i) * *diff_ptr.add(i) - } else { - delta_f32 * (abs_val - 0.5 * delta_f32) - }; - } - }); - } - DataType::Float64 => { - let abs_data = abs_diff.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from abs_diff") - })?; - let diff_data = diff.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from diff") - })?; - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output") - })?; - - output_slice - .par_chunks_mut(CHUNK) - .zip(abs_data.par_chunks(CHUNK).zip(diff_data.par_chunks(CHUNK))) - .for_each(|(out, (abs_chunk, diff_chunk))| unsafe { - let abs_ptr = abs_chunk.as_ptr(); - let diff_ptr = diff_chunk.as_ptr(); - let out_ptr = out.as_mut_ptr(); - for i in 0..out.len() { - let abs_val = *abs_ptr.add(i); - *out_ptr.add(i) = if abs_val <= delta { - 0.5 * *diff_ptr.add(i) * *diff_ptr.add(i) - } else { - delta * (abs_val - 0.5 * delta) - }; - } - }); - } - _ => { - return Err(MinitensorError::invalid_operation( - "Huber loss only supported for floating point tensors", - )); - } - } - - Ok(Tensor::new( - Arc::new(output_data), - abs_diff.shape().clone(), - abs_diff.dtype(), - abs_diff.device(), - abs_diff.requires_grad(), - )) -} - -/// Compute natural logarithm of tensor elements -fn log(tensor: &Tensor) -> Result { - let mut output_data = - TensorData::zeros_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from tensor") - })?; - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output") - })?; - - for (i, &val) in input_data.iter().enumerate() { - if val <= 0.0 { - output_slice[i] = f32::NEG_INFINITY; - } else { - output_slice[i] = val.ln(); - } - } - } - DataType::Float64 => { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from tensor") - })?; - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output") - })?; - - for (i, &val) in input_data.iter().enumerate() { - if val <= 0.0 { - output_slice[i] = f64::NEG_INFINITY; - } else { - output_slice[i] = val.ln(); - } - } - } - _ => { - return Err(MinitensorError::invalid_operation( - "Logarithm only supported for floating point tensors", - )); - } - } - - Ok(Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )) -} - -/// Negate tensor elements -fn negate(tensor: &Tensor) -> Result { - let mut output_data = - TensorData::zeros_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from tensor") - })?; - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output") - })?; - - for (i, &val) in input_data.iter().enumerate() { - output_slice[i] = -val; - } - } - DataType::Float64 => { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from tensor") - })?; - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output") - })?; - - for (i, &val) in input_data.iter().enumerate() { - output_slice[i] = -val; - } - } - _ => { - return Err(MinitensorError::invalid_operation( - "Negation only supported for floating point tensors", - )); - } - } - - Ok(Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )) -} - -/// Compute negative log likelihood for classification -fn negative_log_likelihood(log_predictions: &Tensor, targets: &Tensor) -> Result { - // Simplified implementation - multiply log predictions by targets and negate - let likelihood = mul(log_predictions, targets)?; - negate(&likelihood) -} - -/// Raise tensor elements to a power -fn power(tensor: &Tensor, exponent: f64) -> Result { - let mut output_data = - TensorData::zeros_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => { - let input_data = tensor.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from tensor") - })?; - let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice from output") - })?; - - let exp_f32 = exponent as f32; - for (i, &val) in input_data.iter().enumerate() { - output_slice[i] = val.powf(exp_f32); - } - } - DataType::Float64 => { - let input_data = tensor.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from tensor") - })?; - let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice from output") - })?; - - for (i, &val) in input_data.iter().enumerate() { - output_slice[i] = val.powf(exponent); - } - } - _ => { - return Err(MinitensorError::invalid_operation( - "Power operation only supported for floating point tensors", - )); - } - } - - Ok(Tensor::new( - Arc::new(output_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )) -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::device::Device; - - fn create_test_tensor_f32(data: Vec, shape: Vec, requires_grad: bool) -> Tensor { - let shape_obj = Shape::new(shape); - let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Float32); - - if let Some(slice) = tensor_data.as_f32_slice_mut() { - slice.copy_from_slice(&data); - } - - Tensor::new( - Arc::new(tensor_data), - shape_obj, - DataType::Float32, - Device::cpu(), - requires_grad, - ) - } - - #[test] - fn test_mse_loss_mean() { - let predictions = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], false); - let targets = create_test_tensor_f32(vec![1.5, 2.5, 2.5], vec![3], false); - - let loss = mse_loss(&predictions, &targets, "mean").unwrap(); - let loss_data = loss.data().as_f32_slice().unwrap(); - - // Expected: ((1.0-1.5)² + (2.0-2.5)² + (3.0-2.5)²) / 3 = (0.25 + 0.25 + 0.25) / 3 = 0.25 - assert!((loss_data[0] - 0.25).abs() < 1e-6); + let scalar_shape = Shape::scalar(); + let mut output_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from tensor") + })?; + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output") + })?; + + let sum: f32 = input_data + .par_chunks(CHUNK) + .map(|chunk| unsafe { + let mut acc = 0f32; + let ptr = chunk.as_ptr(); + for i in 0..chunk.len() { + acc += *ptr.add(i); + } + acc + }) + .sum(); + output_slice[0] = sum; + } + DataType::Float64 => { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from tensor") + })?; + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output") + })?; + + let sum: f64 = input_data + .par_chunks(CHUNK) + .map(|chunk| unsafe { + let mut acc = 0f64; + let ptr = chunk.as_ptr(); + for i in 0..chunk.len() { + acc += *ptr.add(i); + } + acc + }) + .sum(); + output_slice[0] = sum; + } + _ => { + return Err(MinitensorError::invalid_operation( + "Sum only supported for floating point tensors", + )); + } + } + + Ok(Tensor::new( + Arc::new(output_data), + scalar_shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )) +} + +/// Divide tensor by a scalar value +pub(crate) fn divide_by_scalar(tensor: &Tensor, scalar: f64) -> Result { + let mut output_data = + TensorData::zeros_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from tensor") + })?; + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output") + })?; + + let scalar_f32 = scalar as f32; + output_slice + .par_chunks_mut(CHUNK) + .zip(input_data.par_chunks(CHUNK)) + .for_each(|(out, inp)| unsafe { + let in_ptr = inp.as_ptr(); + let out_ptr = out.as_mut_ptr(); + for i in 0..out.len() { + *out_ptr.add(i) = *in_ptr.add(i) / scalar_f32; + } + }); + } + DataType::Float64 => { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from tensor") + })?; + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output") + })?; + + output_slice + .par_chunks_mut(CHUNK) + .zip(input_data.par_chunks(CHUNK)) + .for_each(|(out, inp)| unsafe { + let in_ptr = inp.as_ptr(); + let out_ptr = out.as_mut_ptr(); + for i in 0..out.len() { + *out_ptr.add(i) = *in_ptr.add(i) / scalar; + } + }); + } + _ => { + return Err(MinitensorError::invalid_operation( + "Division only supported for floating point tensors", + )); + } + } + + Ok(Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )) +} + +/// Create a scalar tensor with the given value +pub(crate) fn create_scalar_tensor( + value: f64, + dtype: DataType, + device: crate::device::Device, +) -> Result { + let scalar_shape = Shape::new(vec![1]); + let mut tensor_data = TensorData::zeros_on_device(1, dtype, device); + + match dtype { + DataType::Float32 => { + let slice = tensor_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?; + slice[0] = value as f32; + } + DataType::Float64 => { + let slice = tensor_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?; + slice[0] = value; + } + _ => { + return Err(MinitensorError::invalid_operation( + "Scalar tensor creation only supported for floating point types", + )); + } + } + + Ok(Tensor::new( + Arc::new(tensor_data), + scalar_shape, + dtype, + device, + false, + )) +} + +/// Compute Huber loss element-wise +pub(crate) fn compute_huber_elementwise( + abs_diff: &Tensor, + diff: &Tensor, + _delta_tensor: &Tensor, + delta: f64, +) -> Result { + let mut output_data = + TensorData::zeros_on_device(abs_diff.numel(), abs_diff.dtype(), abs_diff.device()); + + match abs_diff.dtype() { + DataType::Float32 => { + let abs_data = abs_diff.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from abs_diff") + })?; + let diff_data = diff.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from diff") + })?; + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output") + })?; + + let delta_f32 = delta as f32; + output_slice + .par_chunks_mut(CHUNK) + .zip(abs_data.par_chunks(CHUNK).zip(diff_data.par_chunks(CHUNK))) + .for_each(|(out, (abs_chunk, diff_chunk))| unsafe { + let abs_ptr = abs_chunk.as_ptr(); + let diff_ptr = diff_chunk.as_ptr(); + let out_ptr = out.as_mut_ptr(); + for i in 0..out.len() { + let abs_val = *abs_ptr.add(i); + *out_ptr.add(i) = if abs_val <= delta_f32 { + 0.5 * *diff_ptr.add(i) * *diff_ptr.add(i) + } else { + delta_f32 * (abs_val - 0.5 * delta_f32) + }; + } + }); + } + DataType::Float64 => { + let abs_data = abs_diff.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from abs_diff") + })?; + let diff_data = diff.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from diff") + })?; + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output") + })?; + + output_slice + .par_chunks_mut(CHUNK) + .zip(abs_data.par_chunks(CHUNK).zip(diff_data.par_chunks(CHUNK))) + .for_each(|(out, (abs_chunk, diff_chunk))| unsafe { + let abs_ptr = abs_chunk.as_ptr(); + let diff_ptr = diff_chunk.as_ptr(); + let out_ptr = out.as_mut_ptr(); + for i in 0..out.len() { + let abs_val = *abs_ptr.add(i); + *out_ptr.add(i) = if abs_val <= delta { + 0.5 * *diff_ptr.add(i) * *diff_ptr.add(i) + } else { + delta * (abs_val - 0.5 * delta) + }; + } + }); + } + _ => { + return Err(MinitensorError::invalid_operation( + "Huber loss only supported for floating point tensors", + )); + } + } + + Ok(Tensor::new( + Arc::new(output_data), + abs_diff.shape().clone(), + abs_diff.dtype(), + abs_diff.device(), + abs_diff.requires_grad(), + )) +} + +/// Compute natural logarithm of tensor elements +pub(crate) fn log_tensor(tensor: &Tensor) -> Result { + let mut output_data = + TensorData::zeros_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from tensor") + })?; + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output") + })?; + + for (i, &val) in input_data.iter().enumerate() { + if val <= 0.0 { + output_slice[i] = f32::NEG_INFINITY; + } else { + output_slice[i] = val.ln(); + } + } + } + DataType::Float64 => { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from tensor") + })?; + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output") + })?; + + for (i, &val) in input_data.iter().enumerate() { + if val <= 0.0 { + output_slice[i] = f64::NEG_INFINITY; + } else { + output_slice[i] = val.ln(); + } + } + } + _ => { + return Err(MinitensorError::invalid_operation( + "Logarithm only supported for floating point tensors", + )); + } + } + + Ok(Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )) +} + +/// Negate tensor elements +fn negate(tensor: &Tensor) -> Result { + let mut output_data = + TensorData::zeros_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from tensor") + })?; + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output") + })?; + + for (i, &val) in input_data.iter().enumerate() { + output_slice[i] = -val; + } + } + DataType::Float64 => { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from tensor") + })?; + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output") + })?; + + for (i, &val) in input_data.iter().enumerate() { + output_slice[i] = -val; + } + } + _ => { + return Err(MinitensorError::invalid_operation( + "Negation only supported for floating point tensors", + )); + } + } + + Ok(Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )) +} + +/// Compute negative log likelihood for classification +pub(crate) fn negative_log_likelihood( + log_predictions: &Tensor, + targets: &Tensor, +) -> Result { + // Simplified implementation - multiply log predictions by targets and negate + let likelihood = mul(log_predictions, targets)?; + negate(&likelihood) +} + +/// Raise tensor elements to a power +pub(crate) fn power(tensor: &Tensor, exponent: f64) -> Result { + let mut output_data = + TensorData::zeros_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => { + let input_data = tensor.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from tensor") + })?; + let output_slice = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice from output") + })?; + + let exp_f32 = exponent as f32; + for (i, &val) in input_data.iter().enumerate() { + output_slice[i] = val.powf(exp_f32); + } + } + DataType::Float64 => { + let input_data = tensor.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from tensor") + })?; + let output_slice = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice from output") + })?; + + for (i, &val) in input_data.iter().enumerate() { + output_slice[i] = val.powf(exponent); + } + } + _ => { + return Err(MinitensorError::invalid_operation( + "Power operation only supported for floating point tensors", + )); + } + } + + Ok(Tensor::new( + Arc::new(output_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::device::Device; + + fn create_test_tensor_f32(data: Vec, shape: Vec, requires_grad: bool) -> Tensor { + let shape_obj = Shape::new(shape); + let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Float32); + + if let Some(slice) = tensor_data.as_f32_slice_mut() { + slice.copy_from_slice(&data); + } + + Tensor::new( + Arc::new(tensor_data), + shape_obj, + DataType::Float32, + Device::cpu(), + requires_grad, + ) + } + + #[test] + fn test_mse_loss_mean() { + let predictions = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], false); + let targets = create_test_tensor_f32(vec![1.5, 2.5, 2.5], vec![3], false); + + let loss = mse_loss(&predictions, &targets, "mean").unwrap(); + let loss_data = loss.data().as_f32_slice().unwrap(); + + // Expected: ((1.0-1.5)² + (2.0-2.5)² + (3.0-2.5)²) / 3 = (0.25 + 0.25 + 0.25) / 3 = 0.25 + assert!((loss_data[0] - 0.25).abs() < 1e-6); // A reduced loss is a 0-dim scalar. - assert_eq!(loss.shape().dims(), &[] as &[usize]); - } - - #[test] - fn test_mse_loss_sum() { - let predictions = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - let targets = create_test_tensor_f32(vec![2.0, 3.0], vec![2], false); - - let loss = mse_loss(&predictions, &targets, "sum").unwrap(); - let loss_data = loss.data().as_f32_slice().unwrap(); - - // Expected: (1.0-2.0)² + (2.0-3.0)² = 1.0 + 1.0 = 2.0 - assert!((loss_data[0] - 2.0).abs() < 1e-6); - } - - #[test] - fn test_mse_loss_none() { - let predictions = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - let targets = create_test_tensor_f32(vec![2.0, 3.0], vec![2], false); - - let loss = mse_loss(&predictions, &targets, "none").unwrap(); - let loss_data = loss.data().as_f32_slice().unwrap(); - - // Expected: [(1.0-2.0)², (2.0-3.0)²] = [1.0, 1.0] - assert!((loss_data[0] - 1.0).abs() < 1e-6); - assert!((loss_data[1] - 1.0).abs() < 1e-6); - assert_eq!(loss.shape().dims(), &[2]); - } - - #[test] - fn test_mae_loss_mean() { - let predictions = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], false); - let targets = create_test_tensor_f32(vec![1.5, 2.5, 2.0], vec![3], false); - - let loss = mae_loss(&predictions, &targets, "mean").unwrap(); - let loss_data = loss.data().as_f32_slice().unwrap(); - - // Expected: (|1.0-1.5| + |2.0-2.5| + |3.0-2.0|) / 3 = (0.5 + 0.5 + 1.0) / 3 = 2.0/3 ≈ 0.667 - assert!((loss_data[0] - (2.0 / 3.0)).abs() < 1e-6); - } - - #[test] - fn test_huber_loss_quadratic_region() { - let predictions = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - let targets = create_test_tensor_f32(vec![1.2, 2.3], vec![2], false); - - // Delta = 1.0, differences are 0.2 and 0.3, both <= 1.0, so quadratic - let loss = huber_loss(&predictions, &targets, 1.0, "none").unwrap(); - let loss_data = loss.data().as_f32_slice().unwrap(); - - // Expected: [0.5 * 0.2², 0.5 * 0.3²] = [0.02, 0.045] - assert!((loss_data[0] - 0.02).abs() < 1e-6); - assert!((loss_data[1] - 0.045).abs() < 1e-6); - } - - #[test] - fn test_huber_loss_linear_region() { - let predictions = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - let targets = create_test_tensor_f32(vec![3.0, 0.0], vec![2], false); - - // Delta = 1.0, differences are 2.0 and 2.0, both > 1.0, so linear - let loss = huber_loss(&predictions, &targets, 1.0, "none").unwrap(); - let loss_data = loss.data().as_f32_slice().unwrap(); - - // Expected: [1.0 * (2.0 - 0.5 * 1.0), 1.0 * (2.0 - 0.5 * 1.0)] = [1.5, 1.5] - assert!((loss_data[0] - 1.5).abs() < 1e-6); - assert!((loss_data[1] - 1.5).abs() < 1e-6); - } - - #[test] - fn test_bce_loss_mean_and_backward() { - let predictions = create_test_tensor_f32(vec![0.8, 0.2], vec![2], true); - let targets = create_test_tensor_f32(vec![1.0, 0.0], vec![2], false); - - let loss = binary_cross_entropy_loss(&predictions, &targets, "mean").unwrap(); - let loss_val = loss.data().as_f32_slice().unwrap()[0]; - let expected = -((0.8f32).ln() + (0.8f32).ln()) / 2.0; - assert!((loss_val - expected).abs() < 1e-6); - - let grads = crate::autograd::backward(&loss, None).unwrap(); - let grad = grads.get(&predictions.id()).unwrap(); - let grad_slice = grad.data().as_f32_slice().unwrap(); - let expected_grad = [-(1.0 / 0.8) / 2.0, (1.0 / 0.8) / 2.0]; - assert!((grad_slice[0] - expected_grad[0]).abs() < 1e-6); - assert!((grad_slice[1] - expected_grad[1]).abs() < 1e-6); - } - - #[test] - fn test_kl_div_loss_mean_and_backward() { - let predictions = create_test_tensor_f32(vec![0.4, 0.6], vec![2], true); - let targets = create_test_tensor_f32(vec![0.5, 0.5], vec![2], false); - - let loss = kl_div_loss(&predictions, &targets, "mean").unwrap(); - let loss_val = loss.data().as_f32_slice().unwrap()[0]; - let expected = 0.5 * ((0.5f32.ln() - 0.4f32.ln()) + (0.5f32.ln() - 0.6f32.ln())); - assert!((loss_val - expected).abs() < 1e-6); - - let grads = crate::autograd::backward(&loss, None).unwrap(); - let grad = grads.get(&predictions.id()).unwrap(); - let grad_slice = grad.data().as_f32_slice().unwrap(); - let expected_grad = [-(0.5 / 0.4) / 2.0, -(0.5 / 0.6) / 2.0]; - assert!((grad_slice[0] - expected_grad[0]).abs() < 1e-6); - assert!((grad_slice[1] - expected_grad[1]).abs() < 1e-6); - } - - #[test] - fn test_loss_gradient_tracking() { - let predictions = create_test_tensor_f32(vec![1.0, 2.0], vec![2], true); - let targets = create_test_tensor_f32(vec![1.5, 2.5], vec![2], false); - - let loss = mse_loss(&predictions, &targets, "mean").unwrap(); - - assert!(loss.requires_grad()); - assert!(loss.grad_fn().is_some()); - } - - #[test] - fn test_loss_input_validation() { - let predictions = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - let targets = create_test_tensor_f32(vec![1.5, 2.5, 3.5], vec![3], false); - - // Shape mismatch should fail - let result = mse_loss(&predictions, &targets, "mean"); - assert!(result.is_err()); - } - - #[test] - fn test_invalid_reduction_mode() { - let predictions = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - let targets = create_test_tensor_f32(vec![1.5, 2.5], vec![2], false); - - let result = mse_loss(&predictions, &targets, "invalid"); - assert!(result.is_err()); - } - - #[test] - fn test_huber_loss_invalid_delta() { - let predictions = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - let targets = create_test_tensor_f32(vec![1.5, 2.5], vec![2], false); - - let result = huber_loss(&predictions, &targets, -1.0, "mean"); - assert!(result.is_err()); - } - - #[test] - fn test_smooth_l1_loss_matches_huber() { - let predictions = create_test_tensor_f32(vec![0.5, 2.0], vec![2], false); - let targets = create_test_tensor_f32(vec![0.0, 0.0], vec![2], false); - - let smooth = smooth_l1_loss(&predictions, &targets, "none").unwrap(); - let huber = huber_loss(&predictions, &targets, 1.0, "none").unwrap(); - - let smooth_data = smooth.data().as_f32_slice().unwrap(); - let huber_data = huber.data().as_f32_slice().unwrap(); - assert!((smooth_data[0] - huber_data[0]).abs() < 1e-6); - assert!((smooth_data[1] - huber_data[1]).abs() < 1e-6); - } - - #[test] - fn test_log_cosh_loss_mean() { - let predictions = create_test_tensor_f32(vec![0.0, 1.0], vec![2], false); - let targets = create_test_tensor_f32(vec![0.0, 0.0], vec![2], false); - - let loss = log_cosh_loss(&predictions, &targets, "mean").unwrap(); - let loss_data = loss.data().as_f32_slice().unwrap(); - - let expected = (0.0f32.cosh().ln() + 1.0f32.cosh().ln()) / 2.0; - assert!((loss_data[0] - expected).abs() < 1e-6); - } - - #[test] - fn test_log_cosh_loss_invalid_reduction() { - let predictions = create_test_tensor_f32(vec![0.0], vec![1], false); - let targets = create_test_tensor_f32(vec![0.0], vec![1], false); - - let result = log_cosh_loss(&predictions, &targets, "invalid"); - assert!(result.is_err()); - } -} + assert_eq!(loss.shape().dims(), &[] as &[usize]); + } + + #[test] + fn test_mse_loss_sum() { + let predictions = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + let targets = create_test_tensor_f32(vec![2.0, 3.0], vec![2], false); + + let loss = mse_loss(&predictions, &targets, "sum").unwrap(); + let loss_data = loss.data().as_f32_slice().unwrap(); + + // Expected: (1.0-2.0)² + (2.0-3.0)² = 1.0 + 1.0 = 2.0 + assert!((loss_data[0] - 2.0).abs() < 1e-6); + } + + #[test] + fn test_mse_loss_none() { + let predictions = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + let targets = create_test_tensor_f32(vec![2.0, 3.0], vec![2], false); + + let loss = mse_loss(&predictions, &targets, "none").unwrap(); + let loss_data = loss.data().as_f32_slice().unwrap(); + + // Expected: [(1.0-2.0)², (2.0-3.0)²] = [1.0, 1.0] + assert!((loss_data[0] - 1.0).abs() < 1e-6); + assert!((loss_data[1] - 1.0).abs() < 1e-6); + assert_eq!(loss.shape().dims(), &[2]); + } + + #[test] + fn test_mae_loss_mean() { + let predictions = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], false); + let targets = create_test_tensor_f32(vec![1.5, 2.5, 2.0], vec![3], false); + + let loss = mae_loss(&predictions, &targets, "mean").unwrap(); + let loss_data = loss.data().as_f32_slice().unwrap(); + + // Expected: (|1.0-1.5| + |2.0-2.5| + |3.0-2.0|) / 3 = (0.5 + 0.5 + 1.0) / 3 = 2.0/3 ≈ 0.667 + assert!((loss_data[0] - (2.0 / 3.0)).abs() < 1e-6); + } + + #[test] + fn test_huber_loss_quadratic_region() { + let predictions = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + let targets = create_test_tensor_f32(vec![1.2, 2.3], vec![2], false); + + // Delta = 1.0, differences are 0.2 and 0.3, both <= 1.0, so quadratic + let loss = huber_loss(&predictions, &targets, 1.0, "none").unwrap(); + let loss_data = loss.data().as_f32_slice().unwrap(); + + // Expected: [0.5 * 0.2², 0.5 * 0.3²] = [0.02, 0.045] + assert!((loss_data[0] - 0.02).abs() < 1e-6); + assert!((loss_data[1] - 0.045).abs() < 1e-6); + } + + #[test] + fn test_huber_loss_linear_region() { + let predictions = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + let targets = create_test_tensor_f32(vec![3.0, 0.0], vec![2], false); + + // Delta = 1.0, differences are 2.0 and 2.0, both > 1.0, so linear + let loss = huber_loss(&predictions, &targets, 1.0, "none").unwrap(); + let loss_data = loss.data().as_f32_slice().unwrap(); + + // Expected: [1.0 * (2.0 - 0.5 * 1.0), 1.0 * (2.0 - 0.5 * 1.0)] = [1.5, 1.5] + assert!((loss_data[0] - 1.5).abs() < 1e-6); + assert!((loss_data[1] - 1.5).abs() < 1e-6); + } + + #[test] + fn test_bce_loss_mean_and_backward() { + let predictions = create_test_tensor_f32(vec![0.8, 0.2], vec![2], true); + let targets = create_test_tensor_f32(vec![1.0, 0.0], vec![2], false); + + let loss = binary_cross_entropy_loss(&predictions, &targets, "mean").unwrap(); + let loss_val = loss.data().as_f32_slice().unwrap()[0]; + let expected = -((0.8f32).ln() + (0.8f32).ln()) / 2.0; + assert!((loss_val - expected).abs() < 1e-6); + + let grads = crate::autograd::backward_collect(&loss, None).unwrap(); + let grad = grads.get(&predictions.id()).unwrap(); + let grad_slice = grad.data().as_f32_slice().unwrap(); + let expected_grad = [-(1.0 / 0.8) / 2.0, (1.0 / 0.8) / 2.0]; + assert!((grad_slice[0] - expected_grad[0]).abs() < 1e-6); + assert!((grad_slice[1] - expected_grad[1]).abs() < 1e-6); + } + + #[test] + fn test_kl_div_loss_mean_and_backward() { + let predictions = create_test_tensor_f32(vec![0.4, 0.6], vec![2], true); + let targets = create_test_tensor_f32(vec![0.5, 0.5], vec![2], false); + + let loss = kl_div_loss(&predictions, &targets, "mean").unwrap(); + let loss_val = loss.data().as_f32_slice().unwrap()[0]; + let expected = 0.5 * ((0.5f32.ln() - 0.4f32.ln()) + (0.5f32.ln() - 0.6f32.ln())); + assert!((loss_val - expected).abs() < 1e-6); + + let grads = crate::autograd::backward_collect(&loss, None).unwrap(); + let grad = grads.get(&predictions.id()).unwrap(); + let grad_slice = grad.data().as_f32_slice().unwrap(); + let expected_grad = [-(0.5 / 0.4) / 2.0, -(0.5 / 0.6) / 2.0]; + assert!((grad_slice[0] - expected_grad[0]).abs() < 1e-6); + assert!((grad_slice[1] - expected_grad[1]).abs() < 1e-6); + } + + #[test] + fn test_loss_gradient_tracking() { + let predictions = create_test_tensor_f32(vec![1.0, 2.0], vec![2], true); + let targets = create_test_tensor_f32(vec![1.5, 2.5], vec![2], false); + + let loss = mse_loss(&predictions, &targets, "mean").unwrap(); + + assert!(loss.requires_grad()); + assert!(loss.grad_fn().is_some()); + } + + #[test] + fn test_loss_input_validation() { + let predictions = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + let targets = create_test_tensor_f32(vec![1.5, 2.5, 3.5], vec![3], false); + + // Shape mismatch should fail + let result = mse_loss(&predictions, &targets, "mean"); + assert!(result.is_err()); + } + + #[test] + fn test_invalid_reduction_mode() { + let predictions = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + let targets = create_test_tensor_f32(vec![1.5, 2.5], vec![2], false); + + let result = mse_loss(&predictions, &targets, "invalid"); + assert!(result.is_err()); + } + + #[test] + fn test_huber_loss_invalid_delta() { + let predictions = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + let targets = create_test_tensor_f32(vec![1.5, 2.5], vec![2], false); + + let result = huber_loss(&predictions, &targets, -1.0, "mean"); + assert!(result.is_err()); + } + + #[test] + fn test_smooth_l1_loss_matches_huber() { + let predictions = create_test_tensor_f32(vec![0.5, 2.0], vec![2], false); + let targets = create_test_tensor_f32(vec![0.0, 0.0], vec![2], false); + + let smooth = smooth_l1_loss(&predictions, &targets, "none").unwrap(); + let huber = huber_loss(&predictions, &targets, 1.0, "none").unwrap(); + + let smooth_data = smooth.data().as_f32_slice().unwrap(); + let huber_data = huber.data().as_f32_slice().unwrap(); + assert!((smooth_data[0] - huber_data[0]).abs() < 1e-6); + assert!((smooth_data[1] - huber_data[1]).abs() < 1e-6); + } + + #[test] + fn test_log_cosh_loss_mean() { + let predictions = create_test_tensor_f32(vec![0.0, 1.0], vec![2], false); + let targets = create_test_tensor_f32(vec![0.0, 0.0], vec![2], false); + + let loss = log_cosh_loss(&predictions, &targets, "mean").unwrap(); + let loss_data = loss.data().as_f32_slice().unwrap(); + + let expected = (0.0f32.cosh().ln() + 1.0f32.cosh().ln()) / 2.0; + assert!((loss_data[0] - expected).abs() < 1e-6); + } + + #[test] + fn test_log_cosh_loss_invalid_reduction() { + let predictions = create_test_tensor_f32(vec![0.0], vec![1], false); + let targets = create_test_tensor_f32(vec![0.0], vec![1], false); + + let result = log_cosh_loss(&predictions, &targets, "invalid"); + assert!(result.is_err()); + } +} diff --git a/engine/src/operations/loss/regression.rs b/engine/src/operations/loss/regression.rs index fa204d4a..c0937e81 100644 --- a/engine/src/operations/loss/regression.rs +++ b/engine/src/operations/loss/regression.rs @@ -1,921 +1,928 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -use crate::{ - autograd::{ - BCELossBackward, CrossEntropyLossBackward, FocalLossBackward, HuberLossBackward, - KLDivLossBackward, MAELossBackward, MSELossBackward, add_to_graph, - }, - error::{MinitensorError, Result}, - operations::{ - activation::{abs as activation_abs, exp, log_softmax, log1p}, - arithmetic::{add, mul, sub}, - comparison, - reduction::{mean, sum}, - selection::masked_fill_scalar, - }, - tensor::{DataType, Shape, Tensor, TensorData}, -}; -use rayon::prelude::*; -use std::sync::Arc; - -const CHUNK: usize = 1024; - -/// Mean Squared Error (MSE) loss function -/// -/// Computes the mean squared error between predictions and targets: -/// MSE = (1/n) * Σ(predictions - targets)² -/// -/// # Arguments -/// * `predictions` - Model predictions tensor -/// * `targets` - Ground truth targets tensor -/// * `reduction` - How to reduce the loss ("mean", "sum", or "none") -/// -/// # Returns -/// * `Result` - The computed MSE loss -pub fn mse_loss(predictions: &Tensor, targets: &Tensor, reduction: &str) -> Result { - // Validate inputs - validate_loss_inputs(predictions, targets)?; - - // Compute squared differences: (predictions - targets)² - // Also keep the difference for gradient computation - let diff = sub(predictions, targets)?; - let diff_for_grad = diff.clone().detach(); - let squared_diff = mul(&diff, &diff)?; - - // Apply reduction - let loss = match reduction { - "mean" => { - // Compute mean of squared differences - let sum = sum_all_elements(&squared_diff)?; - let n = squared_diff.numel() as f64; - divide_by_scalar(&sum, n)? - } - "sum" => { - // Sum all squared differences - sum_all_elements(&squared_diff)? - } - "none" => { - // Return element-wise squared differences - squared_diff - } - _ => { - return Err(MinitensorError::invalid_operation(format!( - "Invalid reduction mode: {}. Must be 'mean', 'sum', or 'none'", - reduction - ))); - } - }; - - // Set up gradient function if needed - if loss.requires_grad() { - let grad_fn = Arc::new(MSELossBackward { - predictions_shape: predictions.shape().dims().to_vec(), - targets_shape: targets.shape().dims().to_vec(), - input_ids: [predictions.id(), targets.id()], - reduction: reduction.to_string(), - diff: diff_for_grad, - }); - - let mut loss_with_grad = loss; - loss_with_grad.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&loss_with_grad, Some(grad_fn))?; - - Ok(loss_with_grad) - } else { - Ok(loss) - } -} - -/// Mean Absolute Error (MAE) loss function -/// -/// Computes the mean absolute error between predictions and targets: -/// MAE = (1/n) * Σ|predictions - targets| -/// -/// # Arguments -/// * `predictions` - Model predictions tensor -/// * `targets` - Ground truth targets tensor -/// * `reduction` - How to reduce the loss ("mean", "sum", or "none") -/// -/// # Returns -/// * `Result` - The computed MAE loss -pub fn mae_loss(predictions: &Tensor, targets: &Tensor, reduction: &str) -> Result { - // Validate inputs - validate_loss_inputs(predictions, targets)?; - - // Compute absolute differences: |predictions - targets| - // Also compute the sign for gradient computation - let diff = sub(predictions, targets)?; - let sign_diff = sign(&diff)?; - let sign_for_grad = sign_diff.clone().detach(); - let abs_diff = activation_abs(&diff.detach())?; - - // Apply reduction - let loss = match reduction { - "mean" => { - // Compute mean of absolute differences - let sum = sum_all_elements(&abs_diff)?; - let n = abs_diff.numel() as f64; - divide_by_scalar(&sum, n)? - } - "sum" => { - // Sum all absolute differences - sum_all_elements(&abs_diff)? - } - "none" => { - // Return element-wise absolute differences - abs_diff - } - _ => { - return Err(MinitensorError::invalid_operation(format!( - "Invalid reduction mode: {}. Must be 'mean', 'sum', or 'none'", - reduction - ))); - } - }; - - // Set up gradient function if needed. The forward is computed on detached - // data (the exact gradient is provided by MAELossBackward from the stored - // sign), so gate on the inputs and enable grad on the loss explicitly. - if predictions.requires_grad() || targets.requires_grad() { - let grad_fn = Arc::new(MAELossBackward { - predictions_shape: predictions.shape().dims().to_vec(), - targets_shape: targets.shape().dims().to_vec(), - input_ids: [predictions.id(), targets.id()], - reduction: reduction.to_string(), - sign: sign_for_grad, - }); - - let mut loss_with_grad = loss.requires_grad_(true); - loss_with_grad.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&loss_with_grad, Some(grad_fn))?; - - Ok(loss_with_grad) - } else { - Ok(loss) - } -} - -/// Cross Entropy loss function for classification -/// -/// Computes the cross entropy loss between predictions (logits) and targets: -/// CE = -Σ(targets * log(softmax(predictions))) -/// -/// # Arguments -/// * `predictions` - Model predictions (logits) tensor -/// * `targets` - Ground truth targets tensor (class indices or one-hot) -/// * `reduction` - How to reduce the loss ("mean", "sum", or "none") -/// -/// # Returns -/// * `Result` - The computed cross entropy loss -pub fn cross_entropy_loss( - predictions: &Tensor, - targets: &Tensor, - reduction: &str, -) -> Result { - // Validate inputs - validate_classification_inputs(predictions, targets, false)?; - - // Convert class indices to one-hot encoding if needed - let targets_one_hot = prepare_classification_targets(predictions, targets)?; - - // Apply log-softmax to predictions for numerical stability - let log_predictions_base = log_softmax(predictions, None)?; - let softmax_predictions = exp(&log_predictions_base.detach())?; - let eps = match softmax_predictions.dtype() { - DataType::Float32 => 1e-30, - DataType::Float64 => 1e-300, - _ => 0.0, - }; - let eps_template = Tensor::zeros( - softmax_predictions.shape().clone(), - softmax_predictions.dtype(), - softmax_predictions.device(), - false, - ); - let eps_mask = comparison::eq(&eps_template, &eps_template)?; - let eps_tensor = masked_fill_scalar(&eps_template, &eps_mask, eps)?; - let zero_mask = comparison::le(&softmax_predictions, &eps_tensor)?; - let log_predictions = - masked_fill_scalar(&log_predictions_base, &zero_mask, f64::NEG_INFINITY)?; - - // Compute negative log likelihood summed over classes - let nll = negative_log_likelihood(&log_predictions, &targets_one_hot)?; - let per_sample = sum(&nll, Some(vec![1]), false)?; - - // Apply reduction - let loss = match reduction { - "mean" => { - let sum = sum_all_elements(&per_sample)?; - let batch = per_sample.shape().dims().first().copied().unwrap_or(1) as f64; - divide_by_scalar(&sum, batch)? - } - "sum" => sum_all_elements(&per_sample)?, - "none" => per_sample, - _ => { - return Err(MinitensorError::invalid_operation(format!( - "Invalid reduction mode: {}. Must be 'mean', 'sum', or 'none'", - reduction - ))); - } - }; - - // Set up gradient function if needed - if loss.requires_grad() { - let grad_fn = Arc::new(CrossEntropyLossBackward { - predictions_shape: predictions.shape().dims().to_vec(), - targets_shape: targets_one_hot.shape().dims().to_vec(), - input_ids: [predictions.id(), targets.id()], - reduction: reduction.to_string(), - softmax_predictions: softmax_predictions.clone().detach(), - targets: targets_one_hot.clone().detach(), - }); - - let mut loss_with_grad = loss; - loss_with_grad.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&loss_with_grad, Some(grad_fn))?; - - Ok(loss_with_grad) - } else { - Ok(loss) - } -} - -/// Cross entropy loss for tensors with arbitrary shapes and class dimension. -/// -/// This wrapper permutes and flattens the input so that the core -/// `cross_entropy_loss` implementation can operate on ``[N, C]`` shaped -/// tensors entirely in Rust. -pub fn cross_entropy( - input: &Tensor, - target: &Tensor, - reduction: &str, - dim: usize, -) -> Result { - let ndim = input.ndim(); - if dim >= ndim { - return Err(MinitensorError::invalid_operation( - "dim out of range in cross_entropy", - )); - } - - // Move class dimension to the end using successive transposes - let mut pred = input.clone(); - let mut tgt = target.clone(); - if dim != ndim - 1 { - for i in dim..(ndim - 1) { - pred = pred.transpose(i as isize, (i + 1) as isize)?; - if target.ndim() == ndim { - tgt = tgt.transpose(i as isize, (i + 1) as isize)?; - } - } - } - - // Flatten all but the class dimension - let flat_size: usize = pred.shape().dims().iter().take(ndim - 1).product(); - let classes = pred.shape().dims()[ndim - 1]; - let pred_2d = pred.reshape(Shape::new(vec![flat_size, classes]))?; - let tgt_flat = if tgt.ndim() == ndim { - tgt.reshape(Shape::new(vec![flat_size, classes]))? - } else { - tgt.reshape(Shape::new(vec![flat_size]))? - }; - - let loss = cross_entropy_loss(&pred_2d, &tgt_flat, reduction)?; - - if reduction == "none" { - // Restore the original shape without the class dimension - let out_shape: Vec = input - .shape() - .dims() - .iter() - .enumerate() - .filter_map(|(i, &d)| if i != dim { Some(d) } else { None }) - .collect(); - loss.reshape(Shape::new(out_shape)) - } else { - Ok(loss) - } -} - -/// Binary Cross Entropy loss function -/// -/// Computes the binary cross entropy loss between predictions and targets: -/// BCE = -Σ(targets * log(predictions) + (1 - targets) * log(1 - predictions)) -/// -/// # Arguments -/// * `predictions` - Model predictions tensor (probabilities between 0 and 1) -/// * `targets` - Ground truth targets tensor (0 or 1) -/// * `reduction` - How to reduce the loss ("mean", "sum", or "none") -/// -/// # Returns -/// * `Result` - The computed BCE loss -pub fn binary_cross_entropy_loss( - predictions: &Tensor, - targets: &Tensor, - reduction: &str, -) -> Result { - // Validate inputs - validate_loss_inputs(predictions, targets)?; - - // Compute BCE: -[targets * log(predictions) + (1 - targets) * log(1 - predictions)] - let log_predictions = log(predictions)?; - - let ones = Tensor::ones( - predictions.shape().clone(), - predictions.dtype(), - predictions.device(), - false, - ); - let one_minus_targets = sub(&ones, targets)?; - let one_minus_predictions = sub(&ones, predictions)?; - let log_one_minus_predictions = log(&one_minus_predictions)?; - - let term1 = mul(targets, &log_predictions)?; - let term2 = mul(&one_minus_targets, &log_one_minus_predictions)?; - let combined = add(&term1, &term2)?; - let zeros = Tensor::zeros( - combined.shape().clone(), - combined.dtype(), - combined.device(), - combined.requires_grad(), - ); - let negative_bce = sub(&zeros, &combined)?; - - // Apply reduction - let loss = match reduction { - "mean" => { - let sum = sum_all_elements(&negative_bce)?; - let n = negative_bce.numel() as f64; - divide_by_scalar(&sum, n)? - } - "sum" => sum_all_elements(&negative_bce)?, - "none" => negative_bce, - _ => { - return Err(MinitensorError::invalid_operation(format!( - "Invalid reduction mode: {}. Must be 'mean', 'sum', or 'none'", - reduction - ))); - } - }; - - // Set up gradient function if needed - if loss.requires_grad() { - let grad_fn = Arc::new(BCELossBackward { - predictions_shape: predictions.shape().dims().to_vec(), - targets_shape: targets.shape().dims().to_vec(), - input_ids: [predictions.id(), targets.id()], - reduction: reduction.to_string(), - predictions: predictions.clone().detach(), - targets: targets.clone().detach(), - }); - - let mut loss_with_grad = loss; - loss_with_grad.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&loss_with_grad, Some(grad_fn))?; - - Ok(loss_with_grad) - } else { - Ok(loss) - } -} - -/// Kullback-Leibler divergence loss function -/// -/// Computes KL divergence between target and prediction distributions: -/// KL(target || prediction) = Σ target * (log(target) - log(prediction)) -pub fn kl_div_loss(predictions: &Tensor, targets: &Tensor, reduction: &str) -> Result { - // Validate inputs - validate_loss_inputs(predictions, targets)?; - - // Compute elementwise targets * (log(targets) - log(predictions)) - let log_targets = log(targets)?; - let log_predictions = log(predictions)?; - let diff = sub(&log_targets, &log_predictions)?; - let kld = mul(targets, &diff)?; - - // Apply reduction - let loss = match reduction { - "mean" => { - let sum = sum_all_elements(&kld)?; - // Compute mean over the batch dimension if present. - // For 1D tensors (single distribution), the batch size is 1 - let batch = if predictions.shape().dims().len() > 1 { - predictions.shape().dims()[0] as f64 - } else { - 1.0 - }; - divide_by_scalar(&sum, batch)? - } - "sum" => sum_all_elements(&kld)?, - "none" => kld, - _ => { - return Err(MinitensorError::invalid_operation(format!( - "Invalid reduction mode: {}. Must be 'mean', 'sum', or 'none'", - reduction - ))); - } - }; - - // Set up gradient function if needed - if loss.requires_grad() { - let grad_fn = Arc::new(KLDivLossBackward { - predictions_shape: predictions.shape().dims().to_vec(), - targets_shape: targets.shape().dims().to_vec(), - input_ids: [predictions.id(), targets.id()], - reduction: reduction.to_string(), - predictions: predictions.clone().detach(), - targets: targets.clone().detach(), - }); - - let mut loss_with_grad = loss; - loss_with_grad.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&loss_with_grad, Some(grad_fn))?; - - Ok(loss_with_grad) - } else { - Ok(loss) - } -} - -/// Focal loss function for handling class imbalance -/// -/// Computes the focal loss, which is a modified cross entropy loss: -/// FL = -α * (1 - p_t)^γ * log(p_t) -/// where p_t is the predicted probability for the true class -/// -/// # Arguments -/// * `predictions` - Model predictions (logits) tensor -/// * `targets` - Ground truth targets tensor -/// * `alpha` - Weighting factor for rare class (typically 0.25) -/// * `gamma` - Focusing parameter (typically 2.0) -/// * `reduction` - How to reduce the loss ("mean", "sum", or "none") -/// -/// # Returns -/// * `Result` - The computed focal loss -pub fn focal_loss( - predictions: &Tensor, - targets: &Tensor, - alpha: f64, - gamma: f64, - reduction: &str, -) -> Result { - // Validate inputs - validate_classification_inputs(predictions, targets, false)?; - - let targets_one_hot = prepare_classification_targets(predictions, targets)?; - - if alpha <= 0.0 || alpha >= 1.0 { - return Err(MinitensorError::invalid_operation( - "Alpha must be between 0 and 1 for focal loss", - )); - } - - if gamma < 0.0 { - return Err(MinitensorError::invalid_operation( - "Gamma must be non-negative for focal loss", - )); - } - - // Apply log-softmax to predictions for numerical stability - let log_predictions = log_softmax(predictions, None)?; - let softmax_predictions = exp(&log_predictions)?; - let softmax_for_grad = softmax_predictions.clone().detach(); - - // Compute focal loss components - let ones = Tensor::ones( - softmax_predictions.shape().clone(), - softmax_predictions.dtype(), - softmax_predictions.device(), - false, - ); - let one_minus_p = sub(&ones, &softmax_predictions)?; - let focal_weight = power(&one_minus_p, gamma)?; - - // Compute negative log likelihood with focal weighting - let nll = negative_log_likelihood(&log_predictions, &targets_one_hot)?; - let alpha_tensor = create_scalar_tensor(alpha, predictions.dtype(), predictions.device())?; - let weighted_nll = mul(&nll, &focal_weight)?; - let focal_values = mul(&weighted_nll, &alpha_tensor)?; - - // Apply reduction - let loss = match reduction { - "mean" => { - let sum = sum_all_elements(&focal_values)?; - // Average over samples, matching cross_entropy: only the true-class - // term per sample is non-zero, so the denominator is the number of - // samples (numel / num_classes), not the total element count. - let num_classes = predictions.size(predictions.ndim() - 1)?.max(1); - let n = (focal_values.numel() / num_classes) as f64; - divide_by_scalar(&sum, n)? - } - "sum" => sum_all_elements(&focal_values)?, - "none" => focal_values, - _ => { - return Err(MinitensorError::invalid_operation(format!( - "Invalid reduction mode: {}. Must be 'mean', 'sum', or 'none'", - reduction - ))); - } - }; - - // Set up gradient function if needed - if loss.requires_grad() { - let grad_fn = Arc::new(FocalLossBackward { - predictions_shape: predictions.shape().dims().to_vec(), - targets_shape: targets_one_hot.shape().dims().to_vec(), - input_ids: [predictions.id(), targets.id()], - alpha, - gamma, - reduction: reduction.to_string(), - softmax_predictions: softmax_for_grad, - targets: targets_one_hot.clone().detach(), - }); - - let mut loss_with_grad = loss; - loss_with_grad.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&loss_with_grad, Some(grad_fn))?; - - Ok(loss_with_grad) - } else { - Ok(loss) - } -} - -/// Huber loss function for robust regression -/// -/// Combines MSE and MAE for robust regression: -/// - For |x| <= delta: 0.5 * x² -/// - For |x| > delta: delta * (|x| - 0.5 * delta) -/// -/// # Arguments -/// * `predictions` - Model predictions tensor -/// * `targets` - Ground truth targets tensor -/// * `delta` - Threshold for switching between MSE and MAE behavior -/// * `reduction` - How to reduce the loss ("mean", "sum", or "none") -/// -/// # Returns -/// * `Result` - The computed Huber loss -pub fn huber_loss( - predictions: &Tensor, - targets: &Tensor, - delta: f64, - reduction: &str, -) -> Result { - // Validate inputs - validate_loss_inputs(predictions, targets)?; - - if delta <= 0.0 { - return Err(MinitensorError::invalid_operation( - "Delta must be positive for Huber loss", - )); - } - - // Compute absolute differences: |predictions - targets| - let diff = sub(predictions, targets)?; - let diff_for_grad = diff.clone().detach(); - let abs_diff = activation_abs(&diff.detach())?; - - // Create delta tensor for comparison - let delta_tensor = create_scalar_tensor(delta, predictions.dtype(), predictions.device())?; - - // Compute Huber loss element-wise - let huber_values = compute_huber_elementwise(&abs_diff, &diff, &delta_tensor, delta)?; - - // Apply reduction - let loss = match reduction { - "mean" => { - let sum = sum_all_elements(&huber_values)?; - let n = huber_values.numel() as f64; - divide_by_scalar(&sum, n)? - } - "sum" => sum_all_elements(&huber_values)?, - "none" => huber_values, - _ => { - return Err(MinitensorError::invalid_operation(format!( - "Invalid reduction mode: {}. Must be 'mean', 'sum', or 'none'", - reduction - ))); - } - }; - - // Set up gradient function if needed. The forward is computed on detached - // data (the exact gradient is provided by HuberLossBackward from the stored - // diff), so gate on the inputs and enable grad on the loss explicitly. - if predictions.requires_grad() || targets.requires_grad() { - let grad_fn = Arc::new(HuberLossBackward { - predictions_shape: predictions.shape().dims().to_vec(), - targets_shape: targets.shape().dims().to_vec(), - input_ids: [predictions.id(), targets.id()], - delta, - reduction: reduction.to_string(), - diff: diff_for_grad, - }); - - let mut loss_with_grad = loss.requires_grad_(true); - loss_with_grad.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&loss_with_grad, Some(grad_fn))?; - - Ok(loss_with_grad) - } else { - Ok(loss) - } -} - -/// Smooth L1 loss (Huber loss with delta=1.0) -/// -/// Computes Smooth L1 loss between predictions and targets: -/// SmoothL1(x) = 0.5 * x² if |x| < 1, otherwise |x| - 0.5 -/// -/// # Arguments -/// * `predictions` - Model predictions tensor -/// * `targets` - Ground truth targets tensor -/// * `reduction` - How to reduce the loss ("mean", "sum", or "none") -pub fn smooth_l1_loss(predictions: &Tensor, targets: &Tensor, reduction: &str) -> Result { - huber_loss(predictions, targets, 1.0, reduction) -} - -/// Log-cosh loss for robust regression -/// -/// Computes log(cosh(x)) where x = predictions - targets using a numerically -/// stable formulation: |x| + log1p(exp(-2|x|)) - log(2). -/// -/// # Arguments -/// * `predictions` - Model predictions tensor -/// * `targets` - Ground truth targets tensor -/// * `reduction` - How to reduce the loss ("mean", "sum", or "none") -pub fn log_cosh_loss(predictions: &Tensor, targets: &Tensor, reduction: &str) -> Result { - validate_loss_inputs(predictions, targets)?; - - let diff = sub(predictions, targets)?; - let diff_abs = activation_abs(&diff)?; - let neg_two = create_scalar_tensor(-2.0, diff.dtype(), diff.device())?; - let exp_term = exp(&mul(&diff_abs, &neg_two)?)?; - let log1p_term = log1p(&exp_term)?; - let log2 = create_scalar_tensor(std::f64::consts::LN_2, diff.dtype(), diff.device())?; - let log_cosh = sub(&add(&diff_abs, &log1p_term)?, &log2)?; - - match reduction { - "mean" => mean(&log_cosh, None, false), - "sum" => sum(&log_cosh, None, false), - "none" => Ok(log_cosh), - _ => Err(MinitensorError::invalid_operation(format!( - "Invalid reduction mode: {}. Must be 'mean', 'sum', or 'none'", - reduction - ))), - } -} - -// Helper functions - -/// Validate that loss function inputs are compatible -fn validate_loss_inputs(predictions: &Tensor, targets: &Tensor) -> Result<()> { - // Check device compatibility - if predictions.device() != targets.device() { - return Err(MinitensorError::device_mismatch( - format!("{:?}", predictions.device()), - format!("{:?}", targets.device()), - )); - } - - // Check data type compatibility - if predictions.dtype() != targets.dtype() { - return Err(MinitensorError::type_mismatch( - format!("{:?}", predictions.dtype()), - format!("{:?}", targets.dtype()), - )); - } - - // Check shape compatibility - if predictions.shape() != targets.shape() { - return Err(MinitensorError::shape_mismatch( - predictions.shape().dims().to_vec(), - targets.shape().dims().to_vec(), - )); - } - - // Check that tensors contain floating point data (required for loss computation) - match predictions.dtype() { - DataType::Float32 | DataType::Float64 => {} - _ => { - return Err(MinitensorError::invalid_operation( - "Loss functions require floating point tensors", - )); - } - } - - Ok(()) -} - -/// Validate that classification loss function inputs are compatible -fn validate_classification_inputs( - predictions: &Tensor, - targets: &Tensor, - require_same_dtype: bool, -) -> Result<()> { - // Check device compatibility - if predictions.device() != targets.device() { - return Err(MinitensorError::device_mismatch( - format!("{:?}", predictions.device()), - format!("{:?}", targets.device()), - )); - } - - // Optionally enforce data type equality - if require_same_dtype && predictions.dtype() != targets.dtype() { - return Err(MinitensorError::type_mismatch( - format!("{:?}", predictions.dtype()), - format!("{:?}", targets.dtype()), - )); - } - - // Predictions must be at least 2D (batch_size, num_classes) - if predictions.ndim() < 2 { - return Err(MinitensorError::invalid_operation( - "Classification predictions must be at least 2D (batch_size, num_classes)", - )); - } - - // Predictions must be floating point - match predictions.dtype() { - DataType::Float32 | DataType::Float64 => {} - _ => { - return Err(MinitensorError::invalid_operation( - "Classification loss functions require floating point tensors", - )); - } - } - - Ok(()) -} - -fn prepare_classification_targets(predictions: &Tensor, targets: &Tensor) -> Result { - if targets.ndim() + 1 == predictions.ndim() { - let num_classes = predictions.size(predictions.ndim() - 1)?; - let total = targets.numel(); - let mut data = TensorData::zeros_on_device( - total * num_classes, - predictions.dtype(), - predictions.device(), - ); - match (targets.dtype(), predictions.dtype()) { - (DataType::Int32, DataType::Float32) => { - let idx = targets.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from targets") - })?; - let out = data.as_f32_slice_mut().unwrap(); - fill_one_hot_f32(idx, out, num_classes, |val| { - checked_index_from_i64(i64::from(*val), num_classes) - })?; - } - (DataType::Int64, DataType::Float32) => { - let idx = targets.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from targets") - })?; - let out = data.as_f32_slice_mut().unwrap(); - fill_one_hot_f32(idx, out, num_classes, |val| { - checked_index_from_i64(*val, num_classes) - })?; - } - (DataType::Int32, DataType::Float64) => { - let idx = targets.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from targets") - })?; - let out = data.as_f64_slice_mut().unwrap(); - fill_one_hot_f64(idx, out, num_classes, |val| { - checked_index_from_i64(i64::from(*val), num_classes) - })?; - } - (DataType::Int64, DataType::Float64) => { - let idx = targets.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from targets") - })?; - let out = data.as_f64_slice_mut().unwrap(); - fill_one_hot_f64(idx, out, num_classes, |val| { - checked_index_from_i64(*val, num_classes) - })?; - } - (DataType::Float32, DataType::Float32) => { - let idx = targets.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from targets") - })?; - let out = data.as_f32_slice_mut().unwrap(); - fill_one_hot_f32(idx, out, num_classes, |val| { - checked_index_from_f32(*val, num_classes) - })?; - } - (DataType::Float64, DataType::Float64) => { - let idx = targets.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from targets") - })?; - let out = data.as_f64_slice_mut().unwrap(); - fill_one_hot_f64(idx, out, num_classes, |val| { - checked_index_from_f64(*val, num_classes) - })?; - } - _ => { - return Err(MinitensorError::invalid_operation( - "Unsupported target dtype for classification loss", - )); - } - } - let mut dims = targets.shape().dims().to_vec(); - dims.push(num_classes); - Ok(Tensor::new( - Arc::new(data), - Shape::new(dims), - predictions.dtype(), - predictions.device(), - false, - )) - } else if targets.ndim() == predictions.ndim() { - if targets.shape().dims() != predictions.shape().dims() { - return Err(MinitensorError::shape_mismatch( - predictions.shape().dims().to_vec(), - targets.shape().dims().to_vec(), - )); - } - Ok(targets.clone()) - } else { - Err(MinitensorError::shape_mismatch( - predictions.shape().dims().to_vec(), - targets.shape().dims().to_vec(), - )) - } -} - -fn checked_index_from_i64(value: i64, num_classes: usize) -> Result { - if value < 0 { - return Err(MinitensorError::invalid_operation( - "Target class index must be non-negative", - )); - } - let index = value as usize; - if index >= num_classes { - return Err(MinitensorError::invalid_operation( - "Target class index out of range", - )); - } - Ok(index) -} - -fn checked_index_from_f32(value: f32, num_classes: usize) -> Result { - if !value.is_finite() || value.fract() != 0.0 { - return Err(MinitensorError::invalid_operation( - "Target class index must be a finite integer", - )); - } - if value < 0.0 || value >= num_classes as f32 { - return Err(MinitensorError::invalid_operation( - "Target class index out of range", - )); - } - Ok(value as usize) -} - -fn checked_index_from_f64(value: f64, num_classes: usize) -> Result { - if !value.is_finite() || value.fract() != 0.0 { - return Err(MinitensorError::invalid_operation( - "Target class index must be a finite integer", - )); - } - if value < 0.0 || value >= num_classes as f64 { - return Err(MinitensorError::invalid_operation( - "Target class index out of range", - )); - } - Ok(value as usize) -} - -fn fill_one_hot_f32( - indices: &[T], - out: &mut [f32], - num_classes: usize, - to_index: F, -) -> Result<()> -where - F: Fn(&T) -> Result, -{ - for (i, value) in indices.iter().enumerate() { - let class = to_index(value)?; - out[i * num_classes + class] = 1.0; - } - Ok(()) -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; + +use crate::{ + autograd::{ + BCELossBackward, CrossEntropyLossBackward, FocalLossBackward, HuberLossBackward, + KLDivLossBackward, MAELossBackward, MSELossBackward, add_to_graph, + }, + error::{MinitensorError, Result}, + operations::{ + activation::{abs as activation_abs, exp, log_softmax, log1p}, + arithmetic::{add, mul, sub}, + comparison, + reduction::{mean, sum}, + selection::masked_fill_scalar, + }, + tensor::{DataType, Shape, Tensor, TensorData}, +}; +use std::sync::Arc; + +pub(crate) const CHUNK: usize = 1024; + +/// Mean Squared Error (MSE) loss function +/// +/// Computes the mean squared error between predictions and targets: +/// MSE = (1/n) * Σ(predictions - targets)² +/// +/// # Arguments +/// * `predictions` - Model predictions tensor +/// * `targets` - Ground truth targets tensor +/// * `reduction` - How to reduce the loss ("mean", "sum", or "none") +/// +/// # Returns +/// * `Result` - The computed MSE loss +pub fn mse_loss(predictions: &Tensor, targets: &Tensor, reduction: &str) -> Result { + // Validate inputs + validate_loss_inputs(predictions, targets)?; + + // Compute squared differences: (predictions - targets)² + // Also keep the difference for gradient computation + let diff = sub(predictions, targets)?; + let diff_for_grad = diff.clone().detach(); + let squared_diff = mul(&diff, &diff)?; + + // Apply reduction + let loss = match reduction { + "mean" => { + // Compute mean of squared differences + let sum = sum_all_elements(&squared_diff)?; + let n = squared_diff.numel() as f64; + divide_by_scalar(&sum, n)? + } + "sum" => { + // Sum all squared differences + sum_all_elements(&squared_diff)? + } + "none" => { + // Return element-wise squared differences + squared_diff + } + _ => { + return Err(MinitensorError::invalid_operation(format!( + "Invalid reduction mode: {}. Must be 'mean', 'sum', or 'none'", + reduction + ))); + } + }; + + // Set up gradient function if needed + if loss.requires_grad() { + let grad_fn = Arc::new(MSELossBackward { + predictions_shape: predictions.shape().dims().to_vec(), + targets_shape: targets.shape().dims().to_vec(), + input_ids: [predictions.id(), targets.id()], + input_requires_grad: [predictions.requires_grad(), targets.requires_grad()], + reduction: reduction.to_string(), + diff: diff_for_grad, + }); + + let mut loss_with_grad = loss; + loss_with_grad.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&loss_with_grad, Some(grad_fn))?; + + Ok(loss_with_grad) + } else { + Ok(loss) + } +} + +/// Mean Absolute Error (MAE) loss function +/// +/// Computes the mean absolute error between predictions and targets: +/// MAE = (1/n) * Σ|predictions - targets| +/// +/// # Arguments +/// * `predictions` - Model predictions tensor +/// * `targets` - Ground truth targets tensor +/// * `reduction` - How to reduce the loss ("mean", "sum", or "none") +/// +/// # Returns +/// * `Result` - The computed MAE loss +pub fn mae_loss(predictions: &Tensor, targets: &Tensor, reduction: &str) -> Result { + // Validate inputs + validate_loss_inputs(predictions, targets)?; + + // Compute absolute differences: |predictions - targets| + // Also compute the sign for gradient computation + let diff = sub(predictions, targets)?; + let sign_diff = sign_tensor(&diff)?; + let sign_for_grad = sign_diff.clone().detach(); + let abs_diff = activation_abs(&diff.detach())?; + + // Apply reduction + let loss = match reduction { + "mean" => { + // Compute mean of absolute differences + let sum = sum_all_elements(&abs_diff)?; + let n = abs_diff.numel() as f64; + divide_by_scalar(&sum, n)? + } + "sum" => { + // Sum all absolute differences + sum_all_elements(&abs_diff)? + } + "none" => { + // Return element-wise absolute differences + abs_diff + } + _ => { + return Err(MinitensorError::invalid_operation(format!( + "Invalid reduction mode: {}. Must be 'mean', 'sum', or 'none'", + reduction + ))); + } + }; + + // Set up gradient function if needed. The forward is computed on detached + // data (the exact gradient is provided by MAELossBackward from the stored + // sign), so gate on the inputs and enable grad on the loss explicitly. + if predictions.requires_grad() || targets.requires_grad() { + let grad_fn = Arc::new(MAELossBackward { + predictions_shape: predictions.shape().dims().to_vec(), + targets_shape: targets.shape().dims().to_vec(), + input_ids: [predictions.id(), targets.id()], + input_requires_grad: [predictions.requires_grad(), targets.requires_grad()], + reduction: reduction.to_string(), + sign: sign_for_grad, + }); + + let mut loss_with_grad = loss.requires_grad_(true); + loss_with_grad.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&loss_with_grad, Some(grad_fn))?; + + Ok(loss_with_grad) + } else { + Ok(loss) + } +} + +/// Cross Entropy loss function for classification +/// +/// Computes the cross entropy loss between predictions (logits) and targets: +/// CE = -Σ(targets * log_tensor(softmax(predictions))) +/// +/// # Arguments +/// * `predictions` - Model predictions (logits) tensor +/// * `targets` - Ground truth targets tensor (class indices or one-hot) +/// * `reduction` - How to reduce the loss ("mean", "sum", or "none") +/// +/// # Returns +/// * `Result` - The computed cross entropy loss +pub fn cross_entropy_loss( + predictions: &Tensor, + targets: &Tensor, + reduction: &str, +) -> Result { + // Validate inputs + validate_classification_inputs(predictions, targets, false)?; + + // Convert class indices to one-hot encoding if needed + let targets_one_hot = prepare_classification_targets(predictions, targets)?; + + // Apply log-softmax to predictions for numerical stability + let log_predictions_base = log_softmax(predictions, None)?; + let softmax_predictions = exp(&log_predictions_base.detach())?; + let eps = match softmax_predictions.dtype() { + DataType::Float32 => 1e-30, + DataType::Float64 => 1e-300, + _ => 0.0, + }; + let eps_template = Tensor::zeros( + softmax_predictions.shape().clone(), + softmax_predictions.dtype(), + softmax_predictions.device(), + false, + ); + let eps_mask = comparison::eq(&eps_template, &eps_template)?; + let eps_tensor = masked_fill_scalar(&eps_template, &eps_mask, eps)?; + let zero_mask = comparison::le(&softmax_predictions, &eps_tensor)?; + let log_predictions = masked_fill_scalar(&log_predictions_base, &zero_mask, f64::NEG_INFINITY)?; + + // Compute negative log likelihood summed over classes + let nll = negative_log_likelihood(&log_predictions, &targets_one_hot)?; + let per_sample = sum(&nll, Some(vec![1]), false)?; + + // Apply reduction + let loss = match reduction { + "mean" => { + let sum = sum_all_elements(&per_sample)?; + let batch = per_sample.shape().dims().first().copied().unwrap_or(1) as f64; + divide_by_scalar(&sum, batch)? + } + "sum" => sum_all_elements(&per_sample)?, + "none" => per_sample, + _ => { + return Err(MinitensorError::invalid_operation(format!( + "Invalid reduction mode: {}. Must be 'mean', 'sum', or 'none'", + reduction + ))); + } + }; + + // Set up gradient function if needed + if loss.requires_grad() { + let grad_fn = Arc::new(CrossEntropyLossBackward { + predictions_shape: predictions.shape().dims().to_vec(), + targets_shape: targets_one_hot.shape().dims().to_vec(), + input_ids: [predictions.id(), targets.id()], + input_requires_grad: [predictions.requires_grad(), targets.requires_grad()], + reduction: reduction.to_string(), + softmax_predictions: softmax_predictions.clone().detach(), + targets: targets_one_hot.clone().detach(), + }); + + let mut loss_with_grad = loss; + loss_with_grad.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&loss_with_grad, Some(grad_fn))?; + + Ok(loss_with_grad) + } else { + Ok(loss) + } +} + +/// Cross entropy loss for tensors with arbitrary shapes and class dimension. +/// +/// This wrapper permutes and flattens the input so that the core +/// `cross_entropy_loss` implementation can operate on ``[N, C]`` shaped +/// tensors entirely in Rust. +pub fn cross_entropy( + input: &Tensor, + target: &Tensor, + reduction: &str, + dim: usize, +) -> Result { + let ndim = input.ndim(); + if dim >= ndim { + return Err(MinitensorError::invalid_operation( + "dim out of range in cross_entropy", + )); + } + + // Move class dimension to the end using successive transposes + let mut pred = input.clone(); + let mut tgt = target.clone(); + if dim != ndim - 1 { + for i in dim..(ndim - 1) { + pred = pred.transpose(i as isize, (i + 1) as isize)?; + if target.ndim() == ndim { + tgt = tgt.transpose(i as isize, (i + 1) as isize)?; + } + } + } + + // Flatten all but the class dimension + let flat_size: usize = pred.shape().dims().iter().take(ndim - 1).product(); + let classes = pred.shape().dims()[ndim - 1]; + let pred_2d = pred.reshape(Shape::new(vec![flat_size, classes]))?; + let tgt_flat = if tgt.ndim() == ndim { + tgt.reshape(Shape::new(vec![flat_size, classes]))? + } else { + tgt.reshape(Shape::new(vec![flat_size]))? + }; + + let loss = cross_entropy_loss(&pred_2d, &tgt_flat, reduction)?; + + if reduction == "none" { + // Restore the original shape without the class dimension + let out_shape: Vec = input + .shape() + .dims() + .iter() + .enumerate() + .filter_map(|(i, &d)| if i != dim { Some(d) } else { None }) + .collect(); + loss.reshape(Shape::new(out_shape)) + } else { + Ok(loss) + } +} + +/// Binary Cross Entropy loss function +/// +/// Computes the binary cross entropy loss between predictions and targets: +/// BCE = -Σ(targets * log_tensor(predictions) + (1 - targets) * log_tensor(1 - predictions)) +/// +/// # Arguments +/// * `predictions` - Model predictions tensor (probabilities between 0 and 1) +/// * `targets` - Ground truth targets tensor (0 or 1) +/// * `reduction` - How to reduce the loss ("mean", "sum", or "none") +/// +/// # Returns +/// * `Result` - The computed BCE loss +pub fn binary_cross_entropy_loss( + predictions: &Tensor, + targets: &Tensor, + reduction: &str, +) -> Result { + // Validate inputs + validate_loss_inputs(predictions, targets)?; + + // Compute BCE: -[targets * log_tensor(predictions) + (1 - targets) * log_tensor(1 - predictions)] + let log_predictions = log_tensor(predictions)?; + + let ones = Tensor::ones( + predictions.shape().clone(), + predictions.dtype(), + predictions.device(), + false, + ); + let one_minus_targets = sub(&ones, targets)?; + let one_minus_predictions = sub(&ones, predictions)?; + let log_one_minus_predictions = log_tensor(&one_minus_predictions)?; + + let term1 = mul(targets, &log_predictions)?; + let term2 = mul(&one_minus_targets, &log_one_minus_predictions)?; + let combined = add(&term1, &term2)?; + let zeros = Tensor::zeros( + combined.shape().clone(), + combined.dtype(), + combined.device(), + combined.requires_grad(), + ); + let negative_bce = sub(&zeros, &combined)?; + + // Apply reduction + let loss = match reduction { + "mean" => { + let sum = sum_all_elements(&negative_bce)?; + let n = negative_bce.numel() as f64; + divide_by_scalar(&sum, n)? + } + "sum" => sum_all_elements(&negative_bce)?, + "none" => negative_bce, + _ => { + return Err(MinitensorError::invalid_operation(format!( + "Invalid reduction mode: {}. Must be 'mean', 'sum', or 'none'", + reduction + ))); + } + }; + + // Set up gradient function if needed + if loss.requires_grad() { + let grad_fn = Arc::new(BCELossBackward { + predictions_shape: predictions.shape().dims().to_vec(), + targets_shape: targets.shape().dims().to_vec(), + input_ids: [predictions.id(), targets.id()], + input_requires_grad: [predictions.requires_grad(), targets.requires_grad()], + reduction: reduction.to_string(), + predictions: predictions.clone().detach(), + targets: targets.clone().detach(), + }); + + let mut loss_with_grad = loss; + loss_with_grad.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&loss_with_grad, Some(grad_fn))?; + + Ok(loss_with_grad) + } else { + Ok(loss) + } +} + +/// Kullback-Leibler divergence loss function +/// +/// Computes KL divergence between target and prediction distributions: +/// KL(target || prediction) = Σ target * (log_tensor(target) - log_tensor(prediction)) +pub fn kl_div_loss(predictions: &Tensor, targets: &Tensor, reduction: &str) -> Result { + // Validate inputs + validate_loss_inputs(predictions, targets)?; + + // Compute elementwise targets * (log_tensor(targets) - log_tensor(predictions)) + let log_targets = log_tensor(targets)?; + let log_predictions = log_tensor(predictions)?; + let diff = sub(&log_targets, &log_predictions)?; + let kld = mul(targets, &diff)?; + + // Apply reduction + let loss = match reduction { + "mean" => { + let sum = sum_all_elements(&kld)?; + // Compute mean over the batch dimension if present. + // For 1D tensors (single distribution), the batch size is 1 + let batch = if predictions.shape().dims().len() > 1 { + predictions.shape().dims()[0] as f64 + } else { + 1.0 + }; + divide_by_scalar(&sum, batch)? + } + "sum" => sum_all_elements(&kld)?, + "none" => kld, + _ => { + return Err(MinitensorError::invalid_operation(format!( + "Invalid reduction mode: {}. Must be 'mean', 'sum', or 'none'", + reduction + ))); + } + }; + + // Set up gradient function if needed + if loss.requires_grad() { + let grad_fn = Arc::new(KLDivLossBackward { + predictions_shape: predictions.shape().dims().to_vec(), + targets_shape: targets.shape().dims().to_vec(), + input_ids: [predictions.id(), targets.id()], + input_requires_grad: [predictions.requires_grad(), targets.requires_grad()], + reduction: reduction.to_string(), + predictions: predictions.clone().detach(), + targets: targets.clone().detach(), + }); + + let mut loss_with_grad = loss; + loss_with_grad.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&loss_with_grad, Some(grad_fn))?; + + Ok(loss_with_grad) + } else { + Ok(loss) + } +} + +/// Focal loss function for handling class imbalance +/// +/// Computes the focal loss, which is a modified cross entropy loss: +/// FL = -α * (1 - p_t)^γ * log_tensor(p_t) +/// where p_t is the predicted probability for the true class +/// +/// # Arguments +/// * `predictions` - Model predictions (logits) tensor +/// * `targets` - Ground truth targets tensor +/// * `alpha` - Weighting factor for rare class (typically 0.25) +/// * `gamma` - Focusing parameter (typically 2.0) +/// * `reduction` - How to reduce the loss ("mean", "sum", or "none") +/// +/// # Returns +/// * `Result` - The computed focal loss +pub fn focal_loss( + predictions: &Tensor, + targets: &Tensor, + alpha: f64, + gamma: f64, + reduction: &str, +) -> Result { + // Validate inputs + validate_classification_inputs(predictions, targets, false)?; + + let targets_one_hot = prepare_classification_targets(predictions, targets)?; + + if alpha <= 0.0 || alpha >= 1.0 { + return Err(MinitensorError::invalid_operation( + "Alpha must be between 0 and 1 for focal loss", + )); + } + + if gamma < 0.0 { + return Err(MinitensorError::invalid_operation( + "Gamma must be non-negative for focal loss", + )); + } + + // Apply log-softmax to predictions for numerical stability + let log_predictions = log_softmax(predictions, None)?; + let softmax_predictions = exp(&log_predictions)?; + let softmax_for_grad = softmax_predictions.clone().detach(); + + // Compute focal loss components + let ones = Tensor::ones( + softmax_predictions.shape().clone(), + softmax_predictions.dtype(), + softmax_predictions.device(), + false, + ); + let one_minus_p = sub(&ones, &softmax_predictions)?; + let focal_weight = power(&one_minus_p, gamma)?; + + // Compute negative log likelihood with focal weighting + let nll = negative_log_likelihood(&log_predictions, &targets_one_hot)?; + let alpha_tensor = create_scalar_tensor(alpha, predictions.dtype(), predictions.device())?; + let weighted_nll = mul(&nll, &focal_weight)?; + let focal_values = mul(&weighted_nll, &alpha_tensor)?; + + // Apply reduction + let loss = match reduction { + "mean" => { + let sum = sum_all_elements(&focal_values)?; + // Average over samples, matching cross_entropy: only the true-class + // term per sample is non-zero, so the denominator is the number of + // samples (numel / num_classes), not the total element count. + let num_classes = predictions.size(predictions.ndim() - 1)?.max(1); + let n = (focal_values.numel() / num_classes) as f64; + divide_by_scalar(&sum, n)? + } + "sum" => sum_all_elements(&focal_values)?, + "none" => focal_values, + _ => { + return Err(MinitensorError::invalid_operation(format!( + "Invalid reduction mode: {}. Must be 'mean', 'sum', or 'none'", + reduction + ))); + } + }; + + // Set up gradient function if needed + if loss.requires_grad() { + let grad_fn = Arc::new(FocalLossBackward { + predictions_shape: predictions.shape().dims().to_vec(), + targets_shape: targets_one_hot.shape().dims().to_vec(), + input_ids: [predictions.id(), targets.id()], + input_requires_grad: [predictions.requires_grad(), targets.requires_grad()], + alpha, + gamma, + reduction: reduction.to_string(), + softmax_predictions: softmax_for_grad, + targets: targets_one_hot.clone().detach(), + }); + + let mut loss_with_grad = loss; + loss_with_grad.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&loss_with_grad, Some(grad_fn))?; + + Ok(loss_with_grad) + } else { + Ok(loss) + } +} + +/// Huber loss function for robust regression +/// +/// Combines MSE and MAE for robust regression: +/// - For |x| <= delta: 0.5 * x² +/// - For |x| > delta: delta * (|x| - 0.5 * delta) +/// +/// # Arguments +/// * `predictions` - Model predictions tensor +/// * `targets` - Ground truth targets tensor +/// * `delta` - Threshold for switching between MSE and MAE behavior +/// * `reduction` - How to reduce the loss ("mean", "sum", or "none") +/// +/// # Returns +/// * `Result` - The computed Huber loss +pub fn huber_loss( + predictions: &Tensor, + targets: &Tensor, + delta: f64, + reduction: &str, +) -> Result { + // Validate inputs + validate_loss_inputs(predictions, targets)?; + + if delta <= 0.0 { + return Err(MinitensorError::invalid_operation( + "Delta must be positive for Huber loss", + )); + } + + // Compute absolute differences: |predictions - targets| + let diff = sub(predictions, targets)?; + let diff_for_grad = diff.clone().detach(); + let abs_diff = activation_abs(&diff.detach())?; + + // Create delta tensor for comparison + let delta_tensor = create_scalar_tensor(delta, predictions.dtype(), predictions.device())?; + + // Compute Huber loss element-wise + let huber_values = compute_huber_elementwise(&abs_diff, &diff, &delta_tensor, delta)?; + + // Apply reduction + let loss = match reduction { + "mean" => { + let sum = sum_all_elements(&huber_values)?; + let n = huber_values.numel() as f64; + divide_by_scalar(&sum, n)? + } + "sum" => sum_all_elements(&huber_values)?, + "none" => huber_values, + _ => { + return Err(MinitensorError::invalid_operation(format!( + "Invalid reduction mode: {}. Must be 'mean', 'sum', or 'none'", + reduction + ))); + } + }; + + // Set up gradient function if needed. The forward is computed on detached + // data (the exact gradient is provided by HuberLossBackward from the stored + // diff), so gate on the inputs and enable grad on the loss explicitly. + if predictions.requires_grad() || targets.requires_grad() { + let grad_fn = Arc::new(HuberLossBackward { + predictions_shape: predictions.shape().dims().to_vec(), + targets_shape: targets.shape().dims().to_vec(), + input_ids: [predictions.id(), targets.id()], + input_requires_grad: [predictions.requires_grad(), targets.requires_grad()], + delta, + reduction: reduction.to_string(), + diff: diff_for_grad, + }); + + let mut loss_with_grad = loss.requires_grad_(true); + loss_with_grad.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&loss_with_grad, Some(grad_fn))?; + + Ok(loss_with_grad) + } else { + Ok(loss) + } +} + +/// Smooth L1 loss (Huber loss with delta=1.0) +/// +/// Computes Smooth L1 loss between predictions and targets: +/// SmoothL1(x) = 0.5 * x² if |x| < 1, otherwise |x| - 0.5 +/// +/// # Arguments +/// * `predictions` - Model predictions tensor +/// * `targets` - Ground truth targets tensor +/// * `reduction` - How to reduce the loss ("mean", "sum", or "none") +pub fn smooth_l1_loss(predictions: &Tensor, targets: &Tensor, reduction: &str) -> Result { + huber_loss(predictions, targets, 1.0, reduction) +} + +/// Log-cosh loss for robust regression +/// +/// Computes log_tensor(cosh(x)) where x = predictions - targets using a numerically +/// stable formulation: |x| + log1p(exp(-2|x|)) - log_tensor(2). +/// +/// # Arguments +/// * `predictions` - Model predictions tensor +/// * `targets` - Ground truth targets tensor +/// * `reduction` - How to reduce the loss ("mean", "sum", or "none") +pub fn log_cosh_loss(predictions: &Tensor, targets: &Tensor, reduction: &str) -> Result { + validate_loss_inputs(predictions, targets)?; + + let diff = sub(predictions, targets)?; + let diff_abs = activation_abs(&diff)?; + let neg_two = create_scalar_tensor(-2.0, diff.dtype(), diff.device())?; + let exp_term = exp(&mul(&diff_abs, &neg_two)?)?; + let log1p_term = log1p(&exp_term)?; + let log2 = create_scalar_tensor(std::f64::consts::LN_2, diff.dtype(), diff.device())?; + let log_cosh = sub(&add(&diff_abs, &log1p_term)?, &log2)?; + + match reduction { + "mean" => mean(&log_cosh, None, false), + "sum" => sum(&log_cosh, None, false), + "none" => Ok(log_cosh), + _ => Err(MinitensorError::invalid_operation(format!( + "Invalid reduction mode: {}. Must be 'mean', 'sum', or 'none'", + reduction + ))), + } +} + +// Helper functions + +/// Validate that loss function inputs are compatible +fn validate_loss_inputs(predictions: &Tensor, targets: &Tensor) -> Result<()> { + // Check device compatibility + if predictions.device() != targets.device() { + return Err(MinitensorError::device_mismatch( + format!("{:?}", predictions.device()), + format!("{:?}", targets.device()), + )); + } + + // Check data type compatibility + if predictions.dtype() != targets.dtype() { + return Err(MinitensorError::type_mismatch( + format!("{:?}", predictions.dtype()), + format!("{:?}", targets.dtype()), + )); + } + + // Check shape compatibility + if predictions.shape() != targets.shape() { + return Err(MinitensorError::shape_mismatch( + predictions.shape().dims().to_vec(), + targets.shape().dims().to_vec(), + )); + } + + // Check that tensors contain floating point data (required for loss computation) + match predictions.dtype() { + DataType::Float32 | DataType::Float64 => {} + _ => { + return Err(MinitensorError::invalid_operation( + "Loss functions require floating point tensors", + )); + } + } + + Ok(()) +} + +/// Validate that classification loss function inputs are compatible +fn validate_classification_inputs( + predictions: &Tensor, + targets: &Tensor, + require_same_dtype: bool, +) -> Result<()> { + // Check device compatibility + if predictions.device() != targets.device() { + return Err(MinitensorError::device_mismatch( + format!("{:?}", predictions.device()), + format!("{:?}", targets.device()), + )); + } + + // Optionally enforce data type equality + if require_same_dtype && predictions.dtype() != targets.dtype() { + return Err(MinitensorError::type_mismatch( + format!("{:?}", predictions.dtype()), + format!("{:?}", targets.dtype()), + )); + } + + // Predictions must be at least 2D (batch_size, num_classes) + if predictions.ndim() < 2 { + return Err(MinitensorError::invalid_operation( + "Classification predictions must be at least 2D (batch_size, num_classes)", + )); + } + + // Predictions must be floating point + match predictions.dtype() { + DataType::Float32 | DataType::Float64 => {} + _ => { + return Err(MinitensorError::invalid_operation( + "Classification loss functions require floating point tensors", + )); + } + } + + Ok(()) +} + +fn prepare_classification_targets(predictions: &Tensor, targets: &Tensor) -> Result { + if targets.ndim() + 1 == predictions.ndim() { + let num_classes = predictions.size(predictions.ndim() - 1)?; + let total = targets.numel(); + let mut data = TensorData::zeros_on_device( + total * num_classes, + predictions.dtype(), + predictions.device(), + ); + match (targets.dtype(), predictions.dtype()) { + (DataType::Int32, DataType::Float32) => { + let idx = targets.data().as_i32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i32 slice from targets") + })?; + let out = data.as_f32_slice_mut().unwrap(); + fill_one_hot_f32(idx, out, num_classes, |val| { + checked_index_from_i64(i64::from(*val), num_classes) + })?; + } + (DataType::Int64, DataType::Float32) => { + let idx = targets.data().as_i64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i64 slice from targets") + })?; + let out = data.as_f32_slice_mut().unwrap(); + fill_one_hot_f32(idx, out, num_classes, |val| { + checked_index_from_i64(*val, num_classes) + })?; + } + (DataType::Int32, DataType::Float64) => { + let idx = targets.data().as_i32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i32 slice from targets") + })?; + let out = data.as_f64_slice_mut().unwrap(); + fill_one_hot_f64(idx, out, num_classes, |val| { + checked_index_from_i64(i64::from(*val), num_classes) + })?; + } + (DataType::Int64, DataType::Float64) => { + let idx = targets.data().as_i64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i64 slice from targets") + })?; + let out = data.as_f64_slice_mut().unwrap(); + fill_one_hot_f64(idx, out, num_classes, |val| { + checked_index_from_i64(*val, num_classes) + })?; + } + (DataType::Float32, DataType::Float32) => { + let idx = targets.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from targets") + })?; + let out = data.as_f32_slice_mut().unwrap(); + fill_one_hot_f32(idx, out, num_classes, |val| { + checked_index_from_f32(*val, num_classes) + })?; + } + (DataType::Float64, DataType::Float64) => { + let idx = targets.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from targets") + })?; + let out = data.as_f64_slice_mut().unwrap(); + fill_one_hot_f64(idx, out, num_classes, |val| { + checked_index_from_f64(*val, num_classes) + })?; + } + _ => { + return Err(MinitensorError::invalid_operation( + "Unsupported target dtype for classification loss", + )); + } + } + let mut dims = targets.shape().dims().to_vec(); + dims.push(num_classes); + Ok(Tensor::new( + Arc::new(data), + Shape::new(dims), + predictions.dtype(), + predictions.device(), + false, + )) + } else if targets.ndim() == predictions.ndim() { + if targets.shape().dims() != predictions.shape().dims() { + return Err(MinitensorError::shape_mismatch( + predictions.shape().dims().to_vec(), + targets.shape().dims().to_vec(), + )); + } + Ok(targets.clone()) + } else { + Err(MinitensorError::shape_mismatch( + predictions.shape().dims().to_vec(), + targets.shape().dims().to_vec(), + )) + } +} + +fn checked_index_from_i64(value: i64, num_classes: usize) -> Result { + if value < 0 { + return Err(MinitensorError::invalid_operation( + "Target class index must be non-negative", + )); + } + let index = value as usize; + if index >= num_classes { + return Err(MinitensorError::invalid_operation( + "Target class index out of range", + )); + } + Ok(index) +} + +fn checked_index_from_f32(value: f32, num_classes: usize) -> Result { + if !value.is_finite() || value.fract() != 0.0 { + return Err(MinitensorError::invalid_operation( + "Target class index must be a finite integer", + )); + } + if value < 0.0 || value >= num_classes as f32 { + return Err(MinitensorError::invalid_operation( + "Target class index out of range", + )); + } + Ok(value as usize) +} + +fn checked_index_from_f64(value: f64, num_classes: usize) -> Result { + if !value.is_finite() || value.fract() != 0.0 { + return Err(MinitensorError::invalid_operation( + "Target class index must be a finite integer", + )); + } + if value < 0.0 || value >= num_classes as f64 { + return Err(MinitensorError::invalid_operation( + "Target class index out of range", + )); + } + Ok(value as usize) +} + +fn fill_one_hot_f32( + indices: &[T], + out: &mut [f32], + num_classes: usize, + to_index: F, +) -> Result<()> +where + F: Fn(&T) -> Result, +{ + for (i, value) in indices.iter().enumerate() { + let class = to_index(value)?; + out[i * num_classes + class] = 1.0; + } + Ok(()) +} diff --git a/engine/src/operations/mod.rs b/engine/src/operations/mod.rs index cb67c412..98130419 100644 --- a/engine/src/operations/mod.rs +++ b/engine/src/operations/mod.rs @@ -9,7 +9,6 @@ pub mod arithmetic; pub mod binary; pub mod comparison; pub mod conv; -pub mod fusion; pub mod linalg; pub mod loss; pub mod minmax; @@ -24,7 +23,6 @@ pub use activation::*; pub use arithmetic::*; pub use comparison::*; pub use conv::*; -pub use fusion::*; pub use linalg::*; pub use loss::*; pub use minmax::*; diff --git a/engine/src/operations/normalization.rs b/engine/src/operations/normalization.rs index ee6b8e74..71b479db 100644 --- a/engine/src/operations/normalization.rs +++ b/engine/src/operations/normalization.rs @@ -1,460 +1,458 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -use crate::autograd::{LayerNormBackward, TensorId, add_to_graph}; -use crate::device::Device; -use crate::error::{MinitensorError, Result}; -use crate::tensor::{DataType, Shape, Tensor, TensorData}; -use smallvec::SmallVec; -use std::sync::Arc; - -fn scalar_tensor(value: f64, dtype: DataType, device: Device) -> Result { - let mut data = TensorData::zeros_on_device(1, dtype, device); - match dtype { - DataType::Float32 => { - let slice = data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from scalar tensor", - ) - })?; - slice[0] = value as f32; - } - DataType::Float64 => { - let slice = data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from scalar tensor", - ) - })?; - slice[0] = value; - } - _ => { - return Err(MinitensorError::invalid_operation( - "Normalization operations only support floating point tensors".to_string(), - )); - } - } - - Ok(Tensor::new( - Arc::new(data), - Shape::new(vec![1]), - dtype, - device, - false, - )) -} - -/// Functional batch normalization. -/// -/// Normalizes the input tensor using batch statistics during training or -/// running estimates during evaluation. -/// -/// * `input` - Input tensor of shape `[N, C, ...]` where the second dimension -/// is interpreted as the feature/channel dimension. -/// * `running_mean` - Optional running mean buffer updated during training. -/// * `running_var` - Optional running variance buffer updated during training. -/// * `weight` - Optional learnable scale parameter (gamma). -/// * `bias` - Optional learnable shift parameter (beta). -/// * `training` - When true, use batch statistics and update running stats. -/// * `momentum` - Momentum factor for running statistics update. -/// * `eps` - Small epsilon added to variance for numerical stability. -#[allow(clippy::too_many_arguments)] -pub fn batch_norm( - input: &Tensor, - running_mean: Option<&mut Tensor>, - running_var: Option<&mut Tensor>, - weight: Option<&Tensor>, - bias: Option<&Tensor>, - training: bool, - momentum: f64, - eps: f64, -) -> Result { - if input.ndim() < 2 { - return Err(MinitensorError::invalid_operation( - "batch_norm expects input with at least 2 dimensions", - )); - } - - let num_features = input.size(1)?; - - // Validate parameter shapes - if let Some(w) = weight { - if w.ndim() != 1 || w.size(0)? != num_features { - return Err(MinitensorError::shape_mismatch( - vec![num_features], - vec![w.size(0)?], - )); - } - } - if let Some(b) = bias { - if b.ndim() != 1 || b.size(0)? != num_features { - return Err(MinitensorError::shape_mismatch( - vec![num_features], - vec![b.size(0)?], - )); - } - } - if let Some(rm) = &running_mean { - if rm.ndim() != 1 || rm.size(0)? != num_features { - return Err(MinitensorError::shape_mismatch( - vec![num_features], - vec![rm.size(0)?], - )); - } - } - if let Some(rv) = &running_var { - if rv.ndim() != 1 || rv.size(0)? != num_features { - return Err(MinitensorError::shape_mismatch( - vec![num_features], - vec![rv.size(0)?], - )); - } - } - - // Dimensions along which to compute statistics (all except channel dim) - let axes: Vec = (0..input.ndim()).filter(|&d| d != 1).collect(); - let axes_isize: Vec = axes.iter().map(|&d| d as isize).collect(); - - // Compute batch statistics only when they are actually used: during - // training, or in eval mode when no running estimates are available. - let use_batch_stats = training || running_mean.is_none() || running_var.is_none(); - let (mean_used, var_used, centered) = if use_batch_stats { - let batch_mean = input.mean(Some(axes_isize.clone()), true)?; // [1, C, ...] - let centered = crate::operations::arithmetic::sub(input, &batch_mean)?; - let batch_var = crate::operations::arithmetic::mul(¢ered, ¢ered)? - .mean(Some(axes_isize.clone()), true)?; - (batch_mean, batch_var, centered) - } else if let (Some(rm), Some(rv)) = (running_mean.as_ref(), running_var.as_ref()) { - // Use running statistics (reshape for broadcasting) - let mut rm_view = (*rm).clone().unsqueeze(0)?; // [1, C] - let mut rv_view = (*rv).clone().unsqueeze(0)?; - for _ in 2..input.ndim() { - rm_view = rm_view.unsqueeze(rm_view.ndim() as isize)?; - rv_view = rv_view.unsqueeze(rv_view.ndim() as isize)?; - } - let centered = crate::operations::arithmetic::sub(input, &rm_view)?; - (rm_view, rv_view, centered) - } else { - unreachable!("running stats checked") - }; - - // Prepare epsilon tensor - let eps_tensor = scalar_tensor(eps, input.dtype(), input.device())?; - - let var_eps = crate::operations::arithmetic::add(&var_used, &eps_tensor)?; - let std = crate::operations::activation::sqrt(&var_eps)?; - let mut output = crate::operations::arithmetic::div(¢ered, &std)?; - - // Scale and shift - if let Some(w) = weight { - let mut w_view = w.clone().unsqueeze(0)?; - for _ in 2..input.ndim() { - w_view = w_view.unsqueeze(w_view.ndim() as isize)?; - } - output = crate::operations::arithmetic::mul(&output, &w_view)?; - } - if let Some(b) = bias { - let mut b_view = b.clone().unsqueeze(0)?; - for _ in 2..input.ndim() { - b_view = b_view.unsqueeze(b_view.ndim() as isize)?; - } - output = crate::operations::arithmetic::add(&output, &b_view)?; - } - - // Update running statistics if training (mean_used/var_used hold the - // batch statistics whenever training is true) - if training { - if let (Some(rm), Some(rv)) = (running_mean, running_var) { - let mean_flat = mean_used.view(Shape::new(vec![num_features]))?.detach(); - let var_flat = var_used.view(Shape::new(vec![num_features]))?.detach(); - - let m_tensor = scalar_tensor(momentum, input.dtype(), input.device())?; - let one_minus_tensor = scalar_tensor(1.0 - momentum, input.dtype(), input.device())?; - - *rm = crate::operations::arithmetic::add( - &crate::operations::arithmetic::mul(rm, &one_minus_tensor)?, - &crate::operations::arithmetic::mul(&mean_flat, &m_tensor)?, - )?; - *rv = crate::operations::arithmetic::add( - &crate::operations::arithmetic::mul(rv, &one_minus_tensor)?, - &crate::operations::arithmetic::mul(&var_flat, &m_tensor)?, - )?; - } - } - - Ok(output) -} - -/// Apply layer normalization to the input tensor. -pub fn layer_norm( - input: &Tensor, - normalized_shape: &[usize], - weight: Option<&Tensor>, - bias: Option<&Tensor>, - eps: f64, -) -> Result { - if normalized_shape.is_empty() { - return Err(MinitensorError::invalid_argument( - "layer_norm requires at least one normalized dimension".to_string(), - )); - } - - if normalized_shape.len() > input.ndim() { - return Err(MinitensorError::invalid_operation( - "normalized_shape rank cannot exceed input rank for layer_norm".to_string(), - )); - } - - match input.dtype() { - DataType::Float32 | DataType::Float64 => {} - _ => { - return Err(MinitensorError::invalid_operation( - "layer_norm only supports floating point tensors".to_string(), - )); - } - } - - let axis_start = input.ndim() - normalized_shape.len(); - for (i, &expected) in normalized_shape.iter().enumerate() { - let dim = axis_start + i; - let actual = input.size(dim)?; - if actual != expected { - return Err(MinitensorError::shape_mismatch( - vec![expected], - vec![actual], - )); - } - } - - if let Some(w) = weight { - if w.dtype() != input.dtype() { - return Err(MinitensorError::type_mismatch( - input.dtype().to_string(), - w.dtype().to_string(), - )); - } - if w.device() != input.device() { - return Err(MinitensorError::device_mismatch( - input.device().to_string(), - w.device().to_string(), - )); - } - if w.shape().dims() != normalized_shape { - return Err(MinitensorError::shape_mismatch( - normalized_shape.to_vec(), - w.shape().dims().to_vec(), - )); - } - } - - if let Some(b) = bias { - if b.dtype() != input.dtype() { - return Err(MinitensorError::type_mismatch( - input.dtype().to_string(), - b.dtype().to_string(), - )); - } - if b.device() != input.device() { - return Err(MinitensorError::device_mismatch( - input.device().to_string(), - b.device().to_string(), - )); - } - if b.shape().dims() != normalized_shape { - return Err(MinitensorError::shape_mismatch( - normalized_shape.to_vec(), - b.shape().dims().to_vec(), - )); - } - } - - let axes: Vec = (axis_start..input.ndim()).collect(); - let axes_isize: Vec = axes.iter().map(|&d| d as isize).collect(); - let mean = input.mean(Some(axes_isize.clone()), true)?; - let centered = crate::operations::arithmetic::sub(input, &mean)?; - let var = crate::operations::arithmetic::mul(¢ered, ¢ered)? - .mean(Some(axes_isize.clone()), true)?; - let eps_tensor = scalar_tensor(eps, input.dtype(), input.device())?; - let var_eps = crate::operations::arithmetic::add(&var, &eps_tensor)?; - let std = crate::operations::activation::sqrt(&var_eps)?; - let ones = Tensor::ones(std.shape().clone(), std.dtype(), std.device(), false); - let inv_std = crate::operations::arithmetic::div(&ones, &std)?; - let normalized = crate::operations::arithmetic::mul(¢ered, &inv_std)?; - - let mut output = normalized.clone(); - let mut weight_broadcast: Option = None; - if let Some(w) = weight { - let mut view = w.clone(); - for _ in 0..axis_start { - view = view.unsqueeze(0)?; - } - output = crate::operations::arithmetic::mul(&output, &view)?; - weight_broadcast = Some(view.detach()); - } - if let Some(b) = bias { - let mut view = b.clone(); - for _ in 0..axis_start { - view = view.unsqueeze(0)?; - } - output = crate::operations::arithmetic::add(&output, &view)?; - } - - let requires_grad = input.requires_grad() - || weight.map(|w| w.requires_grad()).unwrap_or(false) - || bias.map(|b| b.requires_grad()).unwrap_or(false); - - if !requires_grad { - return Ok(output); - } - - let mut input_ids: SmallVec<[TensorId; 3]> = SmallVec::new(); - input_ids.push(input.id()); - if let Some(w) = weight { - input_ids.push(w.id()); - } - if let Some(b) = bias { - input_ids.push(b.id()); - } - - let grad_fn = Arc::new(LayerNormBackward { - input_ids, - input_id: input.id(), - weight_id: weight.map(|w| w.id()), - bias_id: bias.map(|b| b.id()), - normalized: normalized.detach(), - inv_std: inv_std.detach(), - weight_broadcast, - normalized_shape: normalized_shape.to_vec(), - axis_start, - element_count: normalized_shape.iter().product(), - input_requires_grad: input.requires_grad(), - weight_requires_grad: weight.map(|w| w.requires_grad()).unwrap_or(false), - bias_requires_grad: bias.map(|b| b.requires_grad()).unwrap_or(false), - }); - - let mut output_with_grad = output; - output_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output_with_grad, Some(grad_fn))?; - Ok(output_with_grad) -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::autograd; - use crate::device::Device; - use crate::tensor::{DataType, TensorData}; - use std::sync::Arc; - - fn tensor_from_vec(data: Vec, shape: Vec, requires_grad: bool) -> Tensor { - Tensor::new( - Arc::new(TensorData::from_vec_f32(data, Device::cpu())), - Shape::new(shape), - DataType::Float32, - Device::cpu(), - requires_grad, - ) - } - - #[test] - fn test_layer_norm_forward_zero_mean_unit_var() { - let input = tensor_from_vec(vec![1.0, 2.0, 3.0, -1.0, 0.0, 4.0], vec![2, 3], false); - let result = layer_norm(&input, &[3], None, None, 1e-5).unwrap(); - let data = result.data().as_f32_slice().unwrap(); - - for row in 0..2 { - let start = row * 3; - let slice = &data[start..start + 3]; - let mean: f32 = slice.iter().sum::() / 3.0; - assert!(mean.abs() < 1e-5); - let var: f32 = slice - .iter() - .map(|v| { - let diff = *v - mean; - diff * diff - }) - .sum::() - / 3.0; - assert!((var - 1.0).abs() < 1e-4); - } - } - - #[test] - fn test_layer_norm_backward_matches_manual_gradients() { - let input_vals = vec![1.2f32, -0.5, 2.0, 0.7, -1.3, 0.25]; - let weight_vals = vec![1.5f32, 0.75, -0.25]; - let bias_vals = vec![0.1f32, -0.2, 0.05]; - - let input = tensor_from_vec(input_vals.clone(), vec![2, 3], true); - let weight = tensor_from_vec(weight_vals.clone(), vec![3], true); - let bias = tensor_from_vec(bias_vals.clone(), vec![3], true); - - let result = layer_norm(&input, &[3], Some(&weight), Some(&bias), 1e-5).unwrap(); - let ones = Tensor::ones( - result.shape().clone(), - result.dtype(), - result.device(), - false, - ); - let grads = autograd::backward(&result, Some(ones)).unwrap(); - - let grad_input = grads.get(&input.id()).unwrap(); - let grad_weight = grads.get(&weight.id()).unwrap(); - let grad_bias = grads.get(&bias.id()).unwrap(); - - let mut expected_input_grad = vec![0.0f32; input_vals.len()]; - let mut expected_weight_grad = vec![0.0f32; weight_vals.len()]; - let mut expected_bias_grad = vec![0.0f32; bias_vals.len()]; - let eps = 1e-5f32; - let m = 3.0f32; - - for row in 0..2 { - let start = row * 3; - let x = &input_vals[start..start + 3]; - let mean = x.iter().sum::() / m; - let centered: Vec = x.iter().map(|v| *v - mean).collect(); - let var = centered.iter().map(|v| v * v).sum::() / m; - let inv_std = 1.0 / (var + eps).sqrt(); - let normalized: Vec = centered.iter().map(|v| v * inv_std).collect(); - - let grad_output = [1.0f32; 3]; - let grad_output_hat: Vec = grad_output - .iter() - .zip(weight_vals.iter()) - .map(|(g, w)| g * *w) - .collect(); - - let sum_grad = grad_output_hat.iter().sum::(); - let sum_grad_norm = grad_output_hat - .iter() - .zip(normalized.iter()) - .map(|(g, n)| g * n) - .sum::(); - - for i in 0..3 { - let numerator = grad_output_hat[i] * m - sum_grad - normalized[i] * sum_grad_norm; - expected_input_grad[start + i] += numerator * inv_std / m; - expected_weight_grad[i] += grad_output[i] * normalized[i]; - expected_bias_grad[i] += grad_output[i]; - } - } - - let input_grad_vals = grad_input.data().as_f32_slice().unwrap(); - let weight_grad_vals = grad_weight.data().as_f32_slice().unwrap(); - let bias_grad_vals = grad_bias.data().as_f32_slice().unwrap(); - - for (actual, expected) in input_grad_vals.iter().zip(expected_input_grad.iter()) { - assert!((actual - expected).abs() < 1e-5); - } - - for (actual, expected) in weight_grad_vals.iter().zip(expected_weight_grad.iter()) { - assert!((actual - expected).abs() < 1e-5); - } - - for (actual, expected) in bias_grad_vals.iter().zip(expected_bias_grad.iter()) { - assert!((actual - expected).abs() < 1e-6); - } - } -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use crate::autograd::{LayerNormBackward, TensorId, add_to_graph}; +use crate::device::Device; +use crate::error::{MinitensorError, Result}; +use crate::tensor::{DataType, Shape, Tensor, TensorData}; +use smallvec::SmallVec; +use std::sync::Arc; + +fn scalar_tensor(value: f64, dtype: DataType, device: Device) -> Result { + let mut data = TensorData::zeros_on_device(1, dtype, device); + match dtype { + DataType::Float32 => { + let slice = data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from scalar tensor", + ) + })?; + slice[0] = value as f32; + } + DataType::Float64 => { + let slice = data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from scalar tensor", + ) + })?; + slice[0] = value; + } + _ => { + return Err(MinitensorError::invalid_operation( + "Normalization operations only support floating point tensors".to_string(), + )); + } + } + + Ok(Tensor::new( + Arc::new(data), + Shape::new(vec![1]), + dtype, + device, + false, + )) +} + +/// Functional batch normalization. +/// +/// Normalizes the input tensor using batch statistics during training or +/// running estimates during evaluation. +/// +/// * `input` - Input tensor of shape `[N, C, ...]` where the second dimension +/// is interpreted as the feature/channel dimension. +/// * `running_mean` - Optional running mean buffer updated during training. +/// * `running_var` - Optional running variance buffer updated during training. +/// * `weight` - Optional learnable scale parameter (gamma). +/// * `bias` - Optional learnable shift parameter (beta). +/// * `training` - When true, use batch statistics and update running stats. +/// * `momentum` - Momentum factor for running statistics update. +/// * `eps` - Small epsilon added to variance for numerical stability. +#[allow(clippy::too_many_arguments)] +pub fn batch_norm( + input: &Tensor, + running_mean: Option<&mut Tensor>, + running_var: Option<&mut Tensor>, + weight: Option<&Tensor>, + bias: Option<&Tensor>, + training: bool, + momentum: f64, + eps: f64, +) -> Result { + if input.ndim() < 2 { + return Err(MinitensorError::invalid_operation( + "batch_norm expects input with at least 2 dimensions", + )); + } + + let num_features = input.size(1)?; + + // Validate parameter shapes + if let Some(w) = weight + && (w.ndim() != 1 || w.size(0)? != num_features) + { + return Err(MinitensorError::shape_mismatch( + vec![num_features], + vec![w.size(0)?], + )); + } + if let Some(b) = bias + && (b.ndim() != 1 || b.size(0)? != num_features) + { + return Err(MinitensorError::shape_mismatch( + vec![num_features], + vec![b.size(0)?], + )); + } + if let Some(rm) = &running_mean + && (rm.ndim() != 1 || rm.size(0)? != num_features) + { + return Err(MinitensorError::shape_mismatch( + vec![num_features], + vec![rm.size(0)?], + )); + } + if let Some(rv) = &running_var + && (rv.ndim() != 1 || rv.size(0)? != num_features) + { + return Err(MinitensorError::shape_mismatch( + vec![num_features], + vec![rv.size(0)?], + )); + } + + // Dimensions along which to compute statistics (all except channel dim) + let axes: Vec = (0..input.ndim()).filter(|&d| d != 1).collect(); + let axes_isize: Vec = axes.iter().map(|&d| d as isize).collect(); + + // Compute batch statistics only when they are actually used: during + // training, or in eval mode when no running estimates are available. + let use_batch_stats = training || running_mean.is_none() || running_var.is_none(); + let (mean_used, var_used, centered) = if use_batch_stats { + let batch_mean = input.mean(Some(axes_isize.clone()), true)?; // [1, C, ...] + let centered = crate::operations::arithmetic::sub(input, &batch_mean)?; + let batch_var = crate::operations::arithmetic::mul(¢ered, ¢ered)? + .mean(Some(axes_isize.clone()), true)?; + (batch_mean, batch_var, centered) + } else if let (Some(rm), Some(rv)) = (running_mean.as_ref(), running_var.as_ref()) { + // Use running statistics (reshape for broadcasting) + let mut rm_view = (*rm).clone().unsqueeze(0)?; // [1, C] + let mut rv_view = (*rv).clone().unsqueeze(0)?; + for _ in 2..input.ndim() { + rm_view = rm_view.unsqueeze(rm_view.ndim() as isize)?; + rv_view = rv_view.unsqueeze(rv_view.ndim() as isize)?; + } + let centered = crate::operations::arithmetic::sub(input, &rm_view)?; + (rm_view, rv_view, centered) + } else { + unreachable!("running stats checked") + }; + + // Prepare epsilon tensor + let eps_tensor = scalar_tensor(eps, input.dtype(), input.device())?; + + let var_eps = crate::operations::arithmetic::add(&var_used, &eps_tensor)?; + let std = crate::operations::activation::sqrt(&var_eps)?; + let mut output = crate::operations::arithmetic::div(¢ered, &std)?; + + // Scale and shift + if let Some(w) = weight { + let mut w_view = w.clone().unsqueeze(0)?; + for _ in 2..input.ndim() { + w_view = w_view.unsqueeze(w_view.ndim() as isize)?; + } + output = crate::operations::arithmetic::mul(&output, &w_view)?; + } + if let Some(b) = bias { + let mut b_view = b.clone().unsqueeze(0)?; + for _ in 2..input.ndim() { + b_view = b_view.unsqueeze(b_view.ndim() as isize)?; + } + output = crate::operations::arithmetic::add(&output, &b_view)?; + } + + // Update running statistics if training (mean_used/var_used hold the + // batch statistics whenever training is true) + if training && let (Some(rm), Some(rv)) = (running_mean, running_var) { + let mean_flat = mean_used.view(Shape::new(vec![num_features]))?.detach(); + let var_flat = var_used.view(Shape::new(vec![num_features]))?.detach(); + + let m_tensor = scalar_tensor(momentum, input.dtype(), input.device())?; + let one_minus_tensor = scalar_tensor(1.0 - momentum, input.dtype(), input.device())?; + + *rm = crate::operations::arithmetic::add( + &crate::operations::arithmetic::mul(rm, &one_minus_tensor)?, + &crate::operations::arithmetic::mul(&mean_flat, &m_tensor)?, + )?; + *rv = crate::operations::arithmetic::add( + &crate::operations::arithmetic::mul(rv, &one_minus_tensor)?, + &crate::operations::arithmetic::mul(&var_flat, &m_tensor)?, + )?; + } + + Ok(output) +} + +/// Apply layer normalization to the input tensor. +pub fn layer_norm( + input: &Tensor, + normalized_shape: &[usize], + weight: Option<&Tensor>, + bias: Option<&Tensor>, + eps: f64, +) -> Result { + if normalized_shape.is_empty() { + return Err(MinitensorError::invalid_argument( + "layer_norm requires at least one normalized dimension".to_string(), + )); + } + + if normalized_shape.len() > input.ndim() { + return Err(MinitensorError::invalid_operation( + "normalized_shape rank cannot exceed input rank for layer_norm".to_string(), + )); + } + + match input.dtype() { + DataType::Float32 | DataType::Float64 => {} + _ => { + return Err(MinitensorError::invalid_operation( + "layer_norm only supports floating point tensors".to_string(), + )); + } + } + + let axis_start = input.ndim() - normalized_shape.len(); + for (i, &expected) in normalized_shape.iter().enumerate() { + let dim = axis_start + i; + let actual = input.size(dim)?; + if actual != expected { + return Err(MinitensorError::shape_mismatch( + vec![expected], + vec![actual], + )); + } + } + + if let Some(w) = weight { + if w.dtype() != input.dtype() { + return Err(MinitensorError::type_mismatch( + input.dtype().to_string(), + w.dtype().to_string(), + )); + } + if w.device() != input.device() { + return Err(MinitensorError::device_mismatch( + input.device().to_string(), + w.device().to_string(), + )); + } + if w.shape().dims() != normalized_shape { + return Err(MinitensorError::shape_mismatch( + normalized_shape.to_vec(), + w.shape().dims().to_vec(), + )); + } + } + + if let Some(b) = bias { + if b.dtype() != input.dtype() { + return Err(MinitensorError::type_mismatch( + input.dtype().to_string(), + b.dtype().to_string(), + )); + } + if b.device() != input.device() { + return Err(MinitensorError::device_mismatch( + input.device().to_string(), + b.device().to_string(), + )); + } + if b.shape().dims() != normalized_shape { + return Err(MinitensorError::shape_mismatch( + normalized_shape.to_vec(), + b.shape().dims().to_vec(), + )); + } + } + + let axes: Vec = (axis_start..input.ndim()).collect(); + let axes_isize: Vec = axes.iter().map(|&d| d as isize).collect(); + let mean = input.mean(Some(axes_isize.clone()), true)?; + let centered = crate::operations::arithmetic::sub(input, &mean)?; + let var = crate::operations::arithmetic::mul(¢ered, ¢ered)? + .mean(Some(axes_isize.clone()), true)?; + let eps_tensor = scalar_tensor(eps, input.dtype(), input.device())?; + let var_eps = crate::operations::arithmetic::add(&var, &eps_tensor)?; + let std = crate::operations::activation::sqrt(&var_eps)?; + let ones = Tensor::ones(std.shape().clone(), std.dtype(), std.device(), false); + let inv_std = crate::operations::arithmetic::div(&ones, &std)?; + let normalized = crate::operations::arithmetic::mul(¢ered, &inv_std)?; + + let mut output = normalized.clone(); + let mut weight_broadcast: Option = None; + if let Some(w) = weight { + let mut view = w.clone(); + for _ in 0..axis_start { + view = view.unsqueeze(0)?; + } + output = crate::operations::arithmetic::mul(&output, &view)?; + weight_broadcast = Some(view.detach()); + } + if let Some(b) = bias { + let mut view = b.clone(); + for _ in 0..axis_start { + view = view.unsqueeze(0)?; + } + output = crate::operations::arithmetic::add(&output, &view)?; + } + + let requires_grad = input.requires_grad() + || weight.map(|w| w.requires_grad()).unwrap_or(false) + || bias.map(|b| b.requires_grad()).unwrap_or(false); + + if !requires_grad { + return Ok(output); + } + + let mut input_ids: SmallVec<[TensorId; 3]> = SmallVec::new(); + input_ids.push(input.id()); + if let Some(w) = weight { + input_ids.push(w.id()); + } + if let Some(b) = bias { + input_ids.push(b.id()); + } + + let grad_fn = Arc::new(LayerNormBackward { + input_ids, + input_id: input.id(), + weight_id: weight.map(|w| w.id()), + bias_id: bias.map(|b| b.id()), + normalized: normalized.detach(), + inv_std: inv_std.detach(), + weight_broadcast, + normalized_shape: normalized_shape.to_vec(), + axis_start, + element_count: normalized_shape.iter().product(), + input_requires_grad: input.requires_grad(), + weight_requires_grad: weight.map(|w| w.requires_grad()).unwrap_or(false), + bias_requires_grad: bias.map(|b| b.requires_grad()).unwrap_or(false), + }); + + let mut output_with_grad = output; + output_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output_with_grad, Some(grad_fn))?; + Ok(output_with_grad) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::autograd; + use crate::device::Device; + use crate::tensor::{DataType, TensorData}; + use std::sync::Arc; + + fn tensor_from_vec(data: Vec, shape: Vec, requires_grad: bool) -> Tensor { + Tensor::new( + Arc::new(TensorData::from_vec_f32(data, Device::cpu())), + Shape::new(shape), + DataType::Float32, + Device::cpu(), + requires_grad, + ) + } + + #[test] + fn test_layer_norm_forward_zero_mean_unit_var() { + let input = tensor_from_vec(vec![1.0, 2.0, 3.0, -1.0, 0.0, 4.0], vec![2, 3], false); + let result = layer_norm(&input, &[3], None, None, 1e-5).unwrap(); + let data = result.data().as_f32_slice().unwrap(); + + for row in 0..2 { + let start = row * 3; + let slice = &data[start..start + 3]; + let mean: f32 = slice.iter().sum::() / 3.0; + assert!(mean.abs() < 1e-5); + let var: f32 = slice + .iter() + .map(|v| { + let diff = *v - mean; + diff * diff + }) + .sum::() + / 3.0; + assert!((var - 1.0).abs() < 1e-4); + } + } + + #[test] + fn test_layer_norm_backward_matches_manual_gradients() { + let input_vals = vec![1.2f32, -0.5, 2.0, 0.7, -1.3, 0.25]; + let weight_vals = vec![1.5f32, 0.75, -0.25]; + let bias_vals = vec![0.1f32, -0.2, 0.05]; + + let input = tensor_from_vec(input_vals.clone(), vec![2, 3], true); + let weight = tensor_from_vec(weight_vals.clone(), vec![3], true); + let bias = tensor_from_vec(bias_vals.clone(), vec![3], true); + + let result = layer_norm(&input, &[3], Some(&weight), Some(&bias), 1e-5).unwrap(); + let ones = Tensor::ones( + result.shape().clone(), + result.dtype(), + result.device(), + false, + ); + let grads = autograd::backward_collect(&result, Some(ones)).unwrap(); + + let grad_input = grads.get(&input.id()).unwrap(); + let grad_weight = grads.get(&weight.id()).unwrap(); + let grad_bias = grads.get(&bias.id()).unwrap(); + + let mut expected_input_grad = vec![0.0f32; input_vals.len()]; + let mut expected_weight_grad = vec![0.0f32; weight_vals.len()]; + let mut expected_bias_grad = vec![0.0f32; bias_vals.len()]; + let eps = 1e-5f32; + let m = 3.0f32; + + for row in 0..2 { + let start = row * 3; + let x = &input_vals[start..start + 3]; + let mean = x.iter().sum::() / m; + let centered: Vec = x.iter().map(|v| *v - mean).collect(); + let var = centered.iter().map(|v| v * v).sum::() / m; + let inv_std = 1.0 / (var + eps).sqrt(); + let normalized: Vec = centered.iter().map(|v| v * inv_std).collect(); + + let grad_output = [1.0f32; 3]; + let grad_output_hat: Vec = grad_output + .iter() + .zip(weight_vals.iter()) + .map(|(g, w)| g * *w) + .collect(); + + let sum_grad = grad_output_hat.iter().sum::(); + let sum_grad_norm = grad_output_hat + .iter() + .zip(normalized.iter()) + .map(|(g, n)| g * n) + .sum::(); + + for i in 0..3 { + let numerator = grad_output_hat[i] * m - sum_grad - normalized[i] * sum_grad_norm; + expected_input_grad[start + i] += numerator * inv_std / m; + expected_weight_grad[i] += grad_output[i] * normalized[i]; + expected_bias_grad[i] += grad_output[i]; + } + } + + let input_grad_vals = grad_input.data().as_f32_slice().unwrap(); + let weight_grad_vals = grad_weight.data().as_f32_slice().unwrap(); + let bias_grad_vals = grad_bias.data().as_f32_slice().unwrap(); + + for (actual, expected) in input_grad_vals.iter().zip(expected_input_grad.iter()) { + assert!((actual - expected).abs() < 1e-5); + } + + for (actual, expected) in weight_grad_vals.iter().zip(expected_weight_grad.iter()) { + assert!((actual - expected).abs() < 1e-5); + } + + for (actual, expected) in bias_grad_vals.iter().zip(expected_bias_grad.iter()) { + assert!((actual - expected).abs() < 1e-6); + } + } +} diff --git a/engine/src/operations/reduction.rs b/engine/src/operations/reduction.rs index a9387fc3..d270f432 100644 --- a/engine/src/operations/reduction.rs +++ b/engine/src/operations/reduction.rs @@ -4,13 +4,34 @@ // This source code is licensed under the Apache-style license found in the // LICENSE file in the root directory of this source tree. -include!("reduction/core.rs"); -include!("reduction/quantile.rs"); -include!("reduction/nanquantile.rs"); -include!("reduction/logsumexp.rs"); -include!("reduction/boolean.rs"); -include!("reduction/sort.rs"); -include!("reduction/sum_prod.rs"); -include!("reduction/nan_minmax.rs"); -include!("reduction/minmax_indices.rs"); -include!("reduction/argminmax.rs"); +#[path = "reduction/argminmax.rs"] +mod argminmax_impl; +#[path = "reduction/boolean.rs"] +mod boolean_impl; +#[path = "reduction/core.rs"] +mod core_impl; +#[path = "reduction/logsumexp.rs"] +mod logsumexp_impl; +#[path = "reduction/minmax_indices.rs"] +mod minmax_indices_impl; +#[path = "reduction/nan_minmax.rs"] +mod nan_minmax_impl; +#[path = "reduction/nanquantile.rs"] +mod nanquantile_impl; +#[path = "reduction/quantile.rs"] +mod quantile_impl; +#[path = "reduction/sort.rs"] +mod sort_impl; +#[path = "reduction/sum_prod.rs"] +mod sum_prod_impl; + +pub(crate) use self::argminmax_impl::*; +pub use self::boolean_impl::*; +pub use self::core_impl::*; +pub use self::logsumexp_impl::*; +pub(crate) use self::minmax_indices_impl::*; +pub(crate) use self::nan_minmax_impl::*; +pub use self::nanquantile_impl::*; +pub(crate) use self::quantile_impl::*; +pub use self::sort_impl::*; +pub use self::sum_prod_impl::*; diff --git a/engine/src/operations/reduction/argminmax.rs b/engine/src/operations/reduction/argminmax.rs index 17198405..abb31d8f 100644 --- a/engine/src/operations/reduction/argminmax.rs +++ b/engine/src/operations/reduction/argminmax.rs @@ -1,940 +1,961 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -fn argmin_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { - let layout = reduction_layout(tensor, dim, keepdim)?; - let mut result_data = TensorData::zeros_on_device( - layout.output_shape.numel(), - DataType::Int64, - tensor.device(), - ); - - let output = result_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - match tensor.dtype() { - DataType::Float32 => { - let input = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut min_val = f32::INFINITY; - let mut min_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if val.is_nan() { - min_idx = d; - break; - } - if val < min_val { - min_val = val; - min_idx = d; - } - } - output[o * layout.inner + r] = min_idx as i64; - } - } - } - DataType::Float64 => { - let input = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut min_val = f64::INFINITY; - let mut min_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if val.is_nan() { - min_idx = d; - break; - } - if val < min_val { - min_val = val; - min_idx = d; - } - } - output[o * layout.inner + r] = min_idx as i64; - } - } - } - DataType::Int32 => { - let input = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut min_val = i32::MAX; - let mut min_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if val < min_val { - min_val = val; - min_idx = d; - } - } - output[o * layout.inner + r] = min_idx as i64; - } - } - } - DataType::Int64 => { - let input = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut min_val = i64::MAX; - let mut min_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if val < min_val { - min_val = val; - min_idx = d; - } - } - output[o * layout.inner + r] = min_idx as i64; - } - } - } - DataType::Bool => { - let input = tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut min_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - if !input[idx] { - min_idx = d; - break; - } - } - output[o * layout.inner + r] = min_idx as i64; - } - } - } - } - - Ok(Tensor::new( - Arc::new(result_data), - layout.output_shape, - DataType::Int64, - tensor.device(), - false, - )) -} - -macro_rules! cumprod_forward { - ($name:ident, $get:ident, $get_mut:ident, $t:ty) => { - fn $name(tensor: &Tensor, result_data: &mut TensorData, dim: usize) -> Result<()> { - let input_data = tensor - .data() - .$get() - .ok_or_else(|| MinitensorError::internal_error("Failed to get slice"))?; - let output = result_data - .$get_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable slice"))?; - let shape = tensor.shape().dims(); - - if tensor.ndim() == 1 { - if dim != 0 { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - let mut acc: $t = 1 as $t; - for i in 0..input_data.len() { - acc *= input_data[i]; - output[i] = acc; - } - } else if tensor.ndim() == 2 { - let rows = shape[0]; - let cols = shape[1]; - match dim { - 0 => { - let out_ptr = output.as_mut_ptr() as usize; - (0..cols).into_par_iter().for_each(|c| { - let out_ptr = out_ptr as *mut $t; - let mut acc: $t = 1 as $t; - for r in 0..rows { - let idx = r * cols + c; - acc *= input_data[idx]; - unsafe { - *out_ptr.add(idx) = acc; - } - } - }); - } - 1 => { - input_data - .par_chunks_exact(cols) - .zip(output.par_chunks_mut(cols)) - .for_each(|(in_row, out_row)| { - let mut acc: $t = 1 as $t; - for i in 0..cols { - acc *= in_row[i]; - out_row[i] = acc; - } - }); - } - _ => return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())), - } - } else { - let dim_size = shape[dim]; - let inner = shape[dim + 1..].iter().product::(); - let outer = shape[..dim].iter().product::(); - let total = outer * inner; - let out_ptr = output.as_mut_ptr() as usize; - (0..total).into_par_iter().for_each(|idx| { - let out_ptr = out_ptr as *mut $t; - let o = idx / inner; - let r = idx % inner; - let mut acc: $t = 1 as $t; - let mut base = o * dim_size * inner + r; - for _ in 0..dim_size { - acc *= input_data[base]; - unsafe { - *out_ptr.add(base) = acc; - } - base += inner; - } - }); - } - Ok(()) - } - }; -} - -macro_rules! cumprod_backward { - ($name:ident, $get:ident, $get_mut:ident, $t:ty) => { - fn $name( - input: &Tensor, - output: &Tensor, - grad: &Tensor, - result_data: &mut TensorData, - dim: usize, - ) -> Result<()> { - let input_data = input - .data() - .$get() - .ok_or_else(|| MinitensorError::internal_error("Failed to get slice"))?; - let out_data = output - .data() - .$get() - .ok_or_else(|| MinitensorError::internal_error("Failed to get slice"))?; - let grad_data = grad - .data() - .$get() - .ok_or_else(|| MinitensorError::internal_error("Failed to get slice"))?; - let output = result_data - .$get_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable slice"))?; - let shape = input.shape().dims(); - - if input.ndim() == 1 { - if dim != 0 { - return Err(MinitensorError::index_error(dim as isize, 0, input.ndim())); - } - let len = input_data.len(); - // count zeros and index - let mut zero_count = 0; - let mut zero_idx = 0; - for i in 0..len { - if input_data[i] == 0 as $t { - zero_count += 1; - if zero_count == 1 { - zero_idx = i; - } - } - } - if zero_count == 0 { - let mut s: $t = 0 as $t; - for i in (0..len).rev() { - s += grad_data[i] * out_data[i]; - output[i] = s / input_data[i]; - } - } else if zero_count == 1 { - let mut s: $t = 0 as $t; - for i in (0..zero_idx).rev() { - s += grad_data[i] * out_data[i]; - output[i] = s / input_data[i]; - } - let mut prefix: $t = 1 as $t; - for i in 0..zero_idx { - prefix *= input_data[i]; - } - let mut prod_suffix: $t = 1 as $t; - let mut grad_zero: $t = 0 as $t; - for j in zero_idx..len { - grad_zero += grad_data[j] * prod_suffix; - if j + 1 < len { - prod_suffix *= input_data[j + 1]; - } - } - output[zero_idx] = grad_zero * prefix; - for i in zero_idx + 1..len { - output[i] = 0 as $t; - } - } else { - for i in 0..len { - output[i] = 0 as $t; - } - } - } else if input.ndim() == 2 { - let rows = shape[0]; - let cols = shape[1]; - match dim { - 0 => { - for c in 0..cols { - let mut zero_count = 0; - let mut zero_idx = 0; - for r in 0..rows { - let idx = r * cols + c; - if input_data[idx] == 0 as $t { - zero_count += 1; - if zero_count == 1 { - zero_idx = r; - } - } - } - if zero_count == 0 { - let mut s: $t = 0 as $t; - for r in (0..rows).rev() { - let idx = r * cols + c; - s += grad_data[idx] * out_data[idx]; - output[idx] = s / input_data[idx]; - } - } else if zero_count == 1 { - let mut s: $t = 0 as $t; - for r in (0..zero_idx).rev() { - let idx = r * cols + c; - s += grad_data[idx] * out_data[idx]; - output[idx] = s / input_data[idx]; - } - let mut prefix: $t = 1 as $t; - for r in 0..zero_idx { - prefix *= input_data[r * cols + c]; - } - let mut prod_suffix: $t = 1 as $t; - let mut grad_zero: $t = 0 as $t; - for r in zero_idx..rows { - let idx = r * cols + c; - grad_zero += grad_data[idx] * prod_suffix; - if r + 1 < rows { - prod_suffix *= input_data[(r + 1) * cols + c]; - } - } - let zero_index = zero_idx * cols + c; - output[zero_index] = grad_zero * prefix; - for r in zero_idx + 1..rows { - let idx = r * cols + c; - output[idx] = 0 as $t; - } - } else { - for r in 0..rows { - let idx = r * cols + c; - output[idx] = 0 as $t; - } - } - } - } - 1 => { - for r in 0..rows { - let base = r * cols; - let mut zero_count = 0; - let mut zero_idx = 0; - for c in 0..cols { - let idx = base + c; - if input_data[idx] == 0 as $t { - zero_count += 1; - if zero_count == 1 { - zero_idx = c; - } - } - } - if zero_count == 0 { - let mut s: $t = 0 as $t; - for c in (0..cols).rev() { - let idx = base + c; - s += grad_data[idx] * out_data[idx]; - output[idx] = s / input_data[idx]; - } - } else if zero_count == 1 { - let mut s: $t = 0 as $t; - for c in (0..zero_idx).rev() { - let idx = base + c; - s += grad_data[idx] * out_data[idx]; - output[idx] = s / input_data[idx]; - } - let mut prefix: $t = 1 as $t; - for c in 0..zero_idx { - prefix *= input_data[base + c]; - } - let mut prod_suffix: $t = 1 as $t; - let mut grad_zero: $t = 0 as $t; - for c in zero_idx..cols { - let idx = base + c; - grad_zero += grad_data[idx] * prod_suffix; - if c + 1 < cols { - prod_suffix *= input_data[base + c + 1]; - } - } - output[base + zero_idx] = grad_zero * prefix; - for c in zero_idx + 1..cols { - output[base + c] = 0 as $t; - } - } else { - for c in 0..cols { - output[base + c] = 0 as $t; - } - } - } - } - _ => return Err(MinitensorError::index_error(dim as isize, 0, input.ndim())), - } - } else { - let dim_size = shape[dim]; - let inner = shape[dim + 1..].iter().product::(); - let outer = shape[..dim].iter().product::(); - let total = outer * inner; - for idx in 0..total { - let o = idx / inner; - let r = idx % inner; - let base = o * dim_size * inner + r; - let mut zero_count = 0; - let mut zero_idx = 0; - for d in 0..dim_size { - let i = base + d * inner; - if input_data[i] == 0 as $t { - zero_count += 1; - if zero_count == 1 { - zero_idx = d; - } - } - } - if zero_count == 0 { - let mut s: $t = 0 as $t; - for d in (0..dim_size).rev() { - let i = base + d * inner; - s += grad_data[i] * out_data[i]; - output[i] = s / input_data[i]; - } - } else if zero_count == 1 { - let mut s: $t = 0 as $t; - for d in (0..zero_idx).rev() { - let i = base + d * inner; - s += grad_data[i] * out_data[i]; - output[i] = s / input_data[i]; - } - let mut prefix: $t = 1 as $t; - for d in 0..zero_idx { - prefix *= input_data[base + d * inner]; - } - let mut prod_suffix: $t = 1 as $t; - let mut grad_zero: $t = 0 as $t; - for d in zero_idx..dim_size { - let i = base + d * inner; - grad_zero += grad_data[i] * prod_suffix; - if d + 1 < dim_size { - prod_suffix *= input_data[base + (d + 1) * inner]; - } - } - let zero_index = base + zero_idx * inner; - output[zero_index] = grad_zero * prefix; - for d in zero_idx + 1..dim_size { - output[base + d * inner] = 0 as $t; - } - } else { - for d in 0..dim_size { - output[base + d * inner] = 0 as $t; - } - } - } - } - Ok(()) - } - }; -} - -macro_rules! cumsum_forward { - ($name:ident, $get:ident, $get_mut:ident, $t:ty) => { - fn $name(tensor: &Tensor, result_data: &mut TensorData, dim: usize) -> Result<()> { - let input_data = tensor - .data() - .$get() - .ok_or_else(|| MinitensorError::internal_error("Failed to get slice"))?; - let output = result_data - .$get_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable slice"))?; - let shape = tensor.shape().dims(); - - if tensor.ndim() == 1 { - if dim != 0 { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - let mut acc: $t = 0 as $t; - for i in 0..input_data.len() { - acc += input_data[i]; - output[i] = acc; - } - } else if tensor.ndim() == 2 { - let rows = shape[0]; - let cols = shape[1]; - match dim { - 0 => { - let out_ptr = output.as_mut_ptr() as usize; - (0..cols).into_par_iter().for_each(|c| { - let out_ptr = out_ptr as *mut $t; - let mut acc: $t = 0 as $t; - for r in 0..rows { - let idx = r * cols + c; - acc += input_data[idx]; - unsafe { - *out_ptr.add(idx) = acc; - } - } - }); - } - 1 => { - input_data - .par_chunks_exact(cols) - .zip(output.par_chunks_mut(cols)) - .for_each(|(in_row, out_row)| { - let mut acc: $t = 0 as $t; - for i in 0..cols { - acc += in_row[i]; - out_row[i] = acc; - } - }); - } - _ => return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())), - } - } else { - let dim_size = shape[dim]; - let inner = shape[dim + 1..].iter().product::(); - let outer = shape[..dim].iter().product::(); - let total = outer * inner; - let out_ptr = output.as_mut_ptr() as usize; - (0..total).into_par_iter().for_each(|idx| { - let out_ptr = out_ptr as *mut $t; - let o = idx / inner; - let r = idx % inner; - let mut acc: $t = 0 as $t; - let mut base = o * dim_size * inner + r; - for _ in 0..dim_size { - acc += input_data[base]; - unsafe { - *out_ptr.add(base) = acc; - } - base += inner; - } - }); - } - Ok(()) - } - }; -} - -macro_rules! cumsum_backward { - ($name:ident, $get:ident, $get_mut:ident, $t:ty) => { - fn $name(tensor: &Tensor, result_data: &mut TensorData, dim: usize) -> Result<()> { - let input_data = tensor - .data() - .$get() - .ok_or_else(|| MinitensorError::internal_error("Failed to get slice"))?; - let output = result_data - .$get_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable slice"))?; - let shape = tensor.shape().dims(); - - if tensor.ndim() == 1 { - if dim != 0 { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - let mut acc: $t = 0 as $t; - for i in (0..input_data.len()).rev() { - acc += input_data[i]; - output[i] = acc; - } - } else if tensor.ndim() == 2 { - let rows = shape[0]; - let cols = shape[1]; - match dim { - 0 => { - let out_ptr = output.as_mut_ptr() as usize; - (0..cols).into_par_iter().for_each(|c| { - let out_ptr = out_ptr as *mut $t; - let mut acc: $t = 0 as $t; - for r in (0..rows).rev() { - let idx = r * cols + c; - acc += input_data[idx]; - unsafe { - *out_ptr.add(idx) = acc; - } - } - }); - } - 1 => { - input_data - .par_chunks_exact(cols) - .zip(output.par_chunks_mut(cols)) - .for_each(|(in_row, out_row)| { - let mut acc: $t = 0 as $t; - for i in (0..cols).rev() { - acc += in_row[i]; - out_row[i] = acc; - } - }); - } - _ => return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())), - } - } else { - let dim_size = shape[dim]; - let inner = shape[dim + 1..].iter().product::(); - let outer = shape[..dim].iter().product::(); - let total = outer * inner; - let out_ptr = output.as_mut_ptr() as usize; - (0..total).into_par_iter().for_each(|idx| { - let out_ptr = out_ptr as *mut $t; - let o = idx / inner; - let r = idx % inner; - let mut acc: $t = 0 as $t; - let mut base = o * dim_size * inner + r + (dim_size - 1) * inner; - for _ in 0..dim_size { - acc += input_data[base]; - unsafe { - *out_ptr.add(base) = acc; - } - if base >= inner { - base -= inner; - } - } - }); - } - Ok(()) - } - }; -} - -cumprod_forward!(cumprod_f32, as_f32_slice, as_f32_slice_mut, f32); -cumprod_forward!(cumprod_f64, as_f64_slice, as_f64_slice_mut, f64); -cumprod_forward!(cumprod_i32, as_i32_slice, as_i32_slice_mut, i32); -cumprod_forward!(cumprod_i64, as_i64_slice, as_i64_slice_mut, i64); - -cumprod_backward!(cumprod_backward_f32, as_f32_slice, as_f32_slice_mut, f32); -cumprod_backward!(cumprod_backward_f64, as_f64_slice, as_f64_slice_mut, f64); - -cumsum_forward!(cumsum_f32, as_f32_slice, as_f32_slice_mut, f32); -cumsum_forward!(cumsum_f64, as_f64_slice, as_f64_slice_mut, f64); -cumsum_forward!(cumsum_i32, as_i32_slice, as_i32_slice_mut, i32); -cumsum_forward!(cumsum_i64, as_i64_slice, as_i64_slice_mut, i64); - -cumsum_backward!(cumsum_backward_f32, as_f32_slice, as_f32_slice_mut, f32); -cumsum_backward!(cumsum_backward_f64, as_f64_slice, as_f64_slice_mut, f64); -cumsum_backward!(cumsum_backward_i32, as_i32_slice, as_i32_slice_mut, i32); -cumsum_backward!(cumsum_backward_i64, as_i64_slice, as_i64_slice_mut, i64); - -#[cfg(test)] -mod tests { - use super::*; - use crate::device::Device; - use std::sync::Arc; - - fn create_tensor_f32(data: Vec, shape: Vec) -> Tensor { - let shape_obj = Shape::new(shape.clone()); - let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Float32); - tensor_data - .as_f32_slice_mut() - .unwrap() - .copy_from_slice(&data); - Tensor::new( - Arc::new(tensor_data), - shape_obj, - DataType::Float32, - Device::cpu(), - false, - ) - } - - fn create_tensor_i32(data: Vec, shape: Vec) -> Tensor { - let shape_obj = Shape::new(shape.clone()); - let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Int32); - tensor_data - .as_i32_slice_mut() - .unwrap() - .copy_from_slice(&data); - Tensor::new( - Arc::new(tensor_data), - shape_obj, - DataType::Int32, - Device::cpu(), - false, - ) - } - - fn create_tensor_bool(data: Vec, shape: Vec) -> Tensor { - let shape_obj = Shape::new(shape.clone()); - let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Bool); - tensor_data - .as_bool_slice_mut() - .unwrap() - .copy_from_slice(&data); - Tensor::new( - Arc::new(tensor_data), - shape_obj, - DataType::Bool, - Device::cpu(), - false, - ) - } - - #[test] - fn test_median_global_even_length() { - let t = create_tensor_f32(vec![3.0, 1.0, 4.0, 2.0], vec![4]); - let (value, indices) = median(&t, None, false).unwrap(); - assert!(indices.is_none()); - assert!(value.shape().is_scalar()); - let result = value.data().as_f32_slice().unwrap(); - assert_eq!(result, &[2.0]); - } - - #[test] - fn test_median_with_dim_returns_indices() { - let t = create_tensor_f32(vec![1.0, 3.0, 2.0, 4.0, 6.0, 5.0], vec![2, 3]); - let (values, indices_opt) = median(&t, Some(1), false).unwrap(); - let indices = indices_opt.unwrap(); - assert_eq!(values.shape().dims(), &[2]); - assert_eq!(indices.shape().dims(), &[2]); - let values_slice = values.data().as_f32_slice().unwrap(); - let indices_slice = indices.data().as_i64_slice().unwrap(); - assert_eq!(values_slice, &[2.0, 5.0]); - assert_eq!(indices_slice, &[2, 2]); - } - - #[test] - fn test_median_keepdim_preserves_rank() { - let t = create_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]); - let (values, indices_opt) = median(&t, Some(1), true).unwrap(); - let indices = indices_opt.unwrap(); - assert_eq!(values.shape().dims(), &[2, 1]); - assert_eq!(indices.shape().dims(), &[2, 1]); - assert_eq!(values.data().as_f32_slice().unwrap(), &[1.0, 3.0]); - assert_eq!(indices.data().as_i64_slice().unwrap(), &[0, 0]); - } - - #[test] - fn test_median_empty_tensor_errors() { - let t = create_tensor_f32(vec![], vec![0]); - assert!(median(&t, None, false).is_err()); - } - - #[test] - fn test_quantiles_all_multiple_probs() { - let t = create_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![4]); - let result = quantiles( - &t, - &[0.25, 0.75], - None, - false, - QuantileInterpolation::Linear, - ) - .unwrap(); - assert_eq!(result.shape().dims(), &[2]); - let values = result.data().as_f32_slice().unwrap(); - assert!((values[0] - 1.75).abs() < 1e-6); - assert!((values[1] - 3.25).abs() < 1e-6); - } - - #[test] - fn test_quantiles_dim_keepdim_layout() { - let t = create_tensor_f32(vec![1.0, 3.0, 2.0, 4.0, 6.0, 5.0], vec![2, 3]); - let result = quantiles( - &t, - &[0.5, 0.9], - Some(1), - true, - QuantileInterpolation::Linear, - ) - .unwrap(); - assert_eq!(result.shape().dims(), &[2, 2, 1]); - let values = result.data().as_f32_slice().unwrap(); - let expected = [2.0, 5.0, 2.8, 5.8]; - for (value, target) in values.iter().zip(expected.iter()) { - assert!((*value - *target).abs() < 1e-6); - } - } - - #[test] - fn test_argmax_along_dim() { - let t = create_tensor_f32(vec![1.0, 5.0, 3.0, 4.0, 2.0, 6.0], vec![2, 3]); - let result = argmax(&t, Some(1), false).unwrap(); - let res = result.data().as_i64_slice().unwrap(); - assert_eq!(res, &[1, 2]); - } - - #[test] - fn test_argmin_along_dim_keepdim() { - let t = create_tensor_f32(vec![1.0, 5.0, 3.0, 4.0, 2.0, 6.0], vec![2, 3]); - let result = argmin(&t, Some(1), true).unwrap(); - assert_eq!(result.shape().dims(), &[2, 1]); - let res = result.data().as_i64_slice().unwrap(); - assert_eq!(res, &[0, 1]); - } - - #[test] - fn test_all_any_global() { - let t = create_tensor_i32(vec![1, 0, 2, 3], vec![2, 2]); - let all_res = all(&t, None, false).unwrap(); - let any_res = any(&t, None, false).unwrap(); - assert_eq!(all_res.data().as_bool_slice().unwrap()[0], false); - assert_eq!(any_res.data().as_bool_slice().unwrap()[0], true); - } - - #[test] - fn test_all_along_dim() { - let t = create_tensor_bool(vec![true, false, true, true], vec![2, 2]); - let res = all(&t, Some(1), false).unwrap(); - assert_eq!(res.data().as_bool_slice().unwrap(), &[false, true]); - } - - #[test] - fn test_sum_global_and_keepdim() { - let t = create_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]); - let s = sum(&t, None, false).unwrap(); - assert_eq!(s.shape().dims(), &[] as &[usize]); - assert_eq!(s.data().as_f32_slice().unwrap()[0], 10.0); - let s_keep = sum(&t, None, true).unwrap(); - assert_eq!(s_keep.shape().dims(), &[1, 1]); - assert_eq!(s_keep.data().as_f32_slice().unwrap()[0], 10.0); - } - - #[test] - fn test_topk_largest_float() { - let t = create_tensor_f32(vec![1.0, 3.0, 2.0, 4.0, -1.0, 5.0], vec![2, 3]); - let (values, indices) = topk(&t, 2, Some(1), true, true).unwrap(); - assert_eq!(values.shape().dims(), &[2, 2]); - assert_eq!(indices.shape().dims(), &[2, 2]); - let values_slice = values.data().as_f32_slice().unwrap(); - let indices_slice = indices.data().as_i64_slice().unwrap(); - assert_eq!(values_slice, &[3.0, 2.0, 5.0, 4.0]); - assert_eq!(indices_slice, &[1, 2, 2, 0]); - } - - #[test] - fn test_topk_smallest_unsorted() { - let t = create_tensor_f32(vec![1.0, -2.0, 3.5, 0.0], vec![4]); - let (values, indices) = topk(&t, 2, None, false, false).unwrap(); - assert_eq!(values.shape().dims(), &[2]); - let mut pairs: Vec<(i64, f32)> = indices - .data() - .as_i64_slice() - .unwrap() - .iter() - .zip(values.data().as_f32_slice().unwrap()) - .map(|(&i, &v)| (i, v)) - .collect(); - pairs.sort_by_key(|p| p.0); - assert_eq!(pairs, vec![(1, -2.0), (3, 0.0)]); - } - - #[test] - fn test_topk_sorted_partial_ties_by_first_index() { - let t = create_tensor_f32(vec![1.0, 5.0, 5.0, 4.0, 3.0, 2.0], vec![6]); - let (values, indices) = topk(&t, 3, None, true, true).unwrap(); - - assert_eq!(values.data().as_f32_slice().unwrap(), &[5.0, 5.0, 4.0]); - assert_eq!(indices.data().as_i64_slice().unwrap(), &[1, 2, 3]); - } - - #[test] - fn test_topk_sorted_smallest_bool() { - let t = create_tensor_bool(vec![true, false, true, false], vec![4]); - let (values, indices) = topk(&t, 2, None, false, true).unwrap(); - - assert_eq!(values.data().as_bool_slice().unwrap(), &[false, false]); - assert_eq!(indices.data().as_i64_slice().unwrap(), &[1, 3]); - } - - #[test] - fn test_topk_sorted_partial_nan_ordering() { - let t = create_tensor_f32(vec![1.0, f32::NAN, 3.0, f32::NAN, 2.0], vec![5]); - let (values, indices) = topk(&t, 3, None, true, true).unwrap(); - let values = values.data().as_f32_slice().unwrap(); - - assert!(values[0].is_nan()); - assert!(values[1].is_nan()); - assert_eq!(values[2], 3.0); - assert_eq!(indices.data().as_i64_slice().unwrap(), &[1, 3, 2]); - } - - #[test] - fn test_sum_along_dim() { - let t = create_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]); - let res = sum(&t, Some(vec![0]), false).unwrap(); - assert_eq!(res.shape().dims(), &[2]); - assert_eq!(res.data().as_f32_slice().unwrap(), &[4.0, 6.0]); - } - - #[test] - fn test_sum_bool_error() { - let t = create_tensor_bool(vec![true, false, true, true], vec![2, 2]); - assert!(sum(&t, Some(vec![0]), false).is_err()); - } - - #[test] - fn test_sum_multi_dim() { - let t = create_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]); - let res = sum(&t, Some(vec![0, 1]), false).unwrap(); - assert!(res.shape().is_scalar()); - assert_eq!(res.data().as_f32_slice().unwrap()[0], 10.0); - let res_keep = sum(&t, Some(vec![0, 1]), true).unwrap(); - assert_eq!(res_keep.shape().dims(), &[1, 1]); - assert_eq!(res_keep.data().as_f32_slice().unwrap()[0], 10.0); - } - - #[test] - fn test_prod_global_and_keepdim() { - let t = create_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]); - let p = prod(&t, None, false).unwrap(); - assert_eq!(p.data().as_f32_slice().unwrap()[0], 24.0); - let p_keep = prod(&t, None, true).unwrap(); - assert_eq!(p_keep.shape().dims(), &[1, 1]); - assert_eq!(p_keep.data().as_f32_slice().unwrap()[0], 24.0); - } - - #[test] - fn test_prod_along_dim() { - let t = create_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]); - let res = prod(&t, Some(vec![0]), false).unwrap(); - assert_eq!(res.data().as_f32_slice().unwrap(), &[3.0, 8.0]); - } - - #[test] - fn test_mean_along_dim() { - let t = create_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]); - let res = mean(&t, Some(vec![1]), true).unwrap(); - assert_eq!(res.shape().dims(), &[2, 1]); - assert_eq!(res.data().as_f32_slice().unwrap(), &[1.5, 3.5]); - } - - #[test] - fn test_mean_int_support() { - let t = create_tensor_i32(vec![1, 2, 3, 4], vec![2, 2]); - let res = mean(&t, Some(vec![0isize]), false).unwrap(); - assert_eq!(res.dtype(), DataType::Float32); - assert_eq!(res.data().as_f32_slice().unwrap(), &[2.0, 3.0]); - } -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::{ + error::{MinitensorError, Result}, + tensor::{DataType, Tensor, TensorData}, +}; +use rayon::prelude::*; +use std::sync::Arc; + +pub(crate) fn argmin_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { + let layout = reduction_layout(tensor, dim, keepdim)?; + let mut result_data = TensorData::zeros_on_device( + layout.output_shape.numel(), + DataType::Int64, + tensor.device(), + ); + + let output = result_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + match tensor.dtype() { + DataType::Float32 => { + let input = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut min_val = f32::INFINITY; + let mut min_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if val.is_nan() { + min_idx = d; + break; + } + if val < min_val { + min_val = val; + min_idx = d; + } + } + output[o * layout.inner + r] = min_idx as i64; + } + } + } + DataType::Float64 => { + let input = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut min_val = f64::INFINITY; + let mut min_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if val.is_nan() { + min_idx = d; + break; + } + if val < min_val { + min_val = val; + min_idx = d; + } + } + output[o * layout.inner + r] = min_idx as i64; + } + } + } + DataType::Int32 => { + let input = tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut min_val = i32::MAX; + let mut min_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if val < min_val { + min_val = val; + min_idx = d; + } + } + output[o * layout.inner + r] = min_idx as i64; + } + } + } + DataType::Int64 => { + let input = tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut min_val = i64::MAX; + let mut min_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if val < min_val { + min_val = val; + min_idx = d; + } + } + output[o * layout.inner + r] = min_idx as i64; + } + } + } + DataType::Bool => { + let input = tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut min_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + if !input[idx] { + min_idx = d; + break; + } + } + output[o * layout.inner + r] = min_idx as i64; + } + } + } + } + + Ok(Tensor::new( + Arc::new(result_data), + layout.output_shape, + DataType::Int64, + tensor.device(), + false, + )) +} + +macro_rules! cumprod_forward { + ($name:ident, $get:ident, $get_mut:ident, $t:ty) => { + pub(crate) fn $name( + tensor: &Tensor, + result_data: &mut TensorData, + dim: usize, + ) -> Result<()> { + let input_data = tensor + .data() + .$get() + .ok_or_else(|| MinitensorError::internal_error("Failed to get slice"))?; + let output = result_data + .$get_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable slice"))?; + let shape = tensor.shape().dims(); + + if tensor.ndim() == 1 { + if dim != 0 { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + let mut acc: $t = 1 as $t; + for i in 0..input_data.len() { + acc *= input_data[i]; + output[i] = acc; + } + } else if tensor.ndim() == 2 { + let rows = shape[0]; + let cols = shape[1]; + match dim { + 0 => { + let out_ptr = output.as_mut_ptr() as usize; + (0..cols).into_par_iter().for_each(|c| { + let out_ptr = out_ptr as *mut $t; + let mut acc: $t = 1 as $t; + for r in 0..rows { + let idx = r * cols + c; + acc *= input_data[idx]; + unsafe { + *out_ptr.add(idx) = acc; + } + } + }); + } + 1 => { + input_data + .par_chunks_exact(cols) + .zip(output.par_chunks_mut(cols)) + .for_each(|(in_row, out_row)| { + let mut acc: $t = 1 as $t; + for i in 0..cols { + acc *= in_row[i]; + out_row[i] = acc; + } + }); + } + _ => return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())), + } + } else { + let dim_size = shape[dim]; + let inner = shape[dim + 1..].iter().product::(); + let outer = shape[..dim].iter().product::(); + let total = outer * inner; + let out_ptr = output.as_mut_ptr() as usize; + (0..total).into_par_iter().for_each(|idx| { + let out_ptr = out_ptr as *mut $t; + let o = idx / inner; + let r = idx % inner; + let mut acc: $t = 1 as $t; + let mut base = o * dim_size * inner + r; + for _ in 0..dim_size { + acc *= input_data[base]; + unsafe { + *out_ptr.add(base) = acc; + } + base += inner; + } + }); + } + Ok(()) + } + }; +} + +macro_rules! cumprod_backward { + ($name:ident, $get:ident, $get_mut:ident, $t:ty) => { + pub(crate) fn $name( + input: &Tensor, + output: &Tensor, + grad: &Tensor, + result_data: &mut TensorData, + dim: usize, + ) -> Result<()> { + let input_data = input + .data() + .$get() + .ok_or_else(|| MinitensorError::internal_error("Failed to get slice"))?; + let out_data = output + .data() + .$get() + .ok_or_else(|| MinitensorError::internal_error("Failed to get slice"))?; + let grad_data = grad + .data() + .$get() + .ok_or_else(|| MinitensorError::internal_error("Failed to get slice"))?; + let output = result_data + .$get_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable slice"))?; + let shape = input.shape().dims(); + + if input.ndim() == 1 { + if dim != 0 { + return Err(MinitensorError::index_error(dim as isize, 0, input.ndim())); + } + let len = input_data.len(); + // count zeros and index + let mut zero_count = 0; + let mut zero_idx = 0; + for i in 0..len { + if input_data[i] == 0 as $t { + zero_count += 1; + if zero_count == 1 { + zero_idx = i; + } + } + } + if zero_count == 0 { + let mut s: $t = 0 as $t; + for i in (0..len).rev() { + s += grad_data[i] * out_data[i]; + output[i] = s / input_data[i]; + } + } else if zero_count == 1 { + let mut s: $t = 0 as $t; + for i in (0..zero_idx).rev() { + s += grad_data[i] * out_data[i]; + output[i] = s / input_data[i]; + } + let mut prefix: $t = 1 as $t; + for i in 0..zero_idx { + prefix *= input_data[i]; + } + let mut prod_suffix: $t = 1 as $t; + let mut grad_zero: $t = 0 as $t; + for j in zero_idx..len { + grad_zero += grad_data[j] * prod_suffix; + if j + 1 < len { + prod_suffix *= input_data[j + 1]; + } + } + output[zero_idx] = grad_zero * prefix; + for i in zero_idx + 1..len { + output[i] = 0 as $t; + } + } else { + for i in 0..len { + output[i] = 0 as $t; + } + } + } else if input.ndim() == 2 { + let rows = shape[0]; + let cols = shape[1]; + match dim { + 0 => { + for c in 0..cols { + let mut zero_count = 0; + let mut zero_idx = 0; + for r in 0..rows { + let idx = r * cols + c; + if input_data[idx] == 0 as $t { + zero_count += 1; + if zero_count == 1 { + zero_idx = r; + } + } + } + if zero_count == 0 { + let mut s: $t = 0 as $t; + for r in (0..rows).rev() { + let idx = r * cols + c; + s += grad_data[idx] * out_data[idx]; + output[idx] = s / input_data[idx]; + } + } else if zero_count == 1 { + let mut s: $t = 0 as $t; + for r in (0..zero_idx).rev() { + let idx = r * cols + c; + s += grad_data[idx] * out_data[idx]; + output[idx] = s / input_data[idx]; + } + let mut prefix: $t = 1 as $t; + for r in 0..zero_idx { + prefix *= input_data[r * cols + c]; + } + let mut prod_suffix: $t = 1 as $t; + let mut grad_zero: $t = 0 as $t; + for r in zero_idx..rows { + let idx = r * cols + c; + grad_zero += grad_data[idx] * prod_suffix; + if r + 1 < rows { + prod_suffix *= input_data[(r + 1) * cols + c]; + } + } + let zero_index = zero_idx * cols + c; + output[zero_index] = grad_zero * prefix; + for r in zero_idx + 1..rows { + let idx = r * cols + c; + output[idx] = 0 as $t; + } + } else { + for r in 0..rows { + let idx = r * cols + c; + output[idx] = 0 as $t; + } + } + } + } + 1 => { + for r in 0..rows { + let base = r * cols; + let mut zero_count = 0; + let mut zero_idx = 0; + for c in 0..cols { + let idx = base + c; + if input_data[idx] == 0 as $t { + zero_count += 1; + if zero_count == 1 { + zero_idx = c; + } + } + } + if zero_count == 0 { + let mut s: $t = 0 as $t; + for c in (0..cols).rev() { + let idx = base + c; + s += grad_data[idx] * out_data[idx]; + output[idx] = s / input_data[idx]; + } + } else if zero_count == 1 { + let mut s: $t = 0 as $t; + for c in (0..zero_idx).rev() { + let idx = base + c; + s += grad_data[idx] * out_data[idx]; + output[idx] = s / input_data[idx]; + } + let mut prefix: $t = 1 as $t; + for c in 0..zero_idx { + prefix *= input_data[base + c]; + } + let mut prod_suffix: $t = 1 as $t; + let mut grad_zero: $t = 0 as $t; + for c in zero_idx..cols { + let idx = base + c; + grad_zero += grad_data[idx] * prod_suffix; + if c + 1 < cols { + prod_suffix *= input_data[base + c + 1]; + } + } + output[base + zero_idx] = grad_zero * prefix; + for c in zero_idx + 1..cols { + output[base + c] = 0 as $t; + } + } else { + for c in 0..cols { + output[base + c] = 0 as $t; + } + } + } + } + _ => return Err(MinitensorError::index_error(dim as isize, 0, input.ndim())), + } + } else { + let dim_size = shape[dim]; + let inner = shape[dim + 1..].iter().product::(); + let outer = shape[..dim].iter().product::(); + let total = outer * inner; + for idx in 0..total { + let o = idx / inner; + let r = idx % inner; + let base = o * dim_size * inner + r; + let mut zero_count = 0; + let mut zero_idx = 0; + for d in 0..dim_size { + let i = base + d * inner; + if input_data[i] == 0 as $t { + zero_count += 1; + if zero_count == 1 { + zero_idx = d; + } + } + } + if zero_count == 0 { + let mut s: $t = 0 as $t; + for d in (0..dim_size).rev() { + let i = base + d * inner; + s += grad_data[i] * out_data[i]; + output[i] = s / input_data[i]; + } + } else if zero_count == 1 { + let mut s: $t = 0 as $t; + for d in (0..zero_idx).rev() { + let i = base + d * inner; + s += grad_data[i] * out_data[i]; + output[i] = s / input_data[i]; + } + let mut prefix: $t = 1 as $t; + for d in 0..zero_idx { + prefix *= input_data[base + d * inner]; + } + let mut prod_suffix: $t = 1 as $t; + let mut grad_zero: $t = 0 as $t; + for d in zero_idx..dim_size { + let i = base + d * inner; + grad_zero += grad_data[i] * prod_suffix; + if d + 1 < dim_size { + prod_suffix *= input_data[base + (d + 1) * inner]; + } + } + let zero_index = base + zero_idx * inner; + output[zero_index] = grad_zero * prefix; + for d in zero_idx + 1..dim_size { + output[base + d * inner] = 0 as $t; + } + } else { + for d in 0..dim_size { + output[base + d * inner] = 0 as $t; + } + } + } + } + Ok(()) + } + }; +} + +macro_rules! cumsum_forward { + ($name:ident, $get:ident, $get_mut:ident, $t:ty) => { + pub(crate) fn $name( + tensor: &Tensor, + result_data: &mut TensorData, + dim: usize, + ) -> Result<()> { + let input_data = tensor + .data() + .$get() + .ok_or_else(|| MinitensorError::internal_error("Failed to get slice"))?; + let output = result_data + .$get_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable slice"))?; + let shape = tensor.shape().dims(); + + if tensor.ndim() == 1 { + if dim != 0 { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + let mut acc: $t = 0 as $t; + for i in 0..input_data.len() { + acc += input_data[i]; + output[i] = acc; + } + } else if tensor.ndim() == 2 { + let rows = shape[0]; + let cols = shape[1]; + match dim { + 0 => { + let out_ptr = output.as_mut_ptr() as usize; + (0..cols).into_par_iter().for_each(|c| { + let out_ptr = out_ptr as *mut $t; + let mut acc: $t = 0 as $t; + for r in 0..rows { + let idx = r * cols + c; + acc += input_data[idx]; + unsafe { + *out_ptr.add(idx) = acc; + } + } + }); + } + 1 => { + input_data + .par_chunks_exact(cols) + .zip(output.par_chunks_mut(cols)) + .for_each(|(in_row, out_row)| { + let mut acc: $t = 0 as $t; + for i in 0..cols { + acc += in_row[i]; + out_row[i] = acc; + } + }); + } + _ => return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())), + } + } else { + let dim_size = shape[dim]; + let inner = shape[dim + 1..].iter().product::(); + let outer = shape[..dim].iter().product::(); + let total = outer * inner; + let out_ptr = output.as_mut_ptr() as usize; + (0..total).into_par_iter().for_each(|idx| { + let out_ptr = out_ptr as *mut $t; + let o = idx / inner; + let r = idx % inner; + let mut acc: $t = 0 as $t; + let mut base = o * dim_size * inner + r; + for _ in 0..dim_size { + acc += input_data[base]; + unsafe { + *out_ptr.add(base) = acc; + } + base += inner; + } + }); + } + Ok(()) + } + }; +} + +macro_rules! cumsum_backward { + ($name:ident, $get:ident, $get_mut:ident, $t:ty) => { + pub(crate) fn $name( + tensor: &Tensor, + result_data: &mut TensorData, + dim: usize, + ) -> Result<()> { + let input_data = tensor + .data() + .$get() + .ok_or_else(|| MinitensorError::internal_error("Failed to get slice"))?; + let output = result_data + .$get_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable slice"))?; + let shape = tensor.shape().dims(); + + if tensor.ndim() == 1 { + if dim != 0 { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + let mut acc: $t = 0 as $t; + for i in (0..input_data.len()).rev() { + acc += input_data[i]; + output[i] = acc; + } + } else if tensor.ndim() == 2 { + let rows = shape[0]; + let cols = shape[1]; + match dim { + 0 => { + let out_ptr = output.as_mut_ptr() as usize; + (0..cols).into_par_iter().for_each(|c| { + let out_ptr = out_ptr as *mut $t; + let mut acc: $t = 0 as $t; + for r in (0..rows).rev() { + let idx = r * cols + c; + acc += input_data[idx]; + unsafe { + *out_ptr.add(idx) = acc; + } + } + }); + } + 1 => { + input_data + .par_chunks_exact(cols) + .zip(output.par_chunks_mut(cols)) + .for_each(|(in_row, out_row)| { + let mut acc: $t = 0 as $t; + for i in (0..cols).rev() { + acc += in_row[i]; + out_row[i] = acc; + } + }); + } + _ => return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())), + } + } else { + let dim_size = shape[dim]; + let inner = shape[dim + 1..].iter().product::(); + let outer = shape[..dim].iter().product::(); + let total = outer * inner; + let out_ptr = output.as_mut_ptr() as usize; + (0..total).into_par_iter().for_each(|idx| { + let out_ptr = out_ptr as *mut $t; + let o = idx / inner; + let r = idx % inner; + let mut acc: $t = 0 as $t; + let mut base = o * dim_size * inner + r + (dim_size - 1) * inner; + for _ in 0..dim_size { + acc += input_data[base]; + unsafe { + *out_ptr.add(base) = acc; + } + if base >= inner { + base -= inner; + } + } + }); + } + Ok(()) + } + }; +} + +cumprod_forward!(cumprod_f32, as_f32_slice, as_f32_slice_mut, f32); +cumprod_forward!(cumprod_f64, as_f64_slice, as_f64_slice_mut, f64); +cumprod_forward!(cumprod_i32, as_i32_slice, as_i32_slice_mut, i32); +cumprod_forward!(cumprod_i64, as_i64_slice, as_i64_slice_mut, i64); + +cumprod_backward!(cumprod_backward_f32, as_f32_slice, as_f32_slice_mut, f32); +cumprod_backward!(cumprod_backward_f64, as_f64_slice, as_f64_slice_mut, f64); + +cumsum_forward!(cumsum_f32, as_f32_slice, as_f32_slice_mut, f32); +cumsum_forward!(cumsum_f64, as_f64_slice, as_f64_slice_mut, f64); +cumsum_forward!(cumsum_i32, as_i32_slice, as_i32_slice_mut, i32); +cumsum_forward!(cumsum_i64, as_i64_slice, as_i64_slice_mut, i64); + +cumsum_backward!(cumsum_backward_f32, as_f32_slice, as_f32_slice_mut, f32); +cumsum_backward!(cumsum_backward_f64, as_f64_slice, as_f64_slice_mut, f64); +cumsum_backward!(cumsum_backward_i32, as_i32_slice, as_i32_slice_mut, i32); +cumsum_backward!(cumsum_backward_i64, as_i64_slice, as_i64_slice_mut, i64); + +#[cfg(test)] +mod tests { + use super::*; + use crate::device::Device; + use crate::tensor::Shape; + use std::sync::Arc; + + fn create_tensor_f32(data: Vec, shape: Vec) -> Tensor { + let shape_obj = Shape::new(shape.clone()); + let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Float32); + tensor_data + .as_f32_slice_mut() + .unwrap() + .copy_from_slice(&data); + Tensor::new( + Arc::new(tensor_data), + shape_obj, + DataType::Float32, + Device::cpu(), + false, + ) + } + + fn create_tensor_i32(data: Vec, shape: Vec) -> Tensor { + let shape_obj = Shape::new(shape.clone()); + let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Int32); + tensor_data + .as_i32_slice_mut() + .unwrap() + .copy_from_slice(&data); + Tensor::new( + Arc::new(tensor_data), + shape_obj, + DataType::Int32, + Device::cpu(), + false, + ) + } + + fn create_tensor_bool(data: Vec, shape: Vec) -> Tensor { + let shape_obj = Shape::new(shape.clone()); + let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Bool); + tensor_data + .as_bool_slice_mut() + .unwrap() + .copy_from_slice(&data); + Tensor::new( + Arc::new(tensor_data), + shape_obj, + DataType::Bool, + Device::cpu(), + false, + ) + } + + #[test] + fn test_median_global_even_length() { + let t = create_tensor_f32(vec![3.0, 1.0, 4.0, 2.0], vec![4]); + let (value, indices) = median(&t, None, false).unwrap(); + assert!(indices.is_none()); + assert!(value.shape().is_scalar()); + let result = value.data().as_f32_slice().unwrap(); + assert_eq!(result, &[2.0]); + } + + #[test] + fn test_median_with_dim_returns_indices() { + let t = create_tensor_f32(vec![1.0, 3.0, 2.0, 4.0, 6.0, 5.0], vec![2, 3]); + let (values, indices_opt) = median(&t, Some(1), false).unwrap(); + let indices = indices_opt.unwrap(); + assert_eq!(values.shape().dims(), &[2]); + assert_eq!(indices.shape().dims(), &[2]); + let values_slice = values.data().as_f32_slice().unwrap(); + let indices_slice = indices.data().as_i64_slice().unwrap(); + assert_eq!(values_slice, &[2.0, 5.0]); + assert_eq!(indices_slice, &[2, 2]); + } + + #[test] + fn test_median_keepdim_preserves_rank() { + let t = create_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]); + let (values, indices_opt) = median(&t, Some(1), true).unwrap(); + let indices = indices_opt.unwrap(); + assert_eq!(values.shape().dims(), &[2, 1]); + assert_eq!(indices.shape().dims(), &[2, 1]); + assert_eq!(values.data().as_f32_slice().unwrap(), &[1.0, 3.0]); + assert_eq!(indices.data().as_i64_slice().unwrap(), &[0, 0]); + } + + #[test] + fn test_median_empty_tensor_errors() { + let t = create_tensor_f32(vec![], vec![0]); + assert!(median(&t, None, false).is_err()); + } + + #[test] + fn test_quantiles_all_multiple_probs() { + let t = create_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![4]); + let result = quantiles( + &t, + &[0.25, 0.75], + None, + false, + QuantileInterpolation::Linear, + ) + .unwrap(); + assert_eq!(result.shape().dims(), &[2]); + let values = result.data().as_f32_slice().unwrap(); + assert!((values[0] - 1.75).abs() < 1e-6); + assert!((values[1] - 3.25).abs() < 1e-6); + } + + #[test] + fn test_quantiles_dim_keepdim_layout() { + let t = create_tensor_f32(vec![1.0, 3.0, 2.0, 4.0, 6.0, 5.0], vec![2, 3]); + let result = quantiles( + &t, + &[0.5, 0.9], + Some(1), + true, + QuantileInterpolation::Linear, + ) + .unwrap(); + assert_eq!(result.shape().dims(), &[2, 2, 1]); + let values = result.data().as_f32_slice().unwrap(); + let expected = [2.0, 5.0, 2.8, 5.8]; + for (value, target) in values.iter().zip(expected.iter()) { + assert!((*value - *target).abs() < 1e-6); + } + } + + #[test] + fn test_argmax_along_dim() { + let t = create_tensor_f32(vec![1.0, 5.0, 3.0, 4.0, 2.0, 6.0], vec![2, 3]); + let result = argmax(&t, Some(1), false).unwrap(); + let res = result.data().as_i64_slice().unwrap(); + assert_eq!(res, &[1, 2]); + } + + #[test] + fn test_argmin_along_dim_keepdim() { + let t = create_tensor_f32(vec![1.0, 5.0, 3.0, 4.0, 2.0, 6.0], vec![2, 3]); + let result = argmin(&t, Some(1), true).unwrap(); + assert_eq!(result.shape().dims(), &[2, 1]); + let res = result.data().as_i64_slice().unwrap(); + assert_eq!(res, &[0, 1]); + } + + #[test] + fn test_all_any_global() { + let t = create_tensor_i32(vec![1, 0, 2, 3], vec![2, 2]); + let all_res = all(&t, None, false).unwrap(); + let any_res = any(&t, None, false).unwrap(); + assert!(!all_res.data().as_bool_slice().unwrap()[0]); + assert!(any_res.data().as_bool_slice().unwrap()[0]); + } + + #[test] + fn test_all_along_dim() { + let t = create_tensor_bool(vec![true, false, true, true], vec![2, 2]); + let res = all(&t, Some(1), false).unwrap(); + assert_eq!(res.data().as_bool_slice().unwrap(), &[false, true]); + } + + #[test] + fn test_sum_global_and_keepdim() { + let t = create_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]); + let s = sum(&t, None, false).unwrap(); + assert_eq!(s.shape().dims(), &[] as &[usize]); + assert_eq!(s.data().as_f32_slice().unwrap()[0], 10.0); + let s_keep = sum(&t, None, true).unwrap(); + assert_eq!(s_keep.shape().dims(), &[1, 1]); + assert_eq!(s_keep.data().as_f32_slice().unwrap()[0], 10.0); + } + + #[test] + fn test_topk_largest_float() { + let t = create_tensor_f32(vec![1.0, 3.0, 2.0, 4.0, -1.0, 5.0], vec![2, 3]); + let (values, indices) = topk(&t, 2, Some(1), true, true).unwrap(); + assert_eq!(values.shape().dims(), &[2, 2]); + assert_eq!(indices.shape().dims(), &[2, 2]); + let values_slice = values.data().as_f32_slice().unwrap(); + let indices_slice = indices.data().as_i64_slice().unwrap(); + assert_eq!(values_slice, &[3.0, 2.0, 5.0, 4.0]); + assert_eq!(indices_slice, &[1, 2, 2, 0]); + } + + #[test] + fn test_topk_smallest_unsorted() { + let t = create_tensor_f32(vec![1.0, -2.0, 3.5, 0.0], vec![4]); + let (values, indices) = topk(&t, 2, None, false, false).unwrap(); + assert_eq!(values.shape().dims(), &[2]); + let mut pairs: Vec<(i64, f32)> = indices + .data() + .as_i64_slice() + .unwrap() + .iter() + .zip(values.data().as_f32_slice().unwrap()) + .map(|(&i, &v)| (i, v)) + .collect(); + pairs.sort_by_key(|p| p.0); + assert_eq!(pairs, vec![(1, -2.0), (3, 0.0)]); + } + + #[test] + fn test_topk_sorted_partial_ties_by_first_index() { + let t = create_tensor_f32(vec![1.0, 5.0, 5.0, 4.0, 3.0, 2.0], vec![6]); + let (values, indices) = topk(&t, 3, None, true, true).unwrap(); + + assert_eq!(values.data().as_f32_slice().unwrap(), &[5.0, 5.0, 4.0]); + assert_eq!(indices.data().as_i64_slice().unwrap(), &[1, 2, 3]); + } + + #[test] + fn test_topk_sorted_smallest_bool() { + let t = create_tensor_bool(vec![true, false, true, false], vec![4]); + let (values, indices) = topk(&t, 2, None, false, true).unwrap(); + + assert_eq!(values.data().as_bool_slice().unwrap(), &[false, false]); + assert_eq!(indices.data().as_i64_slice().unwrap(), &[1, 3]); + } + + #[test] + fn test_topk_sorted_partial_nan_ordering() { + let t = create_tensor_f32(vec![1.0, f32::NAN, 3.0, f32::NAN, 2.0], vec![5]); + let (values, indices) = topk(&t, 3, None, true, true).unwrap(); + let values = values.data().as_f32_slice().unwrap(); + + assert!(values[0].is_nan()); + assert!(values[1].is_nan()); + assert_eq!(values[2], 3.0); + assert_eq!(indices.data().as_i64_slice().unwrap(), &[1, 3, 2]); + } + + #[test] + fn test_sum_along_dim() { + let t = create_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]); + let res = sum(&t, Some(vec![0]), false).unwrap(); + assert_eq!(res.shape().dims(), &[2]); + assert_eq!(res.data().as_f32_slice().unwrap(), &[4.0, 6.0]); + } + + #[test] + fn test_sum_bool_error() { + let t = create_tensor_bool(vec![true, false, true, true], vec![2, 2]); + assert!(sum(&t, Some(vec![0]), false).is_err()); + } + + #[test] + fn test_sum_multi_dim() { + let t = create_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]); + let res = sum(&t, Some(vec![0, 1]), false).unwrap(); + assert!(res.shape().is_scalar()); + assert_eq!(res.data().as_f32_slice().unwrap()[0], 10.0); + let res_keep = sum(&t, Some(vec![0, 1]), true).unwrap(); + assert_eq!(res_keep.shape().dims(), &[1, 1]); + assert_eq!(res_keep.data().as_f32_slice().unwrap()[0], 10.0); + } + + #[test] + fn test_prod_global_and_keepdim() { + let t = create_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]); + let p = prod(&t, None, false).unwrap(); + assert_eq!(p.data().as_f32_slice().unwrap()[0], 24.0); + let p_keep = prod(&t, None, true).unwrap(); + assert_eq!(p_keep.shape().dims(), &[1, 1]); + assert_eq!(p_keep.data().as_f32_slice().unwrap()[0], 24.0); + } + + #[test] + fn test_prod_along_dim() { + let t = create_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]); + let res = prod(&t, Some(vec![0]), false).unwrap(); + assert_eq!(res.data().as_f32_slice().unwrap(), &[3.0, 8.0]); + } + + #[test] + fn test_mean_along_dim() { + let t = create_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]); + let res = mean(&t, Some(vec![1]), true).unwrap(); + assert_eq!(res.shape().dims(), &[2, 1]); + assert_eq!(res.data().as_f32_slice().unwrap(), &[1.5, 3.5]); + } + + #[test] + fn test_mean_int_support() { + let t = create_tensor_i32(vec![1, 2, 3, 4], vec![2, 2]); + let res = mean(&t, Some(vec![0isize]), false).unwrap(); + assert_eq!(res.dtype(), DataType::Float32); + assert_eq!(res.data().as_f32_slice().unwrap(), &[2.0, 3.0]); + } +} diff --git a/engine/src/operations/reduction/boolean.rs b/engine/src/operations/reduction/boolean.rs index e332fb3f..e058376a 100644 --- a/engine/src/operations/reduction/boolean.rs +++ b/engine/src/operations/reduction/boolean.rs @@ -1,776 +1,788 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -fn any_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { - if dim >= tensor.ndim() { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - - let input_shape = tensor.shape().dims(); - let mut output_shape = input_shape.to_vec(); - if keepdim { - output_shape[dim] = 1; - } else { - output_shape.remove(dim); - } - let output_shape_obj = Shape::new(output_shape.clone()); - let mut result_data = - TensorData::zeros_on_device(output_shape_obj.numel(), DataType::Bool, tensor.device()); - - let dim_size = input_shape[dim]; - let _outer = input_shape[..dim].iter().product::(); - let inner = input_shape[dim + 1..].iter().product::(); - let outer_stride = dim_size * inner; - - match tensor.dtype() { - DataType::Float32 => { - let input = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let output = result_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice") - })?; - output.par_iter_mut().enumerate().for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut val = false; - for d in 0..dim_size { - let in_idx = o * outer_stride + d * inner + r; - if input[in_idx] != 0.0 { - val = true; - break; - } - } - *out = val; - }); - } - DataType::Float64 => { - let input = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let output = result_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice") - })?; - output.par_iter_mut().enumerate().for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut val = false; - for d in 0..dim_size { - let in_idx = o * outer_stride + d * inner + r; - if input[in_idx] != 0.0 { - val = true; - break; - } - } - *out = val; - }); - } - DataType::Int32 => { - let input = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - let output = result_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice") - })?; - output.par_iter_mut().enumerate().for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut val = false; - for d in 0..dim_size { - let in_idx = o * outer_stride + d * inner + r; - if input[in_idx] != 0 { - val = true; - break; - } - } - *out = val; - }); - } - DataType::Int64 => { - let input = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - let output = result_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice") - })?; - output.par_iter_mut().enumerate().for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut val = false; - for d in 0..dim_size { - let in_idx = o * outer_stride + d * inner + r; - if input[in_idx] != 0 { - val = true; - break; - } - } - *out = val; - }); - } - DataType::Bool => { - let input = tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - let output = result_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice") - })?; - output.par_iter_mut().enumerate().for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut val = false; - for d in 0..dim_size { - let in_idx = o * outer_stride + d * inner + r; - if input[in_idx] { - val = true; - break; - } - } - *out = val; - }); - } - } - - Ok(Tensor::new( - Arc::new(result_data), - output_shape_obj, - DataType::Bool, - tensor.device(), - false, - )) -} - -/// Maximum value along specified dimension -pub fn max(tensor: &Tensor, dim: Option, keepdim: bool) -> Result { - let (output, norm_dim) = match dim { - None => { - // Find global maximum - let result_shape = if keepdim { - Shape::new(vec![1; tensor.ndim()]) - } else { - Shape::scalar() - }; - - let mut result_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => max_all_f32(tensor, &mut result_data)?, - DataType::Float64 => max_all_f64(tensor, &mut result_data)?, - DataType::Int32 => max_all_i32(tensor, &mut result_data)?, - DataType::Int64 => max_all_i64(tensor, &mut result_data)?, - DataType::Bool => max_all_bool(tensor, &mut result_data)?, - } - - ( - Tensor::new( - Arc::new(result_data), - result_shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ), - None, - ) - } - Some(d) => { - let d = normalize_dim(d, tensor.ndim())?; - (max_along_dim(tensor, d, keepdim)?, Some(d)) - } - }; - attach_minmax_grad(output, tensor, norm_dim, keepdim, true, false) -} - -/// Attach a [`GatherBackward`] gradient to a value tensor that was formed by -/// gathering the input along `dim` at `indices` (`sort`/`topk`). The forward is -/// `values = gather(input, dim, indices)`, so the backward scatters the gradient -/// straight back to the selected source positions. -fn attach_gather_like_grad( - values: Tensor, - input: &Tensor, - dim: usize, - indices: &Tensor, -) -> Result { - if !input.requires_grad() || !input.dtype().is_float() { - return Ok(values); - } - let index = indices - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("selection indices must be int64"))? - .to_vec(); - let grad_fn = Arc::new(GatherBackward { - input_id: input.id(), - input_shape: input.shape().dims().to_vec(), - dim, - index, - }); - let mut values = values; - values.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&values, Some(grad_fn))?; - Ok(values) -} - -/// Attach a [`MinMaxBackward`] gradient to a `min`/`max`/`nanmax`/`nanmin` value -/// reduction (`nan_aware` selects the NaN-ignoring recompute in the backward). -fn attach_minmax_grad( - output: Tensor, - input: &Tensor, - dim: Option, - keepdim: bool, - is_max: bool, - nan_aware: bool, -) -> Result { - if !input.requires_grad() || !input.dtype().is_float() { - return Ok(output); - } - let grad_fn = Arc::new(MinMaxBackward { - input_id: input.id(), - input: input.detach(), - dim, - keepdim, - is_max, - nan_aware, - }); - let mut output = output; - output.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output, Some(grad_fn))?; - Ok(output) -} - -/// Minimum value along specified dimension -pub fn min(tensor: &Tensor, dim: Option, keepdim: bool) -> Result { - let (output, norm_dim) = match dim { - None => { - // Find global minimum - let result_shape = if keepdim { - Shape::new(vec![1; tensor.ndim()]) - } else { - Shape::scalar() - }; - - let mut result_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => min_all_f32(tensor, &mut result_data)?, - DataType::Float64 => min_all_f64(tensor, &mut result_data)?, - DataType::Int32 => min_all_i32(tensor, &mut result_data)?, - DataType::Int64 => min_all_i64(tensor, &mut result_data)?, - DataType::Bool => min_all_bool(tensor, &mut result_data)?, - } - - ( - Tensor::new( - Arc::new(result_data), - result_shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ), - None, - ) - } - Some(d) => { - let d = normalize_dim(d, tensor.ndim())?; - (min_along_dim(tensor, d, keepdim)?, Some(d)) - } - }; - attach_minmax_grad(output, tensor, norm_dim, keepdim, false, false) -} - -/// NaN-aware maximum value along specified dimension -pub fn nanmax(tensor: &Tensor, dim: Option, keepdim: bool) -> Result { - if !tensor.dtype().is_float() { - return max(tensor, dim, keepdim); - } - - let (output, norm_dim) = match dim { - None => { - let result_shape = if keepdim { - Shape::new(vec![1; tensor.ndim()]) - } else { - Shape::scalar() - }; - - let mut result_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => nanmax_all_f32(tensor, &mut result_data)?, - DataType::Float64 => nanmax_all_f64(tensor, &mut result_data)?, - _ => unreachable!("nanmax only supports floating point tensors"), - } - - ( - Tensor::new( - Arc::new(result_data), - result_shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ), - None, - ) - } - Some(d) => { - let d = normalize_dim(d, tensor.ndim())?; - let (values, _) = nanmax_along_dim_with_indices(tensor, d, keepdim)?; - (values, Some(d)) - } - }; - attach_minmax_grad(output, tensor, norm_dim, keepdim, true, true) -} - -/// NaN-aware minimum value along specified dimension -pub fn nanmin(tensor: &Tensor, dim: Option, keepdim: bool) -> Result { - if !tensor.dtype().is_float() { - return min(tensor, dim, keepdim); - } - - let (output, norm_dim) = match dim { - None => { - let result_shape = if keepdim { - Shape::new(vec![1; tensor.ndim()]) - } else { - Shape::scalar() - }; - - let mut result_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => nanmin_all_f32(tensor, &mut result_data)?, - DataType::Float64 => nanmin_all_f64(tensor, &mut result_data)?, - _ => unreachable!("nanmin only supports floating point tensors"), - } - - ( - Tensor::new( - Arc::new(result_data), - result_shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ), - None, - ) - } - Some(d) => { - let d = normalize_dim(d, tensor.ndim())?; - let (values, _) = nanmin_along_dim_with_indices(tensor, d, keepdim)?; - (values, Some(d)) - } - }; - attach_minmax_grad(output, tensor, norm_dim, keepdim, false, true) -} - -/// Maximum values and their indices along specified dimension -pub fn max_with_indices(tensor: &Tensor, dim: isize, keepdim: bool) -> Result<(Tensor, Tensor)> { - let d = normalize_dim(dim, tensor.ndim())?; - let (values, indices) = max_along_dim_with_indices(tensor, d, keepdim)?; - let values = attach_minmax_grad(values, tensor, Some(d), keepdim, true, false)?; - Ok((values, indices)) -} - -/// NaN-aware maximum values and their indices along specified dimension -pub fn nanmax_with_indices(tensor: &Tensor, dim: isize, keepdim: bool) -> Result<(Tensor, Tensor)> { - if !tensor.dtype().is_float() { - return max_with_indices(tensor, dim, keepdim); - } - - let d = normalize_dim(dim, tensor.ndim())?; - let (values, indices) = nanmax_along_dim_with_indices(tensor, d, keepdim)?; - let values = attach_minmax_grad(values, tensor, Some(d), keepdim, true, true)?; - Ok((values, indices)) -} - -/// Minimum values and their indices along specified dimension -pub fn min_with_indices(tensor: &Tensor, dim: isize, keepdim: bool) -> Result<(Tensor, Tensor)> { - let d = normalize_dim(dim, tensor.ndim())?; - let (values, indices) = min_along_dim_with_indices(tensor, d, keepdim)?; - let values = attach_minmax_grad(values, tensor, Some(d), keepdim, false, false)?; - Ok((values, indices)) -} - -/// NaN-aware minimum values and their indices along specified dimension -pub fn nanmin_with_indices(tensor: &Tensor, dim: isize, keepdim: bool) -> Result<(Tensor, Tensor)> { - if !tensor.dtype().is_float() { - return min_with_indices(tensor, dim, keepdim); - } - - let d = normalize_dim(dim, tensor.ndim())?; - let (values, indices) = nanmin_along_dim_with_indices(tensor, d, keepdim)?; - let values = attach_minmax_grad(values, tensor, Some(d), keepdim, false, true)?; - Ok((values, indices)) -} - -/// Argument of maximum value along specified dimension -pub fn argmax(tensor: &Tensor, dim: Option, keepdim: bool) -> Result { - match dim { - None => { - // Find global argmax - let result_shape = if keepdim { - Shape::new(vec![1; tensor.ndim()]) - } else { - Shape::scalar() - }; - - let mut result_data = TensorData::zeros_on_device(1, DataType::Int64, tensor.device()); - - match tensor.dtype() { - DataType::Float32 => argmax_all_f32(tensor, &mut result_data)?, - DataType::Float64 => argmax_all_f64(tensor, &mut result_data)?, - DataType::Int32 => argmax_all_i32(tensor, &mut result_data)?, - DataType::Int64 => argmax_all_i64(tensor, &mut result_data)?, - DataType::Bool => argmax_all_bool(tensor, &mut result_data)?, - } - - Ok(Tensor::new( - Arc::new(result_data), - result_shape, - DataType::Int64, - tensor.device(), - false, // argmax doesn't require gradients - )) - } - Some(d) => { - let d = normalize_dim(d, tensor.ndim())?; - argmax_along_dim(tensor, d, keepdim) - } - } -} - -/// Argument of minimum value along specified dimension -pub fn argmin(tensor: &Tensor, dim: Option, keepdim: bool) -> Result { - match dim { - None => { - // Find global argmin - let result_shape = if keepdim { - Shape::new(vec![1; tensor.ndim()]) - } else { - Shape::scalar() - }; - - let mut result_data = TensorData::zeros_on_device(1, DataType::Int64, tensor.device()); - - match tensor.dtype() { - DataType::Float32 => argmin_all_f32(tensor, &mut result_data)?, - DataType::Float64 => argmin_all_f64(tensor, &mut result_data)?, - DataType::Int32 => argmin_all_i32(tensor, &mut result_data)?, - DataType::Int64 => argmin_all_i64(tensor, &mut result_data)?, - DataType::Bool => argmin_all_bool(tensor, &mut result_data)?, - } - - Ok(Tensor::new( - Arc::new(result_data), - result_shape, - DataType::Int64, - tensor.device(), - false, // argmin doesn't require gradients - )) - } - Some(d) => { - let d = normalize_dim(d, tensor.ndim())?; - argmin_along_dim(tensor, d, keepdim) - } - } -} - -#[inline] -fn select_topk_entries( - entries: &mut [(usize, T)], - k: usize, - sorted: bool, - compare: fn(&(usize, T), &(usize, T)) -> Ordering, -) { - if k == 0 || entries.is_empty() { - return; - } - - if k < entries.len() { - entries.select_nth_unstable_by(k - 1, compare); - if sorted { - entries[..k].sort_by(compare); - } - } else if sorted { - entries.sort_by(compare); - } -} - -/// Return the top-``k`` values and their indices along ``dim`` -pub fn topk( - tensor: &Tensor, - k: usize, - dim: Option, - largest: bool, - sorted: bool, -) -> Result<(Tensor, Tensor)> { - let ndim = tensor.ndim(); - - let axis = if ndim == 0 { - match dim { - Some(d) if d == 0 || d == -1 => 0, - Some(d) => return Err(MinitensorError::index_error(d, 0, 1)), - None => 0, - } - } else { - let dim_value = dim.unwrap_or(-1); - normalize_dim(dim_value, ndim)? - }; - - let dims = tensor.shape().dims(); - let dim_size = if dims.is_empty() { 1 } else { dims[axis] }; - - if k > dim_size { - return Err(MinitensorError::invalid_argument(format!( - "selected index k out of range for dimension {axis} with size {dim_size}" - ))); - } - - let output_dims = if dims.is_empty() { - vec![k] - } else { - let mut dims_vec = dims.to_vec(); - dims_vec[axis] = k; - dims_vec - }; - - let values_shape = Shape::new(output_dims.clone()); - let indices_shape = Shape::new(output_dims); - - let num_out = values_shape.numel(); - let mut values_data = TensorData::zeros_on_device(num_out, tensor.dtype(), tensor.device()); - let mut indices_data = TensorData::zeros_on_device(num_out, DataType::Int64, tensor.device()); - - if k == 0 || num_out == 0 { - let values = Tensor::new( - Arc::new(values_data), - values_shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - let indices = Tensor::new( - Arc::new(indices_data), - indices_shape, - DataType::Int64, - tensor.device(), - false, - ); - return Ok((values, indices)); - } - - let outer = if dims.is_empty() || axis == 0 { - 1 - } else { - dims[..axis].iter().product() - }; - let inner = if dims.is_empty() || axis + 1 >= dims.len() { - 1 - } else { - dims[axis + 1..].iter().product() - }; - let outer_stride = dim_size * inner; - - match tensor.dtype() { - DataType::Float32 => { - let input = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let values = values_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?; - let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - - let mut entries = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - entries.clear(); - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - entries.push((d, input[idx])); - } - - let compare = if largest { cmp_f32_desc } else { cmp_f32_asc }; - select_topk_entries(&mut entries, k, sorted, compare); - - // Output shape is (outer, k, inner); write row-major so a - // non-trailing reduction axis (inner > 1) lands correctly. - for j in 0..k { - let (index, value) = entries[j]; - let pos = o * k * inner + j * inner + r; - values[pos] = value; - indices[pos] = index as i64; - } - } - } - } - DataType::Float64 => { - let input = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let values = values_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?; - let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - - let mut entries = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - entries.clear(); - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - entries.push((d, input[idx])); - } - - let compare = if largest { cmp_f64_desc } else { cmp_f64_asc }; - select_topk_entries(&mut entries, k, sorted, compare); - - // Output shape is (outer, k, inner); write row-major so a - // non-trailing reduction axis (inner > 1) lands correctly. - for j in 0..k { - let (index, value) = entries[j]; - let pos = o * k * inner + j * inner + r; - values[pos] = value; - indices[pos] = index as i64; - } - } - } - } - DataType::Int32 => { - let input = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - let values = values_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i32 slice") - })?; - let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - - let mut entries = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - entries.clear(); - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - entries.push((d, input[idx])); - } - - let compare = if largest { cmp_i32_desc } else { cmp_i32_asc }; - select_topk_entries(&mut entries, k, sorted, compare); - - // Output shape is (outer, k, inner); write row-major so a - // non-trailing reduction axis (inner > 1) lands correctly. - for j in 0..k { - let (index, value) = entries[j]; - let pos = o * k * inner + j * inner + r; - values[pos] = value; - indices[pos] = index as i64; - } - } - } - } - DataType::Int64 => { - let input = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - let values = values_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - - let mut entries = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - entries.clear(); - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - entries.push((d, input[idx])); - } - - let compare = if largest { cmp_i64_desc } else { cmp_i64_asc }; - select_topk_entries(&mut entries, k, sorted, compare); - - // Output shape is (outer, k, inner); write row-major so a - // non-trailing reduction axis (inner > 1) lands correctly. - for j in 0..k { - let (index, value) = entries[j]; - let pos = o * k * inner + j * inner + r; - values[pos] = value; - indices[pos] = index as i64; - } - } - } - } - DataType::Bool => { - let input = tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - let values = values_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice") - })?; - let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - - let mut entries = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - entries.clear(); - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - entries.push((d, input[idx])); - } - - let compare = if largest { cmp_bool_desc } else { cmp_bool_asc }; - select_topk_entries(&mut entries, k, sorted, compare); - - // Output shape is (outer, k, inner); write row-major so a - // non-trailing reduction axis (inner > 1) lands correctly. - for j in 0..k { - let (index, value) = entries[j]; - let pos = o * k * inner + j * inner + r; - values[pos] = value; - indices[pos] = index as i64; - } - } - } - } - } - - let values = Tensor::new( - Arc::new(values_data), - values_shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - let indices = Tensor::new( - Arc::new(indices_data), - indices_shape, - DataType::Int64, - tensor.device(), - false, - ); - - // `values = gather(input, axis, indices)`; scatter the gradient back. - let values = attach_gather_like_grad(values, tensor, axis, &indices)?; - - Ok((values, indices)) -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::autograd::GatherBackward; +use crate::autograd::MinMaxBackward; +use crate::{ + autograd::add_to_graph, + error::{MinitensorError, Result}, + tensor::{DataType, Shape, Tensor, TensorData}, +}; +use rayon::prelude::*; +use std::cmp::Ordering; +use std::sync::Arc; + +pub(crate) fn any_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { + if dim >= tensor.ndim() { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + + let input_shape = tensor.shape().dims(); + let mut output_shape = input_shape.to_vec(); + if keepdim { + output_shape[dim] = 1; + } else { + output_shape.remove(dim); + } + let output_shape_obj = Shape::new(output_shape.clone()); + let mut result_data = + TensorData::zeros_on_device(output_shape_obj.numel(), DataType::Bool, tensor.device()); + + let dim_size = input_shape[dim]; + let _outer = input_shape[..dim].iter().product::(); + let inner = input_shape[dim + 1..].iter().product::(); + let outer_stride = dim_size * inner; + + match tensor.dtype() { + DataType::Float32 => { + let input = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let output = result_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice") + })?; + output.par_iter_mut().enumerate().for_each(|(idx, out)| { + let o = idx / inner; + let r = idx % inner; + let mut val = false; + for d in 0..dim_size { + let in_idx = o * outer_stride + d * inner + r; + if input[in_idx] != 0.0 { + val = true; + break; + } + } + *out = val; + }); + } + DataType::Float64 => { + let input = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let output = result_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice") + })?; + output.par_iter_mut().enumerate().for_each(|(idx, out)| { + let o = idx / inner; + let r = idx % inner; + let mut val = false; + for d in 0..dim_size { + let in_idx = o * outer_stride + d * inner + r; + if input[in_idx] != 0.0 { + val = true; + break; + } + } + *out = val; + }); + } + DataType::Int32 => { + let input = tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; + let output = result_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice") + })?; + output.par_iter_mut().enumerate().for_each(|(idx, out)| { + let o = idx / inner; + let r = idx % inner; + let mut val = false; + for d in 0..dim_size { + let in_idx = o * outer_stride + d * inner + r; + if input[in_idx] != 0 { + val = true; + break; + } + } + *out = val; + }); + } + DataType::Int64 => { + let input = tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; + let output = result_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice") + })?; + output.par_iter_mut().enumerate().for_each(|(idx, out)| { + let o = idx / inner; + let r = idx % inner; + let mut val = false; + for d in 0..dim_size { + let in_idx = o * outer_stride + d * inner + r; + if input[in_idx] != 0 { + val = true; + break; + } + } + *out = val; + }); + } + DataType::Bool => { + let input = tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + let output = result_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice") + })?; + output.par_iter_mut().enumerate().for_each(|(idx, out)| { + let o = idx / inner; + let r = idx % inner; + let mut val = false; + for d in 0..dim_size { + let in_idx = o * outer_stride + d * inner + r; + if input[in_idx] { + val = true; + break; + } + } + *out = val; + }); + } + } + + Ok(Tensor::new( + Arc::new(result_data), + output_shape_obj, + DataType::Bool, + tensor.device(), + false, + )) +} + +/// Maximum value along specified dimension +pub fn max(tensor: &Tensor, dim: Option, keepdim: bool) -> Result { + let (output, norm_dim) = match dim { + None => { + // Find global maximum + let result_shape = if keepdim { + Shape::new(vec![1; tensor.ndim()]) + } else { + Shape::scalar() + }; + + let mut result_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => max_all_f32(tensor, &mut result_data)?, + DataType::Float64 => max_all_f64(tensor, &mut result_data)?, + DataType::Int32 => max_all_i32(tensor, &mut result_data)?, + DataType::Int64 => max_all_i64(tensor, &mut result_data)?, + DataType::Bool => max_all_bool(tensor, &mut result_data)?, + } + + ( + Tensor::new( + Arc::new(result_data), + result_shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ), + None, + ) + } + Some(d) => { + let d = normalize_dim(d, tensor.ndim())?; + (max_along_dim(tensor, d, keepdim)?, Some(d)) + } + }; + attach_minmax_grad(output, tensor, norm_dim, keepdim, true, false) +} + +/// Attach a [`GatherBackward`] gradient to a value tensor that was formed by +/// gathering the input along `dim` at `indices` (`sort`/`topk`). The forward is +/// `values = gather(input, dim, indices)`, so the backward scatters the gradient +/// straight back to the selected source positions. +pub(crate) fn attach_gather_like_grad( + values: Tensor, + input: &Tensor, + dim: usize, + indices: &Tensor, +) -> Result { + if !input.requires_grad() || !input.dtype().is_float() { + return Ok(values); + } + let index = indices + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("selection indices must be int64"))? + .to_vec(); + let grad_fn = Arc::new(GatherBackward { + input_id: input.id(), + input_shape: input.shape().dims().to_vec(), + dim, + index, + }); + let mut values = values; + values.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&values, Some(grad_fn))?; + Ok(values) +} + +/// Attach a [`MinMaxBackward`] gradient to a `min`/`max`/`nanmax`/`nanmin` value +/// reduction (`nan_aware` selects the NaN-ignoring recompute in the backward). +fn attach_minmax_grad( + output: Tensor, + input: &Tensor, + dim: Option, + keepdim: bool, + is_max: bool, + nan_aware: bool, +) -> Result { + if !input.requires_grad() || !input.dtype().is_float() { + return Ok(output); + } + let grad_fn = Arc::new(MinMaxBackward { + input_id: input.id(), + input: input.detach(), + dim, + keepdim, + is_max, + nan_aware, + }); + let mut output = output; + output.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output, Some(grad_fn))?; + Ok(output) +} + +/// Minimum value along specified dimension +pub fn min(tensor: &Tensor, dim: Option, keepdim: bool) -> Result { + let (output, norm_dim) = match dim { + None => { + // Find global minimum + let result_shape = if keepdim { + Shape::new(vec![1; tensor.ndim()]) + } else { + Shape::scalar() + }; + + let mut result_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => min_all_f32(tensor, &mut result_data)?, + DataType::Float64 => min_all_f64(tensor, &mut result_data)?, + DataType::Int32 => min_all_i32(tensor, &mut result_data)?, + DataType::Int64 => min_all_i64(tensor, &mut result_data)?, + DataType::Bool => min_all_bool(tensor, &mut result_data)?, + } + + ( + Tensor::new( + Arc::new(result_data), + result_shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ), + None, + ) + } + Some(d) => { + let d = normalize_dim(d, tensor.ndim())?; + (min_along_dim(tensor, d, keepdim)?, Some(d)) + } + }; + attach_minmax_grad(output, tensor, norm_dim, keepdim, false, false) +} + +/// NaN-aware maximum value along specified dimension +pub fn nanmax(tensor: &Tensor, dim: Option, keepdim: bool) -> Result { + if !tensor.dtype().is_float() { + return max(tensor, dim, keepdim); + } + + let (output, norm_dim) = match dim { + None => { + let result_shape = if keepdim { + Shape::new(vec![1; tensor.ndim()]) + } else { + Shape::scalar() + }; + + let mut result_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => nanmax_all_f32(tensor, &mut result_data)?, + DataType::Float64 => nanmax_all_f64(tensor, &mut result_data)?, + _ => unreachable!("nanmax only supports floating point tensors"), + } + + ( + Tensor::new( + Arc::new(result_data), + result_shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ), + None, + ) + } + Some(d) => { + let d = normalize_dim(d, tensor.ndim())?; + let (values, _) = nanmax_along_dim_with_indices(tensor, d, keepdim)?; + (values, Some(d)) + } + }; + attach_minmax_grad(output, tensor, norm_dim, keepdim, true, true) +} + +/// NaN-aware minimum value along specified dimension +pub fn nanmin(tensor: &Tensor, dim: Option, keepdim: bool) -> Result { + if !tensor.dtype().is_float() { + return min(tensor, dim, keepdim); + } + + let (output, norm_dim) = match dim { + None => { + let result_shape = if keepdim { + Shape::new(vec![1; tensor.ndim()]) + } else { + Shape::scalar() + }; + + let mut result_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => nanmin_all_f32(tensor, &mut result_data)?, + DataType::Float64 => nanmin_all_f64(tensor, &mut result_data)?, + _ => unreachable!("nanmin only supports floating point tensors"), + } + + ( + Tensor::new( + Arc::new(result_data), + result_shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ), + None, + ) + } + Some(d) => { + let d = normalize_dim(d, tensor.ndim())?; + let (values, _) = nanmin_along_dim_with_indices(tensor, d, keepdim)?; + (values, Some(d)) + } + }; + attach_minmax_grad(output, tensor, norm_dim, keepdim, false, true) +} + +/// Maximum values and their indices along specified dimension +pub fn max_with_indices(tensor: &Tensor, dim: isize, keepdim: bool) -> Result<(Tensor, Tensor)> { + let d = normalize_dim(dim, tensor.ndim())?; + let (values, indices) = max_along_dim_with_indices(tensor, d, keepdim)?; + let values = attach_minmax_grad(values, tensor, Some(d), keepdim, true, false)?; + Ok((values, indices)) +} + +/// NaN-aware maximum values and their indices along specified dimension +pub fn nanmax_with_indices(tensor: &Tensor, dim: isize, keepdim: bool) -> Result<(Tensor, Tensor)> { + if !tensor.dtype().is_float() { + return max_with_indices(tensor, dim, keepdim); + } + + let d = normalize_dim(dim, tensor.ndim())?; + let (values, indices) = nanmax_along_dim_with_indices(tensor, d, keepdim)?; + let values = attach_minmax_grad(values, tensor, Some(d), keepdim, true, true)?; + Ok((values, indices)) +} + +/// Minimum values and their indices along specified dimension +pub fn min_with_indices(tensor: &Tensor, dim: isize, keepdim: bool) -> Result<(Tensor, Tensor)> { + let d = normalize_dim(dim, tensor.ndim())?; + let (values, indices) = min_along_dim_with_indices(tensor, d, keepdim)?; + let values = attach_minmax_grad(values, tensor, Some(d), keepdim, false, false)?; + Ok((values, indices)) +} + +/// NaN-aware minimum values and their indices along specified dimension +pub fn nanmin_with_indices(tensor: &Tensor, dim: isize, keepdim: bool) -> Result<(Tensor, Tensor)> { + if !tensor.dtype().is_float() { + return min_with_indices(tensor, dim, keepdim); + } + + let d = normalize_dim(dim, tensor.ndim())?; + let (values, indices) = nanmin_along_dim_with_indices(tensor, d, keepdim)?; + let values = attach_minmax_grad(values, tensor, Some(d), keepdim, false, true)?; + Ok((values, indices)) +} + +/// Argument of maximum value along specified dimension +pub fn argmax(tensor: &Tensor, dim: Option, keepdim: bool) -> Result { + match dim { + None => { + // Find global argmax + let result_shape = if keepdim { + Shape::new(vec![1; tensor.ndim()]) + } else { + Shape::scalar() + }; + + let mut result_data = TensorData::zeros_on_device(1, DataType::Int64, tensor.device()); + + match tensor.dtype() { + DataType::Float32 => argmax_all_f32(tensor, &mut result_data)?, + DataType::Float64 => argmax_all_f64(tensor, &mut result_data)?, + DataType::Int32 => argmax_all_i32(tensor, &mut result_data)?, + DataType::Int64 => argmax_all_i64(tensor, &mut result_data)?, + DataType::Bool => argmax_all_bool(tensor, &mut result_data)?, + } + + Ok(Tensor::new( + Arc::new(result_data), + result_shape, + DataType::Int64, + tensor.device(), + false, // argmax doesn't require gradients + )) + } + Some(d) => { + let d = normalize_dim(d, tensor.ndim())?; + argmax_along_dim(tensor, d, keepdim) + } + } +} + +/// Argument of minimum value along specified dimension +pub fn argmin(tensor: &Tensor, dim: Option, keepdim: bool) -> Result { + match dim { + None => { + // Find global argmin + let result_shape = if keepdim { + Shape::new(vec![1; tensor.ndim()]) + } else { + Shape::scalar() + }; + + let mut result_data = TensorData::zeros_on_device(1, DataType::Int64, tensor.device()); + + match tensor.dtype() { + DataType::Float32 => argmin_all_f32(tensor, &mut result_data)?, + DataType::Float64 => argmin_all_f64(tensor, &mut result_data)?, + DataType::Int32 => argmin_all_i32(tensor, &mut result_data)?, + DataType::Int64 => argmin_all_i64(tensor, &mut result_data)?, + DataType::Bool => argmin_all_bool(tensor, &mut result_data)?, + } + + Ok(Tensor::new( + Arc::new(result_data), + result_shape, + DataType::Int64, + tensor.device(), + false, // argmin doesn't require gradients + )) + } + Some(d) => { + let d = normalize_dim(d, tensor.ndim())?; + argmin_along_dim(tensor, d, keepdim) + } + } +} + +#[inline] +fn select_topk_entries( + entries: &mut [(usize, T)], + k: usize, + sorted: bool, + compare: fn(&(usize, T), &(usize, T)) -> Ordering, +) { + if k == 0 || entries.is_empty() { + return; + } + + if k < entries.len() { + entries.select_nth_unstable_by(k - 1, compare); + if sorted { + entries[..k].sort_by(compare); + } + } else if sorted { + entries.sort_by(compare); + } +} + +/// Return the top-``k`` values and their indices along ``dim`` +pub fn topk( + tensor: &Tensor, + k: usize, + dim: Option, + largest: bool, + sorted: bool, +) -> Result<(Tensor, Tensor)> { + let ndim = tensor.ndim(); + + let axis = if ndim == 0 { + match dim { + Some(d) if d == 0 || d == -1 => 0, + Some(d) => return Err(MinitensorError::index_error(d, 0, 1)), + None => 0, + } + } else { + let dim_value = dim.unwrap_or(-1); + normalize_dim(dim_value, ndim)? + }; + + let dims = tensor.shape().dims(); + let dim_size = if dims.is_empty() { 1 } else { dims[axis] }; + + if k > dim_size { + return Err(MinitensorError::invalid_argument(format!( + "selected index k out of range for dimension {axis} with size {dim_size}" + ))); + } + + let output_dims = if dims.is_empty() { + vec![k] + } else { + let mut dims_vec = dims.to_vec(); + dims_vec[axis] = k; + dims_vec + }; + + let values_shape = Shape::new(output_dims.clone()); + let indices_shape = Shape::new(output_dims); + + let num_out = values_shape.numel(); + let mut values_data = TensorData::zeros_on_device(num_out, tensor.dtype(), tensor.device()); + let mut indices_data = TensorData::zeros_on_device(num_out, DataType::Int64, tensor.device()); + + if k == 0 || num_out == 0 { + let values = Tensor::new( + Arc::new(values_data), + values_shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + let indices = Tensor::new( + Arc::new(indices_data), + indices_shape, + DataType::Int64, + tensor.device(), + false, + ); + return Ok((values, indices)); + } + + let outer = if dims.is_empty() || axis == 0 { + 1 + } else { + dims[..axis].iter().product() + }; + let inner = if dims.is_empty() || axis + 1 >= dims.len() { + 1 + } else { + dims[axis + 1..].iter().product() + }; + let outer_stride = dim_size * inner; + + match tensor.dtype() { + DataType::Float32 => { + let input = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let values = values_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?; + let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + + let mut entries = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + entries.clear(); + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + entries.push((d, input[idx])); + } + + let compare = if largest { cmp_f32_desc } else { cmp_f32_asc }; + select_topk_entries(&mut entries, k, sorted, compare); + + // Output shape is (outer, k, inner); write row-major so a + // non-trailing reduction axis (inner > 1) lands correctly. + for j in 0..k { + let (index, value) = entries[j]; + let pos = o * k * inner + j * inner + r; + values[pos] = value; + indices[pos] = index as i64; + } + } + } + } + DataType::Float64 => { + let input = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let values = values_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?; + let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + + let mut entries = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + entries.clear(); + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + entries.push((d, input[idx])); + } + + let compare = if largest { cmp_f64_desc } else { cmp_f64_asc }; + select_topk_entries(&mut entries, k, sorted, compare); + + // Output shape is (outer, k, inner); write row-major so a + // non-trailing reduction axis (inner > 1) lands correctly. + for j in 0..k { + let (index, value) = entries[j]; + let pos = o * k * inner + j * inner + r; + values[pos] = value; + indices[pos] = index as i64; + } + } + } + } + DataType::Int32 => { + let input = tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; + let values = values_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i32 slice") + })?; + let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + + let mut entries = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + entries.clear(); + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + entries.push((d, input[idx])); + } + + let compare = if largest { cmp_i32_desc } else { cmp_i32_asc }; + select_topk_entries(&mut entries, k, sorted, compare); + + // Output shape is (outer, k, inner); write row-major so a + // non-trailing reduction axis (inner > 1) lands correctly. + for j in 0..k { + let (index, value) = entries[j]; + let pos = o * k * inner + j * inner + r; + values[pos] = value; + indices[pos] = index as i64; + } + } + } + } + DataType::Int64 => { + let input = tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; + let values = values_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + + let mut entries = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + entries.clear(); + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + entries.push((d, input[idx])); + } + + let compare = if largest { cmp_i64_desc } else { cmp_i64_asc }; + select_topk_entries(&mut entries, k, sorted, compare); + + // Output shape is (outer, k, inner); write row-major so a + // non-trailing reduction axis (inner > 1) lands correctly. + for j in 0..k { + let (index, value) = entries[j]; + let pos = o * k * inner + j * inner + r; + values[pos] = value; + indices[pos] = index as i64; + } + } + } + } + DataType::Bool => { + let input = tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + let values = values_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice") + })?; + let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + + let mut entries = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + entries.clear(); + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + entries.push((d, input[idx])); + } + + let compare = if largest { cmp_bool_desc } else { cmp_bool_asc }; + select_topk_entries(&mut entries, k, sorted, compare); + + // Output shape is (outer, k, inner); write row-major so a + // non-trailing reduction axis (inner > 1) lands correctly. + for j in 0..k { + let (index, value) = entries[j]; + let pos = o * k * inner + j * inner + r; + values[pos] = value; + indices[pos] = index as i64; + } + } + } + } + } + + let values = Tensor::new( + Arc::new(values_data), + values_shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + let indices = Tensor::new( + Arc::new(indices_data), + indices_shape, + DataType::Int64, + tensor.device(), + false, + ); + + // `values = gather(input, axis, indices)`; scatter the gradient back. + let values = attach_gather_like_grad(values, tensor, axis, &indices)?; + + Ok((values, indices)) +} diff --git a/engine/src/operations/reduction/core.rs b/engine/src/operations/reduction/core.rs index 7e2a7e96..615f36df 100644 --- a/engine/src/operations/reduction/core.rs +++ b/engine/src/operations/reduction/core.rs @@ -1,1168 +1,1177 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -use crate::{ - autograd::{ - CumprodBackward, CumsumBackward, GatherBackward, MedianBackward, MinMaxBackward, - NanMeanBackward, NanSumBackward, ProdBackward, QuantileBackward, SumBackward, add_to_graph, - }, - error::{MinitensorError, Result}, - operations::{ - activation, arithmetic, shape_ops, - simd::{ - simd_prod_f32, simd_prod_f64, simd_prod_i32, simd_prod_i64, simd_sum_f32, simd_sum_f64, - simd_sum_i32, simd_sum_i64, - }, - }, - tensor::{DataType, Shape, Tensor, TensorData}, -}; -use rayon::prelude::*; -use std::cmp::Ordering; -use std::sync::Arc; - -const NANQUANTILE_ALL_NAN_ERR: &str = "nanquantile() encountered an all-NaN slice"; - -/// Interpolation modes supported by the quantile reduction. -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum QuantileInterpolation { - Linear, - Lower, - Higher, - Midpoint, - Nearest, -} - -impl QuantileInterpolation { - #[inline(always)] - fn interpolate(self, lower: f64, upper: f64, weight: f64) -> f64 { - match self { - QuantileInterpolation::Linear => lower + (upper - lower) * weight, - QuantileInterpolation::Lower => lower, - QuantileInterpolation::Higher => upper, - QuantileInterpolation::Midpoint => 0.5 * (lower + upper), - QuantileInterpolation::Nearest => { - // Index-aware callers should route through `nearest_index_with_tie_even` +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; + +use crate::{ + autograd::{MedianBackward, QuantileBackward, add_to_graph}, + error::{MinitensorError, Result}, + tensor::{DataType, Shape, Tensor, TensorData}, +}; +use rayon::prelude::*; +use std::cmp::Ordering; +use std::sync::Arc; + +pub(crate) const NANQUANTILE_ALL_NAN_ERR: &str = "nanquantile() encountered an all-NaN slice"; + +/// Interpolation modes supported by the quantile reduction. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum QuantileInterpolation { + Linear, + Lower, + Higher, + Midpoint, + Nearest, +} + +impl QuantileInterpolation { + #[inline(always)] + pub(crate) fn interpolate(self, lower: f64, upper: f64, weight: f64) -> f64 { + match self { + QuantileInterpolation::Linear => lower + (upper - lower) * weight, + QuantileInterpolation::Lower => lower, + QuantileInterpolation::Higher => upper, + QuantileInterpolation::Midpoint => 0.5 * (lower + upper), + QuantileInterpolation::Nearest => { + // Index-aware callers should route through `nearest_index_with_tie_even` // for tie-to-even behavior. This fallback remains - // deterministic for non-indexed interpolation use-cases. - if weight < 0.5 { - lower - } else { - upper - } - } - } - } -} - -fn normalize_reduction_dims(dims: Option>, ndim: usize) -> Result>> { - let ndim = ndim as isize; - Ok(match dims { - Some(dims) => { - let mut normalized = Vec::with_capacity(dims.len()); - for d in dims { - let d = if d < 0 { d + ndim } else { d }; - if d < 0 || d >= ndim { - return Err(MinitensorError::index_error(d, 0, ndim as usize)); - } - normalized.push(d as usize); - } - normalized.sort_unstable(); - normalized.dedup(); - Some(normalized) - } - None => None, - }) -} - -fn non_nan_mask(tensor: &Tensor) -> Result { - let numel = tensor.numel(); - let mut mask = vec![false; numel]; - - match tensor.dtype() { - DataType::Float32 => { - let data = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - mask.par_iter_mut() - .zip(data.par_iter()) - .for_each(|(out, &v)| { - *out = !v.is_nan(); - }); - } - DataType::Float64 => { - let data = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - mask.par_iter_mut() - .zip(data.par_iter()) - .for_each(|(out, &v)| { - *out = !v.is_nan(); - }); - } - _ => { - return Err(MinitensorError::invalid_operation( - "nan reductions are only supported for floating point tensors", - )); - } - } - - Ok(Tensor::new( - Arc::new(TensorData::from_vec_bool(mask, tensor.device())), - tensor.shape().clone(), - DataType::Bool, - tensor.device(), - false, - )) -} - -fn cmp_f32_desc(a: &(usize, f32), b: &(usize, f32)) -> Ordering { - match (a.1.is_nan(), b.1.is_nan()) { - (true, true) => a.0.cmp(&b.0), - (true, false) => Ordering::Less, - (false, true) => Ordering::Greater, - (false, false) => match b.1.partial_cmp(&a.1).unwrap_or(Ordering::Equal) { - Ordering::Equal => a.0.cmp(&b.0), - order => order, - }, - } -} - -fn cmp_f32_asc(a: &(usize, f32), b: &(usize, f32)) -> Ordering { - match (a.1.is_nan(), b.1.is_nan()) { - (true, true) => a.0.cmp(&b.0), - (true, false) => Ordering::Greater, - (false, true) => Ordering::Less, - (false, false) => match a.1.partial_cmp(&b.1).unwrap_or(Ordering::Equal) { - Ordering::Equal => a.0.cmp(&b.0), - order => order, - }, - } -} - -fn cmp_f64_desc(a: &(usize, f64), b: &(usize, f64)) -> Ordering { - match (a.1.is_nan(), b.1.is_nan()) { - (true, true) => a.0.cmp(&b.0), - (true, false) => Ordering::Less, - (false, true) => Ordering::Greater, - (false, false) => match b.1.partial_cmp(&a.1).unwrap_or(Ordering::Equal) { - Ordering::Equal => a.0.cmp(&b.0), - order => order, - }, - } -} - -fn cmp_f64_asc(a: &(usize, f64), b: &(usize, f64)) -> Ordering { - match (a.1.is_nan(), b.1.is_nan()) { - (true, true) => a.0.cmp(&b.0), - (true, false) => Ordering::Greater, - (false, true) => Ordering::Less, - (false, false) => match a.1.partial_cmp(&b.1).unwrap_or(Ordering::Equal) { - Ordering::Equal => a.0.cmp(&b.0), - order => order, - }, - } -} - -fn cmp_i32_desc(a: &(usize, i32), b: &(usize, i32)) -> Ordering { - match b.1.cmp(&a.1) { - Ordering::Equal => a.0.cmp(&b.0), - order => order, - } -} - -fn cmp_i32_asc(a: &(usize, i32), b: &(usize, i32)) -> Ordering { - match a.1.cmp(&b.1) { - Ordering::Equal => a.0.cmp(&b.0), - order => order, - } -} - -fn cmp_i64_desc(a: &(usize, i64), b: &(usize, i64)) -> Ordering { - match b.1.cmp(&a.1) { - Ordering::Equal => a.0.cmp(&b.0), - order => order, - } -} - -fn cmp_i64_asc(a: &(usize, i64), b: &(usize, i64)) -> Ordering { - match a.1.cmp(&b.1) { - Ordering::Equal => a.0.cmp(&b.0), - order => order, - } -} - -fn cmp_bool_desc(a: &(usize, bool), b: &(usize, bool)) -> Ordering { - match (a.1, b.1) { - (true, true) | (false, false) => a.0.cmp(&b.0), - (true, false) => Ordering::Less, - (false, true) => Ordering::Greater, - } -} - -fn cmp_bool_asc(a: &(usize, bool), b: &(usize, bool)) -> Ordering { - match (a.1, b.1) { - (true, true) | (false, false) => a.0.cmp(&b.0), - (true, false) => Ordering::Greater, - (false, true) => Ordering::Less, - } -} - -fn ensure_non_empty(numel: usize) -> Result<()> { - if numel == 0 { - Err(MinitensorError::invalid_argument( - "median() does not support empty tensors".to_string(), - )) - } else { - Ok(()) - } -} - -pub fn median( - tensor: &Tensor, - dim: Option, - keepdim: bool, -) -> Result<(Tensor, Option)> { - ensure_non_empty(tensor.numel())?; - - if tensor.ndim() == 0 { - return Ok((tensor.clone(), None)); - } - - let (values, indices, norm_dim) = match dim { - None => { - let (values, indices) = median_all(tensor)?; - (values, indices, None) - } - Some(dim_value) => { - let axis = if tensor.ndim() == 0 { - if dim_value == 0 || dim_value == -1 { - 0 - } else { - return Err(MinitensorError::index_error(dim_value, 0, 1)); - } - } else { - normalize_dim(dim_value, tensor.ndim())? - }; - let (values, indices) = median_along_dim(tensor, axis, keepdim)?; - (values, Some(indices), Some(axis)) - } - }; - let values = attach_median_grad(values, tensor, norm_dim, keepdim, false)?; - Ok((values, indices)) -} - -/// Attach a [`MedianBackward`] gradient to a median value reduction. -fn attach_median_grad( - values: Tensor, - input: &Tensor, - dim: Option, - keepdim: bool, - nan_aware: bool, -) -> Result { - if !input.requires_grad() || !input.dtype().is_float() { - return Ok(values); - } - let grad_fn = Arc::new(MedianBackward { - input_id: input.id(), - input: input.detach(), - dim, - keepdim, - nan_aware, - }); - let mut values = values; - values.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&values, Some(grad_fn))?; - Ok(values) -} - -/// Compute the q-th quantile of the tensor data. -pub fn quantile( - tensor: &Tensor, - q: f64, - dim: Option, - keepdim: bool, - interpolation: QuantileInterpolation, -) -> Result { - if tensor.numel() == 0 { - return Err(MinitensorError::invalid_argument( - "quantile() does not support empty tensors".to_string(), - )); - } - - validate_quantile_value(q)?; - ensure_floating_point_dtype(tensor.dtype())?; - - let (output, norm_dim) = match dim { - None => (quantile_all(tensor, q, keepdim, interpolation)?, None), - Some(dim_value) => { - if tensor.ndim() == 0 { - if dim_value == 0 || dim_value == -1 { - (quantile_all(tensor, q, keepdim, interpolation)?, None) - } else { - return Err(MinitensorError::index_error(dim_value, 0, 1)); - } - } else { - let axis = normalize_dim(dim_value, tensor.ndim())?; - ( - quantile_along_dim(tensor, axis, keepdim, q, interpolation)?, - Some(axis), - ) - } - } - }; - - attach_quantile_grad(output, tensor, norm_dim, q, interpolation, false) -} - -/// Attach a [`QuantileBackward`] gradient to a quantile value reduction. -fn attach_quantile_grad( - output: Tensor, - input: &Tensor, - dim: Option, - q: f64, - interpolation: QuantileInterpolation, - nan_aware: bool, -) -> Result { - if !input.requires_grad() || !input.dtype().is_float() { - return Ok(output); - } - let grad_fn = Arc::new(QuantileBackward { - input_id: input.id(), - input: input.detach(), - dim, - q, - interpolation, - nan_aware, - }); - let mut output = output; - output.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output, Some(grad_fn))?; - Ok(output) -} - -/// Compute multiple quantiles of the tensor data in a single pass. -pub fn quantiles( - tensor: &Tensor, - qs: &[f64], - dim: Option, - keepdim: bool, - interpolation: QuantileInterpolation, -) -> Result { - if tensor.numel() == 0 { - return Err(MinitensorError::invalid_argument( - "quantile() does not support empty tensors".to_string(), - )); - } - - if qs.is_empty() { - return Err(MinitensorError::invalid_argument( - "quantile() expected at least one probability value".to_string(), - )); - } - - for &q in qs { - validate_quantile_value(q)?; - } - - ensure_floating_point_dtype(tensor.dtype())?; - - match dim { - None => quantiles_all(tensor, qs, keepdim, interpolation), - Some(dim_value) => { - if tensor.ndim() == 0 { - if dim_value == 0 || dim_value == -1 { - return quantiles_all(tensor, qs, keepdim, interpolation); - } - return Err(MinitensorError::index_error(dim_value, 0, 1)); - } - - let axis = normalize_dim(dim_value, tensor.ndim())?; - quantiles_along_dim(tensor, axis, qs, keepdim, interpolation) - } - } -} - -/// Compute the q-th quantile of the tensor data while ignoring NaN values. -pub fn nanquantile( - tensor: &Tensor, - q: f64, - dim: Option, - keepdim: bool, - interpolation: QuantileInterpolation, -) -> Result { - if tensor.numel() == 0 { - return Err(MinitensorError::invalid_argument( - "nanquantile() does not support empty tensors".to_string(), - )); - } - - validate_quantile_value(q)?; - ensure_floating_point_dtype(tensor.dtype())?; - - let (output, norm_dim) = match dim { - None => (nanquantile_all(tensor, q, keepdim, interpolation)?, None), - Some(dim_value) => { - if tensor.ndim() == 0 { - if dim_value == 0 || dim_value == -1 { - (nanquantile_all(tensor, q, keepdim, interpolation)?, None) - } else { - return Err(MinitensorError::index_error(dim_value, 0, 1)); - } - } else { - let axis = normalize_dim(dim_value, tensor.ndim())?; - ( - nanquantile_along_dim(tensor, axis, keepdim, q, interpolation)?, - Some(axis), - ) - } - } - }; - attach_quantile_grad(output, tensor, norm_dim, q, interpolation, true) -} - -/// Compute the median while ignoring NaN values. -pub fn nanmedian(tensor: &Tensor, dim: Option, keepdim: bool) -> Result { - ensure_floating_point_dtype_for(tensor.dtype(), "nanmedian")?; - - let (values, norm_dim) = match dim { - None => (nanmedian_all(tensor, keepdim)?, None), - Some(dim_value) => { - if tensor.ndim() == 0 { - if dim_value == 0 || dim_value == -1 { - (nanmedian_all(tensor, keepdim)?, None) - } else { - return Err(MinitensorError::index_error(dim_value, 0, 1)); - } - } else { - let axis = normalize_dim(dim_value, tensor.ndim())?; - (nanmedian_along_dim(tensor, axis, keepdim)?, Some(axis)) - } - } - }; - attach_median_grad(values, tensor, norm_dim, keepdim, true) -} - -/// Compute multiple quantiles of the tensor data in a single pass while ignoring NaN values. -pub fn nanquantiles( - tensor: &Tensor, - qs: &[f64], - dim: Option, - keepdim: bool, - interpolation: QuantileInterpolation, -) -> Result { - if tensor.numel() == 0 { - return Err(MinitensorError::invalid_argument( - "nanquantile() does not support empty tensors".to_string(), - )); - } - - if qs.is_empty() { - return Err(MinitensorError::invalid_argument( - "nanquantile() expected at least one probability value".to_string(), - )); - } - - for &q in qs { - validate_quantile_value(q)?; - } - - ensure_floating_point_dtype(tensor.dtype())?; - - match dim { - None => nanquantiles_all(tensor, qs, keepdim, interpolation), - Some(dim_value) => { - if tensor.ndim() == 0 { - if dim_value == 0 || dim_value == -1 { - return nanquantiles_all(tensor, qs, keepdim, interpolation); - } - return Err(MinitensorError::index_error(dim_value, 0, 1)); - } - - let axis = normalize_dim(dim_value, tensor.ndim())?; - nanquantiles_along_dim(tensor, axis, qs, keepdim, interpolation) - } - } -} - -fn validate_quantile_value(q: f64) -> Result<()> { - if !q.is_finite() { - return Err(MinitensorError::invalid_argument( - "quantile() requires a finite probability in [0, 1]".to_string(), - )); - } - if !(0.0..=1.0).contains(&q) { - return Err(MinitensorError::invalid_argument(format!( - "quantile() expected q in [0, 1], got {q}", - ))); - } - Ok(()) -} - -fn ensure_floating_point_dtype(dtype: DataType) -> Result<()> { - ensure_floating_point_dtype_for(dtype, "quantile") -} - -fn ensure_floating_point_dtype_for(dtype: DataType, operation: &str) -> Result<()> { - match dtype { - DataType::Float32 | DataType::Float64 => Ok(()), - _ => Err(MinitensorError::invalid_operation(format!( - "{operation}() currently supports only floating point tensors" - ))), - } -} - -fn fill_quantile_single_f32( - input: &[f32], - values: &mut [f32], - outer: usize, - inner: usize, - outer_stride: usize, -) { - for o in 0..outer { - for r in 0..inner { - let idx = o * outer_stride + r; - let value = input[idx]; - values[o * inner + r] = if value.is_nan() { f32::NAN } else { value }; - } - } -} - -fn fill_quantile_single_f64( - input: &[f64], - values: &mut [f64], - outer: usize, - inner: usize, - outer_stride: usize, -) { - for o in 0..outer { - for r in 0..inner { - let idx = o * outer_stride + r; - let value = input[idx]; - values[o * inner + r] = if value.is_nan() { f64::NAN } else { value }; - } - } -} - -fn fill_quantiles_single_f32( - input: &[f32], - values: &mut [f32], - outer: usize, - inner: usize, - outer_stride: usize, - q_len: usize, -) { - for o in 0..outer { - for r in 0..inner { - let idx = o * outer_stride + r; - let value = input[idx]; - for qi in 0..q_len { - let out_idx = ((qi * outer) + o) * inner + r; - values[out_idx] = if value.is_nan() { f32::NAN } else { value }; - } - } - } -} - -fn fill_quantiles_single_f64( - input: &[f64], - values: &mut [f64], - outer: usize, - inner: usize, - outer_stride: usize, - q_len: usize, -) { - for o in 0..outer { - for r in 0..inner { - let idx = o * outer_stride + r; - let value = input[idx]; - for qi in 0..q_len { - let out_idx = ((qi * outer) + o) * inner + r; - values[out_idx] = if value.is_nan() { f64::NAN } else { value }; - } - } - } -} - -fn fill_nanquantile_single_f32( - input: &[f32], - values: &mut [f32], - outer: usize, - inner: usize, - outer_stride: usize, -) -> Result<()> { - for o in 0..outer { - for r in 0..inner { - let idx = o * outer_stride + r; - let value = input[idx]; - if value.is_nan() { - return Err(MinitensorError::invalid_argument( - NANQUANTILE_ALL_NAN_ERR.to_string(), - )); - } - values[o * inner + r] = value; - } - } - Ok(()) -} - -fn fill_nanquantile_single_f64( - input: &[f64], - values: &mut [f64], - outer: usize, - inner: usize, - outer_stride: usize, -) -> Result<()> { - for o in 0..outer { - for r in 0..inner { - let idx = o * outer_stride + r; - let value = input[idx]; - if value.is_nan() { - return Err(MinitensorError::invalid_argument( - NANQUANTILE_ALL_NAN_ERR.to_string(), - )); - } - values[o * inner + r] = value; - } - } - Ok(()) -} - -fn fill_nanquantiles_single_f32( - input: &[f32], - values: &mut [f32], - outer: usize, - inner: usize, - outer_stride: usize, - q_len: usize, -) -> Result<()> { - for o in 0..outer { - for r in 0..inner { - let idx = o * outer_stride + r; - let value = input[idx]; - if value.is_nan() { - return Err(MinitensorError::invalid_argument( - NANQUANTILE_ALL_NAN_ERR.to_string(), - )); - } - for qi in 0..q_len { - let out_idx = ((qi * outer) + o) * inner + r; - values[out_idx] = value; - } - } - } - Ok(()) -} - -fn fill_nanquantiles_single_f64( - input: &[f64], - values: &mut [f64], - outer: usize, - inner: usize, - outer_stride: usize, - q_len: usize, -) -> Result<()> { - for o in 0..outer { - for r in 0..inner { - let idx = o * outer_stride + r; - let value = input[idx]; - if value.is_nan() { - return Err(MinitensorError::invalid_argument( - NANQUANTILE_ALL_NAN_ERR.to_string(), - )); - } - for qi in 0..q_len { - let out_idx = ((qi * outer) + o) * inner + r; - values[out_idx] = value; - } - } - } - Ok(()) -} - -fn fill_quantiles_all_single_f32(value: f32, values: &mut [f32]) { - if value.is_nan() { - values.fill(f32::NAN); - } else { - values.fill(value); - } -} - -fn fill_quantiles_all_single_f64(value: f64, values: &mut [f64]) { - if value.is_nan() { - values.fill(f64::NAN); - } else { - values.fill(value); - } -} - -fn fill_nanquantiles_all_single_f32(value: f32, values: &mut [f32]) -> Result<()> { - if value.is_nan() { - return Err(MinitensorError::invalid_argument( - NANQUANTILE_ALL_NAN_ERR.to_string(), - )); - } - values.fill(value); - Ok(()) -} - -fn fill_nanquantiles_all_single_f64(value: f64, values: &mut [f64]) -> Result<()> { - if value.is_nan() { - return Err(MinitensorError::invalid_argument( - NANQUANTILE_ALL_NAN_ERR.to_string(), - )); - } - values.fill(value); - Ok(()) -} - -#[derive(Clone, Copy)] -struct QuantilePosition { - lower_idx: usize, - upper_idx: usize, - nearest_idx: usize, - weight: f64, -} - -fn quantile_positions_for_len(len: usize, qs: &[f64]) -> Vec { - let mut positions = Vec::with_capacity(qs.len()); - for &q in qs { - positions.push(quantile_position_for_len_q(len, q)); - } - positions -} - -#[inline(always)] -fn quantile_position_for_len_q(len: usize, q: f64) -> QuantilePosition { - if len <= 1 { - return QuantilePosition { - lower_idx: 0, - upper_idx: 0, - nearest_idx: 0, - weight: 0.0, - }; - } - - let max_index = (len - 1) as f64; - let pos = (q * max_index).clamp(0.0, max_index); - let lower_idx = pos.floor() as usize; - let upper_idx = pos.ceil() as usize; - let weight = (pos - lower_idx as f64).clamp(0.0, 1.0); - let nearest_idx = nearest_index_with_tie_even(lower_idx, upper_idx, weight); - - QuantilePosition { - lower_idx, - upper_idx, - nearest_idx, - weight, - } -} - -fn quantiles_from_sorted_f32( - values: &[f32], - positions: &[QuantilePosition], - interpolation: QuantileInterpolation, - output: &mut [f32], -) { - for (slot, position) in output.iter_mut().zip(positions.iter()) { - *slot = quantile_from_sorted_position_f32(values, position, interpolation); - } -} - -fn quantiles_from_sorted_f64( - values: &[f64], - positions: &[QuantilePosition], - interpolation: QuantileInterpolation, - output: &mut [f64], -) { - for (slot, position) in output.iter_mut().zip(positions.iter()) { - *slot = quantile_from_sorted_position_f64(values, position, interpolation); - } -} - -#[inline(always)] -fn nearest_index_with_tie_even(lower_idx: usize, upper_idx: usize, weight: f64) -> usize { - debug_assert!((0.0..=1.0).contains(&weight)); - - if upper_idx <= lower_idx { - debug_assert_eq!(upper_idx, lower_idx, "nearest_index_with_tie_even requires upper_idx >= lower_idx"); - return lower_idx; - } - - debug_assert_eq!(upper_idx, lower_idx + 1); - - if weight < 0.5 { - lower_idx - } else if weight > 0.5 { - upper_idx - } else { - lower_idx + (lower_idx & 1) - } -} - -fn quantile_from_sorted_position_f32( - values: &[f32], - position: &QuantilePosition, - interpolation: QuantileInterpolation, -) -> f32 { - match interpolation { - QuantileInterpolation::Lower => values[position.lower_idx], - QuantileInterpolation::Higher => values[position.upper_idx], - QuantileInterpolation::Nearest => values[position.nearest_idx], - QuantileInterpolation::Linear | QuantileInterpolation::Midpoint => { - let lower = values[position.lower_idx] as f64; - let upper = values[position.upper_idx] as f64; - interpolation.interpolate(lower, upper, position.weight) as f32 - } - } -} - -fn quantile_from_sorted_position_f64( - values: &[f64], - position: &QuantilePosition, - interpolation: QuantileInterpolation, -) -> f64 { - match interpolation { - QuantileInterpolation::Lower => values[position.lower_idx], - QuantileInterpolation::Higher => values[position.upper_idx], - QuantileInterpolation::Nearest => values[position.nearest_idx], - QuantileInterpolation::Linear | QuantileInterpolation::Midpoint => { - let lower = values[position.lower_idx]; - let upper = values[position.upper_idx]; - interpolation.interpolate(lower, upper, position.weight) - } - } -} - -fn quantile_all( - tensor: &Tensor, - q: f64, - keepdim: bool, - interpolation: QuantileInterpolation, -) -> Result { - if tensor.ndim() == 0 { - return Ok(tensor.clone()); - } - - let mut result_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => { - let data = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let mut values = Vec::with_capacity(data.len()); - for &value in data { - if value.is_nan() { - result_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?[0] = f32::NAN; - let result_shape = if keepdim { - Shape::new(vec![1; tensor.ndim()]) - } else { - Shape::scalar() - }; - return Ok(Tensor::new( - Arc::new(result_data), - result_shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )); - } - values.push(value); - } - let quant = quantile_from_unsorted_f32(&mut values, q, interpolation); - result_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?[0] = quant; - } - DataType::Float64 => { - let data = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let mut values = Vec::with_capacity(data.len()); - for &value in data { - if value.is_nan() { - result_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?[0] = f64::NAN; - let result_shape = if keepdim { - Shape::new(vec![1; tensor.ndim()]) - } else { - Shape::scalar() - }; - return Ok(Tensor::new( - Arc::new(result_data), - result_shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )); - } - values.push(value); - } - let quant = quantile_from_unsorted_f64(&mut values, q, interpolation); - result_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?[0] = quant; - } - _ => unreachable!("dtype validated"), - } - - let result_shape = if keepdim { - Shape::new(vec![1; tensor.ndim()]) - } else { - Shape::scalar() - }; - - Ok(Tensor::new( - Arc::new(result_data), - result_shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )) -} - -fn nanquantile_all( - tensor: &Tensor, - q: f64, - keepdim: bool, - interpolation: QuantileInterpolation, -) -> Result { - if tensor.ndim() == 0 { - match tensor.dtype() { - DataType::Float32 => { - let value = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?[0]; - if value.is_nan() { - return Err(MinitensorError::invalid_argument( - NANQUANTILE_ALL_NAN_ERR.to_string(), - )); - } - } - DataType::Float64 => { - let value = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?[0]; - if value.is_nan() { - return Err(MinitensorError::invalid_argument( - NANQUANTILE_ALL_NAN_ERR.to_string(), - )); - } - } - _ => unreachable!("dtype validated"), - } - return Ok(tensor.clone()); - } - - let mut result_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => { - let data = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let mut values: Vec = data.iter().copied().filter(|v| !v.is_nan()).collect(); - if values.is_empty() { - return Err(MinitensorError::invalid_argument( - NANQUANTILE_ALL_NAN_ERR.to_string(), - )); - } - let quant = quantile_from_unsorted_f32(&mut values, q, interpolation); - result_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?[0] = quant; - } - DataType::Float64 => { - let data = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let mut values: Vec = data.iter().copied().filter(|v| !v.is_nan()).collect(); - if values.is_empty() { - return Err(MinitensorError::invalid_argument( - NANQUANTILE_ALL_NAN_ERR.to_string(), - )); - } - let quant = quantile_from_unsorted_f64(&mut values, q, interpolation); - result_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?[0] = quant; - } - _ => unreachable!("dtype validated"), - } - - let result_shape = if keepdim { - Shape::new(vec![1; tensor.ndim()]) - } else { - Shape::scalar() - }; - - Ok(Tensor::new( - Arc::new(result_data), - result_shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )) -} - -#[cfg(test)] -mod core_tests { - use super::*; - use crate::Device; - - #[test] - fn test_quantile_interpolation_modes() { - let lower = 1.0; - let upper = 3.0; - let weight = 0.25; - - assert_eq!( - QuantileInterpolation::Linear.interpolate(lower, upper, weight), - 1.5 - ); - assert_eq!( - QuantileInterpolation::Lower.interpolate(lower, upper, weight), - lower - ); - assert_eq!( - QuantileInterpolation::Higher.interpolate(lower, upper, weight), - upper - ); - assert_eq!( - QuantileInterpolation::Midpoint.interpolate(lower, upper, weight), - 2.0 - ); - assert_eq!( - QuantileInterpolation::Nearest.interpolate(lower, upper, 0.49), - lower - ); - assert_eq!( - QuantileInterpolation::Nearest.interpolate(lower, upper, 0.5), - upper - ); - } - - #[test] - fn test_normalize_reduction_dims_sorts_dedups_and_supports_negative_dims() { - let dims = Some(vec![2, -1, 0, 2, -3]); - let normalized = normalize_reduction_dims(dims, 3).unwrap(); - assert_eq!(normalized, Some(vec![0, 2])); - assert_eq!(normalize_reduction_dims(None, 3).unwrap(), None); - } - - #[test] - fn test_normalize_reduction_dims_rejects_out_of_range() { - assert!(normalize_reduction_dims(Some(vec![3]), 3).is_err()); - assert!(normalize_reduction_dims(Some(vec![-4]), 3).is_err()); - } - - #[test] - fn test_non_nan_mask_for_float_tensors_and_invalid_dtype() { - let f32_tensor = Tensor::new( - Arc::new(TensorData::from_vec_f32( - vec![1.0, f32::NAN, -2.0], - Device::cpu(), - )), - Shape::new(vec![3]), - DataType::Float32, - Device::cpu(), - false, - ); - let mask = non_nan_mask(&f32_tensor).unwrap(); - assert_eq!(mask.dtype(), DataType::Bool); - assert_eq!(mask.data().as_bool_slice().unwrap(), &[true, false, true],); - - let f64_tensor = Tensor::new( - Arc::new(TensorData::from_vec_f64( - vec![f64::NAN, 4.0, 0.0], - Device::cpu(), - )), - Shape::new(vec![3]), - DataType::Float64, - Device::cpu(), - false, - ); - let f64_mask = non_nan_mask(&f64_tensor).unwrap(); - assert_eq!(f64_mask.data().as_bool_slice().unwrap(), &[false, true, true],); - - let bool_tensor = Tensor::new( - Arc::new(TensorData::from_vec_bool(vec![true, false], Device::cpu())), - Shape::new(vec![2]), - DataType::Bool, - Device::cpu(), - false, - ); - assert!(non_nan_mask(&bool_tensor).is_err()); - } - - #[test] - fn test_comparator_tie_breaking_and_nan_ordering() { - let mut f32_desc = vec![(1, 2.0f32), (0, 2.0), (2, f32::NAN)]; - f32_desc.sort_by(cmp_f32_desc); - assert!(f32_desc[0].1.is_nan()); - assert_eq!(f32_desc[0].0, 2); - assert_eq!(f32_desc[1..].iter().map(|(i, _)| *i).collect::>(), vec![0, 1]); - - let mut f32_asc = vec![(1, 2.0f32), (0, 2.0), (2, f32::NAN)]; - f32_asc.sort_by(cmp_f32_asc); - assert!(f32_asc[2].1.is_nan()); - assert_eq!(f32_asc[2].0, 2); - assert_eq!(f32_asc[..2].iter().map(|(i, _)| *i).collect::>(), vec![0, 1]); - - let mut f64_desc = vec![(1, 2.0f64), (0, 2.0), (2, f64::NAN)]; - f64_desc.sort_by(cmp_f64_desc); - assert!(f64_desc[0].1.is_nan()); - assert_eq!(f64_desc[0].0, 2); - assert_eq!(f64_desc[1..].iter().map(|(i, _)| *i).collect::>(), vec![0, 1]); - - let mut f64_asc = vec![(1, 2.0f64), (0, 2.0), (2, f64::NAN)]; - f64_asc.sort_by(cmp_f64_asc); - assert!(f64_asc[2].1.is_nan()); - assert_eq!(f64_asc[2].0, 2); - assert_eq!(f64_asc[..2].iter().map(|(i, _)| *i).collect::>(), vec![0, 1]); - - let mut i32_desc = vec![(1, 7_i32), (0, 7_i32), (2, 1_i32)]; - i32_desc.sort_by(cmp_i32_desc); - assert_eq!(i32_desc, vec![(0, 7), (1, 7), (2, 1)]); - - let mut i32_asc = vec![(1, 7_i32), (0, 7_i32), (2, 1_i32)]; - i32_asc.sort_by(cmp_i32_asc); - assert_eq!(i32_asc, vec![(2, 1), (0, 7), (1, 7)]); - - let mut i64_desc = vec![(1, 7_i64), (0, 7_i64), (2, 1_i64)]; - i64_desc.sort_by(cmp_i64_desc); - assert_eq!(i64_desc, vec![(0, 7), (1, 7), (2, 1)]); - - let mut i64_asc = vec![(1, 7_i64), (0, 7_i64), (2, 1_i64)]; - i64_asc.sort_by(cmp_i64_asc); - assert_eq!(i64_asc, vec![(2, 1), (0, 7), (1, 7)]); - - let mut bool_desc = vec![(1, true), (0, true), (2, false)]; - bool_desc.sort_by(cmp_bool_desc); - assert_eq!(bool_desc, vec![(0, true), (1, true), (2, false)]); - - let mut bool_asc = vec![(1, true), (0, true), (2, false)]; - bool_asc.sort_by(cmp_bool_asc); - assert_eq!(bool_asc, vec![(2, false), (0, true), (1, true)]); - } - - #[test] - fn test_ensure_non_empty_guard() { - assert!(ensure_non_empty(0).is_err()); - assert!(ensure_non_empty(1).is_ok()); - } -} + // deterministic for non-indexed interpolation use-cases. + if weight < 0.5 { lower } else { upper } + } + } + } +} + +pub(crate) fn normalize_reduction_dims( + dims: Option>, + ndim: usize, +) -> Result>> { + let ndim = ndim as isize; + Ok(match dims { + Some(dims) => { + let mut normalized = Vec::with_capacity(dims.len()); + for d in dims { + let d = if d < 0 { d + ndim } else { d }; + if d < 0 || d >= ndim { + return Err(MinitensorError::index_error(d, 0, ndim as usize)); + } + normalized.push(d as usize); + } + normalized.sort_unstable(); + normalized.dedup(); + Some(normalized) + } + None => None, + }) +} + +pub(crate) fn non_nan_mask(tensor: &Tensor) -> Result { + let numel = tensor.numel(); + let mut mask = vec![false; numel]; + + match tensor.dtype() { + DataType::Float32 => { + let data = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + mask.par_iter_mut() + .zip(data.par_iter()) + .for_each(|(out, &v)| { + *out = !v.is_nan(); + }); + } + DataType::Float64 => { + let data = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + mask.par_iter_mut() + .zip(data.par_iter()) + .for_each(|(out, &v)| { + *out = !v.is_nan(); + }); + } + _ => { + return Err(MinitensorError::invalid_operation( + "nan reductions are only supported for floating point tensors", + )); + } + } + + Ok(Tensor::new( + Arc::new(TensorData::from_vec_bool(mask, tensor.device())), + tensor.shape().clone(), + DataType::Bool, + tensor.device(), + false, + )) +} + +pub(crate) fn cmp_f32_desc(a: &(usize, f32), b: &(usize, f32)) -> Ordering { + match (a.1.is_nan(), b.1.is_nan()) { + (true, true) => a.0.cmp(&b.0), + (true, false) => Ordering::Less, + (false, true) => Ordering::Greater, + (false, false) => match b.1.partial_cmp(&a.1).unwrap_or(Ordering::Equal) { + Ordering::Equal => a.0.cmp(&b.0), + order => order, + }, + } +} + +pub(crate) fn cmp_f32_asc(a: &(usize, f32), b: &(usize, f32)) -> Ordering { + match (a.1.is_nan(), b.1.is_nan()) { + (true, true) => a.0.cmp(&b.0), + (true, false) => Ordering::Greater, + (false, true) => Ordering::Less, + (false, false) => match a.1.partial_cmp(&b.1).unwrap_or(Ordering::Equal) { + Ordering::Equal => a.0.cmp(&b.0), + order => order, + }, + } +} + +pub(crate) fn cmp_f64_desc(a: &(usize, f64), b: &(usize, f64)) -> Ordering { + match (a.1.is_nan(), b.1.is_nan()) { + (true, true) => a.0.cmp(&b.0), + (true, false) => Ordering::Less, + (false, true) => Ordering::Greater, + (false, false) => match b.1.partial_cmp(&a.1).unwrap_or(Ordering::Equal) { + Ordering::Equal => a.0.cmp(&b.0), + order => order, + }, + } +} + +pub(crate) fn cmp_f64_asc(a: &(usize, f64), b: &(usize, f64)) -> Ordering { + match (a.1.is_nan(), b.1.is_nan()) { + (true, true) => a.0.cmp(&b.0), + (true, false) => Ordering::Greater, + (false, true) => Ordering::Less, + (false, false) => match a.1.partial_cmp(&b.1).unwrap_or(Ordering::Equal) { + Ordering::Equal => a.0.cmp(&b.0), + order => order, + }, + } +} + +pub(crate) fn cmp_i32_desc(a: &(usize, i32), b: &(usize, i32)) -> Ordering { + match b.1.cmp(&a.1) { + Ordering::Equal => a.0.cmp(&b.0), + order => order, + } +} + +pub(crate) fn cmp_i32_asc(a: &(usize, i32), b: &(usize, i32)) -> Ordering { + match a.1.cmp(&b.1) { + Ordering::Equal => a.0.cmp(&b.0), + order => order, + } +} + +pub(crate) fn cmp_i64_desc(a: &(usize, i64), b: &(usize, i64)) -> Ordering { + match b.1.cmp(&a.1) { + Ordering::Equal => a.0.cmp(&b.0), + order => order, + } +} + +pub(crate) fn cmp_i64_asc(a: &(usize, i64), b: &(usize, i64)) -> Ordering { + match a.1.cmp(&b.1) { + Ordering::Equal => a.0.cmp(&b.0), + order => order, + } +} + +pub(crate) fn cmp_bool_desc(a: &(usize, bool), b: &(usize, bool)) -> Ordering { + match (a.1, b.1) { + (true, true) | (false, false) => a.0.cmp(&b.0), + (true, false) => Ordering::Less, + (false, true) => Ordering::Greater, + } +} + +pub(crate) fn cmp_bool_asc(a: &(usize, bool), b: &(usize, bool)) -> Ordering { + match (a.1, b.1) { + (true, true) | (false, false) => a.0.cmp(&b.0), + (true, false) => Ordering::Greater, + (false, true) => Ordering::Less, + } +} + +pub(crate) fn ensure_non_empty(numel: usize) -> Result<()> { + if numel == 0 { + Err(MinitensorError::invalid_argument( + "median() does not support empty tensors".to_string(), + )) + } else { + Ok(()) + } +} + +pub fn median( + tensor: &Tensor, + dim: Option, + keepdim: bool, +) -> Result<(Tensor, Option)> { + ensure_non_empty(tensor.numel())?; + + if tensor.ndim() == 0 { + return Ok((tensor.clone(), None)); + } + + let (values, indices, norm_dim) = match dim { + None => { + let (values, indices) = median_all(tensor)?; + (values, indices, None) + } + Some(dim_value) => { + let axis = if tensor.ndim() == 0 { + if dim_value == 0 || dim_value == -1 { + 0 + } else { + return Err(MinitensorError::index_error(dim_value, 0, 1)); + } + } else { + normalize_dim(dim_value, tensor.ndim())? + }; + let (values, indices) = median_along_dim(tensor, axis, keepdim)?; + (values, Some(indices), Some(axis)) + } + }; + let values = attach_median_grad(values, tensor, norm_dim, keepdim, false)?; + Ok((values, indices)) +} + +/// Attach a [`MedianBackward`] gradient to a median value reduction. +fn attach_median_grad( + values: Tensor, + input: &Tensor, + dim: Option, + keepdim: bool, + nan_aware: bool, +) -> Result { + if !input.requires_grad() || !input.dtype().is_float() { + return Ok(values); + } + let grad_fn = Arc::new(MedianBackward { + input_id: input.id(), + input: input.detach(), + dim, + keepdim, + nan_aware, + }); + let mut values = values; + values.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&values, Some(grad_fn))?; + Ok(values) +} + +/// Compute the q-th quantile of the tensor data. +pub fn quantile( + tensor: &Tensor, + q: f64, + dim: Option, + keepdim: bool, + interpolation: QuantileInterpolation, +) -> Result { + if tensor.numel() == 0 { + return Err(MinitensorError::invalid_argument( + "quantile() does not support empty tensors".to_string(), + )); + } + + validate_quantile_value(q)?; + ensure_floating_point_dtype(tensor.dtype())?; + + let (output, norm_dim) = match dim { + None => (quantile_all(tensor, q, keepdim, interpolation)?, None), + Some(dim_value) => { + if tensor.ndim() == 0 { + if dim_value == 0 || dim_value == -1 { + (quantile_all(tensor, q, keepdim, interpolation)?, None) + } else { + return Err(MinitensorError::index_error(dim_value, 0, 1)); + } + } else { + let axis = normalize_dim(dim_value, tensor.ndim())?; + ( + quantile_along_dim(tensor, axis, keepdim, q, interpolation)?, + Some(axis), + ) + } + } + }; + + attach_quantile_grad(output, tensor, norm_dim, q, interpolation, false) +} + +/// Attach a [`QuantileBackward`] gradient to a quantile value reduction. +fn attach_quantile_grad( + output: Tensor, + input: &Tensor, + dim: Option, + q: f64, + interpolation: QuantileInterpolation, + nan_aware: bool, +) -> Result { + if !input.requires_grad() || !input.dtype().is_float() { + return Ok(output); + } + let grad_fn = Arc::new(QuantileBackward { + input_id: input.id(), + input: input.detach(), + dim, + q, + interpolation, + nan_aware, + }); + let mut output = output; + output.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output, Some(grad_fn))?; + Ok(output) +} + +/// Compute multiple quantiles of the tensor data in a single pass. +pub fn quantiles( + tensor: &Tensor, + qs: &[f64], + dim: Option, + keepdim: bool, + interpolation: QuantileInterpolation, +) -> Result { + if tensor.numel() == 0 { + return Err(MinitensorError::invalid_argument( + "quantile() does not support empty tensors".to_string(), + )); + } + + if qs.is_empty() { + return Err(MinitensorError::invalid_argument( + "quantile() expected at least one probability value".to_string(), + )); + } + + for &q in qs { + validate_quantile_value(q)?; + } + + ensure_floating_point_dtype(tensor.dtype())?; + + match dim { + None => quantiles_all(tensor, qs, keepdim, interpolation), + Some(dim_value) => { + if tensor.ndim() == 0 { + if dim_value == 0 || dim_value == -1 { + return quantiles_all(tensor, qs, keepdim, interpolation); + } + return Err(MinitensorError::index_error(dim_value, 0, 1)); + } + + let axis = normalize_dim(dim_value, tensor.ndim())?; + quantiles_along_dim(tensor, axis, qs, keepdim, interpolation) + } + } +} + +/// Compute the q-th quantile of the tensor data while ignoring NaN values. +pub fn nanquantile( + tensor: &Tensor, + q: f64, + dim: Option, + keepdim: bool, + interpolation: QuantileInterpolation, +) -> Result { + if tensor.numel() == 0 { + return Err(MinitensorError::invalid_argument( + "nanquantile() does not support empty tensors".to_string(), + )); + } + + validate_quantile_value(q)?; + ensure_floating_point_dtype(tensor.dtype())?; + + let (output, norm_dim) = match dim { + None => (nanquantile_all(tensor, q, keepdim, interpolation)?, None), + Some(dim_value) => { + if tensor.ndim() == 0 { + if dim_value == 0 || dim_value == -1 { + (nanquantile_all(tensor, q, keepdim, interpolation)?, None) + } else { + return Err(MinitensorError::index_error(dim_value, 0, 1)); + } + } else { + let axis = normalize_dim(dim_value, tensor.ndim())?; + ( + nanquantile_along_dim(tensor, axis, keepdim, q, interpolation)?, + Some(axis), + ) + } + } + }; + attach_quantile_grad(output, tensor, norm_dim, q, interpolation, true) +} + +/// Compute the median while ignoring NaN values. +pub fn nanmedian(tensor: &Tensor, dim: Option, keepdim: bool) -> Result { + ensure_floating_point_dtype_for(tensor.dtype(), "nanmedian")?; + + let (values, norm_dim) = match dim { + None => (nanmedian_all(tensor, keepdim)?, None), + Some(dim_value) => { + if tensor.ndim() == 0 { + if dim_value == 0 || dim_value == -1 { + (nanmedian_all(tensor, keepdim)?, None) + } else { + return Err(MinitensorError::index_error(dim_value, 0, 1)); + } + } else { + let axis = normalize_dim(dim_value, tensor.ndim())?; + (nanmedian_along_dim(tensor, axis, keepdim)?, Some(axis)) + } + } + }; + attach_median_grad(values, tensor, norm_dim, keepdim, true) +} + +/// Compute multiple quantiles of the tensor data in a single pass while ignoring NaN values. +pub fn nanquantiles( + tensor: &Tensor, + qs: &[f64], + dim: Option, + keepdim: bool, + interpolation: QuantileInterpolation, +) -> Result { + if tensor.numel() == 0 { + return Err(MinitensorError::invalid_argument( + "nanquantile() does not support empty tensors".to_string(), + )); + } + + if qs.is_empty() { + return Err(MinitensorError::invalid_argument( + "nanquantile() expected at least one probability value".to_string(), + )); + } + + for &q in qs { + validate_quantile_value(q)?; + } + + ensure_floating_point_dtype(tensor.dtype())?; + + match dim { + None => nanquantiles_all(tensor, qs, keepdim, interpolation), + Some(dim_value) => { + if tensor.ndim() == 0 { + if dim_value == 0 || dim_value == -1 { + return nanquantiles_all(tensor, qs, keepdim, interpolation); + } + return Err(MinitensorError::index_error(dim_value, 0, 1)); + } + + let axis = normalize_dim(dim_value, tensor.ndim())?; + nanquantiles_along_dim(tensor, axis, qs, keepdim, interpolation) + } + } +} + +fn validate_quantile_value(q: f64) -> Result<()> { + if !q.is_finite() { + return Err(MinitensorError::invalid_argument( + "quantile() requires a finite probability in [0, 1]".to_string(), + )); + } + if !(0.0..=1.0).contains(&q) { + return Err(MinitensorError::invalid_argument(format!( + "quantile() expected q in [0, 1], got {q}", + ))); + } + Ok(()) +} + +fn ensure_floating_point_dtype(dtype: DataType) -> Result<()> { + ensure_floating_point_dtype_for(dtype, "quantile") +} + +fn ensure_floating_point_dtype_for(dtype: DataType, operation: &str) -> Result<()> { + match dtype { + DataType::Float32 | DataType::Float64 => Ok(()), + _ => Err(MinitensorError::invalid_operation(format!( + "{operation}() currently supports only floating point tensors" + ))), + } +} + +pub(crate) fn fill_quantile_single_f32( + input: &[f32], + values: &mut [f32], + outer: usize, + inner: usize, + outer_stride: usize, +) { + for o in 0..outer { + for r in 0..inner { + let idx = o * outer_stride + r; + let value = input[idx]; + values[o * inner + r] = if value.is_nan() { f32::NAN } else { value }; + } + } +} + +pub(crate) fn fill_quantile_single_f64( + input: &[f64], + values: &mut [f64], + outer: usize, + inner: usize, + outer_stride: usize, +) { + for o in 0..outer { + for r in 0..inner { + let idx = o * outer_stride + r; + let value = input[idx]; + values[o * inner + r] = if value.is_nan() { f64::NAN } else { value }; + } + } +} + +pub(crate) fn fill_quantiles_single_f32( + input: &[f32], + values: &mut [f32], + outer: usize, + inner: usize, + outer_stride: usize, + q_len: usize, +) { + for o in 0..outer { + for r in 0..inner { + let idx = o * outer_stride + r; + let value = input[idx]; + for qi in 0..q_len { + let out_idx = ((qi * outer) + o) * inner + r; + values[out_idx] = if value.is_nan() { f32::NAN } else { value }; + } + } + } +} + +pub(crate) fn fill_quantiles_single_f64( + input: &[f64], + values: &mut [f64], + outer: usize, + inner: usize, + outer_stride: usize, + q_len: usize, +) { + for o in 0..outer { + for r in 0..inner { + let idx = o * outer_stride + r; + let value = input[idx]; + for qi in 0..q_len { + let out_idx = ((qi * outer) + o) * inner + r; + values[out_idx] = if value.is_nan() { f64::NAN } else { value }; + } + } + } +} + +pub(crate) fn fill_nanquantile_single_f32( + input: &[f32], + values: &mut [f32], + outer: usize, + inner: usize, + outer_stride: usize, +) -> Result<()> { + for o in 0..outer { + for r in 0..inner { + let idx = o * outer_stride + r; + let value = input[idx]; + if value.is_nan() { + return Err(MinitensorError::invalid_argument( + NANQUANTILE_ALL_NAN_ERR.to_string(), + )); + } + values[o * inner + r] = value; + } + } + Ok(()) +} + +pub(crate) fn fill_nanquantile_single_f64( + input: &[f64], + values: &mut [f64], + outer: usize, + inner: usize, + outer_stride: usize, +) -> Result<()> { + for o in 0..outer { + for r in 0..inner { + let idx = o * outer_stride + r; + let value = input[idx]; + if value.is_nan() { + return Err(MinitensorError::invalid_argument( + NANQUANTILE_ALL_NAN_ERR.to_string(), + )); + } + values[o * inner + r] = value; + } + } + Ok(()) +} + +pub(crate) fn fill_nanquantiles_single_f32( + input: &[f32], + values: &mut [f32], + outer: usize, + inner: usize, + outer_stride: usize, + q_len: usize, +) -> Result<()> { + for o in 0..outer { + for r in 0..inner { + let idx = o * outer_stride + r; + let value = input[idx]; + if value.is_nan() { + return Err(MinitensorError::invalid_argument( + NANQUANTILE_ALL_NAN_ERR.to_string(), + )); + } + for qi in 0..q_len { + let out_idx = ((qi * outer) + o) * inner + r; + values[out_idx] = value; + } + } + } + Ok(()) +} + +pub(crate) fn fill_nanquantiles_single_f64( + input: &[f64], + values: &mut [f64], + outer: usize, + inner: usize, + outer_stride: usize, + q_len: usize, +) -> Result<()> { + for o in 0..outer { + for r in 0..inner { + let idx = o * outer_stride + r; + let value = input[idx]; + if value.is_nan() { + return Err(MinitensorError::invalid_argument( + NANQUANTILE_ALL_NAN_ERR.to_string(), + )); + } + for qi in 0..q_len { + let out_idx = ((qi * outer) + o) * inner + r; + values[out_idx] = value; + } + } + } + Ok(()) +} + +pub(crate) fn fill_quantiles_all_single_f32(value: f32, values: &mut [f32]) { + if value.is_nan() { + values.fill(f32::NAN); + } else { + values.fill(value); + } +} + +pub(crate) fn fill_quantiles_all_single_f64(value: f64, values: &mut [f64]) { + if value.is_nan() { + values.fill(f64::NAN); + } else { + values.fill(value); + } +} + +pub(crate) fn fill_nanquantiles_all_single_f32(value: f32, values: &mut [f32]) -> Result<()> { + if value.is_nan() { + return Err(MinitensorError::invalid_argument( + NANQUANTILE_ALL_NAN_ERR.to_string(), + )); + } + values.fill(value); + Ok(()) +} + +pub(crate) fn fill_nanquantiles_all_single_f64(value: f64, values: &mut [f64]) -> Result<()> { + if value.is_nan() { + return Err(MinitensorError::invalid_argument( + NANQUANTILE_ALL_NAN_ERR.to_string(), + )); + } + values.fill(value); + Ok(()) +} + +#[derive(Clone, Copy)] +pub(crate) struct QuantilePosition { + pub(crate) lower_idx: usize, + pub(crate) upper_idx: usize, + pub(crate) nearest_idx: usize, + pub(crate) weight: f64, +} + +pub(crate) fn quantile_positions_for_len(len: usize, qs: &[f64]) -> Vec { + let mut positions = Vec::with_capacity(qs.len()); + for &q in qs { + positions.push(quantile_position_for_len_q(len, q)); + } + positions +} + +#[inline(always)] +pub(crate) fn quantile_position_for_len_q(len: usize, q: f64) -> QuantilePosition { + if len <= 1 { + return QuantilePosition { + lower_idx: 0, + upper_idx: 0, + nearest_idx: 0, + weight: 0.0, + }; + } + + let max_index = (len - 1) as f64; + let pos = (q * max_index).clamp(0.0, max_index); + let lower_idx = pos.floor() as usize; + let upper_idx = pos.ceil() as usize; + let weight = (pos - lower_idx as f64).clamp(0.0, 1.0); + let nearest_idx = nearest_index_with_tie_even(lower_idx, upper_idx, weight); + + QuantilePosition { + lower_idx, + upper_idx, + nearest_idx, + weight, + } +} + +pub(crate) fn quantiles_from_sorted_f32( + values: &[f32], + positions: &[QuantilePosition], + interpolation: QuantileInterpolation, + output: &mut [f32], +) { + for (slot, position) in output.iter_mut().zip(positions.iter()) { + *slot = quantile_from_sorted_position_f32(values, position, interpolation); + } +} + +pub(crate) fn quantiles_from_sorted_f64( + values: &[f64], + positions: &[QuantilePosition], + interpolation: QuantileInterpolation, + output: &mut [f64], +) { + for (slot, position) in output.iter_mut().zip(positions.iter()) { + *slot = quantile_from_sorted_position_f64(values, position, interpolation); + } +} + +#[inline(always)] +fn nearest_index_with_tie_even(lower_idx: usize, upper_idx: usize, weight: f64) -> usize { + debug_assert!((0.0..=1.0).contains(&weight)); + + if upper_idx <= lower_idx { + debug_assert_eq!( + upper_idx, lower_idx, + "nearest_index_with_tie_even requires upper_idx >= lower_idx" + ); + return lower_idx; + } + + debug_assert_eq!(upper_idx, lower_idx + 1); + + if weight < 0.5 { + lower_idx + } else if weight > 0.5 { + upper_idx + } else { + lower_idx + (lower_idx & 1) + } +} + +pub(crate) fn quantile_from_sorted_position_f32( + values: &[f32], + position: &QuantilePosition, + interpolation: QuantileInterpolation, +) -> f32 { + match interpolation { + QuantileInterpolation::Lower => values[position.lower_idx], + QuantileInterpolation::Higher => values[position.upper_idx], + QuantileInterpolation::Nearest => values[position.nearest_idx], + QuantileInterpolation::Linear | QuantileInterpolation::Midpoint => { + let lower = values[position.lower_idx] as f64; + let upper = values[position.upper_idx] as f64; + interpolation.interpolate(lower, upper, position.weight) as f32 + } + } +} + +pub(crate) fn quantile_from_sorted_position_f64( + values: &[f64], + position: &QuantilePosition, + interpolation: QuantileInterpolation, +) -> f64 { + match interpolation { + QuantileInterpolation::Lower => values[position.lower_idx], + QuantileInterpolation::Higher => values[position.upper_idx], + QuantileInterpolation::Nearest => values[position.nearest_idx], + QuantileInterpolation::Linear | QuantileInterpolation::Midpoint => { + let lower = values[position.lower_idx]; + let upper = values[position.upper_idx]; + interpolation.interpolate(lower, upper, position.weight) + } + } +} + +fn quantile_all( + tensor: &Tensor, + q: f64, + keepdim: bool, + interpolation: QuantileInterpolation, +) -> Result { + if tensor.ndim() == 0 { + return Ok(tensor.clone()); + } + + let mut result_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => { + let data = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let mut values = Vec::with_capacity(data.len()); + for &value in data { + if value.is_nan() { + result_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?[0] = f32::NAN; + let result_shape = if keepdim { + Shape::new(vec![1; tensor.ndim()]) + } else { + Shape::scalar() + }; + return Ok(Tensor::new( + Arc::new(result_data), + result_shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )); + } + values.push(value); + } + let quant = quantile_from_unsorted_f32(&mut values, q, interpolation); + result_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?[0] = quant; + } + DataType::Float64 => { + let data = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let mut values = Vec::with_capacity(data.len()); + for &value in data { + if value.is_nan() { + result_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?[0] = f64::NAN; + let result_shape = if keepdim { + Shape::new(vec![1; tensor.ndim()]) + } else { + Shape::scalar() + }; + return Ok(Tensor::new( + Arc::new(result_data), + result_shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )); + } + values.push(value); + } + let quant = quantile_from_unsorted_f64(&mut values, q, interpolation); + result_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?[0] = quant; + } + _ => unreachable!("dtype validated"), + } + + let result_shape = if keepdim { + Shape::new(vec![1; tensor.ndim()]) + } else { + Shape::scalar() + }; + + Ok(Tensor::new( + Arc::new(result_data), + result_shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )) +} + +fn nanquantile_all( + tensor: &Tensor, + q: f64, + keepdim: bool, + interpolation: QuantileInterpolation, +) -> Result { + if tensor.ndim() == 0 { + match tensor.dtype() { + DataType::Float32 => { + let value = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?[0]; + if value.is_nan() { + return Err(MinitensorError::invalid_argument( + NANQUANTILE_ALL_NAN_ERR.to_string(), + )); + } + } + DataType::Float64 => { + let value = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?[0]; + if value.is_nan() { + return Err(MinitensorError::invalid_argument( + NANQUANTILE_ALL_NAN_ERR.to_string(), + )); + } + } + _ => unreachable!("dtype validated"), + } + return Ok(tensor.clone()); + } + + let mut result_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => { + let data = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let mut values: Vec = data.iter().copied().filter(|v| !v.is_nan()).collect(); + if values.is_empty() { + return Err(MinitensorError::invalid_argument( + NANQUANTILE_ALL_NAN_ERR.to_string(), + )); + } + let quant = quantile_from_unsorted_f32(&mut values, q, interpolation); + result_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?[0] = quant; + } + DataType::Float64 => { + let data = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let mut values: Vec = data.iter().copied().filter(|v| !v.is_nan()).collect(); + if values.is_empty() { + return Err(MinitensorError::invalid_argument( + NANQUANTILE_ALL_NAN_ERR.to_string(), + )); + } + let quant = quantile_from_unsorted_f64(&mut values, q, interpolation); + result_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?[0] = quant; + } + _ => unreachable!("dtype validated"), + } + + let result_shape = if keepdim { + Shape::new(vec![1; tensor.ndim()]) + } else { + Shape::scalar() + }; + + Ok(Tensor::new( + Arc::new(result_data), + result_shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )) +} + +#[cfg(test)] +mod core_tests { + use super::*; + use crate::Device; + + #[test] + fn test_quantile_interpolation_modes() { + let lower = 1.0; + let upper = 3.0; + let weight = 0.25; + + assert_eq!( + QuantileInterpolation::Linear.interpolate(lower, upper, weight), + 1.5 + ); + assert_eq!( + QuantileInterpolation::Lower.interpolate(lower, upper, weight), + lower + ); + assert_eq!( + QuantileInterpolation::Higher.interpolate(lower, upper, weight), + upper + ); + assert_eq!( + QuantileInterpolation::Midpoint.interpolate(lower, upper, weight), + 2.0 + ); + assert_eq!( + QuantileInterpolation::Nearest.interpolate(lower, upper, 0.49), + lower + ); + assert_eq!( + QuantileInterpolation::Nearest.interpolate(lower, upper, 0.5), + upper + ); + } + + #[test] + fn test_normalize_reduction_dims_sorts_dedups_and_supports_negative_dims() { + let dims = Some(vec![2, -1, 0, 2, -3]); + let normalized = normalize_reduction_dims(dims, 3).unwrap(); + assert_eq!(normalized, Some(vec![0, 2])); + assert_eq!(normalize_reduction_dims(None, 3).unwrap(), None); + } + + #[test] + fn test_normalize_reduction_dims_rejects_out_of_range() { + assert!(normalize_reduction_dims(Some(vec![3]), 3).is_err()); + assert!(normalize_reduction_dims(Some(vec![-4]), 3).is_err()); + } + + #[test] + fn test_non_nan_mask_for_float_tensors_and_invalid_dtype() { + let f32_tensor = Tensor::new( + Arc::new(TensorData::from_vec_f32( + vec![1.0, f32::NAN, -2.0], + Device::cpu(), + )), + Shape::new(vec![3]), + DataType::Float32, + Device::cpu(), + false, + ); + let mask = non_nan_mask(&f32_tensor).unwrap(); + assert_eq!(mask.dtype(), DataType::Bool); + assert_eq!(mask.data().as_bool_slice().unwrap(), &[true, false, true],); + + let f64_tensor = Tensor::new( + Arc::new(TensorData::from_vec_f64( + vec![f64::NAN, 4.0, 0.0], + Device::cpu(), + )), + Shape::new(vec![3]), + DataType::Float64, + Device::cpu(), + false, + ); + let f64_mask = non_nan_mask(&f64_tensor).unwrap(); + assert_eq!( + f64_mask.data().as_bool_slice().unwrap(), + &[false, true, true], + ); + + let bool_tensor = Tensor::new( + Arc::new(TensorData::from_vec_bool(vec![true, false], Device::cpu())), + Shape::new(vec![2]), + DataType::Bool, + Device::cpu(), + false, + ); + assert!(non_nan_mask(&bool_tensor).is_err()); + } + + #[test] + fn test_comparator_tie_breaking_and_nan_ordering() { + let mut f32_desc = [(1, 2.0f32), (0, 2.0), (2, f32::NAN)]; + f32_desc.sort_by(cmp_f32_desc); + assert!(f32_desc[0].1.is_nan()); + assert_eq!(f32_desc[0].0, 2); + assert_eq!( + f32_desc[1..].iter().map(|(i, _)| *i).collect::>(), + vec![0, 1] + ); + + let mut f32_asc = [(1, 2.0f32), (0, 2.0), (2, f32::NAN)]; + f32_asc.sort_by(cmp_f32_asc); + assert!(f32_asc[2].1.is_nan()); + assert_eq!(f32_asc[2].0, 2); + assert_eq!( + f32_asc[..2].iter().map(|(i, _)| *i).collect::>(), + vec![0, 1] + ); + + let mut f64_desc = [(1, 2.0f64), (0, 2.0), (2, f64::NAN)]; + f64_desc.sort_by(cmp_f64_desc); + assert!(f64_desc[0].1.is_nan()); + assert_eq!(f64_desc[0].0, 2); + assert_eq!( + f64_desc[1..].iter().map(|(i, _)| *i).collect::>(), + vec![0, 1] + ); + + let mut f64_asc = [(1, 2.0f64), (0, 2.0), (2, f64::NAN)]; + f64_asc.sort_by(cmp_f64_asc); + assert!(f64_asc[2].1.is_nan()); + assert_eq!(f64_asc[2].0, 2); + assert_eq!( + f64_asc[..2].iter().map(|(i, _)| *i).collect::>(), + vec![0, 1] + ); + + let mut i32_desc = vec![(1, 7_i32), (0, 7_i32), (2, 1_i32)]; + i32_desc.sort_by(cmp_i32_desc); + assert_eq!(i32_desc, vec![(0, 7), (1, 7), (2, 1)]); + + let mut i32_asc = vec![(1, 7_i32), (0, 7_i32), (2, 1_i32)]; + i32_asc.sort_by(cmp_i32_asc); + assert_eq!(i32_asc, vec![(2, 1), (0, 7), (1, 7)]); + + let mut i64_desc = vec![(1, 7_i64), (0, 7_i64), (2, 1_i64)]; + i64_desc.sort_by(cmp_i64_desc); + assert_eq!(i64_desc, vec![(0, 7), (1, 7), (2, 1)]); + + let mut i64_asc = vec![(1, 7_i64), (0, 7_i64), (2, 1_i64)]; + i64_asc.sort_by(cmp_i64_asc); + assert_eq!(i64_asc, vec![(2, 1), (0, 7), (1, 7)]); + + let mut bool_desc = vec![(1, true), (0, true), (2, false)]; + bool_desc.sort_by(cmp_bool_desc); + assert_eq!(bool_desc, vec![(0, true), (1, true), (2, false)]); + + let mut bool_asc = vec![(1, true), (0, true), (2, false)]; + bool_asc.sort_by(cmp_bool_asc); + assert_eq!(bool_asc, vec![(2, false), (0, true), (1, true)]); + } + + #[test] + fn test_ensure_non_empty_guard() { + assert!(ensure_non_empty(0).is_err()); + assert!(ensure_non_empty(1).is_ok()); + } +} diff --git a/engine/src/operations/reduction/logsumexp.rs b/engine/src/operations/reduction/logsumexp.rs index 4223093b..d2566402 100644 --- a/engine/src/operations/reduction/logsumexp.rs +++ b/engine/src/operations/reduction/logsumexp.rs @@ -1,790 +1,804 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -/// Numerically stable log-sum-exp reduction along specified dimensions -pub fn logsumexp(tensor: &Tensor, dim: Option>, keepdim: bool) -> Result { - match tensor.dtype() { - DataType::Float32 | DataType::Float64 => {} - _ => { - return Err(MinitensorError::invalid_operation( - "Logsumexp only supported for floating point tensors", - )); - } - } - - let ndim = tensor.ndim() as isize; - let dims = match dim { - Some(dims) => { - if dims.is_empty() { - Vec::new() - } else { - let mut normalized = Vec::with_capacity(dims.len()); - for d in dims { - let d = if d < 0 { d + ndim } else { d }; - if d < 0 || d >= ndim { - return Err(MinitensorError::index_error(d, 0, tensor.ndim())); - } - normalized.push(d as usize); - } - normalized.sort_unstable(); - normalized.dedup(); - normalized - } - } - None => (0..tensor.ndim()).collect(), - }; - - if dims.is_empty() { - return Ok(tensor.clone()); - } - - let mut max_tensor = tensor.clone(); - for &d in &dims { - max_tensor = max_along_dim(&max_tensor, d, true)?; - } - let max_tensor = max_tensor.detach(); - - let shifted = arithmetic::sub(tensor, &max_tensor)?; - let exp_shifted = activation::exp(&shifted)?; - let dims_isize: Vec = dims.iter().map(|&d| d as isize).collect(); - let sum_exp = sum(&exp_shifted, Some(dims_isize), true)?; - let log_sum = activation::log(&sum_exp)?; - let mut result = arithmetic::add(&max_tensor, &log_sum)?; - - if !keepdim { - let mut new_dims = Vec::with_capacity(result.ndim() - dims.len()); - for (idx, &size) in result.shape().dims().iter().enumerate() { - if dims.binary_search(&idx).is_err() { - new_dims.push(size); - } - } - - let target_shape = if new_dims.is_empty() { - Shape::scalar() - } else { - Shape::new(new_dims) - }; - - result = shape_ops::reshape(&result, target_shape)?; - } - - Ok(result) -} - -/// Product reduction along specified dimensions -pub fn prod(tensor: &Tensor, dim: Option>, keepdim: bool) -> Result { - // Normalise negative dimensions and deduplicate - let ndim = tensor.ndim() as isize; - let dim = match dim { - Some(dims) => { - let mut normalized = Vec::with_capacity(dims.len()); - for d in dims { - let d = if d < 0 { d + ndim } else { d }; - if d < 0 || d >= ndim { - return Err(MinitensorError::index_error(d, 0, tensor.ndim())); - } - normalized.push(d as usize); - } - normalized.sort_unstable(); - normalized.dedup(); - Some(normalized) - } - None => None, - }; - let dims_clone = dim.clone(); - - let result = match dim { - None => { - let result_shape = if keepdim { - Shape::new(vec![1; tensor.ndim()]) - } else { - Shape::scalar() - }; - - let mut result_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); - match tensor.dtype() { - DataType::Float32 => prod_all_f32(tensor, &mut result_data)?, - DataType::Float64 => prod_all_f64(tensor, &mut result_data)?, - DataType::Int32 => prod_all_i32(tensor, &mut result_data)?, - DataType::Int64 => prod_all_i64(tensor, &mut result_data)?, - DataType::Bool => prod_all_bool(tensor, &mut result_data)?, - } - - let requires_grad = tensor.requires_grad() && tensor.dtype() != DataType::Bool; - Tensor::new( - Arc::new(result_data), - result_shape, - tensor.dtype(), - tensor.device(), - requires_grad, - ) - } - Some(dims) => { - if dims.is_empty() { - tensor.clone() - } else { - let mut result = tensor.clone(); - if keepdim { - for &d in &dims { - result = prod_along_dim(&result, d, true)?; - } - } else { - for &d in dims.iter().rev() { - result = prod_along_dim(&result, d, false)?; - } - } - result - } - } - }; - - if result.requires_grad() { - let grad_fn = Arc::new(ProdBackward { - input: tensor.detach(), - result: result.clone(), - input_id: tensor.id(), - dims: dims_clone, - keepdim, - }); - let mut result_with_grad = result; - result_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&result_with_grad, Some(grad_fn))?; - Ok(result_with_grad) - } else { - Ok(result) - } -} - -/// Cumulative sum along a specified dimension -pub fn cumsum(tensor: &Tensor, dim: isize) -> Result { - let dim = normalize_dim(dim, tensor.ndim())?; - - let mut result_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => cumsum_f32(tensor, &mut result_data, dim)?, - DataType::Float64 => cumsum_f64(tensor, &mut result_data, dim)?, - DataType::Int32 => cumsum_i32(tensor, &mut result_data, dim)?, - DataType::Int64 => cumsum_i64(tensor, &mut result_data, dim)?, - DataType::Bool => { - return Err(MinitensorError::invalid_operation( - "Cumsum not supported for boolean tensors", - )); - } - } - - let result = Tensor::new( - Arc::new(result_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - if result.requires_grad() { - let grad_fn = Arc::new(CumsumBackward { - input_id: tensor.id(), - dim, - }); - let mut result_with_grad = result; - result_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&result_with_grad, Some(grad_fn))?; - Ok(result_with_grad) - } else { - Ok(result) - } -} - -/// Backward helper for cumulative sum -pub fn cumsum_backward(tensor: &Tensor, dim: usize) -> Result { - if dim >= tensor.ndim() { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - - let mut result_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => cumsum_backward_f32(tensor, &mut result_data, dim)?, - DataType::Float64 => cumsum_backward_f64(tensor, &mut result_data, dim)?, - DataType::Int32 => cumsum_backward_i32(tensor, &mut result_data, dim)?, - DataType::Int64 => cumsum_backward_i64(tensor, &mut result_data, dim)?, - DataType::Bool => { - return Err(MinitensorError::invalid_operation( - "Cumsum not supported for boolean tensors", - )); - } - } - - Ok(Tensor::new( - Arc::new(result_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - false, - )) -} - -/// Cumulative product along a specified dimension -pub fn cumprod(tensor: &Tensor, dim: isize) -> Result { - let dim = normalize_dim(dim, tensor.ndim())?; - - let mut result_data = - TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => cumprod_f32(tensor, &mut result_data, dim)?, - DataType::Float64 => cumprod_f64(tensor, &mut result_data, dim)?, - DataType::Int32 => cumprod_i32(tensor, &mut result_data, dim)?, - DataType::Int64 => cumprod_i64(tensor, &mut result_data, dim)?, - DataType::Bool => { - return Err(MinitensorError::invalid_operation( - "Cumprod not supported for boolean tensors", - )); - } - } - - let requires_grad = - tensor.requires_grad() && matches!(tensor.dtype(), DataType::Float32 | DataType::Float64); - - let result = Tensor::new( - Arc::new(result_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - requires_grad, - ); - - if result.requires_grad() { - let grad_fn = Arc::new(CumprodBackward { - input_id: tensor.id(), - input: tensor.clone(), - output: result.clone(), - dim, - }); - let mut result_with_grad = result; - result_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&result_with_grad, Some(grad_fn))?; - Ok(result_with_grad) - } else { - Ok(result) - } -} - -/// Backward helper for cumulative product -pub fn cumprod_backward( - input: &Tensor, - output: &Tensor, - grad: &Tensor, - dim: usize, -) -> Result { - if dim >= input.ndim() { - return Err(MinitensorError::index_error(dim as isize, 0, input.ndim())); - } - - let mut result_data = TensorData::zeros_on_device(input.numel(), input.dtype(), input.device()); - - match input.dtype() { - DataType::Float32 => cumprod_backward_f32(input, output, grad, &mut result_data, dim)?, - DataType::Float64 => cumprod_backward_f64(input, output, grad, &mut result_data, dim)?, - _ => { - return Err(MinitensorError::invalid_operation( - "Cumprod backward only supported for floating point tensors", - )); - } - } - - Ok(Tensor::new( - Arc::new(result_data), - input.shape().clone(), - input.dtype(), - input.device(), - false, - )) -} - -/// Mean reduction along specified dimensions -pub fn mean(tensor: &Tensor, dim: Option>, keepdim: bool) -> Result { - // Normalise negative dimensions and deduplicate - let ndim = tensor.ndim() as isize; - let normalized = match dim { - Some(dims) => { - let mut normalized = Vec::with_capacity(dims.len()); - for d in dims { - let d = if d < 0 { d + ndim } else { d }; - if d < 0 || d >= ndim { - return Err(MinitensorError::index_error(d, 0, tensor.ndim())); - } - normalized.push(d as usize); - } - normalized.sort_unstable(); - normalized.dedup(); - Some(normalized) - } - None => None, - }; - - let sum_result = sum( - tensor, - normalized - .clone() - .map(|d| d.iter().map(|&x| x as isize).collect()), - keepdim, - )?; - - // Compute the number of elements being averaged - let num_elements = match &normalized { - None => tensor.numel() as f64, - Some(dims) => { - if dims.is_empty() { - return Ok(tensor.clone()); - } - - let mut count = 1.0; - for &d in dims { - count *= tensor.shape().dims()[d] as f64; - } - count - } - }; - - // Prepare sum tensor and divisor for division - let (sum_tensor, divisor) = match tensor.dtype() { - DataType::Float32 => ( - sum_result, - Tensor::new( - Arc::new(TensorData::from_vec( - vec![num_elements as f32], - DataType::Float32, - tensor.device(), - )), - Shape::scalar(), - DataType::Float32, - tensor.device(), - false, - ), - ), - DataType::Float64 => ( - sum_result, - Tensor::new( - Arc::new(TensorData::from_vec( - vec![num_elements], - DataType::Float64, - tensor.device(), - )), - Shape::scalar(), - DataType::Float64, - tensor.device(), - false, - ), - ), - DataType::Int32 => ( - sum_result.astype(DataType::Float32)?, - Tensor::new( - Arc::new(TensorData::from_vec( - vec![num_elements as f32], - DataType::Float32, - tensor.device(), - )), - Shape::scalar(), - DataType::Float32, - tensor.device(), - false, - ), - ), - DataType::Int64 => ( - sum_result.astype(DataType::Float64)?, - Tensor::new( - Arc::new(TensorData::from_vec( - vec![num_elements], - DataType::Float64, - tensor.device(), - )), - Shape::scalar(), - DataType::Float64, - tensor.device(), - false, - ), - ), - DataType::Bool => { - return Err(MinitensorError::invalid_operation( - "Mean not supported for boolean tensors", - )); - } - }; - - crate::operations::arithmetic::div(&sum_tensor, &divisor) -} - -/// NaN-aware mean reduction along specified dimensions -pub fn nanmean(tensor: &Tensor, dim: Option>, keepdim: bool) -> Result { - if !tensor.dtype().is_float() { - return mean(tensor, dim, keepdim); - } - - let dim = normalize_reduction_dims(dim, tensor.ndim())?; - let dims_clone = dim.clone(); - let needs_mask = - tensor.requires_grad() || dim.as_ref().map(|dims| !dims.is_empty()).unwrap_or(false); - let mask = if needs_mask { - Some(non_nan_mask(tensor)?) - } else { - None - }; - - if let Some(dims) = &dim { - if dims.is_empty() { - return Ok(tensor.clone()); - } - } - - let (sum, count) = match dim { - None => { - let result_shape = if keepdim { - Shape::new(vec![1; tensor.ndim()]) - } else { - Shape::scalar() - }; - let mut sum_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); - let mut count_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => nanmean_all_f32(tensor, &mut sum_data, &mut count_data)?, - DataType::Float64 => nanmean_all_f64(tensor, &mut sum_data, &mut count_data)?, - _ => unreachable!("nanmean only supports floating point tensors"), - } - - ( - Tensor::new( - Arc::new(sum_data), - result_shape.clone(), - tensor.dtype(), - tensor.device(), - false, - ), - Tensor::new( - Arc::new(count_data), - result_shape, - tensor.dtype(), - tensor.device(), - false, - ), - ) - } - Some(dims) => { - let mask = mask.as_ref().ok_or_else(|| { - MinitensorError::internal_error("nanmean expected mask for count computation") - })?; - let mut sum = tensor.clone(); - let mut count = mask.astype(tensor.dtype())?; - - if keepdim { - for &d in &dims { - sum = nansum_along_dim(&sum, d, true)?; - count = sum_along_dim(&count, d, true)?; - } - } else { - for &d in dims.iter().rev() { - sum = nansum_along_dim(&sum, d, false)?; - count = sum_along_dim(&count, d, false)?; - } - } - (sum, count) - } - }; - - let result = nanmean_from_sum_count(&sum, &count, tensor.requires_grad())?; - - if result.requires_grad() { - let mask = mask.ok_or_else(|| { - MinitensorError::internal_error("nanmean expected mask for gradient computation") - })?; - let grad_fn = Arc::new(NanMeanBackward { - input_id: tensor.id(), - input_shape: tensor.shape().dims().to_vec(), - dims: dims_clone, - keepdim, - mask, - count, - }); - let mut result_with_grad = result; - result_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&result_with_grad, Some(grad_fn))?; - Ok(result_with_grad) - } else { - Ok(result) - } -} - -/// Logical all reduction along specified dimension -pub fn all(tensor: &Tensor, dim: Option, keepdim: bool) -> Result { - match dim { - None => all_all(tensor, keepdim), - Some(d) => { - let d = normalize_dim(d, tensor.ndim())?; - all_along_dim(tensor, d, keepdim) - } - } -} - -/// Logical any reduction along specified dimension -pub fn any(tensor: &Tensor, dim: Option, keepdim: bool) -> Result { - match dim { - None => any_all(tensor, keepdim), - Some(d) => { - let d = normalize_dim(d, tensor.ndim())?; - any_along_dim(tensor, d, keepdim) - } - } -} - -fn all_all(tensor: &Tensor, keepdim: bool) -> Result { - let result_shape = if keepdim { - Shape::new(vec![1; tensor.ndim()]) - } else { - Shape::scalar() - }; - let mut result_data = TensorData::zeros_on_device(1, DataType::Bool, tensor.device()); - let out_slice = result_data - .as_bool_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - let all_true = match tensor.dtype() { - DataType::Float32 => tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected f32 data"))? - .par_iter() - .all(|&x| x != 0.0), - DataType::Float64 => tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected f64 data"))? - .par_iter() - .all(|&x| x != 0.0), - DataType::Int32 => tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected i32 data"))? - .par_iter() - .all(|&x| x != 0), - DataType::Int64 => tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected i64 data"))? - .par_iter() - .all(|&x| x != 0), - DataType::Bool => tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected bool data"))? - .par_iter() - .all(|&x| x), - }; - out_slice[0] = all_true; - Ok(Tensor::new( - Arc::new(result_data), - result_shape, - DataType::Bool, - tensor.device(), - false, - )) -} - -fn any_all(tensor: &Tensor, keepdim: bool) -> Result { - let result_shape = if keepdim { - Shape::new(vec![1; tensor.ndim()]) - } else { - Shape::scalar() - }; - let mut result_data = TensorData::zeros_on_device(1, DataType::Bool, tensor.device()); - let out_slice = result_data - .as_bool_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - let any_true = match tensor.dtype() { - DataType::Float32 => tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected f32 data"))? - .par_iter() - .any(|&x| x != 0.0), - DataType::Float64 => tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected f64 data"))? - .par_iter() - .any(|&x| x != 0.0), - DataType::Int32 => tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected i32 data"))? - .par_iter() - .any(|&x| x != 0), - DataType::Int64 => tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected i64 data"))? - .par_iter() - .any(|&x| x != 0), - DataType::Bool => tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected bool data"))? - .par_iter() - .any(|&x| x), - }; - out_slice[0] = any_true; - Ok(Tensor::new( - Arc::new(result_data), - result_shape, - DataType::Bool, - tensor.device(), - false, - )) -} - -fn all_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { - if dim >= tensor.ndim() { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - - let input_shape = tensor.shape().dims(); - let mut output_shape = input_shape.to_vec(); - if keepdim { - output_shape[dim] = 1; - } else { - output_shape.remove(dim); - } - let output_shape_obj = Shape::new(output_shape.clone()); - let mut result_data = - TensorData::zeros_on_device(output_shape_obj.numel(), DataType::Bool, tensor.device()); - - let dim_size = input_shape[dim]; - let _outer = input_shape[..dim].iter().product::(); - let inner = input_shape[dim + 1..].iter().product::(); - let outer_stride = dim_size * inner; - - match tensor.dtype() { - DataType::Float32 => { - let input = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let output = result_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice") - })?; - output.par_iter_mut().enumerate().for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut val = true; - for d in 0..dim_size { - let in_idx = o * outer_stride + d * inner + r; - if input[in_idx] == 0.0 { - val = false; - break; - } - } - *out = val; - }); - } - DataType::Float64 => { - let input = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let output = result_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice") - })?; - output.par_iter_mut().enumerate().for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut val = true; - for d in 0..dim_size { - let in_idx = o * outer_stride + d * inner + r; - if input[in_idx] == 0.0 { - val = false; - break; - } - } - *out = val; - }); - } - DataType::Int32 => { - let input = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - let output = result_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice") - })?; - output.par_iter_mut().enumerate().for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut val = true; - for d in 0..dim_size { - let in_idx = o * outer_stride + d * inner + r; - if input[in_idx] == 0 { - val = false; - break; - } - } - *out = val; - }); - } - DataType::Int64 => { - let input = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - let output = result_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice") - })?; - output.par_iter_mut().enumerate().for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut val = true; - for d in 0..dim_size { - let in_idx = o * outer_stride + d * inner + r; - if input[in_idx] == 0 { - val = false; - break; - } - } - *out = val; - }); - } - DataType::Bool => { - let input = tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - let output = result_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice") - })?; - output.par_iter_mut().enumerate().for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut val = true; - for d in 0..dim_size { - let in_idx = o * outer_stride + d * inner + r; - if !input[in_idx] { - val = false; - break; - } - } - *out = val; - }); - } - } - - Ok(Tensor::new( - Arc::new(result_data), - output_shape_obj, - DataType::Bool, - tensor.device(), - false, - )) -} +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::autograd::CumprodBackward; +use crate::autograd::CumsumBackward; +use crate::autograd::NanMeanBackward; +use crate::autograd::ProdBackward; +use crate::operations::{activation, arithmetic, shape_ops}; +use crate::{ + autograd::add_to_graph, + error::{MinitensorError, Result}, + tensor::{DataType, Shape, Tensor, TensorData}, +}; +use rayon::prelude::*; +use std::sync::Arc; + +/// Numerically stable log-sum-exp reduction along specified dimensions +pub fn logsumexp(tensor: &Tensor, dim: Option>, keepdim: bool) -> Result { + match tensor.dtype() { + DataType::Float32 | DataType::Float64 => {} + _ => { + return Err(MinitensorError::invalid_operation( + "Logsumexp only supported for floating point tensors", + )); + } + } + + let ndim = tensor.ndim() as isize; + let dims = match dim { + Some(dims) => { + if dims.is_empty() { + Vec::new() + } else { + let mut normalized = Vec::with_capacity(dims.len()); + for d in dims { + let d = if d < 0 { d + ndim } else { d }; + if d < 0 || d >= ndim { + return Err(MinitensorError::index_error(d, 0, tensor.ndim())); + } + normalized.push(d as usize); + } + normalized.sort_unstable(); + normalized.dedup(); + normalized + } + } + None => (0..tensor.ndim()).collect(), + }; + + if dims.is_empty() { + return Ok(tensor.clone()); + } + + let mut max_tensor = tensor.clone(); + for &d in &dims { + max_tensor = max_along_dim(&max_tensor, d, true)?; + } + let max_tensor = max_tensor.detach(); + + let shifted = arithmetic::sub(tensor, &max_tensor)?; + let exp_shifted = activation::exp(&shifted)?; + let dims_isize: Vec = dims.iter().map(|&d| d as isize).collect(); + let sum_exp = sum(&exp_shifted, Some(dims_isize), true)?; + let log_sum = activation::log(&sum_exp)?; + let mut result = arithmetic::add(&max_tensor, &log_sum)?; + + if !keepdim { + let mut new_dims = Vec::with_capacity(result.ndim() - dims.len()); + for (idx, &size) in result.shape().dims().iter().enumerate() { + if dims.binary_search(&idx).is_err() { + new_dims.push(size); + } + } + + let target_shape = if new_dims.is_empty() { + Shape::scalar() + } else { + Shape::new(new_dims) + }; + + result = shape_ops::reshape(&result, target_shape)?; + } + + Ok(result) +} + +/// Product reduction along specified dimensions +pub fn prod(tensor: &Tensor, dim: Option>, keepdim: bool) -> Result { + // Normalise negative dimensions and deduplicate + let ndim = tensor.ndim() as isize; + let dim = match dim { + Some(dims) => { + let mut normalized = Vec::with_capacity(dims.len()); + for d in dims { + let d = if d < 0 { d + ndim } else { d }; + if d < 0 || d >= ndim { + return Err(MinitensorError::index_error(d, 0, tensor.ndim())); + } + normalized.push(d as usize); + } + normalized.sort_unstable(); + normalized.dedup(); + Some(normalized) + } + None => None, + }; + let dims_clone = dim.clone(); + + let result = match dim { + None => { + let result_shape = if keepdim { + Shape::new(vec![1; tensor.ndim()]) + } else { + Shape::scalar() + }; + + let mut result_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); + match tensor.dtype() { + DataType::Float32 => prod_all_f32(tensor, &mut result_data)?, + DataType::Float64 => prod_all_f64(tensor, &mut result_data)?, + DataType::Int32 => prod_all_i32(tensor, &mut result_data)?, + DataType::Int64 => prod_all_i64(tensor, &mut result_data)?, + DataType::Bool => prod_all_bool(tensor, &mut result_data)?, + } + + let requires_grad = tensor.requires_grad() && tensor.dtype() != DataType::Bool; + Tensor::new( + Arc::new(result_data), + result_shape, + tensor.dtype(), + tensor.device(), + requires_grad, + ) + } + Some(dims) => { + if dims.is_empty() { + tensor.clone() + } else { + let mut result = tensor.clone(); + if keepdim { + for &d in &dims { + result = prod_along_dim(&result, d, true)?; + } + } else { + for &d in dims.iter().rev() { + result = prod_along_dim(&result, d, false)?; + } + } + result + } + } + }; + + if result.requires_grad() { + let grad_fn = Arc::new(ProdBackward { + input: tensor.detach(), + result: result.clone(), + input_id: tensor.id(), + dims: dims_clone, + keepdim, + }); + let mut result_with_grad = result; + result_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&result_with_grad, Some(grad_fn))?; + Ok(result_with_grad) + } else { + Ok(result) + } +} + +/// Cumulative sum along a specified dimension +pub fn cumsum(tensor: &Tensor, dim: isize) -> Result { + let dim = normalize_dim(dim, tensor.ndim())?; + + let mut result_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => cumsum_f32(tensor, &mut result_data, dim)?, + DataType::Float64 => cumsum_f64(tensor, &mut result_data, dim)?, + DataType::Int32 => cumsum_i32(tensor, &mut result_data, dim)?, + DataType::Int64 => cumsum_i64(tensor, &mut result_data, dim)?, + DataType::Bool => { + return Err(MinitensorError::invalid_operation( + "Cumsum not supported for boolean tensors", + )); + } + } + + let result = Tensor::new( + Arc::new(result_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + if result.requires_grad() { + let grad_fn = Arc::new(CumsumBackward { + input_id: tensor.id(), + dim, + }); + let mut result_with_grad = result; + result_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&result_with_grad, Some(grad_fn))?; + Ok(result_with_grad) + } else { + Ok(result) + } +} + +/// Backward helper for cumulative sum +pub fn cumsum_backward(tensor: &Tensor, dim: usize) -> Result { + if dim >= tensor.ndim() { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + + let mut result_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => cumsum_backward_f32(tensor, &mut result_data, dim)?, + DataType::Float64 => cumsum_backward_f64(tensor, &mut result_data, dim)?, + DataType::Int32 => cumsum_backward_i32(tensor, &mut result_data, dim)?, + DataType::Int64 => cumsum_backward_i64(tensor, &mut result_data, dim)?, + DataType::Bool => { + return Err(MinitensorError::invalid_operation( + "Cumsum not supported for boolean tensors", + )); + } + } + + Ok(Tensor::new( + Arc::new(result_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + false, + )) +} + +/// Cumulative product along a specified dimension +pub fn cumprod(tensor: &Tensor, dim: isize) -> Result { + let dim = normalize_dim(dim, tensor.ndim())?; + + let mut result_data = + TensorData::uninitialized_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => cumprod_f32(tensor, &mut result_data, dim)?, + DataType::Float64 => cumprod_f64(tensor, &mut result_data, dim)?, + DataType::Int32 => cumprod_i32(tensor, &mut result_data, dim)?, + DataType::Int64 => cumprod_i64(tensor, &mut result_data, dim)?, + DataType::Bool => { + return Err(MinitensorError::invalid_operation( + "Cumprod not supported for boolean tensors", + )); + } + } + + let requires_grad = + tensor.requires_grad() && matches!(tensor.dtype(), DataType::Float32 | DataType::Float64); + + let result = Tensor::new( + Arc::new(result_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + requires_grad, + ); + + if result.requires_grad() { + let grad_fn = Arc::new(CumprodBackward { + input_id: tensor.id(), + input: tensor.clone(), + output: result.clone(), + dim, + }); + let mut result_with_grad = result; + result_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&result_with_grad, Some(grad_fn))?; + Ok(result_with_grad) + } else { + Ok(result) + } +} + +/// Backward helper for cumulative product +pub fn cumprod_backward( + input: &Tensor, + output: &Tensor, + grad: &Tensor, + dim: usize, +) -> Result { + if dim >= input.ndim() { + return Err(MinitensorError::index_error(dim as isize, 0, input.ndim())); + } + + let mut result_data = TensorData::zeros_on_device(input.numel(), input.dtype(), input.device()); + + match input.dtype() { + DataType::Float32 => cumprod_backward_f32(input, output, grad, &mut result_data, dim)?, + DataType::Float64 => cumprod_backward_f64(input, output, grad, &mut result_data, dim)?, + _ => { + return Err(MinitensorError::invalid_operation( + "Cumprod backward only supported for floating point tensors", + )); + } + } + + Ok(Tensor::new( + Arc::new(result_data), + input.shape().clone(), + input.dtype(), + input.device(), + false, + )) +} + +/// Mean reduction along specified dimensions +pub fn mean(tensor: &Tensor, dim: Option>, keepdim: bool) -> Result { + // Normalise negative dimensions and deduplicate + let ndim = tensor.ndim() as isize; + let normalized = match dim { + Some(dims) => { + let mut normalized = Vec::with_capacity(dims.len()); + for d in dims { + let d = if d < 0 { d + ndim } else { d }; + if d < 0 || d >= ndim { + return Err(MinitensorError::index_error(d, 0, tensor.ndim())); + } + normalized.push(d as usize); + } + normalized.sort_unstable(); + normalized.dedup(); + Some(normalized) + } + None => None, + }; + + let sum_result = sum( + tensor, + normalized + .clone() + .map(|d| d.iter().map(|&x| x as isize).collect()), + keepdim, + )?; + + // Compute the number of elements being averaged + let num_elements = match &normalized { + None => tensor.numel() as f64, + Some(dims) => { + if dims.is_empty() { + return Ok(tensor.clone()); + } + + let mut count = 1.0; + for &d in dims { + count *= tensor.shape().dims()[d] as f64; + } + count + } + }; + + // Prepare sum tensor and divisor for division + let (sum_tensor, divisor) = match tensor.dtype() { + DataType::Float32 => ( + sum_result, + Tensor::new( + Arc::new(TensorData::from_vec( + vec![num_elements as f32], + DataType::Float32, + tensor.device(), + )), + Shape::scalar(), + DataType::Float32, + tensor.device(), + false, + ), + ), + DataType::Float64 => ( + sum_result, + Tensor::new( + Arc::new(TensorData::from_vec( + vec![num_elements], + DataType::Float64, + tensor.device(), + )), + Shape::scalar(), + DataType::Float64, + tensor.device(), + false, + ), + ), + DataType::Int32 => ( + sum_result.astype(DataType::Float32)?, + Tensor::new( + Arc::new(TensorData::from_vec( + vec![num_elements as f32], + DataType::Float32, + tensor.device(), + )), + Shape::scalar(), + DataType::Float32, + tensor.device(), + false, + ), + ), + DataType::Int64 => ( + sum_result.astype(DataType::Float64)?, + Tensor::new( + Arc::new(TensorData::from_vec( + vec![num_elements], + DataType::Float64, + tensor.device(), + )), + Shape::scalar(), + DataType::Float64, + tensor.device(), + false, + ), + ), + DataType::Bool => { + return Err(MinitensorError::invalid_operation( + "Mean not supported for boolean tensors", + )); + } + }; + + crate::operations::arithmetic::div(&sum_tensor, &divisor) +} + +/// NaN-aware mean reduction along specified dimensions +pub fn nanmean(tensor: &Tensor, dim: Option>, keepdim: bool) -> Result { + if !tensor.dtype().is_float() { + return mean(tensor, dim, keepdim); + } + + let dim = normalize_reduction_dims(dim, tensor.ndim())?; + let dims_clone = dim.clone(); + let needs_mask = + tensor.requires_grad() || dim.as_ref().map(|dims| !dims.is_empty()).unwrap_or(false); + let mask = if needs_mask { + Some(non_nan_mask(tensor)?) + } else { + None + }; + + if let Some(dims) = &dim + && dims.is_empty() + { + return Ok(tensor.clone()); + } + + let (sum, count) = match dim { + None => { + let result_shape = if keepdim { + Shape::new(vec![1; tensor.ndim()]) + } else { + Shape::scalar() + }; + let mut sum_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); + let mut count_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => nanmean_all_f32(tensor, &mut sum_data, &mut count_data)?, + DataType::Float64 => nanmean_all_f64(tensor, &mut sum_data, &mut count_data)?, + _ => unreachable!("nanmean only supports floating point tensors"), + } + + ( + Tensor::new( + Arc::new(sum_data), + result_shape.clone(), + tensor.dtype(), + tensor.device(), + false, + ), + Tensor::new( + Arc::new(count_data), + result_shape, + tensor.dtype(), + tensor.device(), + false, + ), + ) + } + Some(dims) => { + let mask = mask.as_ref().ok_or_else(|| { + MinitensorError::internal_error("nanmean expected mask for count computation") + })?; + let mut sum = tensor.clone(); + let mut count = mask.astype(tensor.dtype())?; + + if keepdim { + for &d in &dims { + sum = nansum_along_dim(&sum, d, true)?; + count = sum_along_dim(&count, d, true)?; + } + } else { + for &d in dims.iter().rev() { + sum = nansum_along_dim(&sum, d, false)?; + count = sum_along_dim(&count, d, false)?; + } + } + (sum, count) + } + }; + + let result = nanmean_from_sum_count(&sum, &count, tensor.requires_grad())?; + + if result.requires_grad() { + let mask = mask.ok_or_else(|| { + MinitensorError::internal_error("nanmean expected mask for gradient computation") + })?; + let grad_fn = Arc::new(NanMeanBackward { + input_id: tensor.id(), + input_shape: tensor.shape().dims().to_vec(), + dims: dims_clone, + keepdim, + mask, + count, + }); + let mut result_with_grad = result; + result_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&result_with_grad, Some(grad_fn))?; + Ok(result_with_grad) + } else { + Ok(result) + } +} + +/// Logical all reduction along specified dimension +pub fn all(tensor: &Tensor, dim: Option, keepdim: bool) -> Result { + match dim { + None => all_all(tensor, keepdim), + Some(d) => { + let d = normalize_dim(d, tensor.ndim())?; + all_along_dim(tensor, d, keepdim) + } + } +} + +/// Logical any reduction along specified dimension +pub fn any(tensor: &Tensor, dim: Option, keepdim: bool) -> Result { + match dim { + None => any_all(tensor, keepdim), + Some(d) => { + let d = normalize_dim(d, tensor.ndim())?; + any_along_dim(tensor, d, keepdim) + } + } +} + +fn all_all(tensor: &Tensor, keepdim: bool) -> Result { + let result_shape = if keepdim { + Shape::new(vec![1; tensor.ndim()]) + } else { + Shape::scalar() + }; + let mut result_data = TensorData::zeros_on_device(1, DataType::Bool, tensor.device()); + let out_slice = result_data + .as_bool_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + let all_true = match tensor.dtype() { + DataType::Float32 => tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected f32 data"))? + .par_iter() + .all(|&x| x != 0.0), + DataType::Float64 => tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected f64 data"))? + .par_iter() + .all(|&x| x != 0.0), + DataType::Int32 => tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected i32 data"))? + .par_iter() + .all(|&x| x != 0), + DataType::Int64 => tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected i64 data"))? + .par_iter() + .all(|&x| x != 0), + DataType::Bool => tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected bool data"))? + .par_iter() + .all(|&x| x), + }; + out_slice[0] = all_true; + Ok(Tensor::new( + Arc::new(result_data), + result_shape, + DataType::Bool, + tensor.device(), + false, + )) +} + +fn any_all(tensor: &Tensor, keepdim: bool) -> Result { + let result_shape = if keepdim { + Shape::new(vec![1; tensor.ndim()]) + } else { + Shape::scalar() + }; + let mut result_data = TensorData::zeros_on_device(1, DataType::Bool, tensor.device()); + let out_slice = result_data + .as_bool_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + let any_true = match tensor.dtype() { + DataType::Float32 => tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected f32 data"))? + .par_iter() + .any(|&x| x != 0.0), + DataType::Float64 => tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected f64 data"))? + .par_iter() + .any(|&x| x != 0.0), + DataType::Int32 => tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected i32 data"))? + .par_iter() + .any(|&x| x != 0), + DataType::Int64 => tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected i64 data"))? + .par_iter() + .any(|&x| x != 0), + DataType::Bool => tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected bool data"))? + .par_iter() + .any(|&x| x), + }; + out_slice[0] = any_true; + Ok(Tensor::new( + Arc::new(result_data), + result_shape, + DataType::Bool, + tensor.device(), + false, + )) +} + +fn all_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { + if dim >= tensor.ndim() { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + + let input_shape = tensor.shape().dims(); + let mut output_shape = input_shape.to_vec(); + if keepdim { + output_shape[dim] = 1; + } else { + output_shape.remove(dim); + } + let output_shape_obj = Shape::new(output_shape.clone()); + let mut result_data = + TensorData::zeros_on_device(output_shape_obj.numel(), DataType::Bool, tensor.device()); + + let dim_size = input_shape[dim]; + let _outer = input_shape[..dim].iter().product::(); + let inner = input_shape[dim + 1..].iter().product::(); + let outer_stride = dim_size * inner; + + match tensor.dtype() { + DataType::Float32 => { + let input = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let output = result_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice") + })?; + output.par_iter_mut().enumerate().for_each(|(idx, out)| { + let o = idx / inner; + let r = idx % inner; + let mut val = true; + for d in 0..dim_size { + let in_idx = o * outer_stride + d * inner + r; + if input[in_idx] == 0.0 { + val = false; + break; + } + } + *out = val; + }); + } + DataType::Float64 => { + let input = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let output = result_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice") + })?; + output.par_iter_mut().enumerate().for_each(|(idx, out)| { + let o = idx / inner; + let r = idx % inner; + let mut val = true; + for d in 0..dim_size { + let in_idx = o * outer_stride + d * inner + r; + if input[in_idx] == 0.0 { + val = false; + break; + } + } + *out = val; + }); + } + DataType::Int32 => { + let input = tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; + let output = result_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice") + })?; + output.par_iter_mut().enumerate().for_each(|(idx, out)| { + let o = idx / inner; + let r = idx % inner; + let mut val = true; + for d in 0..dim_size { + let in_idx = o * outer_stride + d * inner + r; + if input[in_idx] == 0 { + val = false; + break; + } + } + *out = val; + }); + } + DataType::Int64 => { + let input = tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; + let output = result_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice") + })?; + output.par_iter_mut().enumerate().for_each(|(idx, out)| { + let o = idx / inner; + let r = idx % inner; + let mut val = true; + for d in 0..dim_size { + let in_idx = o * outer_stride + d * inner + r; + if input[in_idx] == 0 { + val = false; + break; + } + } + *out = val; + }); + } + DataType::Bool => { + let input = tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + let output = result_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice") + })?; + output.par_iter_mut().enumerate().for_each(|(idx, out)| { + let o = idx / inner; + let r = idx % inner; + let mut val = true; + for d in 0..dim_size { + let in_idx = o * outer_stride + d * inner + r; + if !input[in_idx] { + val = false; + break; + } + } + *out = val; + }); + } + } + + Ok(Tensor::new( + Arc::new(result_data), + output_shape_obj, + DataType::Bool, + tensor.device(), + false, + )) +} diff --git a/engine/src/operations/reduction/minmax_indices.rs b/engine/src/operations/reduction/minmax_indices.rs index 014167ba..48690075 100644 --- a/engine/src/operations/reduction/minmax_indices.rs +++ b/engine/src/operations/reduction/minmax_indices.rs @@ -1,415 +1,422 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -fn min_along_dim_with_indices( - tensor: &Tensor, - dim: usize, - keepdim: bool, -) -> Result<(Tensor, Tensor)> { - let layout = reduction_layout(tensor, dim, keepdim)?; - let mut values_data = - TensorData::zeros_on_device(layout.output_shape.numel(), tensor.dtype(), tensor.device()); - let mut indices_data = TensorData::zeros_on_device( - layout.output_shape.numel(), - DataType::Int64, - tensor.device(), - ); - - let indices = indices_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - match tensor.dtype() { - DataType::Float32 => { - let input = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let values = values_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut min_val = f32::INFINITY; - let mut min_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if val.is_nan() { - min_val = f32::NAN; - min_idx = d; - break; - } - if val < min_val { - min_val = val; - min_idx = d; - } - } - let out_idx = o * layout.inner + r; - values[out_idx] = min_val; - indices[out_idx] = min_idx as i64; - } - } - } - DataType::Float64 => { - let input = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let values = values_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut min_val = f64::INFINITY; - let mut min_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if val.is_nan() { - min_val = f64::NAN; - min_idx = d; - break; - } - if val < min_val { - min_val = val; - min_idx = d; - } - } - let out_idx = o * layout.inner + r; - values[out_idx] = min_val; - indices[out_idx] = min_idx as i64; - } - } - } - DataType::Int32 => { - let input = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - let values = values_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i32 slice") - })?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut min_val = i32::MAX; - let mut min_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if val < min_val { - min_val = val; - min_idx = d; - } - } - let out_idx = o * layout.inner + r; - values[out_idx] = min_val; - indices[out_idx] = min_idx as i64; - } - } - } - DataType::Int64 => { - let input = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - let values = values_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut min_val = i64::MAX; - let mut min_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if val < min_val { - min_val = val; - min_idx = d; - } - } - let out_idx = o * layout.inner + r; - values[out_idx] = min_val; - indices[out_idx] = min_idx as i64; - } - } - } - DataType::Bool => { - let input = tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - let values = values_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice") - })?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut min_val = true; - let mut min_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - if !input[idx] { - min_val = false; - min_idx = d; - break; - } - } - let out_idx = o * layout.inner + r; - values[out_idx] = min_val; - indices[out_idx] = min_idx as i64; - } - } - } - } - - Ok(( - Tensor::new( - Arc::new(values_data), - layout.output_shape.clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ), - Tensor::new( - Arc::new(indices_data), - layout.output_shape, - DataType::Int64, - tensor.device(), - false, - ), - )) -} - -fn nanmin_along_dim_with_indices( - tensor: &Tensor, - dim: usize, - keepdim: bool, -) -> Result<(Tensor, Tensor)> { - let layout = reduction_layout(tensor, dim, keepdim)?; - let mut values_data = - TensorData::zeros_on_device(layout.output_shape.numel(), tensor.dtype(), tensor.device()); - let mut indices_data = TensorData::zeros_on_device( - layout.output_shape.numel(), - DataType::Int64, - tensor.device(), - ); - - let indices = indices_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - match tensor.dtype() { - DataType::Float32 => { - let input = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let values = values_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut min_val = f32::NAN; - let mut min_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if !val.is_nan() && (min_val.is_nan() || val < min_val) { - min_val = val; - min_idx = d; - } - } - let out_idx = o * layout.inner + r; - values[out_idx] = min_val; - indices[out_idx] = min_idx as i64; - } - } - } - DataType::Float64 => { - let input = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let values = values_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut min_val = f64::NAN; - let mut min_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if !val.is_nan() && (min_val.is_nan() || val < min_val) { - min_val = val; - min_idx = d; - } - } - let out_idx = o * layout.inner + r; - values[out_idx] = min_val; - indices[out_idx] = min_idx as i64; - } - } - } - _ => { - return Err(MinitensorError::invalid_operation( - "nanmin only supports floating point tensors", - )); - } - } - - Ok(( - Tensor::new( - Arc::new(values_data), - layout.output_shape.clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ), - Tensor::new( - Arc::new(indices_data), - layout.output_shape, - DataType::Int64, - tensor.device(), - false, - ), - )) -} - -fn argmax_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { - let layout = reduction_layout(tensor, dim, keepdim)?; - let mut result_data = TensorData::zeros_on_device( - layout.output_shape.numel(), - DataType::Int64, - tensor.device(), - ); - - let output = result_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - match tensor.dtype() { - DataType::Float32 => { - let input = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut max_val = f32::NEG_INFINITY; - let mut max_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if val.is_nan() { - max_idx = d; - break; - } - if val > max_val { - max_val = val; - max_idx = d; - } - } - output[o * layout.inner + r] = max_idx as i64; - } - } - } - DataType::Float64 => { - let input = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut max_val = f64::NEG_INFINITY; - let mut max_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if val.is_nan() { - max_idx = d; - break; - } - if val > max_val { - max_val = val; - max_idx = d; - } - } - output[o * layout.inner + r] = max_idx as i64; - } - } - } - DataType::Int32 => { - let input = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut max_val = i32::MIN; - let mut max_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if val > max_val { - max_val = val; - max_idx = d; - } - } - output[o * layout.inner + r] = max_idx as i64; - } - } - } - DataType::Int64 => { - let input = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut max_val = i64::MIN; - let mut max_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if val > max_val { - max_val = val; - max_idx = d; - } - } - output[o * layout.inner + r] = max_idx as i64; - } - } - } - DataType::Bool => { - let input = tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut max_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - if input[idx] { - max_idx = d; - break; - } - } - output[o * layout.inner + r] = max_idx as i64; - } - } - } - } - - Ok(Tensor::new( - Arc::new(result_data), - layout.output_shape, - DataType::Int64, - tensor.device(), - false, - )) -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::{ + error::{MinitensorError, Result}, + tensor::{DataType, Tensor, TensorData}, +}; +use std::sync::Arc; + +pub(crate) fn min_along_dim_with_indices( + tensor: &Tensor, + dim: usize, + keepdim: bool, +) -> Result<(Tensor, Tensor)> { + let layout = reduction_layout(tensor, dim, keepdim)?; + let mut values_data = + TensorData::zeros_on_device(layout.output_shape.numel(), tensor.dtype(), tensor.device()); + let mut indices_data = TensorData::zeros_on_device( + layout.output_shape.numel(), + DataType::Int64, + tensor.device(), + ); + + let indices = indices_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + match tensor.dtype() { + DataType::Float32 => { + let input = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let values = values_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut min_val = f32::INFINITY; + let mut min_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if val.is_nan() { + min_val = f32::NAN; + min_idx = d; + break; + } + if val < min_val { + min_val = val; + min_idx = d; + } + } + let out_idx = o * layout.inner + r; + values[out_idx] = min_val; + indices[out_idx] = min_idx as i64; + } + } + } + DataType::Float64 => { + let input = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let values = values_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut min_val = f64::INFINITY; + let mut min_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if val.is_nan() { + min_val = f64::NAN; + min_idx = d; + break; + } + if val < min_val { + min_val = val; + min_idx = d; + } + } + let out_idx = o * layout.inner + r; + values[out_idx] = min_val; + indices[out_idx] = min_idx as i64; + } + } + } + DataType::Int32 => { + let input = tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; + let values = values_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i32 slice") + })?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut min_val = i32::MAX; + let mut min_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if val < min_val { + min_val = val; + min_idx = d; + } + } + let out_idx = o * layout.inner + r; + values[out_idx] = min_val; + indices[out_idx] = min_idx as i64; + } + } + } + DataType::Int64 => { + let input = tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; + let values = values_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut min_val = i64::MAX; + let mut min_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if val < min_val { + min_val = val; + min_idx = d; + } + } + let out_idx = o * layout.inner + r; + values[out_idx] = min_val; + indices[out_idx] = min_idx as i64; + } + } + } + DataType::Bool => { + let input = tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + let values = values_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice") + })?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut min_val = true; + let mut min_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + if !input[idx] { + min_val = false; + min_idx = d; + break; + } + } + let out_idx = o * layout.inner + r; + values[out_idx] = min_val; + indices[out_idx] = min_idx as i64; + } + } + } + } + + Ok(( + Tensor::new( + Arc::new(values_data), + layout.output_shape.clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ), + Tensor::new( + Arc::new(indices_data), + layout.output_shape, + DataType::Int64, + tensor.device(), + false, + ), + )) +} + +pub(crate) fn nanmin_along_dim_with_indices( + tensor: &Tensor, + dim: usize, + keepdim: bool, +) -> Result<(Tensor, Tensor)> { + let layout = reduction_layout(tensor, dim, keepdim)?; + let mut values_data = + TensorData::zeros_on_device(layout.output_shape.numel(), tensor.dtype(), tensor.device()); + let mut indices_data = TensorData::zeros_on_device( + layout.output_shape.numel(), + DataType::Int64, + tensor.device(), + ); + + let indices = indices_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + match tensor.dtype() { + DataType::Float32 => { + let input = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let values = values_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut min_val = f32::NAN; + let mut min_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if !val.is_nan() && (min_val.is_nan() || val < min_val) { + min_val = val; + min_idx = d; + } + } + let out_idx = o * layout.inner + r; + values[out_idx] = min_val; + indices[out_idx] = min_idx as i64; + } + } + } + DataType::Float64 => { + let input = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let values = values_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut min_val = f64::NAN; + let mut min_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if !val.is_nan() && (min_val.is_nan() || val < min_val) { + min_val = val; + min_idx = d; + } + } + let out_idx = o * layout.inner + r; + values[out_idx] = min_val; + indices[out_idx] = min_idx as i64; + } + } + } + _ => { + return Err(MinitensorError::invalid_operation( + "nanmin only supports floating point tensors", + )); + } + } + + Ok(( + Tensor::new( + Arc::new(values_data), + layout.output_shape.clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ), + Tensor::new( + Arc::new(indices_data), + layout.output_shape, + DataType::Int64, + tensor.device(), + false, + ), + )) +} + +pub(crate) fn argmax_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { + let layout = reduction_layout(tensor, dim, keepdim)?; + let mut result_data = TensorData::zeros_on_device( + layout.output_shape.numel(), + DataType::Int64, + tensor.device(), + ); + + let output = result_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + match tensor.dtype() { + DataType::Float32 => { + let input = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut max_val = f32::NEG_INFINITY; + let mut max_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if val.is_nan() { + max_idx = d; + break; + } + if val > max_val { + max_val = val; + max_idx = d; + } + } + output[o * layout.inner + r] = max_idx as i64; + } + } + } + DataType::Float64 => { + let input = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut max_val = f64::NEG_INFINITY; + let mut max_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if val.is_nan() { + max_idx = d; + break; + } + if val > max_val { + max_val = val; + max_idx = d; + } + } + output[o * layout.inner + r] = max_idx as i64; + } + } + } + DataType::Int32 => { + let input = tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut max_val = i32::MIN; + let mut max_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if val > max_val { + max_val = val; + max_idx = d; + } + } + output[o * layout.inner + r] = max_idx as i64; + } + } + } + DataType::Int64 => { + let input = tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut max_val = i64::MIN; + let mut max_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if val > max_val { + max_val = val; + max_idx = d; + } + } + output[o * layout.inner + r] = max_idx as i64; + } + } + } + DataType::Bool => { + let input = tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut max_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + if input[idx] { + max_idx = d; + break; + } + } + output[o * layout.inner + r] = max_idx as i64; + } + } + } + } + + Ok(Tensor::new( + Arc::new(result_data), + layout.output_shape, + DataType::Int64, + tensor.device(), + false, + )) +} diff --git a/engine/src/operations/reduction/nan_minmax.rs b/engine/src/operations/reduction/nan_minmax.rs index 822c935f..2b02b2b3 100644 --- a/engine/src/operations/reduction/nan_minmax.rs +++ b/engine/src/operations/reduction/nan_minmax.rs @@ -1,952 +1,963 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -fn nanmax_all_f64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - - let (max_val, found) = data - .par_iter() - .map(|&v| { - if v.is_nan() { - (f64::NEG_INFINITY, false) - } else { - (v, true) - } - }) - .reduce( - || (f64::NEG_INFINITY, false), - |(a_val, a_found), (b_val, b_found)| match (a_found, b_found) { - (true, true) => (a_val.max(b_val), true), - (true, false) => (a_val, true), - (false, true) => (b_val, true), - (false, false) => (f64::NEG_INFINITY, false), - }, - ); - - let result_slice = result_data - .as_f64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; - - result_slice[0] = if found { max_val } else { f64::NAN }; - Ok(()) -} - -fn nanmin_all_f32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - - let (min_val, found) = data - .par_iter() - .map(|&v| { - if v.is_nan() { - (f32::INFINITY, false) - } else { - (v, true) - } - }) - .reduce( - || (f32::INFINITY, false), - |(a_val, a_found), (b_val, b_found)| match (a_found, b_found) { - (true, true) => (a_val.min(b_val), true), - (true, false) => (a_val, true), - (false, true) => (b_val, true), - (false, false) => (f32::INFINITY, false), - }, - ); - - let result_slice = result_data - .as_f32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; - - result_slice[0] = if found { min_val } else { f32::NAN }; - Ok(()) -} - -fn nanmin_all_f64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - - let (min_val, found) = data - .par_iter() - .map(|&v| { - if v.is_nan() { - (f64::INFINITY, false) - } else { - (v, true) - } - }) - .reduce( - || (f64::INFINITY, false), - |(a_val, a_found), (b_val, b_found)| match (a_found, b_found) { - (true, true) => (a_val.min(b_val), true), - (true, false) => (a_val, true), - (false, true) => (b_val, true), - (false, false) => (f64::INFINITY, false), - }, - ); - - let result_slice = result_data - .as_f64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; - - result_slice[0] = if found { min_val } else { f64::NAN }; - Ok(()) -} - -// Placeholder implementations for argmax/argmin -fn argmax_all_f32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - - let (argmax_idx, _) = data.par_iter().enumerate().map(|(i, &v)| (i, v)).reduce( - || (0, f32::NEG_INFINITY), - |(i1, v1), (i2, v2)| match (v1.is_nan(), v2.is_nan()) { - (true, true) => { - if i1 <= i2 { - (i1, v1) - } else { - (i2, v2) - } - } - (true, false) => (i1, v1), - (false, true) => (i2, v2), - (false, false) => { - if v1 > v2 { - (i1, v1) - } else if v2 > v1 { - (i2, v2) - } else if i1 <= i2 { - (i1, v1) - } else { - (i2, v2) - } - } - }, - ); - - let result_slice = result_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - result_slice[0] = argmax_idx as i64; - Ok(()) -} - -fn argmax_all_f64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - - let (argmax_idx, _) = data.par_iter().enumerate().map(|(i, &v)| (i, v)).reduce( - || (0, f64::NEG_INFINITY), - |(i1, v1), (i2, v2)| match (v1.is_nan(), v2.is_nan()) { - (true, true) => { - if i1 <= i2 { - (i1, v1) - } else { - (i2, v2) - } - } - (true, false) => (i1, v1), - (false, true) => (i2, v2), - (false, false) => { - if v1 > v2 { - (i1, v1) - } else if v2 > v1 { - (i2, v2) - } else if i1 <= i2 { - (i1, v1) - } else { - (i2, v2) - } - } - }, - ); - - let result_slice = result_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - result_slice[0] = argmax_idx as i64; - Ok(()) -} - -fn argmax_all_i32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - - let (argmax_idx, _) = data.par_iter().enumerate().map(|(i, &v)| (i, v)).reduce( - || (0, i32::MIN), - |(i1, v1), (i2, v2)| { - if v1 >= v2 { (i1, v1) } else { (i2, v2) } - }, - ); - - let result_slice = result_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - result_slice[0] = argmax_idx as i64; - Ok(()) -} - -fn argmax_all_i64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - - let (argmax_idx, _) = data.par_iter().enumerate().map(|(i, &v)| (i, v)).reduce( - || (0, i64::MIN), - |(i1, v1), (i2, v2)| { - if v1 >= v2 { (i1, v1) } else { (i2, v2) } - }, - ); - - let result_slice = result_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - result_slice[0] = argmax_idx as i64; - Ok(()) -} - -fn argmax_all_bool(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - - let argmax_idx = data.iter().position(|&x| x).unwrap_or(0); - - let result_slice = result_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - result_slice[0] = argmax_idx as i64; - Ok(()) -} - -// Similar implementations for argmin -fn argmin_all_f32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - - let (argmin_idx, _) = data.par_iter().enumerate().map(|(i, &v)| (i, v)).reduce( - || (0, f32::INFINITY), - |(i1, v1), (i2, v2)| match (v1.is_nan(), v2.is_nan()) { - (true, true) => { - if i1 <= i2 { - (i1, v1) - } else { - (i2, v2) - } - } - (true, false) => (i1, v1), - (false, true) => (i2, v2), - (false, false) => { - if v1 < v2 { - (i1, v1) - } else if v2 < v1 { - (i2, v2) - } else if i1 <= i2 { - (i1, v1) - } else { - (i2, v2) - } - } - }, - ); - - let result_slice = result_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - result_slice[0] = argmin_idx as i64; - Ok(()) -} - -fn argmin_all_f64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - - let (argmin_idx, _) = data.par_iter().enumerate().map(|(i, &v)| (i, v)).reduce( - || (0, f64::INFINITY), - |(i1, v1), (i2, v2)| match (v1.is_nan(), v2.is_nan()) { - (true, true) => { - if i1 <= i2 { - (i1, v1) - } else { - (i2, v2) - } - } - (true, false) => (i1, v1), - (false, true) => (i2, v2), - (false, false) => { - if v1 < v2 { - (i1, v1) - } else if v2 < v1 { - (i2, v2) - } else if i1 <= i2 { - (i1, v1) - } else { - (i2, v2) - } - } - }, - ); - - let result_slice = result_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - result_slice[0] = argmin_idx as i64; - Ok(()) -} - -fn argmin_all_i32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - - let (argmin_idx, _) = data.par_iter().enumerate().map(|(i, &v)| (i, v)).reduce( - || (0, i32::MAX), - |(i1, v1), (i2, v2)| { - if v1 <= v2 { (i1, v1) } else { (i2, v2) } - }, - ); - - let result_slice = result_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - result_slice[0] = argmin_idx as i64; - Ok(()) -} - -fn argmin_all_i64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - - let (argmin_idx, _) = data.par_iter().enumerate().map(|(i, &v)| (i, v)).reduce( - || (0, i64::MAX), - |(i1, v1), (i2, v2)| { - if v1 <= v2 { (i1, v1) } else { (i2, v2) } - }, - ); - - let result_slice = result_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - result_slice[0] = argmin_idx as i64; - Ok(()) -} - -fn argmin_all_bool(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - - let argmin_idx = data.par_iter().position_first(|&x| !x).unwrap_or(0); - - let result_slice = result_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - result_slice[0] = argmin_idx as i64; - Ok(()) -} - -struct DimReductionLayout { - output_shape: Shape, - dim_size: usize, - outer: usize, - inner: usize, - outer_stride: usize, -} - -fn reduction_layout(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { - if dim >= tensor.ndim() { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - - let input_shape = tensor.shape().dims(); - let mut output_shape = input_shape.to_vec(); - if keepdim { - output_shape[dim] = 1; - } else { - output_shape.remove(dim); - } - let dim_size = input_shape[dim]; - let outer = input_shape[..dim].iter().product::(); - let inner = input_shape[dim + 1..].iter().product::(); - let outer_stride = dim_size * inner; - - Ok(DimReductionLayout { - output_shape: Shape::new(output_shape), - dim_size, - outer, - inner, - outer_stride, - }) -} - -// Placeholder implementations for dimensional operations -fn max_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { - let layout = reduction_layout(tensor, dim, keepdim)?; - let mut result_data = - TensorData::zeros_on_device(layout.output_shape.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => { - let input = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let output = result_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?; - - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut max_val = f32::NEG_INFINITY; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if val.is_nan() { - max_val = f32::NAN; - break; - } - max_val = max_val.max(val); - } - output[o * layout.inner + r] = max_val; - } - } - } - DataType::Float64 => { - let input = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let output = result_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?; - - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut max_val = f64::NEG_INFINITY; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if val.is_nan() { - max_val = f64::NAN; - break; - } - max_val = max_val.max(val); - } - output[o * layout.inner + r] = max_val; - } - } - } - DataType::Int32 => { - let input = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - let output = result_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i32 slice") - })?; - - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut max_val = i32::MIN; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - max_val = max_val.max(input[idx]); - } - output[o * layout.inner + r] = max_val; - } - } - } - DataType::Int64 => { - let input = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - let output = result_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut max_val = i64::MIN; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - max_val = max_val.max(input[idx]); - } - output[o * layout.inner + r] = max_val; - } - } - } - DataType::Bool => { - let input = tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - let output = result_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice") - })?; - - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut max_val = false; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - max_val |= input[idx]; - if max_val { - break; - } - } - output[o * layout.inner + r] = max_val; - } - } - } - } - - Ok(Tensor::new( - Arc::new(result_data), - layout.output_shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )) -} - -fn min_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { - let layout = reduction_layout(tensor, dim, keepdim)?; - let mut result_data = - TensorData::zeros_on_device(layout.output_shape.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => { - let input = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let output = result_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?; - - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut min_val = f32::INFINITY; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if val.is_nan() { - min_val = f32::NAN; - break; - } - min_val = min_val.min(val); - } - output[o * layout.inner + r] = min_val; - } - } - } - DataType::Float64 => { - let input = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let output = result_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?; - - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut min_val = f64::INFINITY; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if val.is_nan() { - min_val = f64::NAN; - break; - } - min_val = min_val.min(val); - } - output[o * layout.inner + r] = min_val; - } - } - } - DataType::Int32 => { - let input = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - let output = result_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i32 slice") - })?; - - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut min_val = i32::MAX; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - min_val = min_val.min(input[idx]); - } - output[o * layout.inner + r] = min_val; - } - } - } - DataType::Int64 => { - let input = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - let output = result_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut min_val = i64::MAX; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - min_val = min_val.min(input[idx]); - } - output[o * layout.inner + r] = min_val; - } - } - } - DataType::Bool => { - let input = tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - let output = result_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice") - })?; - - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut min_val = true; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - min_val &= input[idx]; - if !min_val { - break; - } - } - output[o * layout.inner + r] = min_val; - } - } - } - } - - Ok(Tensor::new( - Arc::new(result_data), - layout.output_shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )) -} - -fn max_along_dim_with_indices( - tensor: &Tensor, - dim: usize, - keepdim: bool, -) -> Result<(Tensor, Tensor)> { - let layout = reduction_layout(tensor, dim, keepdim)?; - let mut values_data = - TensorData::zeros_on_device(layout.output_shape.numel(), tensor.dtype(), tensor.device()); - let mut indices_data = TensorData::zeros_on_device( - layout.output_shape.numel(), - DataType::Int64, - tensor.device(), - ); - - let indices = indices_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - match tensor.dtype() { - DataType::Float32 => { - let input = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let values = values_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut max_val = f32::NEG_INFINITY; - let mut max_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if val.is_nan() { - max_val = f32::NAN; - max_idx = d; - break; - } - if val > max_val { - max_val = val; - max_idx = d; - } - } - let out_idx = o * layout.inner + r; - values[out_idx] = max_val; - indices[out_idx] = max_idx as i64; - } - } - } - DataType::Float64 => { - let input = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let values = values_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut max_val = f64::NEG_INFINITY; - let mut max_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if val.is_nan() { - max_val = f64::NAN; - max_idx = d; - break; - } - if val > max_val { - max_val = val; - max_idx = d; - } - } - let out_idx = o * layout.inner + r; - values[out_idx] = max_val; - indices[out_idx] = max_idx as i64; - } - } - } - DataType::Int32 => { - let input = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - let values = values_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i32 slice") - })?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut max_val = i32::MIN; - let mut max_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if val > max_val { - max_val = val; - max_idx = d; - } - } - let out_idx = o * layout.inner + r; - values[out_idx] = max_val; - indices[out_idx] = max_idx as i64; - } - } - } - DataType::Int64 => { - let input = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - let values = values_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut max_val = i64::MIN; - let mut max_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if val > max_val { - max_val = val; - max_idx = d; - } - } - let out_idx = o * layout.inner + r; - values[out_idx] = max_val; - indices[out_idx] = max_idx as i64; - } - } - } - DataType::Bool => { - let input = tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - let values = values_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice") - })?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut max_val = false; - let mut max_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - if input[idx] { - max_val = true; - max_idx = d; - break; - } - } - let out_idx = o * layout.inner + r; - values[out_idx] = max_val; - indices[out_idx] = max_idx as i64; - } - } - } - } - - Ok(( - Tensor::new( - Arc::new(values_data), - layout.output_shape.clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ), - Tensor::new( - Arc::new(indices_data), - layout.output_shape, - DataType::Int64, - tensor.device(), - false, - ), - )) -} - -fn nanmax_along_dim_with_indices( - tensor: &Tensor, - dim: usize, - keepdim: bool, -) -> Result<(Tensor, Tensor)> { - let layout = reduction_layout(tensor, dim, keepdim)?; - let mut values_data = - TensorData::zeros_on_device(layout.output_shape.numel(), tensor.dtype(), tensor.device()); - let mut indices_data = TensorData::zeros_on_device( - layout.output_shape.numel(), - DataType::Int64, - tensor.device(), - ); - - let indices = indices_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - match tensor.dtype() { - DataType::Float32 => { - let input = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let values = values_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut max_val = f32::NAN; - let mut max_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if !val.is_nan() && (max_val.is_nan() || val > max_val) { - max_val = val; - max_idx = d; - } - } - let out_idx = o * layout.inner + r; - values[out_idx] = max_val; - indices[out_idx] = max_idx as i64; - } - } - } - DataType::Float64 => { - let input = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let values = values_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?; - for o in 0..layout.outer { - for r in 0..layout.inner { - let mut max_val = f64::NAN; - let mut max_idx = 0usize; - for d in 0..layout.dim_size { - let idx = o * layout.outer_stride + d * layout.inner + r; - let val = input[idx]; - if !val.is_nan() && (max_val.is_nan() || val > max_val) { - max_val = val; - max_idx = d; - } - } - let out_idx = o * layout.inner + r; - values[out_idx] = max_val; - indices[out_idx] = max_idx as i64; - } - } - } - _ => { - return Err(MinitensorError::invalid_operation( - "nanmax only supports floating point tensors", - )); - } - } - - Ok(( - Tensor::new( - Arc::new(values_data), - layout.output_shape.clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ), - Tensor::new( - Arc::new(indices_data), - layout.output_shape, - DataType::Int64, - tensor.device(), - false, - ), - )) -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use crate::{ + error::{MinitensorError, Result}, + tensor::{DataType, Shape, Tensor, TensorData}, +}; +use rayon::prelude::*; +use std::sync::Arc; + +pub(crate) fn nanmax_all_f64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + + let (max_val, found) = data + .par_iter() + .map(|&v| { + if v.is_nan() { + (f64::NEG_INFINITY, false) + } else { + (v, true) + } + }) + .reduce( + || (f64::NEG_INFINITY, false), + |(a_val, a_found), (b_val, b_found)| match (a_found, b_found) { + (true, true) => (a_val.max(b_val), true), + (true, false) => (a_val, true), + (false, true) => (b_val, true), + (false, false) => (f64::NEG_INFINITY, false), + }, + ); + + let result_slice = result_data + .as_f64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; + + result_slice[0] = if found { max_val } else { f64::NAN }; + Ok(()) +} + +pub(crate) fn nanmin_all_f32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + + let (min_val, found) = data + .par_iter() + .map(|&v| { + if v.is_nan() { + (f32::INFINITY, false) + } else { + (v, true) + } + }) + .reduce( + || (f32::INFINITY, false), + |(a_val, a_found), (b_val, b_found)| match (a_found, b_found) { + (true, true) => (a_val.min(b_val), true), + (true, false) => (a_val, true), + (false, true) => (b_val, true), + (false, false) => (f32::INFINITY, false), + }, + ); + + let result_slice = result_data + .as_f32_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; + + result_slice[0] = if found { min_val } else { f32::NAN }; + Ok(()) +} + +pub(crate) fn nanmin_all_f64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + + let (min_val, found) = data + .par_iter() + .map(|&v| { + if v.is_nan() { + (f64::INFINITY, false) + } else { + (v, true) + } + }) + .reduce( + || (f64::INFINITY, false), + |(a_val, a_found), (b_val, b_found)| match (a_found, b_found) { + (true, true) => (a_val.min(b_val), true), + (true, false) => (a_val, true), + (false, true) => (b_val, true), + (false, false) => (f64::INFINITY, false), + }, + ); + + let result_slice = result_data + .as_f64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; + + result_slice[0] = if found { min_val } else { f64::NAN }; + Ok(()) +} + +// Placeholder implementations for argmax/argmin +pub(crate) fn argmax_all_f32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + + let (argmax_idx, _) = data.par_iter().enumerate().map(|(i, &v)| (i, v)).reduce( + || (0, f32::NEG_INFINITY), + |(i1, v1), (i2, v2)| match (v1.is_nan(), v2.is_nan()) { + (true, true) => { + if i1 <= i2 { + (i1, v1) + } else { + (i2, v2) + } + } + (true, false) => (i1, v1), + (false, true) => (i2, v2), + (false, false) => { + if v1 > v2 { + (i1, v1) + } else if v2 > v1 { + (i2, v2) + } else if i1 <= i2 { + (i1, v1) + } else { + (i2, v2) + } + } + }, + ); + + let result_slice = result_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + result_slice[0] = argmax_idx as i64; + Ok(()) +} + +pub(crate) fn argmax_all_f64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + + let (argmax_idx, _) = data.par_iter().enumerate().map(|(i, &v)| (i, v)).reduce( + || (0, f64::NEG_INFINITY), + |(i1, v1), (i2, v2)| match (v1.is_nan(), v2.is_nan()) { + (true, true) => { + if i1 <= i2 { + (i1, v1) + } else { + (i2, v2) + } + } + (true, false) => (i1, v1), + (false, true) => (i2, v2), + (false, false) => { + if v1 > v2 { + (i1, v1) + } else if v2 > v1 { + (i2, v2) + } else if i1 <= i2 { + (i1, v1) + } else { + (i2, v2) + } + } + }, + ); + + let result_slice = result_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + result_slice[0] = argmax_idx as i64; + Ok(()) +} + +pub(crate) fn argmax_all_i32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; + + let (argmax_idx, _) = data.par_iter().enumerate().map(|(i, &v)| (i, v)).reduce( + || (0, i32::MIN), + |(i1, v1), (i2, v2)| { + if v1 >= v2 { (i1, v1) } else { (i2, v2) } + }, + ); + + let result_slice = result_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + result_slice[0] = argmax_idx as i64; + Ok(()) +} + +pub(crate) fn argmax_all_i64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; + + let (argmax_idx, _) = data.par_iter().enumerate().map(|(i, &v)| (i, v)).reduce( + || (0, i64::MIN), + |(i1, v1), (i2, v2)| { + if v1 >= v2 { (i1, v1) } else { (i2, v2) } + }, + ); + + let result_slice = result_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + result_slice[0] = argmax_idx as i64; + Ok(()) +} + +pub(crate) fn argmax_all_bool(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + + let argmax_idx = data.iter().position(|&x| x).unwrap_or(0); + + let result_slice = result_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + result_slice[0] = argmax_idx as i64; + Ok(()) +} + +// Similar implementations for argmin +pub(crate) fn argmin_all_f32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + + let (argmin_idx, _) = data.par_iter().enumerate().map(|(i, &v)| (i, v)).reduce( + || (0, f32::INFINITY), + |(i1, v1), (i2, v2)| match (v1.is_nan(), v2.is_nan()) { + (true, true) => { + if i1 <= i2 { + (i1, v1) + } else { + (i2, v2) + } + } + (true, false) => (i1, v1), + (false, true) => (i2, v2), + (false, false) => { + if v1 < v2 { + (i1, v1) + } else if v2 < v1 { + (i2, v2) + } else if i1 <= i2 { + (i1, v1) + } else { + (i2, v2) + } + } + }, + ); + + let result_slice = result_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + result_slice[0] = argmin_idx as i64; + Ok(()) +} + +pub(crate) fn argmin_all_f64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + + let (argmin_idx, _) = data.par_iter().enumerate().map(|(i, &v)| (i, v)).reduce( + || (0, f64::INFINITY), + |(i1, v1), (i2, v2)| match (v1.is_nan(), v2.is_nan()) { + (true, true) => { + if i1 <= i2 { + (i1, v1) + } else { + (i2, v2) + } + } + (true, false) => (i1, v1), + (false, true) => (i2, v2), + (false, false) => { + if v1 < v2 { + (i1, v1) + } else if v2 < v1 { + (i2, v2) + } else if i1 <= i2 { + (i1, v1) + } else { + (i2, v2) + } + } + }, + ); + + let result_slice = result_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + result_slice[0] = argmin_idx as i64; + Ok(()) +} + +pub(crate) fn argmin_all_i32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; + + let (argmin_idx, _) = data.par_iter().enumerate().map(|(i, &v)| (i, v)).reduce( + || (0, i32::MAX), + |(i1, v1), (i2, v2)| { + if v1 <= v2 { (i1, v1) } else { (i2, v2) } + }, + ); + + let result_slice = result_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + result_slice[0] = argmin_idx as i64; + Ok(()) +} + +pub(crate) fn argmin_all_i64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; + + let (argmin_idx, _) = data.par_iter().enumerate().map(|(i, &v)| (i, v)).reduce( + || (0, i64::MAX), + |(i1, v1), (i2, v2)| { + if v1 <= v2 { (i1, v1) } else { (i2, v2) } + }, + ); + + let result_slice = result_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + result_slice[0] = argmin_idx as i64; + Ok(()) +} + +pub(crate) fn argmin_all_bool(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + + let argmin_idx = data.par_iter().position_first(|&x| !x).unwrap_or(0); + + let result_slice = result_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + result_slice[0] = argmin_idx as i64; + Ok(()) +} + +pub(crate) struct DimReductionLayout { + pub(crate) output_shape: Shape, + pub(crate) dim_size: usize, + pub(crate) outer: usize, + pub(crate) inner: usize, + pub(crate) outer_stride: usize, +} + +pub(crate) fn reduction_layout( + tensor: &Tensor, + dim: usize, + keepdim: bool, +) -> Result { + if dim >= tensor.ndim() { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + + let input_shape = tensor.shape().dims(); + let mut output_shape = input_shape.to_vec(); + if keepdim { + output_shape[dim] = 1; + } else { + output_shape.remove(dim); + } + let dim_size = input_shape[dim]; + let outer = input_shape[..dim].iter().product::(); + let inner = input_shape[dim + 1..].iter().product::(); + let outer_stride = dim_size * inner; + + Ok(DimReductionLayout { + output_shape: Shape::new(output_shape), + dim_size, + outer, + inner, + outer_stride, + }) +} + +// Placeholder implementations for dimensional operations +pub(crate) fn max_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { + let layout = reduction_layout(tensor, dim, keepdim)?; + let mut result_data = + TensorData::zeros_on_device(layout.output_shape.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => { + let input = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let output = result_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?; + + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut max_val = f32::NEG_INFINITY; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if val.is_nan() { + max_val = f32::NAN; + break; + } + max_val = max_val.max(val); + } + output[o * layout.inner + r] = max_val; + } + } + } + DataType::Float64 => { + let input = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let output = result_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?; + + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut max_val = f64::NEG_INFINITY; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if val.is_nan() { + max_val = f64::NAN; + break; + } + max_val = max_val.max(val); + } + output[o * layout.inner + r] = max_val; + } + } + } + DataType::Int32 => { + let input = tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; + let output = result_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i32 slice") + })?; + + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut max_val = i32::MIN; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + max_val = max_val.max(input[idx]); + } + output[o * layout.inner + r] = max_val; + } + } + } + DataType::Int64 => { + let input = tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; + let output = result_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut max_val = i64::MIN; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + max_val = max_val.max(input[idx]); + } + output[o * layout.inner + r] = max_val; + } + } + } + DataType::Bool => { + let input = tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + let output = result_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice") + })?; + + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut max_val = false; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + max_val |= input[idx]; + if max_val { + break; + } + } + output[o * layout.inner + r] = max_val; + } + } + } + } + + Ok(Tensor::new( + Arc::new(result_data), + layout.output_shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )) +} + +pub(crate) fn min_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { + let layout = reduction_layout(tensor, dim, keepdim)?; + let mut result_data = + TensorData::zeros_on_device(layout.output_shape.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => { + let input = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let output = result_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?; + + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut min_val = f32::INFINITY; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if val.is_nan() { + min_val = f32::NAN; + break; + } + min_val = min_val.min(val); + } + output[o * layout.inner + r] = min_val; + } + } + } + DataType::Float64 => { + let input = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let output = result_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?; + + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut min_val = f64::INFINITY; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if val.is_nan() { + min_val = f64::NAN; + break; + } + min_val = min_val.min(val); + } + output[o * layout.inner + r] = min_val; + } + } + } + DataType::Int32 => { + let input = tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; + let output = result_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i32 slice") + })?; + + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut min_val = i32::MAX; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + min_val = min_val.min(input[idx]); + } + output[o * layout.inner + r] = min_val; + } + } + } + DataType::Int64 => { + let input = tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; + let output = result_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut min_val = i64::MAX; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + min_val = min_val.min(input[idx]); + } + output[o * layout.inner + r] = min_val; + } + } + } + DataType::Bool => { + let input = tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + let output = result_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice") + })?; + + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut min_val = true; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + min_val &= input[idx]; + if !min_val { + break; + } + } + output[o * layout.inner + r] = min_val; + } + } + } + } + + Ok(Tensor::new( + Arc::new(result_data), + layout.output_shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )) +} + +pub(crate) fn max_along_dim_with_indices( + tensor: &Tensor, + dim: usize, + keepdim: bool, +) -> Result<(Tensor, Tensor)> { + let layout = reduction_layout(tensor, dim, keepdim)?; + let mut values_data = + TensorData::zeros_on_device(layout.output_shape.numel(), tensor.dtype(), tensor.device()); + let mut indices_data = TensorData::zeros_on_device( + layout.output_shape.numel(), + DataType::Int64, + tensor.device(), + ); + + let indices = indices_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + match tensor.dtype() { + DataType::Float32 => { + let input = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let values = values_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut max_val = f32::NEG_INFINITY; + let mut max_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if val.is_nan() { + max_val = f32::NAN; + max_idx = d; + break; + } + if val > max_val { + max_val = val; + max_idx = d; + } + } + let out_idx = o * layout.inner + r; + values[out_idx] = max_val; + indices[out_idx] = max_idx as i64; + } + } + } + DataType::Float64 => { + let input = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let values = values_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut max_val = f64::NEG_INFINITY; + let mut max_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if val.is_nan() { + max_val = f64::NAN; + max_idx = d; + break; + } + if val > max_val { + max_val = val; + max_idx = d; + } + } + let out_idx = o * layout.inner + r; + values[out_idx] = max_val; + indices[out_idx] = max_idx as i64; + } + } + } + DataType::Int32 => { + let input = tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; + let values = values_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i32 slice") + })?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut max_val = i32::MIN; + let mut max_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if val > max_val { + max_val = val; + max_idx = d; + } + } + let out_idx = o * layout.inner + r; + values[out_idx] = max_val; + indices[out_idx] = max_idx as i64; + } + } + } + DataType::Int64 => { + let input = tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; + let values = values_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut max_val = i64::MIN; + let mut max_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if val > max_val { + max_val = val; + max_idx = d; + } + } + let out_idx = o * layout.inner + r; + values[out_idx] = max_val; + indices[out_idx] = max_idx as i64; + } + } + } + DataType::Bool => { + let input = tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + let values = values_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice") + })?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut max_val = false; + let mut max_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + if input[idx] { + max_val = true; + max_idx = d; + break; + } + } + let out_idx = o * layout.inner + r; + values[out_idx] = max_val; + indices[out_idx] = max_idx as i64; + } + } + } + } + + Ok(( + Tensor::new( + Arc::new(values_data), + layout.output_shape.clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ), + Tensor::new( + Arc::new(indices_data), + layout.output_shape, + DataType::Int64, + tensor.device(), + false, + ), + )) +} + +pub(crate) fn nanmax_along_dim_with_indices( + tensor: &Tensor, + dim: usize, + keepdim: bool, +) -> Result<(Tensor, Tensor)> { + let layout = reduction_layout(tensor, dim, keepdim)?; + let mut values_data = + TensorData::zeros_on_device(layout.output_shape.numel(), tensor.dtype(), tensor.device()); + let mut indices_data = TensorData::zeros_on_device( + layout.output_shape.numel(), + DataType::Int64, + tensor.device(), + ); + + let indices = indices_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + match tensor.dtype() { + DataType::Float32 => { + let input = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let values = values_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut max_val = f32::NAN; + let mut max_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if !val.is_nan() && (max_val.is_nan() || val > max_val) { + max_val = val; + max_idx = d; + } + } + let out_idx = o * layout.inner + r; + values[out_idx] = max_val; + indices[out_idx] = max_idx as i64; + } + } + } + DataType::Float64 => { + let input = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let values = values_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?; + for o in 0..layout.outer { + for r in 0..layout.inner { + let mut max_val = f64::NAN; + let mut max_idx = 0usize; + for d in 0..layout.dim_size { + let idx = o * layout.outer_stride + d * layout.inner + r; + let val = input[idx]; + if !val.is_nan() && (max_val.is_nan() || val > max_val) { + max_val = val; + max_idx = d; + } + } + let out_idx = o * layout.inner + r; + values[out_idx] = max_val; + indices[out_idx] = max_idx as i64; + } + } + } + _ => { + return Err(MinitensorError::invalid_operation( + "nanmax only supports floating point tensors", + )); + } + } + + Ok(( + Tensor::new( + Arc::new(values_data), + layout.output_shape.clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ), + Tensor::new( + Arc::new(indices_data), + layout.output_shape, + DataType::Int64, + tensor.device(), + false, + ), + )) +} diff --git a/engine/src/operations/reduction/nanquantile.rs b/engine/src/operations/reduction/nanquantile.rs index 38f9048a..b264fe14 100644 --- a/engine/src/operations/reduction/nanquantile.rs +++ b/engine/src/operations/reduction/nanquantile.rs @@ -1,854 +1,868 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -fn nanquantiles_along_dim( - tensor: &Tensor, - dim: usize, - qs: &[f64], - keepdim: bool, - interpolation: QuantileInterpolation, -) -> Result { - let dims = tensor.shape().dims(); - let dim_size = if dims.is_empty() { 1 } else { dims[dim] }; - - if dim_size == 0 { - return Err(MinitensorError::invalid_argument( - "nanquantile() does not support empty slices".to_string(), - )); - } - - let q_len = qs.len(); - - let mut out_dims = Vec::with_capacity(dims.len() + 2); - out_dims.push(q_len); - if !dims.is_empty() { - out_dims.extend_from_slice(&dims[..dim]); - if keepdim { - out_dims.push(1); - } - out_dims.extend_from_slice(&dims[dim + 1..]); - } else if keepdim { - out_dims.push(1); - } - - let shape = Shape::new(out_dims); - let mut values_data = - TensorData::zeros_on_device(shape.numel(), tensor.dtype(), tensor.device()); - - let outer = if dims.is_empty() || dim == 0 { - 1 - } else { - dims[..dim].iter().product() - }; - let inner = if dims.is_empty() || dim + 1 >= dims.len() { - 1 - } else { - dims[dim + 1..].iter().product() - }; - let outer_stride = dim_size * inner; - - match tensor.dtype() { - DataType::Float32 => { - let input = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let values = values_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?; - - if dim_size == 1 { - fill_nanquantiles_single_f32(input, values, outer, inner, outer_stride, q_len)?; - } else { - let mut buffer = Vec::with_capacity(dim_size); - let mut cached_positions: Option<(usize, Vec)> = None; - if q_len == 1 { - let q_value = qs[0]; - for o in 0..outer { - for r in 0..inner { - buffer.clear(); - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - let val = input[idx]; - if !val.is_nan() { - buffer.push(val); - } - } - - if buffer.is_empty() { - return Err(MinitensorError::invalid_argument( - NANQUANTILE_ALL_NAN_ERR.to_string(), - )); - } - - let out_idx = o * inner + r; - values[out_idx] = - quantile_from_unsorted_f32(&mut buffer, q_value, interpolation); - } - } - } else { - for o in 0..outer { - for r in 0..inner { - buffer.clear(); - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - let val = input[idx]; - if !val.is_nan() { - buffer.push(val); - } - } - - if buffer.is_empty() { - return Err(MinitensorError::invalid_argument( - NANQUANTILE_ALL_NAN_ERR.to_string(), - )); - } - - buffer.sort_by(|a, b| a.total_cmp(b)); - let positions = match cached_positions { - Some((len, ref positions)) if len == buffer.len() => positions, - _ => { - let positions = quantile_positions_for_len(buffer.len(), qs); - cached_positions = Some((buffer.len(), positions)); - &cached_positions.as_ref().expect("positions cached").1 - } - }; - for (qi, position) in positions.iter().enumerate() { - let out_idx = ((qi * outer) + o) * inner + r; - values[out_idx] = quantile_from_sorted_position_f32( - &buffer, - position, - interpolation, - ); - } - } - } - } - } - } - DataType::Float64 => { - let input = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let values = values_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?; - - if dim_size == 1 { - fill_nanquantiles_single_f64(input, values, outer, inner, outer_stride, q_len)?; - } else { - let mut buffer = Vec::with_capacity(dim_size); - let mut cached_positions: Option<(usize, Vec)> = None; - if q_len == 1 { - let q_value = qs[0]; - for o in 0..outer { - for r in 0..inner { - buffer.clear(); - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - let val = input[idx]; - if !val.is_nan() { - buffer.push(val); - } - } - - if buffer.is_empty() { - return Err(MinitensorError::invalid_argument( - NANQUANTILE_ALL_NAN_ERR.to_string(), - )); - } - - let out_idx = o * inner + r; - values[out_idx] = - quantile_from_unsorted_f64(&mut buffer, q_value, interpolation); - } - } - } else { - for o in 0..outer { - for r in 0..inner { - buffer.clear(); - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - let val = input[idx]; - if !val.is_nan() { - buffer.push(val); - } - } - - if buffer.is_empty() { - return Err(MinitensorError::invalid_argument( - NANQUANTILE_ALL_NAN_ERR.to_string(), - )); - } - - buffer.sort_by(|a, b| a.total_cmp(b)); - let positions = match cached_positions { - Some((len, ref positions)) if len == buffer.len() => positions, - _ => { - let positions = quantile_positions_for_len(buffer.len(), qs); - cached_positions = Some((buffer.len(), positions)); - &cached_positions.as_ref().expect("positions cached").1 - } - }; - for (qi, position) in positions.iter().enumerate() { - let out_idx = ((qi * outer) + o) * inner + r; - values[out_idx] = quantile_from_sorted_position_f64( - &buffer, - position, - interpolation, - ); - } - } - } - } - } - } - _ => unreachable!("dtype validated"), - } - - Ok(Tensor::new( - Arc::new(values_data), - shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )) -} - -fn quantile_from_unsorted_f32( - values: &mut [f32], - q: f64, - interpolation: QuantileInterpolation, -) -> f32 { - if values.len() == 1 { - return values[0]; - } - - let position = quantile_position_for_len_q(values.len(), q); - let lower_idx = position.lower_idx; - let upper_idx = position.upper_idx; - let weight = position.weight; - - if lower_idx == upper_idx { - return select_quantile_at_f32(values, lower_idx); - } - - let value = match interpolation { - QuantileInterpolation::Lower => select_quantile_at_f32(values, lower_idx) as f64, - QuantileInterpolation::Higher => select_quantile_at_f32(values, upper_idx) as f64, - QuantileInterpolation::Nearest => { - let idx = position.nearest_idx; - select_quantile_at_f32(values, idx) as f64 - } - QuantileInterpolation::Linear | QuantileInterpolation::Midpoint => { - let (lower, upper) = select_quantile_bounds_f32(values, lower_idx, upper_idx); - interpolation.interpolate(lower as f64, upper as f64, weight) - } - }; - - value as f32 -} - -fn quantile_from_unsorted_f64( - values: &mut [f64], - q: f64, - interpolation: QuantileInterpolation, -) -> f64 { - if values.len() == 1 { - return values[0]; - } - - let position = quantile_position_for_len_q(values.len(), q); - let lower_idx = position.lower_idx; - let upper_idx = position.upper_idx; - let weight = position.weight; - - if lower_idx == upper_idx { - return select_quantile_at_f64(values, lower_idx); - } - - match interpolation { - QuantileInterpolation::Lower => select_quantile_at_f64(values, lower_idx), - QuantileInterpolation::Higher => select_quantile_at_f64(values, upper_idx), - QuantileInterpolation::Nearest => { - let idx = position.nearest_idx; - select_quantile_at_f64(values, idx) - } - QuantileInterpolation::Linear | QuantileInterpolation::Midpoint => { - let (lower, upper) = select_quantile_bounds_f64(values, lower_idx, upper_idx); - interpolation.interpolate(lower, upper, weight) - } - } -} - -fn select_quantile_at_f32(values: &mut [f32], idx: usize) -> f32 { - let (_, pivot, _) = values.select_nth_unstable_by(idx, |a, b| a.total_cmp(b)); - *pivot -} - -fn select_quantile_bounds_f32( - values: &mut [f32], - lower_idx: usize, - upper_idx: usize, -) -> (f32, f32) { - if lower_idx == upper_idx { - let value = select_quantile_at_f32(values, lower_idx); - return (value, value); - } - - let (_, upper_pivot, _) = values.select_nth_unstable_by(upper_idx, |a, b| a.total_cmp(b)); - let upper = *upper_pivot; - let (_, lower_pivot, _) = - values[..upper_idx].select_nth_unstable_by(lower_idx, |a, b| a.total_cmp(b)); - let lower = *lower_pivot; - (lower, upper) -} - -fn select_quantile_at_f64(values: &mut [f64], idx: usize) -> f64 { - let (_, pivot, _) = values.select_nth_unstable_by(idx, |a, b| a.total_cmp(b)); - *pivot -} - -fn select_quantile_bounds_f64( - values: &mut [f64], - lower_idx: usize, - upper_idx: usize, -) -> (f64, f64) { - if lower_idx == upper_idx { - let value = select_quantile_at_f64(values, lower_idx); - return (value, value); - } - - let (_, upper_pivot, _) = values.select_nth_unstable_by(upper_idx, |a, b| a.total_cmp(b)); - let upper = *upper_pivot; - let (_, lower_pivot, _) = - values[..upper_idx].select_nth_unstable_by(lower_idx, |a, b| a.total_cmp(b)); - let lower = *lower_pivot; - (lower, upper) -} - -fn median_all(tensor: &Tensor) -> Result<(Tensor, Option)> { - let mut result_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => { - let data = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let mut values = Vec::with_capacity(data.len()); - for &value in data { - if value.is_nan() { - result_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?[0] = f32::NAN; - return Ok(( - Tensor::new( - Arc::new(result_data), - Shape::scalar(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ), - None, - )); - } - values.push(value); - } - let median_index = (values.len() - 1) / 2; - values.select_nth_unstable_by(median_index, |a, b| a.total_cmp(b)); - let median = values[median_index]; - result_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?[0] = median; - } - DataType::Float64 => { - let data = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let mut values = Vec::with_capacity(data.len()); - for &value in data { - if value.is_nan() { - result_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?[0] = f64::NAN; - return Ok(( - Tensor::new( - Arc::new(result_data), - Shape::scalar(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ), - None, - )); - } - values.push(value); - } - let median_index = (values.len() - 1) / 2; - values.select_nth_unstable_by(median_index, |a, b| a.total_cmp(b)); - let median = values[median_index]; - result_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?[0] = median; - } - DataType::Int32 => { - let data = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - let mut values: Vec = data.to_vec(); - let median_index = (values.len() - 1) / 2; - values.select_nth_unstable(median_index); - let median = values[median_index]; - result_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i32 slice") - })?[0] = median; - } - DataType::Int64 => { - let data = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - let mut values: Vec = data.to_vec(); - let median_index = (values.len() - 1) / 2; - values.select_nth_unstable(median_index); - let median = values[median_index]; - result_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?[0] = median; - } - DataType::Bool => { - let data = tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - let mut values: Vec = data.to_vec(); - let median_index = (values.len() - 1) / 2; - values.select_nth_unstable(median_index); - let median = values[median_index]; - result_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice") - })?[0] = median; - } - } - - let value = Tensor::new( - Arc::new(result_data), - Shape::scalar(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - Ok((value, None)) -} - -fn median_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result<(Tensor, Tensor)> { - let dims = tensor.shape().dims(); - let dim_size = if dims.is_empty() { 1 } else { dims[dim] }; - - ensure_non_empty(dim_size)?; - - let mut out_dims = if dims.is_empty() { - vec![1] - } else { - dims.to_vec() - }; - - if keepdim { - if !out_dims.is_empty() { - out_dims[dim] = 1; - } - } else if !out_dims.is_empty() { - out_dims.remove(dim); - } - - let values_shape = Shape::new(out_dims); - let num_out = values_shape.numel(); - - let mut values_data = TensorData::zeros_on_device(num_out, tensor.dtype(), tensor.device()); - let mut indices_data = TensorData::zeros_on_device(num_out, DataType::Int64, tensor.device()); - - let outer = if dims.is_empty() || dim == 0 { - 1 - } else { - dims[..dim].iter().product() - }; - let inner = if dims.is_empty() || dim + 1 >= dims.len() { - 1 - } else { - dims[dim + 1..].iter().product() - }; - let outer_stride = dim_size * inner; - let median_pos = (dim_size - 1) / 2; - - match tensor.dtype() { - DataType::Float32 => { - let input = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let values = values_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?; - let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - - let mut entries = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - entries.clear(); - let mut has_nan = false; - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - let value = input[idx]; - if value.is_nan() { - has_nan = true; - break; - } - entries.push((d, value)); - } - - let base = o * inner + r; - if has_nan { - values[base] = f32::NAN; - continue; - } - - entries.select_nth_unstable_by(median_pos, cmp_f32_asc); - let (index, value) = entries[median_pos]; - values[base] = value; - indices[base] = index as i64; - } - } - } - DataType::Float64 => { - let input = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let values = values_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?; - let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - - let mut entries = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - entries.clear(); - let mut has_nan = false; - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - let value = input[idx]; - if value.is_nan() { - has_nan = true; - break; - } - entries.push((d, value)); - } - - let base = o * inner + r; - if has_nan { - values[base] = f64::NAN; - continue; - } - - entries.select_nth_unstable_by(median_pos, cmp_f64_asc); - let (index, value) = entries[median_pos]; - values[base] = value; - indices[base] = index as i64; - } - } - } - DataType::Int32 => { - let input = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - let values = values_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i32 slice") - })?; - let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - - let mut entries = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - entries.clear(); - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - entries.push((d, input[idx])); - } - - entries.select_nth_unstable_by(median_pos, cmp_i32_asc); - let (index, value) = entries[median_pos]; - let base = o * inner + r; - values[base] = value; - indices[base] = index as i64; - } - } - } - DataType::Int64 => { - let input = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - let values = values_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - - let mut entries = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - entries.clear(); - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - entries.push((d, input[idx])); - } - - entries.select_nth_unstable_by(median_pos, cmp_i64_asc); - let (index, value) = entries[median_pos]; - let base = o * inner + r; - values[base] = value; - indices[base] = index as i64; - } - } - } - DataType::Bool => { - let input = tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - let values = values_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice") - })?; - let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - - let mut entries = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - entries.clear(); - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - entries.push((d, input[idx])); - } - - entries.select_nth_unstable_by(median_pos, cmp_bool_asc); - let (index, value) = entries[median_pos]; - let base = o * inner + r; - values[base] = value; - indices[base] = index as i64; - } - } - } - } - - let values = Tensor::new( - Arc::new(values_data), - values_shape.clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - - let indices = Tensor::new( - Arc::new(indices_data), - values_shape, - DataType::Int64, - tensor.device(), - false, - ); - - Ok((values, indices)) -} - -fn normalize_dim(dim: isize, ndim: usize) -> Result { - let dim = if dim < 0 { dim + ndim as isize } else { dim }; - if dim < 0 || dim >= ndim as isize { - Err(MinitensorError::index_error(dim, 0, ndim)) - } else { - Ok(dim as usize) - } -} - -/// Sum reduction along specified dimensions -pub fn sum(tensor: &Tensor, dim: Option>, keepdim: bool) -> Result { - // Normalise negative dimensions and deduplicate - let ndim = tensor.ndim() as isize; - let dim = match dim { - Some(dims) => { - let mut normalized = Vec::with_capacity(dims.len()); - for d in dims { - let d = if d < 0 { d + ndim } else { d }; - if d < 0 || d >= ndim { - return Err(MinitensorError::index_error(d, 0, tensor.ndim())); - } - normalized.push(d as usize); - } - normalized.sort_unstable(); - normalized.dedup(); - Some(normalized) - } - None => None, - }; - let dims_clone = dim.clone(); - - let result = match dim { - None => { - // Sum all elements - let result_shape = if keepdim { - Shape::new(vec![1; tensor.ndim()]) - } else { - Shape::scalar() - }; - - let mut result_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => sum_all_f32(tensor, &mut result_data)?, - DataType::Float64 => sum_all_f64(tensor, &mut result_data)?, - DataType::Int32 => sum_all_i32(tensor, &mut result_data)?, - DataType::Int64 => sum_all_i64(tensor, &mut result_data)?, - DataType::Bool => { - return Err(MinitensorError::invalid_operation( - "Sum not supported for boolean tensors", - )); - } - } - - Tensor::new( - Arc::new(result_data), - result_shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ) - } - Some(dims) => { - // Sum along specific dimensions - if dims.is_empty() { - tensor.clone() - } else { - let mut result = tensor.clone(); - if keepdim { - for &d in &dims { - result = sum_along_dim(&result, d, true)?; - } - } else { - for &d in dims.iter().rev() { - result = sum_along_dim(&result, d, false)?; - } - } - result - } - } - }; - - if result.requires_grad() { - let grad_fn = Arc::new(SumBackward { - input_id: tensor.id(), - input_shape: tensor.shape().dims().to_vec(), - dims: dims_clone, - keepdim, - }); - let mut result_with_grad = result; - result_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&result_with_grad, Some(grad_fn))?; - Ok(result_with_grad) - } else { - Ok(result) - } -} - -/// NaN-aware sum reduction along specified dimensions -pub fn nansum(tensor: &Tensor, dim: Option>, keepdim: bool) -> Result { - if !tensor.dtype().is_float() { - return sum(tensor, dim, keepdim); - } - - let dim = normalize_reduction_dims(dim, tensor.ndim())?; - let dims_clone = dim.clone(); - let needs_mask = - tensor.requires_grad() || dim.as_ref().map(|dims| !dims.is_empty()).unwrap_or(false); - let mask = if needs_mask { - Some(non_nan_mask(tensor)?) - } else { - None - }; - - let result = match dim { - None => { - let result_shape = if keepdim { - Shape::new(vec![1; tensor.ndim()]) - } else { - Shape::scalar() - }; - - let mut result_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); - match tensor.dtype() { - DataType::Float32 => nansum_all_f32(tensor, &mut result_data)?, - DataType::Float64 => nansum_all_f64(tensor, &mut result_data)?, - _ => unreachable!("nansum only supports floating point tensors"), - } - - Tensor::new( - Arc::new(result_data), - result_shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ) - } - Some(dims) => { - if dims.is_empty() { - tensor.clone() - } else { - let mut result = tensor.clone(); - if keepdim { - for &d in &dims { - result = nansum_along_dim(&result, d, true)?; - } - } else { - for &d in dims.iter().rev() { - result = nansum_along_dim(&result, d, false)?; - } - } - result - } - } - }; - - if result.requires_grad() { - let mask = mask.ok_or_else(|| { - MinitensorError::internal_error("nansum expected mask for gradient computation") - })?; - let grad_fn = Arc::new(NanSumBackward { - input_id: tensor.id(), - input_shape: tensor.shape().dims().to_vec(), - dims: dims_clone, - keepdim, - mask, - }); - let mut result_with_grad = result; - result_with_grad.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&result_with_grad, Some(grad_fn))?; - Ok(result_with_grad) - } else { - Ok(result) - } -} +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::autograd::NanSumBackward; +use crate::autograd::SumBackward; +use crate::{ + autograd::add_to_graph, + error::{MinitensorError, Result}, + tensor::{DataType, Shape, Tensor, TensorData}, +}; +use std::sync::Arc; + +pub(crate) fn nanquantiles_along_dim( + tensor: &Tensor, + dim: usize, + qs: &[f64], + keepdim: bool, + interpolation: QuantileInterpolation, +) -> Result { + let dims = tensor.shape().dims(); + let dim_size = if dims.is_empty() { 1 } else { dims[dim] }; + + if dim_size == 0 { + return Err(MinitensorError::invalid_argument( + "nanquantile() does not support empty slices".to_string(), + )); + } + + let q_len = qs.len(); + + let mut out_dims = Vec::with_capacity(dims.len() + 2); + out_dims.push(q_len); + if !dims.is_empty() { + out_dims.extend_from_slice(&dims[..dim]); + if keepdim { + out_dims.push(1); + } + out_dims.extend_from_slice(&dims[dim + 1..]); + } else if keepdim { + out_dims.push(1); + } + + let shape = Shape::new(out_dims); + let mut values_data = + TensorData::zeros_on_device(shape.numel(), tensor.dtype(), tensor.device()); + + let outer = if dims.is_empty() || dim == 0 { + 1 + } else { + dims[..dim].iter().product() + }; + let inner = if dims.is_empty() || dim + 1 >= dims.len() { + 1 + } else { + dims[dim + 1..].iter().product() + }; + let outer_stride = dim_size * inner; + + match tensor.dtype() { + DataType::Float32 => { + let input = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let values = values_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?; + + if dim_size == 1 { + fill_nanquantiles_single_f32(input, values, outer, inner, outer_stride, q_len)?; + } else { + let mut buffer = Vec::with_capacity(dim_size); + let mut cached_positions: Option<(usize, Vec)> = None; + if q_len == 1 { + let q_value = qs[0]; + for o in 0..outer { + for r in 0..inner { + buffer.clear(); + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + let val = input[idx]; + if !val.is_nan() { + buffer.push(val); + } + } + + if buffer.is_empty() { + return Err(MinitensorError::invalid_argument( + NANQUANTILE_ALL_NAN_ERR.to_string(), + )); + } + + let out_idx = o * inner + r; + values[out_idx] = + quantile_from_unsorted_f32(&mut buffer, q_value, interpolation); + } + } + } else { + for o in 0..outer { + for r in 0..inner { + buffer.clear(); + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + let val = input[idx]; + if !val.is_nan() { + buffer.push(val); + } + } + + if buffer.is_empty() { + return Err(MinitensorError::invalid_argument( + NANQUANTILE_ALL_NAN_ERR.to_string(), + )); + } + + buffer.sort_by(|a, b| a.total_cmp(b)); + let positions = match cached_positions { + Some((len, ref positions)) if len == buffer.len() => positions, + _ => { + let positions = quantile_positions_for_len(buffer.len(), qs); + cached_positions = Some((buffer.len(), positions)); + &cached_positions.as_ref().expect("positions cached").1 + } + }; + for (qi, position) in positions.iter().enumerate() { + let out_idx = ((qi * outer) + o) * inner + r; + values[out_idx] = quantile_from_sorted_position_f32( + &buffer, + position, + interpolation, + ); + } + } + } + } + } + } + DataType::Float64 => { + let input = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let values = values_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?; + + if dim_size == 1 { + fill_nanquantiles_single_f64(input, values, outer, inner, outer_stride, q_len)?; + } else { + let mut buffer = Vec::with_capacity(dim_size); + let mut cached_positions: Option<(usize, Vec)> = None; + if q_len == 1 { + let q_value = qs[0]; + for o in 0..outer { + for r in 0..inner { + buffer.clear(); + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + let val = input[idx]; + if !val.is_nan() { + buffer.push(val); + } + } + + if buffer.is_empty() { + return Err(MinitensorError::invalid_argument( + NANQUANTILE_ALL_NAN_ERR.to_string(), + )); + } + + let out_idx = o * inner + r; + values[out_idx] = + quantile_from_unsorted_f64(&mut buffer, q_value, interpolation); + } + } + } else { + for o in 0..outer { + for r in 0..inner { + buffer.clear(); + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + let val = input[idx]; + if !val.is_nan() { + buffer.push(val); + } + } + + if buffer.is_empty() { + return Err(MinitensorError::invalid_argument( + NANQUANTILE_ALL_NAN_ERR.to_string(), + )); + } + + buffer.sort_by(|a, b| a.total_cmp(b)); + let positions = match cached_positions { + Some((len, ref positions)) if len == buffer.len() => positions, + _ => { + let positions = quantile_positions_for_len(buffer.len(), qs); + cached_positions = Some((buffer.len(), positions)); + &cached_positions.as_ref().expect("positions cached").1 + } + }; + for (qi, position) in positions.iter().enumerate() { + let out_idx = ((qi * outer) + o) * inner + r; + values[out_idx] = quantile_from_sorted_position_f64( + &buffer, + position, + interpolation, + ); + } + } + } + } + } + } + _ => unreachable!("dtype validated"), + } + + Ok(Tensor::new( + Arc::new(values_data), + shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )) +} + +pub(crate) fn quantile_from_unsorted_f32( + values: &mut [f32], + q: f64, + interpolation: QuantileInterpolation, +) -> f32 { + if values.len() == 1 { + return values[0]; + } + + let position = quantile_position_for_len_q(values.len(), q); + let lower_idx = position.lower_idx; + let upper_idx = position.upper_idx; + let weight = position.weight; + + if lower_idx == upper_idx { + return select_quantile_at_f32(values, lower_idx); + } + + let value = match interpolation { + QuantileInterpolation::Lower => select_quantile_at_f32(values, lower_idx) as f64, + QuantileInterpolation::Higher => select_quantile_at_f32(values, upper_idx) as f64, + QuantileInterpolation::Nearest => { + let idx = position.nearest_idx; + select_quantile_at_f32(values, idx) as f64 + } + QuantileInterpolation::Linear | QuantileInterpolation::Midpoint => { + let (lower, upper) = select_quantile_bounds_f32(values, lower_idx, upper_idx); + interpolation.interpolate(lower as f64, upper as f64, weight) + } + }; + + value as f32 +} + +pub(crate) fn quantile_from_unsorted_f64( + values: &mut [f64], + q: f64, + interpolation: QuantileInterpolation, +) -> f64 { + if values.len() == 1 { + return values[0]; + } + + let position = quantile_position_for_len_q(values.len(), q); + let lower_idx = position.lower_idx; + let upper_idx = position.upper_idx; + let weight = position.weight; + + if lower_idx == upper_idx { + return select_quantile_at_f64(values, lower_idx); + } + + match interpolation { + QuantileInterpolation::Lower => select_quantile_at_f64(values, lower_idx), + QuantileInterpolation::Higher => select_quantile_at_f64(values, upper_idx), + QuantileInterpolation::Nearest => { + let idx = position.nearest_idx; + select_quantile_at_f64(values, idx) + } + QuantileInterpolation::Linear | QuantileInterpolation::Midpoint => { + let (lower, upper) = select_quantile_bounds_f64(values, lower_idx, upper_idx); + interpolation.interpolate(lower, upper, weight) + } + } +} + +fn select_quantile_at_f32(values: &mut [f32], idx: usize) -> f32 { + let (_, pivot, _) = values.select_nth_unstable_by(idx, |a, b| a.total_cmp(b)); + *pivot +} + +fn select_quantile_bounds_f32( + values: &mut [f32], + lower_idx: usize, + upper_idx: usize, +) -> (f32, f32) { + if lower_idx == upper_idx { + let value = select_quantile_at_f32(values, lower_idx); + return (value, value); + } + + let (_, upper_pivot, _) = values.select_nth_unstable_by(upper_idx, |a, b| a.total_cmp(b)); + let upper = *upper_pivot; + let (_, lower_pivot, _) = + values[..upper_idx].select_nth_unstable_by(lower_idx, |a, b| a.total_cmp(b)); + let lower = *lower_pivot; + (lower, upper) +} + +fn select_quantile_at_f64(values: &mut [f64], idx: usize) -> f64 { + let (_, pivot, _) = values.select_nth_unstable_by(idx, |a, b| a.total_cmp(b)); + *pivot +} + +fn select_quantile_bounds_f64( + values: &mut [f64], + lower_idx: usize, + upper_idx: usize, +) -> (f64, f64) { + if lower_idx == upper_idx { + let value = select_quantile_at_f64(values, lower_idx); + return (value, value); + } + + let (_, upper_pivot, _) = values.select_nth_unstable_by(upper_idx, |a, b| a.total_cmp(b)); + let upper = *upper_pivot; + let (_, lower_pivot, _) = + values[..upper_idx].select_nth_unstable_by(lower_idx, |a, b| a.total_cmp(b)); + let lower = *lower_pivot; + (lower, upper) +} + +pub(crate) fn median_all(tensor: &Tensor) -> Result<(Tensor, Option)> { + let mut result_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => { + let data = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let mut values = Vec::with_capacity(data.len()); + for &value in data { + if value.is_nan() { + result_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?[0] = f32::NAN; + return Ok(( + Tensor::new( + Arc::new(result_data), + Shape::scalar(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ), + None, + )); + } + values.push(value); + } + let median_index = (values.len() - 1) / 2; + values.select_nth_unstable_by(median_index, |a, b| a.total_cmp(b)); + let median = values[median_index]; + result_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?[0] = median; + } + DataType::Float64 => { + let data = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let mut values = Vec::with_capacity(data.len()); + for &value in data { + if value.is_nan() { + result_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?[0] = f64::NAN; + return Ok(( + Tensor::new( + Arc::new(result_data), + Shape::scalar(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ), + None, + )); + } + values.push(value); + } + let median_index = (values.len() - 1) / 2; + values.select_nth_unstable_by(median_index, |a, b| a.total_cmp(b)); + let median = values[median_index]; + result_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?[0] = median; + } + DataType::Int32 => { + let data = tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; + let mut values: Vec = data.to_vec(); + let median_index = (values.len() - 1) / 2; + values.select_nth_unstable(median_index); + let median = values[median_index]; + result_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i32 slice") + })?[0] = median; + } + DataType::Int64 => { + let data = tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; + let mut values: Vec = data.to_vec(); + let median_index = (values.len() - 1) / 2; + values.select_nth_unstable(median_index); + let median = values[median_index]; + result_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?[0] = median; + } + DataType::Bool => { + let data = tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + let mut values: Vec = data.to_vec(); + let median_index = (values.len() - 1) / 2; + values.select_nth_unstable(median_index); + let median = values[median_index]; + result_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice") + })?[0] = median; + } + } + + let value = Tensor::new( + Arc::new(result_data), + Shape::scalar(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + Ok((value, None)) +} + +pub(crate) fn median_along_dim( + tensor: &Tensor, + dim: usize, + keepdim: bool, +) -> Result<(Tensor, Tensor)> { + let dims = tensor.shape().dims(); + let dim_size = if dims.is_empty() { 1 } else { dims[dim] }; + + ensure_non_empty(dim_size)?; + + let mut out_dims = if dims.is_empty() { + vec![1] + } else { + dims.to_vec() + }; + + if keepdim { + if !out_dims.is_empty() { + out_dims[dim] = 1; + } + } else if !out_dims.is_empty() { + out_dims.remove(dim); + } + + let values_shape = Shape::new(out_dims); + let num_out = values_shape.numel(); + + let mut values_data = TensorData::zeros_on_device(num_out, tensor.dtype(), tensor.device()); + let mut indices_data = TensorData::zeros_on_device(num_out, DataType::Int64, tensor.device()); + + let outer = if dims.is_empty() || dim == 0 { + 1 + } else { + dims[..dim].iter().product() + }; + let inner = if dims.is_empty() || dim + 1 >= dims.len() { + 1 + } else { + dims[dim + 1..].iter().product() + }; + let outer_stride = dim_size * inner; + let median_pos = (dim_size - 1) / 2; + + match tensor.dtype() { + DataType::Float32 => { + let input = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let values = values_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?; + let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + + let mut entries = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + entries.clear(); + let mut has_nan = false; + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + let value = input[idx]; + if value.is_nan() { + has_nan = true; + break; + } + entries.push((d, value)); + } + + let base = o * inner + r; + if has_nan { + values[base] = f32::NAN; + continue; + } + + entries.select_nth_unstable_by(median_pos, cmp_f32_asc); + let (index, value) = entries[median_pos]; + values[base] = value; + indices[base] = index as i64; + } + } + } + DataType::Float64 => { + let input = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let values = values_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?; + let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + + let mut entries = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + entries.clear(); + let mut has_nan = false; + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + let value = input[idx]; + if value.is_nan() { + has_nan = true; + break; + } + entries.push((d, value)); + } + + let base = o * inner + r; + if has_nan { + values[base] = f64::NAN; + continue; + } + + entries.select_nth_unstable_by(median_pos, cmp_f64_asc); + let (index, value) = entries[median_pos]; + values[base] = value; + indices[base] = index as i64; + } + } + } + DataType::Int32 => { + let input = tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; + let values = values_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i32 slice") + })?; + let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + + let mut entries = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + entries.clear(); + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + entries.push((d, input[idx])); + } + + entries.select_nth_unstable_by(median_pos, cmp_i32_asc); + let (index, value) = entries[median_pos]; + let base = o * inner + r; + values[base] = value; + indices[base] = index as i64; + } + } + } + DataType::Int64 => { + let input = tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; + let values = values_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + + let mut entries = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + entries.clear(); + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + entries.push((d, input[idx])); + } + + entries.select_nth_unstable_by(median_pos, cmp_i64_asc); + let (index, value) = entries[median_pos]; + let base = o * inner + r; + values[base] = value; + indices[base] = index as i64; + } + } + } + DataType::Bool => { + let input = tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + let values = values_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice") + })?; + let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + + let mut entries = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + entries.clear(); + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + entries.push((d, input[idx])); + } + + entries.select_nth_unstable_by(median_pos, cmp_bool_asc); + let (index, value) = entries[median_pos]; + let base = o * inner + r; + values[base] = value; + indices[base] = index as i64; + } + } + } + } + + let values = Tensor::new( + Arc::new(values_data), + values_shape.clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + + let indices = Tensor::new( + Arc::new(indices_data), + values_shape, + DataType::Int64, + tensor.device(), + false, + ); + + Ok((values, indices)) +} + +pub(crate) fn normalize_dim(dim: isize, ndim: usize) -> Result { + let dim = if dim < 0 { dim + ndim as isize } else { dim }; + if dim < 0 || dim >= ndim as isize { + Err(MinitensorError::index_error(dim, 0, ndim)) + } else { + Ok(dim as usize) + } +} + +/// Sum reduction along specified dimensions +pub fn sum(tensor: &Tensor, dim: Option>, keepdim: bool) -> Result { + // Normalise negative dimensions and deduplicate + let ndim = tensor.ndim() as isize; + let dim = match dim { + Some(dims) => { + let mut normalized = Vec::with_capacity(dims.len()); + for d in dims { + let d = if d < 0 { d + ndim } else { d }; + if d < 0 || d >= ndim { + return Err(MinitensorError::index_error(d, 0, tensor.ndim())); + } + normalized.push(d as usize); + } + normalized.sort_unstable(); + normalized.dedup(); + Some(normalized) + } + None => None, + }; + let dims_clone = dim.clone(); + + let result = match dim { + None => { + // Sum all elements + let result_shape = if keepdim { + Shape::new(vec![1; tensor.ndim()]) + } else { + Shape::scalar() + }; + + let mut result_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => sum_all_f32(tensor, &mut result_data)?, + DataType::Float64 => sum_all_f64(tensor, &mut result_data)?, + DataType::Int32 => sum_all_i32(tensor, &mut result_data)?, + DataType::Int64 => sum_all_i64(tensor, &mut result_data)?, + DataType::Bool => { + return Err(MinitensorError::invalid_operation( + "Sum not supported for boolean tensors", + )); + } + } + + Tensor::new( + Arc::new(result_data), + result_shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ) + } + Some(dims) => { + // Sum along specific dimensions + if dims.is_empty() { + tensor.clone() + } else { + let mut result = tensor.clone(); + if keepdim { + for &d in &dims { + result = sum_along_dim(&result, d, true)?; + } + } else { + for &d in dims.iter().rev() { + result = sum_along_dim(&result, d, false)?; + } + } + result + } + } + }; + + if result.requires_grad() { + let grad_fn = Arc::new(SumBackward { + input_id: tensor.id(), + input_shape: tensor.shape().dims().to_vec(), + dims: dims_clone, + keepdim, + }); + let mut result_with_grad = result; + result_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&result_with_grad, Some(grad_fn))?; + Ok(result_with_grad) + } else { + Ok(result) + } +} + +/// NaN-aware sum reduction along specified dimensions +pub fn nansum(tensor: &Tensor, dim: Option>, keepdim: bool) -> Result { + if !tensor.dtype().is_float() { + return sum(tensor, dim, keepdim); + } + + let dim = normalize_reduction_dims(dim, tensor.ndim())?; + let dims_clone = dim.clone(); + let needs_mask = + tensor.requires_grad() || dim.as_ref().map(|dims| !dims.is_empty()).unwrap_or(false); + let mask = if needs_mask { + Some(non_nan_mask(tensor)?) + } else { + None + }; + + let result = match dim { + None => { + let result_shape = if keepdim { + Shape::new(vec![1; tensor.ndim()]) + } else { + Shape::scalar() + }; + + let mut result_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); + match tensor.dtype() { + DataType::Float32 => nansum_all_f32(tensor, &mut result_data)?, + DataType::Float64 => nansum_all_f64(tensor, &mut result_data)?, + _ => unreachable!("nansum only supports floating point tensors"), + } + + Tensor::new( + Arc::new(result_data), + result_shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ) + } + Some(dims) => { + if dims.is_empty() { + tensor.clone() + } else { + let mut result = tensor.clone(); + if keepdim { + for &d in &dims { + result = nansum_along_dim(&result, d, true)?; + } + } else { + for &d in dims.iter().rev() { + result = nansum_along_dim(&result, d, false)?; + } + } + result + } + } + }; + + if result.requires_grad() { + let mask = mask.ok_or_else(|| { + MinitensorError::internal_error("nansum expected mask for gradient computation") + })?; + let grad_fn = Arc::new(NanSumBackward { + input_id: tensor.id(), + input_shape: tensor.shape().dims().to_vec(), + dims: dims_clone, + keepdim, + mask, + }); + let mut result_with_grad = result; + result_with_grad.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&result_with_grad, Some(grad_fn))?; + Ok(result_with_grad) + } else { + Ok(result) + } +} diff --git a/engine/src/operations/reduction/quantile.rs b/engine/src/operations/reduction/quantile.rs index 1ae5c9a0..9d228626 100644 --- a/engine/src/operations/reduction/quantile.rs +++ b/engine/src/operations/reduction/quantile.rs @@ -1,891 +1,896 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -fn quantile_along_dim( - tensor: &Tensor, - dim: usize, - keepdim: bool, - q: f64, - interpolation: QuantileInterpolation, -) -> Result { - let dims = tensor.shape().dims(); - let dim_size = if dims.is_empty() { 1 } else { dims[dim] }; - - if dim_size == 0 { - return Err(MinitensorError::invalid_argument( - "quantile() does not support reductions over empty dimensions".to_string(), - )); - } - - let mut out_dims = if dims.is_empty() { - vec![1] - } else { - dims.to_vec() - }; - - if keepdim { - if !out_dims.is_empty() { - out_dims[dim] = 1; - } - } else if !out_dims.is_empty() { - out_dims.remove(dim); - } - - let values_shape = Shape::new(out_dims); - let num_out = values_shape.numel(); - let mut values_data = TensorData::zeros_on_device(num_out, tensor.dtype(), tensor.device()); - - let outer = if dims.is_empty() || dim == 0 { - 1 - } else { - dims[..dim].iter().product() - }; - let inner = if dims.is_empty() || dim + 1 >= dims.len() { - 1 - } else { - dims[dim + 1..].iter().product() - }; - let outer_stride = dim_size * inner; - - match tensor.dtype() { - DataType::Float32 => { - let input = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let values = values_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?; - - if dim_size == 1 { - fill_quantile_single_f32(input, values, outer, inner, outer_stride); - } else { - let mut buffer = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - buffer.clear(); - let mut has_nan = false; - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - let value = input[idx]; - if value.is_nan() { - has_nan = true; - break; - } - buffer.push(value); - } - - if has_nan { - values[o * inner + r] = f32::NAN; - continue; - } - - values[o * inner + r] = - quantile_from_unsorted_f32(&mut buffer, q, interpolation); - } - } - } - } - DataType::Float64 => { - let input = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let values = values_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?; - - if dim_size == 1 { - fill_quantile_single_f64(input, values, outer, inner, outer_stride); - } else { - let mut buffer = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - buffer.clear(); - let mut has_nan = false; - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - let value = input[idx]; - if value.is_nan() { - has_nan = true; - break; - } - buffer.push(value); - } - - if has_nan { - values[o * inner + r] = f64::NAN; - continue; - } - - values[o * inner + r] = - quantile_from_unsorted_f64(&mut buffer, q, interpolation); - } - } - } - } - _ => unreachable!("dtype validated"), - } - - Ok(Tensor::new( - Arc::new(values_data), - values_shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )) -} - -fn nanquantile_along_dim( - tensor: &Tensor, - dim: usize, - keepdim: bool, - q: f64, - interpolation: QuantileInterpolation, -) -> Result { - let dims = tensor.shape().dims(); - let dim_size = if dims.is_empty() { 1 } else { dims[dim] }; - - if dim_size == 0 { - return Err(MinitensorError::invalid_argument( - "nanquantile() does not support reductions over empty dimensions".to_string(), - )); - } - - let mut out_dims = if dims.is_empty() { - vec![1] - } else { - dims.to_vec() - }; - - if keepdim { - if !out_dims.is_empty() { - out_dims[dim] = 1; - } - } else if !out_dims.is_empty() { - out_dims.remove(dim); - } - - let values_shape = Shape::new(out_dims); - let num_out = values_shape.numel(); - let mut values_data = TensorData::zeros_on_device(num_out, tensor.dtype(), tensor.device()); - - let outer = if dims.is_empty() || dim == 0 { - 1 - } else { - dims[..dim].iter().product() - }; - let inner = if dims.is_empty() || dim + 1 >= dims.len() { - 1 - } else { - dims[dim + 1..].iter().product() - }; - let outer_stride = dim_size * inner; - - match tensor.dtype() { - DataType::Float32 => { - let input = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let values = values_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?; - - if dim_size == 1 { - fill_nanquantile_single_f32(input, values, outer, inner, outer_stride)?; - } else { - let mut buffer = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - buffer.clear(); - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - let val = input[idx]; - if !val.is_nan() { - buffer.push(val); - } - } - - if buffer.is_empty() { - return Err(MinitensorError::invalid_argument( - NANQUANTILE_ALL_NAN_ERR.to_string(), - )); - } - - let quant = quantile_from_unsorted_f32(&mut buffer, q, interpolation); - values[o * inner + r] = quant; - } - } - } - } - DataType::Float64 => { - let input = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let values = values_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?; - - if dim_size == 1 { - fill_nanquantile_single_f64(input, values, outer, inner, outer_stride)?; - } else { - let mut buffer = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - buffer.clear(); - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - let val = input[idx]; - if !val.is_nan() { - buffer.push(val); - } - } - - if buffer.is_empty() { - return Err(MinitensorError::invalid_argument( - NANQUANTILE_ALL_NAN_ERR.to_string(), - )); - } - - let quant = quantile_from_unsorted_f64(&mut buffer, q, interpolation); - values[o * inner + r] = quant; - } - } - } - } - _ => unreachable!("dtype validated"), - } - - Ok(Tensor::new( - Arc::new(values_data), - values_shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )) -} - -fn nanmedian_all(tensor: &Tensor, keepdim: bool) -> Result { - let output_dims = if keepdim && tensor.ndim() > 0 { - vec![1; tensor.ndim()] - } else { - Vec::new() - }; - let shape = Shape::new(output_dims); - let mut values_data = - TensorData::zeros_on_device(shape.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => { - let data = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let values = values_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?; - let mut buffer: Vec = data.iter().copied().filter(|v| !v.is_nan()).collect(); - values[0] = if buffer.is_empty() { - f32::NAN - } else { - quantile_from_unsorted_f32(&mut buffer, 0.5, QuantileInterpolation::Linear) - }; - } - DataType::Float64 => { - let data = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let values = values_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?; - let mut buffer: Vec = data.iter().copied().filter(|v| !v.is_nan()).collect(); - values[0] = if buffer.is_empty() { - f64::NAN - } else { - quantile_from_unsorted_f64(&mut buffer, 0.5, QuantileInterpolation::Linear) - }; - } - _ => unreachable!("dtype validated"), - } - - Ok(Tensor::new( - Arc::new(values_data), - shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )) -} - -fn nanmedian_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { - let dims = tensor.shape().dims(); - let dim_size = if dims.is_empty() { 1 } else { dims[dim] }; - - let mut out_dims = if dims.is_empty() { - vec![1] - } else { - dims.to_vec() - }; - - if keepdim { - if !out_dims.is_empty() { - out_dims[dim] = 1; - } - } else if !out_dims.is_empty() { - out_dims.remove(dim); - } - - let values_shape = Shape::new(out_dims); - let num_out = values_shape.numel(); - let mut values_data = TensorData::zeros_on_device(num_out, tensor.dtype(), tensor.device()); - - let outer = if dims.is_empty() || dim == 0 { - 1 - } else { - dims[..dim].iter().product() - }; - let inner = if dims.is_empty() || dim + 1 >= dims.len() { - 1 - } else { - dims[dim + 1..].iter().product() - }; - let outer_stride = dim_size * inner; - - match tensor.dtype() { - DataType::Float32 => { - let input = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let values = values_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?; - let mut buffer = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - buffer.clear(); - for d in 0..dim_size { - let value = input[o * outer_stride + d * inner + r]; - if !value.is_nan() { - buffer.push(value); - } - } - let out_idx = o * inner + r; - values[out_idx] = if buffer.is_empty() { - f32::NAN - } else { - quantile_from_unsorted_f32( - &mut buffer, - 0.5, - QuantileInterpolation::Linear, - ) - }; - } - } - } - DataType::Float64 => { - let input = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let values = values_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?; - let mut buffer = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - buffer.clear(); - for d in 0..dim_size { - let value = input[o * outer_stride + d * inner + r]; - if !value.is_nan() { - buffer.push(value); - } - } - let out_idx = o * inner + r; - values[out_idx] = if buffer.is_empty() { - f64::NAN - } else { - quantile_from_unsorted_f64( - &mut buffer, - 0.5, - QuantileInterpolation::Linear, - ) - }; - } - } - } - _ => unreachable!("dtype validated"), - } - - Ok(Tensor::new( - Arc::new(values_data), - values_shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )) -} - -fn quantiles_all( - tensor: &Tensor, - qs: &[f64], - keepdim: bool, - interpolation: QuantileInterpolation, -) -> Result { - let q_len = qs.len(); - let output_dims = quantiles_output_dims(tensor.ndim(), q_len, keepdim); - - let shape = Shape::new(output_dims); - let mut values_data = - TensorData::zeros_on_device(shape.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => { - let data = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let values = values_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?; - if data.len() == 1 { - fill_quantiles_all_single_f32(data[0], values); - return Ok(Tensor::new( - Arc::new(values_data), - shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )); - } - - let Some(mut buffer) = copy_or_none_if_nan(data) else { - values.fill(f32::NAN); - return Ok(Tensor::new( - Arc::new(values_data), - shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )); - }; - - if q_len == 1 { - values[0] = quantile_from_unsorted_f32(&mut buffer, qs[0], interpolation); - return Ok(Tensor::new( - Arc::new(values_data), - shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )); - } - - let positions = quantile_positions_for_len(buffer.len(), qs); - buffer.sort_by(|a, b| a.total_cmp(b)); - quantiles_from_sorted_f32(&buffer, &positions, interpolation, values); - } - DataType::Float64 => { - let data = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let values = values_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?; - if data.len() == 1 { - fill_quantiles_all_single_f64(data[0], values); - return Ok(Tensor::new( - Arc::new(values_data), - shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )); - } - - let Some(mut buffer) = copy_or_none_if_nan(data) else { - values.fill(f64::NAN); - return Ok(Tensor::new( - Arc::new(values_data), - shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )); - }; - - if q_len == 1 { - values[0] = quantile_from_unsorted_f64(&mut buffer, qs[0], interpolation); - return Ok(Tensor::new( - Arc::new(values_data), - shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )); - } - - let positions = quantile_positions_for_len(buffer.len(), qs); - buffer.sort_by(|a, b| a.total_cmp(b)); - quantiles_from_sorted_f64(&buffer, &positions, interpolation, values); - } - _ => unreachable!("dtype validated"), - } - - Ok(Tensor::new( - Arc::new(values_data), - shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )) -} - -fn copy_or_none_if_nan(data: &[T]) -> Option> { - let mut out = Vec::with_capacity(data.len()); - for &value in data { - if value.is_nan() { - return None; - } - out.push(value); - } - Some(out) -} - -fn nanquantiles_all( - tensor: &Tensor, - qs: &[f64], - keepdim: bool, - interpolation: QuantileInterpolation, -) -> Result { - let q_len = qs.len(); - let output_dims = quantiles_output_dims(tensor.ndim(), q_len, keepdim); - - let shape = Shape::new(output_dims); - let mut values_data = - TensorData::zeros_on_device(shape.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => { - let data = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let values = values_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?; - if data.len() == 1 { - fill_nanquantiles_all_single_f32(data[0], values)?; - return Ok(Tensor::new( - Arc::new(values_data), - shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )); - } - if q_len == 1 { - let mut buffer: Vec = data.iter().copied().filter(|v| !v.is_nan()).collect(); - if buffer.is_empty() { - return Err(MinitensorError::invalid_argument( - NANQUANTILE_ALL_NAN_ERR.to_string(), - )); - } - values[0] = quantile_from_unsorted_f32(&mut buffer, qs[0], interpolation); - return Ok(Tensor::new( - Arc::new(values_data), - shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )); - } - let mut sorted: Vec = data.iter().copied().filter(|v| !v.is_nan()).collect(); - if sorted.is_empty() { - return Err(MinitensorError::invalid_argument( - NANQUANTILE_ALL_NAN_ERR.to_string(), - )); - } - let positions = quantile_positions_for_len(sorted.len(), qs); - sorted.sort_by(|a, b| a.total_cmp(b)); - quantiles_from_sorted_f32(&sorted, &positions, interpolation, values); - } - DataType::Float64 => { - let data = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let values = values_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?; - if data.len() == 1 { - fill_nanquantiles_all_single_f64(data[0], values)?; - return Ok(Tensor::new( - Arc::new(values_data), - shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )); - } - if q_len == 1 { - let mut buffer: Vec = data.iter().copied().filter(|v| !v.is_nan()).collect(); - if buffer.is_empty() { - return Err(MinitensorError::invalid_argument( - NANQUANTILE_ALL_NAN_ERR.to_string(), - )); - } - values[0] = quantile_from_unsorted_f64(&mut buffer, qs[0], interpolation); - return Ok(Tensor::new( - Arc::new(values_data), - shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )); - } - let mut sorted: Vec = data.iter().copied().filter(|v| !v.is_nan()).collect(); - if sorted.is_empty() { - return Err(MinitensorError::invalid_argument( - NANQUANTILE_ALL_NAN_ERR.to_string(), - )); - } - let positions = quantile_positions_for_len(sorted.len(), qs); - sorted.sort_by(|a, b| a.total_cmp(b)); - quantiles_from_sorted_f64(&sorted, &positions, interpolation, values); - } - _ => unreachable!("dtype validated"), - } - - Ok(Tensor::new( - Arc::new(values_data), - shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )) -} - -fn quantiles_output_dims(tensor_ndim: usize, q_len: usize, keepdim: bool) -> Vec { - if keepdim && tensor_ndim > 0 { - let mut dims = vec![1; tensor_ndim + 1]; - dims[0] = q_len; - dims - } else { - vec![q_len] - } -} - -fn quantiles_along_dim( - tensor: &Tensor, - dim: usize, - qs: &[f64], - keepdim: bool, - interpolation: QuantileInterpolation, -) -> Result { - let dims = tensor.shape().dims(); - let dim_size = if dims.is_empty() { 1 } else { dims[dim] }; - - if dim_size == 0 { - return Err(MinitensorError::invalid_argument( - "quantile() does not support empty slices".to_string(), - )); - } - - let q_len = qs.len(); - - let mut out_dims = Vec::with_capacity(dims.len() + 2); - out_dims.push(q_len); - if !dims.is_empty() { - out_dims.extend_from_slice(&dims[..dim]); - if keepdim { - out_dims.push(1); - } - out_dims.extend_from_slice(&dims[dim + 1..]); - } else if keepdim { - out_dims.push(1); - } - - let shape = Shape::new(out_dims); - let mut values_data = - TensorData::zeros_on_device(shape.numel(), tensor.dtype(), tensor.device()); - - let outer = if dims.is_empty() || dim == 0 { - 1 - } else { - dims[..dim].iter().product() - }; - let inner = if dims.is_empty() || dim + 1 >= dims.len() { - 1 - } else { - dims[dim + 1..].iter().product() - }; - let outer_stride = dim_size * inner; - - match tensor.dtype() { - DataType::Float32 => { - let input = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let values = values_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?; - - if dim_size == 1 { - fill_quantiles_single_f32(input, values, outer, inner, outer_stride, q_len); - } else { - let mut buffer = Vec::with_capacity(dim_size); - if q_len == 1 { - let q_value = qs[0]; - for o in 0..outer { - for r in 0..inner { - buffer.clear(); - let mut has_nan = false; - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - let value = input[idx]; - if value.is_nan() { - has_nan = true; - break; - } - buffer.push(value); - } - - let out_idx = o * inner + r; - if has_nan { - values[out_idx] = f32::NAN; - continue; - } - - values[out_idx] = - quantile_from_unsorted_f32(&mut buffer, q_value, interpolation); - } - } - } else { - let positions = quantile_positions_for_len(dim_size, qs); - for o in 0..outer { - for r in 0..inner { - buffer.clear(); - let mut has_nan = false; - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - let value = input[idx]; - if value.is_nan() { - has_nan = true; - break; - } - buffer.push(value); - } - - if has_nan { - for qi in 0..q_len { - let out_idx = ((qi * outer) + o) * inner + r; - values[out_idx] = f32::NAN; - } - continue; - } - - buffer.sort_by(|a, b| a.total_cmp(b)); - for (qi, position) in positions.iter().enumerate() { - let out_idx = ((qi * outer) + o) * inner + r; - values[out_idx] = - quantile_from_sorted_position_f32(&buffer, position, interpolation); - } - } - } - } - } - } - DataType::Float64 => { - let input = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let values = values_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?; - - if dim_size == 1 { - fill_quantiles_single_f64(input, values, outer, inner, outer_stride, q_len); - } else { - let mut buffer = Vec::with_capacity(dim_size); - if q_len == 1 { - let q_value = qs[0]; - for o in 0..outer { - for r in 0..inner { - buffer.clear(); - let mut has_nan = false; - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - let value = input[idx]; - if value.is_nan() { - has_nan = true; - break; - } - buffer.push(value); - } - - let out_idx = o * inner + r; - if has_nan { - values[out_idx] = f64::NAN; - continue; - } - - values[out_idx] = - quantile_from_unsorted_f64(&mut buffer, q_value, interpolation); - } - } - } else { - let positions = quantile_positions_for_len(dim_size, qs); - for o in 0..outer { - for r in 0..inner { - buffer.clear(); - let mut has_nan = false; - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - let value = input[idx]; - if value.is_nan() { - has_nan = true; - break; - } - buffer.push(value); - } - - if has_nan { - for qi in 0..q_len { - let out_idx = ((qi * outer) + o) * inner + r; - values[out_idx] = f64::NAN; - } - continue; - } - - buffer.sort_by(|a, b| a.total_cmp(b)); - for (qi, position) in positions.iter().enumerate() { - let out_idx = ((qi * outer) + o) * inner + r; - values[out_idx] = - quantile_from_sorted_position_f64(&buffer, position, interpolation); - } - } - } - } - } - } - _ => unreachable!("dtype validated"), - } - - Ok(Tensor::new( - Arc::new(values_data), - shape, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )) -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::{ + error::{MinitensorError, Result}, + tensor::{DataType, Shape, Tensor, TensorData}, +}; +use std::sync::Arc; + +pub(crate) fn quantile_along_dim( + tensor: &Tensor, + dim: usize, + keepdim: bool, + q: f64, + interpolation: QuantileInterpolation, +) -> Result { + let dims = tensor.shape().dims(); + let dim_size = if dims.is_empty() { 1 } else { dims[dim] }; + + if dim_size == 0 { + return Err(MinitensorError::invalid_argument( + "quantile() does not support reductions over empty dimensions".to_string(), + )); + } + + let mut out_dims = if dims.is_empty() { + vec![1] + } else { + dims.to_vec() + }; + + if keepdim { + if !out_dims.is_empty() { + out_dims[dim] = 1; + } + } else if !out_dims.is_empty() { + out_dims.remove(dim); + } + + let values_shape = Shape::new(out_dims); + let num_out = values_shape.numel(); + let mut values_data = TensorData::zeros_on_device(num_out, tensor.dtype(), tensor.device()); + + let outer = if dims.is_empty() || dim == 0 { + 1 + } else { + dims[..dim].iter().product() + }; + let inner = if dims.is_empty() || dim + 1 >= dims.len() { + 1 + } else { + dims[dim + 1..].iter().product() + }; + let outer_stride = dim_size * inner; + + match tensor.dtype() { + DataType::Float32 => { + let input = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let values = values_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?; + + if dim_size == 1 { + fill_quantile_single_f32(input, values, outer, inner, outer_stride); + } else { + let mut buffer = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + buffer.clear(); + let mut has_nan = false; + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + let value = input[idx]; + if value.is_nan() { + has_nan = true; + break; + } + buffer.push(value); + } + + if has_nan { + values[o * inner + r] = f32::NAN; + continue; + } + + values[o * inner + r] = + quantile_from_unsorted_f32(&mut buffer, q, interpolation); + } + } + } + } + DataType::Float64 => { + let input = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let values = values_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?; + + if dim_size == 1 { + fill_quantile_single_f64(input, values, outer, inner, outer_stride); + } else { + let mut buffer = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + buffer.clear(); + let mut has_nan = false; + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + let value = input[idx]; + if value.is_nan() { + has_nan = true; + break; + } + buffer.push(value); + } + + if has_nan { + values[o * inner + r] = f64::NAN; + continue; + } + + values[o * inner + r] = + quantile_from_unsorted_f64(&mut buffer, q, interpolation); + } + } + } + } + _ => unreachable!("dtype validated"), + } + + Ok(Tensor::new( + Arc::new(values_data), + values_shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )) +} + +pub(crate) fn nanquantile_along_dim( + tensor: &Tensor, + dim: usize, + keepdim: bool, + q: f64, + interpolation: QuantileInterpolation, +) -> Result { + let dims = tensor.shape().dims(); + let dim_size = if dims.is_empty() { 1 } else { dims[dim] }; + + if dim_size == 0 { + return Err(MinitensorError::invalid_argument( + "nanquantile() does not support reductions over empty dimensions".to_string(), + )); + } + + let mut out_dims = if dims.is_empty() { + vec![1] + } else { + dims.to_vec() + }; + + if keepdim { + if !out_dims.is_empty() { + out_dims[dim] = 1; + } + } else if !out_dims.is_empty() { + out_dims.remove(dim); + } + + let values_shape = Shape::new(out_dims); + let num_out = values_shape.numel(); + let mut values_data = TensorData::zeros_on_device(num_out, tensor.dtype(), tensor.device()); + + let outer = if dims.is_empty() || dim == 0 { + 1 + } else { + dims[..dim].iter().product() + }; + let inner = if dims.is_empty() || dim + 1 >= dims.len() { + 1 + } else { + dims[dim + 1..].iter().product() + }; + let outer_stride = dim_size * inner; + + match tensor.dtype() { + DataType::Float32 => { + let input = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let values = values_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?; + + if dim_size == 1 { + fill_nanquantile_single_f32(input, values, outer, inner, outer_stride)?; + } else { + let mut buffer = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + buffer.clear(); + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + let val = input[idx]; + if !val.is_nan() { + buffer.push(val); + } + } + + if buffer.is_empty() { + return Err(MinitensorError::invalid_argument( + NANQUANTILE_ALL_NAN_ERR.to_string(), + )); + } + + let quant = quantile_from_unsorted_f32(&mut buffer, q, interpolation); + values[o * inner + r] = quant; + } + } + } + } + DataType::Float64 => { + let input = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let values = values_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?; + + if dim_size == 1 { + fill_nanquantile_single_f64(input, values, outer, inner, outer_stride)?; + } else { + let mut buffer = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + buffer.clear(); + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + let val = input[idx]; + if !val.is_nan() { + buffer.push(val); + } + } + + if buffer.is_empty() { + return Err(MinitensorError::invalid_argument( + NANQUANTILE_ALL_NAN_ERR.to_string(), + )); + } + + let quant = quantile_from_unsorted_f64(&mut buffer, q, interpolation); + values[o * inner + r] = quant; + } + } + } + } + _ => unreachable!("dtype validated"), + } + + Ok(Tensor::new( + Arc::new(values_data), + values_shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )) +} + +pub(crate) fn nanmedian_all(tensor: &Tensor, keepdim: bool) -> Result { + let output_dims = if keepdim && tensor.ndim() > 0 { + vec![1; tensor.ndim()] + } else { + Vec::new() + }; + let shape = Shape::new(output_dims); + let mut values_data = + TensorData::zeros_on_device(shape.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => { + let data = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let values = values_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?; + let mut buffer: Vec = data.iter().copied().filter(|v| !v.is_nan()).collect(); + values[0] = if buffer.is_empty() { + f32::NAN + } else { + quantile_from_unsorted_f32(&mut buffer, 0.5, QuantileInterpolation::Linear) + }; + } + DataType::Float64 => { + let data = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let values = values_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?; + let mut buffer: Vec = data.iter().copied().filter(|v| !v.is_nan()).collect(); + values[0] = if buffer.is_empty() { + f64::NAN + } else { + quantile_from_unsorted_f64(&mut buffer, 0.5, QuantileInterpolation::Linear) + }; + } + _ => unreachable!("dtype validated"), + } + + Ok(Tensor::new( + Arc::new(values_data), + shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )) +} + +pub(crate) fn nanmedian_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { + let dims = tensor.shape().dims(); + let dim_size = if dims.is_empty() { 1 } else { dims[dim] }; + + let mut out_dims = if dims.is_empty() { + vec![1] + } else { + dims.to_vec() + }; + + if keepdim { + if !out_dims.is_empty() { + out_dims[dim] = 1; + } + } else if !out_dims.is_empty() { + out_dims.remove(dim); + } + + let values_shape = Shape::new(out_dims); + let num_out = values_shape.numel(); + let mut values_data = TensorData::zeros_on_device(num_out, tensor.dtype(), tensor.device()); + + let outer = if dims.is_empty() || dim == 0 { + 1 + } else { + dims[..dim].iter().product() + }; + let inner = if dims.is_empty() || dim + 1 >= dims.len() { + 1 + } else { + dims[dim + 1..].iter().product() + }; + let outer_stride = dim_size * inner; + + match tensor.dtype() { + DataType::Float32 => { + let input = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let values = values_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?; + let mut buffer = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + buffer.clear(); + for d in 0..dim_size { + let value = input[o * outer_stride + d * inner + r]; + if !value.is_nan() { + buffer.push(value); + } + } + let out_idx = o * inner + r; + values[out_idx] = if buffer.is_empty() { + f32::NAN + } else { + quantile_from_unsorted_f32(&mut buffer, 0.5, QuantileInterpolation::Linear) + }; + } + } + } + DataType::Float64 => { + let input = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let values = values_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?; + let mut buffer = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + buffer.clear(); + for d in 0..dim_size { + let value = input[o * outer_stride + d * inner + r]; + if !value.is_nan() { + buffer.push(value); + } + } + let out_idx = o * inner + r; + values[out_idx] = if buffer.is_empty() { + f64::NAN + } else { + quantile_from_unsorted_f64(&mut buffer, 0.5, QuantileInterpolation::Linear) + }; + } + } + } + _ => unreachable!("dtype validated"), + } + + Ok(Tensor::new( + Arc::new(values_data), + values_shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )) +} + +pub(crate) fn quantiles_all( + tensor: &Tensor, + qs: &[f64], + keepdim: bool, + interpolation: QuantileInterpolation, +) -> Result { + let q_len = qs.len(); + let output_dims = quantiles_output_dims(tensor.ndim(), q_len, keepdim); + + let shape = Shape::new(output_dims); + let mut values_data = + TensorData::zeros_on_device(shape.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => { + let data = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let values = values_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?; + if data.len() == 1 { + fill_quantiles_all_single_f32(data[0], values); + return Ok(Tensor::new( + Arc::new(values_data), + shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )); + } + + let Some(mut buffer) = copy_or_none_if_nan(data) else { + values.fill(f32::NAN); + return Ok(Tensor::new( + Arc::new(values_data), + shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )); + }; + + if q_len == 1 { + values[0] = quantile_from_unsorted_f32(&mut buffer, qs[0], interpolation); + return Ok(Tensor::new( + Arc::new(values_data), + shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )); + } + + let positions = quantile_positions_for_len(buffer.len(), qs); + buffer.sort_by(|a, b| a.total_cmp(b)); + quantiles_from_sorted_f32(&buffer, &positions, interpolation, values); + } + DataType::Float64 => { + let data = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let values = values_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?; + if data.len() == 1 { + fill_quantiles_all_single_f64(data[0], values); + return Ok(Tensor::new( + Arc::new(values_data), + shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )); + } + + let Some(mut buffer) = copy_or_none_if_nan(data) else { + values.fill(f64::NAN); + return Ok(Tensor::new( + Arc::new(values_data), + shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )); + }; + + if q_len == 1 { + values[0] = quantile_from_unsorted_f64(&mut buffer, qs[0], interpolation); + return Ok(Tensor::new( + Arc::new(values_data), + shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )); + } + + let positions = quantile_positions_for_len(buffer.len(), qs); + buffer.sort_by(|a, b| a.total_cmp(b)); + quantiles_from_sorted_f64(&buffer, &positions, interpolation, values); + } + _ => unreachable!("dtype validated"), + } + + Ok(Tensor::new( + Arc::new(values_data), + shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )) +} + +fn copy_or_none_if_nan(data: &[T]) -> Option> { + let mut out = Vec::with_capacity(data.len()); + for &value in data { + if value.is_nan() { + return None; + } + out.push(value); + } + Some(out) +} + +pub(crate) fn nanquantiles_all( + tensor: &Tensor, + qs: &[f64], + keepdim: bool, + interpolation: QuantileInterpolation, +) -> Result { + let q_len = qs.len(); + let output_dims = quantiles_output_dims(tensor.ndim(), q_len, keepdim); + + let shape = Shape::new(output_dims); + let mut values_data = + TensorData::zeros_on_device(shape.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => { + let data = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let values = values_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?; + if data.len() == 1 { + fill_nanquantiles_all_single_f32(data[0], values)?; + return Ok(Tensor::new( + Arc::new(values_data), + shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )); + } + if q_len == 1 { + let mut buffer: Vec = data.iter().copied().filter(|v| !v.is_nan()).collect(); + if buffer.is_empty() { + return Err(MinitensorError::invalid_argument( + NANQUANTILE_ALL_NAN_ERR.to_string(), + )); + } + values[0] = quantile_from_unsorted_f32(&mut buffer, qs[0], interpolation); + return Ok(Tensor::new( + Arc::new(values_data), + shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )); + } + let mut sorted: Vec = data.iter().copied().filter(|v| !v.is_nan()).collect(); + if sorted.is_empty() { + return Err(MinitensorError::invalid_argument( + NANQUANTILE_ALL_NAN_ERR.to_string(), + )); + } + let positions = quantile_positions_for_len(sorted.len(), qs); + sorted.sort_by(|a, b| a.total_cmp(b)); + quantiles_from_sorted_f32(&sorted, &positions, interpolation, values); + } + DataType::Float64 => { + let data = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let values = values_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?; + if data.len() == 1 { + fill_nanquantiles_all_single_f64(data[0], values)?; + return Ok(Tensor::new( + Arc::new(values_data), + shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )); + } + if q_len == 1 { + let mut buffer: Vec = data.iter().copied().filter(|v| !v.is_nan()).collect(); + if buffer.is_empty() { + return Err(MinitensorError::invalid_argument( + NANQUANTILE_ALL_NAN_ERR.to_string(), + )); + } + values[0] = quantile_from_unsorted_f64(&mut buffer, qs[0], interpolation); + return Ok(Tensor::new( + Arc::new(values_data), + shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )); + } + let mut sorted: Vec = data.iter().copied().filter(|v| !v.is_nan()).collect(); + if sorted.is_empty() { + return Err(MinitensorError::invalid_argument( + NANQUANTILE_ALL_NAN_ERR.to_string(), + )); + } + let positions = quantile_positions_for_len(sorted.len(), qs); + sorted.sort_by(|a, b| a.total_cmp(b)); + quantiles_from_sorted_f64(&sorted, &positions, interpolation, values); + } + _ => unreachable!("dtype validated"), + } + + Ok(Tensor::new( + Arc::new(values_data), + shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )) +} + +fn quantiles_output_dims(tensor_ndim: usize, q_len: usize, keepdim: bool) -> Vec { + if keepdim && tensor_ndim > 0 { + let mut dims = vec![1; tensor_ndim + 1]; + dims[0] = q_len; + dims + } else { + vec![q_len] + } +} + +pub(crate) fn quantiles_along_dim( + tensor: &Tensor, + dim: usize, + qs: &[f64], + keepdim: bool, + interpolation: QuantileInterpolation, +) -> Result { + let dims = tensor.shape().dims(); + let dim_size = if dims.is_empty() { 1 } else { dims[dim] }; + + if dim_size == 0 { + return Err(MinitensorError::invalid_argument( + "quantile() does not support empty slices".to_string(), + )); + } + + let q_len = qs.len(); + + let mut out_dims = Vec::with_capacity(dims.len() + 2); + out_dims.push(q_len); + if !dims.is_empty() { + out_dims.extend_from_slice(&dims[..dim]); + if keepdim { + out_dims.push(1); + } + out_dims.extend_from_slice(&dims[dim + 1..]); + } else if keepdim { + out_dims.push(1); + } + + let shape = Shape::new(out_dims); + let mut values_data = + TensorData::zeros_on_device(shape.numel(), tensor.dtype(), tensor.device()); + + let outer = if dims.is_empty() || dim == 0 { + 1 + } else { + dims[..dim].iter().product() + }; + let inner = if dims.is_empty() || dim + 1 >= dims.len() { + 1 + } else { + dims[dim + 1..].iter().product() + }; + let outer_stride = dim_size * inner; + + match tensor.dtype() { + DataType::Float32 => { + let input = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let values = values_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?; + + if dim_size == 1 { + fill_quantiles_single_f32(input, values, outer, inner, outer_stride, q_len); + } else { + let mut buffer = Vec::with_capacity(dim_size); + if q_len == 1 { + let q_value = qs[0]; + for o in 0..outer { + for r in 0..inner { + buffer.clear(); + let mut has_nan = false; + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + let value = input[idx]; + if value.is_nan() { + has_nan = true; + break; + } + buffer.push(value); + } + + let out_idx = o * inner + r; + if has_nan { + values[out_idx] = f32::NAN; + continue; + } + + values[out_idx] = + quantile_from_unsorted_f32(&mut buffer, q_value, interpolation); + } + } + } else { + let positions = quantile_positions_for_len(dim_size, qs); + for o in 0..outer { + for r in 0..inner { + buffer.clear(); + let mut has_nan = false; + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + let value = input[idx]; + if value.is_nan() { + has_nan = true; + break; + } + buffer.push(value); + } + + if has_nan { + for qi in 0..q_len { + let out_idx = ((qi * outer) + o) * inner + r; + values[out_idx] = f32::NAN; + } + continue; + } + + buffer.sort_by(|a, b| a.total_cmp(b)); + for (qi, position) in positions.iter().enumerate() { + let out_idx = ((qi * outer) + o) * inner + r; + values[out_idx] = quantile_from_sorted_position_f32( + &buffer, + position, + interpolation, + ); + } + } + } + } + } + } + DataType::Float64 => { + let input = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let values = values_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?; + + if dim_size == 1 { + fill_quantiles_single_f64(input, values, outer, inner, outer_stride, q_len); + } else { + let mut buffer = Vec::with_capacity(dim_size); + if q_len == 1 { + let q_value = qs[0]; + for o in 0..outer { + for r in 0..inner { + buffer.clear(); + let mut has_nan = false; + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + let value = input[idx]; + if value.is_nan() { + has_nan = true; + break; + } + buffer.push(value); + } + + let out_idx = o * inner + r; + if has_nan { + values[out_idx] = f64::NAN; + continue; + } + + values[out_idx] = + quantile_from_unsorted_f64(&mut buffer, q_value, interpolation); + } + } + } else { + let positions = quantile_positions_for_len(dim_size, qs); + for o in 0..outer { + for r in 0..inner { + buffer.clear(); + let mut has_nan = false; + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + let value = input[idx]; + if value.is_nan() { + has_nan = true; + break; + } + buffer.push(value); + } + + if has_nan { + for qi in 0..q_len { + let out_idx = ((qi * outer) + o) * inner + r; + values[out_idx] = f64::NAN; + } + continue; + } + + buffer.sort_by(|a, b| a.total_cmp(b)); + for (qi, position) in positions.iter().enumerate() { + let out_idx = ((qi * outer) + o) * inner + r; + values[out_idx] = quantile_from_sorted_position_f64( + &buffer, + position, + interpolation, + ); + } + } + } + } + } + } + _ => unreachable!("dtype validated"), + } + + Ok(Tensor::new( + Arc::new(values_data), + shape, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )) +} diff --git a/engine/src/operations/reduction/sort.rs b/engine/src/operations/reduction/sort.rs index 26fcff8d..3a35679c 100644 --- a/engine/src/operations/reduction/sort.rs +++ b/engine/src/operations/reduction/sort.rs @@ -1,897 +1,929 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -pub fn sort( - tensor: &Tensor, - dim: Option, - descending: bool, - stable: bool, -) -> Result<(Tensor, Tensor)> { - let ndim = tensor.ndim(); - - let axis = if ndim == 0 { - match dim { - Some(d) if d == 0 || d == -1 => 0, - Some(d) => return Err(MinitensorError::index_error(d, 0, 1)), - None => 0, - } - } else { - let dim_value = dim.unwrap_or(-1); - normalize_dim(dim_value, ndim)? - }; - - if tensor.shape().dims().is_empty() { - let mut values_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); - let mut indices_data = TensorData::zeros_on_device(1, DataType::Int64, tensor.device()); - - match tensor.dtype() { - DataType::Float32 => { - let src = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let dst = values_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?; - dst[0] = src[0]; - } - DataType::Float64 => { - let src = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let dst = values_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?; - dst[0] = src[0]; - } - DataType::Int32 => { - let src = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - let dst = values_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i32 slice") - })?; - dst[0] = src[0]; - } - DataType::Int64 => { - let src = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - let dst = values_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - dst[0] = src[0]; - } - DataType::Bool => { - let src = tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - let dst = values_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice") - })?; - dst[0] = src[0]; - } - } - - let indices = indices_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - indices[0] = 0; - - let values = Tensor::new( - Arc::new(values_data), - Shape::scalar(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - let indices = Tensor::new( - Arc::new(indices_data), - Shape::scalar(), - DataType::Int64, - tensor.device(), - false, - ); - return Ok((values, indices)); - } - - let dims = tensor.shape().dims(); - let dim_size = dims[axis]; - - let mut values_data = - TensorData::zeros_on_device(tensor.numel(), tensor.dtype(), tensor.device()); - let mut indices_data = - TensorData::zeros_on_device(tensor.numel(), DataType::Int64, tensor.device()); - - let outer = if axis == 0 { - 1 - } else { - dims[..axis].iter().product() - }; - let inner = if axis + 1 >= dims.len() { - 1 - } else { - dims[axis + 1..].iter().product() - }; - let outer_stride = dim_size * inner; - - match tensor.dtype() { - DataType::Float32 => { - let input = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let values = values_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f32 slice") - })?; - let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - - let mut entries = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - entries.clear(); - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - entries.push((d, input[idx])); - } - - if stable { - if descending { - entries.sort_by(cmp_f32_desc); - } else { - entries.sort_by(cmp_f32_asc); - } - } else if descending { - entries.sort_unstable_by(cmp_f32_desc); - } else { - entries.sort_unstable_by(cmp_f32_asc); - } - - let base = o * outer_stride + r; - for (j, (index, value)) in entries.iter().enumerate() { - let offset = base + j * inner; - values[offset] = *value; - indices[offset] = *index as i64; - } - } - } - } - DataType::Float64 => { - let input = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let values = values_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable f64 slice") - })?; - let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - - let mut entries = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - entries.clear(); - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - entries.push((d, input[idx])); - } - - if stable { - if descending { - entries.sort_by(cmp_f64_desc); - } else { - entries.sort_by(cmp_f64_asc); - } - } else if descending { - entries.sort_unstable_by(cmp_f64_desc); - } else { - entries.sort_unstable_by(cmp_f64_asc); - } - - let base = o * outer_stride + r; - for (j, (index, value)) in entries.iter().enumerate() { - let offset = base + j * inner; - values[offset] = *value; - indices[offset] = *index as i64; - } - } - } - } - DataType::Int32 => { - let input = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - let values = values_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i32 slice") - })?; - let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - - let mut entries = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - entries.clear(); - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - entries.push((d, input[idx])); - } - - if stable { - if descending { - entries.sort_by(cmp_i32_desc); - } else { - entries.sort_by(cmp_i32_asc); - } - } else if descending { - entries.sort_unstable_by(cmp_i32_desc); - } else { - entries.sort_unstable_by(cmp_i32_asc); - } - - let base = o * outer_stride + r; - for (j, (index, value)) in entries.iter().enumerate() { - let offset = base + j * inner; - values[offset] = *value; - indices[offset] = *index as i64; - } - } - } - } - DataType::Int64 => { - let input = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - let values = values_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - - let mut entries = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - entries.clear(); - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - entries.push((d, input[idx])); - } - - if stable { - if descending { - entries.sort_by(cmp_i64_desc); - } else { - entries.sort_by(cmp_i64_asc); - } - } else if descending { - entries.sort_unstable_by(cmp_i64_desc); - } else { - entries.sort_unstable_by(cmp_i64_asc); - } - - let base = o * outer_stride + r; - for (j, (index, value)) in entries.iter().enumerate() { - let offset = base + j * inner; - values[offset] = *value; - indices[offset] = *index as i64; - } - } - } - } - DataType::Bool => { - let input = tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - let values = values_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable bool slice") - })?; - let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error("Failed to get mutable i64 slice") - })?; - - let mut entries = Vec::with_capacity(dim_size); - for o in 0..outer { - for r in 0..inner { - entries.clear(); - for d in 0..dim_size { - let idx = o * outer_stride + d * inner + r; - entries.push((d, input[idx])); - } - - if stable { - if descending { - entries.sort_by(cmp_bool_desc); - } else { - entries.sort_by(cmp_bool_asc); - } - } else if descending { - entries.sort_unstable_by(cmp_bool_desc); - } else { - entries.sort_unstable_by(cmp_bool_asc); - } - - let base = o * outer_stride + r; - for (j, (index, value)) in entries.iter().enumerate() { - let offset = base + j * inner; - values[offset] = *value; - indices[offset] = *index as i64; - } - } - } - } - } - - let values = Tensor::new( - Arc::new(values_data), - tensor.shape().clone(), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - ); - let indices = Tensor::new( - Arc::new(indices_data), - tensor.shape().clone(), - DataType::Int64, - tensor.device(), - false, - ); - - // `values = gather(input, axis, indices)`; scatter the gradient back. - let values = attach_gather_like_grad(values, tensor, axis, &indices)?; - - Ok((values, indices)) -} - -pub fn argsort( - tensor: &Tensor, - dim: Option, - descending: bool, - stable: bool, -) -> Result { - let (_, indices) = sort(tensor, dim, descending, stable)?; - Ok(indices) -} - -/// Standard deviation along specified dimensions -pub fn std(tensor: &Tensor, dim: Option>, keepdim: bool, unbiased: bool) -> Result { - let variance = var(tensor, dim, keepdim, unbiased)?; - crate::operations::activation::sqrt(&variance) -} - -/// Variance along specified dimensions -pub fn var(tensor: &Tensor, dim: Option>, keepdim: bool, unbiased: bool) -> Result { - if !tensor.dtype().is_float() { - return Err(MinitensorError::invalid_operation( - "Variance only supported for floating point tensors", - )); - } - - let dims = match dim { - Some(dims) => { - let ndim = tensor.ndim() as isize; - let mut normalized = Vec::with_capacity(dims.len()); - for d in dims { - let d = if d < 0 { d + ndim } else { d }; - if d < 0 || d >= ndim { - return Err(MinitensorError::index_error(d, 0, tensor.ndim())); - } - normalized.push(d as usize); - } - normalized.sort_unstable(); - normalized.dedup(); - Some(normalized) - } - None => None, - }; - - if matches!(dims, Some(ref dims) if dims.is_empty()) { - return Ok(tensor.clone()); - } - - let reduction_dims: Vec = dims - .clone() - .unwrap_or_else(|| (0..tensor.ndim()).collect()); - let reduction_dims_isize: Vec = reduction_dims.iter().map(|&d| d as isize).collect(); - - // Keep reduced axes while computing deviations so broadcasting is unambiguous for - // both single-axis and multi-axis reductions. - let mean_tensor = mean(tensor, Some(reduction_dims_isize.clone()), true)?; - let diff = crate::operations::arithmetic::sub(tensor, &mean_tensor)?; - let squared_diff = crate::operations::arithmetic::mul(&diff, &diff)?; - let mut variance = mean(&squared_diff, Some(reduction_dims_isize), true)?; - - let sample_count = reduction_dims - .iter() - .map(|&axis| tensor.shape().dims()[axis]) - .product::(); - - if unbiased { - if sample_count <= 1 { - let nan_count = variance.numel(); - let nan_data = match variance.dtype() { - DataType::Float32 => TensorData::from_vec_f32(vec![f32::NAN; nan_count], variance.device()), - DataType::Float64 => TensorData::from_vec_f64(vec![f64::NAN; nan_count], variance.device()), - _ => unreachable!("variance is only defined for floating point tensors"), - }; - variance = Tensor::new( - Arc::new(nan_data), - variance.shape().clone(), - variance.dtype(), - variance.device(), - variance.requires_grad(), - ); - } else { - let correction = sample_count as f64 / (sample_count - 1) as f64; - let correction_tensor = match variance.dtype() { - DataType::Float32 => Tensor::new( - Arc::new(TensorData::from_vec_f32(vec![correction as f32], variance.device())), - Shape::scalar(), - DataType::Float32, - variance.device(), - false, - ), - DataType::Float64 => Tensor::new( - Arc::new(TensorData::from_vec_f64(vec![correction], variance.device())), - Shape::scalar(), - DataType::Float64, - variance.device(), - false, - ), - _ => unreachable!("variance is only defined for floating point tensors"), - }; - variance = crate::operations::arithmetic::mul(&variance, &correction_tensor)?; - } - } - - if keepdim { - return Ok(variance); - } - - let mut new_dims = Vec::with_capacity(variance.ndim().saturating_sub(reduction_dims.len())); - for (idx, &size) in variance.shape().dims().iter().enumerate() { - if reduction_dims.binary_search(&idx).is_err() { - new_dims.push(size); - } - } - let target_shape = if new_dims.is_empty() { - Shape::scalar() - } else { - Shape::new(new_dims) - }; - shape_ops::reshape(&variance, target_shape) -} - -// Helper functions for type-specific operations - -fn prod_all_f32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - - let prod: f32 = if data.len() >= 1024 { - data.par_chunks(8192).map(simd_prod_f32).product::() - } else { - simd_prod_f32(data) - }; - - let result_slice = result_data - .as_f32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; - - result_slice[0] = prod; - Ok(()) -} - -fn prod_all_f64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - - let prod: f64 = if data.len() >= 1024 { - data.par_chunks(8192).map(simd_prod_f64).product::() - } else { - simd_prod_f64(data) - }; - - let result_slice = result_data - .as_f64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; - - result_slice[0] = prod; - Ok(()) -} - -fn prod_all_i32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - - let prod: i32 = if data.len() >= 1024 { - data.par_chunks(8192).map(simd_prod_i32).product::() - } else { - simd_prod_i32(data) - }; - - let result_slice = result_data - .as_i32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i32 slice"))?; - - result_slice[0] = prod; - Ok(()) -} - -fn prod_all_i64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - - let prod: i64 = if data.len() >= 1024 { - data.par_chunks(8192).map(simd_prod_i64).product::() - } else { - simd_prod_i64(data) - }; - - let result_slice = result_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - result_slice[0] = prod; - Ok(()) -} - -fn prod_all_bool(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - - let prod = data.par_iter().all(|&x| x); - - let result_slice = result_data - .as_bool_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable bool slice"))?; - - result_slice[0] = prod; - Ok(()) -} - -fn sum_all_f32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - - let sum: f32 = if data.len() >= 1024 { - data.par_chunks(8192).map(simd_sum_f32).sum::() - } else { - simd_sum_f32(data) - }; - - let result_slice = result_data - .as_f32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; - - result_slice[0] = sum; - Ok(()) -} - -fn sum_all_f64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - - let sum: f64 = if data.len() >= 1024 { - data.par_chunks(8192).map(simd_sum_f64).sum::() - } else { - simd_sum_f64(data) - }; - - let result_slice = result_data - .as_f64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; - - result_slice[0] = sum; - Ok(()) -} - -fn sum_all_i32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - - let sum: i32 = if data.len() >= 1024 { - data.par_chunks(8192).map(simd_sum_i32).sum::() - } else { - simd_sum_i32(data) - }; - - let result_slice = result_data - .as_i32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i32 slice"))?; - - result_slice[0] = sum; - Ok(()) -} - -fn sum_all_i64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - - let sum: i64 = if data.len() >= 1024 { - data.par_chunks(8192).map(simd_sum_i64).sum::() - } else { - simd_sum_i64(data) - }; - - let result_slice = result_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - result_slice[0] = sum; - Ok(()) -} - -fn nansum_all_f32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - - let sum: f32 = data - .par_iter() - .map(|&v| if v.is_nan() { 0.0 } else { v }) - .sum(); - - let result_slice = result_data - .as_f32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; - result_slice[0] = sum; - Ok(()) -} - -fn nansum_all_f64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - - let sum: f64 = data - .par_iter() - .map(|&v| if v.is_nan() { 0.0 } else { v }) - .sum(); - - let result_slice = result_data - .as_f64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; - result_slice[0] = sum; - Ok(()) -} - -fn nanmean_all_f32( - tensor: &Tensor, - sum_data: &mut TensorData, - count_data: &mut TensorData, -) -> Result<()> { - let data = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - - let (sum, count) = data - .par_iter() - .map(|&v| { - if v.is_nan() { - (0.0, 0usize) - } else { - (v, 1usize) - } - }) - .reduce(|| (0.0, 0usize), |(s1, c1), (s2, c2)| (s1 + s2, c1 + c2)); - - let sum_slice = sum_data - .as_f32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; - let count_slice = count_data - .as_f32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; - - sum_slice[0] = sum; - count_slice[0] = count as f32; - Ok(()) -} - -fn nanmean_all_f64( - tensor: &Tensor, - sum_data: &mut TensorData, - count_data: &mut TensorData, -) -> Result<()> { - let data = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - - let (sum, count) = data - .par_iter() - .map(|&v| { - if v.is_nan() { - (0.0, 0usize) - } else { - (v, 1usize) - } - }) - .reduce(|| (0.0, 0usize), |(s1, c1), (s2, c2)| (s1 + s2, c1 + c2)); - - let sum_slice = sum_data - .as_f64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; - let count_slice = count_data - .as_f64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; - - sum_slice[0] = sum; - count_slice[0] = count as f64; - Ok(()) -} - -fn nanmean_from_sum_count(sum: &Tensor, count: &Tensor, requires_grad: bool) -> Result { - if sum.dtype() != count.dtype() || sum.shape() != count.shape() { - return Err(MinitensorError::invalid_operation( - "nanmean requires sum and count tensors with matching dtype and shape", - )); - } - - let numel = sum.numel(); - let mut result_data = TensorData::zeros_on_device(numel, sum.dtype(), sum.device()); - - match sum.dtype() { - DataType::Float32 => { - let sum_slice = sum - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let count_slice = count - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let out = result_data - .as_f32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - out.par_iter_mut() - .zip(sum_slice.par_iter().zip(count_slice.par_iter())) - .for_each(|(dst, (&s, &c))| { - *dst = if c == 0.0 { f32::NAN } else { s / c }; - }); - } - DataType::Float64 => { - let sum_slice = sum - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let count_slice = count - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let out = result_data - .as_f64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - out.par_iter_mut() - .zip(sum_slice.par_iter().zip(count_slice.par_iter())) - .for_each(|(dst, (&s, &c))| { - *dst = if c == 0.0 { f64::NAN } else { s / c }; - }); - } - _ => { - return Err(MinitensorError::invalid_operation( - "nanmean only supports floating point tensors", - )); - } - } - - Ok(Tensor::new( - Arc::new(result_data), - sum.shape().clone(), - sum.dtype(), - sum.device(), - requires_grad, - )) -} - -#[inline] -pub fn nansum_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { - if dim >= tensor.ndim() { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - - let input_shape = tensor.shape().dims(); - let mut output_shape = input_shape.to_vec(); - - if keepdim { - output_shape[dim] = 1; - } else { - output_shape.remove(dim); - } - - let output_shape_obj = Shape::new(output_shape); - let mut result_data = - TensorData::zeros_on_device(output_shape_obj.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => nansum_along_dim_f32(tensor, &mut result_data, dim)?, - DataType::Float64 => nansum_along_dim_f64(tensor, &mut result_data, dim)?, - _ => { - return Err(MinitensorError::invalid_operation( - "nansum only supports floating point tensors", - )); - } - } - - Ok(Tensor::new( - Arc::new(result_data), - output_shape_obj, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )) -} - -#[inline] -pub fn sum_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { - if dim >= tensor.ndim() { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - - let input_shape = tensor.shape().dims(); - let mut output_shape = input_shape.to_vec(); - - if keepdim { - output_shape[dim] = 1; - } else { - output_shape.remove(dim); - } - - let output_shape_obj = Shape::new(output_shape); - let mut result_data = - TensorData::zeros_on_device(output_shape_obj.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => sum_along_dim_f32(tensor, &mut result_data, dim)?, - DataType::Float64 => sum_along_dim_f64(tensor, &mut result_data, dim)?, - DataType::Int32 => sum_along_dim_i32(tensor, &mut result_data, dim)?, - DataType::Int64 => sum_along_dim_i64(tensor, &mut result_data, dim)?, - DataType::Bool => { - return Err(MinitensorError::invalid_operation( - "Sum not supported for boolean tensors", - )); - } - } - - Ok(Tensor::new( - Arc::new(result_data), - output_shape_obj, - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )) -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::operations::shape_ops; +use crate::operations::simd::*; +use crate::{ + error::{MinitensorError, Result}, + tensor::{DataType, Shape, Tensor, TensorData}, +}; +use rayon::prelude::*; +use std::sync::Arc; + +pub fn sort( + tensor: &Tensor, + dim: Option, + descending: bool, + stable: bool, +) -> Result<(Tensor, Tensor)> { + let ndim = tensor.ndim(); + + let axis = if ndim == 0 { + match dim { + Some(d) if d == 0 || d == -1 => 0, + Some(d) => return Err(MinitensorError::index_error(d, 0, 1)), + None => 0, + } + } else { + let dim_value = dim.unwrap_or(-1); + normalize_dim(dim_value, ndim)? + }; + + if tensor.shape().dims().is_empty() { + let mut values_data = TensorData::zeros_on_device(1, tensor.dtype(), tensor.device()); + let mut indices_data = TensorData::zeros_on_device(1, DataType::Int64, tensor.device()); + + match tensor.dtype() { + DataType::Float32 => { + let src = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let dst = values_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?; + dst[0] = src[0]; + } + DataType::Float64 => { + let src = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let dst = values_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?; + dst[0] = src[0]; + } + DataType::Int32 => { + let src = tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; + let dst = values_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i32 slice") + })?; + dst[0] = src[0]; + } + DataType::Int64 => { + let src = tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; + let dst = values_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + dst[0] = src[0]; + } + DataType::Bool => { + let src = tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + let dst = values_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice") + })?; + dst[0] = src[0]; + } + } + + let indices = indices_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + indices[0] = 0; + + let values = Tensor::new( + Arc::new(values_data), + Shape::scalar(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + let indices = Tensor::new( + Arc::new(indices_data), + Shape::scalar(), + DataType::Int64, + tensor.device(), + false, + ); + return Ok((values, indices)); + } + + let dims = tensor.shape().dims(); + let dim_size = dims[axis]; + + let mut values_data = + TensorData::zeros_on_device(tensor.numel(), tensor.dtype(), tensor.device()); + let mut indices_data = + TensorData::zeros_on_device(tensor.numel(), DataType::Int64, tensor.device()); + + let outer = if axis == 0 { + 1 + } else { + dims[..axis].iter().product() + }; + let inner = if axis + 1 >= dims.len() { + 1 + } else { + dims[axis + 1..].iter().product() + }; + let outer_stride = dim_size * inner; + + match tensor.dtype() { + DataType::Float32 => { + let input = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let values = values_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f32 slice") + })?; + let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + + let mut entries = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + entries.clear(); + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + entries.push((d, input[idx])); + } + + if stable { + if descending { + entries.sort_by(cmp_f32_desc); + } else { + entries.sort_by(cmp_f32_asc); + } + } else if descending { + entries.sort_unstable_by(cmp_f32_desc); + } else { + entries.sort_unstable_by(cmp_f32_asc); + } + + let base = o * outer_stride + r; + for (j, (index, value)) in entries.iter().enumerate() { + let offset = base + j * inner; + values[offset] = *value; + indices[offset] = *index as i64; + } + } + } + } + DataType::Float64 => { + let input = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let values = values_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable f64 slice") + })?; + let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + + let mut entries = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + entries.clear(); + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + entries.push((d, input[idx])); + } + + if stable { + if descending { + entries.sort_by(cmp_f64_desc); + } else { + entries.sort_by(cmp_f64_asc); + } + } else if descending { + entries.sort_unstable_by(cmp_f64_desc); + } else { + entries.sort_unstable_by(cmp_f64_asc); + } + + let base = o * outer_stride + r; + for (j, (index, value)) in entries.iter().enumerate() { + let offset = base + j * inner; + values[offset] = *value; + indices[offset] = *index as i64; + } + } + } + } + DataType::Int32 => { + let input = tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; + let values = values_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i32 slice") + })?; + let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + + let mut entries = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + entries.clear(); + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + entries.push((d, input[idx])); + } + + if stable { + if descending { + entries.sort_by(cmp_i32_desc); + } else { + entries.sort_by(cmp_i32_asc); + } + } else if descending { + entries.sort_unstable_by(cmp_i32_desc); + } else { + entries.sort_unstable_by(cmp_i32_asc); + } + + let base = o * outer_stride + r; + for (j, (index, value)) in entries.iter().enumerate() { + let offset = base + j * inner; + values[offset] = *value; + indices[offset] = *index as i64; + } + } + } + } + DataType::Int64 => { + let input = tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; + let values = values_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + + let mut entries = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + entries.clear(); + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + entries.push((d, input[idx])); + } + + if stable { + if descending { + entries.sort_by(cmp_i64_desc); + } else { + entries.sort_by(cmp_i64_asc); + } + } else if descending { + entries.sort_unstable_by(cmp_i64_desc); + } else { + entries.sort_unstable_by(cmp_i64_asc); + } + + let base = o * outer_stride + r; + for (j, (index, value)) in entries.iter().enumerate() { + let offset = base + j * inner; + values[offset] = *value; + indices[offset] = *index as i64; + } + } + } + } + DataType::Bool => { + let input = tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + let values = values_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable bool slice") + })?; + let indices = indices_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error("Failed to get mutable i64 slice") + })?; + + let mut entries = Vec::with_capacity(dim_size); + for o in 0..outer { + for r in 0..inner { + entries.clear(); + for d in 0..dim_size { + let idx = o * outer_stride + d * inner + r; + entries.push((d, input[idx])); + } + + if stable { + if descending { + entries.sort_by(cmp_bool_desc); + } else { + entries.sort_by(cmp_bool_asc); + } + } else if descending { + entries.sort_unstable_by(cmp_bool_desc); + } else { + entries.sort_unstable_by(cmp_bool_asc); + } + + let base = o * outer_stride + r; + for (j, (index, value)) in entries.iter().enumerate() { + let offset = base + j * inner; + values[offset] = *value; + indices[offset] = *index as i64; + } + } + } + } + } + + let values = Tensor::new( + Arc::new(values_data), + tensor.shape().clone(), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + ); + let indices = Tensor::new( + Arc::new(indices_data), + tensor.shape().clone(), + DataType::Int64, + tensor.device(), + false, + ); + + // `values = gather(input, axis, indices)`; scatter the gradient back. + let values = attach_gather_like_grad(values, tensor, axis, &indices)?; + + Ok((values, indices)) +} + +pub fn argsort( + tensor: &Tensor, + dim: Option, + descending: bool, + stable: bool, +) -> Result { + let (_, indices) = sort(tensor, dim, descending, stable)?; + Ok(indices) +} + +/// Standard deviation along specified dimensions +pub fn std( + tensor: &Tensor, + dim: Option>, + keepdim: bool, + unbiased: bool, +) -> Result { + let variance = var(tensor, dim, keepdim, unbiased)?; + crate::operations::activation::sqrt(&variance) +} + +/// Variance along specified dimensions +pub fn var( + tensor: &Tensor, + dim: Option>, + keepdim: bool, + unbiased: bool, +) -> Result { + if !tensor.dtype().is_float() { + return Err(MinitensorError::invalid_operation( + "Variance only supported for floating point tensors", + )); + } + + let dims = match dim { + Some(dims) => { + let ndim = tensor.ndim() as isize; + let mut normalized = Vec::with_capacity(dims.len()); + for d in dims { + let d = if d < 0 { d + ndim } else { d }; + if d < 0 || d >= ndim { + return Err(MinitensorError::index_error(d, 0, tensor.ndim())); + } + normalized.push(d as usize); + } + normalized.sort_unstable(); + normalized.dedup(); + Some(normalized) + } + None => None, + }; + + if matches!(dims, Some(ref dims) if dims.is_empty()) { + return Ok(tensor.clone()); + } + + let reduction_dims: Vec = dims.clone().unwrap_or_else(|| (0..tensor.ndim()).collect()); + let reduction_dims_isize: Vec = reduction_dims.iter().map(|&d| d as isize).collect(); + + // Keep reduced axes while computing deviations so broadcasting is unambiguous for + // both single-axis and multi-axis reductions. + let mean_tensor = mean(tensor, Some(reduction_dims_isize.clone()), true)?; + let diff = crate::operations::arithmetic::sub(tensor, &mean_tensor)?; + let squared_diff = crate::operations::arithmetic::mul(&diff, &diff)?; + let mut variance = mean(&squared_diff, Some(reduction_dims_isize), true)?; + + let sample_count = reduction_dims + .iter() + .map(|&axis| tensor.shape().dims()[axis]) + .product::(); + + if unbiased { + if sample_count <= 1 { + let nan_count = variance.numel(); + let nan_data = match variance.dtype() { + DataType::Float32 => { + TensorData::from_vec_f32(vec![f32::NAN; nan_count], variance.device()) + } + DataType::Float64 => { + TensorData::from_vec_f64(vec![f64::NAN; nan_count], variance.device()) + } + _ => unreachable!("variance is only defined for floating point tensors"), + }; + variance = Tensor::new( + Arc::new(nan_data), + variance.shape().clone(), + variance.dtype(), + variance.device(), + variance.requires_grad(), + ); + } else { + let correction = sample_count as f64 / (sample_count - 1) as f64; + let correction_tensor = match variance.dtype() { + DataType::Float32 => Tensor::new( + Arc::new(TensorData::from_vec_f32( + vec![correction as f32], + variance.device(), + )), + Shape::scalar(), + DataType::Float32, + variance.device(), + false, + ), + DataType::Float64 => Tensor::new( + Arc::new(TensorData::from_vec_f64( + vec![correction], + variance.device(), + )), + Shape::scalar(), + DataType::Float64, + variance.device(), + false, + ), + _ => unreachable!("variance is only defined for floating point tensors"), + }; + variance = crate::operations::arithmetic::mul(&variance, &correction_tensor)?; + } + } + + if keepdim { + return Ok(variance); + } + + let mut new_dims = Vec::with_capacity(variance.ndim().saturating_sub(reduction_dims.len())); + for (idx, &size) in variance.shape().dims().iter().enumerate() { + if reduction_dims.binary_search(&idx).is_err() { + new_dims.push(size); + } + } + let target_shape = if new_dims.is_empty() { + Shape::scalar() + } else { + Shape::new(new_dims) + }; + shape_ops::reshape(&variance, target_shape) +} + +// Helper functions for type-specific operations + +pub(crate) fn prod_all_f32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + + let prod: f32 = if data.len() >= 1024 { + data.par_chunks(8192).map(simd_prod_f32).product::() + } else { + simd_prod_f32(data) + }; + + let result_slice = result_data + .as_f32_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; + + result_slice[0] = prod; + Ok(()) +} + +pub(crate) fn prod_all_f64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + + let prod: f64 = if data.len() >= 1024 { + data.par_chunks(8192).map(simd_prod_f64).product::() + } else { + simd_prod_f64(data) + }; + + let result_slice = result_data + .as_f64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; + + result_slice[0] = prod; + Ok(()) +} + +pub(crate) fn prod_all_i32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; + + let prod: i32 = if data.len() >= 1024 { + data.par_chunks(8192).map(simd_prod_i32).product::() + } else { + simd_prod_i32(data) + }; + + let result_slice = result_data + .as_i32_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i32 slice"))?; + + result_slice[0] = prod; + Ok(()) +} + +pub(crate) fn prod_all_i64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; + + let prod: i64 = if data.len() >= 1024 { + data.par_chunks(8192).map(simd_prod_i64).product::() + } else { + simd_prod_i64(data) + }; + + let result_slice = result_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + result_slice[0] = prod; + Ok(()) +} + +pub(crate) fn prod_all_bool(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + + let prod = data.par_iter().all(|&x| x); + + let result_slice = result_data + .as_bool_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable bool slice"))?; + + result_slice[0] = prod; + Ok(()) +} + +pub(crate) fn sum_all_f32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + + let sum: f32 = if data.len() >= 1024 { + data.par_chunks(8192).map(simd_sum_f32).sum::() + } else { + simd_sum_f32(data) + }; + + let result_slice = result_data + .as_f32_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; + + result_slice[0] = sum; + Ok(()) +} + +pub(crate) fn sum_all_f64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + + let sum: f64 = if data.len() >= 1024 { + data.par_chunks(8192).map(simd_sum_f64).sum::() + } else { + simd_sum_f64(data) + }; + + let result_slice = result_data + .as_f64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; + + result_slice[0] = sum; + Ok(()) +} + +pub(crate) fn sum_all_i32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; + + let sum: i32 = if data.len() >= 1024 { + data.par_chunks(8192).map(simd_sum_i32).sum::() + } else { + simd_sum_i32(data) + }; + + let result_slice = result_data + .as_i32_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i32 slice"))?; + + result_slice[0] = sum; + Ok(()) +} + +pub(crate) fn sum_all_i64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; + + let sum: i64 = if data.len() >= 1024 { + data.par_chunks(8192).map(simd_sum_i64).sum::() + } else { + simd_sum_i64(data) + }; + + let result_slice = result_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + result_slice[0] = sum; + Ok(()) +} + +pub(crate) fn nansum_all_f32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + + let sum: f32 = data + .par_iter() + .map(|&v| if v.is_nan() { 0.0 } else { v }) + .sum(); + + let result_slice = result_data + .as_f32_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; + result_slice[0] = sum; + Ok(()) +} + +pub(crate) fn nansum_all_f64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + + let sum: f64 = data + .par_iter() + .map(|&v| if v.is_nan() { 0.0 } else { v }) + .sum(); + + let result_slice = result_data + .as_f64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; + result_slice[0] = sum; + Ok(()) +} + +pub(crate) fn nanmean_all_f32( + tensor: &Tensor, + sum_data: &mut TensorData, + count_data: &mut TensorData, +) -> Result<()> { + let data = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + + let (sum, count) = data + .par_iter() + .map(|&v| { + if v.is_nan() { + (0.0, 0usize) + } else { + (v, 1usize) + } + }) + .reduce(|| (0.0, 0usize), |(s1, c1), (s2, c2)| (s1 + s2, c1 + c2)); + + let sum_slice = sum_data + .as_f32_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; + let count_slice = count_data + .as_f32_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; + + sum_slice[0] = sum; + count_slice[0] = count as f32; + Ok(()) +} + +pub(crate) fn nanmean_all_f64( + tensor: &Tensor, + sum_data: &mut TensorData, + count_data: &mut TensorData, +) -> Result<()> { + let data = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + + let (sum, count) = data + .par_iter() + .map(|&v| { + if v.is_nan() { + (0.0, 0usize) + } else { + (v, 1usize) + } + }) + .reduce(|| (0.0, 0usize), |(s1, c1), (s2, c2)| (s1 + s2, c1 + c2)); + + let sum_slice = sum_data + .as_f64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; + let count_slice = count_data + .as_f64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; + + sum_slice[0] = sum; + count_slice[0] = count as f64; + Ok(()) +} + +pub(crate) fn nanmean_from_sum_count( + sum: &Tensor, + count: &Tensor, + requires_grad: bool, +) -> Result { + if sum.dtype() != count.dtype() || sum.shape() != count.shape() { + return Err(MinitensorError::invalid_operation( + "nanmean requires sum and count tensors with matching dtype and shape", + )); + } + + let numel = sum.numel(); + let mut result_data = TensorData::zeros_on_device(numel, sum.dtype(), sum.device()); + + match sum.dtype() { + DataType::Float32 => { + let sum_slice = sum + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let count_slice = count + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + let out = result_data + .as_f32_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + out.par_iter_mut() + .zip(sum_slice.par_iter().zip(count_slice.par_iter())) + .for_each(|(dst, (&s, &c))| { + *dst = if c == 0.0 { f32::NAN } else { s / c }; + }); + } + DataType::Float64 => { + let sum_slice = sum + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let count_slice = count + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + let out = result_data + .as_f64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + out.par_iter_mut() + .zip(sum_slice.par_iter().zip(count_slice.par_iter())) + .for_each(|(dst, (&s, &c))| { + *dst = if c == 0.0 { f64::NAN } else { s / c }; + }); + } + _ => { + return Err(MinitensorError::invalid_operation( + "nanmean only supports floating point tensors", + )); + } + } + + Ok(Tensor::new( + Arc::new(result_data), + sum.shape().clone(), + sum.dtype(), + sum.device(), + requires_grad, + )) +} + +#[inline] +pub fn nansum_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { + if dim >= tensor.ndim() { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + + let input_shape = tensor.shape().dims(); + let mut output_shape = input_shape.to_vec(); + + if keepdim { + output_shape[dim] = 1; + } else { + output_shape.remove(dim); + } + + let output_shape_obj = Shape::new(output_shape); + let mut result_data = + TensorData::zeros_on_device(output_shape_obj.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => nansum_along_dim_f32(tensor, &mut result_data, dim)?, + DataType::Float64 => nansum_along_dim_f64(tensor, &mut result_data, dim)?, + _ => { + return Err(MinitensorError::invalid_operation( + "nansum only supports floating point tensors", + )); + } + } + + Ok(Tensor::new( + Arc::new(result_data), + output_shape_obj, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )) +} + +#[inline] +pub fn sum_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { + if dim >= tensor.ndim() { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + + let input_shape = tensor.shape().dims(); + let mut output_shape = input_shape.to_vec(); + + if keepdim { + output_shape[dim] = 1; + } else { + output_shape.remove(dim); + } + + let output_shape_obj = Shape::new(output_shape); + let mut result_data = + TensorData::zeros_on_device(output_shape_obj.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => sum_along_dim_f32(tensor, &mut result_data, dim)?, + DataType::Float64 => sum_along_dim_f64(tensor, &mut result_data, dim)?, + DataType::Int32 => sum_along_dim_i32(tensor, &mut result_data, dim)?, + DataType::Int64 => sum_along_dim_i64(tensor, &mut result_data, dim)?, + DataType::Bool => { + return Err(MinitensorError::invalid_operation( + "Sum not supported for boolean tensors", + )); + } + } + + Ok(Tensor::new( + Arc::new(result_data), + output_shape_obj, + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )) +} diff --git a/engine/src/operations/reduction/sum_prod.rs b/engine/src/operations/reduction/sum_prod.rs index 2d75e5c2..870cc7a9 100644 --- a/engine/src/operations/reduction/sum_prod.rs +++ b/engine/src/operations/reduction/sum_prod.rs @@ -1,905 +1,890 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -fn sum_along_dim_f32(tensor: &Tensor, result_data: &mut TensorData, dim: usize) -> Result<()> { - let input_data = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - - let result_slice = result_data - .as_f32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; - - let input_shape = tensor.shape().dims(); - - if tensor.ndim() == 1 { - if dim != 0 { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - result_slice[0] = simd_sum_f32(input_data); - } else if tensor.ndim() == 2 { - let cols = input_shape[1]; - match dim { - 0 => { - let sums = input_data - .par_chunks_exact(cols) - .fold( - || vec![0f32; cols], - |mut acc, row| { - for (a, &v) in acc.iter_mut().zip(row) { - *a += v; - } - acc - }, - ) - .reduce( - || vec![0f32; cols], - |mut a, b| { - for (x, y) in a.iter_mut().zip(b) { - *x += y; - } - a - }, - ); - result_slice.copy_from_slice(&sums); - } - 1 => { - result_slice - .par_iter_mut() - .zip(input_data.par_chunks_exact(cols)) - .for_each(|(out, row)| { - *out = simd_sum_f32(row); - }); - } - _ => { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - } - } else { - let dim_size = input_shape[dim]; - let inner = input_shape[dim + 1..].iter().product::(); - let outer_stride = dim_size * inner; - - result_slice - .par_iter_mut() - .enumerate() - .for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut sum_val = 0f32; - let mut base = o * outer_stride + r; - for _ in 0..dim_size { - sum_val += input_data[base]; - base += inner; - } - *out = sum_val; - }); - } - - Ok(()) -} - -fn nansum_along_dim_f32(tensor: &Tensor, result_data: &mut TensorData, dim: usize) -> Result<()> { - let input_data = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - - let result_slice = result_data - .as_f32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; - - let input_shape = tensor.shape().dims(); - - if tensor.ndim() == 1 { - if dim != 0 { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - result_slice[0] = input_data.iter().filter(|v| !v.is_nan()).sum::(); - } else if tensor.ndim() == 2 { - let cols = input_shape[1]; - match dim { - 0 => { - let sums = input_data - .par_chunks_exact(cols) - .fold( - || vec![0f32; cols], - |mut acc, row| { - for (a, &v) in acc.iter_mut().zip(row) { - if !v.is_nan() { - *a += v; - } - } - acc - }, - ) - .reduce( - || vec![0f32; cols], - |mut a, b| { - for (x, y) in a.iter_mut().zip(b) { - *x += y; - } - a - }, - ); - result_slice.copy_from_slice(&sums); - } - 1 => { - result_slice - .par_iter_mut() - .zip(input_data.par_chunks_exact(cols)) - .for_each(|(out, row)| { - *out = row.iter().filter(|v| !v.is_nan()).sum::(); - }); - } - _ => { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - } - } else { - let dim_size = input_shape[dim]; - let inner = input_shape[dim + 1..].iter().product::(); - let outer_stride = dim_size * inner; - - result_slice - .par_iter_mut() - .enumerate() - .for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut sum_val = 0f32; - let mut base = o * outer_stride + r; - for _ in 0..dim_size { - let value = input_data[base]; - if !value.is_nan() { - sum_val += value; - } - base += inner; - } - *out = sum_val; - }); - } - - Ok(()) -} - -fn sum_along_dim_f64(tensor: &Tensor, result_data: &mut TensorData, dim: usize) -> Result<()> { - let input_data = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - - let result_slice = result_data - .as_f64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; - - let input_shape = tensor.shape().dims(); - - if tensor.ndim() == 1 { - if dim != 0 { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - result_slice[0] = simd_sum_f64(input_data); - } else if tensor.ndim() == 2 { - let cols = input_shape[1]; - match dim { - 0 => { - let sums = input_data - .par_chunks_exact(cols) - .fold( - || vec![0f64; cols], - |mut acc, row| { - for (a, &v) in acc.iter_mut().zip(row) { - *a += v; - } - acc - }, - ) - .reduce( - || vec![0f64; cols], - |mut a, b| { - for (x, y) in a.iter_mut().zip(b) { - *x += y; - } - a - }, - ); - result_slice.copy_from_slice(&sums); - } - 1 => { - result_slice - .par_iter_mut() - .zip(input_data.par_chunks_exact(cols)) - .for_each(|(out, row)| { - *out = simd_sum_f64(row); - }); - } - _ => { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - } - } else { - let dim_size = input_shape[dim]; - let inner = input_shape[dim + 1..].iter().product::(); - let outer_stride = dim_size * inner; - - result_slice - .par_iter_mut() - .enumerate() - .for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut sum_val = 0f64; - let mut base = o * outer_stride + r; - for _ in 0..dim_size { - sum_val += input_data[base]; - base += inner; - } - *out = sum_val; - }); - } - - Ok(()) -} - -fn nansum_along_dim_f64(tensor: &Tensor, result_data: &mut TensorData, dim: usize) -> Result<()> { - let input_data = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - - let result_slice = result_data - .as_f64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; - - let input_shape = tensor.shape().dims(); - - if tensor.ndim() == 1 { - if dim != 0 { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - result_slice[0] = input_data.iter().filter(|v| !v.is_nan()).sum::(); - } else if tensor.ndim() == 2 { - let cols = input_shape[1]; - match dim { - 0 => { - let sums = input_data - .par_chunks_exact(cols) - .fold( - || vec![0f64; cols], - |mut acc, row| { - for (a, &v) in acc.iter_mut().zip(row) { - if !v.is_nan() { - *a += v; - } - } - acc - }, - ) - .reduce( - || vec![0f64; cols], - |mut a, b| { - for (x, y) in a.iter_mut().zip(b) { - *x += y; - } - a - }, - ); - result_slice.copy_from_slice(&sums); - } - 1 => { - result_slice - .par_iter_mut() - .zip(input_data.par_chunks_exact(cols)) - .for_each(|(out, row)| { - *out = row.iter().filter(|v| !v.is_nan()).sum::(); - }); - } - _ => { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - } - } else { - let dim_size = input_shape[dim]; - let inner = input_shape[dim + 1..].iter().product::(); - let outer_stride = dim_size * inner; - - result_slice - .par_iter_mut() - .enumerate() - .for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut sum_val = 0f64; - let mut base = o * outer_stride + r; - for _ in 0..dim_size { - let value = input_data[base]; - if !value.is_nan() { - sum_val += value; - } - base += inner; - } - *out = sum_val; - }); - } - - Ok(()) -} - -fn sum_along_dim_i32(tensor: &Tensor, result_data: &mut TensorData, dim: usize) -> Result<()> { - let input_data = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - - let result_slice = result_data - .as_i32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i32 slice"))?; - - let input_shape = tensor.shape().dims(); - - if tensor.ndim() == 1 { - if dim != 0 { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - result_slice[0] = simd_sum_i32(input_data); - } else if tensor.ndim() == 2 { - let cols = input_shape[1]; - match dim { - 0 => { - let sums = input_data - .par_chunks_exact(cols) - .fold( - || vec![0i32; cols], - |mut acc, row| { - for (a, &v) in acc.iter_mut().zip(row) { - *a += v; - } - acc - }, - ) - .reduce( - || vec![0i32; cols], - |mut a, b| { - for (x, y) in a.iter_mut().zip(b) { - *x += y; - } - a - }, - ); - result_slice.copy_from_slice(&sums); - } - 1 => { - result_slice - .par_iter_mut() - .zip(input_data.par_chunks_exact(cols)) - .for_each(|(out, row)| { - *out = simd_sum_i32(row); - }); - } - _ => { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - } - } else { - let dim_size = input_shape[dim]; - let inner = input_shape[dim + 1..].iter().product::(); - let outer_stride = dim_size * inner; - - result_slice - .par_iter_mut() - .enumerate() - .for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut sum_val = 0i32; - let mut base = o * outer_stride + r; - for _ in 0..dim_size { - sum_val += input_data[base]; - base += inner; - } - *out = sum_val; - }); - } - - Ok(()) -} - -fn sum_along_dim_i64(tensor: &Tensor, result_data: &mut TensorData, dim: usize) -> Result<()> { - let input_data = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - - let result_slice = result_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - let input_shape = tensor.shape().dims(); - - if tensor.ndim() == 1 { - if dim != 0 { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - result_slice[0] = simd_sum_i64(input_data); - } else if tensor.ndim() == 2 { - let cols = input_shape[1]; - match dim { - 0 => { - let sums = input_data - .par_chunks_exact(cols) - .fold( - || vec![0i64; cols], - |mut acc, row| { - for (a, &v) in acc.iter_mut().zip(row) { - *a += v; - } - acc - }, - ) - .reduce( - || vec![0i64; cols], - |mut a, b| { - for (x, y) in a.iter_mut().zip(b) { - *x += y; - } - a - }, - ); - result_slice.copy_from_slice(&sums); - } - 1 => { - result_slice - .par_iter_mut() - .zip(input_data.par_chunks_exact(cols)) - .for_each(|(out, row)| { - *out = simd_sum_i64(row); - }); - } - _ => { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - } - } else { - let dim_size = input_shape[dim]; - let inner = input_shape[dim + 1..].iter().product::(); - let outer_stride = dim_size * inner; - - result_slice - .par_iter_mut() - .enumerate() - .for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut sum_val = 0i64; - let mut base = o * outer_stride + r; - for _ in 0..dim_size { - sum_val += input_data[base]; - base += inner; - } - *out = sum_val; - }); - } - - Ok(()) -} - -#[inline] -pub fn prod_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { - if dim >= tensor.ndim() { - return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); - } - - let input_shape = tensor.shape().dims(); - let mut output_shape = input_shape.to_vec(); - if keepdim { - output_shape[dim] = 1; - } else { - output_shape.remove(dim); - } - let output_shape_obj = Shape::new(output_shape); - let mut result_data = - TensorData::zeros_on_device(output_shape_obj.numel(), tensor.dtype(), tensor.device()); - - match tensor.dtype() { - DataType::Float32 => prod_along_dim_f32(tensor, &mut result_data, dim)?, - DataType::Float64 => prod_along_dim_f64(tensor, &mut result_data, dim)?, - DataType::Int32 => prod_along_dim_i32(tensor, &mut result_data, dim)?, - DataType::Int64 => prod_along_dim_i64(tensor, &mut result_data, dim)?, - DataType::Bool => prod_along_dim_bool(tensor, &mut result_data, dim)?, - } - - let requires_grad = tensor.requires_grad() && tensor.dtype() != DataType::Bool; - Ok(Tensor::new( - Arc::new(result_data), - output_shape_obj, - tensor.dtype(), - tensor.device(), - requires_grad, - )) -} - -fn prod_along_dim_f32(tensor: &Tensor, result_data: &mut TensorData, dim: usize) -> Result<()> { - let input_data = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - let result_slice = result_data - .as_f32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; - let input_shape = tensor.shape().dims(); - let dim_size = input_shape[dim]; - let inner = input_shape[dim + 1..].iter().product::(); - let outer_stride = dim_size * inner; - result_slice - .par_iter_mut() - .enumerate() - .for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut prod_val = 1f32; - let mut base = o * outer_stride + r; - for _ in 0..dim_size { - prod_val *= input_data[base]; - base += inner; - } - *out = prod_val; - }); - Ok(()) -} - -fn prod_along_dim_f64(tensor: &Tensor, result_data: &mut TensorData, dim: usize) -> Result<()> { - let input_data = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - let result_slice = result_data - .as_f64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; - let input_shape = tensor.shape().dims(); - let dim_size = input_shape[dim]; - let inner = input_shape[dim + 1..].iter().product::(); - let outer_stride = dim_size * inner; - result_slice - .par_iter_mut() - .enumerate() - .for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut prod_val = 1f64; - let mut base = o * outer_stride + r; - for _ in 0..dim_size { - prod_val *= input_data[base]; - base += inner; - } - *out = prod_val; - }); - Ok(()) -} - -fn prod_along_dim_i32(tensor: &Tensor, result_data: &mut TensorData, dim: usize) -> Result<()> { - let input_data = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - let result_slice = result_data - .as_i32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i32 slice"))?; - let input_shape = tensor.shape().dims(); - let dim_size = input_shape[dim]; - let inner = input_shape[dim + 1..].iter().product::(); - let outer_stride = dim_size * inner; - result_slice - .par_iter_mut() - .enumerate() - .for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut prod_val = 1i32; - let mut base = o * outer_stride + r; - for _ in 0..dim_size { - prod_val *= input_data[base]; - base += inner; - } - *out = prod_val; - }); - Ok(()) -} - -fn prod_along_dim_i64(tensor: &Tensor, result_data: &mut TensorData, dim: usize) -> Result<()> { - let input_data = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - let result_slice = result_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - let input_shape = tensor.shape().dims(); - let dim_size = input_shape[dim]; - let inner = input_shape[dim + 1..].iter().product::(); - let outer_stride = dim_size * inner; - result_slice - .par_iter_mut() - .enumerate() - .for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut prod_val = 1i64; - let mut base = o * outer_stride + r; - for _ in 0..dim_size { - prod_val *= input_data[base]; - base += inner; - } - *out = prod_val; - }); - - Ok(()) -} - -fn prod_along_dim_bool(tensor: &Tensor, result_data: &mut TensorData, dim: usize) -> Result<()> { - let input_data = tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - let result_slice = result_data - .as_bool_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable bool slice"))?; - let input_shape = tensor.shape().dims(); - let dim_size = input_shape[dim]; - let inner = input_shape[dim + 1..].iter().product::(); - let outer_stride = dim_size * inner; - result_slice - .par_iter_mut() - .enumerate() - .for_each(|(idx, out)| { - let o = idx / inner; - let r = idx % inner; - let mut val = true; - let mut base = o * outer_stride + r; - for _ in 0..dim_size { - val &= input_data[base]; - if !val { - break; - } - base += inner; - } - *out = val; - }); - - Ok(()) -} - -// Helper implementations for max/min operations -fn max_all_f32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - - let max_val = data.par_iter().cloned().reduce( - || f32::NEG_INFINITY, - |a, b| { - if a.is_nan() || b.is_nan() { - f32::NAN - } else { - a.max(b) - } - }, - ); - - let result_slice = result_data - .as_f32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; - - result_slice[0] = max_val; - Ok(()) -} - -fn max_all_f64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - - let max_val = data.par_iter().cloned().reduce( - || f64::NEG_INFINITY, - |a, b| { - if a.is_nan() || b.is_nan() { - f64::NAN - } else { - a.max(b) - } - }, - ); - - let result_slice = result_data - .as_f64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; - - result_slice[0] = max_val; - Ok(()) -} - -fn max_all_i32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - - let max_val = data.par_iter().copied().max().unwrap_or(i32::MIN); - - let result_slice = result_data - .as_i32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i32 slice"))?; - - result_slice[0] = max_val; - Ok(()) -} - -fn max_all_i64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - - let max_val = data.par_iter().copied().max().unwrap_or(i64::MIN); - - let result_slice = result_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - result_slice[0] = max_val; - Ok(()) -} - -fn max_all_bool(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - - let max_val = data.par_iter().any(|&x| x); - - let result_slice = result_data - .as_bool_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable bool slice"))?; - - result_slice[0] = max_val; - Ok(()) -} - -// Similar implementations for min functions -fn min_all_f32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - - let min_val = data.par_iter().cloned().reduce( - || f32::INFINITY, - |a, b| { - if a.is_nan() || b.is_nan() { - f32::NAN - } else { - a.min(b) - } - }, - ); - - let result_slice = result_data - .as_f32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; - - result_slice[0] = min_val; - Ok(()) -} - -fn min_all_f64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; - - let min_val = data.par_iter().cloned().reduce( - || f64::INFINITY, - |a, b| { - if a.is_nan() || b.is_nan() { - f64::NAN - } else { - a.min(b) - } - }, - ); - - let result_slice = result_data - .as_f64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; - - result_slice[0] = min_val; - Ok(()) -} - -fn min_all_i32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; - - let min_val = data.par_iter().copied().min().unwrap_or(i32::MAX); - - let result_slice = result_data - .as_i32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i32 slice"))?; - - result_slice[0] = min_val; - Ok(()) -} - -fn min_all_i64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; - - let min_val = data.par_iter().copied().min().unwrap_or(i64::MAX); - - let result_slice = result_data - .as_i64_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; - - result_slice[0] = min_val; - Ok(()) -} - -fn min_all_bool(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - - let min_val = data.par_iter().all(|&x| x); - - let result_slice = result_data - .as_bool_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable bool slice"))?; - - result_slice[0] = min_val; - Ok(()) -} - -fn nanmax_all_f32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { - let data = tensor - .data() - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; - - let (max_val, found) = data - .par_iter() - .map(|&v| { - if v.is_nan() { - (f32::NEG_INFINITY, false) - } else { - (v, true) - } - }) - .reduce( - || (f32::NEG_INFINITY, false), - |(a_val, a_found), (b_val, b_found)| match (a_found, b_found) { - (true, true) => (a_val.max(b_val), true), - (true, false) => (a_val, true), - (false, true) => (b_val, true), - (false, false) => (f32::NEG_INFINITY, false), - }, - ); - - let result_slice = result_data - .as_f32_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; - - result_slice[0] = if found { max_val } else { f32::NAN }; - Ok(()) -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use crate::operations::simd::*; +use crate::{ + error::{MinitensorError, Result}, + tensor::{DataType, Shape, Tensor, TensorData}, +}; +use rayon::prelude::*; +use std::sync::Arc; + +pub(crate) fn sum_along_dim_f32( + tensor: &Tensor, + result_data: &mut TensorData, + dim: usize, +) -> Result<()> { + let input_data = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + + let result_slice = result_data + .as_f32_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; + + let input_shape = tensor.shape().dims(); + + if tensor.ndim() == 1 { + if dim != 0 { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + result_slice[0] = simd_sum_f32(input_data); + } else if tensor.ndim() == 2 { + let cols = input_shape[1]; + match dim { + 0 => { + let sums = input_data + .par_chunks_exact(cols) + .fold( + || vec![0f32; cols], + |mut acc, row| { + for (a, &v) in acc.iter_mut().zip(row) { + *a += v; + } + acc + }, + ) + .reduce( + || vec![0f32; cols], + |mut a, b| { + for (x, y) in a.iter_mut().zip(b) { + *x += y; + } + a + }, + ); + result_slice.copy_from_slice(&sums); + } + 1 => { + result_slice + .par_iter_mut() + .zip(input_data.par_chunks_exact(cols)) + .for_each(|(out, row)| { + *out = simd_sum_f32(row); + }); + } + _ => { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + } + } else { + let dim_size = input_shape[dim]; + let inner = input_shape[dim + 1..].iter().product::(); + let outer_stride = dim_size * inner; + + result_slice + .par_iter_mut() + .enumerate() + .for_each(|(idx, out)| { + let o = idx / inner; + let r = idx % inner; + let mut sum_val = 0f32; + let mut base = o * outer_stride + r; + for _ in 0..dim_size { + sum_val += input_data[base]; + base += inner; + } + *out = sum_val; + }); + } + + Ok(()) +} + +pub(crate) fn nansum_along_dim_f32( + tensor: &Tensor, + result_data: &mut TensorData, + dim: usize, +) -> Result<()> { + let input_data = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + + let result_slice = result_data + .as_f32_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; + + let input_shape = tensor.shape().dims(); + + if tensor.ndim() == 1 { + if dim != 0 { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + result_slice[0] = input_data.iter().filter(|v| !v.is_nan()).sum::(); + } else if tensor.ndim() == 2 { + let cols = input_shape[1]; + match dim { + 0 => { + let sums = input_data + .par_chunks_exact(cols) + .fold( + || vec![0f32; cols], + |mut acc, row| { + for (a, &v) in acc.iter_mut().zip(row) { + if !v.is_nan() { + *a += v; + } + } + acc + }, + ) + .reduce( + || vec![0f32; cols], + |mut a, b| { + for (x, y) in a.iter_mut().zip(b) { + *x += y; + } + a + }, + ); + result_slice.copy_from_slice(&sums); + } + 1 => { + result_slice + .par_iter_mut() + .zip(input_data.par_chunks_exact(cols)) + .for_each(|(out, row)| { + *out = row.iter().filter(|v| !v.is_nan()).sum::(); + }); + } + _ => { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + } + } else { + let dim_size = input_shape[dim]; + let inner = input_shape[dim + 1..].iter().product::(); + let outer_stride = dim_size * inner; + + result_slice + .par_iter_mut() + .enumerate() + .for_each(|(idx, out)| { + let o = idx / inner; + let r = idx % inner; + let mut sum_val = 0f32; + let mut base = o * outer_stride + r; + for _ in 0..dim_size { + let value = input_data[base]; + if !value.is_nan() { + sum_val += value; + } + base += inner; + } + *out = sum_val; + }); + } + + Ok(()) +} + +pub(crate) fn sum_along_dim_f64( + tensor: &Tensor, + result_data: &mut TensorData, + dim: usize, +) -> Result<()> { + let input_data = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + + let result_slice = result_data + .as_f64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; + + let input_shape = tensor.shape().dims(); + + if tensor.ndim() == 1 { + if dim != 0 { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + result_slice[0] = simd_sum_f64(input_data); + } else if tensor.ndim() == 2 { + let cols = input_shape[1]; + match dim { + 0 => { + let sums = input_data + .par_chunks_exact(cols) + .fold( + || vec![0f64; cols], + |mut acc, row| { + for (a, &v) in acc.iter_mut().zip(row) { + *a += v; + } + acc + }, + ) + .reduce( + || vec![0f64; cols], + |mut a, b| { + for (x, y) in a.iter_mut().zip(b) { + *x += y; + } + a + }, + ); + result_slice.copy_from_slice(&sums); + } + 1 => { + result_slice + .par_iter_mut() + .zip(input_data.par_chunks_exact(cols)) + .for_each(|(out, row)| { + *out = simd_sum_f64(row); + }); + } + _ => { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + } + } else { + let dim_size = input_shape[dim]; + let inner = input_shape[dim + 1..].iter().product::(); + let outer_stride = dim_size * inner; + + result_slice + .par_iter_mut() + .enumerate() + .for_each(|(idx, out)| { + let o = idx / inner; + let r = idx % inner; + let mut sum_val = 0f64; + let mut base = o * outer_stride + r; + for _ in 0..dim_size { + sum_val += input_data[base]; + base += inner; + } + *out = sum_val; + }); + } + + Ok(()) +} + +pub(crate) fn nansum_along_dim_f64( + tensor: &Tensor, + result_data: &mut TensorData, + dim: usize, +) -> Result<()> { + let input_data = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + + let result_slice = result_data + .as_f64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; + + let input_shape = tensor.shape().dims(); + + if tensor.ndim() == 1 { + if dim != 0 { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + result_slice[0] = input_data.iter().filter(|v| !v.is_nan()).sum::(); + } else if tensor.ndim() == 2 { + let cols = input_shape[1]; + match dim { + 0 => { + let sums = input_data + .par_chunks_exact(cols) + .fold( + || vec![0f64; cols], + |mut acc, row| { + for (a, &v) in acc.iter_mut().zip(row) { + if !v.is_nan() { + *a += v; + } + } + acc + }, + ) + .reduce( + || vec![0f64; cols], + |mut a, b| { + for (x, y) in a.iter_mut().zip(b) { + *x += y; + } + a + }, + ); + result_slice.copy_from_slice(&sums); + } + 1 => { + result_slice + .par_iter_mut() + .zip(input_data.par_chunks_exact(cols)) + .for_each(|(out, row)| { + *out = row.iter().filter(|v| !v.is_nan()).sum::(); + }); + } + _ => { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + } + } else { + let dim_size = input_shape[dim]; + let inner = input_shape[dim + 1..].iter().product::(); + let outer_stride = dim_size * inner; + + result_slice + .par_iter_mut() + .enumerate() + .for_each(|(idx, out)| { + let o = idx / inner; + let r = idx % inner; + let mut sum_val = 0f64; + let mut base = o * outer_stride + r; + for _ in 0..dim_size { + let value = input_data[base]; + if !value.is_nan() { + sum_val += value; + } + base += inner; + } + *out = sum_val; + }); + } + + Ok(()) +} + +pub(crate) fn sum_along_dim_i32( + tensor: &Tensor, + result_data: &mut TensorData, + dim: usize, +) -> Result<()> { + let input_data = tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; + + let result_slice = result_data + .as_i32_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i32 slice"))?; + + let input_shape = tensor.shape().dims(); + + if tensor.ndim() == 1 { + if dim != 0 { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + result_slice[0] = simd_sum_i32(input_data); + } else if tensor.ndim() == 2 { + let cols = input_shape[1]; + match dim { + 0 => { + let sums = input_data + .par_chunks_exact(cols) + .fold( + || vec![0i32; cols], + |mut acc, row| { + for (a, &v) in acc.iter_mut().zip(row) { + *a += v; + } + acc + }, + ) + .reduce( + || vec![0i32; cols], + |mut a, b| { + for (x, y) in a.iter_mut().zip(b) { + *x += y; + } + a + }, + ); + result_slice.copy_from_slice(&sums); + } + 1 => { + result_slice + .par_iter_mut() + .zip(input_data.par_chunks_exact(cols)) + .for_each(|(out, row)| { + *out = simd_sum_i32(row); + }); + } + _ => { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + } + } else { + let dim_size = input_shape[dim]; + let inner = input_shape[dim + 1..].iter().product::(); + let outer_stride = dim_size * inner; + + result_slice + .par_iter_mut() + .enumerate() + .for_each(|(idx, out)| { + let o = idx / inner; + let r = idx % inner; + let mut sum_val = 0i32; + let mut base = o * outer_stride + r; + for _ in 0..dim_size { + sum_val += input_data[base]; + base += inner; + } + *out = sum_val; + }); + } + + Ok(()) +} + +pub(crate) fn sum_along_dim_i64( + tensor: &Tensor, + result_data: &mut TensorData, + dim: usize, +) -> Result<()> { + let input_data = tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; + + let result_slice = result_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + let input_shape = tensor.shape().dims(); + + if tensor.ndim() == 1 { + if dim != 0 { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + result_slice[0] = simd_sum_i64(input_data); + } else if tensor.ndim() == 2 { + let cols = input_shape[1]; + match dim { + 0 => { + let sums = input_data + .par_chunks_exact(cols) + .fold( + || vec![0i64; cols], + |mut acc, row| { + for (a, &v) in acc.iter_mut().zip(row) { + *a += v; + } + acc + }, + ) + .reduce( + || vec![0i64; cols], + |mut a, b| { + for (x, y) in a.iter_mut().zip(b) { + *x += y; + } + a + }, + ); + result_slice.copy_from_slice(&sums); + } + 1 => { + result_slice + .par_iter_mut() + .zip(input_data.par_chunks_exact(cols)) + .for_each(|(out, row)| { + *out = simd_sum_i64(row); + }); + } + _ => { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + } + } else { + let dim_size = input_shape[dim]; + let inner = input_shape[dim + 1..].iter().product::(); + let outer_stride = dim_size * inner; + + result_slice + .par_iter_mut() + .enumerate() + .for_each(|(idx, out)| { + let o = idx / inner; + let r = idx % inner; + let mut sum_val = 0i64; + let mut base = o * outer_stride + r; + for _ in 0..dim_size { + sum_val += input_data[base]; + base += inner; + } + *out = sum_val; + }); + } + + Ok(()) +} + +#[inline] +pub fn prod_along_dim(tensor: &Tensor, dim: usize, keepdim: bool) -> Result { + if dim >= tensor.ndim() { + return Err(MinitensorError::index_error(dim as isize, 0, tensor.ndim())); + } + + let input_shape = tensor.shape().dims(); + let mut output_shape = input_shape.to_vec(); + if keepdim { + output_shape[dim] = 1; + } else { + output_shape.remove(dim); + } + let output_shape_obj = Shape::new(output_shape); + let mut result_data = + TensorData::zeros_on_device(output_shape_obj.numel(), tensor.dtype(), tensor.device()); + + match tensor.dtype() { + DataType::Float32 => prod_along_dim_f32(tensor, &mut result_data, dim)?, + DataType::Float64 => prod_along_dim_f64(tensor, &mut result_data, dim)?, + DataType::Int32 => prod_along_dim_i32(tensor, &mut result_data, dim)?, + DataType::Int64 => prod_along_dim_i64(tensor, &mut result_data, dim)?, + DataType::Bool => prod_along_dim_bool(tensor, &mut result_data, dim)?, + } + + let requires_grad = tensor.requires_grad() && tensor.dtype() != DataType::Bool; + Ok(Tensor::new( + Arc::new(result_data), + output_shape_obj, + tensor.dtype(), + tensor.device(), + requires_grad, + )) +} + +/// Generates a product-along-dim reduction kernel. Body is identical across +/// numeric dtypes; only the element type and multiplicative identity differ. +macro_rules! prod_along_dim_kernel { + ($name:ident, $accessor:ident, $accessor_mut:ident, $tyname:literal, $one:expr) => { + fn $name(tensor: &Tensor, result_data: &mut TensorData, dim: usize) -> Result<()> { + let input_data = tensor.data().$accessor().ok_or_else(|| { + MinitensorError::internal_error(concat!("Failed to get ", $tyname, " slice")) + })?; + let result_slice = result_data.$accessor_mut().ok_or_else(|| { + MinitensorError::internal_error(concat!( + "Failed to get mutable ", + $tyname, + " slice" + )) + })?; + let input_shape = tensor.shape().dims(); + let dim_size = input_shape[dim]; + let inner = input_shape[dim + 1..].iter().product::(); + let outer_stride = dim_size * inner; + result_slice + .par_iter_mut() + .enumerate() + .for_each(|(idx, out)| { + let o = idx / inner; + let r = idx % inner; + let mut prod_val = $one; + let mut base = o * outer_stride + r; + for _ in 0..dim_size { + prod_val *= input_data[base]; + base += inner; + } + *out = prod_val; + }); + Ok(()) + } + }; +} + +prod_along_dim_kernel!( + prod_along_dim_f32, + as_f32_slice, + as_f32_slice_mut, + "f32", + 1f32 +); + +prod_along_dim_kernel!( + prod_along_dim_f64, + as_f64_slice, + as_f64_slice_mut, + "f64", + 1f64 +); + +prod_along_dim_kernel!( + prod_along_dim_i32, + as_i32_slice, + as_i32_slice_mut, + "i32", + 1i32 +); + +prod_along_dim_kernel!( + prod_along_dim_i64, + as_i64_slice, + as_i64_slice_mut, + "i64", + 1i64 +); + +fn prod_along_dim_bool(tensor: &Tensor, result_data: &mut TensorData, dim: usize) -> Result<()> { + let input_data = tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + let result_slice = result_data + .as_bool_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable bool slice"))?; + let input_shape = tensor.shape().dims(); + let dim_size = input_shape[dim]; + let inner = input_shape[dim + 1..].iter().product::(); + let outer_stride = dim_size * inner; + result_slice + .par_iter_mut() + .enumerate() + .for_each(|(idx, out)| { + let o = idx / inner; + let r = idx % inner; + let mut val = true; + let mut base = o * outer_stride + r; + for _ in 0..dim_size { + val &= input_data[base]; + if !val { + break; + } + base += inner; + } + *out = val; + }); + + Ok(()) +} + +// Helper implementations for max/min operations +pub(crate) fn max_all_f32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + + let max_val = data.par_iter().cloned().reduce( + || f32::NEG_INFINITY, + |a, b| { + if a.is_nan() || b.is_nan() { + f32::NAN + } else { + a.max(b) + } + }, + ); + + let result_slice = result_data + .as_f32_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; + + result_slice[0] = max_val; + Ok(()) +} + +pub(crate) fn max_all_f64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + + let max_val = data.par_iter().cloned().reduce( + || f64::NEG_INFINITY, + |a, b| { + if a.is_nan() || b.is_nan() { + f64::NAN + } else { + a.max(b) + } + }, + ); + + let result_slice = result_data + .as_f64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; + + result_slice[0] = max_val; + Ok(()) +} + +pub(crate) fn max_all_i32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; + + let max_val = data.par_iter().copied().max().unwrap_or(i32::MIN); + + let result_slice = result_data + .as_i32_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i32 slice"))?; + + result_slice[0] = max_val; + Ok(()) +} + +pub(crate) fn max_all_i64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; + + let max_val = data.par_iter().copied().max().unwrap_or(i64::MIN); + + let result_slice = result_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + result_slice[0] = max_val; + Ok(()) +} + +pub(crate) fn max_all_bool(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + + let max_val = data.par_iter().any(|&x| x); + + let result_slice = result_data + .as_bool_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable bool slice"))?; + + result_slice[0] = max_val; + Ok(()) +} + +// Similar implementations for min functions +pub(crate) fn min_all_f32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + + let min_val = data.par_iter().cloned().reduce( + || f32::INFINITY, + |a, b| { + if a.is_nan() || b.is_nan() { + f32::NAN + } else { + a.min(b) + } + }, + ); + + let result_slice = result_data + .as_f32_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; + + result_slice[0] = min_val; + Ok(()) +} + +pub(crate) fn min_all_f64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f64 slice"))?; + + let min_val = data.par_iter().cloned().reduce( + || f64::INFINITY, + |a, b| { + if a.is_nan() || b.is_nan() { + f64::NAN + } else { + a.min(b) + } + }, + ); + + let result_slice = result_data + .as_f64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f64 slice"))?; + + result_slice[0] = min_val; + Ok(()) +} + +pub(crate) fn min_all_i32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i32 slice"))?; + + let min_val = data.par_iter().copied().min().unwrap_or(i32::MAX); + + let result_slice = result_data + .as_i32_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i32 slice"))?; + + result_slice[0] = min_val; + Ok(()) +} + +pub(crate) fn min_all_i64(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get i64 slice"))?; + + let min_val = data.par_iter().copied().min().unwrap_or(i64::MAX); + + let result_slice = result_data + .as_i64_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable i64 slice"))?; + + result_slice[0] = min_val; + Ok(()) +} + +pub(crate) fn min_all_bool(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + + let min_val = data.par_iter().all(|&x| x); + + let result_slice = result_data + .as_bool_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable bool slice"))?; + + result_slice[0] = min_val; + Ok(()) +} + +pub(crate) fn nanmax_all_f32(tensor: &Tensor, result_data: &mut TensorData) -> Result<()> { + let data = tensor + .data() + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Failed to get f32 slice"))?; + + let (max_val, found) = data + .par_iter() + .map(|&v| { + if v.is_nan() { + (f32::NEG_INFINITY, false) + } else { + (v, true) + } + }) + .reduce( + || (f32::NEG_INFINITY, false), + |(a_val, a_found), (b_val, b_found)| match (a_found, b_found) { + (true, true) => (a_val.max(b_val), true), + (true, false) => (a_val, true), + (false, true) => (b_val, true), + (false, false) => (f32::NEG_INFINITY, false), + }, + ); + + let result_slice = result_data + .as_f32_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get mutable f32 slice"))?; + + result_slice[0] = if found { max_val } else { f32::NAN }; + Ok(()) +} diff --git a/engine/src/operations/shape_ops.rs b/engine/src/operations/shape_ops.rs index 869df71a..bbadfe53 100644 --- a/engine/src/operations/shape_ops.rs +++ b/engine/src/operations/shape_ops.rs @@ -4,5 +4,10 @@ // This source code is licensed under the Apache-style license found in the // LICENSE file in the root directory of this source tree. -include!("shape_ops/reshape.rs"); -include!("shape_ops/indexing.rs"); +#[path = "shape_ops/indexing.rs"] +mod indexing_impl; +#[path = "shape_ops/reshape.rs"] +mod reshape_impl; + +pub use self::indexing_impl::*; +pub use self::reshape_impl::*; diff --git a/engine/src/operations/shape_ops/indexing.rs b/engine/src/operations/shape_ops/indexing.rs index 838e2954..3701f624 100644 --- a/engine/src/operations/shape_ops/indexing.rs +++ b/engine/src/operations/shape_ops/indexing.rs @@ -1,386 +1,396 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -fn expand_repeats(spec: RepeatInterleaveSpec<'_>, dim_size: usize) -> Result> { - match spec { - RepeatInterleaveSpec::Scalar(value) => Ok(vec![value; dim_size]), - RepeatInterleaveSpec::Slice(values) => { - if values.len() == dim_size { - Ok(values.to_vec()) - } else if values.len() == 1 { - if dim_size == 0 { - Ok(Vec::new()) - } else { - Ok(vec![values[0]; dim_size]) - } - } else if values.is_empty() && dim_size == 0 { - Ok(Vec::new()) - } else { - Err(MinitensorError::invalid_operation( - "repeat_interleave: repeats must be a single value or match tensor size along dim" - .to_string(), - )) - } - } - RepeatInterleaveSpec::Tensor(tensor) => collect_repeats_from_tensor(tensor, dim_size), - } -} - -fn build_empty_repeat_result(tensor: &Tensor, dim: usize, target: usize) -> Result { - let mut out_shape = tensor.shape().dims().to_vec(); - out_shape[dim] = target; - let shape = Shape::new(out_shape); - let dtype = tensor.dtype(); - let device = tensor.device(); - let data = TensorData::zeros_on_device(shape.numel(), dtype, device); - Ok(Tensor::new( - Arc::new(data), - shape, - dtype, - device, - tensor.requires_grad(), - )) -} - -/// Repeat elements of ``tensor`` according to ``repeats`` along ``dim``. -pub fn repeat_interleave( - tensor: &Tensor, - repeats: RepeatInterleaveSpec<'_>, - dim: Option, - output_size: Option, -) -> Result { - if dim.is_none() { - // Grad-aware flatten so the [numel] gradient is reshaped back to the - // input shape (a bare `flatten_all` view aliases the input's id and would - // attribute a wrongly-shaped gradient to it). - let flat = flatten(tensor, 0, -1)?; - return repeat_interleave(&flat, repeats, Some(0), output_size); - } - - if !tensor.device().is_cpu() { - return Err(MinitensorError::invalid_operation( - "repeat_interleave currently supports only CPU tensors".to_string(), - )); - } - - let dim = normalize_dim(dim.unwrap(), tensor.ndim())?; - let dims = tensor.shape().dims(); - let dim_size = dims[dim]; - let reps = expand_repeats(repeats, dim_size)?; - let total_repeats: usize = reps.iter().sum(); - - if let Some(expected) = output_size { - if expected != total_repeats { - return Err(MinitensorError::invalid_argument(format!( - "repeat_interleave: output_size ({expected}) must equal sum of repeats ({total_repeats})" - ))); - } - } - - let dtype = tensor.dtype(); - let device = tensor.device(); - let requires_grad = tensor.requires_grad(); - - let target_dim = output_size.unwrap_or(total_repeats); - let mut output_shape = dims.to_vec(); - output_shape[dim] = target_dim; - let output_shape_obj = Shape::new(output_shape); - let output_numel = output_shape_obj.numel(); - - let inner: usize = dims[dim + 1..].iter().product(); - let outer: usize = if dim == 0 { - 1 - } else { - dims[..dim].iter().product() - }; - - let build_grad_fn = |repeats: Vec| { - Arc::new(RepeatInterleaveBackward { - input_shape: dims.to_vec(), - repeats, - input_id: tensor.id(), - dim, - }) - }; - - if target_dim == 0 || output_numel == 0 || inner == 0 || outer == 0 { - let mut result = build_empty_repeat_result(tensor, dim, target_dim)?; - if requires_grad { - let grad_fn = build_grad_fn(reps); - result.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&result, Some(grad_fn))?; - } - return Ok(result); - } - - macro_rules! repeat_impl { - ($ty:ty, $slice:ident, $from_vec:ident) => {{ - let src = tensor.data().$slice().ok_or_else(|| { - MinitensorError::invalid_operation( - "repeat_interleave: tensor data access failed".to_string(), - ) - })?; - let mut out = vec![<$ty>::default(); output_numel]; - out.par_chunks_mut(target_dim * inner).enumerate().for_each( - |(outer_idx, out_chunk)| { - let mut dst_offset = 0; - let base = outer_idx * dim_size * inner; - for (i, &rep) in reps.iter().enumerate() { - if rep == 0 { - continue; - } - let src_start = base + i * inner; - let src_slice = &src[src_start..src_start + inner]; - for _ in 0..rep { - let end = dst_offset + inner; - out_chunk[dst_offset..end].copy_from_slice(src_slice); - dst_offset = end; - } - } - }, - ); - TensorData::$from_vec(out, device) - }}; - } - - let data = match dtype { - DataType::Float32 => repeat_impl!(f32, as_f32_slice, from_vec_f32), - DataType::Float64 => repeat_impl!(f64, as_f64_slice, from_vec_f64), - DataType::Int32 => repeat_impl!(i32, as_i32_slice, from_vec_i32), - DataType::Int64 => repeat_impl!(i64, as_i64_slice, from_vec_i64), - DataType::Bool => repeat_impl!(bool, as_bool_slice, from_vec_bool), - }; - - let mut result = Tensor::new( - Arc::new(data), - output_shape_obj, - dtype, - device, - requires_grad, - ); - - if requires_grad { - let grad_fn = build_grad_fn(reps); - result.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&result, Some(grad_fn))?; - } - - Ok(result) -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::{ - device::Device, - tensor::{DataType, TensorData}, - }; - - fn create_test_tensor_f32(data: Vec, shape: Vec, requires_grad: bool) -> Tensor { - let shape_obj = Shape::new(shape); - let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Float32); - - if let Some(slice) = tensor_data.as_f32_slice_mut() { - slice.copy_from_slice(&data); - } - - Tensor::new( - Arc::new(tensor_data), - shape_obj, - DataType::Float32, - Device::cpu(), - requires_grad, - ) - } - - #[test] - fn test_reshape_basic() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3], false); - - let reshaped = reshape(&tensor, Shape::new(vec![3, 2])).unwrap(); - - assert_eq!(reshaped.shape().dims(), &[3, 2]); - assert_eq!(reshaped.numel(), 6); - - let data = reshaped.data().as_f32_slice().unwrap(); - assert_eq!(data, &[1.0, 2.0, 3.0, 4.0, 5.0, 6.0]); - } - - #[test] - fn test_reshape_invalid_size() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); - - let result = reshape(&tensor, Shape::new(vec![2, 3])); - assert!(result.is_err()); - } - - #[test] - fn test_reshape_infer_dim() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![6], false); - let reshaped = reshape_with_inference(&tensor, vec![2, -1]).unwrap(); - assert_eq!(reshaped.shape().dims(), &[2, 3]); - } - - #[test] - fn test_reshape_multiple_negative_one_error() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![4], false); - let result = reshape_with_inference(&tensor, vec![-1, -1]); - assert!(result.is_err()); - } - - #[test] - fn test_reshape_infer_mismatch_error() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0, 5.0], vec![5], false); - let result = reshape_with_inference(&tensor, vec![4, -1]); - assert!(result.is_err()); - } - - #[test] - fn test_reshape_zero_dim_with_inference_error() { - let tensor = create_test_tensor_f32(vec![], vec![0], false); - let result = reshape_with_inference(&tensor, vec![-1, 0]); - assert!(result.is_err()); - } - - #[test] - fn test_squeeze_specific_dim() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![1, 4, 1], false); - - let s0 = squeeze(&tensor, Some(0)).unwrap(); - assert_eq!(s0.shape().dims(), &[4, 1]); - - let s1 = squeeze(&s0, Some(1)).unwrap(); - assert_eq!(s1.shape().dims(), &[4]); - - let s_neg = squeeze(&tensor, Some(-1)).unwrap(); - assert_eq!(s_neg.shape().dims(), &[1, 4]); - } - - #[test] - fn test_squeeze_all() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![1, 4, 1], false); - - let squeezed = squeeze(&tensor, None).unwrap(); - assert_eq!(squeezed.shape().dims(), &[4]); - - let scalar = create_test_tensor_f32(vec![1.0], vec![1, 1], false); - let s = squeeze(&scalar, None).unwrap(); - assert!(s.shape().dims().is_empty()); - } - - #[test] - fn test_squeeze_out_of_range() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); - - assert!(squeeze(&tensor, Some(2)).is_err()); - assert!(squeeze(&tensor, Some(-3)).is_err()); - } - - #[test] - fn test_unsqueeze() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![4], false); - - let u0 = unsqueeze(&tensor, 0).unwrap(); - assert_eq!(u0.shape().dims(), &[1, 4]); - - let u1 = unsqueeze(&tensor, 1).unwrap(); - assert_eq!(u1.shape().dims(), &[4, 1]); - - let u_neg = unsqueeze(&tensor, -1).unwrap(); - assert_eq!(u_neg.shape().dims(), &[4, 1]); - } - - #[test] - fn test_gradient_tracking() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], true); - - let reshaped = reshape(&tensor, Shape::new(vec![4])).unwrap(); - - assert!(reshaped.requires_grad()); - assert!(reshaped.grad_fn().is_some()); - } - - #[test] - fn test_concatenate_validation() { - let tensor1 = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - let tensor2 = create_test_tensor_f32(vec![3.0, 4.0], vec![2], false); - - let result = concatenate(&[&tensor1, &tensor2], 0).unwrap(); - assert_eq!(result.shape().dims(), &[4]); - let data = result.data().as_f32_slice().unwrap(); - assert_eq!(data, &[1.0, 2.0, 3.0, 4.0]); - } - - #[test] - fn test_index_select_validation() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3], false); - - let result = index_select(&tensor, 1, &[0, 2]).unwrap(); - assert_eq!(result.shape().dims(), &[2, 2]); - let data = result.data().as_f32_slice().unwrap(); - assert_eq!(data, &[1.0, 3.0, 4.0, 6.0]); - } - - #[test] - fn test_slice_empty_range() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); - - let result = slice(&tensor, 1, 1, 1, 1).unwrap(); - assert_eq!(result.shape().dims(), &[2, 0]); - assert_eq!(result.numel(), 0); - } - - #[test] - fn test_slice_empty_at_end() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); - - let result = slice(&tensor, 0, 2, 2, 1).unwrap(); - assert_eq!(result.shape().dims(), &[0, 2]); - assert_eq!(result.numel(), 0); - } - - #[test] - fn test_index_select_empty_indices() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); - - let result = index_select(&tensor, 1, &[]).unwrap(); - assert_eq!(result.shape().dims(), &[2, 0]); - assert_eq!(result.numel(), 0); - } - - #[test] - fn test_slice_validation() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3], false); - - let result = slice(&tensor, 1, 0, 2, 1).unwrap(); - assert_eq!(result.shape().dims(), &[2, 2]); - let data = result.data().as_f32_slice().unwrap(); - assert_eq!(data, &[1.0, 2.0, 4.0, 5.0]); - } - - #[test] - fn test_repeat_basic() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - let repeated = repeat(&tensor, &[3]).unwrap(); - assert_eq!(repeated.shape().dims(), &[6]); - let data = repeated.data().as_f32_slice().unwrap(); - assert_eq!(data, &[1.0, 2.0, 1.0, 2.0, 1.0, 2.0]); - } - - #[test] - fn test_repeat_zero_numel_shape() { - let tensor = create_test_tensor_f32(vec![], vec![0, 2], false); - let repeated = repeat(&tensor, &[2, 3]).unwrap(); - assert_eq!(repeated.shape().dims(), &[0, 6]); - assert_eq!(repeated.numel(), 0); - } - - #[test] - fn test_repeat_dim_mismatch_error() { - let tensor = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); - assert!(repeat(&tensor, &[]).is_err()); - } -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::autograd::RepeatInterleaveBackward; +use crate::{ + autograd::add_to_graph, + error::{MinitensorError, Result}, + tensor::{DataType, Shape, Tensor, TensorData}, +}; +use rayon::prelude::*; +use std::sync::Arc; + +fn expand_repeats(spec: RepeatInterleaveSpec<'_>, dim_size: usize) -> Result> { + match spec { + RepeatInterleaveSpec::Scalar(value) => Ok(vec![value; dim_size]), + RepeatInterleaveSpec::Slice(values) => { + if values.len() == dim_size { + Ok(values.to_vec()) + } else if values.len() == 1 { + if dim_size == 0 { + Ok(Vec::new()) + } else { + Ok(vec![values[0]; dim_size]) + } + } else if values.is_empty() && dim_size == 0 { + Ok(Vec::new()) + } else { + Err(MinitensorError::invalid_operation( + "repeat_interleave: repeats must be a single value or match tensor size along dim" + .to_string(), + )) + } + } + RepeatInterleaveSpec::Tensor(tensor) => collect_repeats_from_tensor(tensor, dim_size), + } +} + +fn build_empty_repeat_result(tensor: &Tensor, dim: usize, target: usize) -> Result { + let mut out_shape = tensor.shape().dims().to_vec(); + out_shape[dim] = target; + let shape = Shape::new(out_shape); + let dtype = tensor.dtype(); + let device = tensor.device(); + let data = TensorData::zeros_on_device(shape.numel(), dtype, device); + Ok(Tensor::new( + Arc::new(data), + shape, + dtype, + device, + tensor.requires_grad(), + )) +} + +/// Repeat elements of ``tensor`` according to ``repeats`` along ``dim``. +pub fn repeat_interleave( + tensor: &Tensor, + repeats: RepeatInterleaveSpec<'_>, + dim: Option, + output_size: Option, +) -> Result { + if dim.is_none() { + // Grad-aware flatten so the [numel] gradient is reshaped back to the + // input shape (a bare `flatten_all` view aliases the input's id and would + // attribute a wrongly-shaped gradient to it). + let flat = flatten(tensor, 0, -1)?; + return repeat_interleave(&flat, repeats, Some(0), output_size); + } + + if !tensor.device().is_cpu() { + return Err(MinitensorError::invalid_operation( + "repeat_interleave currently supports only CPU tensors".to_string(), + )); + } + + let dim = normalize_dim(dim.unwrap(), tensor.ndim())?; + let dims = tensor.shape().dims(); + let dim_size = dims[dim]; + let reps = expand_repeats(repeats, dim_size)?; + let total_repeats: usize = reps.iter().sum(); + + if let Some(expected) = output_size + && expected != total_repeats + { + return Err(MinitensorError::invalid_argument(format!( + "repeat_interleave: output_size ({expected}) must equal sum of repeats ({total_repeats})" + ))); + } + + let dtype = tensor.dtype(); + let device = tensor.device(); + let requires_grad = tensor.requires_grad(); + + let target_dim = output_size.unwrap_or(total_repeats); + let mut output_shape = dims.to_vec(); + output_shape[dim] = target_dim; + let output_shape_obj = Shape::new(output_shape); + let output_numel = output_shape_obj.numel(); + + let inner: usize = dims[dim + 1..].iter().product(); + let outer: usize = if dim == 0 { + 1 + } else { + dims[..dim].iter().product() + }; + + let build_grad_fn = |repeats: Vec| { + Arc::new(RepeatInterleaveBackward { + input_shape: dims.to_vec(), + repeats, + input_id: tensor.id(), + dim, + }) + }; + + if target_dim == 0 || output_numel == 0 || inner == 0 || outer == 0 { + let mut result = build_empty_repeat_result(tensor, dim, target_dim)?; + if requires_grad { + let grad_fn = build_grad_fn(reps); + result.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&result, Some(grad_fn))?; + } + return Ok(result); + } + + macro_rules! repeat_impl { + ($ty:ty, $slice:ident, $from_vec:ident) => {{ + let src = tensor.data().$slice().ok_or_else(|| { + MinitensorError::invalid_operation( + "repeat_interleave: tensor data access failed".to_string(), + ) + })?; + let mut out = vec![<$ty>::default(); output_numel]; + out.par_chunks_mut(target_dim * inner).enumerate().for_each( + |(outer_idx, out_chunk)| { + let mut dst_offset = 0; + let base = outer_idx * dim_size * inner; + for (i, &rep) in reps.iter().enumerate() { + if rep == 0 { + continue; + } + let src_start = base + i * inner; + let src_slice = &src[src_start..src_start + inner]; + for _ in 0..rep { + let end = dst_offset + inner; + out_chunk[dst_offset..end].copy_from_slice(src_slice); + dst_offset = end; + } + } + }, + ); + TensorData::$from_vec(out, device) + }}; + } + + let data = match dtype { + DataType::Float32 => repeat_impl!(f32, as_f32_slice, from_vec_f32), + DataType::Float64 => repeat_impl!(f64, as_f64_slice, from_vec_f64), + DataType::Int32 => repeat_impl!(i32, as_i32_slice, from_vec_i32), + DataType::Int64 => repeat_impl!(i64, as_i64_slice, from_vec_i64), + DataType::Bool => repeat_impl!(bool, as_bool_slice, from_vec_bool), + }; + + let mut result = Tensor::new( + Arc::new(data), + output_shape_obj, + dtype, + device, + requires_grad, + ); + + if requires_grad { + let grad_fn = build_grad_fn(reps); + result.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&result, Some(grad_fn))?; + } + + Ok(result) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::{ + device::Device, + tensor::{DataType, TensorData}, + }; + + fn create_test_tensor_f32(data: Vec, shape: Vec, requires_grad: bool) -> Tensor { + let shape_obj = Shape::new(shape); + let mut tensor_data = TensorData::zeros(shape_obj.numel(), DataType::Float32); + + if let Some(slice) = tensor_data.as_f32_slice_mut() { + slice.copy_from_slice(&data); + } + + Tensor::new( + Arc::new(tensor_data), + shape_obj, + DataType::Float32, + Device::cpu(), + requires_grad, + ) + } + + #[test] + fn test_reshape_basic() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3], false); + + let reshaped = reshape(&tensor, Shape::new(vec![3, 2])).unwrap(); + + assert_eq!(reshaped.shape().dims(), &[3, 2]); + assert_eq!(reshaped.numel(), 6); + + let data = reshaped.data().as_f32_slice().unwrap(); + assert_eq!(data, &[1.0, 2.0, 3.0, 4.0, 5.0, 6.0]); + } + + #[test] + fn test_reshape_invalid_size() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); + + let result = reshape(&tensor, Shape::new(vec![2, 3])); + assert!(result.is_err()); + } + + #[test] + fn test_reshape_infer_dim() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![6], false); + let reshaped = reshape_with_inference(&tensor, vec![2, -1]).unwrap(); + assert_eq!(reshaped.shape().dims(), &[2, 3]); + } + + #[test] + fn test_reshape_multiple_negative_one_error() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![4], false); + let result = reshape_with_inference(&tensor, vec![-1, -1]); + assert!(result.is_err()); + } + + #[test] + fn test_reshape_infer_mismatch_error() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0, 5.0], vec![5], false); + let result = reshape_with_inference(&tensor, vec![4, -1]); + assert!(result.is_err()); + } + + #[test] + fn test_reshape_zero_dim_with_inference_error() { + let tensor = create_test_tensor_f32(vec![], vec![0], false); + let result = reshape_with_inference(&tensor, vec![-1, 0]); + assert!(result.is_err()); + } + + #[test] + fn test_squeeze_specific_dim() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![1, 4, 1], false); + + let s0 = squeeze(&tensor, Some(0)).unwrap(); + assert_eq!(s0.shape().dims(), &[4, 1]); + + let s1 = squeeze(&s0, Some(1)).unwrap(); + assert_eq!(s1.shape().dims(), &[4]); + + let s_neg = squeeze(&tensor, Some(-1)).unwrap(); + assert_eq!(s_neg.shape().dims(), &[1, 4]); + } + + #[test] + fn test_squeeze_all() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![1, 4, 1], false); + + let squeezed = squeeze(&tensor, None).unwrap(); + assert_eq!(squeezed.shape().dims(), &[4]); + + let scalar = create_test_tensor_f32(vec![1.0], vec![1, 1], false); + let s = squeeze(&scalar, None).unwrap(); + assert!(s.shape().dims().is_empty()); + } + + #[test] + fn test_squeeze_out_of_range() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); + + assert!(squeeze(&tensor, Some(2)).is_err()); + assert!(squeeze(&tensor, Some(-3)).is_err()); + } + + #[test] + fn test_unsqueeze() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![4], false); + + let u0 = unsqueeze(&tensor, 0).unwrap(); + assert_eq!(u0.shape().dims(), &[1, 4]); + + let u1 = unsqueeze(&tensor, 1).unwrap(); + assert_eq!(u1.shape().dims(), &[4, 1]); + + let u_neg = unsqueeze(&tensor, -1).unwrap(); + assert_eq!(u_neg.shape().dims(), &[4, 1]); + } + + #[test] + fn test_gradient_tracking() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], true); + + let reshaped = reshape(&tensor, Shape::new(vec![4])).unwrap(); + + assert!(reshaped.requires_grad()); + assert!(reshaped.grad_fn().is_some()); + } + + #[test] + fn test_concatenate_validation() { + let tensor1 = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + let tensor2 = create_test_tensor_f32(vec![3.0, 4.0], vec![2], false); + + let result = concatenate(&[&tensor1, &tensor2], 0).unwrap(); + assert_eq!(result.shape().dims(), &[4]); + let data = result.data().as_f32_slice().unwrap(); + assert_eq!(data, &[1.0, 2.0, 3.0, 4.0]); + } + + #[test] + fn test_index_select_validation() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3], false); + + let result = index_select(&tensor, 1, &[0, 2]).unwrap(); + assert_eq!(result.shape().dims(), &[2, 2]); + let data = result.data().as_f32_slice().unwrap(); + assert_eq!(data, &[1.0, 3.0, 4.0, 6.0]); + } + + #[test] + fn test_slice_empty_range() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); + + let result = slice(&tensor, 1, 1, 1, 1).unwrap(); + assert_eq!(result.shape().dims(), &[2, 0]); + assert_eq!(result.numel(), 0); + } + + #[test] + fn test_slice_empty_at_end() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); + + let result = slice(&tensor, 0, 2, 2, 1).unwrap(); + assert_eq!(result.shape().dims(), &[0, 2]); + assert_eq!(result.numel(), 0); + } + + #[test] + fn test_index_select_empty_indices() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], false); + + let result = index_select(&tensor, 1, &[]).unwrap(); + assert_eq!(result.shape().dims(), &[2, 0]); + assert_eq!(result.numel(), 0); + } + + #[test] + fn test_slice_validation() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3], false); + + let result = slice(&tensor, 1, 0, 2, 1).unwrap(); + assert_eq!(result.shape().dims(), &[2, 2]); + let data = result.data().as_f32_slice().unwrap(); + assert_eq!(data, &[1.0, 2.0, 4.0, 5.0]); + } + + #[test] + fn test_repeat_basic() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + let repeated = repeat(&tensor, &[3]).unwrap(); + assert_eq!(repeated.shape().dims(), &[6]); + let data = repeated.data().as_f32_slice().unwrap(); + assert_eq!(data, &[1.0, 2.0, 1.0, 2.0, 1.0, 2.0]); + } + + #[test] + fn test_repeat_zero_numel_shape() { + let tensor = create_test_tensor_f32(vec![], vec![0, 2], false); + let repeated = repeat(&tensor, &[2, 3]).unwrap(); + assert_eq!(repeated.shape().dims(), &[0, 6]); + assert_eq!(repeated.numel(), 0); + } + + #[test] + fn test_repeat_dim_mismatch_error() { + let tensor = create_test_tensor_f32(vec![1.0, 2.0], vec![2], false); + assert!(repeat(&tensor, &[]).is_err()); + } +} diff --git a/engine/src/operations/shape_ops/reshape.rs b/engine/src/operations/shape_ops/reshape.rs index c21ebeba..91b0bb69 100644 --- a/engine/src/operations/shape_ops/reshape.rs +++ b/engine/src/operations/shape_ops/reshape.rs @@ -1,1172 +1,1216 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -use crate::{ - autograd::{ - ConcatBackward, GatherBackward, IndexSelectBackward, RepeatBackward, - RepeatInterleaveBackward, ReshapeBackward, add_to_graph, - }, - device::Device, - error::{MinitensorError, Result}, - tensor::{DataType, Shape, Tensor, TensorData}, -}; -use rayon::prelude::*; -use std::sync::Arc; - -fn normalize_dim(dim: isize, ndim: usize) -> Result { - let dim = if dim < 0 { dim + ndim as isize } else { dim }; - if dim < 0 || dim >= ndim as isize { - Err(MinitensorError::index_error(dim, 0, ndim)) - } else { - Ok(dim as usize) - } -} - -fn empty_tensor(shape: Shape, dtype: DataType, device: Device, requires_grad: bool) -> Tensor { - Tensor::new( - Arc::new(TensorData::zeros_on_device(0, dtype, device)), - shape, - dtype, - device, - requires_grad, - ) -} - -fn checked_repeat_dim(size: usize, repeat: usize) -> Result { - size.checked_mul(repeat).ok_or_else(|| { - MinitensorError::invalid_operation( - "repeat output dimensions exceed supported size".to_string(), - ) - }) -} - -fn checked_repeat_numel(dims: &[usize]) -> Result { - if dims.is_empty() { - return Ok(1); - } - - dims.iter().try_fold(1usize, |numel, &dim| { - numel.checked_mul(dim).ok_or_else(|| { - MinitensorError::invalid_operation( - "repeat output dimensions exceed supported size".to_string(), - ) - }) - }) -} - -fn attach_repeat_backward(mut output: Tensor, input: &Tensor, repeats: &[usize]) -> Result { - if input.requires_grad() && input.dtype().is_float() { - output.refresh_autograd_metadata(); - let mut output = output.requires_grad_(true); - let grad_fn = Arc::new(RepeatBackward { - input_id: input.id(), - input_shape: input.shape().dims().to_vec(), - repeats: repeats.to_vec(), - }); - output.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output, Some(grad_fn))?; - Ok(output) - } else { - Ok(output) - } -} - -/// Reshape operation with gradient support -pub fn reshape(tensor: &Tensor, new_shape: Shape) -> Result { - // Check if the total number of elements matches - if tensor.numel() != new_shape.numel() { - return Err(MinitensorError::shape_mismatch( - vec![tensor.numel()], - vec![new_shape.numel()], - )); - } - - // Use the tensor's built-in view method for reshaping and refresh metadata - let mut reshaped = tensor.view(new_shape.clone())?; - reshaped.refresh_autograd_metadata(); - - // Set up gradient function if needed - if reshaped.requires_grad() { - let grad_fn = Arc::new(ReshapeBackward { - input_shape: tensor.shape().dims().to_vec(), - input_id: tensor.id(), - }); - - reshaped.set_grad_fn(Some(grad_fn.clone())); - - // Add to computation graph - add_to_graph(&reshaped, Some(grad_fn))?; - - Ok(reshaped) - } else { - Ok(reshaped) - } -} - -/// This wrapper performs validation and inference for a single ``-1`` -/// dimension before delegating to [`reshape`]. -pub fn reshape_with_inference(tensor: &Tensor, dims: Vec) -> Result { - let mut out_dims = Vec::with_capacity(dims.len()); - let mut inferred_index: Option = None; - let mut known_product: usize = 1; - - for (index, &dim) in dims.iter().enumerate() { - if dim == -1 { - if inferred_index.is_some() { - return Err(MinitensorError::invalid_operation( - "can only specify one -1 dimension in reshape".to_string(), - )); - } - inferred_index = Some(index); - out_dims.push(0); - continue; - } - - if dim < 0 { - return Err(MinitensorError::invalid_operation( - "invalid negative dimension".to_string(), - )); - } - - let dim_usize = dim as usize; - known_product = known_product.checked_mul(dim_usize).ok_or_else(|| { - MinitensorError::invalid_operation("reshape dimensions exceed supported size") - })?; - out_dims.push(dim_usize); - } - - let total_elements = tensor.numel(); - if let Some(index) = inferred_index { - if known_product == 0 { - return Err(MinitensorError::invalid_operation( - "cannot reshape tensor with -1 and 0 dimensions".to_string(), - )); - } - - if total_elements % known_product != 0 { - return Err(MinitensorError::invalid_operation( - "cannot infer reshape dimension".to_string(), - )); - } - - out_dims[index] = total_elements / known_product; - } else if known_product != total_elements { - return Err(MinitensorError::shape_mismatch( - vec![total_elements], - vec![known_product], - )); - } - - reshape(tensor, Shape::new(out_dims)) -} - - -/// Squeeze operation - remove dimensions of size 1. -/// -/// Routed through [`reshape`] so the result is a first-class differentiable node -/// (the plain `Tensor::squeeze` view shares the input's id and would attribute a -/// wrongly-shaped gradient to it). -pub fn squeeze(tensor: &Tensor, dim: Option) -> Result { - let dims = tensor.shape().dims(); - let new_dims: Vec = match dim { - None => dims.iter().copied().filter(|&d| d != 1).collect(), - Some(d) => { - let ndim = tensor.ndim() as isize; - let d = if d < 0 { d + ndim } else { d }; - if d < 0 || d >= ndim { - return Err(MinitensorError::index_error(d, 0, tensor.ndim())); - } - let d = d as usize; - if dims[d] != 1 { - // A non-unit axis remains untouched. - dims.to_vec() - } else { - let mut v = dims.to_vec(); - v.remove(d); - v - } - } - }; - reshape(tensor, Shape::new(new_dims)) -} - -/// Unsqueeze operation - add a dimension of size 1. See [`squeeze`] for why this -/// goes through [`reshape`] rather than the view-based `Tensor::unsqueeze`. -pub fn unsqueeze(tensor: &Tensor, dim: isize) -> Result { - let ndim = tensor.ndim() as isize; - let d = if dim < 0 { dim + ndim + 1 } else { dim }; - if d < 0 || d > ndim { - return Err(MinitensorError::index_error(d, 0, (ndim + 1) as usize)); - } - let mut new_dims = tensor.shape().dims().to_vec(); - new_dims.insert(d as usize, 1); - reshape(tensor, Shape::new(new_dims)) -} - -/// Flatten dimensions `start_dim..=end_dim` into one. Routed through [`reshape`] -/// so gradients flow (see [`squeeze`]). -pub fn flatten(tensor: &Tensor, start_dim: isize, end_dim: isize) -> Result { - let ndim = tensor.ndim() as isize; - let start = if start_dim < 0 { - start_dim + ndim - } else { - start_dim - }; - let end = if end_dim < 0 { end_dim + ndim } else { end_dim }; - if start < 0 || start >= ndim { - return Err(MinitensorError::index_error(start, 0, tensor.ndim())); - } - if end < 0 || end >= ndim { - return Err(MinitensorError::index_error(end, 0, tensor.ndim())); - } - if start > end { - return Err(MinitensorError::invalid_argument( - "start_dim must be less than or equal to end_dim", - )); - } - - let (start, end) = (start as usize, end as usize); - let dims = tensor.shape().dims(); - let mut new_dims = dims[..start].to_vec(); - new_dims.push(dims[start..=end].iter().product()); - new_dims.extend_from_slice(&dims[end + 1..]); - reshape(tensor, Shape::new(new_dims)) -} - -/// Permute tensor dimensions according to `dims` -pub fn permute(tensor: &Tensor, dims: Vec) -> Result { - let ndim = tensor.ndim(); - - // Validate number of dimensions - if dims.len() != ndim { - return Err(MinitensorError::invalid_operation( - "dims must match number of dimensions".to_string(), - )); - } - - // Normalise negative dimensions and validate range - let mut normalized = Vec::with_capacity(ndim); - for &d in &dims { - let d = if d < 0 { d + ndim as isize } else { d }; - if d < 0 || d >= ndim as isize { - return Err(MinitensorError::index_error(d, 0, ndim)); - } - normalized.push(d as usize); - } - // Check that dims form a proper permutation - let mut sorted = normalized.clone(); - sorted.sort_unstable(); - if sorted != (0..ndim).collect::>() { - return Err(MinitensorError::invalid_operation( - "dims must be a permutation of dimensions".to_string(), - )); - } - - // Apply sequence of transposes to achieve the permutation - let mut result = tensor.clone(); - let mut current: Vec = (0..ndim).collect(); - for i in 0..ndim { - let target = normalized[i]; - let j = current.iter().position(|&x| x == target).unwrap(); - if i != j { - result = result.transpose(i as isize, j as isize)?; - current.swap(i, j); - } - } - - Ok(result) -} - -/// Move tensor dimensions to new positions, keeping relative order of other dims -pub fn movedim(tensor: &Tensor, source: &[isize], destination: &[isize]) -> Result { - let ndim = tensor.ndim(); - - if source.len() != destination.len() { - return Err(MinitensorError::invalid_operation( - "movedim: source and destination must have the same length".to_string(), - )); - } - - let mut src_seen = vec![false; ndim]; - let mut dst_seen = vec![false; ndim]; - let mut pairs: Vec<(usize, usize)> = Vec::with_capacity(source.len()); - - for (&s, &d) in source.iter().zip(destination.iter()) { - let s = if s < 0 { s + ndim as isize } else { s }; - if s < 0 || s >= ndim as isize { - return Err(MinitensorError::index_error(s, 0, ndim)); - } - let s = s as usize; - if src_seen[s] { - return Err(MinitensorError::invalid_operation( - "movedim: duplicate dimensions in source".to_string(), - )); - } - src_seen[s] = true; - let d = if d < 0 { d + ndim as isize } else { d }; - if d < 0 || d >= ndim as isize { - return Err(MinitensorError::index_error(d, 0, ndim)); - } - let d = d as usize; - if dst_seen[d] { - return Err(MinitensorError::invalid_operation( - "movedim: duplicate dimensions in destination".to_string(), - )); - } - dst_seen[d] = true; - pairs.push((d, s)); - } - - // Build permutation order - let mut order: Vec = (0..ndim).filter(|&i| !src_seen[i]).collect(); - pairs.sort_by_key(|&(d, _)| d); - for (d, s) in pairs { - order.insert(d, s); - } - let order_isize: Vec = order.into_iter().map(|v| v as isize).collect(); - permute(tensor, order_isize) -} - -/// Concatenate tensors along a specified dimension -pub fn concatenate(tensors: &[&Tensor], dim: isize) -> Result { - if tensors.is_empty() { - return Err(MinitensorError::invalid_operation( - "Cannot concatenate empty list of tensors", - )); - } - - let first_tensor = tensors[0]; - - // Validate that all tensors have the same number of dimensions - for tensor in tensors.iter().skip(1) { - if tensor.ndim() != first_tensor.ndim() { - return Err(MinitensorError::shape_mismatch( - vec![first_tensor.ndim()], - vec![tensor.ndim()], - )); - } - - // Check device compatibility - if tensor.device() != first_tensor.device() { - return Err(MinitensorError::device_mismatch( - format!("{:?}", first_tensor.device()), - format!("{:?}", tensor.device()), - )); - } - - // Check data type compatibility - if tensor.dtype() != first_tensor.dtype() { - return Err(MinitensorError::type_mismatch( - format!("{:?}", first_tensor.dtype()), - format!("{:?}", tensor.dtype()), - )); - } - } - - // Validate concatenation dimension - let dim = normalize_dim(dim, first_tensor.ndim())?; - - // Validate that all dimensions except the concatenation dimension match - for tensor in tensors.iter().skip(1) { - for (i, (&size1, &size2)) in first_tensor - .shape() - .dims() - .iter() - .zip(tensor.shape().dims().iter()) - .enumerate() - { - if i != dim && size1 != size2 { - return Err(MinitensorError::shape_mismatch( - first_tensor.shape().dims().to_vec(), - tensor.shape().dims().to_vec(), - )); - } - } - } - - if !first_tensor.device().is_cpu() { - return Err(MinitensorError::invalid_operation( - "concatenate currently supports only CPU tensors", - )); - } - - // Compute output shape - let mut output_shape = first_tensor.shape().dims().to_vec(); - output_shape[dim] = tensors.iter().map(|t| t.shape().dims()[dim]).sum(); - let output_shape_obj = Shape::new(output_shape); - - let dtype = first_tensor.dtype(); - let device = first_tensor.device(); - let requires_grad = tensors.iter().any(|t| t.requires_grad()); - - let dims = first_tensor.shape().dims(); - let inner: usize = dims[dim + 1..].iter().product(); - let _outer: usize = dims[..dim].iter().product(); - - if output_shape_obj.numel() == 0 { - let data = TensorData::zeros_on_device(0, dtype, device); - return Ok(Tensor::new( - Arc::new(data), - output_shape_obj, - dtype, - device, - requires_grad, - )); - } - - macro_rules! concat_impl { - ($ty:ty, $slice:ident, $from_vec:ident) => {{ - let mut sources: Vec<&[$ty]> = Vec::with_capacity(tensors.len()); - let mut dim_sizes: Vec = Vec::with_capacity(tensors.len()); - for t in tensors { - let src = t.data().$slice().ok_or_else(|| { - MinitensorError::invalid_operation("Tensor data access failed for concatenate") - })?; - sources.push(src); - dim_sizes.push(t.shape().dims()[dim]); - } - let src_strides: Vec = dim_sizes.iter().map(|&d| d * inner).collect(); - - let mut out = vec![<$ty>::default(); output_shape_obj.numel()]; - let chunk_size = output_shape_obj.dims()[dim] * inner; - out.par_chunks_mut(chunk_size) - .enumerate() - .for_each(|(o, out_chunk)| { - let mut dst_offset = 0; - for (src, &src_stride) in sources.iter().zip(src_strides.iter()) { - let src_start = o * src_stride; - let src_len = src_stride; - out_chunk[dst_offset..dst_offset + src_len] - .copy_from_slice(&src[src_start..src_start + src_len]); - dst_offset += src_len; - } - }); - TensorData::$from_vec(out, device) - }}; - } - - let data = match dtype { - DataType::Float32 => concat_impl!(f32, as_f32_slice, from_vec_f32), - DataType::Float64 => concat_impl!(f64, as_f64_slice, from_vec_f64), - DataType::Int32 => concat_impl!(i32, as_i32_slice, from_vec_i32), - DataType::Int64 => concat_impl!(i64, as_i64_slice, from_vec_i64), - DataType::Bool => concat_impl!(bool, as_bool_slice, from_vec_bool), - }; - - let output = Tensor::new(Arc::new(data), output_shape_obj, dtype, device, requires_grad); - - if requires_grad && dtype.is_float() { - let grad_fn = Arc::new(ConcatBackward { - input_ids: tensors.iter().map(|t| t.id()).collect(), - sizes: tensors.iter().map(|t| t.shape().dims()[dim]).collect(), - dim, - }); - let mut output = output; - output.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output, Some(grad_fn))?; - return Ok(output); - } - - Ok(output) -} - -/// Repeat `tensor` according to `repeats` along each dimension. -pub fn repeat(tensor: &Tensor, repeats: &[usize]) -> Result { - if repeats.len() < tensor.ndim() { - return Err(MinitensorError::invalid_operation( - "number of dimensions of repeat dims can not be smaller than number of dimensions of tensor", - )); - } - - // Tile on a detached view so the intermediate per-dimension copies never - // create graph nodes; a single RepeatBackward maps the final result straight - // back to the original input. - let mut result = tensor.detach(); - - if repeats.len() > result.ndim() { - let mut new_shape = vec![1; repeats.len() - result.ndim()]; - new_shape.extend_from_slice(result.shape().dims()); - result = result.reshape(Shape::new(new_shape))?; - } - - if repeats.iter().any(|&r| r == 0) { - let mut out_shape = result.shape().dims().to_vec(); - for (dim, &rep) in repeats.iter().enumerate() { - out_shape[dim] = checked_repeat_dim(out_shape[dim], rep)?; - } - - let output = empty_tensor( - Shape::new(out_shape), - result.dtype(), - result.device(), - false, - ); - return attach_repeat_backward(output, tensor, repeats); - } - - for (dim, &rep) in repeats.iter().enumerate() { - if rep == 1 { - continue; - } - let dims = result.shape().dims().to_vec(); - let dim_size = dims[dim]; - let inner: usize = dims[dim + 1..].iter().product(); - let repeated_dim = checked_repeat_dim(dim_size, rep)?; - let chunk_size = checked_repeat_dim(repeated_dim, inner)?; - let src_chunk_size = checked_repeat_dim(dim_size, inner)?; - let mut output_shape = dims.clone(); - output_shape[dim] = repeated_dim; - let output_numel = checked_repeat_numel(&output_shape)?; - let output_shape_obj = Shape::new(output_shape); - - let dtype = result.dtype(); - let device = result.device(); - let requires_grad = result.requires_grad(); - - if output_numel == 0 { - result = empty_tensor(output_shape_obj, dtype, device, requires_grad); - continue; - } - - macro_rules! repeat_impl { - ($ty:ty, $slice:ident, $from_vec:ident) => {{ - let src = result.data().$slice().ok_or_else(|| { - MinitensorError::invalid_operation("Tensor data access failed for repeat") - })?; - let mut out = vec![<$ty>::default(); output_numel]; - out.par_chunks_mut(chunk_size) - .enumerate() - .for_each(|(o, out_chunk)| { - let src_start = o * src_chunk_size; - let src_chunk = &src[src_start..src_start + src_chunk_size]; - for r in 0..rep { - let dst_start = r * src_chunk_size; - out_chunk[dst_start..dst_start + src_chunk_size] - .copy_from_slice(src_chunk); - } - }); - TensorData::$from_vec(out, device) - }}; - } - - let data = match dtype { - DataType::Float32 => repeat_impl!(f32, as_f32_slice, from_vec_f32), - DataType::Float64 => repeat_impl!(f64, as_f64_slice, from_vec_f64), - DataType::Int32 => repeat_impl!(i32, as_i32_slice, from_vec_i32), - DataType::Int64 => repeat_impl!(i64, as_i64_slice, from_vec_i64), - DataType::Bool => repeat_impl!(bool, as_bool_slice, from_vec_bool), - }; - - result = Tensor::new( - Arc::new(data), - output_shape_obj, - dtype, - device, - requires_grad, - ); - } - - attach_repeat_backward(result, tensor, repeats) -} - -/// Indexing operation - select elements along specified dimensions -pub fn index_select(tensor: &Tensor, dim: isize, indices: &[usize]) -> Result { - let dim = normalize_dim(dim, tensor.ndim())?; - - let dim_size = tensor.shape().dims()[dim]; - - // Validate indices - for &idx in indices { - if idx >= dim_size { - return Err(MinitensorError::index_error(idx as isize, 0, dim_size)); - } - } - - if !tensor.device().is_cpu() { - return Err(MinitensorError::invalid_operation( - "index_select currently supports only CPU tensors", - )); - } - - // Compute output shape - let mut output_shape = tensor.shape().dims().to_vec(); - output_shape[dim] = indices.len(); - let output_shape_vec = output_shape.clone(); - let output_shape_obj = Shape::new(output_shape); - - let dtype = tensor.dtype(); - let device = tensor.device(); - let requires_grad = tensor.requires_grad(); - - if output_shape_obj.numel() == 0 { - return Ok(empty_tensor(output_shape_obj, dtype, device, requires_grad)); - } - - let dims = tensor.shape().dims(); - let inner: usize = dims[dim + 1..].iter().product(); - let _outer: usize = dims[..dim].iter().product(); - - macro_rules! index_impl { - ($ty:ty, $slice:ident, $from_vec:ident) => {{ - let src = tensor.data().$slice().ok_or_else(|| { - MinitensorError::invalid_operation("Tensor data access failed for index_select") - })?; - let mut out = vec![<$ty>::default(); output_shape_obj.numel()]; - out.par_chunks_mut(output_shape_vec[dim] * inner) - .enumerate() - .for_each(|(o, out_chunk)| { - for (i, &idx) in indices.iter().enumerate() { - let src_start = o * dims[dim] * inner + idx * inner; - let dst_start = i * inner; - out_chunk[dst_start..dst_start + inner] - .copy_from_slice(&src[src_start..src_start + inner]); - } - }); - TensorData::$from_vec(out, device) - }}; - } - - let data = match dtype { - DataType::Float32 => index_impl!(f32, as_f32_slice, from_vec_f32), - DataType::Float64 => index_impl!(f64, as_f64_slice, from_vec_f64), - DataType::Int32 => index_impl!(i32, as_i32_slice, from_vec_i32), - DataType::Int64 => index_impl!(i64, as_i64_slice, from_vec_i64), - DataType::Bool => index_impl!(bool, as_bool_slice, from_vec_bool), - }; - - let output = Tensor::new(Arc::new(data), output_shape_obj, dtype, device, requires_grad); - - if requires_grad && dtype.is_float() { - let grad_fn = Arc::new(IndexSelectBackward { - input_id: tensor.id(), - input_shape: tensor.shape().dims().to_vec(), - dim, - indices: indices.to_vec(), - }); - let mut output = output; - output.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output, Some(grad_fn))?; - return Ok(output); - } - - Ok(output) -} - -/// Gather operation - collect elements along a dimension using an index tensor -pub fn gather(tensor: &Tensor, dim: isize, index: &Tensor) -> Result { - let dim = normalize_dim(dim, tensor.ndim())?; - - if index.ndim() != tensor.ndim() { - return Err(MinitensorError::invalid_operation( - "gather index tensor must have the same number of dimensions as input", - )); - } - - if index.dtype() != DataType::Int64 { - return Err(MinitensorError::invalid_operation( - "gather indices must be int64", - )); - } - - let input_dims = tensor.shape().dims(); - let index_dims = index.shape().dims(); - - // Validate shapes except at gather dimension - for (i, (&idx_d, &in_d)) in index_dims.iter().zip(input_dims.iter()).enumerate() { - if i != dim && idx_d != in_d { - return Err(MinitensorError::shape_mismatch( - input_dims.to_vec(), - index_dims.to_vec(), - )); - } - } - - let dim_size = input_dims[dim]; - - // Validate indices - let idx_slice = index - .data() - .as_i64_slice() - .ok_or_else(|| MinitensorError::invalid_operation("gather indices must be int64"))?; - for &v in idx_slice { - if v < 0 || v as usize >= dim_size { - return Err(MinitensorError::index_error(v as isize, 0, dim_size)); - } - } - - if !tensor.device().is_cpu() { - return Err(MinitensorError::invalid_operation( - "gather currently supports only CPU tensors", - )); - } - - let inner: usize = input_dims[dim + 1..].iter().product(); - let idx_dim = index_dims[dim]; - - let dtype = tensor.dtype(); - let device = tensor.device(); - let requires_grad = tensor.requires_grad(); - let output_shape_obj = Shape::new(index_dims.to_vec()); - let output_numel = idx_slice.len(); - - if output_numel == 0 { - return Ok(empty_tensor(output_shape_obj, dtype, device, requires_grad)); - } - - macro_rules! gather_impl { - ($ty:ty, $slice:ident, $from_vec:ident) => {{ - let src = tensor.data().$slice().ok_or_else(|| { - MinitensorError::invalid_operation("Tensor data access failed for gather") - })?; - let idx = idx_slice; - let mut out = vec![<$ty>::default(); output_numel]; - let chunk_size = idx_dim * inner; - if output_numel % chunk_size != 0 { - return Err(MinitensorError::internal_error(format!( - "gather output length ({output_numel}) is not divisible by chunk size ({chunk_size})" - ))); - } - out.par_chunks_mut(chunk_size) - .enumerate() - .for_each(|(o, out_chunk)| { - let base = o * dim_size * inner; - let idx_chunk = &idx[o * chunk_size..(o + 1) * chunk_size]; - for i in 0..idx_dim { - let idx_row = &idx_chunk[i * inner..(i + 1) * inner]; - let dst_row = &mut out_chunk[i * inner..(i + 1) * inner]; - for (j, &gather_val) in idx_row.iter().enumerate() { - let gather_idx = gather_val as usize; - dst_row[j] = src[base + gather_idx * inner + j]; - } - } - }); - TensorData::$from_vec(out, device) - }}; - } - - let data = match dtype { - DataType::Float32 => gather_impl!(f32, as_f32_slice, from_vec_f32), - DataType::Float64 => gather_impl!(f64, as_f64_slice, from_vec_f64), - DataType::Int32 => gather_impl!(i32, as_i32_slice, from_vec_i32), - DataType::Int64 => gather_impl!(i64, as_i64_slice, from_vec_i64), - DataType::Bool => gather_impl!(bool, as_bool_slice, from_vec_bool), - }; - - let output = Tensor::new(Arc::new(data), output_shape_obj, dtype, device, requires_grad); - - if requires_grad && dtype.is_float() { - let grad_fn = Arc::new(GatherBackward { - input_id: tensor.id(), - input_shape: input_dims.to_vec(), - dim, - index: idx_slice.to_vec(), - }); - let mut output = output; - output.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output, Some(grad_fn))?; - return Ok(output); - } - - Ok(output) -} - -/// Slicing operation - select a contiguous range of elements -pub fn slice(tensor: &Tensor, dim: isize, start: usize, end: usize, step: usize) -> Result { - let dim = normalize_dim(dim, tensor.ndim())?; - - let dim_size = tensor.shape().dims()[dim]; - - if start > dim_size || end > dim_size || start > end { - return Err(MinitensorError::invalid_operation(format!( - "Invalid slice range: start={}, end={}, dim_size={}", - start, end, dim_size - ))); - } - - if step == 0 { - return Err(MinitensorError::invalid_operation( - "Slice step cannot be zero", - )); - } - - if !tensor.device().is_cpu() { - return Err(MinitensorError::invalid_operation( - "slice currently supports only CPU tensors", - )); - } - - // Compute output shape - let mut output_shape = tensor.shape().dims().to_vec(); - output_shape[dim] = (end - start).div_ceil(step); - let output_shape_obj = Shape::new(output_shape); - - let dtype = tensor.dtype(); - let device = tensor.device(); - let requires_grad = tensor.requires_grad(); - - if output_shape_obj.numel() == 0 { - return Ok(empty_tensor(output_shape_obj, dtype, device, requires_grad)); - } - - let dims = tensor.shape().dims(); - let inner: usize = dims[dim + 1..].iter().product(); - let count = output_shape_obj.dims()[dim]; - - macro_rules! slice_impl { - ($ty:ty, $slice:ident, $from_vec:ident) => {{ - let src = tensor.data().$slice().ok_or_else(|| { - MinitensorError::invalid_operation("Tensor data access failed for slice") - })?; - let mut out = vec![<$ty>::default(); output_shape_obj.numel()]; - out.par_chunks_mut(count * inner) - .enumerate() - .for_each(|(o, out_chunk)| { - for i in 0..count { - let src_idx = start + i * step; - let src_start = o * dims[dim] * inner + src_idx * inner; - let dst_start = i * inner; - out_chunk[dst_start..dst_start + inner] - .copy_from_slice(&src[src_start..src_start + inner]); - } - }); - TensorData::$from_vec(out, device) - }}; - } - - let data = match dtype { - DataType::Float32 => slice_impl!(f32, as_f32_slice, from_vec_f32), - DataType::Float64 => slice_impl!(f64, as_f64_slice, from_vec_f64), - DataType::Int32 => slice_impl!(i32, as_i32_slice, from_vec_i32), - DataType::Int64 => slice_impl!(i64, as_i64_slice, from_vec_i64), - DataType::Bool => slice_impl!(bool, as_bool_slice, from_vec_bool), - }; - - let output = Tensor::new(Arc::new(data), output_shape_obj, dtype, device, requires_grad); - - if requires_grad && dtype.is_float() { - // A slice selects source positions `start, start+step, ...` along `dim`; - // its backward scatters the gradient back to exactly those positions. - let indices: Vec = (0..count).map(|i| start + i * step).collect(); - let grad_fn = Arc::new(IndexSelectBackward { - input_id: tensor.id(), - input_shape: tensor.shape().dims().to_vec(), - dim, - indices, - }); - let mut output = output; - output.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output, Some(grad_fn))?; - return Ok(output); - } - - Ok(output) -} - -/// Narrow tensor along a dimension starting at `start` for `length` elements -pub fn narrow(tensor: &Tensor, dim: isize, start: usize, length: usize) -> Result { - let dim = normalize_dim(dim, tensor.ndim())?; - let dim_size = tensor.shape().dims()[dim]; - - if start > dim_size { - return Err(MinitensorError::index_error(start as isize, 0, dim_size)); - } - if start + length > dim_size { - return Err(MinitensorError::index_error( - (start + length) as isize, - 0, - dim_size, - )); - } - - if length == 0 { - let mut out_shape = tensor.shape().dims().to_vec(); - out_shape[dim] = 0; - return Ok(Tensor::zeros( - Shape::new(out_shape), - tensor.dtype(), - tensor.device(), - tensor.requires_grad(), - )); - } - - slice(tensor, dim as isize, start, start + length, 1) -} - -/// Flip tensor elements along specified dimensions. -pub fn flip(tensor: &Tensor, dims: &[isize]) -> Result { - let ndim = tensor.ndim(); - let mut normalized = Vec::with_capacity(dims.len()); - for &d in dims { - let dim = normalize_dim(d, ndim)?; - if normalized.contains(&dim) { - return Err(MinitensorError::invalid_operation( - "dims must be unique".to_string(), - )); - } - normalized.push(dim); - } - - let mut result = tensor.clone(); - for &dim in &normalized { - let size = result.shape().dims()[dim]; - let indices: Vec = (0..size).rev().collect(); - result = index_select(&result, dim as isize, &indices)?; - } - - Ok(result) -} - -/// Roll tensor elements along specified dimensions with wrap-around -pub fn roll(tensor: &Tensor, shifts: &[isize], dims: Option<&[isize]>) -> Result { - // Compute the roll on a detached view so the internal slice/concatenate steps - // (which flatten to a storage-sharing view in the `dims == None` case) never - // build gradient edges; a single RollBackward inverts the whole operation. - let track_grad = tensor.requires_grad() && tensor.dtype().is_float(); - let base = if track_grad { - tensor.detach() - } else { - tensor.clone() - }; - let output = roll_forward(&base, shifts, dims)?; - - if track_grad { - let mut output = output; - output.refresh_autograd_metadata(); - let mut output = output.requires_grad_(true); - let grad_fn = Arc::new(crate::autograd::RollBackward { - input_id: tensor.id(), - shifts: shifts.to_vec(), - dims: dims.map(|d| d.to_vec()), - }); - output.set_grad_fn(Some(grad_fn.clone())); - add_to_graph(&output, Some(grad_fn))?; - return Ok(output); - } - - Ok(output) -} - -fn roll_forward(tensor: &Tensor, shifts: &[isize], dims: Option<&[isize]>) -> Result { - if let Some(dims) = dims { - if shifts.len() != dims.len() { - return Err(MinitensorError::invalid_operation( - "shifts and dims must have the same length".to_string(), - )); - } - let mut result = tensor.clone(); - for (&shift, &dim) in shifts.iter().zip(dims.iter()) { - let dim = normalize_dim(dim, result.ndim())?; - let size = result.shape().dims()[dim] as isize; - if size == 0 { - continue; - } - let k = ((shift % size) + size) % size; - if k == 0 { - continue; - } - let split_point = (size - k) as usize; - let first = slice(&result, dim as isize, 0, split_point, 1)?; - let second = slice(&result, dim as isize, split_point, size as usize, 1)?; - result = concatenate(&[&second, &first], dim as isize)?; - } - Ok(result) - } else { - if shifts.len() != 1 { - return Err(MinitensorError::invalid_operation( - "shifts must contain a single value when dims is None".to_string(), - )); - } - let shift = shifts[0]; - let flat = tensor.flatten_all()?; - let size = flat.shape().dims()[0] as isize; - if size == 0 { - return flat.reshape(tensor.shape().clone()); - } - let k = ((shift % size) + size) % size; - if k == 0 { - return flat.reshape(tensor.shape().clone()); - } - let split_point = (size - k) as usize; - let first = slice(&flat, 0, 0, split_point, 1)?; - let second = slice(&flat, 0, split_point, size as usize, 1)?; - let rolled = concatenate(&[&second, &first], 0)?; - rolled.reshape(tensor.shape().clone()) - } -} - -/// Specification of repeat counts accepted by [`repeat_interleave`]. -#[derive(Clone, Copy)] -pub enum RepeatInterleaveSpec<'a> { - /// A single repeat value applied to every element along ``dim``. - Scalar(usize), - /// Explicit repeat counts provided as a slice. - Slice(&'a [usize]), - /// Repeat counts provided as a tensor (must contain integer data). - Tensor(&'a Tensor), -} - -fn collect_repeats_from_values(len: usize, values: I) -> Result> -where - I: IntoIterator, -{ - let mut out = Vec::with_capacity(len); - for value in values { - if value < 0 { - return Err(MinitensorError::invalid_operation( - "repeat_interleave: repeats must be non-negative".to_string(), - )); - } - out.push(value as usize); - } - Ok(out) -} - -fn collect_repeats_from_tensor(tensor: &Tensor, dim_size: usize) -> Result> { - if !tensor.device().is_cpu() { - return Err(MinitensorError::invalid_operation( - "repeat_interleave: repeats tensor must reside on CPU".to_string(), - )); - } - - if tensor.numel() != dim_size { - return Err(MinitensorError::invalid_operation( - "repeat_interleave: repeats tensor must have the same number of elements as the selected dimension" - .to_string(), - )); - } - - match tensor.dtype() { - DataType::Int32 => { - let slice = tensor.data().as_i32_slice().ok_or_else(|| { - MinitensorError::invalid_operation( - "repeat_interleave: repeats tensor must be contiguous".to_string(), - ) - })?; - collect_repeats_from_values(slice.len(), slice.iter().map(|&value| value as i64)) - } - DataType::Int64 => { - let slice = tensor.data().as_i64_slice().ok_or_else(|| { - MinitensorError::invalid_operation( - "repeat_interleave: repeats tensor must be contiguous".to_string(), - ) - })?; - collect_repeats_from_values(slice.len(), slice.iter().copied()) - } - other => Err(MinitensorError::type_mismatch( - "integral tensor", - format!("{:?}", other), - )), - } -} - -#[cfg(test)] -mod reshape_tests { - use super::*; - - #[test] - fn reshape_with_inference_rejects_overflowing_known_product() { - let tensor = Tensor::zeros(Shape::new(vec![1]), DataType::Float32, Device::cpu(), false); - - let result = reshape_with_inference(&tensor, vec![isize::MAX, isize::MAX, -1]); - - assert!(result.is_err()); - assert!( - result - .unwrap_err() - .to_string() - .contains("reshape dimensions exceed supported size") - ); - } - - #[test] - fn reshape_with_inference_rejects_multiple_inferred_dimensions() { - let tensor = Tensor::zeros(Shape::new(vec![4]), DataType::Float32, Device::cpu(), false); - - let result = reshape_with_inference(&tensor, vec![-1, -1]); - - assert!(result.is_err()); - assert!( - result - .unwrap_err() - .to_string() - .contains("can only specify one -1 dimension in reshape") - ); - } - - #[test] - fn reshape_with_inference_rejects_invalid_negative_dimension() { - let tensor = Tensor::zeros(Shape::new(vec![1]), DataType::Float32, Device::cpu(), false); - - let result = reshape_with_inference(&tensor, vec![-2, 1]); - - assert!(result.is_err()); - assert!( - result - .unwrap_err() - .to_string() - .contains("invalid negative dimension") - ); - } - - #[test] - fn reshape_with_inference_rejects_overflowing_shape_without_inference() { - let tensor = Tensor::zeros(Shape::new(vec![1]), DataType::Float32, Device::cpu(), false); - - let result = reshape_with_inference(&tensor, vec![isize::MAX, isize::MAX]); - - assert!(result.is_err()); - assert!( - result - .unwrap_err() - .to_string() - .contains("reshape dimensions exceed supported size") - ); - } - - #[test] - fn reshape_with_inference_rejects_zero_dimension_with_inference() { - let tensor = Tensor::zeros(Shape::new(vec![0]), DataType::Float32, Device::cpu(), false); - - let result = reshape_with_inference(&tensor, vec![-1, 0]); - - assert!(result.is_err()); - assert!( - result - .unwrap_err() - .to_string() - .contains("cannot reshape tensor with -1 and 0 dimensions") - ); - } - - #[test] - fn reshape_with_inference_no_inference_shape_mismatch() { - let tensor = Tensor::zeros(Shape::new(vec![5]), DataType::Float32, Device::cpu(), false); - - let result = reshape_with_inference(&tensor, vec![2, 2]); - - assert!(result.is_err()); - assert!(result.unwrap_err().to_string().contains("Shape mismatch")); - } - - #[test] - fn reshape_with_inference_no_inference_shape_match() { - let tensor = Tensor::zeros(Shape::new(vec![6]), DataType::Float32, Device::cpu(), false); - - let result = reshape_with_inference(&tensor, vec![2, 3]); - - assert!(result.is_ok()); - assert_eq!(result.expect("reshape should succeed").shape().dims(), &[2, 3]); - } - - #[test] - fn reshape_with_inference_infers_single_negative_dimension() { - let tensor = Tensor::zeros(Shape::new(vec![12]), DataType::Float32, Device::cpu(), false); - - let reshaped = reshape_with_inference(&tensor, vec![3, -1]).expect("reshape should work"); - - assert_eq!(reshaped.shape().dims(), &[3, 4]); - } -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use crate::{ + autograd::{ + ConcatBackward, GatherBackward, IndexSelectBackward, RepeatBackward, ReshapeBackward, + add_to_graph, + }, + device::Device, + error::{MinitensorError, Result}, + tensor::{DataType, Shape, Tensor, TensorData}, +}; +use rayon::prelude::*; +use std::sync::Arc; + +pub(crate) fn normalize_dim(dim: isize, ndim: usize) -> Result { + let dim = if dim < 0 { dim + ndim as isize } else { dim }; + if dim < 0 || dim >= ndim as isize { + Err(MinitensorError::index_error(dim, 0, ndim)) + } else { + Ok(dim as usize) + } +} + +fn empty_tensor(shape: Shape, dtype: DataType, device: Device, requires_grad: bool) -> Tensor { + Tensor::new( + Arc::new(TensorData::zeros_on_device(0, dtype, device)), + shape, + dtype, + device, + requires_grad, + ) +} + +fn checked_repeat_dim(size: usize, repeat: usize) -> Result { + size.checked_mul(repeat).ok_or_else(|| { + MinitensorError::invalid_operation( + "repeat output dimensions exceed supported size".to_string(), + ) + }) +} + +fn checked_repeat_numel(dims: &[usize]) -> Result { + if dims.is_empty() { + return Ok(1); + } + + dims.iter().try_fold(1usize, |numel, &dim| { + numel.checked_mul(dim).ok_or_else(|| { + MinitensorError::invalid_operation( + "repeat output dimensions exceed supported size".to_string(), + ) + }) + }) +} + +fn attach_repeat_backward(mut output: Tensor, input: &Tensor, repeats: &[usize]) -> Result { + if input.requires_grad() && input.dtype().is_float() { + output.refresh_autograd_metadata(); + let mut output = output.requires_grad_(true); + let grad_fn = Arc::new(RepeatBackward { + input_id: input.id(), + input_shape: input.shape().dims().to_vec(), + repeats: repeats.to_vec(), + }); + output.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output, Some(grad_fn))?; + Ok(output) + } else { + Ok(output) + } +} + +/// Reshape operation with gradient support +pub fn reshape(tensor: &Tensor, new_shape: Shape) -> Result { + // Check if the total number of elements matches + if tensor.numel() != new_shape.numel() { + return Err(MinitensorError::shape_mismatch( + vec![tensor.numel()], + vec![new_shape.numel()], + )); + } + + // Reinterpret the buffer when possible; materialise a contiguous copy for + // non-contiguous inputs (e.g. results of `expand`) so the new shape always + // describes real storage. The copy is made outside of autograd because the + // ReshapeBackward node attached below already routes gradients straight to + // the original tensor. + let mut reshaped = if tensor.is_contiguous() { + tensor.view(new_shape.clone())? + } else { + tensor + .detach() + .contiguous()? + .view(new_shape.clone())? + .requires_grad_(tensor.requires_grad()) + }; + reshaped.refresh_autograd_metadata(); + + // Set up gradient function if needed + if reshaped.requires_grad() { + let grad_fn = Arc::new(ReshapeBackward { + input_shape: tensor.shape().dims().to_vec(), + input_id: tensor.id(), + }); + + reshaped.set_grad_fn(Some(grad_fn.clone())); + + // Add to computation graph + add_to_graph(&reshaped, Some(grad_fn))?; + + Ok(reshaped) + } else { + Ok(reshaped) + } +} + +/// This wrapper performs validation and inference for a single ``-1`` +/// dimension before delegating to [`reshape`]. +pub fn reshape_with_inference(tensor: &Tensor, dims: Vec) -> Result { + let mut out_dims = Vec::with_capacity(dims.len()); + let mut inferred_index: Option = None; + let mut known_product: usize = 1; + + for (index, &dim) in dims.iter().enumerate() { + if dim == -1 { + if inferred_index.is_some() { + return Err(MinitensorError::invalid_operation( + "can only specify one -1 dimension in reshape".to_string(), + )); + } + inferred_index = Some(index); + out_dims.push(0); + continue; + } + + if dim < 0 { + return Err(MinitensorError::invalid_operation( + "invalid negative dimension".to_string(), + )); + } + + let dim_usize = dim as usize; + known_product = known_product.checked_mul(dim_usize).ok_or_else(|| { + MinitensorError::invalid_operation("reshape dimensions exceed supported size") + })?; + out_dims.push(dim_usize); + } + + let total_elements = tensor.numel(); + if let Some(index) = inferred_index { + if known_product == 0 { + return Err(MinitensorError::invalid_operation( + "cannot reshape tensor with -1 and 0 dimensions".to_string(), + )); + } + + if !total_elements.is_multiple_of(known_product) { + return Err(MinitensorError::invalid_operation( + "cannot infer reshape dimension".to_string(), + )); + } + + out_dims[index] = total_elements / known_product; + } else if known_product != total_elements { + return Err(MinitensorError::shape_mismatch( + vec![total_elements], + vec![known_product], + )); + } + + reshape(tensor, Shape::new(out_dims)) +} + +/// Squeeze operation - remove dimensions of size 1. +/// +/// Routed through [`reshape`] so the result is a first-class differentiable node +/// (the plain `Tensor::squeeze` view shares the input's id and would attribute a +/// wrongly-shaped gradient to it). +pub fn squeeze(tensor: &Tensor, dim: Option) -> Result { + let dims = tensor.shape().dims(); + let new_dims: Vec = match dim { + None => dims.iter().copied().filter(|&d| d != 1).collect(), + Some(d) => { + let ndim = tensor.ndim() as isize; + let d = if d < 0 { d + ndim } else { d }; + if d < 0 || d >= ndim { + return Err(MinitensorError::index_error(d, 0, tensor.ndim())); + } + let d = d as usize; + if dims[d] != 1 { + // A non-unit axis remains untouched. + dims.to_vec() + } else { + let mut v = dims.to_vec(); + v.remove(d); + v + } + } + }; + reshape(tensor, Shape::new(new_dims)) +} + +/// Unsqueeze operation - add a dimension of size 1. See [`squeeze`] for why this +/// goes through [`reshape`] rather than the view-based `Tensor::unsqueeze`. +pub fn unsqueeze(tensor: &Tensor, dim: isize) -> Result { + let ndim = tensor.ndim() as isize; + let d = if dim < 0 { dim + ndim + 1 } else { dim }; + if d < 0 || d > ndim { + return Err(MinitensorError::index_error(d, 0, (ndim + 1) as usize)); + } + let mut new_dims = tensor.shape().dims().to_vec(); + new_dims.insert(d as usize, 1); + reshape(tensor, Shape::new(new_dims)) +} + +/// Flatten dimensions `start_dim..=end_dim` into one. Routed through [`reshape`] +/// so gradients flow (see [`squeeze`]). +pub fn flatten(tensor: &Tensor, start_dim: isize, end_dim: isize) -> Result { + let ndim = tensor.ndim() as isize; + let start = if start_dim < 0 { + start_dim + ndim + } else { + start_dim + }; + let end = if end_dim < 0 { end_dim + ndim } else { end_dim }; + if start < 0 || start >= ndim { + return Err(MinitensorError::index_error(start, 0, tensor.ndim())); + } + if end < 0 || end >= ndim { + return Err(MinitensorError::index_error(end, 0, tensor.ndim())); + } + if start > end { + return Err(MinitensorError::invalid_argument( + "start_dim must be less than or equal to end_dim", + )); + } + + let (start, end) = (start as usize, end as usize); + let dims = tensor.shape().dims(); + let mut new_dims = dims[..start].to_vec(); + new_dims.push(dims[start..=end].iter().product()); + new_dims.extend_from_slice(&dims[end + 1..]); + reshape(tensor, Shape::new(new_dims)) +} + +/// Permute tensor dimensions according to `dims` +pub fn permute(tensor: &Tensor, dims: Vec) -> Result { + let ndim = tensor.ndim(); + + // Validate number of dimensions + if dims.len() != ndim { + return Err(MinitensorError::invalid_operation( + "dims must match number of dimensions".to_string(), + )); + } + + // Normalise negative dimensions and validate range + let mut normalized = Vec::with_capacity(ndim); + for &d in &dims { + let d = if d < 0 { d + ndim as isize } else { d }; + if d < 0 || d >= ndim as isize { + return Err(MinitensorError::index_error(d, 0, ndim)); + } + normalized.push(d as usize); + } + // Check that dims form a proper permutation + let mut sorted = normalized.clone(); + sorted.sort_unstable(); + if sorted != (0..ndim).collect::>() { + return Err(MinitensorError::invalid_operation( + "dims must be a permutation of dimensions".to_string(), + )); + } + + // Apply sequence of transposes to achieve the permutation + let mut result = tensor.clone(); + let mut current: Vec = (0..ndim).collect(); + for i in 0..ndim { + let target = normalized[i]; + let j = current.iter().position(|&x| x == target).unwrap(); + if i != j { + result = result.transpose(i as isize, j as isize)?; + current.swap(i, j); + } + } + + Ok(result) +} + +/// Move tensor dimensions to new positions, keeping relative order of other dims +pub fn movedim(tensor: &Tensor, source: &[isize], destination: &[isize]) -> Result { + let ndim = tensor.ndim(); + + if source.len() != destination.len() { + return Err(MinitensorError::invalid_operation( + "movedim: source and destination must have the same length".to_string(), + )); + } + + let mut src_seen = vec![false; ndim]; + let mut dst_seen = vec![false; ndim]; + let mut pairs: Vec<(usize, usize)> = Vec::with_capacity(source.len()); + + for (&s, &d) in source.iter().zip(destination.iter()) { + let s = if s < 0 { s + ndim as isize } else { s }; + if s < 0 || s >= ndim as isize { + return Err(MinitensorError::index_error(s, 0, ndim)); + } + let s = s as usize; + if src_seen[s] { + return Err(MinitensorError::invalid_operation( + "movedim: duplicate dimensions in source".to_string(), + )); + } + src_seen[s] = true; + let d = if d < 0 { d + ndim as isize } else { d }; + if d < 0 || d >= ndim as isize { + return Err(MinitensorError::index_error(d, 0, ndim)); + } + let d = d as usize; + if dst_seen[d] { + return Err(MinitensorError::invalid_operation( + "movedim: duplicate dimensions in destination".to_string(), + )); + } + dst_seen[d] = true; + pairs.push((d, s)); + } + + // Build permutation order + let mut order: Vec = (0..ndim).filter(|&i| !src_seen[i]).collect(); + pairs.sort_by_key(|&(d, _)| d); + for (d, s) in pairs { + order.insert(d, s); + } + let order_isize: Vec = order.into_iter().map(|v| v as isize).collect(); + permute(tensor, order_isize) +} + +/// Concatenate tensors along a specified dimension +pub fn concatenate(tensors: &[&Tensor], dim: isize) -> Result { + if tensors.is_empty() { + return Err(MinitensorError::invalid_operation( + "Cannot concatenate empty list of tensors", + )); + } + + let first_tensor = tensors[0]; + + // Validate that all tensors have the same number of dimensions + for tensor in tensors.iter().skip(1) { + if tensor.ndim() != first_tensor.ndim() { + return Err(MinitensorError::shape_mismatch( + vec![first_tensor.ndim()], + vec![tensor.ndim()], + )); + } + + // Check device compatibility + if tensor.device() != first_tensor.device() { + return Err(MinitensorError::device_mismatch( + format!("{:?}", first_tensor.device()), + format!("{:?}", tensor.device()), + )); + } + + // Check data type compatibility + if tensor.dtype() != first_tensor.dtype() { + return Err(MinitensorError::type_mismatch( + format!("{:?}", first_tensor.dtype()), + format!("{:?}", tensor.dtype()), + )); + } + } + + // Validate concatenation dimension + let dim = normalize_dim(dim, first_tensor.ndim())?; + + // Validate that all dimensions except the concatenation dimension match + for tensor in tensors.iter().skip(1) { + for (i, (&size1, &size2)) in first_tensor + .shape() + .dims() + .iter() + .zip(tensor.shape().dims().iter()) + .enumerate() + { + if i != dim && size1 != size2 { + return Err(MinitensorError::shape_mismatch( + first_tensor.shape().dims().to_vec(), + tensor.shape().dims().to_vec(), + )); + } + } + } + + if !first_tensor.device().is_cpu() { + return Err(MinitensorError::invalid_operation( + "concatenate currently supports only CPU tensors", + )); + } + + // Compute output shape + let mut output_shape = first_tensor.shape().dims().to_vec(); + output_shape[dim] = tensors.iter().map(|t| t.shape().dims()[dim]).sum(); + let output_shape_obj = Shape::new(output_shape); + + let dtype = first_tensor.dtype(); + let device = first_tensor.device(); + let requires_grad = tensors.iter().any(|t| t.requires_grad()); + + let dims = first_tensor.shape().dims(); + let inner: usize = dims[dim + 1..].iter().product(); + let _outer: usize = dims[..dim].iter().product(); + + if output_shape_obj.numel() == 0 { + let data = TensorData::zeros_on_device(0, dtype, device); + return Ok(Tensor::new( + Arc::new(data), + output_shape_obj, + dtype, + device, + requires_grad, + )); + } + + macro_rules! concat_impl { + ($ty:ty, $slice:ident, $from_vec:ident) => {{ + let mut sources: Vec<&[$ty]> = Vec::with_capacity(tensors.len()); + let mut dim_sizes: Vec = Vec::with_capacity(tensors.len()); + for t in tensors { + let src = t.data().$slice().ok_or_else(|| { + MinitensorError::invalid_operation("Tensor data access failed for concatenate") + })?; + sources.push(src); + dim_sizes.push(t.shape().dims()[dim]); + } + let src_strides: Vec = dim_sizes.iter().map(|&d| d * inner).collect(); + + let mut out = vec![<$ty>::default(); output_shape_obj.numel()]; + let chunk_size = output_shape_obj.dims()[dim] * inner; + out.par_chunks_mut(chunk_size) + .enumerate() + .for_each(|(o, out_chunk)| { + let mut dst_offset = 0; + for (src, &src_stride) in sources.iter().zip(src_strides.iter()) { + let src_start = o * src_stride; + let src_len = src_stride; + out_chunk[dst_offset..dst_offset + src_len] + .copy_from_slice(&src[src_start..src_start + src_len]); + dst_offset += src_len; + } + }); + TensorData::$from_vec(out, device) + }}; + } + + let data = match dtype { + DataType::Float32 => concat_impl!(f32, as_f32_slice, from_vec_f32), + DataType::Float64 => concat_impl!(f64, as_f64_slice, from_vec_f64), + DataType::Int32 => concat_impl!(i32, as_i32_slice, from_vec_i32), + DataType::Int64 => concat_impl!(i64, as_i64_slice, from_vec_i64), + DataType::Bool => concat_impl!(bool, as_bool_slice, from_vec_bool), + }; + + let output = Tensor::new( + Arc::new(data), + output_shape_obj, + dtype, + device, + requires_grad, + ); + + if requires_grad && dtype.is_float() { + let grad_fn = Arc::new(ConcatBackward { + input_ids: tensors.iter().map(|t| t.id()).collect(), + sizes: tensors.iter().map(|t| t.shape().dims()[dim]).collect(), + dim, + input_requires_grad: tensors.iter().map(|t| t.requires_grad()).collect(), + }); + let mut output = output; + output.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output, Some(grad_fn))?; + return Ok(output); + } + + Ok(output) +} + +/// Repeat `tensor` according to `repeats` along each dimension. +pub fn repeat(tensor: &Tensor, repeats: &[usize]) -> Result { + if repeats.len() < tensor.ndim() { + return Err(MinitensorError::invalid_operation( + "number of dimensions of repeat dims can not be smaller than number of dimensions of tensor", + )); + } + + // Tile on a detached view so the intermediate per-dimension copies never + // create graph nodes; a single RepeatBackward maps the final result straight + // back to the original input. + let mut result = tensor.detach(); + + if repeats.len() > result.ndim() { + let mut new_shape = vec![1; repeats.len() - result.ndim()]; + new_shape.extend_from_slice(result.shape().dims()); + result = result.reshape(Shape::new(new_shape))?; + } + + if repeats.contains(&0) { + let mut out_shape = result.shape().dims().to_vec(); + for (dim, &rep) in repeats.iter().enumerate() { + out_shape[dim] = checked_repeat_dim(out_shape[dim], rep)?; + } + + let output = empty_tensor( + Shape::new(out_shape), + result.dtype(), + result.device(), + false, + ); + return attach_repeat_backward(output, tensor, repeats); + } + + for (dim, &rep) in repeats.iter().enumerate() { + if rep == 1 { + continue; + } + let dims = result.shape().dims().to_vec(); + let dim_size = dims[dim]; + let inner: usize = dims[dim + 1..].iter().product(); + let repeated_dim = checked_repeat_dim(dim_size, rep)?; + let chunk_size = checked_repeat_dim(repeated_dim, inner)?; + let src_chunk_size = checked_repeat_dim(dim_size, inner)?; + let mut output_shape = dims.clone(); + output_shape[dim] = repeated_dim; + let output_numel = checked_repeat_numel(&output_shape)?; + let output_shape_obj = Shape::new(output_shape); + + let dtype = result.dtype(); + let device = result.device(); + let requires_grad = result.requires_grad(); + + if output_numel == 0 { + result = empty_tensor(output_shape_obj, dtype, device, requires_grad); + continue; + } + + macro_rules! repeat_impl { + ($ty:ty, $slice:ident, $from_vec:ident) => {{ + let src = result.data().$slice().ok_or_else(|| { + MinitensorError::invalid_operation("Tensor data access failed for repeat") + })?; + let mut out = vec![<$ty>::default(); output_numel]; + out.par_chunks_mut(chunk_size) + .enumerate() + .for_each(|(o, out_chunk)| { + let src_start = o * src_chunk_size; + let src_chunk = &src[src_start..src_start + src_chunk_size]; + for r in 0..rep { + let dst_start = r * src_chunk_size; + out_chunk[dst_start..dst_start + src_chunk_size] + .copy_from_slice(src_chunk); + } + }); + TensorData::$from_vec(out, device) + }}; + } + + let data = match dtype { + DataType::Float32 => repeat_impl!(f32, as_f32_slice, from_vec_f32), + DataType::Float64 => repeat_impl!(f64, as_f64_slice, from_vec_f64), + DataType::Int32 => repeat_impl!(i32, as_i32_slice, from_vec_i32), + DataType::Int64 => repeat_impl!(i64, as_i64_slice, from_vec_i64), + DataType::Bool => repeat_impl!(bool, as_bool_slice, from_vec_bool), + }; + + result = Tensor::new( + Arc::new(data), + output_shape_obj, + dtype, + device, + requires_grad, + ); + } + + attach_repeat_backward(result, tensor, repeats) +} + +/// Indexing operation - select elements along specified dimensions +pub fn index_select(tensor: &Tensor, dim: isize, indices: &[usize]) -> Result { + let dim = normalize_dim(dim, tensor.ndim())?; + + let dim_size = tensor.shape().dims()[dim]; + + // Validate indices + for &idx in indices { + if idx >= dim_size { + return Err(MinitensorError::index_error(idx as isize, 0, dim_size)); + } + } + + if !tensor.device().is_cpu() { + return Err(MinitensorError::invalid_operation( + "index_select currently supports only CPU tensors", + )); + } + + // Compute output shape + let mut output_shape = tensor.shape().dims().to_vec(); + output_shape[dim] = indices.len(); + let output_shape_vec = output_shape.clone(); + let output_shape_obj = Shape::new(output_shape); + + let dtype = tensor.dtype(); + let device = tensor.device(); + let requires_grad = tensor.requires_grad(); + + if output_shape_obj.numel() == 0 { + return Ok(empty_tensor(output_shape_obj, dtype, device, requires_grad)); + } + + let dims = tensor.shape().dims(); + let inner: usize = dims[dim + 1..].iter().product(); + let _outer: usize = dims[..dim].iter().product(); + + macro_rules! index_impl { + ($ty:ty, $slice:ident, $from_vec:ident) => {{ + let src = tensor.data().$slice().ok_or_else(|| { + MinitensorError::invalid_operation("Tensor data access failed for index_select") + })?; + let mut out = vec![<$ty>::default(); output_shape_obj.numel()]; + out.par_chunks_mut(output_shape_vec[dim] * inner) + .enumerate() + .for_each(|(o, out_chunk)| { + for (i, &idx) in indices.iter().enumerate() { + let src_start = o * dims[dim] * inner + idx * inner; + let dst_start = i * inner; + out_chunk[dst_start..dst_start + inner] + .copy_from_slice(&src[src_start..src_start + inner]); + } + }); + TensorData::$from_vec(out, device) + }}; + } + + let data = match dtype { + DataType::Float32 => index_impl!(f32, as_f32_slice, from_vec_f32), + DataType::Float64 => index_impl!(f64, as_f64_slice, from_vec_f64), + DataType::Int32 => index_impl!(i32, as_i32_slice, from_vec_i32), + DataType::Int64 => index_impl!(i64, as_i64_slice, from_vec_i64), + DataType::Bool => index_impl!(bool, as_bool_slice, from_vec_bool), + }; + + let output = Tensor::new( + Arc::new(data), + output_shape_obj, + dtype, + device, + requires_grad, + ); + + if requires_grad && dtype.is_float() { + let grad_fn = Arc::new(IndexSelectBackward { + input_id: tensor.id(), + input_shape: tensor.shape().dims().to_vec(), + dim, + indices: indices.to_vec(), + }); + let mut output = output; + output.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output, Some(grad_fn))?; + return Ok(output); + } + + Ok(output) +} + +/// Gather operation - collect elements along a dimension using an index tensor +pub fn gather(tensor: &Tensor, dim: isize, index: &Tensor) -> Result { + let dim = normalize_dim(dim, tensor.ndim())?; + + if index.ndim() != tensor.ndim() { + return Err(MinitensorError::invalid_operation( + "gather index tensor must have the same number of dimensions as input", + )); + } + + if index.dtype() != DataType::Int64 { + return Err(MinitensorError::invalid_operation( + "gather indices must be int64", + )); + } + + let input_dims = tensor.shape().dims(); + let index_dims = index.shape().dims(); + + // Validate shapes except at gather dimension + for (i, (&idx_d, &in_d)) in index_dims.iter().zip(input_dims.iter()).enumerate() { + if i != dim && idx_d != in_d { + return Err(MinitensorError::shape_mismatch( + input_dims.to_vec(), + index_dims.to_vec(), + )); + } + } + + let dim_size = input_dims[dim]; + + // Validate indices + let idx_slice = index + .data() + .as_i64_slice() + .ok_or_else(|| MinitensorError::invalid_operation("gather indices must be int64"))?; + for &v in idx_slice { + if v < 0 || v as usize >= dim_size { + return Err(MinitensorError::index_error(v as isize, 0, dim_size)); + } + } + + if !tensor.device().is_cpu() { + return Err(MinitensorError::invalid_operation( + "gather currently supports only CPU tensors", + )); + } + + let inner: usize = input_dims[dim + 1..].iter().product(); + let idx_dim = index_dims[dim]; + + let dtype = tensor.dtype(); + let device = tensor.device(); + let requires_grad = tensor.requires_grad(); + let output_shape_obj = Shape::new(index_dims.to_vec()); + let output_numel = idx_slice.len(); + + if output_numel == 0 { + return Ok(empty_tensor(output_shape_obj, dtype, device, requires_grad)); + } + + macro_rules! gather_impl { + ($ty:ty, $slice:ident, $from_vec:ident) => {{ + let src = tensor.data().$slice().ok_or_else(|| { + MinitensorError::invalid_operation("Tensor data access failed for gather") + })?; + let idx = idx_slice; + let mut out = vec![<$ty>::default(); output_numel]; + let chunk_size = idx_dim * inner; + if output_numel % chunk_size != 0 { + return Err(MinitensorError::internal_error(format!( + "gather output length ({output_numel}) is not divisible by chunk size ({chunk_size})" + ))); + } + out.par_chunks_mut(chunk_size) + .enumerate() + .for_each(|(o, out_chunk)| { + let base = o * dim_size * inner; + let idx_chunk = &idx[o * chunk_size..(o + 1) * chunk_size]; + for i in 0..idx_dim { + let idx_row = &idx_chunk[i * inner..(i + 1) * inner]; + let dst_row = &mut out_chunk[i * inner..(i + 1) * inner]; + for (j, &gather_val) in idx_row.iter().enumerate() { + let gather_idx = gather_val as usize; + dst_row[j] = src[base + gather_idx * inner + j]; + } + } + }); + TensorData::$from_vec(out, device) + }}; + } + + let data = match dtype { + DataType::Float32 => gather_impl!(f32, as_f32_slice, from_vec_f32), + DataType::Float64 => gather_impl!(f64, as_f64_slice, from_vec_f64), + DataType::Int32 => gather_impl!(i32, as_i32_slice, from_vec_i32), + DataType::Int64 => gather_impl!(i64, as_i64_slice, from_vec_i64), + DataType::Bool => gather_impl!(bool, as_bool_slice, from_vec_bool), + }; + + let output = Tensor::new( + Arc::new(data), + output_shape_obj, + dtype, + device, + requires_grad, + ); + + if requires_grad && dtype.is_float() { + let grad_fn = Arc::new(GatherBackward { + input_id: tensor.id(), + input_shape: input_dims.to_vec(), + dim, + index: idx_slice.to_vec(), + }); + let mut output = output; + output.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output, Some(grad_fn))?; + return Ok(output); + } + + Ok(output) +} + +/// Slicing operation - select a contiguous range of elements +pub fn slice(tensor: &Tensor, dim: isize, start: usize, end: usize, step: usize) -> Result { + let dim = normalize_dim(dim, tensor.ndim())?; + + let dim_size = tensor.shape().dims()[dim]; + + if start > dim_size || end > dim_size || start > end { + return Err(MinitensorError::invalid_operation(format!( + "Invalid slice range: start={}, end={}, dim_size={}", + start, end, dim_size + ))); + } + + if step == 0 { + return Err(MinitensorError::invalid_operation( + "Slice step cannot be zero", + )); + } + + if !tensor.device().is_cpu() { + return Err(MinitensorError::invalid_operation( + "slice currently supports only CPU tensors", + )); + } + + // Compute output shape + let mut output_shape = tensor.shape().dims().to_vec(); + output_shape[dim] = (end - start).div_ceil(step); + let output_shape_obj = Shape::new(output_shape); + + let dtype = tensor.dtype(); + let device = tensor.device(); + let requires_grad = tensor.requires_grad(); + + if output_shape_obj.numel() == 0 { + return Ok(empty_tensor(output_shape_obj, dtype, device, requires_grad)); + } + + let dims = tensor.shape().dims(); + let inner: usize = dims[dim + 1..].iter().product(); + let count = output_shape_obj.dims()[dim]; + + macro_rules! slice_impl { + ($ty:ty, $slice:ident, $from_vec:ident) => {{ + let src = tensor.data().$slice().ok_or_else(|| { + MinitensorError::invalid_operation("Tensor data access failed for slice") + })?; + let mut out = vec![<$ty>::default(); output_shape_obj.numel()]; + out.par_chunks_mut(count * inner) + .enumerate() + .for_each(|(o, out_chunk)| { + for i in 0..count { + let src_idx = start + i * step; + let src_start = o * dims[dim] * inner + src_idx * inner; + let dst_start = i * inner; + out_chunk[dst_start..dst_start + inner] + .copy_from_slice(&src[src_start..src_start + inner]); + } + }); + TensorData::$from_vec(out, device) + }}; + } + + let data = match dtype { + DataType::Float32 => slice_impl!(f32, as_f32_slice, from_vec_f32), + DataType::Float64 => slice_impl!(f64, as_f64_slice, from_vec_f64), + DataType::Int32 => slice_impl!(i32, as_i32_slice, from_vec_i32), + DataType::Int64 => slice_impl!(i64, as_i64_slice, from_vec_i64), + DataType::Bool => slice_impl!(bool, as_bool_slice, from_vec_bool), + }; + + let output = Tensor::new( + Arc::new(data), + output_shape_obj, + dtype, + device, + requires_grad, + ); + + if requires_grad && dtype.is_float() { + // A slice selects source positions `start, start+step, ...` along `dim`; + // its backward scatters the gradient back to exactly those positions. + let indices: Vec = (0..count).map(|i| start + i * step).collect(); + let grad_fn = Arc::new(IndexSelectBackward { + input_id: tensor.id(), + input_shape: tensor.shape().dims().to_vec(), + dim, + indices, + }); + let mut output = output; + output.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output, Some(grad_fn))?; + return Ok(output); + } + + Ok(output) +} + +/// Narrow tensor along a dimension starting at `start` for `length` elements +pub fn narrow(tensor: &Tensor, dim: isize, start: usize, length: usize) -> Result { + let dim = normalize_dim(dim, tensor.ndim())?; + let dim_size = tensor.shape().dims()[dim]; + + if start > dim_size { + return Err(MinitensorError::index_error(start as isize, 0, dim_size)); + } + if start + length > dim_size { + return Err(MinitensorError::index_error( + (start + length) as isize, + 0, + dim_size, + )); + } + + if length == 0 { + let mut out_shape = tensor.shape().dims().to_vec(); + out_shape[dim] = 0; + return Ok(Tensor::zeros( + Shape::new(out_shape), + tensor.dtype(), + tensor.device(), + tensor.requires_grad(), + )); + } + + slice(tensor, dim as isize, start, start + length, 1) +} + +/// Flip tensor elements along specified dimensions. +pub fn flip(tensor: &Tensor, dims: &[isize]) -> Result { + let ndim = tensor.ndim(); + let mut normalized = Vec::with_capacity(dims.len()); + for &d in dims { + let dim = normalize_dim(d, ndim)?; + if normalized.contains(&dim) { + return Err(MinitensorError::invalid_operation( + "dims must be unique".to_string(), + )); + } + normalized.push(dim); + } + + let mut result = tensor.clone(); + for &dim in &normalized { + let size = result.shape().dims()[dim]; + let indices: Vec = (0..size).rev().collect(); + result = index_select(&result, dim as isize, &indices)?; + } + + Ok(result) +} + +/// Roll tensor elements along specified dimensions with wrap-around +pub fn roll(tensor: &Tensor, shifts: &[isize], dims: Option<&[isize]>) -> Result { + // Compute the roll on a detached view so the internal slice/concatenate steps + // (which flatten to a storage-sharing view in the `dims == None` case) never + // build gradient edges; a single RollBackward inverts the whole operation. + let track_grad = tensor.requires_grad() && tensor.dtype().is_float(); + let base = if track_grad { + tensor.detach() + } else { + tensor.clone() + }; + let output = roll_forward(&base, shifts, dims)?; + + if track_grad { + let mut output = output; + output.refresh_autograd_metadata(); + let mut output = output.requires_grad_(true); + let grad_fn = Arc::new(crate::autograd::RollBackward { + input_id: tensor.id(), + shifts: shifts.to_vec(), + dims: dims.map(|d| d.to_vec()), + }); + output.set_grad_fn(Some(grad_fn.clone())); + add_to_graph(&output, Some(grad_fn))?; + return Ok(output); + } + + Ok(output) +} + +fn roll_forward(tensor: &Tensor, shifts: &[isize], dims: Option<&[isize]>) -> Result { + if let Some(dims) = dims { + if shifts.len() != dims.len() { + return Err(MinitensorError::invalid_operation( + "shifts and dims must have the same length".to_string(), + )); + } + let mut result = tensor.clone(); + for (&shift, &dim) in shifts.iter().zip(dims.iter()) { + let dim = normalize_dim(dim, result.ndim())?; + let size = result.shape().dims()[dim] as isize; + if size == 0 { + continue; + } + let k = ((shift % size) + size) % size; + if k == 0 { + continue; + } + let split_point = (size - k) as usize; + let first = slice(&result, dim as isize, 0, split_point, 1)?; + let second = slice(&result, dim as isize, split_point, size as usize, 1)?; + result = concatenate(&[&second, &first], dim as isize)?; + } + Ok(result) + } else { + if shifts.len() != 1 { + return Err(MinitensorError::invalid_operation( + "shifts must contain a single value when dims is None".to_string(), + )); + } + let shift = shifts[0]; + let flat = tensor.flatten_all()?; + let size = flat.shape().dims()[0] as isize; + if size == 0 { + return flat.reshape(tensor.shape().clone()); + } + let k = ((shift % size) + size) % size; + if k == 0 { + return flat.reshape(tensor.shape().clone()); + } + let split_point = (size - k) as usize; + let first = slice(&flat, 0, 0, split_point, 1)?; + let second = slice(&flat, 0, split_point, size as usize, 1)?; + let rolled = concatenate(&[&second, &first], 0)?; + rolled.reshape(tensor.shape().clone()) + } +} + +/// Specification of repeat counts accepted by [`repeat_interleave`]. +#[derive(Clone, Copy)] +pub enum RepeatInterleaveSpec<'a> { + /// A single repeat value applied to every element along ``dim``. + Scalar(usize), + /// Explicit repeat counts provided as a slice. + Slice(&'a [usize]), + /// Repeat counts provided as a tensor (must contain integer data). + Tensor(&'a Tensor), +} + +fn collect_repeats_from_values(len: usize, values: I) -> Result> +where + I: IntoIterator, +{ + let mut out = Vec::with_capacity(len); + for value in values { + if value < 0 { + return Err(MinitensorError::invalid_operation( + "repeat_interleave: repeats must be non-negative".to_string(), + )); + } + out.push(value as usize); + } + Ok(out) +} + +pub(crate) fn collect_repeats_from_tensor(tensor: &Tensor, dim_size: usize) -> Result> { + if !tensor.device().is_cpu() { + return Err(MinitensorError::invalid_operation( + "repeat_interleave: repeats tensor must reside on CPU".to_string(), + )); + } + + if tensor.numel() != dim_size { + return Err(MinitensorError::invalid_operation( + "repeat_interleave: repeats tensor must have the same number of elements as the selected dimension" + .to_string(), + )); + } + + match tensor.dtype() { + DataType::Int32 => { + let slice = tensor.data().as_i32_slice().ok_or_else(|| { + MinitensorError::invalid_operation( + "repeat_interleave: repeats tensor must be contiguous".to_string(), + ) + })?; + collect_repeats_from_values(slice.len(), slice.iter().map(|&value| value as i64)) + } + DataType::Int64 => { + let slice = tensor.data().as_i64_slice().ok_or_else(|| { + MinitensorError::invalid_operation( + "repeat_interleave: repeats tensor must be contiguous".to_string(), + ) + })?; + collect_repeats_from_values(slice.len(), slice.iter().copied()) + } + other => Err(MinitensorError::type_mismatch( + "integral tensor", + format!("{:?}", other), + )), + } +} + +#[cfg(test)] +mod reshape_tests { + use super::*; + + #[test] + fn reshape_with_inference_rejects_overflowing_known_product() { + let tensor = Tensor::zeros(Shape::new(vec![1]), DataType::Float32, Device::cpu(), false); + + let result = reshape_with_inference(&tensor, vec![isize::MAX, isize::MAX, -1]); + + assert!(result.is_err()); + assert!( + result + .unwrap_err() + .to_string() + .contains("reshape dimensions exceed supported size") + ); + } + + #[test] + fn reshape_with_inference_rejects_multiple_inferred_dimensions() { + let tensor = Tensor::zeros(Shape::new(vec![4]), DataType::Float32, Device::cpu(), false); + + let result = reshape_with_inference(&tensor, vec![-1, -1]); + + assert!(result.is_err()); + assert!( + result + .unwrap_err() + .to_string() + .contains("can only specify one -1 dimension in reshape") + ); + } + + #[test] + fn reshape_with_inference_rejects_invalid_negative_dimension() { + let tensor = Tensor::zeros(Shape::new(vec![1]), DataType::Float32, Device::cpu(), false); + + let result = reshape_with_inference(&tensor, vec![-2, 1]); + + assert!(result.is_err()); + assert!( + result + .unwrap_err() + .to_string() + .contains("invalid negative dimension") + ); + } + + #[test] + fn reshape_with_inference_rejects_overflowing_shape_without_inference() { + let tensor = Tensor::zeros(Shape::new(vec![1]), DataType::Float32, Device::cpu(), false); + + let result = reshape_with_inference(&tensor, vec![isize::MAX, isize::MAX]); + + assert!(result.is_err()); + assert!( + result + .unwrap_err() + .to_string() + .contains("reshape dimensions exceed supported size") + ); + } + + #[test] + fn reshape_with_inference_rejects_zero_dimension_with_inference() { + let tensor = Tensor::zeros(Shape::new(vec![0]), DataType::Float32, Device::cpu(), false); + + let result = reshape_with_inference(&tensor, vec![-1, 0]); + + assert!(result.is_err()); + assert!( + result + .unwrap_err() + .to_string() + .contains("cannot reshape tensor with -1 and 0 dimensions") + ); + } + + #[test] + fn reshape_with_inference_no_inference_shape_mismatch() { + let tensor = Tensor::zeros(Shape::new(vec![5]), DataType::Float32, Device::cpu(), false); + + let result = reshape_with_inference(&tensor, vec![2, 2]); + + assert!(result.is_err()); + assert!(result.unwrap_err().to_string().contains("Shape mismatch")); + } + + #[test] + fn reshape_with_inference_no_inference_shape_match() { + let tensor = Tensor::zeros(Shape::new(vec![6]), DataType::Float32, Device::cpu(), false); + + let result = reshape_with_inference(&tensor, vec![2, 3]); + + assert!(result.is_ok()); + assert_eq!( + result.expect("reshape should succeed").shape().dims(), + &[2, 3] + ); + } + + #[test] + fn reshape_with_inference_infers_single_negative_dimension() { + let tensor = Tensor::zeros( + Shape::new(vec![12]), + DataType::Float32, + Device::cpu(), + false, + ); + + let reshaped = reshape_with_inference(&tensor, vec![3, -1]).expect("reshape should work"); + + assert_eq!(reshaped.shape().dims(), &[3, 4]); + } +} diff --git a/engine/src/operations/simd.rs b/engine/src/operations/simd.rs index ccb128d7..90a61446 100644 --- a/engine/src/operations/simd.rs +++ b/engine/src/operations/simd.rs @@ -4,5 +4,10 @@ // This source code is licensed under the Apache-style license found in the // LICENSE file in the root directory of this source tree. -include!("simd/kernels.rs"); -include!("simd/utils.rs"); +#[path = "simd/kernels.rs"] +mod kernels_impl; +#[path = "simd/utils.rs"] +mod utils_impl; + +pub use self::kernels_impl::*; +pub(crate) use self::utils_impl::*; diff --git a/engine/src/operations/simd/kernels.rs b/engine/src/operations/simd/kernels.rs index 1e020b2d..4e590baf 100644 --- a/engine/src/operations/simd/kernels.rs +++ b/engine/src/operations/simd/kernels.rs @@ -1,987 +1,989 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -use crate::{ - error::{MinitensorError, Result}, - tensor::Shape, -}; - -#[cfg(target_arch = "x86_64")] -use std::arch::x86_64::*; - -#[cfg(target_arch = "aarch64")] -use std::arch::aarch64::*; - -/// SIMD capabilities detected at runtime -#[derive(Debug, Clone, Copy)] -pub struct SimdCapabilities { - pub avx2: bool, - pub avx512: bool, - pub sse4_1: bool, - pub neon: bool, - pub sve: bool, -} - -impl SimdCapabilities { - /// Detect SIMD capabilities at runtime - pub fn detect() -> Self { - Self { - #[cfg(target_arch = "x86_64")] - avx2: is_x86_feature_detected!("avx2"), - #[cfg(target_arch = "x86_64")] - avx512: is_x86_feature_detected!("avx512f"), - #[cfg(target_arch = "x86_64")] - sse4_1: is_x86_feature_detected!("sse4.1"), - #[cfg(not(target_arch = "x86_64"))] - avx2: false, - #[cfg(not(target_arch = "x86_64"))] - avx512: false, - #[cfg(not(target_arch = "x86_64"))] - sse4_1: false, - - #[cfg(target_arch = "aarch64")] - neon: std::arch::is_aarch64_feature_detected!("neon"), - #[cfg(target_arch = "aarch64")] - sve: std::arch::is_aarch64_feature_detected!("sve"), - #[cfg(not(target_arch = "aarch64"))] - neon: false, - #[cfg(not(target_arch = "aarch64"))] - sve: false, - } - } -} - -/// Global SIMD capabilities (detected once at startup) -static SIMD_CAPS: std::sync::OnceLock = std::sync::OnceLock::new(); - -/// Get the detected SIMD capabilities -pub fn simd_capabilities() -> SimdCapabilities { - *SIMD_CAPS.get_or_init(SimdCapabilities::detect) -} - -/// SIMD-optimized element-wise addition for f32 arrays -pub fn simd_add_f32(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { - if lhs.len() != rhs.len() || lhs.len() != output.len() { - return Err(MinitensorError::invalid_operation( - "Array lengths must match for SIMD operations", - )); - } - - let caps = simd_capabilities(); - - #[cfg(target_arch = "x86_64")] - { - if caps.avx2 { - return unsafe { simd_add_f32_avx2(lhs, rhs, output) }; - } else if caps.sse4_1 { - return unsafe { simd_add_f32_sse(lhs, rhs, output) }; - } - } - - #[cfg(target_arch = "aarch64")] - { - if caps.neon { - return unsafe { simd_add_f32_neon(lhs, rhs, output) }; - } - } - - // Fallback to scalar implementation - simd_add_f32_scalar(lhs, rhs, output) -} - -/// SIMD-optimized element-wise subtraction for f32 arrays -pub fn simd_sub_f32(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { - if lhs.len() != rhs.len() || lhs.len() != output.len() { - return Err(MinitensorError::invalid_operation( - "Array lengths must match for SIMD operations", - )); - } - - let caps = simd_capabilities(); - - #[cfg(target_arch = "x86_64")] - { - if caps.avx2 { - return unsafe { simd_sub_f32_avx2(lhs, rhs, output) }; - } else if caps.sse4_1 { - return unsafe { simd_sub_f32_sse(lhs, rhs, output) }; - } - } - - #[cfg(target_arch = "aarch64")] - { - if caps.neon { - return unsafe { simd_sub_f32_neon(lhs, rhs, output) }; - } - } - - // Fallback to scalar implementation - simd_sub_f32_scalar(lhs, rhs, output) -} - -/// SIMD-optimized element-wise multiplication for f32 arrays -pub fn simd_mul_f32(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { - if lhs.len() != rhs.len() || lhs.len() != output.len() { - return Err(MinitensorError::invalid_operation( - "Array lengths must match for SIMD operations", - )); - } - - let caps = simd_capabilities(); - - #[cfg(target_arch = "x86_64")] - { - if caps.avx2 { - return unsafe { simd_mul_f32_avx2(lhs, rhs, output) }; - } else if caps.sse4_1 { - return unsafe { simd_mul_f32_sse(lhs, rhs, output) }; - } - } - - #[cfg(target_arch = "aarch64")] - { - if caps.neon { - return unsafe { simd_mul_f32_neon(lhs, rhs, output) }; - } - } - - // Fallback to scalar implementation - simd_mul_f32_scalar(lhs, rhs, output) -} - -/// SIMD-optimized element-wise division for f32 arrays -pub fn simd_div_f32(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { - if lhs.len() != rhs.len() || lhs.len() != output.len() { - return Err(MinitensorError::invalid_operation( - "Array lengths must match for SIMD operations", - )); - } - - let caps = simd_capabilities(); - - #[cfg(target_arch = "x86_64")] - { - if caps.avx2 { - return unsafe { simd_div_f32_avx2(lhs, rhs, output) }; - } else if caps.sse4_1 { - return unsafe { simd_div_f32_sse(lhs, rhs, output) }; - } - } - - #[cfg(target_arch = "aarch64")] - { - if caps.neon { - return unsafe { simd_div_f32_neon(lhs, rhs, output) }; - } - } - - // Fallback to scalar implementation - simd_div_f32_scalar(lhs, rhs, output) -} - -// Scalar fallback implementations -fn simd_add_f32_scalar(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { - for i in 0..lhs.len() { - output[i] = lhs[i] + rhs[i]; - } - Ok(()) -} - -/// Unrolled sum for f32 slices to leverage auto-vectorization -pub fn simd_sum_f32(data: &[f32]) -> f32 { - let mut sums = [0f32; 8]; - let chunks = data.chunks_exact(8); - let rem = chunks.remainder(); - for chunk in chunks { - sums[0] += chunk[0]; - sums[1] += chunk[1]; - sums[2] += chunk[2]; - sums[3] += chunk[3]; - sums[4] += chunk[4]; - sums[5] += chunk[5]; - sums[6] += chunk[6]; - sums[7] += chunk[7]; - } - let mut total: f32 = sums.iter().sum(); - total += rem.iter().copied().sum::(); - total -} - -/// Unrolled sum for f64 slices to leverage auto-vectorization -pub fn simd_sum_f64(data: &[f64]) -> f64 { - let mut sums = [0f64; 4]; - let chunks = data.chunks_exact(4); - let rem = chunks.remainder(); - for chunk in chunks { - sums[0] += chunk[0]; - sums[1] += chunk[1]; - sums[2] += chunk[2]; - sums[3] += chunk[3]; - } - let mut total: f64 = sums.iter().sum(); - total += rem.iter().copied().sum::(); - total -} - -/// Unrolled sum for i32 slices to leverage auto-vectorization -pub fn simd_sum_i32(data: &[i32]) -> i32 { - let mut sums = [0i32; 8]; - let chunks = data.chunks_exact(8); - let rem = chunks.remainder(); - for chunk in chunks { - sums[0] += chunk[0]; - sums[1] += chunk[1]; - sums[2] += chunk[2]; - sums[3] += chunk[3]; - sums[4] += chunk[4]; - sums[5] += chunk[5]; - sums[6] += chunk[6]; - sums[7] += chunk[7]; - } - let mut total: i32 = sums.iter().sum(); - total += rem.iter().copied().sum::(); - total -} - -/// Unrolled sum for i64 slices to leverage auto-vectorization -pub fn simd_sum_i64(data: &[i64]) -> i64 { - let mut sums = [0i64; 4]; - let chunks = data.chunks_exact(4); - let rem = chunks.remainder(); - for chunk in chunks { - sums[0] += chunk[0]; - sums[1] += chunk[1]; - sums[2] += chunk[2]; - sums[3] += chunk[3]; - } - let mut total: i64 = sums.iter().sum(); - total += rem.iter().copied().sum::(); - total -} - -/// Unrolled product for f32 slices to leverage auto-vectorization -pub fn simd_prod_f32(data: &[f32]) -> f32 { - let mut prods = [1f32; 8]; - let chunks = data.chunks_exact(8); - let rem = chunks.remainder(); - for chunk in chunks { - prods[0] *= chunk[0]; - prods[1] *= chunk[1]; - prods[2] *= chunk[2]; - prods[3] *= chunk[3]; - prods[4] *= chunk[4]; - prods[5] *= chunk[5]; - prods[6] *= chunk[6]; - prods[7] *= chunk[7]; - } - let mut total: f32 = prods.iter().product(); - total *= rem.iter().copied().product::(); - total -} - -/// Unrolled product for f64 slices to leverage auto-vectorization -pub fn simd_prod_f64(data: &[f64]) -> f64 { - let mut prods = [1f64; 4]; - let chunks = data.chunks_exact(4); - let rem = chunks.remainder(); - for chunk in chunks { - prods[0] *= chunk[0]; - prods[1] *= chunk[1]; - prods[2] *= chunk[2]; - prods[3] *= chunk[3]; - } - let mut total: f64 = prods.iter().product(); - total *= rem.iter().copied().product::(); - total -} - -/// Unrolled product for i32 slices to leverage auto-vectorization -pub fn simd_prod_i32(data: &[i32]) -> i32 { - let mut prods = [1i32; 8]; - let chunks = data.chunks_exact(8); - let rem = chunks.remainder(); - for chunk in chunks { - prods[0] *= chunk[0]; - prods[1] *= chunk[1]; - prods[2] *= chunk[2]; - prods[3] *= chunk[3]; - prods[4] *= chunk[4]; - prods[5] *= chunk[5]; - prods[6] *= chunk[6]; - prods[7] *= chunk[7]; - } - let mut total: i32 = prods.iter().product(); - total *= rem.iter().copied().product::(); - total -} - -/// Unrolled product for i64 slices to leverage auto-vectorization -pub fn simd_prod_i64(data: &[i64]) -> i64 { - let mut prods = [1i64; 4]; - let chunks = data.chunks_exact(4); - let rem = chunks.remainder(); - for chunk in chunks { - prods[0] *= chunk[0]; - prods[1] *= chunk[1]; - prods[2] *= chunk[2]; - prods[3] *= chunk[3]; - } - let mut total: i64 = prods.iter().product(); - total *= rem.iter().copied().product::(); - total -} - -fn simd_sub_f32_scalar(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { - for i in 0..lhs.len() { - output[i] = lhs[i] - rhs[i]; - } - Ok(()) -} - -fn simd_mul_f32_scalar(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { - for i in 0..lhs.len() { - output[i] = lhs[i] * rhs[i]; - } - Ok(()) -} - -fn simd_div_f32_scalar(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { - for i in 0..lhs.len() { - output[i] = lhs[i] / rhs[i]; - } - Ok(()) -} - -// x86_64 AVX2 implementations -#[cfg(target_arch = "x86_64")] -#[target_feature(enable = "avx2")] -unsafe fn simd_add_f32_avx2(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { - const SIMD_WIDTH: usize = 8; // AVX2 processes 8 f32s at once - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - // Process SIMD_WIDTH elements at a time - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = _mm256_loadu_ps(lhs.as_ptr().add(i)); - let b = _mm256_loadu_ps(rhs.as_ptr().add(i)); - let result = _mm256_add_ps(a, b); - _mm256_storeu_ps(output.as_mut_ptr().add(i), result); - } - } - - // Handle remaining elements - for i in simd_len..len { - output[i] = lhs[i] + rhs[i]; - } - - Ok(()) -} - -#[cfg(target_arch = "x86_64")] -#[target_feature(enable = "avx2")] -unsafe fn simd_sub_f32_avx2(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { - const SIMD_WIDTH: usize = 8; - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = _mm256_loadu_ps(lhs.as_ptr().add(i)); - let b = _mm256_loadu_ps(rhs.as_ptr().add(i)); - let result = _mm256_sub_ps(a, b); - _mm256_storeu_ps(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] - rhs[i]; - } - - Ok(()) -} - -#[cfg(target_arch = "x86_64")] -#[target_feature(enable = "avx2")] -unsafe fn simd_mul_f32_avx2(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { - const SIMD_WIDTH: usize = 8; - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = _mm256_loadu_ps(lhs.as_ptr().add(i)); - let b = _mm256_loadu_ps(rhs.as_ptr().add(i)); - let result = _mm256_mul_ps(a, b); - _mm256_storeu_ps(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] * rhs[i]; - } - - Ok(()) -} - -#[cfg(target_arch = "x86_64")] -#[target_feature(enable = "avx2")] -unsafe fn simd_div_f32_avx2(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { - const SIMD_WIDTH: usize = 8; - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = _mm256_loadu_ps(lhs.as_ptr().add(i)); - let b = _mm256_loadu_ps(rhs.as_ptr().add(i)); - let result = _mm256_div_ps(a, b); - _mm256_storeu_ps(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] / rhs[i]; - } - - Ok(()) -} - -// x86_64 SSE implementations -#[cfg(target_arch = "x86_64")] -#[target_feature(enable = "sse4.1")] -unsafe fn simd_add_f32_sse(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { - const SIMD_WIDTH: usize = 4; // SSE processes 4 f32s at once - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = _mm_loadu_ps(lhs.as_ptr().add(i)); - let b = _mm_loadu_ps(rhs.as_ptr().add(i)); - let result = _mm_add_ps(a, b); - _mm_storeu_ps(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] + rhs[i]; - } - - Ok(()) -} - -#[cfg(target_arch = "x86_64")] -#[target_feature(enable = "sse4.1")] -unsafe fn simd_sub_f32_sse(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { - const SIMD_WIDTH: usize = 4; - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = _mm_loadu_ps(lhs.as_ptr().add(i)); - let b = _mm_loadu_ps(rhs.as_ptr().add(i)); - let result = _mm_sub_ps(a, b); - _mm_storeu_ps(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] - rhs[i]; - } - - Ok(()) -} - -#[cfg(target_arch = "x86_64")] -#[target_feature(enable = "sse4.1")] -unsafe fn simd_mul_f32_sse(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { - const SIMD_WIDTH: usize = 4; - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = _mm_loadu_ps(lhs.as_ptr().add(i)); - let b = _mm_loadu_ps(rhs.as_ptr().add(i)); - let result = _mm_mul_ps(a, b); - _mm_storeu_ps(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] * rhs[i]; - } - - Ok(()) -} - -#[cfg(target_arch = "x86_64")] -#[target_feature(enable = "sse4.1")] -unsafe fn simd_div_f32_sse(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { - const SIMD_WIDTH: usize = 4; - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = _mm_loadu_ps(lhs.as_ptr().add(i)); - let b = _mm_loadu_ps(rhs.as_ptr().add(i)); - let result = _mm_div_ps(a, b); - _mm_storeu_ps(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] / rhs[i]; - } - - Ok(()) -} - -// ARM NEON implementations -#[cfg(target_arch = "aarch64")] -#[target_feature(enable = "neon")] -unsafe fn simd_add_f32_neon(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { - const SIMD_WIDTH: usize = 4; // NEON processes 4 f32s at once - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = vld1q_f32(lhs.as_ptr().add(i)); - let b = vld1q_f32(rhs.as_ptr().add(i)); - let result = vaddq_f32(a, b); - vst1q_f32(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] + rhs[i]; - } - - Ok(()) -} - -#[cfg(target_arch = "aarch64")] -#[target_feature(enable = "neon")] -unsafe fn simd_sub_f32_neon(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { - const SIMD_WIDTH: usize = 4; - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = vld1q_f32(lhs.as_ptr().add(i)); - let b = vld1q_f32(rhs.as_ptr().add(i)); - let result = vsubq_f32(a, b); - vst1q_f32(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] - rhs[i]; - } - - Ok(()) -} - -#[cfg(target_arch = "aarch64")] -#[target_feature(enable = "neon")] -unsafe fn simd_mul_f32_neon(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { - const SIMD_WIDTH: usize = 4; - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = vld1q_f32(lhs.as_ptr().add(i)); - let b = vld1q_f32(rhs.as_ptr().add(i)); - let result = vmulq_f32(a, b); - vst1q_f32(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] * rhs[i]; - } - - Ok(()) -} - -#[cfg(target_arch = "aarch64")] -#[target_feature(enable = "neon")] -unsafe fn simd_div_f32_neon(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { - const SIMD_WIDTH: usize = 4; - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = vld1q_f32(lhs.as_ptr().add(i)); - let b = vld1q_f32(rhs.as_ptr().add(i)); - let result = vdivq_f32(a, b); - vst1q_f32(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] / rhs[i]; - } - - Ok(()) -} - -/// Check if two tensors can use optimized SIMD operations (same shape, contiguous) -pub fn can_use_simd_fast_path(lhs_shape: &Shape, rhs_shape: &Shape, output_shape: &Shape) -> bool { - // For now, only optimize when all shapes are identical (no broadcasting) - // This ensures contiguous memory access patterns optimal for SIMD - lhs_shape.dims() == rhs_shape.dims() - && lhs_shape.dims() == output_shape.dims() - && lhs_shape.numel() >= 16 // Only use SIMD for reasonably sized arrays -} - -/// SIMD-optimized element-wise addition for f64 arrays -pub fn simd_add_f64(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { - if lhs.len() != rhs.len() || lhs.len() != output.len() { - return Err(MinitensorError::invalid_operation( - "Array lengths must match for SIMD operations", - )); - } - - let caps = simd_capabilities(); - - #[cfg(target_arch = "x86_64")] - { - if caps.avx2 { - return unsafe { simd_add_f64_avx2(lhs, rhs, output) }; - } else if caps.sse4_1 { - return unsafe { simd_add_f64_sse(lhs, rhs, output) }; - } - } - - #[cfg(target_arch = "aarch64")] - { - if caps.neon { - return unsafe { simd_add_f64_neon(lhs, rhs, output) }; - } - } - - // Fallback to scalar implementation - simd_add_f64_scalar(lhs, rhs, output) -} - -/// SIMD-optimized element-wise subtraction for f64 arrays -pub fn simd_sub_f64(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { - if lhs.len() != rhs.len() || lhs.len() != output.len() { - return Err(MinitensorError::invalid_operation( - "Array lengths must match for SIMD operations", - )); - } - - let caps = simd_capabilities(); - - #[cfg(target_arch = "x86_64")] - { - if caps.avx2 { - return unsafe { simd_sub_f64_avx2(lhs, rhs, output) }; - } else if caps.sse4_1 { - return unsafe { simd_sub_f64_sse(lhs, rhs, output) }; - } - } - - #[cfg(target_arch = "aarch64")] - { - if caps.neon { - return unsafe { simd_sub_f64_neon(lhs, rhs, output) }; - } - } - - // Fallback to scalar implementation - simd_sub_f64_scalar(lhs, rhs, output) -} - -/// SIMD-optimized element-wise multiplication for f64 arrays -pub fn simd_mul_f64(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { - if lhs.len() != rhs.len() || lhs.len() != output.len() { - return Err(MinitensorError::invalid_operation( - "Array lengths must match for SIMD operations", - )); - } - - let caps = simd_capabilities(); - - #[cfg(target_arch = "x86_64")] - { - if caps.avx2 { - return unsafe { simd_mul_f64_avx2(lhs, rhs, output) }; - } else if caps.sse4_1 { - return unsafe { simd_mul_f64_sse(lhs, rhs, output) }; - } - } - - #[cfg(target_arch = "aarch64")] - { - if caps.neon { - return unsafe { simd_mul_f64_neon(lhs, rhs, output) }; - } - } - - // Fallback to scalar implementation - simd_mul_f64_scalar(lhs, rhs, output) -} - -/// SIMD-optimized element-wise division for f64 arrays -pub fn simd_div_f64(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { - if lhs.len() != rhs.len() || lhs.len() != output.len() { - return Err(MinitensorError::invalid_operation( - "Array lengths must match for SIMD operations", - )); - } - - let caps = simd_capabilities(); - - #[cfg(target_arch = "x86_64")] - { - if caps.avx2 { - return unsafe { simd_div_f64_avx2(lhs, rhs, output) }; - } else if caps.sse4_1 { - return unsafe { simd_div_f64_sse(lhs, rhs, output) }; - } - } - - #[cfg(target_arch = "aarch64")] - { - if caps.neon { - return unsafe { simd_div_f64_neon(lhs, rhs, output) }; - } - } - - // Fallback to scalar implementation - simd_div_f64_scalar(lhs, rhs, output) -} - -// f64 scalar fallback implementations -fn simd_add_f64_scalar(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { - for i in 0..lhs.len() { - output[i] = lhs[i] + rhs[i]; - } - Ok(()) -} - -fn simd_sub_f64_scalar(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { - for i in 0..lhs.len() { - output[i] = lhs[i] - rhs[i]; - } - Ok(()) -} - -fn simd_mul_f64_scalar(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { - for i in 0..lhs.len() { - output[i] = lhs[i] * rhs[i]; - } - Ok(()) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_simd_capabilities_detection() { - let caps = simd_capabilities(); - // Just ensure it doesn't panic and returns something reasonable - println!("SIMD capabilities: {:?}", caps); - } - - #[test] - fn test_simd_add_f32() { - let a = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0]; - let b = vec![8.0, 7.0, 6.0, 5.0, 4.0, 3.0, 2.0, 1.0]; - let mut result = vec![0.0; 8]; - - simd_add_f32(&a, &b, &mut result).unwrap(); - - for i in 0..8 { - assert_eq!(result[i], 9.0); - } - } - - #[test] - fn test_simd_mul_f32() { - let a = vec![2.0, 3.0, 4.0, 5.0]; - let b = vec![3.0, 4.0, 5.0, 6.0]; - let mut result = vec![0.0; 4]; - - simd_mul_f32(&a, &b, &mut result).unwrap(); - - assert_eq!(result, vec![6.0, 12.0, 20.0, 30.0]); - } - - #[test] - fn test_simd_div_f32() { - let a = vec![12.0, 15.0, 20.0, 24.0]; - let b = vec![3.0, 5.0, 4.0, 6.0]; - let mut result = vec![0.0; 4]; - - simd_div_f32(&a, &b, &mut result).unwrap(); - - assert_eq!(result, vec![4.0, 3.0, 5.0, 4.0]); - } - - #[test] - fn test_simd_div_by_zero() { - let a = vec![1.0, 2.0]; - let b = vec![0.0, 2.0]; - let mut result = vec![0.0; 2]; - - simd_div_f32(&a, &b, &mut result).unwrap(); - - assert_eq!(result[0], f32::INFINITY); - assert_eq!(result[1], 1.0); - } - - #[test] - fn test_simd_div_by_zero_ieee_semantics_including_tail() { - let len = 19; - let a: Vec = (0..len) - .map(|i| match i % 3 { - 0 => -1.0, - 1 => 0.0, - _ => 1.0, - }) - .collect(); - let b = vec![0.0_f32; len]; - let mut result = vec![0.0_f32; len]; - simd_div_f32(&a, &b, &mut result).unwrap(); - for i in 0..len { - match i % 3 { - 0 => assert_eq!(result[i], f32::NEG_INFINITY, "index {i}"), - 1 => assert!(result[i].is_nan(), "index {i}"), - _ => assert_eq!(result[i], f32::INFINITY, "index {i}"), - } - } - - let a64: Vec = a.iter().map(|&v| v as f64).collect(); - let b64 = vec![0.0_f64; len]; - let mut result64 = vec![0.0_f64; len]; - simd_div_f64(&a64, &b64, &mut result64).unwrap(); - for i in 0..len { - match i % 3 { - 0 => assert_eq!(result64[i], f64::NEG_INFINITY, "index {i}"), - 1 => assert!(result64[i].is_nan(), "index {i}"), - _ => assert_eq!(result64[i], f64::INFINITY, "index {i}"), - } - } - } - - #[test] - fn test_simd_f32_length_mismatch_errors() { - let a = [1.0_f32, 2.0, 3.0]; - let b = [4.0_f32, 5.0]; - let mut out = [0.0_f32; 3]; - - let err = simd_add_f32(&a, &b, &mut out).unwrap_err(); - assert!( - err.to_string() - .contains("Array lengths must match for SIMD operations") - ); - - let err = simd_sub_f32(&a, &b, &mut out).unwrap_err(); - assert!( - err.to_string() - .contains("Array lengths must match for SIMD operations") - ); - - let err = simd_mul_f32(&a, &b, &mut out).unwrap_err(); - assert!( - err.to_string() - .contains("Array lengths must match for SIMD operations") - ); - - let err = simd_div_f32(&a, &b, &mut out).unwrap_err(); - assert!( - err.to_string() - .contains("Array lengths must match for SIMD operations") - ); - } - - #[test] - fn test_simd_f64_all_ops_with_remainder_and_division_by_zero() { - let a = vec![10.0_f64, -9.0, 8.0, -7.0, 6.0]; - let b = vec![2.0_f64, -3.0, 4.0, -7.0, 0.0]; - let mut out = vec![0.0_f64; a.len()]; - - simd_add_f64(&a, &b, &mut out).unwrap(); - assert_eq!(out, vec![12.0, -12.0, 12.0, -14.0, 6.0]); - - simd_sub_f64(&a, &b, &mut out).unwrap(); - assert_eq!(out, vec![8.0, -6.0, 4.0, 0.0, 6.0]); - - simd_mul_f64(&a, &b, &mut out).unwrap(); - assert_eq!(out, vec![20.0, 27.0, 32.0, 49.0, 0.0]); - - simd_div_f64(&a, &b, &mut out).unwrap(); - assert_eq!(out[..4], [5.0, 3.0, 2.0, 1.0]); - assert_eq!(out[4], f64::INFINITY); - } - - #[test] - fn test_simd_f64_length_mismatch_errors() { - let a = [1.0_f64, 2.0, 3.0]; - let b = [4.0_f64, 5.0]; - let mut out = [0.0_f64; 3]; - - let err = simd_add_f64(&a, &b, &mut out).unwrap_err(); - assert!( - err.to_string() - .contains("Array lengths must match for SIMD operations") - ); - - let err = simd_sub_f64(&a, &b, &mut out).unwrap_err(); - assert!( - err.to_string() - .contains("Array lengths must match for SIMD operations") - ); - - let err = simd_mul_f64(&a, &b, &mut out).unwrap_err(); - assert!( - err.to_string() - .contains("Array lengths must match for SIMD operations") - ); - - let err = simd_div_f64(&a, &b, &mut out).unwrap_err(); - assert!( - err.to_string() - .contains("Array lengths must match for SIMD operations") - ); - } - - #[test] - fn test_can_use_simd_fast_path_shape_conditions() { - let same = Shape::new(vec![2, 8]); - let different = Shape::new(vec![4, 4]); - let too_small = Shape::new(vec![2, 4]); - - assert!(can_use_simd_fast_path(&same, &same, &same)); - assert!(!can_use_simd_fast_path(&same, &different, &same)); - assert!(!can_use_simd_fast_path(&same, &same, &different)); - assert!(!can_use_simd_fast_path(&too_small, &too_small, &too_small)); - } -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; + +use crate::{ + error::{MinitensorError, Result}, + tensor::Shape, +}; + +#[cfg(target_arch = "x86_64")] +use std::arch::x86_64::*; + +#[cfg(target_arch = "aarch64")] +use std::arch::aarch64::*; + +/// SIMD capabilities detected at runtime +#[derive(Debug, Clone, Copy)] +pub struct SimdCapabilities { + pub avx2: bool, + pub avx512: bool, + pub sse4_1: bool, + pub neon: bool, + pub sve: bool, +} + +impl SimdCapabilities { + /// Detect SIMD capabilities at runtime + pub fn detect() -> Self { + Self { + #[cfg(target_arch = "x86_64")] + avx2: is_x86_feature_detected!("avx2"), + #[cfg(target_arch = "x86_64")] + avx512: is_x86_feature_detected!("avx512f"), + #[cfg(target_arch = "x86_64")] + sse4_1: is_x86_feature_detected!("sse4.1"), + #[cfg(not(target_arch = "x86_64"))] + avx2: false, + #[cfg(not(target_arch = "x86_64"))] + avx512: false, + #[cfg(not(target_arch = "x86_64"))] + sse4_1: false, + + #[cfg(target_arch = "aarch64")] + neon: std::arch::is_aarch64_feature_detected!("neon"), + #[cfg(target_arch = "aarch64")] + sve: std::arch::is_aarch64_feature_detected!("sve"), + #[cfg(not(target_arch = "aarch64"))] + neon: false, + #[cfg(not(target_arch = "aarch64"))] + sve: false, + } + } +} + +/// Global SIMD capabilities (detected once at startup) +static SIMD_CAPS: std::sync::OnceLock = std::sync::OnceLock::new(); + +/// Get the detected SIMD capabilities +pub fn simd_capabilities() -> SimdCapabilities { + *SIMD_CAPS.get_or_init(SimdCapabilities::detect) +} + +/// SIMD-optimized element-wise addition for f32 arrays +pub fn simd_add_f32(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { + if lhs.len() != rhs.len() || lhs.len() != output.len() { + return Err(MinitensorError::invalid_operation( + "Array lengths must match for SIMD operations", + )); + } + + let caps = simd_capabilities(); + + #[cfg(target_arch = "x86_64")] + { + if caps.avx2 { + return unsafe { simd_add_f32_avx2(lhs, rhs, output) }; + } else if caps.sse4_1 { + return unsafe { simd_add_f32_sse(lhs, rhs, output) }; + } + } + + #[cfg(target_arch = "aarch64")] + { + if caps.neon { + return unsafe { simd_add_f32_neon(lhs, rhs, output) }; + } + } + + // Fallback to scalar implementation + simd_add_f32_scalar(lhs, rhs, output) +} + +/// SIMD-optimized element-wise subtraction for f32 arrays +pub fn simd_sub_f32(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { + if lhs.len() != rhs.len() || lhs.len() != output.len() { + return Err(MinitensorError::invalid_operation( + "Array lengths must match for SIMD operations", + )); + } + + let caps = simd_capabilities(); + + #[cfg(target_arch = "x86_64")] + { + if caps.avx2 { + return unsafe { simd_sub_f32_avx2(lhs, rhs, output) }; + } else if caps.sse4_1 { + return unsafe { simd_sub_f32_sse(lhs, rhs, output) }; + } + } + + #[cfg(target_arch = "aarch64")] + { + if caps.neon { + return unsafe { simd_sub_f32_neon(lhs, rhs, output) }; + } + } + + // Fallback to scalar implementation + simd_sub_f32_scalar(lhs, rhs, output) +} + +/// SIMD-optimized element-wise multiplication for f32 arrays +pub fn simd_mul_f32(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { + if lhs.len() != rhs.len() || lhs.len() != output.len() { + return Err(MinitensorError::invalid_operation( + "Array lengths must match for SIMD operations", + )); + } + + let caps = simd_capabilities(); + + #[cfg(target_arch = "x86_64")] + { + if caps.avx2 { + return unsafe { simd_mul_f32_avx2(lhs, rhs, output) }; + } else if caps.sse4_1 { + return unsafe { simd_mul_f32_sse(lhs, rhs, output) }; + } + } + + #[cfg(target_arch = "aarch64")] + { + if caps.neon { + return unsafe { simd_mul_f32_neon(lhs, rhs, output) }; + } + } + + // Fallback to scalar implementation + simd_mul_f32_scalar(lhs, rhs, output) +} + +/// SIMD-optimized element-wise division for f32 arrays +pub fn simd_div_f32(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { + if lhs.len() != rhs.len() || lhs.len() != output.len() { + return Err(MinitensorError::invalid_operation( + "Array lengths must match for SIMD operations", + )); + } + + let caps = simd_capabilities(); + + #[cfg(target_arch = "x86_64")] + { + if caps.avx2 { + return unsafe { simd_div_f32_avx2(lhs, rhs, output) }; + } else if caps.sse4_1 { + return unsafe { simd_div_f32_sse(lhs, rhs, output) }; + } + } + + #[cfg(target_arch = "aarch64")] + { + if caps.neon { + return unsafe { simd_div_f32_neon(lhs, rhs, output) }; + } + } + + // Fallback to scalar implementation + simd_div_f32_scalar(lhs, rhs, output) +} + +// Scalar fallback implementations +fn simd_add_f32_scalar(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { + for i in 0..lhs.len() { + output[i] = lhs[i] + rhs[i]; + } + Ok(()) +} + +/// Unrolled sum for f32 slices to leverage auto-vectorization +pub fn simd_sum_f32(data: &[f32]) -> f32 { + let mut sums = [0f32; 8]; + let chunks = data.chunks_exact(8); + let rem = chunks.remainder(); + for chunk in chunks { + sums[0] += chunk[0]; + sums[1] += chunk[1]; + sums[2] += chunk[2]; + sums[3] += chunk[3]; + sums[4] += chunk[4]; + sums[5] += chunk[5]; + sums[6] += chunk[6]; + sums[7] += chunk[7]; + } + let mut total: f32 = sums.iter().sum(); + total += rem.iter().copied().sum::(); + total +} + +/// Unrolled sum for f64 slices to leverage auto-vectorization +pub fn simd_sum_f64(data: &[f64]) -> f64 { + let mut sums = [0f64; 4]; + let chunks = data.chunks_exact(4); + let rem = chunks.remainder(); + for chunk in chunks { + sums[0] += chunk[0]; + sums[1] += chunk[1]; + sums[2] += chunk[2]; + sums[3] += chunk[3]; + } + let mut total: f64 = sums.iter().sum(); + total += rem.iter().copied().sum::(); + total +} + +/// Unrolled sum for i32 slices to leverage auto-vectorization +pub fn simd_sum_i32(data: &[i32]) -> i32 { + let mut sums = [0i32; 8]; + let chunks = data.chunks_exact(8); + let rem = chunks.remainder(); + for chunk in chunks { + sums[0] += chunk[0]; + sums[1] += chunk[1]; + sums[2] += chunk[2]; + sums[3] += chunk[3]; + sums[4] += chunk[4]; + sums[5] += chunk[5]; + sums[6] += chunk[6]; + sums[7] += chunk[7]; + } + let mut total: i32 = sums.iter().sum(); + total += rem.iter().copied().sum::(); + total +} + +/// Unrolled sum for i64 slices to leverage auto-vectorization +pub fn simd_sum_i64(data: &[i64]) -> i64 { + let mut sums = [0i64; 4]; + let chunks = data.chunks_exact(4); + let rem = chunks.remainder(); + for chunk in chunks { + sums[0] += chunk[0]; + sums[1] += chunk[1]; + sums[2] += chunk[2]; + sums[3] += chunk[3]; + } + let mut total: i64 = sums.iter().sum(); + total += rem.iter().copied().sum::(); + total +} + +/// Unrolled product for f32 slices to leverage auto-vectorization +pub fn simd_prod_f32(data: &[f32]) -> f32 { + let mut prods = [1f32; 8]; + let chunks = data.chunks_exact(8); + let rem = chunks.remainder(); + for chunk in chunks { + prods[0] *= chunk[0]; + prods[1] *= chunk[1]; + prods[2] *= chunk[2]; + prods[3] *= chunk[3]; + prods[4] *= chunk[4]; + prods[5] *= chunk[5]; + prods[6] *= chunk[6]; + prods[7] *= chunk[7]; + } + let mut total: f32 = prods.iter().product(); + total *= rem.iter().copied().product::(); + total +} + +/// Unrolled product for f64 slices to leverage auto-vectorization +pub fn simd_prod_f64(data: &[f64]) -> f64 { + let mut prods = [1f64; 4]; + let chunks = data.chunks_exact(4); + let rem = chunks.remainder(); + for chunk in chunks { + prods[0] *= chunk[0]; + prods[1] *= chunk[1]; + prods[2] *= chunk[2]; + prods[3] *= chunk[3]; + } + let mut total: f64 = prods.iter().product(); + total *= rem.iter().copied().product::(); + total +} + +/// Unrolled product for i32 slices to leverage auto-vectorization +pub fn simd_prod_i32(data: &[i32]) -> i32 { + let mut prods = [1i32; 8]; + let chunks = data.chunks_exact(8); + let rem = chunks.remainder(); + for chunk in chunks { + prods[0] *= chunk[0]; + prods[1] *= chunk[1]; + prods[2] *= chunk[2]; + prods[3] *= chunk[3]; + prods[4] *= chunk[4]; + prods[5] *= chunk[5]; + prods[6] *= chunk[6]; + prods[7] *= chunk[7]; + } + let mut total: i32 = prods.iter().product(); + total *= rem.iter().copied().product::(); + total +} + +/// Unrolled product for i64 slices to leverage auto-vectorization +pub fn simd_prod_i64(data: &[i64]) -> i64 { + let mut prods = [1i64; 4]; + let chunks = data.chunks_exact(4); + let rem = chunks.remainder(); + for chunk in chunks { + prods[0] *= chunk[0]; + prods[1] *= chunk[1]; + prods[2] *= chunk[2]; + prods[3] *= chunk[3]; + } + let mut total: i64 = prods.iter().product(); + total *= rem.iter().copied().product::(); + total +} + +fn simd_sub_f32_scalar(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { + for i in 0..lhs.len() { + output[i] = lhs[i] - rhs[i]; + } + Ok(()) +} + +fn simd_mul_f32_scalar(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { + for i in 0..lhs.len() { + output[i] = lhs[i] * rhs[i]; + } + Ok(()) +} + +fn simd_div_f32_scalar(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { + for i in 0..lhs.len() { + output[i] = lhs[i] / rhs[i]; + } + Ok(()) +} + +// x86_64 AVX2 implementations +#[cfg(target_arch = "x86_64")] +#[target_feature(enable = "avx2")] +unsafe fn simd_add_f32_avx2(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { + const SIMD_WIDTH: usize = 8; // AVX2 processes 8 f32s at once + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + // Process SIMD_WIDTH elements at a time + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = _mm256_loadu_ps(lhs.as_ptr().add(i)); + let b = _mm256_loadu_ps(rhs.as_ptr().add(i)); + let result = _mm256_add_ps(a, b); + _mm256_storeu_ps(output.as_mut_ptr().add(i), result); + } + } + + // Handle remaining elements + for i in simd_len..len { + output[i] = lhs[i] + rhs[i]; + } + + Ok(()) +} + +#[cfg(target_arch = "x86_64")] +#[target_feature(enable = "avx2")] +unsafe fn simd_sub_f32_avx2(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { + const SIMD_WIDTH: usize = 8; + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = _mm256_loadu_ps(lhs.as_ptr().add(i)); + let b = _mm256_loadu_ps(rhs.as_ptr().add(i)); + let result = _mm256_sub_ps(a, b); + _mm256_storeu_ps(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] - rhs[i]; + } + + Ok(()) +} + +#[cfg(target_arch = "x86_64")] +#[target_feature(enable = "avx2")] +unsafe fn simd_mul_f32_avx2(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { + const SIMD_WIDTH: usize = 8; + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = _mm256_loadu_ps(lhs.as_ptr().add(i)); + let b = _mm256_loadu_ps(rhs.as_ptr().add(i)); + let result = _mm256_mul_ps(a, b); + _mm256_storeu_ps(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] * rhs[i]; + } + + Ok(()) +} + +#[cfg(target_arch = "x86_64")] +#[target_feature(enable = "avx2")] +unsafe fn simd_div_f32_avx2(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { + const SIMD_WIDTH: usize = 8; + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = _mm256_loadu_ps(lhs.as_ptr().add(i)); + let b = _mm256_loadu_ps(rhs.as_ptr().add(i)); + let result = _mm256_div_ps(a, b); + _mm256_storeu_ps(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] / rhs[i]; + } + + Ok(()) +} + +// x86_64 SSE implementations +#[cfg(target_arch = "x86_64")] +#[target_feature(enable = "sse4.1")] +unsafe fn simd_add_f32_sse(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { + const SIMD_WIDTH: usize = 4; // SSE processes 4 f32s at once + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = _mm_loadu_ps(lhs.as_ptr().add(i)); + let b = _mm_loadu_ps(rhs.as_ptr().add(i)); + let result = _mm_add_ps(a, b); + _mm_storeu_ps(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] + rhs[i]; + } + + Ok(()) +} + +#[cfg(target_arch = "x86_64")] +#[target_feature(enable = "sse4.1")] +unsafe fn simd_sub_f32_sse(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { + const SIMD_WIDTH: usize = 4; + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = _mm_loadu_ps(lhs.as_ptr().add(i)); + let b = _mm_loadu_ps(rhs.as_ptr().add(i)); + let result = _mm_sub_ps(a, b); + _mm_storeu_ps(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] - rhs[i]; + } + + Ok(()) +} + +#[cfg(target_arch = "x86_64")] +#[target_feature(enable = "sse4.1")] +unsafe fn simd_mul_f32_sse(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { + const SIMD_WIDTH: usize = 4; + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = _mm_loadu_ps(lhs.as_ptr().add(i)); + let b = _mm_loadu_ps(rhs.as_ptr().add(i)); + let result = _mm_mul_ps(a, b); + _mm_storeu_ps(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] * rhs[i]; + } + + Ok(()) +} + +#[cfg(target_arch = "x86_64")] +#[target_feature(enable = "sse4.1")] +unsafe fn simd_div_f32_sse(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { + const SIMD_WIDTH: usize = 4; + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = _mm_loadu_ps(lhs.as_ptr().add(i)); + let b = _mm_loadu_ps(rhs.as_ptr().add(i)); + let result = _mm_div_ps(a, b); + _mm_storeu_ps(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] / rhs[i]; + } + + Ok(()) +} + +// ARM NEON implementations +#[cfg(target_arch = "aarch64")] +#[target_feature(enable = "neon")] +unsafe fn simd_add_f32_neon(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { + const SIMD_WIDTH: usize = 4; // NEON processes 4 f32s at once + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = vld1q_f32(lhs.as_ptr().add(i)); + let b = vld1q_f32(rhs.as_ptr().add(i)); + let result = vaddq_f32(a, b); + vst1q_f32(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] + rhs[i]; + } + + Ok(()) +} + +#[cfg(target_arch = "aarch64")] +#[target_feature(enable = "neon")] +unsafe fn simd_sub_f32_neon(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { + const SIMD_WIDTH: usize = 4; + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = vld1q_f32(lhs.as_ptr().add(i)); + let b = vld1q_f32(rhs.as_ptr().add(i)); + let result = vsubq_f32(a, b); + vst1q_f32(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] - rhs[i]; + } + + Ok(()) +} + +#[cfg(target_arch = "aarch64")] +#[target_feature(enable = "neon")] +unsafe fn simd_mul_f32_neon(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { + const SIMD_WIDTH: usize = 4; + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = vld1q_f32(lhs.as_ptr().add(i)); + let b = vld1q_f32(rhs.as_ptr().add(i)); + let result = vmulq_f32(a, b); + vst1q_f32(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] * rhs[i]; + } + + Ok(()) +} + +#[cfg(target_arch = "aarch64")] +#[target_feature(enable = "neon")] +unsafe fn simd_div_f32_neon(lhs: &[f32], rhs: &[f32], output: &mut [f32]) -> Result<()> { + const SIMD_WIDTH: usize = 4; + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = vld1q_f32(lhs.as_ptr().add(i)); + let b = vld1q_f32(rhs.as_ptr().add(i)); + let result = vdivq_f32(a, b); + vst1q_f32(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] / rhs[i]; + } + + Ok(()) +} + +/// Check if two tensors can use optimized SIMD operations (same shape, contiguous) +pub fn can_use_simd_fast_path(lhs_shape: &Shape, rhs_shape: &Shape, output_shape: &Shape) -> bool { + // For now, only optimize when all shapes are identical (no broadcasting) + // This ensures contiguous memory access patterns optimal for SIMD + lhs_shape.dims() == rhs_shape.dims() + && lhs_shape.dims() == output_shape.dims() + && lhs_shape.numel() >= 16 // Only use SIMD for reasonably sized arrays +} + +/// SIMD-optimized element-wise addition for f64 arrays +pub fn simd_add_f64(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { + if lhs.len() != rhs.len() || lhs.len() != output.len() { + return Err(MinitensorError::invalid_operation( + "Array lengths must match for SIMD operations", + )); + } + + let caps = simd_capabilities(); + + #[cfg(target_arch = "x86_64")] + { + if caps.avx2 { + return unsafe { simd_add_f64_avx2(lhs, rhs, output) }; + } else if caps.sse4_1 { + return unsafe { simd_add_f64_sse(lhs, rhs, output) }; + } + } + + #[cfg(target_arch = "aarch64")] + { + if caps.neon { + return unsafe { simd_add_f64_neon(lhs, rhs, output) }; + } + } + + // Fallback to scalar implementation + simd_add_f64_scalar(lhs, rhs, output) +} + +/// SIMD-optimized element-wise subtraction for f64 arrays +pub fn simd_sub_f64(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { + if lhs.len() != rhs.len() || lhs.len() != output.len() { + return Err(MinitensorError::invalid_operation( + "Array lengths must match for SIMD operations", + )); + } + + let caps = simd_capabilities(); + + #[cfg(target_arch = "x86_64")] + { + if caps.avx2 { + return unsafe { simd_sub_f64_avx2(lhs, rhs, output) }; + } else if caps.sse4_1 { + return unsafe { simd_sub_f64_sse(lhs, rhs, output) }; + } + } + + #[cfg(target_arch = "aarch64")] + { + if caps.neon { + return unsafe { simd_sub_f64_neon(lhs, rhs, output) }; + } + } + + // Fallback to scalar implementation + simd_sub_f64_scalar(lhs, rhs, output) +} + +/// SIMD-optimized element-wise multiplication for f64 arrays +pub fn simd_mul_f64(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { + if lhs.len() != rhs.len() || lhs.len() != output.len() { + return Err(MinitensorError::invalid_operation( + "Array lengths must match for SIMD operations", + )); + } + + let caps = simd_capabilities(); + + #[cfg(target_arch = "x86_64")] + { + if caps.avx2 { + return unsafe { simd_mul_f64_avx2(lhs, rhs, output) }; + } else if caps.sse4_1 { + return unsafe { simd_mul_f64_sse(lhs, rhs, output) }; + } + } + + #[cfg(target_arch = "aarch64")] + { + if caps.neon { + return unsafe { simd_mul_f64_neon(lhs, rhs, output) }; + } + } + + // Fallback to scalar implementation + simd_mul_f64_scalar(lhs, rhs, output) +} + +/// SIMD-optimized element-wise division for f64 arrays +pub fn simd_div_f64(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { + if lhs.len() != rhs.len() || lhs.len() != output.len() { + return Err(MinitensorError::invalid_operation( + "Array lengths must match for SIMD operations", + )); + } + + let caps = simd_capabilities(); + + #[cfg(target_arch = "x86_64")] + { + if caps.avx2 { + return unsafe { simd_div_f64_avx2(lhs, rhs, output) }; + } else if caps.sse4_1 { + return unsafe { simd_div_f64_sse(lhs, rhs, output) }; + } + } + + #[cfg(target_arch = "aarch64")] + { + if caps.neon { + return unsafe { simd_div_f64_neon(lhs, rhs, output) }; + } + } + + // Fallback to scalar implementation + simd_div_f64_scalar(lhs, rhs, output) +} + +// f64 scalar fallback implementations +fn simd_add_f64_scalar(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { + for i in 0..lhs.len() { + output[i] = lhs[i] + rhs[i]; + } + Ok(()) +} + +fn simd_sub_f64_scalar(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { + for i in 0..lhs.len() { + output[i] = lhs[i] - rhs[i]; + } + Ok(()) +} + +fn simd_mul_f64_scalar(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { + for i in 0..lhs.len() { + output[i] = lhs[i] * rhs[i]; + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_simd_capabilities_detection() { + let caps = simd_capabilities(); + // Just ensure it doesn't panic and returns something reasonable + println!("SIMD capabilities: {:?}", caps); + } + + #[test] + fn test_simd_add_f32() { + let a = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0]; + let b = vec![8.0, 7.0, 6.0, 5.0, 4.0, 3.0, 2.0, 1.0]; + let mut result = vec![0.0; 8]; + + simd_add_f32(&a, &b, &mut result).unwrap(); + + for i in 0..8 { + assert_eq!(result[i], 9.0); + } + } + + #[test] + fn test_simd_mul_f32() { + let a = vec![2.0, 3.0, 4.0, 5.0]; + let b = vec![3.0, 4.0, 5.0, 6.0]; + let mut result = vec![0.0; 4]; + + simd_mul_f32(&a, &b, &mut result).unwrap(); + + assert_eq!(result, vec![6.0, 12.0, 20.0, 30.0]); + } + + #[test] + fn test_simd_div_f32() { + let a = vec![12.0, 15.0, 20.0, 24.0]; + let b = vec![3.0, 5.0, 4.0, 6.0]; + let mut result = vec![0.0; 4]; + + simd_div_f32(&a, &b, &mut result).unwrap(); + + assert_eq!(result, vec![4.0, 3.0, 5.0, 4.0]); + } + + #[test] + fn test_simd_div_by_zero() { + let a = vec![1.0, 2.0]; + let b = vec![0.0, 2.0]; + let mut result = vec![0.0; 2]; + + simd_div_f32(&a, &b, &mut result).unwrap(); + + assert_eq!(result[0], f32::INFINITY); + assert_eq!(result[1], 1.0); + } + + #[test] + fn test_simd_div_by_zero_ieee_semantics_including_tail() { + let len = 19; + let a: Vec = (0..len) + .map(|i| match i % 3 { + 0 => -1.0, + 1 => 0.0, + _ => 1.0, + }) + .collect(); + let b = vec![0.0_f32; len]; + let mut result = vec![0.0_f32; len]; + simd_div_f32(&a, &b, &mut result).unwrap(); + for i in 0..len { + match i % 3 { + 0 => assert_eq!(result[i], f32::NEG_INFINITY, "index {i}"), + 1 => assert!(result[i].is_nan(), "index {i}"), + _ => assert_eq!(result[i], f32::INFINITY, "index {i}"), + } + } + + let a64: Vec = a.iter().map(|&v| v as f64).collect(); + let b64 = vec![0.0_f64; len]; + let mut result64 = vec![0.0_f64; len]; + simd_div_f64(&a64, &b64, &mut result64).unwrap(); + for i in 0..len { + match i % 3 { + 0 => assert_eq!(result64[i], f64::NEG_INFINITY, "index {i}"), + 1 => assert!(result64[i].is_nan(), "index {i}"), + _ => assert_eq!(result64[i], f64::INFINITY, "index {i}"), + } + } + } + + #[test] + fn test_simd_f32_length_mismatch_errors() { + let a = [1.0_f32, 2.0, 3.0]; + let b = [4.0_f32, 5.0]; + let mut out = [0.0_f32; 3]; + + let err = simd_add_f32(&a, &b, &mut out).unwrap_err(); + assert!( + err.to_string() + .contains("Array lengths must match for SIMD operations") + ); + + let err = simd_sub_f32(&a, &b, &mut out).unwrap_err(); + assert!( + err.to_string() + .contains("Array lengths must match for SIMD operations") + ); + + let err = simd_mul_f32(&a, &b, &mut out).unwrap_err(); + assert!( + err.to_string() + .contains("Array lengths must match for SIMD operations") + ); + + let err = simd_div_f32(&a, &b, &mut out).unwrap_err(); + assert!( + err.to_string() + .contains("Array lengths must match for SIMD operations") + ); + } + + #[test] + fn test_simd_f64_all_ops_with_remainder_and_division_by_zero() { + let a = vec![10.0_f64, -9.0, 8.0, -7.0, 6.0]; + let b = vec![2.0_f64, -3.0, 4.0, -7.0, 0.0]; + let mut out = vec![0.0_f64; a.len()]; + + simd_add_f64(&a, &b, &mut out).unwrap(); + assert_eq!(out, vec![12.0, -12.0, 12.0, -14.0, 6.0]); + + simd_sub_f64(&a, &b, &mut out).unwrap(); + assert_eq!(out, vec![8.0, -6.0, 4.0, 0.0, 6.0]); + + simd_mul_f64(&a, &b, &mut out).unwrap(); + assert_eq!(out, vec![20.0, 27.0, 32.0, 49.0, 0.0]); + + simd_div_f64(&a, &b, &mut out).unwrap(); + assert_eq!(out[..4], [5.0, 3.0, 2.0, 1.0]); + assert_eq!(out[4], f64::INFINITY); + } + + #[test] + fn test_simd_f64_length_mismatch_errors() { + let a = [1.0_f64, 2.0, 3.0]; + let b = [4.0_f64, 5.0]; + let mut out = [0.0_f64; 3]; + + let err = simd_add_f64(&a, &b, &mut out).unwrap_err(); + assert!( + err.to_string() + .contains("Array lengths must match for SIMD operations") + ); + + let err = simd_sub_f64(&a, &b, &mut out).unwrap_err(); + assert!( + err.to_string() + .contains("Array lengths must match for SIMD operations") + ); + + let err = simd_mul_f64(&a, &b, &mut out).unwrap_err(); + assert!( + err.to_string() + .contains("Array lengths must match for SIMD operations") + ); + + let err = simd_div_f64(&a, &b, &mut out).unwrap_err(); + assert!( + err.to_string() + .contains("Array lengths must match for SIMD operations") + ); + } + + #[test] + fn test_can_use_simd_fast_path_shape_conditions() { + let same = Shape::new(vec![2, 8]); + let different = Shape::new(vec![4, 4]); + let too_small = Shape::new(vec![2, 4]); + + assert!(can_use_simd_fast_path(&same, &same, &same)); + assert!(!can_use_simd_fast_path(&same, &different, &same)); + assert!(!can_use_simd_fast_path(&same, &same, &different)); + assert!(!can_use_simd_fast_path(&too_small, &too_small, &too_small)); + } +} diff --git a/engine/src/operations/simd/utils.rs b/engine/src/operations/simd/utils.rs index 921745e9..6f1e5a0d 100644 --- a/engine/src/operations/simd/utils.rs +++ b/engine/src/operations/simd/utils.rs @@ -1,303 +1,307 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -fn simd_div_f64_scalar(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { - for i in 0..lhs.len() { - output[i] = lhs[i] / rhs[i]; - } - Ok(()) -} - -// x86_64 AVX2 f64 implementations -#[cfg(target_arch = "x86_64")] -#[target_feature(enable = "avx2")] -unsafe fn simd_add_f64_avx2(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { - const SIMD_WIDTH: usize = 4; // AVX2 processes 4 f64s at once - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = _mm256_loadu_pd(lhs.as_ptr().add(i)); - let b = _mm256_loadu_pd(rhs.as_ptr().add(i)); - let result = _mm256_add_pd(a, b); - _mm256_storeu_pd(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] + rhs[i]; - } - - Ok(()) -} - -#[cfg(target_arch = "x86_64")] -#[target_feature(enable = "avx2")] -unsafe fn simd_sub_f64_avx2(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { - const SIMD_WIDTH: usize = 4; - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = _mm256_loadu_pd(lhs.as_ptr().add(i)); - let b = _mm256_loadu_pd(rhs.as_ptr().add(i)); - let result = _mm256_sub_pd(a, b); - _mm256_storeu_pd(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] - rhs[i]; - } - - Ok(()) -} - -#[cfg(target_arch = "x86_64")] -#[target_feature(enable = "avx2")] -unsafe fn simd_mul_f64_avx2(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { - const SIMD_WIDTH: usize = 4; - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = _mm256_loadu_pd(lhs.as_ptr().add(i)); - let b = _mm256_loadu_pd(rhs.as_ptr().add(i)); - let result = _mm256_mul_pd(a, b); - _mm256_storeu_pd(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] * rhs[i]; - } - - Ok(()) -} - -#[cfg(target_arch = "x86_64")] -#[target_feature(enable = "avx2")] -unsafe fn simd_div_f64_avx2(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { - const SIMD_WIDTH: usize = 4; - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = _mm256_loadu_pd(lhs.as_ptr().add(i)); - let b = _mm256_loadu_pd(rhs.as_ptr().add(i)); - let result = _mm256_div_pd(a, b); - _mm256_storeu_pd(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] / rhs[i]; - } - - Ok(()) -} - -// x86_64 SSE f64 implementations -#[cfg(target_arch = "x86_64")] -#[target_feature(enable = "sse4.1")] -unsafe fn simd_add_f64_sse(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { - const SIMD_WIDTH: usize = 2; // SSE processes 2 f64s at once - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = _mm_loadu_pd(lhs.as_ptr().add(i)); - let b = _mm_loadu_pd(rhs.as_ptr().add(i)); - let result = _mm_add_pd(a, b); - _mm_storeu_pd(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] + rhs[i]; - } - - Ok(()) -} - -#[cfg(target_arch = "x86_64")] -#[target_feature(enable = "sse4.1")] -unsafe fn simd_sub_f64_sse(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { - const SIMD_WIDTH: usize = 2; - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = _mm_loadu_pd(lhs.as_ptr().add(i)); - let b = _mm_loadu_pd(rhs.as_ptr().add(i)); - let result = _mm_sub_pd(a, b); - _mm_storeu_pd(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] - rhs[i]; - } - - Ok(()) -} - -#[cfg(target_arch = "x86_64")] -#[target_feature(enable = "sse4.1")] -unsafe fn simd_mul_f64_sse(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { - const SIMD_WIDTH: usize = 2; - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = _mm_loadu_pd(lhs.as_ptr().add(i)); - let b = _mm_loadu_pd(rhs.as_ptr().add(i)); - let result = _mm_mul_pd(a, b); - _mm_storeu_pd(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] * rhs[i]; - } - - Ok(()) -} - -#[cfg(target_arch = "x86_64")] -#[target_feature(enable = "sse4.1")] -unsafe fn simd_div_f64_sse(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { - const SIMD_WIDTH: usize = 2; - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = _mm_loadu_pd(lhs.as_ptr().add(i)); - let b = _mm_loadu_pd(rhs.as_ptr().add(i)); - let result = _mm_div_pd(a, b); - _mm_storeu_pd(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] / rhs[i]; - } - - Ok(()) -} - -// ARM NEON f64 implementations -#[cfg(target_arch = "aarch64")] -#[target_feature(enable = "neon")] -unsafe fn simd_add_f64_neon(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { - const SIMD_WIDTH: usize = 2; // NEON processes 2 f64s at once - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = vld1q_f64(lhs.as_ptr().add(i)); - let b = vld1q_f64(rhs.as_ptr().add(i)); - let result = vaddq_f64(a, b); - vst1q_f64(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] + rhs[i]; - } - - Ok(()) -} - -#[cfg(target_arch = "aarch64")] -#[target_feature(enable = "neon")] -unsafe fn simd_sub_f64_neon(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { - const SIMD_WIDTH: usize = 2; - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = vld1q_f64(lhs.as_ptr().add(i)); - let b = vld1q_f64(rhs.as_ptr().add(i)); - let result = vsubq_f64(a, b); - vst1q_f64(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] - rhs[i]; - } - - Ok(()) -} - -#[cfg(target_arch = "aarch64")] -#[target_feature(enable = "neon")] -unsafe fn simd_mul_f64_neon(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { - const SIMD_WIDTH: usize = 2; - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = vld1q_f64(lhs.as_ptr().add(i)); - let b = vld1q_f64(rhs.as_ptr().add(i)); - let result = vmulq_f64(a, b); - vst1q_f64(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] * rhs[i]; - } - - Ok(()) -} - -#[cfg(target_arch = "aarch64")] -#[target_feature(enable = "neon")] -unsafe fn simd_div_f64_neon(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { - const SIMD_WIDTH: usize = 2; - - let len = lhs.len(); - let simd_len = len - (len % SIMD_WIDTH); - - for i in (0..simd_len).step_by(SIMD_WIDTH) { - unsafe { - let a = vld1q_f64(lhs.as_ptr().add(i)); - let b = vld1q_f64(rhs.as_ptr().add(i)); - let result = vdivq_f64(a, b); - vst1q_f64(output.as_mut_ptr().add(i), result); - } - } - - for i in simd_len..len { - output[i] = lhs[i] / rhs[i]; - } - - Ok(()) -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use crate::error::Result; +#[cfg(target_arch = "x86_64")] +use std::arch::x86_64::*; + +pub(crate) fn simd_div_f64_scalar(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { + for i in 0..lhs.len() { + output[i] = lhs[i] / rhs[i]; + } + Ok(()) +} + +// x86_64 AVX2 f64 implementations +#[cfg(target_arch = "x86_64")] +#[target_feature(enable = "avx2")] +pub(crate) unsafe fn simd_add_f64_avx2(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { + const SIMD_WIDTH: usize = 4; // AVX2 processes 4 f64s at once + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = _mm256_loadu_pd(lhs.as_ptr().add(i)); + let b = _mm256_loadu_pd(rhs.as_ptr().add(i)); + let result = _mm256_add_pd(a, b); + _mm256_storeu_pd(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] + rhs[i]; + } + + Ok(()) +} + +#[cfg(target_arch = "x86_64")] +#[target_feature(enable = "avx2")] +pub(crate) unsafe fn simd_sub_f64_avx2(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { + const SIMD_WIDTH: usize = 4; + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = _mm256_loadu_pd(lhs.as_ptr().add(i)); + let b = _mm256_loadu_pd(rhs.as_ptr().add(i)); + let result = _mm256_sub_pd(a, b); + _mm256_storeu_pd(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] - rhs[i]; + } + + Ok(()) +} + +#[cfg(target_arch = "x86_64")] +#[target_feature(enable = "avx2")] +pub(crate) unsafe fn simd_mul_f64_avx2(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { + const SIMD_WIDTH: usize = 4; + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = _mm256_loadu_pd(lhs.as_ptr().add(i)); + let b = _mm256_loadu_pd(rhs.as_ptr().add(i)); + let result = _mm256_mul_pd(a, b); + _mm256_storeu_pd(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] * rhs[i]; + } + + Ok(()) +} + +#[cfg(target_arch = "x86_64")] +#[target_feature(enable = "avx2")] +pub(crate) unsafe fn simd_div_f64_avx2(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { + const SIMD_WIDTH: usize = 4; + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = _mm256_loadu_pd(lhs.as_ptr().add(i)); + let b = _mm256_loadu_pd(rhs.as_ptr().add(i)); + let result = _mm256_div_pd(a, b); + _mm256_storeu_pd(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] / rhs[i]; + } + + Ok(()) +} + +// x86_64 SSE f64 implementations +#[cfg(target_arch = "x86_64")] +#[target_feature(enable = "sse4.1")] +pub(crate) unsafe fn simd_add_f64_sse(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { + const SIMD_WIDTH: usize = 2; // SSE processes 2 f64s at once + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = _mm_loadu_pd(lhs.as_ptr().add(i)); + let b = _mm_loadu_pd(rhs.as_ptr().add(i)); + let result = _mm_add_pd(a, b); + _mm_storeu_pd(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] + rhs[i]; + } + + Ok(()) +} + +#[cfg(target_arch = "x86_64")] +#[target_feature(enable = "sse4.1")] +pub(crate) unsafe fn simd_sub_f64_sse(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { + const SIMD_WIDTH: usize = 2; + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = _mm_loadu_pd(lhs.as_ptr().add(i)); + let b = _mm_loadu_pd(rhs.as_ptr().add(i)); + let result = _mm_sub_pd(a, b); + _mm_storeu_pd(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] - rhs[i]; + } + + Ok(()) +} + +#[cfg(target_arch = "x86_64")] +#[target_feature(enable = "sse4.1")] +pub(crate) unsafe fn simd_mul_f64_sse(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { + const SIMD_WIDTH: usize = 2; + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = _mm_loadu_pd(lhs.as_ptr().add(i)); + let b = _mm_loadu_pd(rhs.as_ptr().add(i)); + let result = _mm_mul_pd(a, b); + _mm_storeu_pd(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] * rhs[i]; + } + + Ok(()) +} + +#[cfg(target_arch = "x86_64")] +#[target_feature(enable = "sse4.1")] +pub(crate) unsafe fn simd_div_f64_sse(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { + const SIMD_WIDTH: usize = 2; + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = _mm_loadu_pd(lhs.as_ptr().add(i)); + let b = _mm_loadu_pd(rhs.as_ptr().add(i)); + let result = _mm_div_pd(a, b); + _mm_storeu_pd(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] / rhs[i]; + } + + Ok(()) +} + +// ARM NEON f64 implementations +#[cfg(target_arch = "aarch64")] +#[target_feature(enable = "neon")] +pub(crate) unsafe fn simd_add_f64_neon(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { + const SIMD_WIDTH: usize = 2; // NEON processes 2 f64s at once + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = vld1q_f64(lhs.as_ptr().add(i)); + let b = vld1q_f64(rhs.as_ptr().add(i)); + let result = vaddq_f64(a, b); + vst1q_f64(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] + rhs[i]; + } + + Ok(()) +} + +#[cfg(target_arch = "aarch64")] +#[target_feature(enable = "neon")] +pub(crate) unsafe fn simd_sub_f64_neon(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { + const SIMD_WIDTH: usize = 2; + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = vld1q_f64(lhs.as_ptr().add(i)); + let b = vld1q_f64(rhs.as_ptr().add(i)); + let result = vsubq_f64(a, b); + vst1q_f64(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] - rhs[i]; + } + + Ok(()) +} + +#[cfg(target_arch = "aarch64")] +#[target_feature(enable = "neon")] +pub(crate) unsafe fn simd_mul_f64_neon(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { + const SIMD_WIDTH: usize = 2; + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = vld1q_f64(lhs.as_ptr().add(i)); + let b = vld1q_f64(rhs.as_ptr().add(i)); + let result = vmulq_f64(a, b); + vst1q_f64(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] * rhs[i]; + } + + Ok(()) +} + +#[cfg(target_arch = "aarch64")] +#[target_feature(enable = "neon")] +pub(crate) unsafe fn simd_div_f64_neon(lhs: &[f64], rhs: &[f64], output: &mut [f64]) -> Result<()> { + const SIMD_WIDTH: usize = 2; + + let len = lhs.len(); + let simd_len = len - (len % SIMD_WIDTH); + + for i in (0..simd_len).step_by(SIMD_WIDTH) { + unsafe { + let a = vld1q_f64(lhs.as_ptr().add(i)); + let b = vld1q_f64(rhs.as_ptr().add(i)); + let result = vdivq_f64(a, b); + vst1q_f64(output.as_mut_ptr().add(i), result); + } + } + + for i in simd_len..len { + output[i] = lhs[i] / rhs[i]; + } + + Ok(()) +} diff --git a/engine/src/optim/optimizer.rs b/engine/src/optim/optimizer.rs index d2eee126..442e6b38 100644 --- a/engine/src/optim/optimizer.rs +++ b/engine/src/optim/optimizer.rs @@ -51,9 +51,10 @@ impl ParameterGroup { } /// Gradient clipping configuration -#[derive(Debug, Clone)] +#[derive(Debug, Clone, Default)] pub enum GradientClipping { /// No gradient clipping + #[default] None, /// Clip gradients by norm ByNorm { max_norm: f64 }, @@ -61,12 +62,6 @@ pub enum GradientClipping { ByValue { min_value: f64, max_value: f64 }, } -impl Default for GradientClipping { - fn default() -> Self { - Self::None - } -} - /// Learning rate scheduler interface pub trait LearningRateScheduler: Send + Sync { /// Get the learning rate for the current step diff --git a/engine/src/optim/rmsprop.rs b/engine/src/optim/rmsprop.rs index e1f0c3b8..797ae8c9 100644 --- a/engine/src/optim/rmsprop.rs +++ b/engine/src/optim/rmsprop.rs @@ -307,7 +307,6 @@ impl RMSprop { let p = param.data_mut().as_f64_slice_mut().unwrap(); let g = grad.data().as_f64_slice().unwrap(); let sq = square_avg.data_mut().as_f64_slice_mut().unwrap(); - let lr = lr; let momentum = self.momentum; match (momentum_buffer_opt, grad_avg_opt) { (Some(mb), Some(ga)) => { diff --git a/engine/src/optim/tests.rs b/engine/src/optim/tests.rs index 1c86ba7f..28b5e411 100644 --- a/engine/src/optim/tests.rs +++ b/engine/src/optim/tests.rs @@ -5,6 +5,7 @@ // LICENSE file in the root directory of this source tree. #[cfg(test)] +#[allow(clippy::module_inception)] // file is already the `tests` module of `optim` mod tests { use crate::optim::optimizer::LearningRateScheduler; use crate::optim::{ diff --git a/engine/src/plugins.rs b/engine/src/plugins.rs index c503cd17..5409c85f 100644 --- a/engine/src/plugins.rs +++ b/engine/src/plugins.rs @@ -136,13 +136,13 @@ impl PluginManager { ))); } - if let Some(max_version) = &plugin_info.max_minitensor_version { - if current_version > *max_version { - return Err(MinitensorError::version_mismatch(format!( - "Plugin '{}' requires minitensor <= {}, but current version is {}", - plugin_info.name, max_version, current_version - ))); - } + if let Some(max_version) = &plugin_info.max_minitensor_version + && current_version > *max_version + { + return Err(MinitensorError::version_mismatch(format!( + "Plugin '{}' requires minitensor <= {}, but current version is {}", + plugin_info.name, max_version, current_version + ))); } // Check for name conflicts diff --git a/engine/src/random.rs b/engine/src/random.rs index 97664b4e..8c53e7a6 100644 --- a/engine/src/random.rs +++ b/engine/src/random.rs @@ -32,7 +32,7 @@ static GLOBAL_RNG: Lazy> = Lazy::new(|| { #[inline] pub fn with_rng(f: impl FnOnce(&mut StdRng) -> T) -> T { let mut guard = GLOBAL_RNG.lock(); - f(&mut *guard) + f(&mut guard) } /// Seed the global RNG with the provided value. diff --git a/engine/src/tensor/data.rs b/engine/src/tensor/data.rs index 0597f753..61d58b30 100644 --- a/engine/src/tensor/data.rs +++ b/engine/src/tensor/data.rs @@ -8,17 +8,36 @@ use crate::{ device::Device, memory::global_allocate, memory::global_deallocate, tensor::dtype::DataType, }; use rayon::prelude::*; -use std::sync::atomic::{AtomicUsize, Ordering}; - -/// Tensor data storage with reference counting -#[derive(Debug)] +use std::cell::UnsafeCell; + +/// Tensor data storage. +/// +/// Sharing is managed exclusively through `Arc`; the struct itself +/// owns its buffer and frees it on drop. +/// +/// The buffer lives in an [`UnsafeCell`] because one mutation path is +/// deliberately allowed through a shared reference: in-place parameter +/// updates, which must stay visible through every `Arc` handle to the +/// parameter (see [`DataMut`]). All other access goes through ordinary +/// `&self`/`&mut self` methods. pub struct TensorData { /// Raw data buffer - buffer: TensorBuffer, + buffer: UnsafeCell, /// Memory layout information layout: MemoryLayout, - /// Reference count for memory management - ref_count: AtomicUsize, +} + +impl std::fmt::Debug for TensorData { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let kind = match self.buffer_ref() { + TensorBuffer::Owned(_) => "owned", + TensorBuffer::Raw { .. } => "raw", + }; + f.debug_struct("TensorData") + .field("buffer", &kind) + .field("layout", &self.layout) + .finish() + } } /// Buffer storage for tensor data @@ -47,7 +66,97 @@ pub struct MemoryLayout { pub device: Device, } +/// Generates the typed slice accessors for one dtype: +/// - `$shared(&self)`: shared read view; +/// - `$exclusive(&mut self)`: exclusive write view; +/// - `$unchecked(&self)`: write view through a shared reference, used only by +/// [`DataMut`] for in-place parameter updates (see its safety contract). +/// +/// Empty tensors hand out well-aligned dangling pointers so callers can treat +/// every dtype uniformly. All accessors return `None` for dtype mismatches or +/// non-CPU storage. +macro_rules! typed_slice_accessors { + ($ty:ty, $variant:ident, $shared:ident, $exclusive:ident, $unchecked:ident) => { + #[inline(always)] + pub fn $shared(&self) -> Option<&[$ty]> { + if self.layout.dtype != DataType::$variant || !self.layout.device.is_cpu() { + return None; + } + let ptr = if self.layout.numel == 0 { + std::ptr::NonNull::<$ty>::dangling().as_ptr() + } else { + self.as_ptr() as *const $ty + }; + Some(unsafe { std::slice::from_raw_parts(ptr, self.layout.numel) }) + } + + #[inline(always)] + pub fn $exclusive(&mut self) -> Option<&mut [$ty]> { + if self.layout.dtype != DataType::$variant || !self.layout.device.is_cpu() { + return None; + } + let ptr = if self.layout.numel == 0 { + std::ptr::NonNull::<$ty>::dangling().as_ptr() + } else { + self.as_mut_ptr() as *mut $ty + }; + Some(unsafe { std::slice::from_raw_parts_mut(ptr, self.layout.numel) }) + } + + /// # Safety + /// See [`TensorData::data_ptr_shared`]: the returned slice must not + /// overlap the lifetime of any other reference to this tensor's + /// element data, and access must be externally synchronized. + // `&self -> &mut` is the point of this accessor: it is the + // UnsafeCell-backed interior-mutability path for in-place parameter + // updates, with the aliasing contract stated above. + #[allow(clippy::mut_from_ref)] + #[inline(always)] + pub(crate) unsafe fn $unchecked(&self) -> Option<&mut [$ty]> { + if self.layout.dtype != DataType::$variant || !self.layout.device.is_cpu() { + return None; + } + let ptr = if self.layout.numel == 0 { + std::ptr::NonNull::<$ty>::dangling().as_ptr() + } else { + unsafe { self.data_ptr_shared() as *mut $ty } + }; + Some(unsafe { std::slice::from_raw_parts_mut(ptr, self.layout.numel) }) + } + }; +} + impl TensorData { + /// Shared view of the buffer. + /// + /// Sound because nothing ever mutates the `TensorBuffer` value itself + /// (enum tag, `Vec` header, raw-pointer fields) through a shared + /// reference — the shared-mutation path in [`DataMut`] only writes the + /// *element bytes* the buffer points to. + #[inline(always)] + fn buffer_ref(&self) -> &TensorBuffer { + unsafe { &*self.buffer.get() } + } + + /// Pointer to the element bytes with write provenance, obtained through + /// the `UnsafeCell` from a shared reference. + /// + /// # Safety + /// Callers must guarantee that writes through the returned pointer do not + /// overlap the lifetime of any other reference to this tensor's element + /// data, and that access is externally synchronized (in practice: the + /// Python GIL serializes optimizer steps, and gradient functions only + /// read their saved operands before the step runs). + #[inline(always)] + unsafe fn data_ptr_shared(&self) -> *mut u8 { + // A transient `&mut TensorBuffer` scoped to this expression is the + // sanctioned way to derive a writable pointer from an `UnsafeCell`. + match unsafe { &mut *self.buffer.get() } { + TensorBuffer::Owned(vec) => vec.as_mut_ptr(), + TensorBuffer::Raw { ptr, .. } => *ptr, + } + } + #[inline(always)] fn validate_from_vec_type(dtype: DataType) { let type_id = std::any::TypeId::of::(); @@ -62,34 +171,33 @@ impl TensorData { assert!(matches, "dtype/type mismatch in TensorData::from_vec"); } - #[inline(always)] - fn owned_buffer_with_len(size_bytes: usize) -> Vec { - let mut vec = Vec::with_capacity(size_bytes); - unsafe { - vec.set_len(size_bytes); - } - vec - } - + /// Allocate a zero-initialized byte buffer (`alloc_zeroed`). #[inline(always)] fn owned_zeroed_buffer(size_bytes: usize) -> Vec { - let mut vec = Self::owned_buffer_with_len(size_bytes); - unsafe { - std::ptr::write_bytes(vec.as_mut_ptr(), 0, size_bytes); - } - vec + vec![0u8; size_bytes] } + /// Allocate a CPU buffer for an operation output. + /// + /// This is always zero-initialized. The kernels immediately overwrite it, + /// so the zeros are never observed — but allocating it uninitialized and + /// handing it out as `&mut [f32]`/`&mut [i64]`/… (which every kernel does) + /// is undefined behavior: a reference of type `&mut T` must point to a + /// *valid* `T`, and a float/integer read from uninitialized memory is not + /// a valid value even though every bit pattern is representable. Per the + /// project's stated priority order (correctness before performance), + /// soundness wins here. + /// + /// Measured cost (interleaved A/B, release, this machine): the extra + /// `memset` adds ~25–35% to the pure element-wise microbenchmark + /// (`add`+`sum` over a 2000×2000 tensor — bound by output write traffic, + /// which the zeroing roughly increases from 3N to 4N), and is within + /// noise on the matmul-dominated training-step benchmark. The zero-cost + /// alternative is `MaybeUninit`-typed writes threaded through every + /// kernel; it is tracked in docs/architecture_review.md as future work. #[inline(always)] - fn owned_buffer_for_dtype(size_bytes: usize, dtype: DataType) -> Vec { - // Bool has strict valid bit-pattern requirements (0 or 1). Creating - // `&mut [bool]` over arbitrary uninitialized bytes is UB, so we keep - // bool storage initialized even for "uninitialized" constructors. - if dtype == DataType::Bool { - Self::owned_zeroed_buffer(size_bytes) - } else { - Self::owned_buffer_with_len(size_bytes) - } + fn owned_buffer_for_dtype(size_bytes: usize, _dtype: DataType) -> Vec { + Self::owned_zeroed_buffer(size_bytes) } #[inline(always)] @@ -143,14 +251,15 @@ impl TensorData { .checked_mul(dtype.size_bytes()) .expect("tensor size overflow"); let mut data = Self { - buffer: TensorBuffer::Owned(Self::owned_buffer_for_dtype(size_bytes, dtype)), + buffer: UnsafeCell::new(TensorBuffer::Owned(Self::owned_buffer_for_dtype( + size_bytes, dtype, + ))), layout: MemoryLayout { dtype, numel, is_contiguous: true, device, }, - ref_count: AtomicUsize::new(1), }; data.fill_with_ones(); @@ -181,9 +290,6 @@ impl TensorData { .expect("tensor size overflow"); let buffer = if device.is_cpu() { - // Use an uninitialized Vec and explicitly zero it. This avoids the - // double-initialization that `vec![0u8; size_bytes]` may perform on - // some platforms and lets the optimizer emit a single `memset`. TensorBuffer::Owned(Self::owned_zeroed_buffer(size_bytes)) } else { // Use custom allocator for GPU @@ -211,18 +317,23 @@ impl TensorData { }; Self { - buffer, + buffer: UnsafeCell::new(buffer), layout: MemoryLayout { dtype, numel, is_contiguous: true, device: actual_device, }, - ref_count: AtomicUsize::new(1), } } - /// Create new tensor data with uninitialized contents on specified device + /// Create new tensor data for use as an operation output buffer. + /// + /// Historically this handed out genuinely uninitialized memory; CPU + /// buffers are now zero-initialized via `alloc_zeroed` (see + /// [`Self::owned_zeroed_buffer`]), which keeps the fast allocation path + /// while making accidental reads defined. The name is kept for API + /// stability; callers should still treat the contents as unspecified. #[inline(always)] pub fn uninitialized_on_device(numel: usize, dtype: DataType, device: Device) -> Self { let size_bytes = numel @@ -230,7 +341,6 @@ impl TensorData { .expect("tensor size overflow"); let buffer = if device.is_cpu() { - // Allocate vector without initializing memory for maximum performance (bool remains zero-initialized for validity) TensorBuffer::Owned(Self::owned_buffer_for_dtype(size_bytes, dtype)) } else { match global_allocate(size_bytes, device) { @@ -251,14 +361,13 @@ impl TensorData { }; Self { - buffer, + buffer: UnsafeCell::new(buffer), layout: MemoryLayout { dtype, numel, is_contiguous: true, device: actual_device, }, - ref_count: AtomicUsize::new(1), } } @@ -266,14 +375,13 @@ impl TensorData { #[inline(always)] pub fn from_bytes(buffer: Vec, dtype: DataType, numel: usize) -> Self { Self { - buffer: TensorBuffer::Owned(buffer), + buffer: UnsafeCell::new(TensorBuffer::Owned(buffer)), layout: MemoryLayout { dtype, numel, is_contiguous: true, device: Device::cpu(), }, - ref_count: AtomicUsize::new(1), } } @@ -350,14 +458,13 @@ impl TensorData { }; Self { - buffer, + buffer: UnsafeCell::new(buffer), layout: MemoryLayout { dtype, numel, is_contiguous: true, device: actual_device, }, - ref_count: AtomicUsize::new(1), } } @@ -401,21 +508,20 @@ impl TensorData { device: Device, ) -> Self { Self { - buffer: TensorBuffer::Raw { ptr, size, device }, + buffer: UnsafeCell::new(TensorBuffer::Raw { ptr, size, device }), layout: MemoryLayout { dtype, numel, is_contiguous: true, device, }, - ref_count: AtomicUsize::new(1), } } /// Get the raw buffer as a slice (CPU only) #[inline(always)] pub fn as_bytes(&self) -> Option<&[u8]> { - match &self.buffer { + match self.buffer_ref() { TensorBuffer::Owned(vec) => Some(vec.as_slice()), TensorBuffer::Raw { ptr, size, device } => { if device.is_cpu() { @@ -434,7 +540,7 @@ impl TensorData { /// Get the raw buffer as a mutable slice (CPU only) #[inline(always)] pub fn as_bytes_mut(&mut self) -> Option<&mut [u8]> { - match &mut self.buffer { + match self.buffer.get_mut() { TensorBuffer::Owned(vec) => Some(vec.as_mut_slice()), TensorBuffer::Raw { ptr, size, device } => { if device.is_cpu() { @@ -453,7 +559,7 @@ impl TensorData { /// Get the raw pointer (for GPU operations) #[inline(always)] pub fn as_ptr(&self) -> *const u8 { - match &self.buffer { + match self.buffer_ref() { TensorBuffer::Owned(vec) => vec.as_ptr(), TensorBuffer::Raw { ptr, .. } => *ptr, } @@ -462,7 +568,7 @@ impl TensorData { /// Get the mutable raw pointer (for GPU operations) #[inline(always)] pub fn as_mut_ptr(&mut self) -> *mut u8 { - match &mut self.buffer { + match self.buffer.get_mut() { TensorBuffer::Owned(vec) => vec.as_mut_ptr(), TensorBuffer::Raw { ptr, .. } => *ptr, } @@ -501,7 +607,7 @@ impl TensorData { /// Get the size in bytes #[inline(always)] pub fn size_bytes(&self) -> usize { - match &self.buffer { + match self.buffer_ref() { TensorBuffer::Owned(vec) => vec.len(), TensorBuffer::Raw { size, .. } => *size, } @@ -519,44 +625,15 @@ impl TensorData { self.layout.is_contiguous } - /// Increment reference count - #[inline(always)] - pub fn inc_ref(&self) { - self.ref_count.fetch_add(1, Ordering::Relaxed); - } - - /// Decrement reference count and return new count - #[inline(always)] - pub fn dec_ref(&self) -> usize { - self.ref_count.fetch_sub(1, Ordering::Relaxed) - 1 - } - - /// Get current reference count - #[inline(always)] - pub fn ref_count(&self) -> usize { - self.ref_count.load(Ordering::Relaxed) - } - /// Create a copy of the tensor data pub fn clone_data(&self) -> Self { - let new_buffer = match &self.buffer { - TensorBuffer::Owned(vec) => { - let mut out = Vec::with_capacity(vec.len()); - unsafe { - out.set_len(vec.len()); - std::ptr::copy_nonoverlapping(vec.as_ptr(), out.as_mut_ptr(), vec.len()); - } - TensorBuffer::Owned(out) - } + let new_buffer = match self.buffer_ref() { + TensorBuffer::Owned(vec) => TensorBuffer::Owned(vec.clone()), TensorBuffer::Raw { ptr, size, device } => { if device.is_cpu() { // Raw CPU pointer: copy into a Vec for safety - let mut out = Vec::with_capacity(*size); - unsafe { - out.set_len(*size); - std::ptr::copy_nonoverlapping(*ptr, out.as_mut_ptr(), *size); - } - TensorBuffer::Owned(out) + let bytes = unsafe { std::slice::from_raw_parts(*ptr, *size) }.to_vec(); + TensorBuffer::Owned(bytes) } else { // For GPU, allocate new memory and copy match global_allocate(*size, *device) { @@ -571,13 +648,9 @@ impl TensorData { } } Err(_) => { - // Fallback to CPU buffer to avoid allocation failure - let mut out = Vec::with_capacity(*size); - unsafe { - out.set_len(*size); - std::ptr::write_bytes(out.as_mut_ptr(), 0, *size); - } - TensorBuffer::Owned(out) + // Fallback to a zeroed CPU buffer to avoid + // failing the clone outright. + TensorBuffer::Owned(Self::owned_zeroed_buffer(*size)) } } } @@ -585,181 +658,95 @@ impl TensorData { }; Self { - buffer: new_buffer, + buffer: UnsafeCell::new(new_buffer), layout: self.layout.clone(), - ref_count: AtomicUsize::new(1), - } - } - - /// Get typed slice for f32 data (CPU only) - #[inline(always)] - pub fn as_f32_slice(&self) -> Option<&[f32]> { - if self.layout.dtype != DataType::Float32 || !self.layout.device.is_cpu() { - return None; - } - - let ptr = self.as_ptr() as *const f32; - let ptr = if self.layout.numel == 0 { - std::ptr::NonNull::::dangling().as_ptr() - } else { - ptr - }; - Some(unsafe { std::slice::from_raw_parts(ptr, self.layout.numel) }) - } - - /// Get mutable typed slice for f32 data (CPU only) - #[inline(always)] - pub fn as_f32_slice_mut(&mut self) -> Option<&mut [f32]> { - if self.layout.dtype != DataType::Float32 || !self.layout.device.is_cpu() { - return None; - } - - let ptr = self.as_mut_ptr() as *mut f32; - let ptr = if self.layout.numel == 0 { - std::ptr::NonNull::::dangling().as_ptr() as *mut f32 - } else { - ptr - }; - Some(unsafe { std::slice::from_raw_parts_mut(ptr, self.layout.numel) }) - } - - /// Get typed slice for f64 data (CPU only) - #[inline(always)] - pub fn as_f64_slice(&self) -> Option<&[f64]> { - if self.layout.dtype != DataType::Float64 || !self.layout.device.is_cpu() { - return None; - } - - let ptr = self.as_ptr() as *const f64; - let ptr = if self.layout.numel == 0 { - std::ptr::NonNull::::dangling().as_ptr() - } else { - ptr - }; - Some(unsafe { std::slice::from_raw_parts(ptr, self.layout.numel) }) - } - - /// Get mutable typed slice for f64 data (CPU only) - #[inline(always)] - pub fn as_f64_slice_mut(&mut self) -> Option<&mut [f64]> { - if self.layout.dtype != DataType::Float64 || !self.layout.device.is_cpu() { - return None; } - - let ptr = self.as_mut_ptr() as *mut f64; - let ptr = if self.layout.numel == 0 { - std::ptr::NonNull::::dangling().as_ptr() as *mut f64 - } else { - ptr - }; - Some(unsafe { std::slice::from_raw_parts_mut(ptr, self.layout.numel) }) } - /// Get typed slice for i32 data (CPU only) - #[inline(always)] - pub fn as_i32_slice(&self) -> Option<&[i32]> { - if self.layout.dtype != DataType::Int32 || !self.layout.device.is_cpu() { - return None; - } - - let ptr = self.as_ptr() as *const i32; - let ptr = if self.layout.numel == 0 { - std::ptr::NonNull::::dangling().as_ptr() - } else { - ptr - }; - Some(unsafe { std::slice::from_raw_parts(ptr, self.layout.numel) }) - } - - /// Get mutable typed slice for i32 data (CPU only) - #[inline(always)] - pub fn as_i32_slice_mut(&mut self) -> Option<&mut [i32]> { - if self.layout.dtype != DataType::Int32 || !self.layout.device.is_cpu() { - return None; - } - - let ptr = self.as_mut_ptr() as *mut i32; - let ptr = if self.layout.numel == 0 { - std::ptr::NonNull::::dangling().as_ptr() as *mut i32 - } else { - ptr - }; - Some(unsafe { std::slice::from_raw_parts_mut(ptr, self.layout.numel) }) - } - - /// Get typed slice for i64 data (CPU only) - #[inline(always)] - pub fn as_i64_slice(&self) -> Option<&[i64]> { - if self.layout.dtype != DataType::Int64 || !self.layout.device.is_cpu() { - return None; - } - - let ptr = self.as_ptr() as *const i64; - let ptr = if self.layout.numel == 0 { - std::ptr::NonNull::::dangling().as_ptr() - } else { - ptr - }; - Some(unsafe { std::slice::from_raw_parts(ptr, self.layout.numel) }) - } - - /// Get mutable typed slice for i64 data (CPU only) - #[inline(always)] - pub fn as_i64_slice_mut(&mut self) -> Option<&mut [i64]> { - if self.layout.dtype != DataType::Int64 || !self.layout.device.is_cpu() { - return None; - } - - let ptr = self.as_mut_ptr() as *mut i64; - let ptr = if self.layout.numel == 0 { - std::ptr::NonNull::::dangling().as_ptr() as *mut i64 - } else { - ptr - }; - Some(unsafe { std::slice::from_raw_parts_mut(ptr, self.layout.numel) }) - } - - /// Get typed slice for bool data (CPU only) - #[inline(always)] - pub fn as_bool_slice(&self) -> Option<&[bool]> { - if self.layout.dtype != DataType::Bool || !self.layout.device.is_cpu() { - return None; - } + typed_slice_accessors!( + f32, + Float32, + as_f32_slice, + as_f32_slice_mut, + as_f32_slice_mut_unchecked + ); + typed_slice_accessors!( + f64, + Float64, + as_f64_slice, + as_f64_slice_mut, + as_f64_slice_mut_unchecked + ); + typed_slice_accessors!( + i32, + Int32, + as_i32_slice, + as_i32_slice_mut, + as_i32_slice_mut_unchecked + ); + typed_slice_accessors!( + i64, + Int64, + as_i64_slice, + as_i64_slice_mut, + as_i64_slice_mut_unchecked + ); + typed_slice_accessors!( + bool, + Bool, + as_bool_slice, + as_bool_slice_mut, + as_bool_slice_mut_unchecked + ); +} - let ptr = self.as_ptr() as *const bool; - let ptr = if self.layout.numel == 0 { - std::ptr::NonNull::::dangling().as_ptr() - } else { - ptr - }; - Some(unsafe { std::slice::from_raw_parts(ptr, self.layout.numel) }) - } +/// Mutable access to a tensor's storage, produced by `Tensor::data_mut`. +/// +/// The `Unique` variant is ordinary exclusive access (storage uniquely owned, +/// possibly after copy-on-write). The `Shared` variant is the one deliberate +/// exception in the crate: in-place updates of leaf parameters whose storage +/// is shared across `Arc` handles, so the update stays visible through every +/// handle (PyTorch in-place parameter semantics). Its safety contract is +/// documented on [`TensorData::data_ptr_shared`] and upheld by the callers: +/// optimizer steps run GIL-serialized, after backward has finished reading +/// saved operands. +pub enum DataMut<'a> { + Unique(&'a mut TensorData), + Shared(&'a TensorData), +} - /// Get mutable typed slice for bool data (CPU only) - #[inline(always)] - pub fn as_bool_slice_mut(&mut self) -> Option<&mut [bool]> { - if self.layout.dtype != DataType::Bool || !self.layout.device.is_cpu() { - return None; +macro_rules! data_mut_accessor { + ($name:ident, $unchecked:ident, $ty:ty) => { + /// Consumes the access token and returns a slice borrowing from the + /// underlying tensor, so `t.data_mut().as_…_slice_mut()` keeps the + /// slice alive for the caller's borrow of `t`. + #[inline(always)] + pub fn $name(self) -> Option<&'a mut [$ty]> { + match self { + DataMut::Unique(data) => data.$name(), + // SAFETY: contract documented on `DataMut` and + // `TensorData::data_ptr_shared`. + DataMut::Shared(data) => unsafe { data.$unchecked() }, + } } + }; +} - let ptr = self.as_mut_ptr() as *mut bool; - let ptr = if self.layout.numel == 0 { - std::ptr::NonNull::::dangling().as_ptr() as *mut bool - } else { - ptr - }; - Some(unsafe { std::slice::from_raw_parts_mut(ptr, self.layout.numel) }) - } +impl<'a> DataMut<'a> { + data_mut_accessor!(as_f32_slice_mut, as_f32_slice_mut_unchecked, f32); + data_mut_accessor!(as_f64_slice_mut, as_f64_slice_mut_unchecked, f64); + data_mut_accessor!(as_i32_slice_mut, as_i32_slice_mut_unchecked, i32); + data_mut_accessor!(as_i64_slice_mut, as_i64_slice_mut_unchecked, i64); + data_mut_accessor!(as_bool_slice_mut, as_bool_slice_mut_unchecked, bool); } impl Drop for TensorData { fn drop(&mut self) { - // Only deallocate if this is the last reference - if self.ref_count.load(Ordering::Relaxed) == 1 { - if let TensorBuffer::Raw { ptr, size, device } = &self.buffer { - // Deallocate GPU memory - let _ = global_deallocate(*ptr, *size, *device); - } + // `TensorData` is shared via `Arc`, so `drop` runs exactly once, when + // the last reference goes away. Owned buffers free themselves; raw + // device buffers are returned to the allocator here. + if let TensorBuffer::Raw { ptr, size, device } = self.buffer.get_mut() { + let _ = global_deallocate(*ptr, *size, *device); } } } @@ -778,7 +765,6 @@ mod tests { assert_eq!(data.dtype(), DataType::Float32); assert_eq!(data.size_bytes(), 40); // 10 * 4 bytes assert!(data.is_contiguous()); - assert_eq!(data.ref_count(), 1); } #[test] @@ -815,16 +801,15 @@ mod tests { } #[test] - fn test_reference_counting() { - let data = TensorData::zeros(5, DataType::Float32); - assert_eq!(data.ref_count(), 1); - - data.inc_ref(); - assert_eq!(data.ref_count(), 2); - - let new_count = data.dec_ref(); - assert_eq!(new_count, 1); - assert_eq!(data.ref_count(), 1); + fn test_arc_sharing() { + use std::sync::Arc; + + let data = Arc::new(TensorData::zeros(5, DataType::Float32)); + let shared = Arc::clone(&data); + assert_eq!(Arc::strong_count(&data), 2); + drop(shared); + assert_eq!(Arc::strong_count(&data), 1); + assert_eq!(data.numel(), 5); } #[test] diff --git a/engine/src/tensor/mod.rs b/engine/src/tensor/mod.rs index 634c385c..1d86e9fb 100644 --- a/engine/src/tensor/mod.rs +++ b/engine/src/tensor/mod.rs @@ -4,11 +4,18 @@ // This source code is licensed under the Apache-style license found in the // LICENSE file in the root directory of this source tree. +//! Tensor type, storage, shapes, and dtype definitions. +//! +//! `core` declares the `Tensor` struct and hosts the method impls (split +//! across its child modules so they retain access to the private fields); +//! everything public is re-exported here so callers keep using +//! `crate::tensor::X`. + pub mod data; pub mod dtype; pub mod shape; -include!("mod/core.rs"); -include!("mod/autograd.rs"); -include!("mod/ops.rs"); -include!("mod/indexing.rs"); -include!("mod/utils.rs"); + +#[path = "mod/core.rs"] +mod core; + +pub use self::core::*; diff --git a/engine/src/tensor/mod/autograd.rs b/engine/src/tensor/mod/autograd.rs index 2539e7a9..08ec59d6 100644 --- a/engine/src/tensor/mod/autograd.rs +++ b/engine/src/tensor/mod/autograd.rs @@ -1,593 +1,596 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -impl Tensor { - /// Unary negation - #[inline(always)] - pub fn neg(&self) -> Result { - use crate::operations::arithmetic::neg; - neg(self) - } - - /// Add two tensors element-wise - #[inline(always)] - pub fn add(&self, other: &Tensor) -> Result { - use crate::operations::arithmetic::add; - add(self, other) - } - - /// Element-wise maximum - #[inline(always)] - pub fn maximum(&self, other: &Tensor) -> Result { - use crate::operations::minmax::maximum; - maximum(self, other) - } - - /// Element-wise minimum - #[inline(always)] - pub fn minimum(&self, other: &Tensor) -> Result { - use crate::operations::minmax::minimum; - minimum(self, other) - } - - /// Select elements from self or other based on a boolean condition tensor - #[inline(always)] - pub fn where_select(&self, condition: &Tensor, other: &Tensor) -> Result { - use crate::operations::selection::where_op; - where_op(condition, self, other) - } - - /// Fill elements specified by `mask` with values from `value`. - #[inline(always)] - pub fn masked_fill(&self, mask: &Tensor, value: &Tensor) -> Result { - crate::operations::selection::masked_fill(self, mask, value) - } - - /// Fill elements specified by `mask` with a scalar. - #[inline(always)] - pub fn masked_fill_scalar(&self, mask: &Tensor, value: f64) -> Result { - crate::operations::selection::masked_fill_scalar(self, mask, value) - } - - /// Dot product between two 1D tensors - #[inline(always)] - pub fn dot(&self, other: &Tensor) -> Result { - crate::operations::linalg::dot(self, other) - } - - /// Matrix multiplication - #[inline(always)] - pub fn matmul(&self, other: &Tensor) -> Result { - use crate::operations::linalg::matmul; - matmul(self, other) - } - - /// Batched matrix multiplication specialised for 3D tensors - #[inline(always)] - pub fn bmm(&self, other: &Tensor) -> Result { - use crate::operations::linalg::bmm; - bmm(self, other) - } - - /// Upper triangular part of the tensor's last two dimensions - #[inline(always)] - pub fn triu(&self, diagonal: i64) -> Result { - use crate::operations::linalg::triu; - triu(self, diagonal) - } - - /// Lower triangular part of the tensor's last two dimensions - #[inline(always)] - pub fn tril(&self, diagonal: i64) -> Result { - use crate::operations::linalg::tril; - tril(self, diagonal) - } - - /// Extract a diagonal along two dimensions. - #[inline(always)] - pub fn diagonal(&self, offset: isize, dim1: isize, dim2: isize) -> Result { - use crate::operations::linalg::diagonal; - diagonal(self, offset, dim1, dim2) - } - - /// Sum the diagonal elements along two dimensions. - #[inline(always)] - pub fn trace(&self, offset: isize, dim1: isize, dim2: isize) -> Result { - use crate::operations::linalg::trace; - trace(self, offset, dim1, dim2) - } - - /// Sum reduction - #[inline(always)] - pub fn sum(&self, dim: Option>, keepdim: bool) -> Result { - use crate::operations::reduction::sum; - sum(self, dim, keepdim) - } - - /// NaN-aware sum reduction - #[inline(always)] - pub fn nansum(&self, dim: Option>, keepdim: bool) -> Result { - use crate::operations::reduction::nansum; - nansum(self, dim, keepdim) - } - - /// Log-sum-exp reduction - #[inline(always)] - pub fn logsumexp(&self, dim: Option>, keepdim: bool) -> Result { - use crate::operations::reduction::logsumexp; - logsumexp(self, dim, keepdim) - } - - /// Product reduction - #[inline(always)] - pub fn prod(&self, dim: Option>, keepdim: bool) -> Result { - use crate::operations::reduction::prod; - prod(self, dim, keepdim) - } - - /// Mean reduction - #[inline(always)] - pub fn mean(&self, dim: Option>, keepdim: bool) -> Result { - use crate::operations::reduction::mean; - mean(self, dim, keepdim) - } - - /// NaN-aware mean reduction - #[inline(always)] - pub fn nanmean(&self, dim: Option>, keepdim: bool) -> Result { - use crate::operations::reduction::nanmean; - nanmean(self, dim, keepdim) - } - - /// Logical all reduction - #[inline(always)] - pub fn all(&self, dim: Option, keepdim: bool) -> Result { - use crate::operations::reduction::all; - all(self, dim, keepdim) - } - - /// Logical any reduction - #[inline(always)] - pub fn any(&self, dim: Option, keepdim: bool) -> Result { - use crate::operations::reduction::any; - any(self, dim, keepdim) - } - - /// Cumulative sum along a dimension - #[inline(always)] - pub fn cumsum(&self, dim: isize) -> Result { - use crate::operations::reduction::cumsum; - cumsum(self, dim) - } - - /// Cumulative product along a dimension - #[inline(always)] - pub fn cumprod(&self, dim: isize) -> Result { - use crate::operations::reduction::cumprod; - cumprod(self, dim) - } - - /// Maximum value - #[inline(always)] - pub fn max(&self, dim: Option, keepdim: bool) -> Result { - use crate::operations::reduction::max; - max(self, dim, keepdim) - } - - /// NaN-aware maximum value - #[inline(always)] - pub fn nanmax(&self, dim: Option, keepdim: bool) -> Result { - use crate::operations::reduction::nanmax; - nanmax(self, dim, keepdim) - } - - /// Minimum value - #[inline(always)] - pub fn min(&self, dim: Option, keepdim: bool) -> Result { - use crate::operations::reduction::min; - min(self, dim, keepdim) - } - - /// NaN-aware minimum value - #[inline(always)] - pub fn nanmin(&self, dim: Option, keepdim: bool) -> Result { - use crate::operations::reduction::nanmin; - nanmin(self, dim, keepdim) - } - - /// Argument of maximum value - #[inline(always)] - pub fn argmax(&self, dim: Option, keepdim: bool) -> Result { - use crate::operations::reduction::argmax; - argmax(self, dim, keepdim) - } - - /// Argument of minimum value - #[inline(always)] - pub fn argmin(&self, dim: Option, keepdim: bool) -> Result { - use crate::operations::reduction::argmin; - argmin(self, dim, keepdim) - } - - /// Maximum values and their indices along a dimension - #[inline(always)] - pub fn max_with_indices(&self, dim: isize, keepdim: bool) -> Result<(Self, Self)> { - use crate::operations::reduction::max_with_indices; - max_with_indices(self, dim, keepdim) - } - - /// NaN-aware maximum values and their indices along a dimension - #[inline(always)] - pub fn nanmax_with_indices(&self, dim: isize, keepdim: bool) -> Result<(Self, Self)> { - use crate::operations::reduction::nanmax_with_indices; - nanmax_with_indices(self, dim, keepdim) - } - - /// Minimum values and their indices along a dimension - #[inline(always)] - pub fn min_with_indices(&self, dim: isize, keepdim: bool) -> Result<(Self, Self)> { - use crate::operations::reduction::min_with_indices; - min_with_indices(self, dim, keepdim) - } - - /// NaN-aware minimum values and their indices along a dimension - #[inline(always)] - pub fn nanmin_with_indices(&self, dim: isize, keepdim: bool) -> Result<(Self, Self)> { - use crate::operations::reduction::nanmin_with_indices; - nanmin_with_indices(self, dim, keepdim) - } - - /// Median value (optionally along a dimension) - #[inline(always)] - pub fn median(&self, dim: Option, keepdim: bool) -> Result<(Self, Option)> { - use crate::operations::reduction::median; - median(self, dim, keepdim) - } - - /// Quantile reduction with configurable interpolation - #[inline(always)] - pub fn quantile( - &self, - q: f64, - dim: Option, - keepdim: bool, - interpolation: QuantileInterpolation, - ) -> Result { - use crate::operations::reduction::quantile; - quantile(self, q, dim, keepdim, interpolation) - } - - /// Quantile reduction that ignores NaN values - #[inline(always)] - pub fn nanquantile( - &self, - q: f64, - dim: Option, - keepdim: bool, - interpolation: QuantileInterpolation, - ) -> Result { - use crate::operations::reduction::nanquantile; - nanquantile(self, q, dim, keepdim, interpolation) - } - - /// Median reduction that ignores NaN values - #[inline(always)] - pub fn nanmedian(&self, dim: Option, keepdim: bool) -> Result { - use crate::operations::reduction::nanmedian; - nanmedian(self, dim, keepdim) - } - - /// Batched quantile reduction for multiple probabilities at once - #[inline(always)] - pub fn quantiles( - &self, - qs: &[f64], - dim: Option, - keepdim: bool, - interpolation: QuantileInterpolation, - ) -> Result { - use crate::operations::reduction::quantiles; - quantiles(self, qs, dim, keepdim, interpolation) - } - - /// Batched quantile reduction that ignores NaN values - #[inline(always)] - pub fn nanquantiles( - &self, - qs: &[f64], - dim: Option, - keepdim: bool, - interpolation: QuantileInterpolation, - ) -> Result { - use crate::operations::reduction::nanquantiles; - nanquantiles(self, qs, dim, keepdim, interpolation) - } - - /// Top-k values and indices along a dimension - #[inline(always)] - pub fn topk( - &self, - k: usize, - dim: Option, - largest: bool, - sorted: bool, - ) -> Result<(Self, Self)> { - use crate::operations::reduction::topk; - topk(self, k, dim, largest, sorted) - } - - /// Sort tensor values along a dimension - #[inline(always)] - pub fn sort(&self, dim: Option, descending: bool, stable: bool) -> Result<(Self, Self)> { - use crate::operations::reduction::sort; - sort(self, dim, descending, stable) - } - - /// Indices that would sort the tensor along a dimension - #[inline(always)] - pub fn argsort(&self, dim: Option, descending: bool, stable: bool) -> Result { - use crate::operations::reduction::argsort; - argsort(self, dim, descending, stable) - } - - /// Element-wise equality comparison - #[inline(always)] - pub fn eq(&self, other: &Tensor) -> Result { - use crate::operations::comparison::eq; - eq(self, other) - } - - /// Element-wise inequality comparison - pub fn ne(&self, other: &Tensor) -> Result { - use crate::operations::comparison::ne; - ne(self, other) - } - - /// Element-wise less-than comparison - #[inline(always)] - pub fn lt(&self, other: &Tensor) -> Result { - use crate::operations::comparison::lt; - lt(self, other) - } - - /// Element-wise less-than-or-equal comparison - #[inline(always)] - pub fn le(&self, other: &Tensor) -> Result { - use crate::operations::comparison::le; - le(self, other) - } - - /// Element-wise greater-than comparison - #[inline(always)] - pub fn gt(&self, other: &Tensor) -> Result { - use crate::operations::comparison::gt; - gt(self, other) - } - - /// Element-wise greater-than-or-equal comparison - #[inline(always)] - pub fn ge(&self, other: &Tensor) -> Result { - use crate::operations::comparison::ge; - ge(self, other) - } - - /// Standard deviation - #[inline(always)] - pub fn std(&self, dim: Option>, keepdim: bool, unbiased: bool) -> Result { - use crate::operations::reduction::std; - std(self, dim, keepdim, unbiased) - } - - /// Variance - #[inline(always)] - pub fn var(&self, dim: Option>, keepdim: bool, unbiased: bool) -> Result { - use crate::operations::reduction::var; - var(self, dim, keepdim, unbiased) - } - - /// Exponential function - #[inline(always)] - pub fn exp(&self) -> Result { - use crate::operations::activation::exp; - exp(self) - } - - /// Natural logarithm - #[inline(always)] - pub fn log(&self) -> Result { - use crate::operations::activation::log; - log(self) - } - - /// log1p (log(1 + x)) - #[inline(always)] - pub fn log1p(&self) -> Result { - use crate::operations::activation::log1p; - log1p(self) - } - - /// expm1 (exp(x) - 1) - #[inline(always)] - pub fn expm1(&self) -> Result { - use crate::operations::activation::expm1; - expm1(self) - } - - /// Sine function - #[inline(always)] - pub fn sin(&self) -> Result { - use crate::operations::activation::sin; - sin(self) - } - - /// Cosine function - #[inline(always)] - pub fn cos(&self) -> Result { - use crate::operations::activation::cos; - cos(self) - } - - /// Tangent function - #[inline(always)] - pub fn tan(&self) -> Result { - use crate::operations::activation::tan; - tan(self) - } - - /// Inverse sine function - #[inline(always)] - pub fn asin(&self) -> Result { - use crate::operations::activation::asin; - asin(self) - } - - /// Inverse cosine function - #[inline(always)] - pub fn acos(&self) -> Result { - use crate::operations::activation::acos; - acos(self) - } - - /// Inverse tangent function - #[inline(always)] - pub fn atan(&self) -> Result { - use crate::operations::activation::atan; - atan(self) - } - - /// Hyperbolic sine - #[inline(always)] - pub fn sinh(&self) -> Result { - use crate::operations::activation::sinh; - sinh(self) - } - - /// Hyperbolic cosine - #[inline(always)] - pub fn cosh(&self) -> Result { - use crate::operations::activation::cosh; - cosh(self) - } - - /// Inverse hyperbolic sine - #[inline(always)] - pub fn asinh(&self) -> Result { - use crate::operations::activation::asinh; - asinh(self) - } - - /// Inverse hyperbolic cosine - #[inline(always)] - pub fn acosh(&self) -> Result { - use crate::operations::activation::acosh; - acosh(self) - } - - /// Inverse hyperbolic tangent - #[inline(always)] - pub fn atanh(&self) -> Result { - use crate::operations::activation::atanh; - atanh(self) - } - - /// Hyperbolic tangent - #[inline(always)] - pub fn tanh(&self) -> Result { - use crate::operations::activation::tanh; - tanh(self) - } - - /// Sigmoid activation - #[inline(always)] - pub fn sigmoid(&self) -> Result { - use crate::operations::activation::sigmoid; - sigmoid(self) - } - - /// Softplus activation - #[inline(always)] - pub fn softplus(&self, beta: f64, threshold: f64) -> Result { - use crate::operations::activation::softplus; - softplus(self, beta, threshold) - } - - /// GELU activation - #[inline(always)] - pub fn gelu(&self, approximate: bool) -> Result { - use crate::operations::activation::gelu; - gelu(self, approximate) - } - - /// ELU activation - #[inline(always)] - pub fn elu(&self, alpha: f64) -> Result { - use crate::operations::activation::elu; - elu(self, alpha) - } - - /// SELU activation - #[inline(always)] - pub fn selu(&self) -> Result { - use crate::operations::activation::selu; - selu(self) - } - - /// SiLU activation - #[inline(always)] - pub fn silu(&self) -> Result { - use crate::operations::activation::silu; - silu(self) - } - - /// Softsign activation - #[inline(always)] - pub fn softsign(&self) -> Result { - use crate::operations::activation::softsign; - softsign(self) - } - - /// ReLU activation - #[inline(always)] - pub fn relu(&self) -> Result { - use crate::operations::activation::relu; - relu(self) - } - - /// Hardshrink activation - #[inline(always)] - pub fn hardshrink(&self, lambd: f64) -> Result { - use crate::operations::activation::hardshrink; - hardshrink(self, lambd) - } - - /// Softmax activation - #[inline(always)] - pub fn softmax(&self, dim: Option) -> Result { - use crate::operations::activation::softmax; - softmax(self, dim) - } - - /// Log-Softmax activation - #[inline(always)] - pub fn log_softmax(&self, dim: Option) -> Result { - use crate::operations::activation::log_softmax; - log_softmax(self, dim) - } - - /// Masked Softmax activation - #[inline(always)] - pub fn masked_softmax(&self, mask: &Tensor, dim: Option) -> Result { - use crate::operations::activation::masked_softmax; - masked_softmax(self, mask, dim) - } - - /// Masked Log-Softmax activation - #[inline(always)] - pub fn masked_log_softmax(&self, mask: &Tensor, dim: Option) -> Result { - use crate::operations::activation::masked_log_softmax; - masked_log_softmax(self, mask, dim) - } -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::error::Result; + +impl Tensor { + /// Unary negation + #[inline(always)] + pub fn neg(&self) -> Result { + use crate::operations::arithmetic::neg; + neg(self) + } + + /// Add two tensors element-wise + #[inline(always)] + pub fn add(&self, other: &Tensor) -> Result { + use crate::operations::arithmetic::add; + add(self, other) + } + + /// Element-wise maximum + #[inline(always)] + pub fn maximum(&self, other: &Tensor) -> Result { + use crate::operations::minmax::maximum; + maximum(self, other) + } + + /// Element-wise minimum + #[inline(always)] + pub fn minimum(&self, other: &Tensor) -> Result { + use crate::operations::minmax::minimum; + minimum(self, other) + } + + /// Select elements from self or other based on a boolean condition tensor + #[inline(always)] + pub fn where_select(&self, condition: &Tensor, other: &Tensor) -> Result { + use crate::operations::selection::where_op; + where_op(condition, self, other) + } + + /// Fill elements specified by `mask` with values from `value`. + #[inline(always)] + pub fn masked_fill(&self, mask: &Tensor, value: &Tensor) -> Result { + crate::operations::selection::masked_fill(self, mask, value) + } + + /// Fill elements specified by `mask` with a scalar. + #[inline(always)] + pub fn masked_fill_scalar(&self, mask: &Tensor, value: f64) -> Result { + crate::operations::selection::masked_fill_scalar(self, mask, value) + } + + /// Dot product between two 1D tensors + #[inline(always)] + pub fn dot(&self, other: &Tensor) -> Result { + crate::operations::linalg::dot(self, other) + } + + /// Matrix multiplication + #[inline(always)] + pub fn matmul(&self, other: &Tensor) -> Result { + use crate::operations::linalg::matmul; + matmul(self, other) + } + + /// Batched matrix multiplication specialised for 3D tensors + #[inline(always)] + pub fn bmm(&self, other: &Tensor) -> Result { + use crate::operations::linalg::bmm; + bmm(self, other) + } + + /// Upper triangular part of the tensor's last two dimensions + #[inline(always)] + pub fn triu(&self, diagonal: i64) -> Result { + use crate::operations::linalg::triu; + triu(self, diagonal) + } + + /// Lower triangular part of the tensor's last two dimensions + #[inline(always)] + pub fn tril(&self, diagonal: i64) -> Result { + use crate::operations::linalg::tril; + tril(self, diagonal) + } + + /// Extract a diagonal along two dimensions. + #[inline(always)] + pub fn diagonal(&self, offset: isize, dim1: isize, dim2: isize) -> Result { + use crate::operations::linalg::diagonal; + diagonal(self, offset, dim1, dim2) + } + + /// Sum the diagonal elements along two dimensions. + #[inline(always)] + pub fn trace(&self, offset: isize, dim1: isize, dim2: isize) -> Result { + use crate::operations::linalg::trace; + trace(self, offset, dim1, dim2) + } + + /// Sum reduction + #[inline(always)] + pub fn sum(&self, dim: Option>, keepdim: bool) -> Result { + use crate::operations::reduction::sum; + sum(self, dim, keepdim) + } + + /// NaN-aware sum reduction + #[inline(always)] + pub fn nansum(&self, dim: Option>, keepdim: bool) -> Result { + use crate::operations::reduction::nansum; + nansum(self, dim, keepdim) + } + + /// Log-sum-exp reduction + #[inline(always)] + pub fn logsumexp(&self, dim: Option>, keepdim: bool) -> Result { + use crate::operations::reduction::logsumexp; + logsumexp(self, dim, keepdim) + } + + /// Product reduction + #[inline(always)] + pub fn prod(&self, dim: Option>, keepdim: bool) -> Result { + use crate::operations::reduction::prod; + prod(self, dim, keepdim) + } + + /// Mean reduction + #[inline(always)] + pub fn mean(&self, dim: Option>, keepdim: bool) -> Result { + use crate::operations::reduction::mean; + mean(self, dim, keepdim) + } + + /// NaN-aware mean reduction + #[inline(always)] + pub fn nanmean(&self, dim: Option>, keepdim: bool) -> Result { + use crate::operations::reduction::nanmean; + nanmean(self, dim, keepdim) + } + + /// Logical all reduction + #[inline(always)] + pub fn all(&self, dim: Option, keepdim: bool) -> Result { + use crate::operations::reduction::all; + all(self, dim, keepdim) + } + + /// Logical any reduction + #[inline(always)] + pub fn any(&self, dim: Option, keepdim: bool) -> Result { + use crate::operations::reduction::any; + any(self, dim, keepdim) + } + + /// Cumulative sum along a dimension + #[inline(always)] + pub fn cumsum(&self, dim: isize) -> Result { + use crate::operations::reduction::cumsum; + cumsum(self, dim) + } + + /// Cumulative product along a dimension + #[inline(always)] + pub fn cumprod(&self, dim: isize) -> Result { + use crate::operations::reduction::cumprod; + cumprod(self, dim) + } + + /// Maximum value + #[inline(always)] + pub fn max(&self, dim: Option, keepdim: bool) -> Result { + use crate::operations::reduction::max; + max(self, dim, keepdim) + } + + /// NaN-aware maximum value + #[inline(always)] + pub fn nanmax(&self, dim: Option, keepdim: bool) -> Result { + use crate::operations::reduction::nanmax; + nanmax(self, dim, keepdim) + } + + /// Minimum value + #[inline(always)] + pub fn min(&self, dim: Option, keepdim: bool) -> Result { + use crate::operations::reduction::min; + min(self, dim, keepdim) + } + + /// NaN-aware minimum value + #[inline(always)] + pub fn nanmin(&self, dim: Option, keepdim: bool) -> Result { + use crate::operations::reduction::nanmin; + nanmin(self, dim, keepdim) + } + + /// Argument of maximum value + #[inline(always)] + pub fn argmax(&self, dim: Option, keepdim: bool) -> Result { + use crate::operations::reduction::argmax; + argmax(self, dim, keepdim) + } + + /// Argument of minimum value + #[inline(always)] + pub fn argmin(&self, dim: Option, keepdim: bool) -> Result { + use crate::operations::reduction::argmin; + argmin(self, dim, keepdim) + } + + /// Maximum values and their indices along a dimension + #[inline(always)] + pub fn max_with_indices(&self, dim: isize, keepdim: bool) -> Result<(Self, Self)> { + use crate::operations::reduction::max_with_indices; + max_with_indices(self, dim, keepdim) + } + + /// NaN-aware maximum values and their indices along a dimension + #[inline(always)] + pub fn nanmax_with_indices(&self, dim: isize, keepdim: bool) -> Result<(Self, Self)> { + use crate::operations::reduction::nanmax_with_indices; + nanmax_with_indices(self, dim, keepdim) + } + + /// Minimum values and their indices along a dimension + #[inline(always)] + pub fn min_with_indices(&self, dim: isize, keepdim: bool) -> Result<(Self, Self)> { + use crate::operations::reduction::min_with_indices; + min_with_indices(self, dim, keepdim) + } + + /// NaN-aware minimum values and their indices along a dimension + #[inline(always)] + pub fn nanmin_with_indices(&self, dim: isize, keepdim: bool) -> Result<(Self, Self)> { + use crate::operations::reduction::nanmin_with_indices; + nanmin_with_indices(self, dim, keepdim) + } + + /// Median value (optionally along a dimension) + #[inline(always)] + pub fn median(&self, dim: Option, keepdim: bool) -> Result<(Self, Option)> { + use crate::operations::reduction::median; + median(self, dim, keepdim) + } + + /// Quantile reduction with configurable interpolation + #[inline(always)] + pub fn quantile( + &self, + q: f64, + dim: Option, + keepdim: bool, + interpolation: QuantileInterpolation, + ) -> Result { + use crate::operations::reduction::quantile; + quantile(self, q, dim, keepdim, interpolation) + } + + /// Quantile reduction that ignores NaN values + #[inline(always)] + pub fn nanquantile( + &self, + q: f64, + dim: Option, + keepdim: bool, + interpolation: QuantileInterpolation, + ) -> Result { + use crate::operations::reduction::nanquantile; + nanquantile(self, q, dim, keepdim, interpolation) + } + + /// Median reduction that ignores NaN values + #[inline(always)] + pub fn nanmedian(&self, dim: Option, keepdim: bool) -> Result { + use crate::operations::reduction::nanmedian; + nanmedian(self, dim, keepdim) + } + + /// Batched quantile reduction for multiple probabilities at once + #[inline(always)] + pub fn quantiles( + &self, + qs: &[f64], + dim: Option, + keepdim: bool, + interpolation: QuantileInterpolation, + ) -> Result { + use crate::operations::reduction::quantiles; + quantiles(self, qs, dim, keepdim, interpolation) + } + + /// Batched quantile reduction that ignores NaN values + #[inline(always)] + pub fn nanquantiles( + &self, + qs: &[f64], + dim: Option, + keepdim: bool, + interpolation: QuantileInterpolation, + ) -> Result { + use crate::operations::reduction::nanquantiles; + nanquantiles(self, qs, dim, keepdim, interpolation) + } + + /// Top-k values and indices along a dimension + #[inline(always)] + pub fn topk( + &self, + k: usize, + dim: Option, + largest: bool, + sorted: bool, + ) -> Result<(Self, Self)> { + use crate::operations::reduction::topk; + topk(self, k, dim, largest, sorted) + } + + /// Sort tensor values along a dimension + #[inline(always)] + pub fn sort(&self, dim: Option, descending: bool, stable: bool) -> Result<(Self, Self)> { + use crate::operations::reduction::sort; + sort(self, dim, descending, stable) + } + + /// Indices that would sort the tensor along a dimension + #[inline(always)] + pub fn argsort(&self, dim: Option, descending: bool, stable: bool) -> Result { + use crate::operations::reduction::argsort; + argsort(self, dim, descending, stable) + } + + /// Element-wise equality comparison + #[inline(always)] + pub fn eq(&self, other: &Tensor) -> Result { + use crate::operations::comparison::eq; + eq(self, other) + } + + /// Element-wise inequality comparison + pub fn ne(&self, other: &Tensor) -> Result { + use crate::operations::comparison::ne; + ne(self, other) + } + + /// Element-wise less-than comparison + #[inline(always)] + pub fn lt(&self, other: &Tensor) -> Result { + use crate::operations::comparison::lt; + lt(self, other) + } + + /// Element-wise less-than-or-equal comparison + #[inline(always)] + pub fn le(&self, other: &Tensor) -> Result { + use crate::operations::comparison::le; + le(self, other) + } + + /// Element-wise greater-than comparison + #[inline(always)] + pub fn gt(&self, other: &Tensor) -> Result { + use crate::operations::comparison::gt; + gt(self, other) + } + + /// Element-wise greater-than-or-equal comparison + #[inline(always)] + pub fn ge(&self, other: &Tensor) -> Result { + use crate::operations::comparison::ge; + ge(self, other) + } + + /// Standard deviation + #[inline(always)] + pub fn std(&self, dim: Option>, keepdim: bool, unbiased: bool) -> Result { + use crate::operations::reduction::std; + std(self, dim, keepdim, unbiased) + } + + /// Variance + #[inline(always)] + pub fn var(&self, dim: Option>, keepdim: bool, unbiased: bool) -> Result { + use crate::operations::reduction::var; + var(self, dim, keepdim, unbiased) + } + + /// Exponential function + #[inline(always)] + pub fn exp(&self) -> Result { + use crate::operations::activation::exp; + exp(self) + } + + /// Natural logarithm + #[inline(always)] + pub fn log(&self) -> Result { + use crate::operations::activation::log; + log(self) + } + + /// log1p (log(1 + x)) + #[inline(always)] + pub fn log1p(&self) -> Result { + use crate::operations::activation::log1p; + log1p(self) + } + + /// expm1 (exp(x) - 1) + #[inline(always)] + pub fn expm1(&self) -> Result { + use crate::operations::activation::expm1; + expm1(self) + } + + /// Sine function + #[inline(always)] + pub fn sin(&self) -> Result { + use crate::operations::activation::sin; + sin(self) + } + + /// Cosine function + #[inline(always)] + pub fn cos(&self) -> Result { + use crate::operations::activation::cos; + cos(self) + } + + /// Tangent function + #[inline(always)] + pub fn tan(&self) -> Result { + use crate::operations::activation::tan; + tan(self) + } + + /// Inverse sine function + #[inline(always)] + pub fn asin(&self) -> Result { + use crate::operations::activation::asin; + asin(self) + } + + /// Inverse cosine function + #[inline(always)] + pub fn acos(&self) -> Result { + use crate::operations::activation::acos; + acos(self) + } + + /// Inverse tangent function + #[inline(always)] + pub fn atan(&self) -> Result { + use crate::operations::activation::atan; + atan(self) + } + + /// Hyperbolic sine + #[inline(always)] + pub fn sinh(&self) -> Result { + use crate::operations::activation::sinh; + sinh(self) + } + + /// Hyperbolic cosine + #[inline(always)] + pub fn cosh(&self) -> Result { + use crate::operations::activation::cosh; + cosh(self) + } + + /// Inverse hyperbolic sine + #[inline(always)] + pub fn asinh(&self) -> Result { + use crate::operations::activation::asinh; + asinh(self) + } + + /// Inverse hyperbolic cosine + #[inline(always)] + pub fn acosh(&self) -> Result { + use crate::operations::activation::acosh; + acosh(self) + } + + /// Inverse hyperbolic tangent + #[inline(always)] + pub fn atanh(&self) -> Result { + use crate::operations::activation::atanh; + atanh(self) + } + + /// Hyperbolic tangent + #[inline(always)] + pub fn tanh(&self) -> Result { + use crate::operations::activation::tanh; + tanh(self) + } + + /// Sigmoid activation + #[inline(always)] + pub fn sigmoid(&self) -> Result { + use crate::operations::activation::sigmoid; + sigmoid(self) + } + + /// Softplus activation + #[inline(always)] + pub fn softplus(&self, beta: f64, threshold: f64) -> Result { + use crate::operations::activation::softplus; + softplus(self, beta, threshold) + } + + /// GELU activation + #[inline(always)] + pub fn gelu(&self, approximate: bool) -> Result { + use crate::operations::activation::gelu; + gelu(self, approximate) + } + + /// ELU activation + #[inline(always)] + pub fn elu(&self, alpha: f64) -> Result { + use crate::operations::activation::elu; + elu(self, alpha) + } + + /// SELU activation + #[inline(always)] + pub fn selu(&self) -> Result { + use crate::operations::activation::selu; + selu(self) + } + + /// SiLU activation + #[inline(always)] + pub fn silu(&self) -> Result { + use crate::operations::activation::silu; + silu(self) + } + + /// Softsign activation + #[inline(always)] + pub fn softsign(&self) -> Result { + use crate::operations::activation::softsign; + softsign(self) + } + + /// ReLU activation + #[inline(always)] + pub fn relu(&self) -> Result { + use crate::operations::activation::relu; + relu(self) + } + + /// Hardshrink activation + #[inline(always)] + pub fn hardshrink(&self, lambd: f64) -> Result { + use crate::operations::activation::hardshrink; + hardshrink(self, lambd) + } + + /// Softmax activation + #[inline(always)] + pub fn softmax(&self, dim: Option) -> Result { + use crate::operations::activation::softmax; + softmax(self, dim) + } + + /// Log-Softmax activation + #[inline(always)] + pub fn log_softmax(&self, dim: Option) -> Result { + use crate::operations::activation::log_softmax; + log_softmax(self, dim) + } + + /// Masked Softmax activation + #[inline(always)] + pub fn masked_softmax(&self, mask: &Tensor, dim: Option) -> Result { + use crate::operations::activation::masked_softmax; + masked_softmax(self, mask, dim) + } + + /// Masked Log-Softmax activation + #[inline(always)] + pub fn masked_log_softmax(&self, mask: &Tensor, dim: Option) -> Result { + use crate::operations::activation::masked_log_softmax; + masked_log_softmax(self, mask, dim) + } +} diff --git a/engine/src/tensor/mod/core.rs b/engine/src/tensor/mod/core.rs index dc1a13ff..0b674dd0 100644 --- a/engine/src/tensor/mod/core.rs +++ b/engine/src/tensor/mod/core.rs @@ -1,506 +1,573 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -pub use data::TensorData; -pub use dtype::DataType; -pub use shape::{Shape, Strides}; - -use crate::{ - autograd::{self, CloneBackward, GradientFunction, TensorId}, - device::Device, - error::{MinitensorError, Result}, - operations::{arithmetic::add, reduction::QuantileInterpolation}, -}; -use rayon::prelude::*; -use std::{borrow::Cow, sync::Arc}; - -/// Core tensor structure for minitensor -#[derive(Clone)] -pub struct Tensor { - /// Tensor data storage - data: Arc, - /// Tensor shape (dimensions) - shape: Shape, - /// Memory strides for each dimension - strides: Strides, - /// Data type of tensor elements - dtype: DataType, - /// Device where tensor is stored - device: Device, - /// Whether this tensor requires gradient computation - requires_grad: bool, - /// Gradient function for automatic differentiation - grad_fn: Option>, - /// Stored gradient for this tensor - grad: Option>, - /// Unique identifier for this tensor - tensor_id: TensorId, -} - -/// Index specification for tensor slicing and indexing -#[derive(Clone, Copy, Debug)] -pub enum TensorIndex { - /// Select a single index along the dimension - Index(usize), - /// Select a range with optional step (step defaults to 1) - Slice { - start: usize, - end: usize, - step: usize, - }, -} - -impl Tensor { - /// Create a new tensor with the given data, shape, and properties - #[inline(always)] - pub fn new( - data: Arc, - shape: Shape, - dtype: DataType, - device: Device, - requires_grad: bool, - ) -> Self { - let strides = Strides::from_shape(&shape); - Self { - data, - shape, - strides, - dtype, - device, - requires_grad, - grad_fn: None, - grad: None, - tensor_id: TensorId::new(), - } - } - - /// Create a tensor with uninitialized data - #[inline(always)] - pub fn empty(shape: Shape, dtype: DataType, device: Device, requires_grad: bool) -> Self { - let data = Arc::new(TensorData::uninitialized_on_device( - shape.numel(), - dtype, - device, - )); - Self::new(data, shape, dtype, device, requires_grad) - } - - /// Create a tensor filled with zeros - #[inline(always)] - pub fn zeros(shape: Shape, dtype: DataType, device: Device, requires_grad: bool) -> Self { - let data = Arc::new(TensorData::zeros_on_device(shape.numel(), dtype, device)); - Self::new(data, shape, dtype, device, requires_grad) - } - - /// Create a tensor filled with ones - #[inline(always)] - pub fn ones(shape: Shape, dtype: DataType, device: Device, requires_grad: bool) -> Self { - let data = Arc::new(TensorData::ones_on_device(shape.numel(), dtype, device)); - Self::new(data, shape, dtype, device, requires_grad) - } - - /// Get the tensor's shape - #[inline(always)] - pub fn shape(&self) -> &Shape { - &self.shape - } - - /// Get the tensor's strides - #[inline(always)] - pub fn strides(&self) -> &Strides { - &self.strides - } - - /// Get the tensor's data type - #[inline(always)] - pub fn dtype(&self) -> DataType { - self.dtype - } - - /// Get the tensor's device - #[inline(always)] - pub fn device(&self) -> Device { - self.device - } - - /// Check if this tensor requires gradients - #[inline(always)] - pub fn requires_grad(&self) -> bool { - self.requires_grad - } - - /// Get the tensor's unique ID - #[inline(always)] - pub fn id(&self) -> TensorId { - self.tensor_id - } - - /// Get the number of dimensions - #[inline(always)] - pub fn ndim(&self) -> usize { - self.shape.ndim() - } - - /// Get the total number of elements - #[inline(always)] - pub fn numel(&self) -> usize { - self.shape.numel() - } - - /// Get the size of a specific dimension - #[inline(always)] - pub fn size(&self, dim: usize) -> Result { - self.shape.size(dim) - } - - /// Check if the tensor is contiguous in memory - #[inline(always)] - pub fn is_contiguous(&self) -> bool { - self.strides.is_contiguous(&self.shape) - } - - /// Get a reference to the tensor data - #[inline(always)] - pub fn data(&self) -> &Arc { - &self.data - } - - /// Get a mutable reference to the tensor data - #[inline(always)] - pub(crate) fn data_mut(&mut self) -> &mut TensorData { - let needs_detach = self.grad_fn.is_some() || !self.requires_grad; - if needs_detach { - if Arc::get_mut(&mut self.data).is_none() { - let cloned = self.data.as_ref().clone_data(); - self.data = Arc::new(cloned); - } - Arc::get_mut(&mut self.data).expect("Tensor data should be uniquely owned") - } else { - let ptr = Arc::as_ptr(&self.data) as *mut TensorData; - unsafe { &mut *ptr } - } - } - - /// Create a deep copy of the tensor data while preserving autograd history. - #[inline] - pub fn deep_clone(&self) -> Result { - let data = Arc::new(self.data.as_ref().clone_data()); - let mut cloned = Tensor::new( - data, - self.shape.clone(), - self.dtype, - self.device, - self.requires_grad, - ); - - if self.requires_grad { - let grad_fn = Arc::new(CloneBackward { - input_id: self.tensor_id, - }); - cloned.set_grad_fn(Some(grad_fn.clone())); - autograd::add_to_graph(&cloned, Some(grad_fn))?; - } - - Ok(cloned) - } - - /// Materialise the tensor into a contiguous layout. - pub fn contiguous(&self) -> Result { - if self.is_contiguous() && self.data.is_contiguous() { - return Ok(self.clone()); - } - - if !self.device.is_cpu() { - return Err(MinitensorError::invalid_operation( - "contiguous currently supports only CPU tensors".to_string(), - )); - } - - let numel = self.numel(); - let dtype = self.dtype; - let device = self.device; - let requires_grad = self.requires_grad; - let shape = self.shape.dims().to_vec(); - let strides = self.strides.as_slice().to_vec(); - - let mut output_data = TensorData::uninitialized_on_device(numel, dtype, device); - - match dtype { - DataType::Float32 => { - let src = self.data.as_f32_slice().ok_or_else(|| { - MinitensorError::invalid_operation( - "failed to access float32 data for contiguous copy".to_string(), - ) - })?; - let dst = output_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::invalid_operation( - "failed to access float32 storage for contiguous copy".to_string(), - ) - })?; - copy_strided_to_contiguous(src, dst, &shape, &strides); - } - DataType::Float64 => { - let src = self.data.as_f64_slice().ok_or_else(|| { - MinitensorError::invalid_operation( - "failed to access float64 data for contiguous copy".to_string(), - ) - })?; - let dst = output_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::invalid_operation( - "failed to access float64 storage for contiguous copy".to_string(), - ) - })?; - copy_strided_to_contiguous(src, dst, &shape, &strides); - } - DataType::Int32 => { - let src = self.data.as_i32_slice().ok_or_else(|| { - MinitensorError::invalid_operation( - "failed to access int32 data for contiguous copy".to_string(), - ) - })?; - let dst = output_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::invalid_operation( - "failed to access int32 storage for contiguous copy".to_string(), - ) - })?; - copy_strided_to_contiguous(src, dst, &shape, &strides); - } - DataType::Int64 => { - let src = self.data.as_i64_slice().ok_or_else(|| { - MinitensorError::invalid_operation( - "failed to access int64 data for contiguous copy".to_string(), - ) - })?; - let dst = output_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::invalid_operation( - "failed to access int64 storage for contiguous copy".to_string(), - ) - })?; - copy_strided_to_contiguous(src, dst, &shape, &strides); - } - DataType::Bool => { - let src = self.data.as_bool_slice().ok_or_else(|| { - MinitensorError::invalid_operation( - "failed to access bool data for contiguous copy".to_string(), - ) - })?; - let dst = output_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::invalid_operation( - "failed to access bool storage for contiguous copy".to_string(), - ) - })?; - copy_strided_to_contiguous(src, dst, &shape, &strides); - } - } - - let mut output = Tensor::new( - Arc::new(output_data), - self.shape.clone(), - dtype, - device, - requires_grad, - ); - - if requires_grad { - let grad_fn = Arc::new(CloneBackward { - input_id: self.tensor_id, - }); - output.set_grad_fn(Some(grad_fn.clone())); - autograd::add_to_graph(&output, Some(grad_fn))?; - } - - Ok(output) - } -} - -impl Tensor { - /// Set the gradient function for this tensor - #[inline(always)] - pub fn set_grad_fn(&mut self, grad_fn: Option>) { - self.grad_fn = grad_fn; - } - - /// Get the gradient function for this tensor - #[inline(always)] - pub fn grad_fn(&self) -> Option<&Arc> { - self.grad_fn.as_ref() - } - - /// Enable gradient computation for this tensor - #[inline(always)] - pub fn requires_grad_(mut self, requires_grad: bool) -> Self { - self.requires_grad = requires_grad; - self - } - - /// Assign a fresh tensor identifier and clear autograd metadata. - #[inline(always)] - pub(crate) fn refresh_autograd_metadata(&mut self) { - self.tensor_id = TensorId::new(); - self.grad_fn = None; - self.grad = None; - } - - /// Get the gradient for this tensor - #[inline(always)] - pub fn grad(&self) -> Option<&Arc> { - self.grad.as_ref() - } - - /// Get mutable access to the gradient if uniquely owned - #[inline(always)] - pub fn grad_mut(&mut self) -> Option<&mut Tensor> { - self.grad.as_mut().and_then(|g| std::sync::Arc::get_mut(g)) - } - - /// Set the gradient for this tensor - #[inline(always)] - pub fn set_grad(&mut self, grad: Option) { - self.grad = grad.map(Arc::new); - } - - /// Accumulate gradient for this tensor - #[inline] - pub fn accumulate_grad(&mut self, grad: Tensor) -> Result<()> { - match &self.grad { - Some(existing) => { - let sum = add(existing.as_ref(), &grad)?; - self.grad = Some(Arc::new(sum)); - } - None => { - self.grad = Some(Arc::new(grad)); - } - } - Ok(()) - } - - /// Clear the gradient for this tensor - #[inline(always)] - pub fn zero_grad(&mut self, set_to_none: bool) { - autograd::zero_gradients(); - if set_to_none { - self.grad = None; - return; - } - - // If gradient exists, zero it in place - if let Some(g) = self.grad_mut() { - match g.dtype() { - DataType::Float32 => { - if let Some(slice) = g.data_mut().as_f32_slice_mut() { - slice.fill(0.0); - } - } - DataType::Float64 => { - if let Some(slice) = g.data_mut().as_f64_slice_mut() { - slice.fill(0.0); - } - } - DataType::Int32 => { - if let Some(slice) = g.data_mut().as_i32_slice_mut() { - slice.fill(0); - } - } - DataType::Int64 => { - if let Some(slice) = g.data_mut().as_i64_slice_mut() { - slice.fill(0); - } - } - DataType::Bool => { - if let Some(slice) = g.data_mut().as_bool_slice_mut() { - slice.fill(false); - } - } - } - } else if self.requires_grad { - // If gradient doesn't exist but is required, create a zero tensor - let zero = Tensor::zeros(self.shape.clone(), self.dtype, self.device, false); - self.grad = Some(Arc::new(zero)); - } else { - self.grad = None; - } - } - - /// Check if this tensor has a gradient - #[inline(always)] - pub fn has_grad(&self) -> bool { - self.grad.is_some() - } - - /// Perform backward pass from this tensor - pub fn backward(&self, gradient: Option) -> Result<()> { - use crate::autograd; - - // If no gradient is provided, create a gradient of ones for scalar tensors - let grad = match gradient { - Some(g) => g, - None => { - if self.numel() != 1 { - return Err(MinitensorError::gradient_error( - "Gradient can only be implicitly created for scalar tensors", - )); - } - // Create a tensor of ones with the same shape as self - Self::ones(self.shape.clone(), self.dtype, self.device, false) - } - }; - - // Perform backward pass through the computation graph - autograd::backward(self, Some(grad)).map(|_| ()) // Convert HashMap result to () - } -} - -impl Tensor { - /// Create a view of this tensor with a new shape - #[inline(always)] - pub fn view(&self, new_shape: Shape) -> Result { - if new_shape.numel() != self.numel() { - return Err(MinitensorError::shape_mismatch( - vec![self.numel()], - vec![new_shape.numel()], - )); - } - - let mut tensor = self.clone(); - tensor.strides = Strides::from_shape(&new_shape); - tensor.shape = new_shape; - Ok(tensor) - } - - /// Reshape the tensor to a new shape - #[inline(always)] - pub fn reshape(&self, new_shape: Shape) -> Result { - self.view(new_shape) - } - - /// Flatten the tensor into a one-dimensional view. - /// This operation avoids data copies when possible. - #[inline(always)] - pub fn flatten_all(&self) -> Result { - let len = self.numel(); - self.reshape(Shape::new(vec![len])) - } - - /// Alias for [`flatten_all`](Self::flatten_all) for backward compatibility. - #[inline(always)] - pub fn ravel(&self) -> Result { - self.flatten_all() - } - - /// Transpose two dimensions of the tensor - #[inline(always)] - pub fn transpose(&self, dim0: isize, dim1: isize) -> Result { - use crate::operations::linalg::transpose; - transpose(self, dim0, dim1) - } - - /// Permute tensor dimensions - #[inline(always)] - pub fn permute(&self, dims: Vec) -> Result { - use crate::operations::shape_ops::permute; - permute(self, dims) - } -} +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +pub use super::data::{DataMut, TensorData}; +pub use super::dtype::DataType; +pub use super::shape::{Shape, Strides}; + +// Method impls split by concern. They are children of this module (rather +// than siblings) so they keep access to `Tensor`'s private fields. +#[path = "autograd.rs"] +mod autograd_methods; +#[path = "indexing.rs"] +mod indexing_methods; +#[path = "ops.rs"] +mod ops_methods; +#[path = "utils.rs"] +mod utils_methods; + +use self::indexing_methods::copy_strided_to_contiguous; + +use crate::{ + autograd::{self, CloneBackward, GradientFunction, TensorId}, + device::Device, + error::{MinitensorError, Result}, + operations::{arithmetic::add, reduction::QuantileInterpolation}, +}; +use rayon::prelude::*; +use std::{borrow::Cow, sync::Arc}; + +/// Core tensor structure for minitensor +#[derive(Clone)] +pub struct Tensor { + /// Tensor data storage + data: Arc, + /// Tensor shape (dimensions) + shape: Shape, + /// Memory strides for each dimension + strides: Strides, + /// Data type of tensor elements + dtype: DataType, + /// Device where tensor is stored + device: Device, + /// Whether this tensor requires gradient computation + requires_grad: bool, + /// Gradient function for automatic differentiation + grad_fn: Option>, + /// Stored gradient for this tensor + grad: Option>, + /// Unique identifier for this tensor + tensor_id: TensorId, +} + +/// Index specification for tensor slicing and indexing +#[derive(Clone, Copy, Debug)] +pub enum TensorIndex { + /// Select a single index along the dimension + Index(usize), + /// Select a range with optional step (step defaults to 1) + Slice { + start: usize, + end: usize, + step: usize, + }, +} + +impl Tensor { + /// Create a new tensor with the given data, shape, and properties + #[inline(always)] + pub fn new( + data: Arc, + shape: Shape, + dtype: DataType, + device: Device, + requires_grad: bool, + ) -> Self { + let strides = Strides::from_shape(&shape); + Self { + data, + shape, + strides, + dtype, + device, + // While grad recording is disabled (`no_grad`), newly created + // tensors never require gradients. Callers that genuinely need a + // trainable leaf inside a no-grad scope can opt back in with + // `requires_grad_(true)`, which expresses explicit intent and is + // not gated. + requires_grad: requires_grad && autograd::is_grad_enabled(), + grad_fn: None, + grad: None, + tensor_id: TensorId::new(), + } + } + + /// Create a tensor with uninitialized data + #[inline(always)] + pub fn empty(shape: Shape, dtype: DataType, device: Device, requires_grad: bool) -> Self { + let data = Arc::new(TensorData::uninitialized_on_device( + shape.numel(), + dtype, + device, + )); + Self::new(data, shape, dtype, device, requires_grad) + } + + /// Create a tensor filled with zeros + #[inline(always)] + pub fn zeros(shape: Shape, dtype: DataType, device: Device, requires_grad: bool) -> Self { + let data = Arc::new(TensorData::zeros_on_device(shape.numel(), dtype, device)); + Self::new(data, shape, dtype, device, requires_grad) + } + + /// Create a tensor filled with ones + #[inline(always)] + pub fn ones(shape: Shape, dtype: DataType, device: Device, requires_grad: bool) -> Self { + let data = Arc::new(TensorData::ones_on_device(shape.numel(), dtype, device)); + Self::new(data, shape, dtype, device, requires_grad) + } + + /// Get the tensor's shape + #[inline(always)] + pub fn shape(&self) -> &Shape { + &self.shape + } + + /// Get the tensor's strides + #[inline(always)] + pub fn strides(&self) -> &Strides { + &self.strides + } + + /// Get the tensor's data type + #[inline(always)] + pub fn dtype(&self) -> DataType { + self.dtype + } + + /// Get the tensor's device + #[inline(always)] + pub fn device(&self) -> Device { + self.device + } + + /// Check if this tensor requires gradients + #[inline(always)] + pub fn requires_grad(&self) -> bool { + self.requires_grad + } + + /// Get the tensor's unique ID + #[inline(always)] + pub fn id(&self) -> TensorId { + self.tensor_id + } + + /// Get the number of dimensions + #[inline(always)] + pub fn ndim(&self) -> usize { + self.shape.ndim() + } + + /// Get the total number of elements + #[inline(always)] + pub fn numel(&self) -> usize { + self.shape.numel() + } + + /// Get the size of a specific dimension + #[inline(always)] + pub fn size(&self, dim: usize) -> Result { + self.shape.size(dim) + } + + /// Check if the tensor is contiguous in memory + #[inline(always)] + pub fn is_contiguous(&self) -> bool { + self.strides.is_contiguous(&self.shape) + } + + /// Get a reference to the tensor data + #[inline(always)] + pub fn data(&self) -> &Arc { + &self.data + } + + /// Get mutable access to the tensor data. + /// + /// Non-leaf tensors (and tensors that do not require gradients) get + /// copy-on-write semantics: if the storage is shared, it is cloned first + /// so in-place mutation cannot corrupt other tensors or saved autograd + /// state. + /// + /// Leaf tensors that require gradients (i.e. parameters) are mutated in + /// place even when the storage is shared, so that optimizer updates stay + /// visible through every handle to the parameter — mirroring PyTorch's + /// in-place parameter update semantics. That path goes through the + /// storage layer's interior mutability (see [`DataMut`]) instead of + /// fabricating a `&mut TensorData` from the shared `Arc`. + #[inline(always)] + pub(crate) fn data_mut(&mut self) -> DataMut<'_> { + let needs_detach = self.grad_fn.is_some() || !self.requires_grad; + if needs_detach { + if Arc::get_mut(&mut self.data).is_none() { + let cloned = self.data.as_ref().clone_data(); + self.data = Arc::new(cloned); + } + return DataMut::Unique( + Arc::get_mut(&mut self.data).expect("Tensor data should be uniquely owned"), + ); + } + // Take the exclusive path whenever the storage is uniquely owned. + if Arc::get_mut(&mut self.data).is_some() { + return DataMut::Unique(Arc::get_mut(&mut self.data).expect("uniqueness just checked")); + } + DataMut::Shared(self.data.as_ref()) + } + + /// Create a deep copy of the tensor data while preserving autograd history. + #[inline] + pub fn deep_clone(&self) -> Result { + let data = Arc::new(self.data.as_ref().clone_data()); + let mut cloned = Tensor::new( + data, + self.shape.clone(), + self.dtype, + self.device, + self.requires_grad, + ); + + if self.requires_grad { + let grad_fn = Arc::new(CloneBackward { + input_id: self.tensor_id, + }); + cloned.set_grad_fn(Some(grad_fn.clone())); + autograd::add_to_graph(&cloned, Some(grad_fn))?; + } + + Ok(cloned) + } + + /// Materialise the tensor into a contiguous layout. + pub fn contiguous(&self) -> Result { + if self.is_contiguous() && self.data.is_contiguous() { + return Ok(self.clone()); + } + + if !self.device.is_cpu() { + return Err(MinitensorError::invalid_operation( + "contiguous currently supports only CPU tensors".to_string(), + )); + } + + let numel = self.numel(); + let dtype = self.dtype; + let device = self.device; + let requires_grad = self.requires_grad; + let shape = self.shape.dims().to_vec(); + let strides = self.strides.as_slice().to_vec(); + + let mut output_data = TensorData::uninitialized_on_device(numel, dtype, device); + + match dtype { + DataType::Float32 => { + let src = self.data.as_f32_slice().ok_or_else(|| { + MinitensorError::invalid_operation( + "failed to access float32 data for contiguous copy".to_string(), + ) + })?; + let dst = output_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::invalid_operation( + "failed to access float32 storage for contiguous copy".to_string(), + ) + })?; + copy_strided_to_contiguous(src, dst, &shape, &strides); + } + DataType::Float64 => { + let src = self.data.as_f64_slice().ok_or_else(|| { + MinitensorError::invalid_operation( + "failed to access float64 data for contiguous copy".to_string(), + ) + })?; + let dst = output_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::invalid_operation( + "failed to access float64 storage for contiguous copy".to_string(), + ) + })?; + copy_strided_to_contiguous(src, dst, &shape, &strides); + } + DataType::Int32 => { + let src = self.data.as_i32_slice().ok_or_else(|| { + MinitensorError::invalid_operation( + "failed to access int32 data for contiguous copy".to_string(), + ) + })?; + let dst = output_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::invalid_operation( + "failed to access int32 storage for contiguous copy".to_string(), + ) + })?; + copy_strided_to_contiguous(src, dst, &shape, &strides); + } + DataType::Int64 => { + let src = self.data.as_i64_slice().ok_or_else(|| { + MinitensorError::invalid_operation( + "failed to access int64 data for contiguous copy".to_string(), + ) + })?; + let dst = output_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::invalid_operation( + "failed to access int64 storage for contiguous copy".to_string(), + ) + })?; + copy_strided_to_contiguous(src, dst, &shape, &strides); + } + DataType::Bool => { + let src = self.data.as_bool_slice().ok_or_else(|| { + MinitensorError::invalid_operation( + "failed to access bool data for contiguous copy".to_string(), + ) + })?; + let dst = output_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::invalid_operation( + "failed to access bool storage for contiguous copy".to_string(), + ) + })?; + copy_strided_to_contiguous(src, dst, &shape, &strides); + } + } + + let mut output = Tensor::new( + Arc::new(output_data), + self.shape.clone(), + dtype, + device, + requires_grad, + ); + + if requires_grad { + let grad_fn = Arc::new(CloneBackward { + input_id: self.tensor_id, + }); + output.set_grad_fn(Some(grad_fn.clone())); + autograd::add_to_graph(&output, Some(grad_fn))?; + } + + Ok(output) + } +} + +impl Tensor { + /// Set the gradient function for this tensor. + /// + /// While grad recording is disabled (`no_grad` scopes and the backward + /// executor), attaching a gradient function is a no-op so operation + /// outputs stay leaves; this mirrors the gating in + /// [`autograd::add_to_graph`], keeping the tensor's metadata and the + /// graph consistent. Clearing (`None`) is always honoured. + #[inline(always)] + pub fn set_grad_fn(&mut self, grad_fn: Option>) { + if grad_fn.is_some() && !autograd::is_grad_enabled() { + return; + } + self.grad_fn = grad_fn; + } + + /// Get the gradient function for this tensor + #[inline(always)] + pub fn grad_fn(&self) -> Option<&Arc> { + self.grad_fn.as_ref() + } + + /// Enable gradient computation for this tensor + #[inline(always)] + pub fn requires_grad_(mut self, requires_grad: bool) -> Self { + self.requires_grad = requires_grad; + self + } + + /// Assign a fresh tensor identifier and clear autograd metadata. + #[inline(always)] + pub(crate) fn refresh_autograd_metadata(&mut self) { + self.tensor_id = TensorId::new(); + self.grad_fn = None; + self.grad = None; + } + + /// Get the gradient for this tensor + #[inline(always)] + pub fn grad(&self) -> Option<&Arc> { + self.grad.as_ref() + } + + /// Get mutable access to the gradient if uniquely owned + #[inline(always)] + pub fn grad_mut(&mut self) -> Option<&mut Tensor> { + self.grad.as_mut().and_then(|g| std::sync::Arc::get_mut(g)) + } + + /// Set the gradient for this tensor + #[inline(always)] + pub fn set_grad(&mut self, grad: Option) { + self.grad = grad.map(Arc::new); + } + + /// Accumulate gradient for this tensor + #[inline] + pub fn accumulate_grad(&mut self, grad: Tensor) -> Result<()> { + match &self.grad { + Some(existing) => { + let sum = add(existing.as_ref(), &grad)?; + self.grad = Some(Arc::new(sum)); + } + None => { + self.grad = Some(Arc::new(grad)); + } + } + Ok(()) + } + + /// Clear the gradient for this tensor. + /// + /// Only this tensor's gradient is affected — both the copy stored on the + /// tensor and its entry in the global gradient map. (Earlier versions + /// wiped every gradient on the thread, so zeroing one tensor silently + /// cleared unrelated models' gradients.) + #[inline(always)] + pub fn zero_grad(&mut self, set_to_none: bool) { + autograd::clear_gradient(self); + if set_to_none { + self.grad = None; + return; + } + + // If gradient exists, zero it in place + if let Some(g) = self.grad_mut() { + match g.dtype() { + DataType::Float32 => { + if let Some(slice) = g.data_mut().as_f32_slice_mut() { + slice.fill(0.0); + } + } + DataType::Float64 => { + if let Some(slice) = g.data_mut().as_f64_slice_mut() { + slice.fill(0.0); + } + } + DataType::Int32 => { + if let Some(slice) = g.data_mut().as_i32_slice_mut() { + slice.fill(0); + } + } + DataType::Int64 => { + if let Some(slice) = g.data_mut().as_i64_slice_mut() { + slice.fill(0); + } + } + DataType::Bool => { + if let Some(slice) = g.data_mut().as_bool_slice_mut() { + slice.fill(false); + } + } + } + } else if self.requires_grad { + // If gradient doesn't exist but is required, create a zero tensor + let zero = Tensor::zeros(self.shape.clone(), self.dtype, self.device, false); + self.grad = Some(Arc::new(zero)); + } else { + self.grad = None; + } + } + + /// Check if this tensor has a gradient + #[inline(always)] + pub fn has_grad(&self) -> bool { + self.grad.is_some() + } + + /// Perform backward pass from this tensor + pub fn backward(&self, gradient: Option) -> Result<()> { + use crate::autograd; + + // If no gradient is provided, create a gradient of ones for scalar tensors + let grad = match gradient { + Some(g) => g, + None => { + if self.numel() != 1 { + return Err(MinitensorError::gradient_error( + "Gradient can only be implicitly created for scalar tensors", + )); + } + // Create a tensor of ones with the same shape as self + Self::ones(self.shape.clone(), self.dtype, self.device, false) + } + }; + + // Perform backward pass through the computation graph + autograd::backward(self, Some(grad)) + } +} + +impl Tensor { + /// Create a view of this tensor with a new shape. + /// + /// The tensor must be contiguous: a view only reinterprets the existing + /// buffer, and re-striding a non-contiguous tensor (e.g. the result of + /// `expand`) would silently associate the new shape with storage that does + /// not contain the tensor's logical elements. Use [`Self::reshape`] to get + /// an automatic copy in that case. + #[inline(always)] + pub fn view(&self, new_shape: Shape) -> Result { + if new_shape.numel() != self.numel() { + return Err(MinitensorError::shape_mismatch( + vec![self.numel()], + vec![new_shape.numel()], + )); + } + + if !self.is_contiguous() { + return Err(MinitensorError::invalid_operation( + "view is not supported for non-contiguous tensors; call contiguous() or reshape() instead", + )); + } + + let mut tensor = self.clone(); + tensor.strides = Strides::from_shape(&new_shape); + tensor.shape = new_shape; + Ok(tensor) + } + + /// Reshape the tensor to a new shape. + /// + /// Returns a zero-copy view when the tensor is contiguous and materialises + /// a contiguous copy otherwise. + #[inline(always)] + pub fn reshape(&self, new_shape: Shape) -> Result { + if self.is_contiguous() { + self.view(new_shape) + } else { + self.contiguous()?.view(new_shape) + } + } + + /// Flatten the tensor into a one-dimensional view. + /// This operation avoids data copies when possible. + #[inline(always)] + pub fn flatten_all(&self) -> Result { + let len = self.numel(); + self.reshape(Shape::new(vec![len])) + } + + /// Alias for [`flatten_all`](Self::flatten_all) for backward compatibility. + #[inline(always)] + pub fn ravel(&self) -> Result { + self.flatten_all() + } + + /// Transpose two dimensions of the tensor + #[inline(always)] + pub fn transpose(&self, dim0: isize, dim1: isize) -> Result { + use crate::operations::linalg::transpose; + transpose(self, dim0, dim1) + } + + /// Permute tensor dimensions + #[inline(always)] + pub fn permute(&self, dims: Vec) -> Result { + use crate::operations::shape_ops::permute; + permute(self, dims) + } +} diff --git a/engine/src/tensor/mod/indexing.rs b/engine/src/tensor/mod/indexing.rs index 2a991ce2..0b0cafbf 100644 --- a/engine/src/tensor/mod/indexing.rs +++ b/engine/src/tensor/mod/indexing.rs @@ -1,888 +1,902 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -impl Tensor { - /// Squeeze dimensions of size 1 - #[inline(always)] - pub fn squeeze(&self) -> Result { - let new_dims: Vec = self - .shape - .dims() - .iter() - .filter(|&&dim| dim != 1) - .copied() - .collect(); - - let new_shape = Shape::new(new_dims); - self.view(new_shape) - } - - /// Squeeze specific dimension if it has size 1. Negative indices are supported. - #[inline(always)] - pub fn squeeze_dim(&self, dim: isize) -> Result { - let ndim = self.ndim() as isize; - let dim = if dim < 0 { dim + ndim } else { dim }; - - if dim < 0 || dim >= ndim { - return Err(MinitensorError::index_error(dim, 0, ndim as usize)); - } - - let dim = dim as usize; - - if self.shape.dims()[dim] != 1 { - return Ok(self.clone()); - } - - let mut new_dims = self.shape.dims().to_vec(); - new_dims.remove(dim); - let new_shape = Shape::new(new_dims); - self.view(new_shape) - } - - /// Add dimension of size 1. Negative indices are supported. - #[inline(always)] - pub fn unsqueeze(&self, dim: isize) -> Result { - let ndim = self.ndim() as isize; - let dim = if dim < 0 { dim + ndim + 1 } else { dim }; - - if dim < 0 || dim > ndim { - return Err(MinitensorError::index_error(dim, 0, (ndim + 1) as usize)); - } - - let dim = dim as usize; - - let mut new_dims = self.shape.dims().to_vec(); - new_dims.insert(dim, 1); - let new_shape = Shape::new(new_dims); - self.view(new_shape) - } - - /// Expand tensor dimensions without allocating new memory - #[inline(always)] - pub fn expand(&self, dims: Vec) -> Result { - let orig_dims = self.shape.dims(); - let orig_strides = self.strides.as_slice(); - let n_orig = orig_dims.len(); - let n_new = dims.len(); - - if n_new < n_orig { - return Err(MinitensorError::invalid_operation( - "cannot expand to fewer dimensions".to_string(), - )); - } - - let mut new_dims = vec![0usize; n_new]; - let mut new_strides = vec![0usize; n_new]; - - for i in 0..n_new { - let size_spec = dims[n_new - 1 - i]; - if size_spec < -1 { - return Err(MinitensorError::invalid_operation( - "invalid negative dimension".to_string(), - )); - } - - let orig_idx_opt = if i < n_orig { - Some(n_orig - 1 - i) - } else { - None - }; - let orig_dim = orig_idx_opt.map(|idx| orig_dims[idx]).unwrap_or(1); - let orig_stride = orig_idx_opt.map(|idx| orig_strides[idx]).unwrap_or(0); - - let target = if size_spec == -1 { - orig_dim - } else { - size_spec as usize - }; - - if let Some(idx) = orig_idx_opt { - if target == orig_dim { - new_dims[n_new - 1 - i] = target; - new_strides[n_new - 1 - i] = orig_stride; - } else if orig_dim == 1 && target > 0 { - new_dims[n_new - 1 - i] = target; - new_strides[n_new - 1 - i] = 0; - } else { - return Err(MinitensorError::invalid_operation(format!( - "cannot expand dimension {} from {} to {}", - idx, orig_dim, target - ))); - } - } else { - if size_spec == -1 { - return Err(MinitensorError::invalid_operation( - "the size -1 is not allowed for a new leading dimension".to_string(), - )); - } - // New leading dimensions broadcast with stride 0. - new_dims[n_new - 1 - i] = target; - new_strides[n_new - 1 - i] = 0; - } - } - - let mut tensor = self.clone(); - tensor.refresh_autograd_metadata(); - tensor.shape = Shape::new(new_dims.clone()); - tensor.strides = Strides::new(new_strides); - - if tensor.requires_grad { - let grad_fn = Arc::new(crate::autograd::ExpandBackward { - input_shape: orig_dims.to_vec(), - input_id: self.id(), - }); - tensor.set_grad_fn(Some(grad_fn.clone())); - autograd::add_to_graph(&tensor, Some(grad_fn))?; - } - - Ok(tensor) - } - - /// Repeat tensor according to `repeats` along each dimension - #[inline(always)] - pub fn repeat(&self, repeats: Vec) -> Result { - crate::operations::shape_ops::repeat(self, &repeats) - } - - /// Flatten tensor from `start_dim` to `end_dim` - pub fn flatten(&self, start_dim: isize, end_dim: isize) -> Result { - let ndim = self.ndim() as isize; - - let start = if start_dim < 0 { - start_dim + ndim - } else { - start_dim - }; - let end = if end_dim < 0 { end_dim + ndim } else { end_dim }; - - if start < 0 || start >= ndim { - return Err(MinitensorError::index_error(start, 0, ndim as usize)); - } - if end < 0 || end >= ndim { - return Err(MinitensorError::index_error(end, 0, ndim as usize)); - } - if start > end { - return Err(MinitensorError::invalid_argument( - "start_dim must be less than or equal to end_dim", - )); - } - - self.flatten_range(start as usize, end as usize) - } - - /// Flatten tensor from start_dim to end_dim - #[inline(always)] - pub fn flatten_range(&self, start_dim: usize, end_dim: usize) -> Result { - if start_dim >= self.ndim() || end_dim >= self.ndim() || start_dim > end_dim { - return Err(MinitensorError::invalid_argument( - "Invalid dimension range for flatten", - )); - } - - let dims = self.shape.dims(); - let mut new_dims = Vec::new(); - - // Add dimensions before start_dim - new_dims.extend_from_slice(&dims[..start_dim]); - - // Compute flattened dimension size - let flattened_size: usize = dims[start_dim..=end_dim].iter().product(); - new_dims.push(flattened_size); - - // Add dimensions after end_dim - if end_dim + 1 < dims.len() { - new_dims.extend_from_slice(&dims[end_dim + 1..]); - } - - let new_shape = Shape::new(new_dims); - self.view(new_shape) - } -} - -impl Tensor { - /// Basic tensor indexing and slicing - #[inline(always)] - pub fn index(&self, indices: &[TensorIndex]) -> Result { - if indices.len() > self.ndim() { - return Err(MinitensorError::invalid_argument( - "Too many indices for tensor", - )); - } - - let shape_dims = self.shape.dims(); - let strides = self.strides.as_slice(); - let mut offset = 0usize; - let mut out_dims = Vec::new(); - let mut orig_dim_map = Vec::new(); - let mut starts = Vec::new(); - let mut steps: Vec = Vec::new(); - - for i in 0..self.ndim() { - let dim_size = shape_dims[i]; - let idx = indices.get(i).cloned().unwrap_or(TensorIndex::Slice { - start: 0, - end: dim_size, - step: 1, - }); - match idx { - TensorIndex::Index(pos) => { - if pos >= dim_size { - return Err(MinitensorError::index_error(pos as isize, 0, dim_size)); - } - offset += pos * strides[i]; - } - TensorIndex::Slice { start, end, step } => { - if start > end || end > dim_size { - return Err(MinitensorError::index_error(end as isize, 0, dim_size)); - } - let size = if end <= start { - 0 - } else { - (end - start).div_ceil(step) - }; - out_dims.push(size); - orig_dim_map.push(i); - starts.push(start); - steps.push(step); - } - } - } - - if out_dims.is_empty() { - let mut result_data = TensorData::zeros_on_device(1, self.dtype, self.device); - match self.dtype { - DataType::Float32 => { - let input = self - .data - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected f32 data"))?; - result_data.as_f32_slice_mut().unwrap()[0] = input[offset]; - } - DataType::Float64 => { - let input = self - .data - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected f64 data"))?; - result_data.as_f64_slice_mut().unwrap()[0] = input[offset]; - } - DataType::Int32 => { - let input = self - .data - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected i32 data"))?; - result_data.as_i32_slice_mut().unwrap()[0] = input[offset]; - } - DataType::Int64 => { - let input = self - .data - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected i64 data"))?; - result_data.as_i64_slice_mut().unwrap()[0] = input[offset]; - } - DataType::Bool => { - let input = self - .data - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected bool data"))?; - result_data.as_bool_slice_mut().unwrap()[0] = input[offset]; - } - } - let output = Tensor::new( - Arc::new(result_data), - Shape::scalar(), - self.dtype, - self.device, - self.requires_grad, - ); - return self.wrap_index_grad(output, offset, Vec::new(), Vec::new(), Vec::new(), Vec::new()); - } - - let out_shape = Shape::new(out_dims.clone()); - let out_strides = Strides::from_shape(&out_shape); - let mut result_data = - TensorData::zeros_on_device(out_shape.numel(), self.dtype, self.device); - - match self.dtype { - DataType::Float32 => { - let input = self - .data - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected f32 data"))?; - let output = result_data.as_f32_slice_mut().unwrap(); - for (idx, out_elem) in output.iter_mut().enumerate() { - let mut rem = idx; - let mut src_idx = offset; - for (j, &stride) in out_strides.as_slice().iter().enumerate() { - let coord = rem / stride; - rem %= stride; - let orig_dim = orig_dim_map[j]; - let step = steps[j]; - src_idx += (starts[j] + coord * step) * strides[orig_dim]; - } - *out_elem = input[src_idx]; - } - } - DataType::Float64 => { - let input = self - .data - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected f64 data"))?; - let output = result_data.as_f64_slice_mut().unwrap(); - for (idx, out_elem) in output.iter_mut().enumerate() { - let mut rem = idx; - let mut src_idx = offset; - for (j, &stride) in out_strides.as_slice().iter().enumerate() { - let coord = rem / stride; - rem %= stride; - let orig_dim = orig_dim_map[j]; - let step = steps[j]; - src_idx += (starts[j] + coord * step) * strides[orig_dim]; - } - *out_elem = input[src_idx]; - } - } - DataType::Int32 => { - let input = self - .data - .as_i32_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected i32 data"))?; - let output = result_data.as_i32_slice_mut().unwrap(); - for (idx, out_elem) in output.iter_mut().enumerate() { - let mut rem = idx; - let mut src_idx = offset; - for (j, &stride) in out_strides.as_slice().iter().enumerate() { - let coord = rem / stride; - rem %= stride; - let orig_dim = orig_dim_map[j]; - let step = steps[j]; - src_idx += (starts[j] + coord * step) * strides[orig_dim]; - } - *out_elem = input[src_idx]; - } - } - DataType::Int64 => { - let input = self - .data - .as_i64_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected i64 data"))?; - let output = result_data.as_i64_slice_mut().unwrap(); - for (idx, out_elem) in output.iter_mut().enumerate() { - let mut rem = idx; - let mut src_idx = offset; - for (j, &stride) in out_strides.as_slice().iter().enumerate() { - let coord = rem / stride; - rem %= stride; - let orig_dim = orig_dim_map[j]; - let step = steps[j]; - src_idx += (starts[j] + coord * step) * strides[orig_dim]; - } - *out_elem = input[src_idx]; - } - } - DataType::Bool => { - let input = self - .data - .as_bool_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected bool data"))?; - let output = result_data.as_bool_slice_mut().unwrap(); - for (idx, out_elem) in output.iter_mut().enumerate() { - let mut rem = idx; - let mut src_idx = offset; - for (j, &stride) in out_strides.as_slice().iter().enumerate() { - let coord = rem / stride; - rem %= stride; - let orig_dim = orig_dim_map[j]; - let step = steps[j]; - src_idx += (starts[j] + coord * step) * strides[orig_dim]; - } - *out_elem = input[src_idx]; - } - } - } - - let output = Tensor::new( - Arc::new(result_data), - out_shape, - self.dtype, - self.device, - self.requires_grad, - ); - self.wrap_index_grad(output, offset, out_dims, orig_dim_map, starts, steps) - } - - /// Attach an [`IndexBackward`] gradient function to a freshly indexed tensor. - /// - /// `out_dims` is empty for a scalar (fully integer-indexed) result. Gradient - /// tracking is only wired for floating-point, contiguous inputs, which is - /// always the case at the Python boundary where indexing is applied. - fn wrap_index_grad( - &self, - output: Tensor, - offset: usize, - out_dims: Vec, - orig_dim_map: Vec, - starts: Vec, - steps: Vec, - ) -> Result { - if !self.requires_grad || !self.dtype.is_float() || !self.is_contiguous() { - return Ok(output); - } - let grad_fn = Arc::new(crate::autograd::IndexBackward { - input_id: self.tensor_id, - input_shape: self.shape.dims().to_vec(), - input_strides: self.strides.as_slice().to_vec(), - offset, - out_dims, - orig_dim_map, - starts, - steps, - }); - let mut output = output; - output.set_grad_fn(Some(grad_fn.clone())); - autograd::add_to_graph(&output, Some(grad_fn))?; - Ok(output) - } - - /// Assign values to tensor slice - #[inline(always)] - pub fn index_assign(&mut self, indices: &[TensorIndex], value: &Tensor) -> Result<()> { - if indices.len() > self.ndim() { - return Err(MinitensorError::invalid_argument( - "Too many indices for tensor", - )); - } - - let shape_dims = self.shape.dims(); - let strides = self.strides.as_slice(); - let mut offset = 0usize; - let mut out_dims = Vec::new(); - let mut orig_dim_map = Vec::new(); - let mut starts = Vec::new(); - let mut steps: Vec = Vec::new(); - - for i in 0..self.ndim() { - let dim_size = shape_dims[i]; - let idx = indices.get(i).cloned().unwrap_or(TensorIndex::Slice { - start: 0, - end: dim_size, - step: 1, - }); - match idx { - TensorIndex::Index(pos) => { - if pos >= dim_size { - return Err(MinitensorError::index_error(pos as isize, 0, dim_size)); - } - offset += pos * strides[i]; - } - TensorIndex::Slice { start, end, step } => { - if start > end || end > dim_size { - return Err(MinitensorError::index_error(end as isize, 0, dim_size)); - } - let size = if end <= start { - 0 - } else { - (end - start).div_ceil(step) - }; - out_dims.push(size); - orig_dim_map.push(i); - starts.push(start); - steps.push(step); - } - } - } - - let out_shape = Shape::new(out_dims.clone()); - if value.numel() != out_shape.numel() && value.numel() != 1 { - return Err(MinitensorError::invalid_argument( - "Assigned value has incompatible shape", - )); - } - - let out_strides = Strides::from_shape(&out_shape); - let data = if let Some(d) = Arc::get_mut(&mut self.data) { - d - } else { - let cloned = self.data.clone_data(); - self.data = Arc::new(cloned); - Arc::get_mut(&mut self.data).unwrap() - }; - - match self.dtype { - DataType::Float32 => { - let slice = data.as_f32_slice_mut().unwrap(); - let val_slice = value.data().as_f32_slice().unwrap(); - for (idx, &val) in val_slice.iter().cycle().take(out_shape.numel()).enumerate() { - let mut rem = idx; - let mut src_idx = offset; - for (j, &stride) in out_strides.as_slice().iter().enumerate() { - let coord = rem / stride; - rem %= stride; - let orig_dim = orig_dim_map[j]; - let step = steps[j]; - src_idx += (starts[j] + coord * step) * strides[orig_dim]; - } - slice[src_idx] = val; - } - } - DataType::Float64 => { - let slice = data.as_f64_slice_mut().unwrap(); - let val_slice = value.data().as_f64_slice().unwrap(); - for (idx, &val) in val_slice.iter().cycle().take(out_shape.numel()).enumerate() { - let mut rem = idx; - let mut src_idx = offset; - for (j, &stride) in out_strides.as_slice().iter().enumerate() { - let coord = rem / stride; - rem %= stride; - let orig_dim = orig_dim_map[j]; - let step = steps[j]; - src_idx += (starts[j] + coord * step) * strides[orig_dim]; - } - slice[src_idx] = val; - } - } - DataType::Int32 => { - let slice = data.as_i32_slice_mut().unwrap(); - let val_slice = value.data().as_i32_slice().unwrap(); - for (idx, &val) in val_slice.iter().cycle().take(out_shape.numel()).enumerate() { - let mut rem = idx; - let mut src_idx = offset; - for (j, &stride) in out_strides.as_slice().iter().enumerate() { - let coord = rem / stride; - rem %= stride; - let orig_dim = orig_dim_map[j]; - let step = steps[j]; - src_idx += (starts[j] + coord * step) * strides[orig_dim]; - } - slice[src_idx] = val; - } - } - DataType::Int64 => { - let slice = data.as_i64_slice_mut().unwrap(); - let val_slice = value.data().as_i64_slice().unwrap(); - for (idx, &val) in val_slice.iter().cycle().take(out_shape.numel()).enumerate() { - let mut rem = idx; - let mut src_idx = offset; - for (j, &stride) in out_strides.as_slice().iter().enumerate() { - let coord = rem / stride; - rem %= stride; - let orig_dim = orig_dim_map[j]; - let step = steps[j]; - src_idx += (starts[j] + coord * step) * strides[orig_dim]; - } - slice[src_idx] = val; - } - } - DataType::Bool => { - let slice = data.as_bool_slice_mut().unwrap(); - let val_slice = value.data().as_bool_slice().unwrap(); - for (idx, &val) in val_slice.iter().cycle().take(out_shape.numel()).enumerate() { - let mut rem = idx; - let mut src_idx = offset; - for (j, &stride) in out_strides.as_slice().iter().enumerate() { - let coord = rem / stride; - rem %= stride; - let orig_dim = orig_dim_map[j]; - let step = steps[j]; - src_idx += (starts[j] + coord * step) * strides[orig_dim]; - } - slice[src_idx] = val; - } - } - } - Ok(()) - } -} - -impl Tensor { - /// Check if tensor contains NaN values - #[inline(always)] - pub fn has_nan(&self) -> bool { - match self.dtype { - DataType::Float32 => { - if let Some(data) = self.data.as_f32_slice() { - data.iter().any(|&x| x.is_nan()) - } else { - false - } - } - DataType::Float64 => { - if let Some(data) = self.data.as_f64_slice() { - data.iter().any(|&x| x.is_nan()) - } else { - false - } - } - _ => false, // Integer and boolean types cannot be NaN - } - } - - /// Check if tensor contains infinite values - #[inline(always)] - pub fn has_inf(&self) -> bool { - match self.dtype { - DataType::Float32 => { - if let Some(data) = self.data.as_f32_slice() { - data.iter().any(|&x| x.is_infinite()) - } else { - false - } - } - DataType::Float64 => { - if let Some(data) = self.data.as_f64_slice() { - data.iter().any(|&x| x.is_infinite()) - } else { - false - } - } - _ => false, // Integer and boolean types cannot be infinite - } - } - - /// Element-wise check for NaN values - #[inline(always)] - pub fn isnan(&self) -> Result { - let mut output = TensorData::zeros_on_device(self.numel(), DataType::Bool, self.device); - match self.dtype { - DataType::Float32 => { - let input = self - .data - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected f32 data"))?; - let out_slice = output - .as_bool_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - for (o, &x) in out_slice.iter_mut().zip(input.iter()) { - *o = x.is_nan(); - } - } - DataType::Float64 => { - let input = self - .data - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected f64 data"))?; - let out_slice = output - .as_bool_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - for (o, &x) in out_slice.iter_mut().zip(input.iter()) { - *o = x.is_nan(); - } - } - _ => { - // Non-floating types cannot be NaN; output already zero - } - } - Ok(Tensor::new( - Arc::new(output), - self.shape.clone(), - DataType::Bool, - self.device, - false, - )) - } - - /// Element-wise check for infinite values - #[inline(always)] - pub fn isinf(&self) -> Result { - let mut output = TensorData::zeros_on_device(self.numel(), DataType::Bool, self.device); - match self.dtype { - DataType::Float32 => { - let input = self - .data - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected f32 data"))?; - let out_slice = output - .as_bool_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - for (o, &x) in out_slice.iter_mut().zip(input.iter()) { - *o = x.is_infinite(); - } - } - DataType::Float64 => { - let input = self - .data - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected f64 data"))?; - let out_slice = output - .as_bool_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - for (o, &x) in out_slice.iter_mut().zip(input.iter()) { - *o = x.is_infinite(); - } - } - _ => { - // Non-floating types cannot be infinite; output remains false - } - } - Ok(Tensor::new( - Arc::new(output), - self.shape.clone(), - DataType::Bool, - self.device, - false, - )) - } - - /// Element-wise check for finite values - #[inline(always)] - pub fn isfinite(&self) -> Result { - let mut output = TensorData::zeros_on_device(self.numel(), DataType::Bool, self.device); - match self.dtype { - DataType::Float32 => { - let input = self - .data - .as_f32_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected f32 data"))?; - let out_slice = output - .as_bool_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - for (o, &x) in out_slice.iter_mut().zip(input.iter()) { - *o = x.is_finite(); - } - } - DataType::Float64 => { - let input = self - .data - .as_f64_slice() - .ok_or_else(|| MinitensorError::internal_error("Expected f64 data"))?; - let out_slice = output - .as_bool_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - for (o, &x) in out_slice.iter_mut().zip(input.iter()) { - *o = x.is_finite(); - } - } - _ => { - // Integer and bool types are always finite - let out_slice = output - .as_bool_slice_mut() - .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; - for o in out_slice.iter_mut() { - *o = true; - } - } - } - Ok(Tensor::new( - Arc::new(output), - self.shape.clone(), - DataType::Bool, - self.device, - false, - )) - } - - /// Get the maximum value in the tensor - #[inline(always)] - pub fn max_value(&self) -> Option { - match self.dtype { - DataType::Float32 => self - .data - .as_f32_slice()? - .iter() - .max_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)) - .map(|&x| x as f64), - DataType::Float64 => self - .data - .as_f64_slice()? - .iter() - .max_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)) - .copied(), - DataType::Int32 => self.data.as_i32_slice()?.iter().max().map(|&x| x as f64), - DataType::Int64 => self.data.as_i64_slice()?.iter().max().map(|&x| x as f64), - DataType::Bool => self - .data - .as_bool_slice()? - .iter() - .max() - .map(|&x| if x { 1.0 } else { 0.0 }), - } - } - - /// Get the minimum value in the tensor - #[inline(always)] - pub fn min_value(&self) -> Option { - match self.dtype { - DataType::Float32 => self - .data - .as_f32_slice()? - .iter() - .min_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)) - .map(|&x| x as f64), - DataType::Float64 => self - .data - .as_f64_slice()? - .iter() - .min_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)) - .copied(), - DataType::Int32 => self.data.as_i32_slice()?.iter().min().map(|&x| x as f64), - DataType::Int64 => self.data.as_i64_slice()?.iter().min().map(|&x| x as f64), - DataType::Bool => self - .data - .as_bool_slice()? - .iter() - .min() - .map(|&x| if x { 1.0 } else { 0.0 }), - } - } - - /// Get memory usage in bytes - #[inline(always)] - pub fn memory_usage_bytes(&self) -> usize { - let element_size = match self.dtype { - DataType::Float32 => 4, - DataType::Float64 => 8, - DataType::Int32 => 4, - DataType::Int64 => 8, - DataType::Bool => 1, - }; - self.numel() * element_size - } - - /// Get the stride information - pub fn stride(&self) -> &Strides { - &self.strides - } - - /// Check if this tensor is a leaf node in the computation graph - #[inline(always)] - pub fn is_leaf(&self) -> bool { - self.grad_fn.is_none() - } -} - -fn copy_strided_to_contiguous( - src: &[T], - dst: &mut [T], - shape: &[usize], - strides: &[usize], -) { - if dst.is_empty() { - return; - } - - if shape.is_empty() { - dst[0] = src[0]; - return; - } - - let ndim = shape.len(); - let mut index = vec![0usize; ndim]; - - for value in dst.iter_mut() { - let mut offset = 0usize; - for (&idx, &stride) in index.iter().zip(strides.iter()) { - offset += idx * stride; - } - *value = src[offset]; - - for dim in (0..ndim).rev() { - index[dim] += 1; - if index[dim] < shape[dim] { - break; - } - index[dim] = 0; - } - } -} +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::{ + autograd::{self}, + error::{MinitensorError, Result}, +}; +use std::sync::Arc; + +impl Tensor { + /// Squeeze dimensions of size 1 + #[inline(always)] + pub fn squeeze(&self) -> Result { + let new_dims: Vec = self + .shape + .dims() + .iter() + .filter(|&&dim| dim != 1) + .copied() + .collect(); + + let new_shape = Shape::new(new_dims); + self.reshape(new_shape) + } + + /// Squeeze specific dimension if it has size 1. Negative indices are supported. + #[inline(always)] + pub fn squeeze_dim(&self, dim: isize) -> Result { + let ndim = self.ndim() as isize; + let dim = if dim < 0 { dim + ndim } else { dim }; + + if dim < 0 || dim >= ndim { + return Err(MinitensorError::index_error(dim, 0, ndim as usize)); + } + + let dim = dim as usize; + + if self.shape.dims()[dim] != 1 { + return Ok(self.clone()); + } + + let mut new_dims = self.shape.dims().to_vec(); + new_dims.remove(dim); + let new_shape = Shape::new(new_dims); + self.reshape(new_shape) + } + + /// Add dimension of size 1. Negative indices are supported. + #[inline(always)] + pub fn unsqueeze(&self, dim: isize) -> Result { + let ndim = self.ndim() as isize; + let dim = if dim < 0 { dim + ndim + 1 } else { dim }; + + if dim < 0 || dim > ndim { + return Err(MinitensorError::index_error(dim, 0, (ndim + 1) as usize)); + } + + let dim = dim as usize; + + let mut new_dims = self.shape.dims().to_vec(); + new_dims.insert(dim, 1); + let new_shape = Shape::new(new_dims); + self.reshape(new_shape) + } + + /// Expand tensor dimensions without allocating new memory + #[inline(always)] + pub fn expand(&self, dims: Vec) -> Result { + let orig_dims = self.shape.dims(); + let orig_strides = self.strides.as_slice(); + let n_orig = orig_dims.len(); + let n_new = dims.len(); + + if n_new < n_orig { + return Err(MinitensorError::invalid_operation( + "cannot expand to fewer dimensions".to_string(), + )); + } + + let mut new_dims = vec![0usize; n_new]; + let mut new_strides = vec![0usize; n_new]; + + for i in 0..n_new { + let size_spec = dims[n_new - 1 - i]; + if size_spec < -1 { + return Err(MinitensorError::invalid_operation( + "invalid negative dimension".to_string(), + )); + } + + let orig_idx_opt = if i < n_orig { + Some(n_orig - 1 - i) + } else { + None + }; + let orig_dim = orig_idx_opt.map(|idx| orig_dims[idx]).unwrap_or(1); + let orig_stride = orig_idx_opt.map(|idx| orig_strides[idx]).unwrap_or(0); + + let target = if size_spec == -1 { + orig_dim + } else { + size_spec as usize + }; + + if let Some(idx) = orig_idx_opt { + if target == orig_dim { + new_dims[n_new - 1 - i] = target; + new_strides[n_new - 1 - i] = orig_stride; + } else if orig_dim == 1 && target > 0 { + new_dims[n_new - 1 - i] = target; + new_strides[n_new - 1 - i] = 0; + } else { + return Err(MinitensorError::invalid_operation(format!( + "cannot expand dimension {} from {} to {}", + idx, orig_dim, target + ))); + } + } else { + if size_spec == -1 { + return Err(MinitensorError::invalid_operation( + "the size -1 is not allowed for a new leading dimension".to_string(), + )); + } + // New leading dimensions broadcast with stride 0. + new_dims[n_new - 1 - i] = target; + new_strides[n_new - 1 - i] = 0; + } + } + + let mut tensor = self.clone(); + tensor.refresh_autograd_metadata(); + tensor.shape = Shape::new(new_dims.clone()); + tensor.strides = Strides::new(new_strides); + + if tensor.requires_grad { + let grad_fn = Arc::new(crate::autograd::ExpandBackward { + input_shape: orig_dims.to_vec(), + input_id: self.id(), + }); + tensor.set_grad_fn(Some(grad_fn.clone())); + autograd::add_to_graph(&tensor, Some(grad_fn))?; + } + + Ok(tensor) + } + + /// Repeat tensor according to `repeats` along each dimension + #[inline(always)] + pub fn repeat(&self, repeats: Vec) -> Result { + crate::operations::shape_ops::repeat(self, &repeats) + } + + /// Flatten tensor from `start_dim` to `end_dim` + pub fn flatten(&self, start_dim: isize, end_dim: isize) -> Result { + let ndim = self.ndim() as isize; + + let start = if start_dim < 0 { + start_dim + ndim + } else { + start_dim + }; + let end = if end_dim < 0 { end_dim + ndim } else { end_dim }; + + if start < 0 || start >= ndim { + return Err(MinitensorError::index_error(start, 0, ndim as usize)); + } + if end < 0 || end >= ndim { + return Err(MinitensorError::index_error(end, 0, ndim as usize)); + } + if start > end { + return Err(MinitensorError::invalid_argument( + "start_dim must be less than or equal to end_dim", + )); + } + + self.flatten_range(start as usize, end as usize) + } + + /// Flatten tensor from start_dim to end_dim + #[inline(always)] + pub fn flatten_range(&self, start_dim: usize, end_dim: usize) -> Result { + if start_dim >= self.ndim() || end_dim >= self.ndim() || start_dim > end_dim { + return Err(MinitensorError::invalid_argument( + "Invalid dimension range for flatten", + )); + } + + let dims = self.shape.dims(); + let mut new_dims = Vec::new(); + + // Add dimensions before start_dim + new_dims.extend_from_slice(&dims[..start_dim]); + + // Compute flattened dimension size + let flattened_size: usize = dims[start_dim..=end_dim].iter().product(); + new_dims.push(flattened_size); + + // Add dimensions after end_dim + if end_dim + 1 < dims.len() { + new_dims.extend_from_slice(&dims[end_dim + 1..]); + } + + let new_shape = Shape::new(new_dims); + self.reshape(new_shape) + } +} + +impl Tensor { + /// Basic tensor indexing and slicing + #[inline(always)] + pub fn index(&self, indices: &[TensorIndex]) -> Result { + if indices.len() > self.ndim() { + return Err(MinitensorError::invalid_argument( + "Too many indices for tensor", + )); + } + + let shape_dims = self.shape.dims(); + let strides = self.strides.as_slice(); + let mut offset = 0usize; + let mut out_dims = Vec::new(); + let mut orig_dim_map = Vec::new(); + let mut starts = Vec::new(); + let mut steps: Vec = Vec::new(); + + for i in 0..self.ndim() { + let dim_size = shape_dims[i]; + let idx = indices.get(i).cloned().unwrap_or(TensorIndex::Slice { + start: 0, + end: dim_size, + step: 1, + }); + match idx { + TensorIndex::Index(pos) => { + if pos >= dim_size { + return Err(MinitensorError::index_error(pos as isize, 0, dim_size)); + } + offset += pos * strides[i]; + } + TensorIndex::Slice { start, end, step } => { + if start > end || end > dim_size { + return Err(MinitensorError::index_error(end as isize, 0, dim_size)); + } + let size = if end <= start { + 0 + } else { + (end - start).div_ceil(step) + }; + out_dims.push(size); + orig_dim_map.push(i); + starts.push(start); + steps.push(step); + } + } + } + + if out_dims.is_empty() { + let mut result_data = TensorData::zeros_on_device(1, self.dtype, self.device); + match self.dtype { + DataType::Float32 => { + let input = self + .data + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected f32 data"))?; + result_data.as_f32_slice_mut().unwrap()[0] = input[offset]; + } + DataType::Float64 => { + let input = self + .data + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected f64 data"))?; + result_data.as_f64_slice_mut().unwrap()[0] = input[offset]; + } + DataType::Int32 => { + let input = self + .data + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected i32 data"))?; + result_data.as_i32_slice_mut().unwrap()[0] = input[offset]; + } + DataType::Int64 => { + let input = self + .data + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected i64 data"))?; + result_data.as_i64_slice_mut().unwrap()[0] = input[offset]; + } + DataType::Bool => { + let input = self + .data + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected bool data"))?; + result_data.as_bool_slice_mut().unwrap()[0] = input[offset]; + } + } + let output = Tensor::new( + Arc::new(result_data), + Shape::scalar(), + self.dtype, + self.device, + self.requires_grad, + ); + return self.wrap_index_grad( + output, + offset, + Vec::new(), + Vec::new(), + Vec::new(), + Vec::new(), + ); + } + + let out_shape = Shape::new(out_dims.clone()); + let out_strides = Strides::from_shape(&out_shape); + let mut result_data = + TensorData::zeros_on_device(out_shape.numel(), self.dtype, self.device); + + match self.dtype { + DataType::Float32 => { + let input = self + .data + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected f32 data"))?; + let output = result_data.as_f32_slice_mut().unwrap(); + for (idx, out_elem) in output.iter_mut().enumerate() { + let mut rem = idx; + let mut src_idx = offset; + for (j, &stride) in out_strides.as_slice().iter().enumerate() { + let coord = rem / stride; + rem %= stride; + let orig_dim = orig_dim_map[j]; + let step = steps[j]; + src_idx += (starts[j] + coord * step) * strides[orig_dim]; + } + *out_elem = input[src_idx]; + } + } + DataType::Float64 => { + let input = self + .data + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected f64 data"))?; + let output = result_data.as_f64_slice_mut().unwrap(); + for (idx, out_elem) in output.iter_mut().enumerate() { + let mut rem = idx; + let mut src_idx = offset; + for (j, &stride) in out_strides.as_slice().iter().enumerate() { + let coord = rem / stride; + rem %= stride; + let orig_dim = orig_dim_map[j]; + let step = steps[j]; + src_idx += (starts[j] + coord * step) * strides[orig_dim]; + } + *out_elem = input[src_idx]; + } + } + DataType::Int32 => { + let input = self + .data + .as_i32_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected i32 data"))?; + let output = result_data.as_i32_slice_mut().unwrap(); + for (idx, out_elem) in output.iter_mut().enumerate() { + let mut rem = idx; + let mut src_idx = offset; + for (j, &stride) in out_strides.as_slice().iter().enumerate() { + let coord = rem / stride; + rem %= stride; + let orig_dim = orig_dim_map[j]; + let step = steps[j]; + src_idx += (starts[j] + coord * step) * strides[orig_dim]; + } + *out_elem = input[src_idx]; + } + } + DataType::Int64 => { + let input = self + .data + .as_i64_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected i64 data"))?; + let output = result_data.as_i64_slice_mut().unwrap(); + for (idx, out_elem) in output.iter_mut().enumerate() { + let mut rem = idx; + let mut src_idx = offset; + for (j, &stride) in out_strides.as_slice().iter().enumerate() { + let coord = rem / stride; + rem %= stride; + let orig_dim = orig_dim_map[j]; + let step = steps[j]; + src_idx += (starts[j] + coord * step) * strides[orig_dim]; + } + *out_elem = input[src_idx]; + } + } + DataType::Bool => { + let input = self + .data + .as_bool_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected bool data"))?; + let output = result_data.as_bool_slice_mut().unwrap(); + for (idx, out_elem) in output.iter_mut().enumerate() { + let mut rem = idx; + let mut src_idx = offset; + for (j, &stride) in out_strides.as_slice().iter().enumerate() { + let coord = rem / stride; + rem %= stride; + let orig_dim = orig_dim_map[j]; + let step = steps[j]; + src_idx += (starts[j] + coord * step) * strides[orig_dim]; + } + *out_elem = input[src_idx]; + } + } + } + + let output = Tensor::new( + Arc::new(result_data), + out_shape, + self.dtype, + self.device, + self.requires_grad, + ); + self.wrap_index_grad(output, offset, out_dims, orig_dim_map, starts, steps) + } + + /// Attach an [`IndexBackward`] gradient function to a freshly indexed tensor. + /// + /// `out_dims` is empty for a scalar (fully integer-indexed) result. Gradient + /// tracking is only wired for floating-point, contiguous inputs, which is + /// always the case at the Python boundary where indexing is applied. + fn wrap_index_grad( + &self, + output: Tensor, + offset: usize, + out_dims: Vec, + orig_dim_map: Vec, + starts: Vec, + steps: Vec, + ) -> Result { + if !self.requires_grad || !self.dtype.is_float() || !self.is_contiguous() { + return Ok(output); + } + let grad_fn = Arc::new(crate::autograd::IndexBackward { + input_id: self.tensor_id, + input_shape: self.shape.dims().to_vec(), + input_strides: self.strides.as_slice().to_vec(), + offset, + out_dims, + orig_dim_map, + starts, + steps, + }); + let mut output = output; + output.set_grad_fn(Some(grad_fn.clone())); + autograd::add_to_graph(&output, Some(grad_fn))?; + Ok(output) + } + + /// Assign values to tensor slice + #[inline(always)] + pub fn index_assign(&mut self, indices: &[TensorIndex], value: &Tensor) -> Result<()> { + if indices.len() > self.ndim() { + return Err(MinitensorError::invalid_argument( + "Too many indices for tensor", + )); + } + + let shape_dims = self.shape.dims(); + let strides = self.strides.as_slice(); + let mut offset = 0usize; + let mut out_dims = Vec::new(); + let mut orig_dim_map = Vec::new(); + let mut starts = Vec::new(); + let mut steps: Vec = Vec::new(); + + for i in 0..self.ndim() { + let dim_size = shape_dims[i]; + let idx = indices.get(i).cloned().unwrap_or(TensorIndex::Slice { + start: 0, + end: dim_size, + step: 1, + }); + match idx { + TensorIndex::Index(pos) => { + if pos >= dim_size { + return Err(MinitensorError::index_error(pos as isize, 0, dim_size)); + } + offset += pos * strides[i]; + } + TensorIndex::Slice { start, end, step } => { + if start > end || end > dim_size { + return Err(MinitensorError::index_error(end as isize, 0, dim_size)); + } + let size = if end <= start { + 0 + } else { + (end - start).div_ceil(step) + }; + out_dims.push(size); + orig_dim_map.push(i); + starts.push(start); + steps.push(step); + } + } + } + + let out_shape = Shape::new(out_dims.clone()); + if value.numel() != out_shape.numel() && value.numel() != 1 { + return Err(MinitensorError::invalid_argument( + "Assigned value has incompatible shape", + )); + } + + let out_strides = Strides::from_shape(&out_shape); + let data = if let Some(d) = Arc::get_mut(&mut self.data) { + d + } else { + let cloned = self.data.clone_data(); + self.data = Arc::new(cloned); + Arc::get_mut(&mut self.data).unwrap() + }; + + match self.dtype { + DataType::Float32 => { + let slice = data.as_f32_slice_mut().unwrap(); + let val_slice = value.data().as_f32_slice().unwrap(); + for (idx, &val) in val_slice.iter().cycle().take(out_shape.numel()).enumerate() { + let mut rem = idx; + let mut src_idx = offset; + for (j, &stride) in out_strides.as_slice().iter().enumerate() { + let coord = rem / stride; + rem %= stride; + let orig_dim = orig_dim_map[j]; + let step = steps[j]; + src_idx += (starts[j] + coord * step) * strides[orig_dim]; + } + slice[src_idx] = val; + } + } + DataType::Float64 => { + let slice = data.as_f64_slice_mut().unwrap(); + let val_slice = value.data().as_f64_slice().unwrap(); + for (idx, &val) in val_slice.iter().cycle().take(out_shape.numel()).enumerate() { + let mut rem = idx; + let mut src_idx = offset; + for (j, &stride) in out_strides.as_slice().iter().enumerate() { + let coord = rem / stride; + rem %= stride; + let orig_dim = orig_dim_map[j]; + let step = steps[j]; + src_idx += (starts[j] + coord * step) * strides[orig_dim]; + } + slice[src_idx] = val; + } + } + DataType::Int32 => { + let slice = data.as_i32_slice_mut().unwrap(); + let val_slice = value.data().as_i32_slice().unwrap(); + for (idx, &val) in val_slice.iter().cycle().take(out_shape.numel()).enumerate() { + let mut rem = idx; + let mut src_idx = offset; + for (j, &stride) in out_strides.as_slice().iter().enumerate() { + let coord = rem / stride; + rem %= stride; + let orig_dim = orig_dim_map[j]; + let step = steps[j]; + src_idx += (starts[j] + coord * step) * strides[orig_dim]; + } + slice[src_idx] = val; + } + } + DataType::Int64 => { + let slice = data.as_i64_slice_mut().unwrap(); + let val_slice = value.data().as_i64_slice().unwrap(); + for (idx, &val) in val_slice.iter().cycle().take(out_shape.numel()).enumerate() { + let mut rem = idx; + let mut src_idx = offset; + for (j, &stride) in out_strides.as_slice().iter().enumerate() { + let coord = rem / stride; + rem %= stride; + let orig_dim = orig_dim_map[j]; + let step = steps[j]; + src_idx += (starts[j] + coord * step) * strides[orig_dim]; + } + slice[src_idx] = val; + } + } + DataType::Bool => { + let slice = data.as_bool_slice_mut().unwrap(); + let val_slice = value.data().as_bool_slice().unwrap(); + for (idx, &val) in val_slice.iter().cycle().take(out_shape.numel()).enumerate() { + let mut rem = idx; + let mut src_idx = offset; + for (j, &stride) in out_strides.as_slice().iter().enumerate() { + let coord = rem / stride; + rem %= stride; + let orig_dim = orig_dim_map[j]; + let step = steps[j]; + src_idx += (starts[j] + coord * step) * strides[orig_dim]; + } + slice[src_idx] = val; + } + } + } + Ok(()) + } +} + +impl Tensor { + /// Check if tensor contains NaN values + #[inline(always)] + pub fn has_nan(&self) -> bool { + match self.dtype { + DataType::Float32 => { + if let Some(data) = self.data.as_f32_slice() { + data.iter().any(|&x| x.is_nan()) + } else { + false + } + } + DataType::Float64 => { + if let Some(data) = self.data.as_f64_slice() { + data.iter().any(|&x| x.is_nan()) + } else { + false + } + } + _ => false, // Integer and boolean types cannot be NaN + } + } + + /// Check if tensor contains infinite values + #[inline(always)] + pub fn has_inf(&self) -> bool { + match self.dtype { + DataType::Float32 => { + if let Some(data) = self.data.as_f32_slice() { + data.iter().any(|&x| x.is_infinite()) + } else { + false + } + } + DataType::Float64 => { + if let Some(data) = self.data.as_f64_slice() { + data.iter().any(|&x| x.is_infinite()) + } else { + false + } + } + _ => false, // Integer and boolean types cannot be infinite + } + } + + /// Element-wise check for NaN values + #[inline(always)] + pub fn isnan(&self) -> Result { + let mut output = TensorData::zeros_on_device(self.numel(), DataType::Bool, self.device); + match self.dtype { + DataType::Float32 => { + let input = self + .data + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected f32 data"))?; + let out_slice = output + .as_bool_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + for (o, &x) in out_slice.iter_mut().zip(input.iter()) { + *o = x.is_nan(); + } + } + DataType::Float64 => { + let input = self + .data + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected f64 data"))?; + let out_slice = output + .as_bool_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + for (o, &x) in out_slice.iter_mut().zip(input.iter()) { + *o = x.is_nan(); + } + } + _ => { + // Non-floating types cannot be NaN; output already zero + } + } + Ok(Tensor::new( + Arc::new(output), + self.shape.clone(), + DataType::Bool, + self.device, + false, + )) + } + + /// Element-wise check for infinite values + #[inline(always)] + pub fn isinf(&self) -> Result { + let mut output = TensorData::zeros_on_device(self.numel(), DataType::Bool, self.device); + match self.dtype { + DataType::Float32 => { + let input = self + .data + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected f32 data"))?; + let out_slice = output + .as_bool_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + for (o, &x) in out_slice.iter_mut().zip(input.iter()) { + *o = x.is_infinite(); + } + } + DataType::Float64 => { + let input = self + .data + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected f64 data"))?; + let out_slice = output + .as_bool_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + for (o, &x) in out_slice.iter_mut().zip(input.iter()) { + *o = x.is_infinite(); + } + } + _ => { + // Non-floating types cannot be infinite; output remains false + } + } + Ok(Tensor::new( + Arc::new(output), + self.shape.clone(), + DataType::Bool, + self.device, + false, + )) + } + + /// Element-wise check for finite values + #[inline(always)] + pub fn isfinite(&self) -> Result { + let mut output = TensorData::zeros_on_device(self.numel(), DataType::Bool, self.device); + match self.dtype { + DataType::Float32 => { + let input = self + .data + .as_f32_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected f32 data"))?; + let out_slice = output + .as_bool_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + for (o, &x) in out_slice.iter_mut().zip(input.iter()) { + *o = x.is_finite(); + } + } + DataType::Float64 => { + let input = self + .data + .as_f64_slice() + .ok_or_else(|| MinitensorError::internal_error("Expected f64 data"))?; + let out_slice = output + .as_bool_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + for (o, &x) in out_slice.iter_mut().zip(input.iter()) { + *o = x.is_finite(); + } + } + _ => { + // Integer and bool types are always finite + let out_slice = output + .as_bool_slice_mut() + .ok_or_else(|| MinitensorError::internal_error("Failed to get bool slice"))?; + for o in out_slice.iter_mut() { + *o = true; + } + } + } + Ok(Tensor::new( + Arc::new(output), + self.shape.clone(), + DataType::Bool, + self.device, + false, + )) + } + + /// Get the maximum value in the tensor + #[inline(always)] + pub fn max_value(&self) -> Option { + match self.dtype { + DataType::Float32 => self + .data + .as_f32_slice()? + .iter() + .max_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)) + .map(|&x| x as f64), + DataType::Float64 => self + .data + .as_f64_slice()? + .iter() + .max_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)) + .copied(), + DataType::Int32 => self.data.as_i32_slice()?.iter().max().map(|&x| x as f64), + DataType::Int64 => self.data.as_i64_slice()?.iter().max().map(|&x| x as f64), + DataType::Bool => self + .data + .as_bool_slice()? + .iter() + .max() + .map(|&x| if x { 1.0 } else { 0.0 }), + } + } + + /// Get the minimum value in the tensor + #[inline(always)] + pub fn min_value(&self) -> Option { + match self.dtype { + DataType::Float32 => self + .data + .as_f32_slice()? + .iter() + .min_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)) + .map(|&x| x as f64), + DataType::Float64 => self + .data + .as_f64_slice()? + .iter() + .min_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)) + .copied(), + DataType::Int32 => self.data.as_i32_slice()?.iter().min().map(|&x| x as f64), + DataType::Int64 => self.data.as_i64_slice()?.iter().min().map(|&x| x as f64), + DataType::Bool => self + .data + .as_bool_slice()? + .iter() + .min() + .map(|&x| if x { 1.0 } else { 0.0 }), + } + } + + /// Get memory usage in bytes + #[inline(always)] + pub fn memory_usage_bytes(&self) -> usize { + let element_size = match self.dtype { + DataType::Float32 => 4, + DataType::Float64 => 8, + DataType::Int32 => 4, + DataType::Int64 => 8, + DataType::Bool => 1, + }; + self.numel() * element_size + } + + /// Get the stride information + pub fn stride(&self) -> &Strides { + &self.strides + } + + /// Check if this tensor is a leaf node in the computation graph + #[inline(always)] + pub fn is_leaf(&self) -> bool { + self.grad_fn.is_none() + } +} + +pub(super) fn copy_strided_to_contiguous( + src: &[T], + dst: &mut [T], + shape: &[usize], + strides: &[usize], +) { + if dst.is_empty() { + return; + } + + if shape.is_empty() { + dst[0] = src[0]; + return; + } + + let ndim = shape.len(); + let mut index = vec![0usize; ndim]; + + for value in dst.iter_mut() { + let mut offset = 0usize; + for (&idx, &stride) in index.iter().zip(strides.iter()) { + offset += idx * stride; + } + *value = src[offset]; + + for dim in (0..ndim).rev() { + index[dim] += 1; + if index[dim] < shape[dim] { + break; + } + index[dim] = 0; + } + } +} diff --git a/engine/src/tensor/mod/ops.rs b/engine/src/tensor/mod/ops.rs index 5a2e330a..c2d70c40 100644 --- a/engine/src/tensor/mod/ops.rs +++ b/engine/src/tensor/mod/ops.rs @@ -1,835 +1,834 @@ -// Copyright (c) Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -#[inline(always)] -fn allclose_f32(a: f32, b: f32, rtol: f32, atol: f32, equal_nan: bool) -> bool { - if a == b { - return true; - } - if equal_nan && a.is_nan() && b.is_nan() { - return true; - } - if !a.is_finite() || !b.is_finite() { - return false; - } - let diff = (a - b).abs(); - diff <= atol + rtol * b.abs() -} - -#[inline(always)] -fn allclose_f64(a: f64, b: f64, rtol: f64, atol: f64, equal_nan: bool) -> bool { - if a == b { - return true; - } - if equal_nan && a.is_nan() && b.is_nan() { - return true; - } - if !a.is_finite() || !b.is_finite() { - return false; - } - let diff = (a - b).abs(); - diff <= atol + rtol * b.abs() -} - -impl Tensor { - /// Solve a linear system `AX = B` for `X` where `self` provides `A`. - pub fn solve(&self, rhs: &Self) -> Result { - use crate::operations::linalg::solve; - solve(self, rhs) - } - - /// Layer normalization - #[inline(always)] - pub fn layer_norm( - &self, - normalized_shape: &[usize], - weight: Option<&Tensor>, - bias: Option<&Tensor>, - eps: f64, - ) -> Result { - use crate::operations::normalization::layer_norm; - layer_norm(self, normalized_shape, weight, bias, eps) - } - - /// Absolute value - #[inline(always)] - pub fn abs(&self) -> Result { - use crate::operations::activation::abs; - abs(self) - } - - /// Element-wise sign (returns -1, 0, or 1 for each value). - #[inline(always)] - pub fn sign(&self) -> Result { - use crate::operations::activation::sign; - sign(self) - } - - /// Clip tensor values to the provided range. - #[inline(always)] - pub fn clip(&self, min_val: Option, max_val: Option) -> Result { - if let (Some(min), Some(max)) = (min_val, max_val) { - if min > max { - return Err(MinitensorError::invalid_argument(format!( - "clip minimum {min} cannot be greater than maximum {max}", - ))); - } - } - - use crate::operations::activation::clip; - clip(self, min_val, max_val) - } - +// Copyright (c) Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; +use crate::{ + autograd::{self}, + error::{MinitensorError, Result}, +}; +use std::sync::Arc; + +#[inline(always)] +fn allclose_f32(a: f32, b: f32, rtol: f32, atol: f32, equal_nan: bool) -> bool { + if a == b { + return true; + } + if equal_nan && a.is_nan() && b.is_nan() { + return true; + } + if !a.is_finite() || !b.is_finite() { + return false; + } + let diff = (a - b).abs(); + diff <= atol + rtol * b.abs() +} + +#[inline(always)] +fn allclose_f64(a: f64, b: f64, rtol: f64, atol: f64, equal_nan: bool) -> bool { + if a == b { + return true; + } + if equal_nan && a.is_nan() && b.is_nan() { + return true; + } + if !a.is_finite() || !b.is_finite() { + return false; + } + let diff = (a - b).abs(); + diff <= atol + rtol * b.abs() +} + +impl Tensor { + /// Solve a linear system `AX = B` for `X` where `self` provides `A`. + pub fn solve(&self, rhs: &Self) -> Result { + use crate::operations::linalg::solve; + solve(self, rhs) + } + + /// Layer normalization + #[inline(always)] + pub fn layer_norm( + &self, + normalized_shape: &[usize], + weight: Option<&Tensor>, + bias: Option<&Tensor>, + eps: f64, + ) -> Result { + use crate::operations::normalization::layer_norm; + layer_norm(self, normalized_shape, weight, bias, eps) + } + + /// Absolute value + #[inline(always)] + pub fn abs(&self) -> Result { + use crate::operations::activation::abs; + abs(self) + } + + /// Element-wise sign (returns -1, 0, or 1 for each value). + #[inline(always)] + pub fn sign(&self) -> Result { + use crate::operations::activation::sign; + sign(self) + } + + /// Clip tensor values to the provided range. + #[inline(always)] + pub fn clip(&self, min_val: Option, max_val: Option) -> Result { + if let (Some(min), Some(max)) = (min_val, max_val) + && min > max + { + return Err(MinitensorError::invalid_argument(format!( + "clip minimum {min} cannot be greater than maximum {max}", + ))); + } + + use crate::operations::activation::clip; + clip(self, min_val, max_val) + } + /// Alias for [`Tensor::clip`]. - #[inline(always)] - pub fn clamp(&self, min_val: Option, max_val: Option) -> Result { - self.clip(min_val, max_val) - } - - /// Clamp tensor values to be no smaller than `min_val`. - #[inline(always)] - pub fn clamp_min(&self, min_val: f64) -> Result { - self.clip(Some(min_val), None) - } - - /// Clamp tensor values to be no larger than `max_val`. - #[inline(always)] - pub fn clamp_max(&self, max_val: f64) -> Result { - self.clip(None, Some(max_val)) - } - - /// Replace NaN with `nan`, positive infinity with `posinf` or dtype max, - /// and negative infinity with `neginf` or dtype min. - #[inline(always)] - pub fn nan_to_num( - &self, - nan: f64, - posinf: Option, - neginf: Option, - ) -> Result { - use crate::operations::activation::nan_to_num; - nan_to_num(self, nan, posinf, neginf) - } - - /// Round tensor values to a specific number of decimal places. - #[inline(always)] - pub fn round(&self, decimals: i32) -> Result { - use crate::operations::activation::round; - round(self, decimals) - } - - /// Floor tensor values element-wise. - #[inline(always)] - pub fn floor(&self) -> Result { - use crate::operations::activation::floor; - floor(self) - } - - /// Ceil tensor values element-wise. - #[inline(always)] - pub fn ceil(&self) -> Result { - use crate::operations::activation::ceil; - ceil(self) - } - - /// Square root - #[inline(always)] - pub fn sqrt(&self) -> Result { - use crate::operations::activation::sqrt; - sqrt(self) - } - - pub fn rsqrt(&self) -> Result { - use crate::operations::activation::rsqrt; - rsqrt(self) - } - - /// Element-wise reciprocal (1/x). - #[inline(always)] - pub fn reciprocal(&self) -> Result { - use crate::operations::activation::reciprocal; - reciprocal(self) - } - - /// Raise tensor elements to a scalar power - #[inline(always)] - pub fn powf(&self, exponent: f64) -> Result { - use crate::operations::activation::powf; - powf(self, exponent) - } - - /// Numerically stable logaddexp - #[inline(always)] - pub fn logaddexp(&self, other: &Tensor) -> Result { - use crate::operations::activation::logaddexp; - logaddexp(self, other) - } - - /// Element-wise power with another tensor - pub fn pow(&self, exponent: &Tensor) -> Result { - use crate::operations::activation::pow; - pow(self, exponent) - } - - /// Move tensor to device - #[inline(always)] - pub fn to(&self, device: Device) -> Result { - if self.device == device { - return Ok(self.clone()); - } - - // For now, just clone the tensor with the new device - // In a full implementation, we'd copy data between devices - let mut new_tensor = self.clone(); - new_tensor.device = device; - Ok(new_tensor) - } - - /// Convert tensor to a different data type - #[inline(always)] - pub fn astype(&self, dtype: DataType) -> Result { - if self.dtype == dtype { - return Ok(self.clone()); - } - - let numel = self.numel(); - let mut new_data = TensorData::zeros_on_device(numel, dtype, self.device); - - // Helper macro to cast between slices using parallel iteration for large buffers - macro_rules! cast { - ($src:expr, $dst:expr, $conv:expr) => {{ - if numel >= 1024 { - $dst.par_iter_mut().zip($src.par_iter()).for_each(|(d, s)| { - *d = $conv(*s); - }); - } else { - for (d, &s) in $dst.iter_mut().zip($src.iter()) { - *d = $conv(s); - } - } - }}; - } - - match (self.dtype, dtype) { - (DataType::Float32, DataType::Float64) => { - let src = self.data.as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from tensor data") - })?; - let dst = new_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from tensor data", - ) - })?; - cast!(src, dst, |v: f32| v as f64); - } - (DataType::Float32, DataType::Int32) => { - let src = self.data.as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from tensor data") - })?; - let dst = new_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable i32 slice from tensor data", - ) - })?; - cast!(src, dst, |v: f32| v as i32); - } - (DataType::Float32, DataType::Int64) => { - let src = self.data.as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from tensor data") - })?; - let dst = new_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable i64 slice from tensor data", - ) - })?; - cast!(src, dst, |v: f32| v as i64); - } - (DataType::Float32, DataType::Bool) => { - let src = self.data.as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f32 slice from tensor data") - })?; - let dst = new_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable bool slice from tensor data", - ) - })?; - cast!(src, dst, |v: f32| v != 0.0); - } - (DataType::Float64, DataType::Float32) => { - let src = self.data.as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from tensor data") - })?; - let dst = new_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from tensor data", - ) - })?; - cast!(src, dst, |v: f64| v as f32); - } - (DataType::Float64, DataType::Int32) => { - let src = self.data.as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from tensor data") - })?; - let dst = new_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable i32 slice from tensor data", - ) - })?; - cast!(src, dst, |v: f64| v as i32); - } - (DataType::Float64, DataType::Int64) => { - let src = self.data.as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from tensor data") - })?; - let dst = new_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable i64 slice from tensor data", - ) - })?; - cast!(src, dst, |v: f64| v as i64); - } - (DataType::Float64, DataType::Bool) => { - let src = self.data.as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get f64 slice from tensor data") - })?; - let dst = new_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable bool slice from tensor data", - ) - })?; - cast!(src, dst, |v: f64| v != 0.0); - } - (DataType::Int32, DataType::Float32) => { - let src = self.data.as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from tensor data") - })?; - let dst = new_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from tensor data", - ) - })?; - cast!(src, dst, |v: i32| v as f32); - } - (DataType::Int32, DataType::Float64) => { - let src = self.data.as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from tensor data") - })?; - let dst = new_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from tensor data", - ) - })?; - cast!(src, dst, |v: i32| v as f64); - } - (DataType::Int32, DataType::Int64) => { - let src = self.data.as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from tensor data") - })?; - let dst = new_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable i64 slice from tensor data", - ) - })?; - cast!(src, dst, |v: i32| v as i64); - } - (DataType::Int32, DataType::Bool) => { - let src = self.data.as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i32 slice from tensor data") - })?; - let dst = new_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable bool slice from tensor data", - ) - })?; - cast!(src, dst, |v: i32| v != 0); - } - (DataType::Int64, DataType::Float32) => { - let src = self.data.as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from tensor data") - })?; - let dst = new_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from tensor data", - ) - })?; - cast!(src, dst, |v: i64| v as f32); - } - (DataType::Int64, DataType::Float64) => { - let src = self.data.as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from tensor data") - })?; - let dst = new_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from tensor data", - ) - })?; - cast!(src, dst, |v: i64| v as f64); - } - (DataType::Int64, DataType::Int32) => { - let src = self.data.as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from tensor data") - })?; - let dst = new_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable i32 slice from tensor data", - ) - })?; - cast!(src, dst, |v: i64| v as i32); - } - (DataType::Int64, DataType::Bool) => { - let src = self.data.as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get i64 slice from tensor data") - })?; - let dst = new_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable bool slice from tensor data", - ) - })?; - cast!(src, dst, |v: i64| v != 0); - } - (DataType::Bool, DataType::Float32) => { - let src = self.data.as_bool_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from tensor data") - })?; - let dst = new_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f32 slice from tensor data", - ) - })?; - cast!(src, dst, |v: bool| if v { 1.0 } else { 0.0 }); - } - (DataType::Bool, DataType::Float64) => { - let src = self.data.as_bool_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from tensor data") - })?; - let dst = new_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable f64 slice from tensor data", - ) - })?; - cast!(src, dst, |v: bool| if v { 1.0 } else { 0.0 }); - } - (DataType::Bool, DataType::Int32) => { - let src = self.data.as_bool_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from tensor data") - })?; - let dst = new_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable i32 slice from tensor data", - ) - })?; - cast!(src, dst, |v: bool| if v { 1 } else { 0 }); - } - (DataType::Bool, DataType::Int64) => { - let src = self.data.as_bool_slice().ok_or_else(|| { - MinitensorError::internal_error("Failed to get bool slice from tensor data") - })?; - let dst = new_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "Failed to get mutable i64 slice from tensor data", - ) - })?; - cast!(src, dst, |v: bool| if v { 1 } else { 0 }); - } - _ => unreachable!("Unhandled dtype conversion"), - } - - Ok(Tensor::new( - Arc::new(new_data), - self.shape.clone(), - dtype, - self.device, - self.requires_grad, - )) - } -} - -impl Tensor { - /// Copy data from ``source`` into this tensor in-place, preserving dtype and device. - pub fn copy_(&mut self, source: &Tensor) -> Result<()> { - if self.shape != *source.shape() { - return Err(MinitensorError::invalid_argument(format!( - "copy_ expected source with shape {:?}, but received {:?}", - self.shape.dims(), - source.shape().dims() - ))); - } - - if !self.device.is_cpu() { - return Err(MinitensorError::invalid_operation( - "copy_ currently supports only CPU tensors".to_string(), - )); - } - - let mut prepared: Cow<'_, Tensor> = Cow::Borrowed(source); - - if prepared.dtype() != self.dtype { - prepared = Cow::Owned(prepared.astype(self.dtype)?); - } - - if prepared.device() != self.device { - prepared = Cow::Owned(prepared.to(self.device)?); - } - - if !prepared.is_contiguous() { - prepared = Cow::Owned(prepared.contiguous()?); - } - - if !self.is_contiguous() { - return Err(MinitensorError::invalid_operation( - "copy_ currently requires the destination tensor to be contiguous".to_string(), - )); - } - - let dtype = self.dtype; - { - let dst_data = self.data_mut(); - match dtype { - DataType::Float32 => { - let dst = dst_data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "failed to obtain mutable float32 slice for copy_".to_string(), - ) - })?; - let src = prepared.data().as_f32_slice().ok_or_else(|| { - MinitensorError::internal_error( - "failed to access float32 source data for copy_".to_string(), - ) - })?; - dst.copy_from_slice(src); - } - DataType::Float64 => { - let dst = dst_data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "failed to obtain mutable float64 slice for copy_".to_string(), - ) - })?; - let src = prepared.data().as_f64_slice().ok_or_else(|| { - MinitensorError::internal_error( - "failed to access float64 source data for copy_".to_string(), - ) - })?; - dst.copy_from_slice(src); - } - DataType::Int32 => { - let dst = dst_data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "failed to obtain mutable int32 slice for copy_".to_string(), - ) - })?; - let src = prepared.data().as_i32_slice().ok_or_else(|| { - MinitensorError::internal_error( - "failed to access int32 source data for copy_".to_string(), - ) - })?; - dst.copy_from_slice(src); - } - DataType::Int64 => { - let dst = dst_data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "failed to obtain mutable int64 slice for copy_".to_string(), - ) - })?; - let src = prepared.data().as_i64_slice().ok_or_else(|| { - MinitensorError::internal_error( - "failed to access int64 source data for copy_".to_string(), - ) - })?; - dst.copy_from_slice(src); - } - DataType::Bool => { - let dst = dst_data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "failed to obtain mutable bool slice for copy_".to_string(), - ) - })?; - let src = prepared.data().as_bool_slice().ok_or_else(|| { - MinitensorError::internal_error( - "failed to access bool source data for copy_".to_string(), - ) - })?; - dst.copy_from_slice(src); - } - } - } - - self.refresh_autograd_metadata(); - if self.requires_grad { - autograd::add_to_graph(self, None)?; - } - - Ok(()) - } - - /// Fill the tensor in-place with ``value`` converted to the tensor dtype. - pub fn fill_(&mut self, value: f64) -> Result<()> { - if !self.device.is_cpu() { - return Err(MinitensorError::invalid_operation( - "fill_ currently supports only CPU tensors".to_string(), - )); - } - - if !self.is_contiguous() { - return Err(MinitensorError::invalid_operation( - "fill_ currently requires contiguous tensors".to_string(), - )); - } - - let dtype = self.dtype; - { - let data = self.data_mut(); - match dtype { - DataType::Float32 => { - let slice = data.as_f32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "failed to obtain mutable float32 slice for fill_".to_string(), - ) - })?; - slice.fill(value as f32); - } - DataType::Float64 => { - let slice = data.as_f64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "failed to obtain mutable float64 slice for fill_".to_string(), - ) - })?; - slice.fill(value); - } - DataType::Int32 => { - let slice = data.as_i32_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "failed to obtain mutable int32 slice for fill_".to_string(), - ) - })?; - slice.fill(value as i32); - } - DataType::Int64 => { - let slice = data.as_i64_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "failed to obtain mutable int64 slice for fill_".to_string(), - ) - })?; - slice.fill(value as i64); - } - DataType::Bool => { - let slice = data.as_bool_slice_mut().ok_or_else(|| { - MinitensorError::internal_error( - "failed to obtain mutable bool slice for fill_".to_string(), - ) - })?; - slice.fill(value != 0.0); - } - } - } - - self.refresh_autograd_metadata(); - if self.requires_grad { - autograd::add_to_graph(self, None)?; - } - - Ok(()) - } -} - -impl Tensor { - /// Detach tensor from computation graph - #[inline(always)] - pub fn detach(&self) -> Self { - let mut detached = self.clone(); - detached.requires_grad = false; - detached.grad_fn = None; - detached.grad = None; - detached - } - - /// Detach tensor from the computation graph in-place - #[inline(always)] - pub fn detach_inplace(&mut self) { - self.requires_grad = false; - self.refresh_autograd_metadata(); - } - - /// Check if tensors are approximately equal - #[inline(always)] - pub fn allclose(&self, other: &Tensor, rtol: f64, atol: f64) -> bool { - self.allclose_with_equal_nan(other, rtol, atol, false) - } - - /// Check if tensors are approximately equal, optionally treating NaNs at - /// matching positions as equal. - #[inline(always)] - pub fn allclose_with_equal_nan( - &self, - other: &Tensor, - rtol: f64, - atol: f64, - equal_nan: bool, - ) -> bool { - if self.shape != other.shape || self.dtype != other.dtype { - return false; - } - - // Fast path: byte-for-byte equality check for contiguous CPU tensors - if (equal_nan || !self.dtype.is_float()) - && self.device.is_cpu() - && other.device.is_cpu() - && self.is_contiguous() - && other.is_contiguous() - { - if let (Some(a), Some(b)) = (self.data.as_bytes(), other.data.as_bytes()) { - if a == b { - return true; - } - } - } - - let numel = self.numel(); - match self.dtype { - DataType::Float32 => { - if let (Some(self_data), Some(other_data)) = - (self.data.as_f32_slice(), other.data.as_f32_slice()) - { - if numel >= 1024 { - self_data - .par_iter() - .zip(other_data.par_iter()) - .all(|(&a, &b)| allclose_f32(a, b, rtol as f32, atol as f32, equal_nan)) - } else { - self_data - .iter() - .zip(other_data.iter()) - .all(|(&a, &b)| allclose_f32(a, b, rtol as f32, atol as f32, equal_nan)) - } - } else { - false - } - } - DataType::Float64 => { - if let (Some(self_data), Some(other_data)) = - (self.data.as_f64_slice(), other.data.as_f64_slice()) - { - if numel >= 1024 { - self_data - .par_iter() - .zip(other_data.par_iter()) - .all(|(&a, &b)| allclose_f64(a, b, rtol, atol, equal_nan)) - } else { - self_data - .iter() - .zip(other_data.iter()) - .all(|(&a, &b)| allclose_f64(a, b, rtol, atol, equal_nan)) - } - } else { - false - } - } - _ => self.array_equal(other), - } - } - - /// Check if tensors are exactly equal - #[inline(always)] - pub fn array_equal(&self, other: &Tensor) -> bool { - if self.shape != other.shape || self.dtype != other.dtype { - return false; - } - - // Fast path for contiguous CPU tensors using raw bytes comparison - if self.device.is_cpu() - && other.device.is_cpu() - && self.is_contiguous() - && other.is_contiguous() - { - if let (Some(a), Some(b)) = (self.data.as_bytes(), other.data.as_bytes()) { - return a == b; - } - } - - let numel = self.numel(); - match self.dtype { - DataType::Float32 => { - if let (Some(self_data), Some(other_data)) = - (self.data.as_f32_slice(), other.data.as_f32_slice()) - { - if numel >= 1024 { - self_data - .par_iter() - .zip(other_data.par_iter()) - .all(|(&a, &b)| a == b) - } else { - self_data == other_data - } - } else { - false - } - } - DataType::Float64 => { - if let (Some(self_data), Some(other_data)) = - (self.data.as_f64_slice(), other.data.as_f64_slice()) - { - if numel >= 1024 { - self_data - .par_iter() - .zip(other_data.par_iter()) - .all(|(&a, &b)| a == b) - } else { - self_data == other_data - } - } else { - false - } - } - DataType::Int32 => { - if let (Some(self_data), Some(other_data)) = - (self.data.as_i32_slice(), other.data.as_i32_slice()) - { - if numel >= 1024 { - self_data - .par_iter() - .zip(other_data.par_iter()) - .all(|(&a, &b)| a == b) - } else { - self_data == other_data - } - } else { - false - } - } - DataType::Int64 => { - if let (Some(self_data), Some(other_data)) = - (self.data.as_i64_slice(), other.data.as_i64_slice()) - { - if numel >= 1024 { - self_data - .par_iter() - .zip(other_data.par_iter()) - .all(|(&a, &b)| a == b) - } else { - self_data == other_data - } - } else { - false - } - } - DataType::Bool => { - if let (Some(self_data), Some(other_data)) = - (self.data.as_bool_slice(), other.data.as_bool_slice()) - { - if numel >= 1024 { - self_data - .par_iter() - .zip(other_data.par_iter()) - .all(|(&a, &b)| a == b) - } else { - self_data == other_data - } - } else { - false - } - } - } - } -} + #[inline(always)] + pub fn clamp(&self, min_val: Option, max_val: Option) -> Result { + self.clip(min_val, max_val) + } + + /// Clamp tensor values to be no smaller than `min_val`. + #[inline(always)] + pub fn clamp_min(&self, min_val: f64) -> Result { + self.clip(Some(min_val), None) + } + + /// Clamp tensor values to be no larger than `max_val`. + #[inline(always)] + pub fn clamp_max(&self, max_val: f64) -> Result { + self.clip(None, Some(max_val)) + } + + /// Replace NaN with `nan`, positive infinity with `posinf` or dtype max, + /// and negative infinity with `neginf` or dtype min. + #[inline(always)] + pub fn nan_to_num(&self, nan: f64, posinf: Option, neginf: Option) -> Result { + use crate::operations::activation::nan_to_num; + nan_to_num(self, nan, posinf, neginf) + } + + /// Round tensor values to a specific number of decimal places. + #[inline(always)] + pub fn round(&self, decimals: i32) -> Result { + use crate::operations::activation::round; + round(self, decimals) + } + + /// Floor tensor values element-wise. + #[inline(always)] + pub fn floor(&self) -> Result { + use crate::operations::activation::floor; + floor(self) + } + + /// Ceil tensor values element-wise. + #[inline(always)] + pub fn ceil(&self) -> Result { + use crate::operations::activation::ceil; + ceil(self) + } + + /// Square root + #[inline(always)] + pub fn sqrt(&self) -> Result { + use crate::operations::activation::sqrt; + sqrt(self) + } + + pub fn rsqrt(&self) -> Result { + use crate::operations::activation::rsqrt; + rsqrt(self) + } + + /// Element-wise reciprocal (1/x). + #[inline(always)] + pub fn reciprocal(&self) -> Result { + use crate::operations::activation::reciprocal; + reciprocal(self) + } + + /// Raise tensor elements to a scalar power + #[inline(always)] + pub fn powf(&self, exponent: f64) -> Result { + use crate::operations::activation::powf; + powf(self, exponent) + } + + /// Numerically stable logaddexp + #[inline(always)] + pub fn logaddexp(&self, other: &Tensor) -> Result { + use crate::operations::activation::logaddexp; + logaddexp(self, other) + } + + /// Element-wise power with another tensor + pub fn pow(&self, exponent: &Tensor) -> Result { + use crate::operations::activation::pow; + pow(self, exponent) + } + + /// Move tensor to device + #[inline(always)] + pub fn to(&self, device: Device) -> Result { + if self.device == device { + return Ok(self.clone()); + } + + // For now, just clone the tensor with the new device + // In a full implementation, we'd copy data between devices + let mut new_tensor = self.clone(); + new_tensor.device = device; + Ok(new_tensor) + } + + /// Convert tensor to a different data type + #[inline(always)] + pub fn astype(&self, dtype: DataType) -> Result { + if self.dtype == dtype { + return Ok(self.clone()); + } + + let numel = self.numel(); + let mut new_data = TensorData::zeros_on_device(numel, dtype, self.device); + + // Helper macro to cast between slices using parallel iteration for large buffers + macro_rules! cast { + ($src:expr, $dst:expr, $conv:expr) => {{ + if numel >= 1024 { + $dst.par_iter_mut().zip($src.par_iter()).for_each(|(d, s)| { + *d = $conv(*s); + }); + } else { + for (d, &s) in $dst.iter_mut().zip($src.iter()) { + *d = $conv(s); + } + } + }}; + } + + match (self.dtype, dtype) { + (DataType::Float32, DataType::Float64) => { + let src = self.data.as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from tensor data") + })?; + let dst = new_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from tensor data", + ) + })?; + cast!(src, dst, |v: f32| v as f64); + } + (DataType::Float32, DataType::Int32) => { + let src = self.data.as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from tensor data") + })?; + let dst = new_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable i32 slice from tensor data", + ) + })?; + cast!(src, dst, |v: f32| v as i32); + } + (DataType::Float32, DataType::Int64) => { + let src = self.data.as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from tensor data") + })?; + let dst = new_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable i64 slice from tensor data", + ) + })?; + cast!(src, dst, |v: f32| v as i64); + } + (DataType::Float32, DataType::Bool) => { + let src = self.data.as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f32 slice from tensor data") + })?; + let dst = new_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable bool slice from tensor data", + ) + })?; + cast!(src, dst, |v: f32| v != 0.0); + } + (DataType::Float64, DataType::Float32) => { + let src = self.data.as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from tensor data") + })?; + let dst = new_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from tensor data", + ) + })?; + cast!(src, dst, |v: f64| v as f32); + } + (DataType::Float64, DataType::Int32) => { + let src = self.data.as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from tensor data") + })?; + let dst = new_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable i32 slice from tensor data", + ) + })?; + cast!(src, dst, |v: f64| v as i32); + } + (DataType::Float64, DataType::Int64) => { + let src = self.data.as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from tensor data") + })?; + let dst = new_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable i64 slice from tensor data", + ) + })?; + cast!(src, dst, |v: f64| v as i64); + } + (DataType::Float64, DataType::Bool) => { + let src = self.data.as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get f64 slice from tensor data") + })?; + let dst = new_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable bool slice from tensor data", + ) + })?; + cast!(src, dst, |v: f64| v != 0.0); + } + (DataType::Int32, DataType::Float32) => { + let src = self.data.as_i32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i32 slice from tensor data") + })?; + let dst = new_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from tensor data", + ) + })?; + cast!(src, dst, |v: i32| v as f32); + } + (DataType::Int32, DataType::Float64) => { + let src = self.data.as_i32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i32 slice from tensor data") + })?; + let dst = new_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from tensor data", + ) + })?; + cast!(src, dst, |v: i32| v as f64); + } + (DataType::Int32, DataType::Int64) => { + let src = self.data.as_i32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i32 slice from tensor data") + })?; + let dst = new_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable i64 slice from tensor data", + ) + })?; + cast!(src, dst, |v: i32| v as i64); + } + (DataType::Int32, DataType::Bool) => { + let src = self.data.as_i32_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i32 slice from tensor data") + })?; + let dst = new_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable bool slice from tensor data", + ) + })?; + cast!(src, dst, |v: i32| v != 0); + } + (DataType::Int64, DataType::Float32) => { + let src = self.data.as_i64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i64 slice from tensor data") + })?; + let dst = new_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from tensor data", + ) + })?; + cast!(src, dst, |v: i64| v as f32); + } + (DataType::Int64, DataType::Float64) => { + let src = self.data.as_i64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i64 slice from tensor data") + })?; + let dst = new_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from tensor data", + ) + })?; + cast!(src, dst, |v: i64| v as f64); + } + (DataType::Int64, DataType::Int32) => { + let src = self.data.as_i64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i64 slice from tensor data") + })?; + let dst = new_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable i32 slice from tensor data", + ) + })?; + cast!(src, dst, |v: i64| v as i32); + } + (DataType::Int64, DataType::Bool) => { + let src = self.data.as_i64_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get i64 slice from tensor data") + })?; + let dst = new_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable bool slice from tensor data", + ) + })?; + cast!(src, dst, |v: i64| v != 0); + } + (DataType::Bool, DataType::Float32) => { + let src = self.data.as_bool_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get bool slice from tensor data") + })?; + let dst = new_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f32 slice from tensor data", + ) + })?; + cast!(src, dst, |v: bool| if v { 1.0 } else { 0.0 }); + } + (DataType::Bool, DataType::Float64) => { + let src = self.data.as_bool_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get bool slice from tensor data") + })?; + let dst = new_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable f64 slice from tensor data", + ) + })?; + cast!(src, dst, |v: bool| if v { 1.0 } else { 0.0 }); + } + (DataType::Bool, DataType::Int32) => { + let src = self.data.as_bool_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get bool slice from tensor data") + })?; + let dst = new_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable i32 slice from tensor data", + ) + })?; + cast!(src, dst, |v: bool| if v { 1 } else { 0 }); + } + (DataType::Bool, DataType::Int64) => { + let src = self.data.as_bool_slice().ok_or_else(|| { + MinitensorError::internal_error("Failed to get bool slice from tensor data") + })?; + let dst = new_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "Failed to get mutable i64 slice from tensor data", + ) + })?; + cast!(src, dst, |v: bool| if v { 1 } else { 0 }); + } + _ => unreachable!("Unhandled dtype conversion"), + } + + Ok(Tensor::new( + Arc::new(new_data), + self.shape.clone(), + dtype, + self.device, + self.requires_grad, + )) + } +} + +impl Tensor { + /// Copy data from ``source`` into this tensor in-place, preserving dtype and device. + pub fn copy_(&mut self, source: &Tensor) -> Result<()> { + if self.shape != *source.shape() { + return Err(MinitensorError::invalid_argument(format!( + "copy_ expected source with shape {:?}, but received {:?}", + self.shape.dims(), + source.shape().dims() + ))); + } + + if !self.device.is_cpu() { + return Err(MinitensorError::invalid_operation( + "copy_ currently supports only CPU tensors".to_string(), + )); + } + + let mut prepared: Cow<'_, Tensor> = Cow::Borrowed(source); + + if prepared.dtype() != self.dtype { + prepared = Cow::Owned(prepared.astype(self.dtype)?); + } + + if prepared.device() != self.device { + prepared = Cow::Owned(prepared.to(self.device)?); + } + + if !prepared.is_contiguous() { + prepared = Cow::Owned(prepared.contiguous()?); + } + + if !self.is_contiguous() { + return Err(MinitensorError::invalid_operation( + "copy_ currently requires the destination tensor to be contiguous".to_string(), + )); + } + + let dtype = self.dtype; + { + let dst_data = self.data_mut(); + match dtype { + DataType::Float32 => { + let dst = dst_data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "failed to obtain mutable float32 slice for copy_".to_string(), + ) + })?; + let src = prepared.data().as_f32_slice().ok_or_else(|| { + MinitensorError::internal_error( + "failed to access float32 source data for copy_".to_string(), + ) + })?; + dst.copy_from_slice(src); + } + DataType::Float64 => { + let dst = dst_data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "failed to obtain mutable float64 slice for copy_".to_string(), + ) + })?; + let src = prepared.data().as_f64_slice().ok_or_else(|| { + MinitensorError::internal_error( + "failed to access float64 source data for copy_".to_string(), + ) + })?; + dst.copy_from_slice(src); + } + DataType::Int32 => { + let dst = dst_data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "failed to obtain mutable int32 slice for copy_".to_string(), + ) + })?; + let src = prepared.data().as_i32_slice().ok_or_else(|| { + MinitensorError::internal_error( + "failed to access int32 source data for copy_".to_string(), + ) + })?; + dst.copy_from_slice(src); + } + DataType::Int64 => { + let dst = dst_data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "failed to obtain mutable int64 slice for copy_".to_string(), + ) + })?; + let src = prepared.data().as_i64_slice().ok_or_else(|| { + MinitensorError::internal_error( + "failed to access int64 source data for copy_".to_string(), + ) + })?; + dst.copy_from_slice(src); + } + DataType::Bool => { + let dst = dst_data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "failed to obtain mutable bool slice for copy_".to_string(), + ) + })?; + let src = prepared.data().as_bool_slice().ok_or_else(|| { + MinitensorError::internal_error( + "failed to access bool source data for copy_".to_string(), + ) + })?; + dst.copy_from_slice(src); + } + } + } + + self.refresh_autograd_metadata(); + if self.requires_grad { + autograd::add_to_graph(self, None)?; + } + + Ok(()) + } + + /// Fill the tensor in-place with ``value`` converted to the tensor dtype. + pub fn fill_(&mut self, value: f64) -> Result<()> { + if !self.device.is_cpu() { + return Err(MinitensorError::invalid_operation( + "fill_ currently supports only CPU tensors".to_string(), + )); + } + + if !self.is_contiguous() { + return Err(MinitensorError::invalid_operation( + "fill_ currently requires contiguous tensors".to_string(), + )); + } + + let dtype = self.dtype; + { + let data = self.data_mut(); + match dtype { + DataType::Float32 => { + let slice = data.as_f32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "failed to obtain mutable float32 slice for fill_".to_string(), + ) + })?; + slice.fill(value as f32); + } + DataType::Float64 => { + let slice = data.as_f64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "failed to obtain mutable float64 slice for fill_".to_string(), + ) + })?; + slice.fill(value); + } + DataType::Int32 => { + let slice = data.as_i32_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "failed to obtain mutable int32 slice for fill_".to_string(), + ) + })?; + slice.fill(value as i32); + } + DataType::Int64 => { + let slice = data.as_i64_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "failed to obtain mutable int64 slice for fill_".to_string(), + ) + })?; + slice.fill(value as i64); + } + DataType::Bool => { + let slice = data.as_bool_slice_mut().ok_or_else(|| { + MinitensorError::internal_error( + "failed to obtain mutable bool slice for fill_".to_string(), + ) + })?; + slice.fill(value != 0.0); + } + } + } + + self.refresh_autograd_metadata(); + if self.requires_grad { + autograd::add_to_graph(self, None)?; + } + + Ok(()) + } +} + +impl Tensor { + /// Detach tensor from computation graph + #[inline(always)] + pub fn detach(&self) -> Self { + let mut detached = self.clone(); + detached.requires_grad = false; + detached.grad_fn = None; + detached.grad = None; + detached + } + + /// Detach tensor from the computation graph in-place + #[inline(always)] + pub fn detach_inplace(&mut self) { + self.requires_grad = false; + self.refresh_autograd_metadata(); + } + + /// Check if tensors are approximately equal + #[inline(always)] + pub fn allclose(&self, other: &Tensor, rtol: f64, atol: f64) -> bool { + self.allclose_with_equal_nan(other, rtol, atol, false) + } + + /// Check if tensors are approximately equal, optionally treating NaNs at + /// matching positions as equal. + #[inline(always)] + pub fn allclose_with_equal_nan( + &self, + other: &Tensor, + rtol: f64, + atol: f64, + equal_nan: bool, + ) -> bool { + if self.shape != other.shape || self.dtype != other.dtype { + return false; + } + + // Fast path: byte-for-byte equality check for contiguous CPU tensors + if (equal_nan || !self.dtype.is_float()) + && self.device.is_cpu() + && other.device.is_cpu() + && self.is_contiguous() + && other.is_contiguous() + && let (Some(a), Some(b)) = (self.data.as_bytes(), other.data.as_bytes()) + && a == b + { + return true; + } + + let numel = self.numel(); + match self.dtype { + DataType::Float32 => { + if let (Some(self_data), Some(other_data)) = + (self.data.as_f32_slice(), other.data.as_f32_slice()) + { + if numel >= 1024 { + self_data + .par_iter() + .zip(other_data.par_iter()) + .all(|(&a, &b)| allclose_f32(a, b, rtol as f32, atol as f32, equal_nan)) + } else { + self_data + .iter() + .zip(other_data.iter()) + .all(|(&a, &b)| allclose_f32(a, b, rtol as f32, atol as f32, equal_nan)) + } + } else { + false + } + } + DataType::Float64 => { + if let (Some(self_data), Some(other_data)) = + (self.data.as_f64_slice(), other.data.as_f64_slice()) + { + if numel >= 1024 { + self_data + .par_iter() + .zip(other_data.par_iter()) + .all(|(&a, &b)| allclose_f64(a, b, rtol, atol, equal_nan)) + } else { + self_data + .iter() + .zip(other_data.iter()) + .all(|(&a, &b)| allclose_f64(a, b, rtol, atol, equal_nan)) + } + } else { + false + } + } + _ => self.array_equal(other), + } + } + + /// Check if tensors are exactly equal + #[inline(always)] + pub fn array_equal(&self, other: &Tensor) -> bool { + if self.shape != other.shape || self.dtype != other.dtype { + return false; + } + + // Fast path for contiguous CPU tensors using raw bytes comparison + if self.device.is_cpu() + && other.device.is_cpu() + && self.is_contiguous() + && other.is_contiguous() + && let (Some(a), Some(b)) = (self.data.as_bytes(), other.data.as_bytes()) + { + return a == b; + } + + let numel = self.numel(); + match self.dtype { + DataType::Float32 => { + if let (Some(self_data), Some(other_data)) = + (self.data.as_f32_slice(), other.data.as_f32_slice()) + { + if numel >= 1024 { + self_data + .par_iter() + .zip(other_data.par_iter()) + .all(|(&a, &b)| a == b) + } else { + self_data == other_data + } + } else { + false + } + } + DataType::Float64 => { + if let (Some(self_data), Some(other_data)) = + (self.data.as_f64_slice(), other.data.as_f64_slice()) + { + if numel >= 1024 { + self_data + .par_iter() + .zip(other_data.par_iter()) + .all(|(&a, &b)| a == b) + } else { + self_data == other_data + } + } else { + false + } + } + DataType::Int32 => { + if let (Some(self_data), Some(other_data)) = + (self.data.as_i32_slice(), other.data.as_i32_slice()) + { + if numel >= 1024 { + self_data + .par_iter() + .zip(other_data.par_iter()) + .all(|(&a, &b)| a == b) + } else { + self_data == other_data + } + } else { + false + } + } + DataType::Int64 => { + if let (Some(self_data), Some(other_data)) = + (self.data.as_i64_slice(), other.data.as_i64_slice()) + { + if numel >= 1024 { + self_data + .par_iter() + .zip(other_data.par_iter()) + .all(|(&a, &b)| a == b) + } else { + self_data == other_data + } + } else { + false + } + } + DataType::Bool => { + if let (Some(self_data), Some(other_data)) = + (self.data.as_bool_slice(), other.data.as_bool_slice()) + { + if numel >= 1024 { + self_data + .par_iter() + .zip(other_data.par_iter()) + .all(|(&a, &b)| a == b) + } else { + self_data == other_data + } + } else { + false + } + } + } + } +} diff --git a/engine/src/tensor/mod/utils.rs b/engine/src/tensor/mod/utils.rs index e0e039c1..16b5fd53 100644 --- a/engine/src/tensor/mod/utils.rs +++ b/engine/src/tensor/mod/utils.rs @@ -1,483 +1,532 @@ -// Copyright (c) 2026 Soumyadip Sarkar. -// All rights reserved. -// -// This source code is licensed under the Apache-style license found in the -// LICENSE file in the root directory of this source tree. - -impl std::fmt::Debug for Tensor { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("Tensor") - .field("shape", &self.shape) - .field("dtype", &self.dtype) - .field("device", &self.device) - .field("requires_grad", &self.requires_grad) - .field("tensor_id", &self.tensor_id) - .field("has_grad_fn", &self.grad_fn.is_some()) - .field("has_grad", &self.grad.is_some()) - .finish() - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::tensor::data::TensorData; - - #[test] - fn test_tensor_creation() { - let shape = Shape::new(vec![2, 3]); - let data = Arc::new(TensorData::zeros(shape.numel(), DataType::Float32)); - let tensor = Tensor::new(data, shape.clone(), DataType::Float32, Device::cpu(), false); - - assert_eq!(tensor.shape(), &shape); - assert_eq!(tensor.dtype(), DataType::Float32); - assert_eq!(tensor.device(), Device::cpu()); - assert!(!tensor.requires_grad()); - assert_eq!(tensor.ndim(), 2); - assert_eq!(tensor.numel(), 6); - } - - #[test] - fn test_tensor_view() { - let shape = Shape::new(vec![2, 3]); - let data = Arc::new(TensorData::zeros(shape.numel(), DataType::Float32)); - let tensor = Tensor::new(data, shape, DataType::Float32, Device::cpu(), false); - - let new_shape = Shape::new(vec![3, 2]); - let reshaped = tensor.view(new_shape.clone()).unwrap(); - assert_eq!(reshaped.shape(), &new_shape); - assert_eq!(reshaped.numel(), 6); - } - - #[test] - fn test_tensor_zeros_and_ones() { - let shape = Shape::new(vec![2, 3]); - - let zeros = Tensor::zeros(shape.clone(), DataType::Float32, Device::cpu(), false); - assert_eq!(zeros.shape(), &shape); - assert_eq!(zeros.dtype(), DataType::Float32); - assert!(!zeros.requires_grad()); - - let ones = Tensor::ones(shape.clone(), DataType::Float32, Device::cpu(), true); - assert_eq!(ones.shape(), &shape); - assert_eq!(ones.dtype(), DataType::Float32); - assert!(ones.requires_grad()); - } - - #[test] - fn test_gradient_management() { - let shape = Shape::new(vec![2, 2]); - let mut tensor = Tensor::zeros(shape.clone(), DataType::Float32, Device::cpu(), true); - - // Initially no gradient - assert!(!tensor.has_grad()); - assert!(tensor.grad().is_none()); - - // Set a gradient - let grad = Tensor::ones(shape.clone(), DataType::Float32, Device::cpu(), false); - tensor.set_grad(Some(grad)); - assert!(tensor.has_grad()); - assert!(tensor.grad().is_some()); - - // Clear gradient (should zero it in place) - tensor.zero_grad(false); - assert!(tensor.has_grad()); - let expected = Tensor::zeros(shape.clone(), DataType::Float32, Device::cpu(), false); - assert!(tensor.grad().unwrap().allclose(&expected, 1e-6, 1e-6)); - } - - #[test] - fn test_gradient_accumulation() { - let shape = Shape::new(vec![2, 2]); - let mut tensor = Tensor::zeros(shape.clone(), DataType::Float32, Device::cpu(), true); - - let grad1 = Tensor::ones(shape.clone(), DataType::Float32, Device::cpu(), false); - let grad2 = Tensor::ones(shape.clone(), DataType::Float32, Device::cpu(), false); - - // Accumulate first gradient - tensor.accumulate_grad(grad1).unwrap(); - assert!(tensor.has_grad()); - - // Accumulate second gradient (should replace for now) - tensor.accumulate_grad(grad2).unwrap(); - assert!(tensor.has_grad()); - } - - #[test] - fn test_backward_scalar_tensor() { - let shape = Shape::new(vec![1]); - let tensor = Tensor::ones(shape, DataType::Float32, Device::cpu(), true); - - // This should work for scalar tensors and produce a gradient - let result = tensor.backward(None); - assert!(result.is_ok()); - } - - #[test] - fn test_backward_non_scalar_error() { - let tensor = Tensor::ones(Shape::new(vec![2]), DataType::Float32, Device::cpu(), true); - let result = tensor.backward(None); - assert!(result.is_err()); - } - - #[test] - fn test_isnan_isinf_isfinite() { - let data = vec![0.0f32, f32::NAN, f32::INFINITY, -5.0]; - let shape = Shape::new(vec![4]); - let tensor = Tensor::new( - Arc::new(TensorData::from_vec_f32(data.clone(), Device::cpu())), - shape.clone(), - DataType::Float32, - Device::cpu(), - false, - ); - - let isnan = tensor.isnan().unwrap(); - let isinf = tensor.isinf().unwrap(); - let isfinite = tensor.isfinite().unwrap(); - - let isnan_data = isnan.data().as_bool_slice().unwrap(); - let isinf_data = isinf.data().as_bool_slice().unwrap(); - let isfinite_data = isfinite.data().as_bool_slice().unwrap(); - - assert_eq!(isnan_data, &[false, true, false, false]); - assert_eq!(isinf_data, &[false, false, true, false]); - assert_eq!(isfinite_data, &[true, false, false, true]); - assert_eq!(isnan.shape(), &shape); - } - - #[test] - fn test_clamp() { - let data = vec![-2.0f32, -0.5, 0.5, 2.0]; - let shape = Shape::new(vec![4]); - let tensor = Tensor::new( - Arc::new(TensorData::from_vec_f32(data.clone(), Device::cpu())), - shape.clone(), - DataType::Float32, - Device::cpu(), - false, - ); - let clamped = tensor.clamp(Some(-1.0), Some(1.0)).unwrap(); - let clamped_data = clamped.data().as_f32_slice().unwrap(); - assert_eq!(clamped_data, &[-1.0, -0.5, 0.5, 1.0]); - assert_eq!(clamped.shape(), &shape); - } - - #[test] - fn test_astype() { - let data = vec![1.5f32, -2.3]; - let shape = Shape::new(vec![2]); - let tensor = Tensor::new( - Arc::new(TensorData::from_vec_f32(data.clone(), Device::cpu())), - shape.clone(), - DataType::Float32, - Device::cpu(), - false, - ); - - let casted = tensor.astype(DataType::Float64).unwrap(); - let casted_data = casted.data().as_f64_slice().unwrap(); - assert!((casted_data[0] - 1.5).abs() < 1e-6); - assert!((casted_data[1] - (-2.3)).abs() < 1e-6); - assert_eq!(casted.shape(), &shape); - - let casted_int = tensor.astype(DataType::Int32).unwrap(); - let casted_int_data = casted_int.data().as_i32_slice().unwrap(); - assert_eq!(casted_int_data, &[1, -2]); - assert_eq!(casted_int.shape(), &shape); - - let casted_bool = tensor.astype(DataType::Bool).unwrap(); - let casted_bool_data = casted_bool.data().as_bool_slice().unwrap(); - assert_eq!(casted_bool_data, &[true, true]); - } - - #[test] - fn test_astype_from_bool() { - let data = vec![true, false, true]; - let shape = Shape::new(vec![3]); - let tensor = Tensor::new( - Arc::new(TensorData::from_vec_bool(data.clone(), Device::cpu())), - shape.clone(), - DataType::Bool, - Device::cpu(), - false, - ); - - let to_float = tensor.astype(DataType::Float32).unwrap(); - assert_eq!(to_float.data().as_f32_slice().unwrap(), &[1.0, 0.0, 1.0]); - - let to_int = tensor.astype(DataType::Int64).unwrap(); - assert_eq!(to_int.data().as_i64_slice().unwrap(), &[1, 0, 1]); - } - - #[test] - fn test_add_scalar_broadcasting() { - let a = Tensor::ones( - Shape::new(vec![2, 3]), - DataType::Float32, - Device::cpu(), - false, - ); - let scalar = Tensor::ones(Shape::scalar(), DataType::Float32, Device::cpu(), false); - let result = a.add(&scalar).unwrap(); - assert_eq!(result.data().as_f32_slice().unwrap(), &[2.0; 6]); - assert_eq!(result.shape(), &Shape::new(vec![2, 3])); - } - - #[test] - fn test_add_incompatible_shapes_error() { - let a = Tensor::ones( - Shape::new(vec![2, 2]), - DataType::Float32, - Device::cpu(), - false, - ); - let b = Tensor::ones( - Shape::new(vec![3, 1]), - DataType::Float32, - Device::cpu(), - false, - ); - assert!(a.add(&b).is_err()); - } - - #[test] - fn test_view_shape_mismatch_error() { - let shape = Shape::new(vec![2, 2]); - let data = Arc::new(TensorData::zeros(shape.numel(), DataType::Float32)); - let tensor = Tensor::new(data, shape, DataType::Float32, Device::cpu(), false); - let bad_shape = Shape::new(vec![3, 1]); - assert!(tensor.view(bad_shape).is_err()); - } - - #[test] - fn test_reshape_scalar_to_vector() { - let scalar = Tensor::ones(Shape::scalar(), DataType::Float32, Device::cpu(), false); - let reshaped = scalar.reshape(Shape::new(vec![1])).unwrap(); - assert_eq!(reshaped.shape().dims(), &[1]); - assert_eq!(reshaped.data().as_f32_slice().unwrap(), &[1.0]); - } - - #[test] - fn test_transpose_basic() { - let data = vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0]; - let tensor = Tensor::new( - Arc::new(TensorData::from_vec_f32(data, Device::cpu())), - Shape::new(vec![2, 3]), - DataType::Float32, - Device::cpu(), - false, - ); - let transposed = tensor.transpose(0, 1).unwrap(); - assert_eq!(transposed.shape().dims(), &[3, 2]); - assert_eq!( - transposed.data().as_f32_slice().unwrap(), - &[1.0, 4.0, 2.0, 5.0, 3.0, 6.0] - ); - } - - #[test] - fn test_transpose_out_of_bounds() { - let tensor = Tensor::ones( - Shape::new(vec![2, 2]), - DataType::Float32, - Device::cpu(), - false, - ); - assert!(tensor.transpose(0, 2).is_err()); - } - - #[test] - fn test_transpose_same_dim_noop() { - let tensor = Tensor::ones( - Shape::new(vec![2, 2]), - DataType::Float32, - Device::cpu(), - false, - ); - let transposed = tensor.transpose(1, 1).unwrap(); - assert_eq!(transposed.data().as_f32_slice().unwrap(), &[1.0; 4]); - assert_eq!(transposed.shape().dims(), &[2, 2]); - } - - #[test] - fn test_astype_multiple_conversions() { - let base = Tensor::new( - Arc::new(TensorData::from_vec_f32( - vec![1.5, -2.0, 0.0], - Device::cpu(), - )), - Shape::new(vec![3]), - DataType::Float32, - Device::cpu(), - false, - ); - - let as_i32 = base.astype(DataType::Int32).unwrap(); - assert_eq!(as_i32.data().as_i32_slice().unwrap(), &[1, -2, 0]); - - let as_bool = as_i32.astype(DataType::Bool).unwrap(); - assert_eq!( - as_bool.data().as_bool_slice().unwrap(), - &[true, true, false] - ); - - let as_f64 = as_bool.astype(DataType::Float64).unwrap(); - assert_eq!(as_f64.data().as_f64_slice().unwrap(), &[1.0, 1.0, 0.0]); - } - - #[test] - fn test_astype_parallel_large_buffer() { - let size = 2048; - let data: Vec = (0..size).map(|v| v as f32).collect(); - let tensor = Tensor::new( - Arc::new(TensorData::from_vec_f32(data, Device::cpu())), - Shape::new(vec![size]), - DataType::Float32, - Device::cpu(), - false, - ); - - let converted = tensor.astype(DataType::Int32).unwrap(); - let expected: Vec = (0..size).map(|v| v as i32).collect(); - assert_eq!( - converted.data().as_i32_slice().unwrap(), - expected.as_slice() - ); - } - - #[test] - fn test_array_equal_fast_path() { - let t1 = Tensor::new( - Arc::new(TensorData::from_vec_f32(vec![1.0, 2.0, 3.0], Device::cpu())), - Shape::new(vec![3]), - DataType::Float32, - Device::cpu(), - false, - ); - let t2 = Tensor::new( - Arc::new(TensorData::from_vec_f32(vec![1.0, 2.0, 3.0], Device::cpu())), - Shape::new(vec![3]), - DataType::Float32, - Device::cpu(), - false, - ); - assert!(t1.array_equal(&t2)); - assert!(t1.allclose(&t2, 0.0, 0.0)); - } - - #[test] - fn test_array_equal_mismatch() { - let t1 = Tensor::new( - Arc::new(TensorData::from_vec_f32(vec![1.0, 2.0], Device::cpu())), - Shape::new(vec![2]), - DataType::Float32, - Device::cpu(), - false, - ); - let t2 = Tensor::new( - Arc::new(TensorData::from_vec_f32(vec![1.0, 2.1], Device::cpu())), - Shape::new(vec![2]), - DataType::Float32, - Device::cpu(), - false, - ); - assert!(!t1.array_equal(&t2)); - assert!(!t1.allclose(&t2, 1e-5, 1e-5)); - } - - #[test] - fn test_array_equal_zero_sized() { - let empty1 = Tensor::new( - Arc::new(TensorData::from_vec_f32(vec![], Device::cpu())), - Shape::new(vec![0]), - DataType::Float32, - Device::cpu(), - false, - ); - let empty2 = Tensor::new( - Arc::new(TensorData::from_vec_f32(vec![], Device::cpu())), - Shape::new(vec![0]), - DataType::Float32, - Device::cpu(), - false, - ); - assert!(empty1.array_equal(&empty2)); - } - - #[test] - fn test_deep_clone_independent_storage() { - let shape = Shape::new(vec![2, 2]); - let data = Arc::new(TensorData::from_vec_f32( - vec![1.0, 2.0, 3.0, 4.0], - Device::cpu(), - )); - let tensor = Tensor::new(data, shape.clone(), DataType::Float32, Device::cpu(), false); - - let mut cloned = tensor.deep_clone().unwrap(); - { - let slice = cloned.data_mut().as_f32_slice_mut().unwrap(); - slice[0] = 42.0; - } - - let original_slice = tensor.data().as_f32_slice().unwrap(); - assert_eq!(original_slice, &[1.0, 2.0, 3.0, 4.0]); - let cloned_slice = cloned.data().as_f32_slice().unwrap(); - assert_eq!(cloned_slice, &[42.0, 2.0, 3.0, 4.0]); - } - - #[test] - fn test_deep_clone_preserves_gradients() { - let shape = Shape::new(vec![3]); - let data = Arc::new(TensorData::from_vec_f32( - vec![1.0, -2.0, 3.0], - Device::cpu(), - )); - let mut tensor = Tensor::new(data, shape.clone(), DataType::Float32, Device::cpu(), true); - tensor.zero_grad(true); - - let cloned = tensor.deep_clone().unwrap(); - assert!(cloned.requires_grad()); - - let grad = Tensor::new( - Arc::new(TensorData::from_vec_f32( - vec![0.5, -1.0, 2.0], - Device::cpu(), - )), - shape, - DataType::Float32, - Device::cpu(), - false, - ); - cloned.backward(Some(grad.clone())).unwrap(); - - let accumulated = autograd::get_gradient(&tensor).expect("gradient should be set"); - assert!(accumulated.allclose(&grad, 1e-6, 1e-6)); - } - - #[test] - fn test_contiguous_materialises_expanded_views() { - let base = Tensor::new( - Arc::new(TensorData::from_vec_f32(vec![1.0, 2.0], Device::cpu())), - Shape::new(vec![2, 1]), - DataType::Float32, - Device::cpu(), - false, - ); - - let expanded = base - .expand(vec![2isize, 3isize]) - .expect("expand should succeed"); - assert!(!expanded.is_contiguous()); - - let contiguous = expanded.contiguous().expect("contiguous should copy data"); - assert!(contiguous.is_contiguous()); - assert_eq!(contiguous.shape().dims(), &[2, 3]); - let values = contiguous - .data() - .as_f32_slice() - .expect("materialised data should be accessible") - .to_vec(); - assert_eq!(values, vec![1.0, 1.0, 1.0, 2.0, 2.0, 2.0]); - } -} +// Copyright (c) 2026 Soumyadip Sarkar. +// All rights reserved. +// +// This source code is licensed under the Apache-style license found in the +// LICENSE file in the root directory of this source tree. + +use super::*; + +impl std::fmt::Debug for Tensor { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("Tensor") + .field("shape", &self.shape) + .field("dtype", &self.dtype) + .field("device", &self.device) + .field("requires_grad", &self.requires_grad) + .field("tensor_id", &self.tensor_id) + .field("has_grad_fn", &self.grad_fn.is_some()) + .field("has_grad", &self.grad.is_some()) + .finish() + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::tensor::data::TensorData; + + #[test] + fn test_tensor_creation() { + let shape = Shape::new(vec![2, 3]); + let data = Arc::new(TensorData::zeros(shape.numel(), DataType::Float32)); + let tensor = Tensor::new(data, shape.clone(), DataType::Float32, Device::cpu(), false); + + assert_eq!(tensor.shape(), &shape); + assert_eq!(tensor.dtype(), DataType::Float32); + assert_eq!(tensor.device(), Device::cpu()); + assert!(!tensor.requires_grad()); + assert_eq!(tensor.ndim(), 2); + assert_eq!(tensor.numel(), 6); + } + + #[test] + fn test_tensor_view() { + let shape = Shape::new(vec![2, 3]); + let data = Arc::new(TensorData::zeros(shape.numel(), DataType::Float32)); + let tensor = Tensor::new(data, shape, DataType::Float32, Device::cpu(), false); + + let new_shape = Shape::new(vec![3, 2]); + let reshaped = tensor.view(new_shape.clone()).unwrap(); + assert_eq!(reshaped.shape(), &new_shape); + assert_eq!(reshaped.numel(), 6); + } + + #[test] + fn test_tensor_zeros_and_ones() { + let shape = Shape::new(vec![2, 3]); + + let zeros = Tensor::zeros(shape.clone(), DataType::Float32, Device::cpu(), false); + assert_eq!(zeros.shape(), &shape); + assert_eq!(zeros.dtype(), DataType::Float32); + assert!(!zeros.requires_grad()); + + let ones = Tensor::ones(shape.clone(), DataType::Float32, Device::cpu(), true); + assert_eq!(ones.shape(), &shape); + assert_eq!(ones.dtype(), DataType::Float32); + assert!(ones.requires_grad()); + } + + #[test] + fn test_gradient_management() { + let shape = Shape::new(vec![2, 2]); + let mut tensor = Tensor::zeros(shape.clone(), DataType::Float32, Device::cpu(), true); + + // Initially no gradient + assert!(!tensor.has_grad()); + assert!(tensor.grad().is_none()); + + // Set a gradient + let grad = Tensor::ones(shape.clone(), DataType::Float32, Device::cpu(), false); + tensor.set_grad(Some(grad)); + assert!(tensor.has_grad()); + assert!(tensor.grad().is_some()); + + // Clear gradient (should zero it in place) + tensor.zero_grad(false); + assert!(tensor.has_grad()); + let expected = Tensor::zeros(shape.clone(), DataType::Float32, Device::cpu(), false); + assert!(tensor.grad().unwrap().allclose(&expected, 1e-6, 1e-6)); + } + + #[test] + fn test_gradient_accumulation() { + let shape = Shape::new(vec![2, 2]); + let mut tensor = Tensor::zeros(shape.clone(), DataType::Float32, Device::cpu(), true); + + let grad1 = Tensor::ones(shape.clone(), DataType::Float32, Device::cpu(), false); + let grad2 = Tensor::ones(shape.clone(), DataType::Float32, Device::cpu(), false); + + // Accumulate first gradient + tensor.accumulate_grad(grad1).unwrap(); + assert!(tensor.has_grad()); + + // Accumulate second gradient (should replace for now) + tensor.accumulate_grad(grad2).unwrap(); + assert!(tensor.has_grad()); + } + + #[test] + fn test_backward_scalar_tensor() { + let shape = Shape::new(vec![1]); + let tensor = Tensor::ones(shape, DataType::Float32, Device::cpu(), true); + + // This should work for scalar tensors and produce a gradient + let result = tensor.backward(None); + assert!(result.is_ok()); + } + + #[test] + fn test_backward_non_scalar_error() { + let tensor = Tensor::ones(Shape::new(vec![2]), DataType::Float32, Device::cpu(), true); + let result = tensor.backward(None); + assert!(result.is_err()); + } + + #[test] + fn test_isnan_isinf_isfinite() { + let data = vec![0.0f32, f32::NAN, f32::INFINITY, -5.0]; + let shape = Shape::new(vec![4]); + let tensor = Tensor::new( + Arc::new(TensorData::from_vec_f32(data.clone(), Device::cpu())), + shape.clone(), + DataType::Float32, + Device::cpu(), + false, + ); + + let isnan = tensor.isnan().unwrap(); + let isinf = tensor.isinf().unwrap(); + let isfinite = tensor.isfinite().unwrap(); + + let isnan_data = isnan.data().as_bool_slice().unwrap(); + let isinf_data = isinf.data().as_bool_slice().unwrap(); + let isfinite_data = isfinite.data().as_bool_slice().unwrap(); + + assert_eq!(isnan_data, &[false, true, false, false]); + assert_eq!(isinf_data, &[false, false, true, false]); + assert_eq!(isfinite_data, &[true, false, false, true]); + assert_eq!(isnan.shape(), &shape); + } + + #[test] + fn test_clamp() { + let data = vec![-2.0f32, -0.5, 0.5, 2.0]; + let shape = Shape::new(vec![4]); + let tensor = Tensor::new( + Arc::new(TensorData::from_vec_f32(data.clone(), Device::cpu())), + shape.clone(), + DataType::Float32, + Device::cpu(), + false, + ); + let clamped = tensor.clamp(Some(-1.0), Some(1.0)).unwrap(); + let clamped_data = clamped.data().as_f32_slice().unwrap(); + assert_eq!(clamped_data, &[-1.0, -0.5, 0.5, 1.0]); + assert_eq!(clamped.shape(), &shape); + } + + #[test] + fn test_astype() { + let data = vec![1.5f32, -2.3]; + let shape = Shape::new(vec![2]); + let tensor = Tensor::new( + Arc::new(TensorData::from_vec_f32(data.clone(), Device::cpu())), + shape.clone(), + DataType::Float32, + Device::cpu(), + false, + ); + + let casted = tensor.astype(DataType::Float64).unwrap(); + let casted_data = casted.data().as_f64_slice().unwrap(); + assert!((casted_data[0] - 1.5).abs() < 1e-6); + assert!((casted_data[1] - (-2.3)).abs() < 1e-6); + assert_eq!(casted.shape(), &shape); + + let casted_int = tensor.astype(DataType::Int32).unwrap(); + let casted_int_data = casted_int.data().as_i32_slice().unwrap(); + assert_eq!(casted_int_data, &[1, -2]); + assert_eq!(casted_int.shape(), &shape); + + let casted_bool = tensor.astype(DataType::Bool).unwrap(); + let casted_bool_data = casted_bool.data().as_bool_slice().unwrap(); + assert_eq!(casted_bool_data, &[true, true]); + } + + #[test] + fn test_astype_from_bool() { + let data = vec![true, false, true]; + let shape = Shape::new(vec![3]); + let tensor = Tensor::new( + Arc::new(TensorData::from_vec_bool(data.clone(), Device::cpu())), + shape.clone(), + DataType::Bool, + Device::cpu(), + false, + ); + + let to_float = tensor.astype(DataType::Float32).unwrap(); + assert_eq!(to_float.data().as_f32_slice().unwrap(), &[1.0, 0.0, 1.0]); + + let to_int = tensor.astype(DataType::Int64).unwrap(); + assert_eq!(to_int.data().as_i64_slice().unwrap(), &[1, 0, 1]); + } + + #[test] + fn test_add_scalar_broadcasting() { + let a = Tensor::ones( + Shape::new(vec![2, 3]), + DataType::Float32, + Device::cpu(), + false, + ); + let scalar = Tensor::ones(Shape::scalar(), DataType::Float32, Device::cpu(), false); + let result = a.add(&scalar).unwrap(); + assert_eq!(result.data().as_f32_slice().unwrap(), &[2.0; 6]); + assert_eq!(result.shape(), &Shape::new(vec![2, 3])); + } + + #[test] + fn test_add_incompatible_shapes_error() { + let a = Tensor::ones( + Shape::new(vec![2, 2]), + DataType::Float32, + Device::cpu(), + false, + ); + let b = Tensor::ones( + Shape::new(vec![3, 1]), + DataType::Float32, + Device::cpu(), + false, + ); + assert!(a.add(&b).is_err()); + } + + #[test] + fn test_view_shape_mismatch_error() { + let shape = Shape::new(vec![2, 2]); + let data = Arc::new(TensorData::zeros(shape.numel(), DataType::Float32)); + let tensor = Tensor::new(data, shape, DataType::Float32, Device::cpu(), false); + let bad_shape = Shape::new(vec![3, 1]); + assert!(tensor.view(bad_shape).is_err()); + } + + #[test] + fn test_view_rejects_non_contiguous_tensor() { + let data = TensorData::from_vec_f32(vec![1.0, 2.0, 3.0], Device::cpu()); + let tensor = Tensor::new( + Arc::new(data), + Shape::new(vec![1, 3]), + DataType::Float32, + Device::cpu(), + false, + ); + let expanded = tensor.expand(vec![4, 3]).unwrap(); + assert!(!expanded.is_contiguous()); + // A raw view would silently pair the new shape with storage that only + // holds 3 elements; it must be rejected. + assert!(expanded.view(Shape::new(vec![12])).is_err()); + } + + #[test] + fn test_reshape_materializes_non_contiguous_tensor() { + let data = TensorData::from_vec_f32(vec![1.0, 2.0, 3.0], Device::cpu()); + let tensor = Tensor::new( + Arc::new(data), + Shape::new(vec![1, 3]), + DataType::Float32, + Device::cpu(), + false, + ); + let expanded = tensor.expand(vec![4, 3]).unwrap(); + let reshaped = expanded.reshape(Shape::new(vec![12])).unwrap(); + assert_eq!(reshaped.shape().dims(), &[12]); + // Storage must now really contain all 12 broadcast elements. + assert_eq!(reshaped.data().numel(), 12); + assert_eq!( + reshaped.data().as_f32_slice().unwrap(), + &[1.0, 2.0, 3.0, 1.0, 2.0, 3.0, 1.0, 2.0, 3.0, 1.0, 2.0, 3.0] + ); + + // The ops-layer reshape must materialise as well. + let via_op = + crate::operations::shape_ops::reshape(&expanded, Shape::new(vec![12])).unwrap(); + assert_eq!(via_op.data().numel(), 12); + assert_eq!( + via_op.data().as_f32_slice().unwrap(), + &[1.0, 2.0, 3.0, 1.0, 2.0, 3.0, 1.0, 2.0, 3.0, 1.0, 2.0, 3.0] + ); + } + + #[test] + fn test_reshape_scalar_to_vector() { + let scalar = Tensor::ones(Shape::scalar(), DataType::Float32, Device::cpu(), false); + let reshaped = scalar.reshape(Shape::new(vec![1])).unwrap(); + assert_eq!(reshaped.shape().dims(), &[1]); + assert_eq!(reshaped.data().as_f32_slice().unwrap(), &[1.0]); + } + + #[test] + fn test_transpose_basic() { + let data = vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0]; + let tensor = Tensor::new( + Arc::new(TensorData::from_vec_f32(data, Device::cpu())), + Shape::new(vec![2, 3]), + DataType::Float32, + Device::cpu(), + false, + ); + let transposed = tensor.transpose(0, 1).unwrap(); + assert_eq!(transposed.shape().dims(), &[3, 2]); + assert_eq!( + transposed.data().as_f32_slice().unwrap(), + &[1.0, 4.0, 2.0, 5.0, 3.0, 6.0] + ); + } + + #[test] + fn test_transpose_out_of_bounds() { + let tensor = Tensor::ones( + Shape::new(vec![2, 2]), + DataType::Float32, + Device::cpu(), + false, + ); + assert!(tensor.transpose(0, 2).is_err()); + } + + #[test] + fn test_transpose_same_dim_noop() { + let tensor = Tensor::ones( + Shape::new(vec![2, 2]), + DataType::Float32, + Device::cpu(), + false, + ); + let transposed = tensor.transpose(1, 1).unwrap(); + assert_eq!(transposed.data().as_f32_slice().unwrap(), &[1.0; 4]); + assert_eq!(transposed.shape().dims(), &[2, 2]); + } + + #[test] + fn test_astype_multiple_conversions() { + let base = Tensor::new( + Arc::new(TensorData::from_vec_f32( + vec![1.5, -2.0, 0.0], + Device::cpu(), + )), + Shape::new(vec![3]), + DataType::Float32, + Device::cpu(), + false, + ); + + let as_i32 = base.astype(DataType::Int32).unwrap(); + assert_eq!(as_i32.data().as_i32_slice().unwrap(), &[1, -2, 0]); + + let as_bool = as_i32.astype(DataType::Bool).unwrap(); + assert_eq!( + as_bool.data().as_bool_slice().unwrap(), + &[true, true, false] + ); + + let as_f64 = as_bool.astype(DataType::Float64).unwrap(); + assert_eq!(as_f64.data().as_f64_slice().unwrap(), &[1.0, 1.0, 0.0]); + } + + #[test] + fn test_astype_parallel_large_buffer() { + let size = 2048; + let data: Vec = (0..size).map(|v| v as f32).collect(); + let tensor = Tensor::new( + Arc::new(TensorData::from_vec_f32(data, Device::cpu())), + Shape::new(vec![size]), + DataType::Float32, + Device::cpu(), + false, + ); + + let converted = tensor.astype(DataType::Int32).unwrap(); + let expected: Vec = (0..size).map(|v| v as i32).collect(); + assert_eq!( + converted.data().as_i32_slice().unwrap(), + expected.as_slice() + ); + } + + #[test] + fn test_array_equal_fast_path() { + let t1 = Tensor::new( + Arc::new(TensorData::from_vec_f32(vec![1.0, 2.0, 3.0], Device::cpu())), + Shape::new(vec![3]), + DataType::Float32, + Device::cpu(), + false, + ); + let t2 = Tensor::new( + Arc::new(TensorData::from_vec_f32(vec![1.0, 2.0, 3.0], Device::cpu())), + Shape::new(vec![3]), + DataType::Float32, + Device::cpu(), + false, + ); + assert!(t1.array_equal(&t2)); + assert!(t1.allclose(&t2, 0.0, 0.0)); + } + + #[test] + fn test_array_equal_mismatch() { + let t1 = Tensor::new( + Arc::new(TensorData::from_vec_f32(vec![1.0, 2.0], Device::cpu())), + Shape::new(vec![2]), + DataType::Float32, + Device::cpu(), + false, + ); + let t2 = Tensor::new( + Arc::new(TensorData::from_vec_f32(vec![1.0, 2.1], Device::cpu())), + Shape::new(vec![2]), + DataType::Float32, + Device::cpu(), + false, + ); + assert!(!t1.array_equal(&t2)); + assert!(!t1.allclose(&t2, 1e-5, 1e-5)); + } + + #[test] + fn test_array_equal_zero_sized() { + let empty1 = Tensor::new( + Arc::new(TensorData::from_vec_f32(vec![], Device::cpu())), + Shape::new(vec![0]), + DataType::Float32, + Device::cpu(), + false, + ); + let empty2 = Tensor::new( + Arc::new(TensorData::from_vec_f32(vec![], Device::cpu())), + Shape::new(vec![0]), + DataType::Float32, + Device::cpu(), + false, + ); + assert!(empty1.array_equal(&empty2)); + } + + #[test] + fn test_deep_clone_independent_storage() { + let shape = Shape::new(vec![2, 2]); + let data = Arc::new(TensorData::from_vec_f32( + vec![1.0, 2.0, 3.0, 4.0], + Device::cpu(), + )); + let tensor = Tensor::new(data, shape.clone(), DataType::Float32, Device::cpu(), false); + + let mut cloned = tensor.deep_clone().unwrap(); + { + let slice = cloned.data_mut().as_f32_slice_mut().unwrap(); + slice[0] = 42.0; + } + + let original_slice = tensor.data().as_f32_slice().unwrap(); + assert_eq!(original_slice, &[1.0, 2.0, 3.0, 4.0]); + let cloned_slice = cloned.data().as_f32_slice().unwrap(); + assert_eq!(cloned_slice, &[42.0, 2.0, 3.0, 4.0]); + } + + #[test] + fn test_deep_clone_preserves_gradients() { + let shape = Shape::new(vec![3]); + let data = Arc::new(TensorData::from_vec_f32( + vec![1.0, -2.0, 3.0], + Device::cpu(), + )); + let mut tensor = Tensor::new(data, shape.clone(), DataType::Float32, Device::cpu(), true); + tensor.zero_grad(true); + + let cloned = tensor.deep_clone().unwrap(); + assert!(cloned.requires_grad()); + + let grad = Tensor::new( + Arc::new(TensorData::from_vec_f32( + vec![0.5, -1.0, 2.0], + Device::cpu(), + )), + shape, + DataType::Float32, + Device::cpu(), + false, + ); + cloned.backward(Some(grad.clone())).unwrap(); + + let accumulated = autograd::get_gradient(&tensor).expect("gradient should be set"); + assert!(accumulated.allclose(&grad, 1e-6, 1e-6)); + } + + #[test] + fn test_contiguous_materialises_expanded_views() { + let base = Tensor::new( + Arc::new(TensorData::from_vec_f32(vec![1.0, 2.0], Device::cpu())), + Shape::new(vec![2, 1]), + DataType::Float32, + Device::cpu(), + false, + ); + + let expanded = base + .expand(vec![2isize, 3isize]) + .expect("expand should succeed"); + assert!(!expanded.is_contiguous()); + + let contiguous = expanded.contiguous().expect("contiguous should copy data"); + assert!(contiguous.is_contiguous()); + assert_eq!(contiguous.shape().dims(), &[2, 3]); + let values = contiguous + .data() + .as_f32_slice() + .expect("materialised data should be accessible") + .to_vec(); + assert_eq!(values, vec![1.0, 1.0, 1.0, 2.0, 2.0, 2.0]); + } +} diff --git a/engine/src/tensor/shape.rs b/engine/src/tensor/shape.rs index db1aec64..014a644c 100644 --- a/engine/src/tensor/shape.rs +++ b/engine/src/tensor/shape.rs @@ -42,14 +42,24 @@ impl Shape { self.dims.len() } - /// Get the total number of elements + /// Get the total number of elements. + /// + /// The product is computed with checked arithmetic in every build + /// profile: a silently wrapped element count would under-allocate + /// storage while indexing code still trusts the individual dimensions, + /// turning an absurd shape into out-of-bounds reads/writes instead of a + /// clean failure. #[inline(always)] pub fn numel(&self) -> usize { - if self.dims.is_empty() { - 1 // scalar - } else { - self.dims.iter().product() - } + self.dims + .iter() + .try_fold(1usize, |acc, &d| acc.checked_mul(d)) + .unwrap_or_else(|| { + panic!( + "tensor shape {:?} has more elements than usize can represent", + self.dims + ) + }) } /// Get the size of a specific dimension @@ -165,11 +175,16 @@ impl Strides { #[inline(always)] pub fn from_shape(shape: &Shape) -> Self { let mut strides = Vec::with_capacity(shape.ndim()); - let mut stride = 1; + let mut stride = 1usize; for &dim in shape.dims().iter().rev() { strides.push(stride); - stride *= dim; + stride = stride.checked_mul(dim).unwrap_or_else(|| { + panic!( + "tensor shape {:?} has more elements than usize can represent", + shape.dims() + ) + }); } strides.reverse(); @@ -238,6 +253,22 @@ mod tests { assert!(scalar.is_scalar()); } + #[test] + #[should_panic(expected = "more elements than usize can represent")] + fn test_numel_overflow_panics() { + // A wrapped product would report a tiny element count for an absurd + // shape and let downstream code under-allocate storage. + let shape = Shape::new(vec![usize::MAX, 2]); + let _ = shape.numel(); + } + + #[test] + #[should_panic(expected = "more elements than usize can represent")] + fn test_strides_overflow_panics() { + let shape = Shape::new(vec![usize::MAX, 4, 2]); + let _ = Strides::from_shape(&shape); + } + #[test] fn test_broadcasting() { let shape1 = Shape::new(vec![3, 1]); @@ -258,7 +289,7 @@ mod tests { assert!(strides.is_contiguous(&shape)); let linear_idx = strides.linear_index(&[1, 2, 3]); - assert_eq!(linear_idx, 1 * 12 + 2 * 4 + 3 * 1); + assert_eq!(linear_idx, 12 + 2 * 4 + 3); } #[test] diff --git a/engine/tests/gradient_tests.rs b/engine/tests/gradient_tests.rs index 743db1d3..541cdd2d 100644 --- a/engine/tests/gradient_tests.rs +++ b/engine/tests/gradient_tests.rs @@ -41,7 +41,7 @@ fn test_mul_backward_correct() { Device::cpu(), false, ); - let grads = autograd::backward(&product, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&product, Some(grad_output)).unwrap(); let grad_a = grads.get(&a.id()).unwrap(); let grad_b = grads.get(&b.id()).unwrap(); assert_eq!(grad_a.data().as_f32_slice().unwrap(), &[3.0, 4.0]); @@ -61,7 +61,7 @@ fn test_sub_backward_correct() { Device::cpu(), false, ); - let grads = autograd::backward(&diff, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&diff, Some(grad_output)).unwrap(); let grad_a = grads.get(&a.id()).unwrap(); let grad_b = grads.get(&b.id()).unwrap(); assert_eq!(grad_a.data().as_f32_slice().unwrap(), &[1.0, 1.0]); @@ -76,7 +76,7 @@ fn test_div_backward_correct() { let b = create_test_tensor_f32(vec![2.0, 3.0], vec![2], true); let quo = arithmetic::div(&a, &b).unwrap(); let grad_output = Tensor::ones(quo.shape().clone(), DataType::Float32, Device::cpu(), false); - let grads = autograd::backward(&quo, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&quo, Some(grad_output)).unwrap(); let grad_a = grads.get(&a.id()).unwrap(); let grad_b = grads.get(&b.id()).unwrap(); let ga = grad_a.data().as_f32_slice().unwrap(); @@ -94,7 +94,7 @@ fn test_neg_backward_correct() { let x = create_test_tensor_f32(vec![1.0, -2.0], vec![2], true); let y = arithmetic::neg(&x).unwrap(); let grad_output = Tensor::ones(y.shape().clone(), DataType::Float32, Device::cpu(), false); - let grads = autograd::backward(&y, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&y, Some(grad_output)).unwrap(); let grad_x = grads.get(&x.id()).unwrap(); assert_eq!(grad_x.data().as_f32_slice().unwrap(), &[-1.0, -1.0]); autograd::clear_graph().unwrap(); @@ -111,7 +111,7 @@ fn test_cos_backward_correct() { Device::cpu(), false, ); - let grads = autograd::backward(&output, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&output, Some(grad_output)).unwrap(); let grad_input = grads.get(&input.id()).unwrap(); let vals = grad_input.data().as_f32_slice().unwrap(); assert_relative_eq!(vals[0], 0.0, epsilon = 1e-6); @@ -158,7 +158,7 @@ fn test_logsumexp_backward_matches_softmax() { Device::cpu(), false, ); - let grads = autograd::backward(&output, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&output, Some(grad_output)).unwrap(); let grad_input = grads.get(&input.id()).unwrap(); let grad_vals = grad_input.data().as_f32_slice().unwrap(); @@ -206,7 +206,7 @@ fn test_log_softmax_backward_matches_manual() { let output = activation::log_softmax(&input, Some(1)).unwrap(); let grad_output = create_test_tensor_f32(grad_values.clone(), vec![2, 3], false); - let grads = autograd::backward(&output, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&output, Some(grad_output)).unwrap(); let grad_input = grads.get(&input.id()).unwrap(); let grad_vals = grad_input.data().as_f32_slice().unwrap(); @@ -242,7 +242,7 @@ fn test_z_leaky_relu_backward_correct() { Device::cpu(), false, ); - let grads = autograd::backward(&output, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&output, Some(grad_output)).unwrap(); let grad_input = grads.get(&input.id()).unwrap(); assert_eq!(grad_input.data().as_f32_slice().unwrap(), &[0.1, 1.0, 1.0]); autograd::clear_graph().unwrap(); @@ -254,7 +254,7 @@ fn test_relu_backward_nan_propagates() { let input = create_test_tensor_f32(vec![-1.0, f32::NAN, 1.0], vec![3], true); let output = activation::relu(&input).unwrap(); let grad_output = create_test_tensor_f32(vec![1.0, f32::NAN, 1.0], vec![3], false); - let grads = autograd::backward(&output, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&output, Some(grad_output)).unwrap(); let grad_input = grads.get(&input.id()).unwrap(); let vals = grad_input.data().as_f32_slice().unwrap(); assert_eq!(vals[0], 0.0); @@ -269,7 +269,7 @@ fn test_sum_backward_correct() { let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0], vec![3], true); let s = reduction::sum(&a, None, false).unwrap(); let grad_output = Tensor::ones(s.shape().clone(), DataType::Float32, Device::cpu(), false); - let grads = autograd::backward(&s, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&s, Some(grad_output)).unwrap(); let grad_a = grads.get(&a.id()).unwrap(); assert_eq!(grad_a.data().as_f32_slice().unwrap(), &[1.0, 1.0, 1.0]); autograd::clear_graph().unwrap(); @@ -281,7 +281,7 @@ fn test_mean_backward_correct() { let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], true); let m = reduction::mean(&a, None, false).unwrap(); let grad_output = Tensor::ones(m.shape().clone(), DataType::Float32, Device::cpu(), false); - let grads = autograd::backward(&m, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&m, Some(grad_output)).unwrap(); let grad_a = grads.get(&a.id()).unwrap(); assert_eq!( grad_a.data().as_f32_slice().unwrap(), @@ -297,7 +297,7 @@ fn test_add_backward_broadcasting() { let b = create_test_tensor_f32(vec![10.0, 20.0], vec![1, 2], true); let sum = arithmetic::add(&a, &b).unwrap(); let grad_output = Tensor::ones(sum.shape().clone(), DataType::Float32, Device::cpu(), false); - let grads = autograd::backward(&sum, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&sum, Some(grad_output)).unwrap(); let grad_a = grads.get(&a.id()).unwrap(); let grad_b = grads.get(&b.id()).unwrap(); assert_eq!(grad_a.data().as_f32_slice().unwrap(), &[2.0, 2.0, 2.0]); @@ -317,7 +317,7 @@ fn test_mul_backward_broadcasting() { Device::cpu(), false, ); - let grads = autograd::backward(&prod, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&prod, Some(grad_output)).unwrap(); let grad_a = grads.get(&a.id()).unwrap(); let grad_b = grads.get(&b.id()).unwrap(); assert_eq!(grad_a.data().as_f32_slice().unwrap(), &[30.0, 30.0, 30.0]); @@ -337,7 +337,7 @@ fn test_sub_backward_broadcasting() { Device::cpu(), false, ); - let grads = autograd::backward(&diff, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&diff, Some(grad_output)).unwrap(); let grad_a = grads.get(&a.id()).unwrap(); let grad_b = grads.get(&b.id()).unwrap(); assert_eq!(grad_a.data().as_f32_slice().unwrap(), &[2.0, 2.0, 2.0]); @@ -352,7 +352,7 @@ fn test_div_backward_broadcasting() { let b = create_test_tensor_f32(vec![2.0, 4.0], vec![1, 2], true); let quo = arithmetic::div(&a, &b).unwrap(); let grad_output = Tensor::ones(quo.shape().clone(), DataType::Float32, Device::cpu(), false); - let grads = autograd::backward(&quo, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&quo, Some(grad_output)).unwrap(); let grad_a = grads.get(&a.id()).unwrap(); let grad_b = grads.get(&b.id()).unwrap(); let ga = grad_a.data().as_f32_slice().unwrap(); @@ -372,7 +372,7 @@ fn test_pow_backward_correct() { let exp = create_test_tensor_f32(vec![3.0, 2.0], vec![2], true); let out = activation::pow(&base, &exp).unwrap(); let grad_output = Tensor::ones(out.shape().clone(), DataType::Float32, Device::cpu(), false); - let grads = autograd::backward(&out, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&out, Some(grad_output)).unwrap(); let grad_base = grads.get(&base.id()).unwrap(); let grad_exp = grads.get(&exp.id()).unwrap(); let gb = grad_base.data().as_f32_slice().unwrap(); @@ -390,7 +390,7 @@ fn test_powf_backward_correct() { let base = create_test_tensor_f32(vec![2.0, 3.0], vec![2], true); let out = activation::powf(&base, 3.0).unwrap(); let grad_output = Tensor::ones(out.shape().clone(), DataType::Float32, Device::cpu(), false); - let grads = autograd::backward(&out, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&out, Some(grad_output)).unwrap(); let grad_base = grads.get(&base.id()).unwrap(); let gb = grad_base.data().as_f32_slice().unwrap(); assert_relative_eq!(gb[0], 3.0 * 2.0f32.powf(2.0), epsilon = 1e-6); @@ -406,7 +406,7 @@ fn test_broadcast_backward_multiple_axes() { let b = create_test_tensor_f32(vec![2.0], vec![1, 1], true); let sum = arithmetic::add(&a, &b).unwrap(); let grad_output = Tensor::ones(sum.shape().clone(), DataType::Float32, Device::cpu(), false); - let grads = autograd::backward(&sum, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&sum, Some(grad_output)).unwrap(); let grad_a = grads.get(&a.id()).unwrap(); let grad_b = grads.get(&b.id()).unwrap(); assert_eq!(grad_a.shape().dims(), &[2, 3]); @@ -422,7 +422,7 @@ fn test_sum_backward_with_dims_keepdim() { let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], true); let s = reduction::sum(&a, Some(vec![1]), false).unwrap(); let grad_output = Tensor::ones(s.shape().clone(), DataType::Float32, Device::cpu(), false); - let grads = autograd::backward(&s, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&s, Some(grad_output)).unwrap(); let grad_a = grads.get(&a.id()).unwrap(); assert_eq!(grad_a.data().as_f32_slice().unwrap(), &[1.0, 1.0, 1.0, 1.0]); autograd::clear_graph().unwrap(); @@ -434,7 +434,7 @@ fn test_sum_backward_with_dims_keepdim() { Device::cpu(), false, ); - let grads = autograd::backward(&s_keep, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&s_keep, Some(grad_output)).unwrap(); let grad_a = grads.get(&a.id()).unwrap(); assert_eq!(grad_a.data().as_f32_slice().unwrap(), &[1.0, 1.0, 1.0, 1.0]); autograd::clear_graph().unwrap(); @@ -446,7 +446,7 @@ fn test_transpose_backward_correct() { let a = create_test_tensor_f32(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2], true); let t = linalg::transpose(&a, 0, 1).unwrap(); let grad_output = Tensor::ones(t.shape().clone(), DataType::Float32, Device::cpu(), false); - let grads = autograd::backward(&t, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&t, Some(grad_output)).unwrap(); let grad_a = grads.get(&a.id()).unwrap(); assert_eq!(grad_a.data().as_f32_slice().unwrap(), &[1.0, 1.0, 1.0, 1.0]); autograd::clear_graph().unwrap(); @@ -460,7 +460,7 @@ fn test_solve_backward_matches_manual() { let solution = linalg::solve(&a, &b).unwrap(); let grad_output = create_test_tensor_f32(vec![1.0, 1.0], vec![2], false); - let grads = autograd::backward(&solution, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&solution, Some(grad_output)).unwrap(); let grad_a = grads.get(&a.id()).expect("gradient for A missing"); let grad_b = grads.get(&b.id()).expect("gradient for B missing"); @@ -490,7 +490,7 @@ fn test_gradient_accumulation_multiple_paths() { let temp2 = arithmetic::add(&a, &c).unwrap(); let d = arithmetic::add(&temp1, &temp2).unwrap(); let grad_output = Tensor::ones(d.shape().clone(), DataType::Float32, Device::cpu(), false); - let grads = autograd::backward(&d, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&d, Some(grad_output)).unwrap(); let grad_a = grads.get(&a.id()).unwrap(); assert_eq!(grad_a.data().as_f32_slice().unwrap(), &[2.0, 2.0]); autograd::clear_graph().unwrap(); @@ -503,7 +503,7 @@ proptest! { let b = create_test_tensor_f32(b_vals.to_vec(), vec![2], true); let product = arithmetic::mul(&a, &b).unwrap(); let grad_output = Tensor::ones(product.shape().clone(), DataType::Float32, Device::cpu(), false); - let grads = autograd::backward(&product, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&product, Some(grad_output)).unwrap(); let grad_a = grads.get(&a.id()).unwrap(); let grad_b = grads.get(&b.id()).unwrap(); let ga = grad_a.data().as_f32_slice().unwrap(); @@ -531,7 +531,7 @@ fn test_repeat_zero_repeat_keeps_differentiable_zero_gradient() { Device::cpu(), false, ); - let grads = autograd::backward(&repeated, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&repeated, Some(grad_output)).unwrap(); let grad_input = grads.get(&input.id()).unwrap(); assert_eq!(grad_input.shape().dims(), &[3]); assert_eq!(grad_input.data().as_f32_slice().unwrap(), &[0.0, 0.0, 0.0]); @@ -553,7 +553,7 @@ fn test_repeat_empty_identity_keeps_differentiable_edge() { Device::cpu(), false, ); - let grads = autograd::backward(&repeated, Some(grad_output)).unwrap(); + let grads = autograd::backward_collect(&repeated, Some(grad_output)).unwrap(); let grad_input = grads.get(&input.id()).unwrap(); assert_eq!(grad_input.shape().dims(), &[0, 2]); assert_eq!(grad_input.numel(), 0); diff --git a/engine/tests/integration_test.rs b/engine/tests/integration_test.rs index 46990d5b..4bd9feac 100644 --- a/engine/tests/integration_test.rs +++ b/engine/tests/integration_test.rs @@ -397,7 +397,7 @@ fn test_softmax_backward_dim1() { let grad_output = create_test_tensor_f32(vec![0.1, 0.2, 0.3, 0.4], vec![2, 2], false); let result = activation::softmax(&input, Some(1)).unwrap(); - let grads = autograd::backward(&result, Some(grad_output.clone())).unwrap(); + let grads = autograd::backward_collect(&result, Some(grad_output.clone())).unwrap(); let grad_data = grads .get(&input.id()) .unwrap() @@ -434,7 +434,7 @@ fn test_softmax_backward_dim0() { let grad_output = create_test_tensor_f32(vec![0.1, 0.2, 0.3, 0.4], vec![2, 2], false); let result = activation::softmax(&input, Some(0)).unwrap(); - let grads = autograd::backward(&result, Some(grad_output.clone())).unwrap(); + let grads = autograd::backward_collect(&result, Some(grad_output.clone())).unwrap(); let grad_data = grads .get(&input.id()) .unwrap() diff --git a/minitensor/__init__.py b/minitensor/__init__.py index 9071a652..4ac799ce 100644 --- a/minitensor/__init__.py +++ b/minitensor/__init__.py @@ -128,6 +128,10 @@ clear_autograd_graph = _C.clear_autograd_graph is_autograd_graph_consumed = _C.is_autograd_graph_consumed mark_autograd_graph_consumed = _C.mark_autograd_graph_consumed +no_grad = _C.no_grad +enable_grad = _C.enable_grad +is_grad_enabled = _C.is_grad_enabled +set_grad_enabled = _C.set_grad_enabled functional = _C.functional _sys.modules[__name__ + ".functional"] = functional @@ -285,6 +289,10 @@ def default_dtype(dtype: str): "clear_autograd_graph", "is_autograd_graph_consumed", "mark_autograd_graph_consumed", + "no_grad", + "enable_grad", + "is_grad_enabled", + "set_grad_enabled", "functional", "nn", "optim", diff --git a/tests/ops/test_creation_and_numpy.py b/tests/ops/test_creation_and_numpy.py index 7ef305fe..582c126d 100644 --- a/tests/ops/test_creation_and_numpy.py +++ b/tests/ops/test_creation_and_numpy.py @@ -249,6 +249,10 @@ def _load_stubbed_module(monkeypatch: pytest.MonkeyPatch): core.clear_autograd_graph = lambda: None core.is_autograd_graph_consumed = lambda: False core.mark_autograd_graph_consumed = lambda: None + core.no_grad = lambda: None + core.enable_grad = lambda: None + core.is_grad_enabled = lambda: True + core.set_grad_enabled = lambda enabled: True core.functional = _DummyModule(f"{module_name}._core.functional") core.nn = _DummyModule(f"{module_name}._core.nn") core.optim = _DummyModule(f"{module_name}._core.optim") diff --git a/tests/tensor/test_grad_mode.py b/tests/tensor/test_grad_mode.py new file mode 100644 index 00000000..88ab5221 --- /dev/null +++ b/tests/tensor/test_grad_mode.py @@ -0,0 +1,114 @@ +# Copyright (c) 2026 Soumyadip Sarkar. +# All rights reserved. +# +# This source code is licensed under the Apache-style license found in the +# LICENSE file in the root directory of this source tree. + +"""Tests for gradient-recording mode (no_grad / enable_grad) and related +per-tensor gradient semantics.""" + +import pytest + +import minitensor as mt + + +class TestGradMode: + def test_grad_enabled_by_default(self): + assert mt.is_grad_enabled() + + def test_no_grad_disables_and_restores(self): + assert mt.is_grad_enabled() + with mt.no_grad(): + assert not mt.is_grad_enabled() + assert mt.is_grad_enabled() + + def test_no_grad_restores_on_exception(self): + with pytest.raises(RuntimeError): + with mt.no_grad(): + assert not mt.is_grad_enabled() + raise RuntimeError("boom") + assert mt.is_grad_enabled() + + def test_nested_enable_grad(self): + with mt.no_grad(): + assert not mt.is_grad_enabled() + with mt.enable_grad(): + assert mt.is_grad_enabled() + assert not mt.is_grad_enabled() + assert mt.is_grad_enabled() + + def test_set_grad_enabled_returns_previous(self): + prev = mt.set_grad_enabled(False) + try: + assert prev is True + assert not mt.is_grad_enabled() + finally: + mt.set_grad_enabled(True) + assert mt.is_grad_enabled() + + def test_op_results_inside_no_grad_are_detached_leaves(self): + x = mt.randn(3, 3) + x.requires_grad_(True) + with mt.no_grad(): + y = x * 2.0 + 1.0 + assert not y.requires_grad + # Backward on a detached result must fail like PyTorch. + with pytest.raises(RuntimeError): + y.sum().backward() + + def test_new_tensors_inside_no_grad_do_not_require_grad(self): + with mt.no_grad(): + t = mt.randn(2, 2) + t2 = t + 1.0 + assert not t2.requires_grad + + def test_explicit_opt_in_inside_no_grad(self): + with mt.no_grad(): + t = mt.randn(2, 2) + t.requires_grad_(True) + assert t.requires_grad + + def test_grad_flows_normally_after_no_grad_block(self): + x = mt.randn(2, 2) + x.requires_grad_(True) + with mt.no_grad(): + frozen = x * 3.0 + assert not frozen.requires_grad + y = (x * x).sum() + y.backward() + grad = mt.get_gradient(x) + assert grad is not None + assert grad.shape == x.shape + + +class TestRequiresGradChaining: + def test_requires_grad_returns_self(self): + x = mt.randn(2, 2).requires_grad_(True) + assert x is not None + assert x.requires_grad + + y = x.requires_grad_(False) + assert y is not None + assert not y.requires_grad + + +class TestPerTensorZeroGrad: + def test_zero_grad_only_clears_own_gradient(self): + a = mt.randn(2, 2) + a.requires_grad_(True) + b = mt.randn(2, 2) + b.requires_grad_(True) + + loss = (a * b).sum() + loss.backward() + + assert mt.get_gradient(a) is not None + assert mt.get_gradient(b) is not None + + a.zero_grad(set_to_none=True) + + assert mt.get_gradient(a) is None + # b's gradient must survive a.zero_grad(). + assert mt.get_gradient(b) is not None + + mt.clear_autograd_graph()