Mila
Deep Neural Network Library
Loading...
Searching...
No Matches
TensorOps.Math.ixx File Reference

Device-dispatching math helpers for tensor arithmetic operations. More...

#include <concepts>
#include <memory>
import Compute.DeviceType;
import Compute.ExecutionContext;
import Dnn.TensorOps.Base;
import Dnn.TensorDataTypeMap;
import Dnn.TensorDataTypeTraits;
import Dnn.TensorDataType;
import Dnn.Tensor;

Namespaces

namespace  Mila
 Mila main API namespace.

Functions

template<TensorDataType TDataType, typename TMemoryResource>
requires isValidTensor<TDataType, TMemoryResource>
void Mila::Dnn::add (const Tensor< TDataType, TMemoryResource > &a, const Tensor< TDataType, TMemoryResource > &b, Tensor< TDataType, TMemoryResource > &result, IExecutionContext *exec_context=nullptr)
 Element-wise addition with optional ExecutionContext (device-dispatched).
template<TensorDataType TDataType, typename TMemoryResource>
requires isValidTensor<TDataType, TMemoryResource>
void Mila::Dnn::divide (const Tensor< TDataType, TMemoryResource > &a, const Tensor< TDataType, TMemoryResource > &b, Tensor< TDataType, TMemoryResource > &result, IExecutionContext *exec_context=nullptr)
 Element-wise division with optional ExecutionContext (device-dispatched).
template<TensorDataType TDataType, typename TMemoryResource>
requires isValidTensor<TDataType, TMemoryResource>
void Mila::Dnn::multiply (const Tensor< TDataType, TMemoryResource > &a, const Tensor< TDataType, TMemoryResource > &b, Tensor< TDataType, TMemoryResource > &result, IExecutionContext *exec_context=nullptr)
 Element-wise multiplication with optional ExecutionContext (device-dispatched).
template<TensorDataType TDataType, typename TMemoryResource>
requires isValidTensor<TDataType, TMemoryResource>
Tensor< TDataType, TMemoryResource > Mila::Dnn::operator* (const Tensor< TDataType, TMemoryResource > &a, const Tensor< TDataType, TMemoryResource > &b)
 Element-wise multiplication operator (always synchronous).
template<TensorDataType TDataType, typename TMemoryResource>
requires isValidTensor<TDataType, TMemoryResource>
Tensor< TDataType, TMemoryResource > Mila::Dnn::operator+ (const Tensor< TDataType, TMemoryResource > &a, const Tensor< TDataType, TMemoryResource > &b)
 Element-wise addition operator (always synchronous).
template<TensorDataType TDataType, typename TMemoryResource>
requires isValidTensor<TDataType, TMemoryResource>
Tensor< TDataType, TMemoryResource > Mila::Dnn::operator- (const Tensor< TDataType, TMemoryResource > &a, const Tensor< TDataType, TMemoryResource > &b)
 Element-wise subtraction operator (always synchronous).
template<TensorDataType TDataType, typename TMemoryResource>
requires isValidTensor<TDataType, TMemoryResource>
Tensor< TDataType, TMemoryResource > Mila::Dnn::operator/ (const Tensor< TDataType, TMemoryResource > &a, const Tensor< TDataType, TMemoryResource > &b)
 Element-wise division operator (always synchronous).
template<TensorDataType TDataType, typename TMemoryResource>
requires isValidTensor<TDataType, TMemoryResource>
void Mila::Dnn::scale (const Tensor< TDataType, TMemoryResource > &input, float scalar, Tensor< TDataType, TMemoryResource > &result, IExecutionContext *exec_context=nullptr)
 Scale a tensor by a scalar: result[i] = input[i] * scalar (device-dispatched).
template<TensorDataType TDataType, typename TMemoryResource>
requires isValidTensor<TDataType, TMemoryResource>
void Mila::Dnn::subtract (const Tensor< TDataType, TMemoryResource > &a, const Tensor< TDataType, TMemoryResource > &b, Tensor< TDataType, TMemoryResource > &result, IExecutionContext *exec_context=nullptr)
 Element-wise subtraction with optional ExecutionContext (device-dispatched).
template<TensorDataType TDataType, typename TMemoryResource>
requires isValidTensor<TDataType, TMemoryResource>
float Mila::Dnn::sum (const Tensor< TDataType, TMemoryResource > &tensor, IExecutionContext *exec_context=nullptr)
 Sum reduction with optional ExecutionContext (device-dispatched).

Detailed Description

Device-dispatching math helpers for tensor arithmetic operations.

This partition provides the high-level, device-agnostic entry points for tensor math operations (e.g., element-wise addition). Each helper forwards to the device-specific TensorOps<ComputeDeviceTag>::... implementation (see CPU and CUDA specializations).

The templates are constrained with isValidTensor<TDataType, TMemoryResource> to ensure the tensor configuration is valid (memory resource compatibility, type traits available, and device accessibility).

ExecutionContext handling:

  • Optional ExecutionContext parameter for stream control (borrowed, not owned)
  • When provided, operations use the context's stream (caller controls sync)
  • When null, operations use default stream and synchronize before returning
  • Raw pointer semantics ensure zero overhead

Usage:

  • Call add(a, b, result) for element-wise addition of two tensors with the same abstract data type and memory resource. The call is automatically dispatched to the appropriate device implementation.
  • Optionally provide ExecutionContext for explicit stream control: add(a, b, result, ctx.get())

Preconditions:

  • All operands must satisfy isValidTensor and have matching shapes.
  • Result tensor must be pre-allocated with matching shape.
  • Device-specific implementations validate shapes and perform operations efficiently.
  • ExecutionContext (if provided) must outlive the function call.

Function Documentation

◆ add()

template<TensorDataType TDataType, typename TMemoryResource>
requires isValidTensor<TDataType, TMemoryResource>
void Mila::Dnn::add ( const Tensor< TDataType, TMemoryResource > & a,
const Tensor< TDataType, TMemoryResource > & b,
Tensor< TDataType, TMemoryResource > & result,
IExecutionContext * exec_context = nullptr )
export

Element-wise addition with optional ExecutionContext (device-dispatched).

Computes result[i] = a[i] + b[i] for all elements. Automatically dispatches to the appropriate device implementation based on memory resource type.

Template Parameters
TDataTypeAbstract tensor data type
TMemoryResourceMemory resource type determining device
Parameters
aFirst input tensor
bSecond input tensor
resultOutput tensor (must be pre-allocated with matching shape)
exec_contextOptional execution context for stream control (borrowed, not owned)
Note
For CUDA tensors, use CudaExecutionContext; for CPU, parameter is ignored
exec_context must outlive this function call
When exec_context provided, caller controls synchronization
When null, uses default stream/execution and synchronizes before returning

Example:

// With explicit context (async)
auto ctx = std::make_unique<CudaExecutionContext>(0);
add(tensor_a, tensor_b, result, ctx.get());
ctx->synchronize();
// Without context (sync)
add(tensor_a, tensor_b, result); // Returns after completion

◆ divide()

template<TensorDataType TDataType, typename TMemoryResource>
requires isValidTensor<TDataType, TMemoryResource>
void Mila::Dnn::divide ( const Tensor< TDataType, TMemoryResource > & a,
const Tensor< TDataType, TMemoryResource > & b,
Tensor< TDataType, TMemoryResource > & result,
IExecutionContext * exec_context = nullptr )
export

Element-wise division with optional ExecutionContext (device-dispatched).

Computes result[i] = a[i] / b[i] for all elements.

Template Parameters
TDataTypeAbstract tensor data type
TMemoryResourceMemory resource type determining device
Parameters
aFirst input tensor (dividend)
bSecond input tensor (divisor)
resultOutput tensor (must be pre-allocated with matching shape)
exec_contextOptional execution context for stream control (borrowed, not owned)

◆ multiply()

template<TensorDataType TDataType, typename TMemoryResource>
requires isValidTensor<TDataType, TMemoryResource>
void Mila::Dnn::multiply ( const Tensor< TDataType, TMemoryResource > & a,
const Tensor< TDataType, TMemoryResource > & b,
Tensor< TDataType, TMemoryResource > & result,
IExecutionContext * exec_context = nullptr )
export

Element-wise multiplication with optional ExecutionContext (device-dispatched).

Computes result[i] = a[i] * b[i] for all elements.

Template Parameters
TDataTypeAbstract tensor data type
TMemoryResourceMemory resource type determining device
Parameters
aFirst input tensor
bSecond input tensor
resultOutput tensor (must be pre-allocated with matching shape)
exec_contextOptional execution context for stream control (borrowed, not owned)

◆ operator*()

template<TensorDataType TDataType, typename TMemoryResource>
requires isValidTensor<TDataType, TMemoryResource>
Tensor< TDataType, TMemoryResource > Mila::Dnn::operator* ( const Tensor< TDataType, TMemoryResource > & a,
const Tensor< TDataType, TMemoryResource > & b )
export

Element-wise multiplication operator (always synchronous).

Note
This operator always uses default execution and synchronizes.
For async operations with stream control, use multiply(a, b, result, ctx).

◆ operator+()

template<TensorDataType TDataType, typename TMemoryResource>
requires isValidTensor<TDataType, TMemoryResource>
Tensor< TDataType, TMemoryResource > Mila::Dnn::operator+ ( const Tensor< TDataType, TMemoryResource > & a,
const Tensor< TDataType, TMemoryResource > & b )
export

Element-wise addition operator (always synchronous).

Note
This operator always uses default execution and synchronizes.
For async operations with stream control, use add(a, b, result, ctx).

◆ operator-()

template<TensorDataType TDataType, typename TMemoryResource>
requires isValidTensor<TDataType, TMemoryResource>
Tensor< TDataType, TMemoryResource > Mila::Dnn::operator- ( const Tensor< TDataType, TMemoryResource > & a,
const Tensor< TDataType, TMemoryResource > & b )
export

Element-wise subtraction operator (always synchronous).

Note
This operator always uses default execution and synchronizes.
For async operations with stream control, use subtract(a, b, result, ctx).

◆ operator/()

template<TensorDataType TDataType, typename TMemoryResource>
requires isValidTensor<TDataType, TMemoryResource>
Tensor< TDataType, TMemoryResource > Mila::Dnn::operator/ ( const Tensor< TDataType, TMemoryResource > & a,
const Tensor< TDataType, TMemoryResource > & b )
export

Element-wise division operator (always synchronous).

Note
This operator always uses default execution and synchronizes.
For async operations with stream control, use divide(a, b, result, ctx).

◆ scale()

template<TensorDataType TDataType, typename TMemoryResource>
requires isValidTensor<TDataType, TMemoryResource>
void Mila::Dnn::scale ( const Tensor< TDataType, TMemoryResource > & input,
float scalar,
Tensor< TDataType, TMemoryResource > & result,
IExecutionContext * exec_context = nullptr )
export

Scale a tensor by a scalar: result[i] = input[i] * scalar (device-dispatched).

Supports in-place (input and result may alias). Added for Gemma 4's per-layer hidden_states *= layer_scalar; default-safe for all backends.

Parameters
inputInput tensor
scalarScalar multiplier (converted to the tensor's native type)
resultOutput tensor (must match input shape; may alias input)
exec_contextOptional execution context for stream control (borrowed, not owned)

◆ subtract()

template<TensorDataType TDataType, typename TMemoryResource>
requires isValidTensor<TDataType, TMemoryResource>
void Mila::Dnn::subtract ( const Tensor< TDataType, TMemoryResource > & a,
const Tensor< TDataType, TMemoryResource > & b,
Tensor< TDataType, TMemoryResource > & result,
IExecutionContext * exec_context = nullptr )
export

Element-wise subtraction with optional ExecutionContext (device-dispatched).

Computes result[i] = a[i] - b[i] for all elements.

Template Parameters
TDataTypeAbstract tensor data type
TMemoryResourceMemory resource type determining device
Parameters
aFirst input tensor (minuend)
bSecond input tensor (subtrahend)
resultOutput tensor (must be pre-allocated with matching shape)
exec_contextOptional execution context for stream control (borrowed, not owned)

◆ sum()

template<TensorDataType TDataType, typename TMemoryResource>
requires isValidTensor<TDataType, TMemoryResource>
float Mila::Dnn::sum ( const Tensor< TDataType, TMemoryResource > & tensor,
IExecutionContext * exec_context = nullptr )
export

Sum reduction with optional ExecutionContext (device-dispatched).

Computes the sum of all elements in the tensor. Always synchronizes before returning the result (even when exec_context is provided).

Template Parameters
TDataTypeAbstract tensor data type
TMemoryResourceMemory resource type determining device
Parameters
tensorInput tensor
exec_contextOptional execution context for stream control (borrowed, not owned)
Returns
Sum of all elements as float
Note
Always returns after synchronization to ensure result validity