use std::fmt;
use std::ops::{Add, Div, Mul, Neg, Sub};
use std::sync::Arc;
use static_assertions::assert_impl_all;
use crate::backend::MapTask;
use super::elementary::MapOperation;
use super::gemm;
use super::layout::{Layout, Strides};
use super::normalized::BatchNormTask;
use super::recordable::{composed_batch_norm, composed_max_pool, composed_windowed_patches};
use super::storage::Storage;
use super::{Differentiable, Element, Elementary, Recordable, Shape};
assert_impl_all!(Tensor<f64>: Send, Sync);
#[derive(Debug, Clone)]
pub struct Tensor<Element> {
storage: Storage<Element>,
}
impl<Element> Tensor<Element> {
pub fn shape(&self) -> Shape {
self.logical_shape().clone()
}
fn logical_shape(&self) -> &Shape {
match &self.storage {
Storage::Dense { layout, .. } => layout.shape(),
Storage::Constant { shape, .. } => shape,
Storage::Selection { shape, .. } => shape,
}
}
pub fn as_constant(&self) -> Option<&Element> {
match &self.storage {
Storage::Constant { value, .. } => Some(value),
_ => None,
}
}
fn selection_indices(&self) -> &[usize] {
match &self.storage {
Storage::Selection { indices, .. } => indices,
_ => panic!("gather requires a selection tensor built with `Tensor::selection`"),
}
}
fn gemm_operand(&self) -> Option<(&[Element], &[usize])> {
match &self.storage {
Storage::Dense { data, layout } if layout.rank() >= 2 => {
Some((&data.as_slice()[layout.offset()..], layout.strides()))
}
_ => None,
}
}
pub fn as_slice(&self) -> Option<&[Element]> {
match &self.storage {
Storage::Dense { data, layout } if layout.is_contiguous() => {
let start = layout.offset();
Some(&data.as_slice()[start..start + layout.volume()])
}
_ => None,
}
}
fn strided_window(&self) -> Option<(&[Element], Layout)> {
let Storage::Dense { data, layout } = &self.storage else {
return None;
};
if layout.is_contiguous() || layout.span() >= layout.volume() {
return None;
}
let start = layout.offset();
Some((
&data.as_slice()[start..start + layout.span()],
layout.rebased(),
))
}
}
impl<Element: Clone> Tensor<Element> {
pub fn iter(&self) -> impl Iterator<Item = Element> + '_ {
match &self.storage {
Storage::Constant { shape, value } => ElementIter::Constant {
value,
remaining: shape.volume(),
},
Storage::Dense { data, layout } if layout.is_contiguous() => {
let start = layout.offset();
ElementIter::Contiguous(data.as_slice()[start..start + layout.volume()].iter())
}
Storage::Dense { data, layout } => ElementIter::Strided {
data: data.as_slice(),
shape: layout.shape().axes(),
strides: layout.strides(),
coordinates: std::iter::repeat_n(0usize, layout.rank()).collect(),
index: layout.offset(),
remaining: layout.volume(),
},
Storage::Selection {
indices,
shape,
zero,
one,
} => ElementIter::Selection {
indices: indices.as_slice(),
vocab: shape.axes()[1],
zero,
one,
position: 0,
total: shape.volume(),
},
}
}
pub fn to_vec(&self) -> Vec<Element> {
self.iter().collect()
}
pub fn convert<Target: From<Element>>(&self) -> Tensor<Target> {
match &self.storage {
Storage::Dense { data, layout } => Tensor {
storage: Storage::Dense {
data: Arc::new(data.iter().cloned().map(Target::from).collect()),
layout: layout.clone(),
},
},
Storage::Constant { shape, value } => Tensor {
storage: Storage::Constant {
shape: shape.clone(),
value: Target::from(value.clone()),
},
},
Storage::Selection {
indices,
shape,
zero,
one,
} => Tensor {
storage: Storage::Selection {
indices: Arc::clone(indices),
shape: shape.clone(),
zero: Target::from(zero.clone()),
one: Target::from(one.clone()),
},
},
}
}
pub fn scalar(&self) -> Element {
assert_eq!(
self.logical_shape().rank(),
0,
"scalar reads a rank-0 tensor, got {}",
self.logical_shape()
);
self.get(0)
}
fn get(&self, position: usize) -> Element {
match &self.storage {
Storage::Dense { data, layout } => data[layout.storage_index(position)].clone(),
Storage::Constant { value, .. } => value.clone(),
Storage::Selection {
indices,
shape,
zero,
one,
} => {
let vocab = shape.axes()[1];
if indices[position / vocab] == position % vocab {
one.clone()
} else {
zero.clone()
}
}
}
}
}
impl<Element: Differentiable> Tensor<Element> {
pub fn new(shape: impl Into<Shape>, elements: impl Into<Vec<Element>>) -> Self {
Self::dense(shape.into(), elements.into())
}
pub fn filled(shape: impl Into<Shape>, element: Element) -> Self {
Self::constant(shape.into(), element)
}
fn dense(shape: Shape, elements: Vec<Element>) -> Self {
assert_eq!(
shape.volume(),
elements.len(),
"tensor shape does not match its number of elements"
);
assert!(
!elements.is_empty(),
"tensors must hold at least one element"
);
Self {
storage: Storage::Dense {
layout: Layout::contiguous(shape),
data: Arc::new(elements),
},
}
}
fn constant(shape: Shape, value: Element) -> Self {
assert!(shape.volume() > 0, "tensors must hold at least one element");
Self {
storage: Storage::Constant { shape, value },
}
}
pub fn selection(indices: impl Into<Vec<usize>>, vocab: usize, one: Element) -> Self {
let indices = indices.into();
assert!(vocab > 0, "a selection needs a non-empty vocabulary");
assert!(
!indices.is_empty(),
"tensors must hold at least one element"
);
assert!(
indices.len().checked_mul(vocab).is_some(),
"shape volume overflows `usize`"
);
for &index in &indices {
assert!(
index < vocab,
"selection index {index} is out of vocabulary {vocab}"
);
}
let zero = Element::zero();
let shape = Shape::new([indices.len(), vocab]);
Self {
storage: Storage::Selection {
indices: Arc::new(indices),
shape,
zero,
one,
},
}
}
fn densify(&self) -> Self {
match &self.storage {
Storage::Dense { layout, .. } if layout.is_contiguous() => self.clone(),
_ => Self::dense(self.logical_shape().clone(), self.to_vec()),
}
}
fn map(&self, transform: impl Fn(&Element) -> Element) -> Self {
if let Storage::Constant { shape, value } = &self.storage {
return Self::constant(shape.clone(), transform(value));
}
if let Some(elements) = self.as_slice() {
return Self::dense(
self.logical_shape().clone(),
elements.iter().map(transform).collect(),
);
}
if let Some((window, layout)) = self.strided_window() {
return Self {
storage: Storage::Dense {
data: Arc::new(window.iter().map(transform).collect()),
layout,
},
};
}
Self::dense(
self.logical_shape().clone(),
self.iter().map(|element| transform(&element)).collect(),
)
}
fn zip(&self, other: &Self, combine: impl Fn(&Element, &Element) -> Element) -> Self {
assert_eq!(
self.logical_shape(),
other.logical_shape(),
"tensors have different shapes"
);
if let (Storage::Constant { value: left, .. }, Storage::Constant { value: right, .. }) =
(&self.storage, &other.storage)
{
return Self::constant(self.logical_shape().clone(), combine(left, right));
}
if let (Some(left), Some(right)) = (self.as_slice(), other.as_slice()) {
return Self::dense(
self.logical_shape().clone(),
left.iter()
.zip(right)
.map(|(left, right)| combine(left, right))
.collect(),
);
}
if let (Storage::Dense { .. }, Storage::Constant { value: right, .. }) =
(&self.storage, &other.storage)
{
return self.map(|left| combine(left, right));
}
if let (Storage::Constant { value: left, .. }, Storage::Dense { .. }) =
(&self.storage, &other.storage)
{
return other.map(|right| combine(left, right));
}
if let (
Storage::Dense {
data: left,
layout: left_layout,
},
Storage::Dense {
data: right,
layout: right_layout,
},
) = (&self.storage, &other.storage)
&& let Some(combined) = zipped_runs(left, left_layout, right, right_layout, &combine)
{
return Self::dense(self.logical_shape().clone(), combined);
}
Self::dense(
self.logical_shape().clone(),
self.iter()
.zip(other.iter())
.map(|(left, right)| combine(&left, &right))
.collect(),
)
}
}
fn zipped_runs<Element: Clone>(
left: &[Element],
left_layout: &Layout,
right: &[Element],
right_layout: &Layout,
combine: impl Fn(&Element, &Element) -> Element,
) -> Option<Vec<Element>> {
let left_stride = left_layout.inner_stride();
let right_stride = right_layout.inner_stride();
if left_stride > 1 || right_stride > 1 {
return None;
}
let extent = left_layout.inner_extent();
let mut combined = Vec::with_capacity(left_layout.volume());
for (left_start, right_start) in left_layout.run_offsets().zip(right_layout.run_offsets()) {
match (left_stride, right_stride) {
(1, 1) => combined.extend(
left[left_start..left_start + extent]
.iter()
.zip(&right[right_start..right_start + extent])
.map(|(left, right)| combine(left, right)),
),
(1, 0) => {
let held = &right[right_start];
combined.extend(
left[left_start..left_start + extent]
.iter()
.map(|left| combine(left, held)),
);
}
(0, 1) => {
let held = &left[left_start];
combined.extend(
right[right_start..right_start + extent]
.iter()
.map(|right| combine(held, right)),
);
}
_ => {
let value = combine(&left[left_start], &right[right_start]);
combined.extend(std::iter::repeat_n(value, extent));
}
}
}
Some(combined)
}
fn transpose_shape(shape: &Shape) -> Shape {
if shape.rank() < 2 {
return shape.clone();
}
assert_eq!(shape.rank(), 2, "transpose supports rank 2 at most");
let axes = shape.axes();
Shape::new([axes[1], axes[0]])
}
enum ElementIter<'tensor, Element> {
Constant {
value: &'tensor Element,
remaining: usize,
},
Contiguous(std::slice::Iter<'tensor, Element>),
Strided {
data: &'tensor [Element],
shape: &'tensor [usize],
strides: &'tensor [usize],
coordinates: Strides,
index: usize,
remaining: usize,
},
Selection {
indices: &'tensor [usize],
vocab: usize,
zero: &'tensor Element,
one: &'tensor Element,
position: usize,
total: usize,
},
}
impl<'tensor, Element: Clone> Iterator for ElementIter<'tensor, Element> {
type Item = Element;
fn next(&mut self) -> Option<Element> {
match self {
ElementIter::Constant { value, remaining } => {
if *remaining == 0 {
return None;
}
*remaining -= 1;
Some((*value).clone())
}
ElementIter::Contiguous(iterator) => iterator.next().cloned(),
ElementIter::Strided {
data,
shape,
strides,
coordinates,
index,
remaining,
} => {
if *remaining == 0 {
return None;
}
let element = data[*index].clone();
*remaining -= 1;
if *remaining > 0 {
for axis in (0..shape.len()).rev() {
coordinates[axis] += 1;
if coordinates[axis] < shape[axis] {
*index += strides[axis];
break;
}
*index -= (shape[axis] - 1) * strides[axis];
coordinates[axis] = 0;
}
}
Some(element)
}
ElementIter::Selection {
indices,
vocab,
zero,
one,
position,
total,
} => {
if *position >= *total {
return None;
}
let row = *position / *vocab;
let column = *position % *vocab;
*position += 1;
Some(if indices[row] == column {
(*one).clone()
} else {
(*zero).clone()
})
}
}
}
}
impl<E: Element> From<E> for Tensor<E> {
fn from(element: E) -> Self {
Self::constant(Shape::scalar(), element)
}
}
impl<Element: Clone + fmt::Display> fmt::Display for Tensor<Element> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
let elements = self.to_vec();
display_level(formatter, self.logical_shape().axes(), &elements)
}
}
fn display_level<Element: fmt::Display>(
formatter: &mut fmt::Formatter<'_>,
axes: &[usize],
elements: &[Element],
) -> fmt::Result {
let Some((&extent, rest)) = axes.split_first() else {
return write!(formatter, "{}", elements[0]);
};
let stride = elements.len() / extent;
write!(formatter, "[")?;
for index in 0..extent {
if index > 0 {
write!(formatter, ", ")?;
}
display_level(
formatter,
rest,
&elements[index * stride..(index + 1) * stride],
)?;
}
write!(formatter, "]")
}
impl<Element: PartialEq + Clone> PartialEq for Tensor<Element> {
fn eq(&self, other: &Self) -> bool {
self.logical_shape() == other.logical_shape() && self.iter().eq(other.iter())
}
}
impl<Element: Differentiable> Add for Tensor<Element> {
type Output = Self;
fn add(self, rhs: Self) -> Self {
self.zip(&rhs, |left, right| left.clone() + right.clone())
}
}
impl<Element: Differentiable> Sub for Tensor<Element> {
type Output = Self;
fn sub(self, rhs: Self) -> Self {
self.zip(&rhs, |left, right| left.clone() - right.clone())
}
}
impl<Element: Differentiable> Mul for Tensor<Element> {
type Output = Self;
fn mul(self, rhs: Self) -> Self {
self.zip(&rhs, |left, right| left.clone() * right.clone())
}
}
impl<Element: Differentiable> Div for Tensor<Element> {
type Output = Self;
fn div(self, rhs: Self) -> Self {
self.zip(&rhs, |left, right| left.clone() / right.clone())
}
}
impl<Element: Differentiable> Neg for Tensor<Element> {
type Output = Self;
fn neg(self) -> Self {
self.map(|element| -element.clone())
}
}
impl<Element: Differentiable> Tensor<Element> {
pub fn zero_like(&self) -> Self {
Self::constant(self.logical_shape().clone(), Element::zero())
}
pub fn one_like(&self) -> Self {
Self::constant(self.logical_shape().clone(), Element::one())
}
pub fn counted(shape: Shape, count: usize) -> Self {
Self::constant(shape, Element::from_count(count))
}
pub fn is_counted(&self, shape: &Shape, count: usize) -> bool {
*self.logical_shape() == *shape && self.iter().all(|element| element.is_count(count))
}
}
impl<Element: Elementary> Tensor<Element> {
fn mapped(&self, operation: MapOperation, fallback: impl Fn(&Element) -> Element) -> Self {
if let Some(elements) = self.as_slice()
&& let Some(mapped) = Element::map(&MapTask::new(operation, elements))
{
assert_eq!(
mapped.len(),
elements.len(),
"the `Elementary::map` contract requires one output per input element"
);
return Self::dense(self.logical_shape().clone(), mapped);
}
if let Some((window, layout)) = self.strided_window()
&& let Some(mapped) = Element::map(&MapTask::new(operation, window))
{
assert_eq!(
mapped.len(),
window.len(),
"the `Elementary::map` contract requires one output per input element"
);
return Self {
storage: Storage::Dense {
data: Arc::new(mapped),
layout,
},
};
}
self.map(fallback)
}
}
impl<Element: Elementary> Tensor<Element> {
pub fn exp(&self) -> Self {
self.mapped(MapOperation::Exp, |element| element.exp())
}
pub fn ln(&self) -> Self {
self.mapped(MapOperation::Ln, |element| element.ln())
}
pub fn sqrt(&self) -> Self {
self.mapped(MapOperation::Sqrt, |element| element.sqrt())
}
pub fn tanh(&self) -> Self {
self.mapped(MapOperation::Tanh, |element| element.tanh())
}
pub fn sin(&self) -> Self {
self.mapped(MapOperation::Sin, |element| element.sin())
}
pub fn cos(&self) -> Self {
self.mapped(MapOperation::Cos, |element| element.cos())
}
pub fn log1p(&self) -> Self {
self.mapped(MapOperation::Log1p, |element| element.log1p())
}
pub fn expm1(&self) -> Self {
self.mapped(MapOperation::Expm1, |element| element.expm1())
}
pub fn erf(&self) -> Self {
self.mapped(MapOperation::Erf, |element| element.erf())
}
pub fn erf_derivative(&self) -> Self {
self.mapped(MapOperation::ErfDerivative, |element| {
element.erf_derivative()
})
}
pub fn powf(&self, exponent: Self) -> Self {
self.zip(&exponent, |element, exponent| {
element.powf(exponent.clone())
})
}
pub fn maximum(&self, other: &Self) -> Self {
self.zip(other, |element, other| element.maximum(other))
}
pub fn step(&self, threshold: &Self) -> Self {
self.zip(threshold, |element, threshold| element.step(threshold))
}
}
impl<Element: Elementary> Tensor<Element> {
pub fn batch_normalized(
&self,
scale: &Self,
shift: &Self,
epsilon: &Self,
) -> (Self, Self, Self) {
let shape = self.logical_shape().clone();
let axes = shape.axes();
assert_eq!(
axes.len(),
2,
"batch_normalized input must be rank 2 [batch, features], got {shape}"
);
let (batch, features) = (axes[0], axes[1]);
assert_eq!(
scale.logical_shape().volume(),
features,
"batch_normalized scale must hold {features} features"
);
assert_eq!(
shift.logical_shape().volume(),
features,
"batch_normalized shift must hold {features} features"
);
assert_eq!(
epsilon.logical_shape().volume(),
1,
"batch_normalized epsilon must hold a single value"
);
if let Some(input) = self.as_slice() {
let scale_elements = scale.to_vec();
let shift_elements = shift.to_vec();
let epsilon_value = epsilon.iter().next().expect("epsilon holds one value");
if let Some(normalized) = Element::batch_norm(&BatchNormTask::new(
input,
&scale_elements,
&shift_elements,
epsilon_value,
batch,
features,
)) {
let feature_shape = Shape::new([features]);
return (
Self::dense(shape, normalized.output),
Self::dense(feature_shape.clone(), normalized.mean),
Self::dense(feature_shape, normalized.variance),
);
}
}
composed_batch_norm(self, scale, shift, epsilon)
}
pub fn max_pooled(&self, size: usize, stride: usize) -> Self {
let shape = self.logical_shape();
let axes = shape.axes();
assert_eq!(
axes.len(),
4,
"max_pooled input must be rank 4 [batch, channels, height, width], got {shape}"
);
assert!(
size > 0 && stride > 0,
"max_pooled needs positive size and stride"
);
let (batch, channels, height, width) = (axes[0], axes[1], axes[2], axes[3]);
assert!(
size <= height && size <= width,
"max_pooled window {size} does not fit {shape}"
);
let Some(elements) = self.as_slice() else {
return composed_max_pool(self, size, stride);
};
let out_height = (height - size) / stride + 1;
let out_width = (width - size) / stride + 1;
let mut pooled = Vec::with_capacity(batch * channels * out_height * out_width);
for image in 0..batch {
for channel in 0..channels {
let plane = (image * channels + channel) * height;
for out_y in 0..out_height {
for out_x in 0..out_width {
let corner = (plane + out_y * stride) * width + out_x * stride;
let mut largest = elements[corner].clone();
for lane_y in 0..size {
let row = corner + lane_y * width;
for lane_x in 0..size {
if lane_y == 0 && lane_x == 0 {
continue;
}
largest = largest.maximum(&elements[row + lane_x]);
}
}
pooled.push(largest);
}
}
}
}
Self::dense(Shape::new([batch, channels, out_height, out_width]), pooled)
}
pub fn matmul(&self, rhs: &Self) -> Self {
let left = self.logical_shape();
let right = rhs.logical_shape();
assert!(left.rank() >= 2, "matmul requires rank-2 or higher tensors");
assert_eq!(
left.rank(),
right.rank(),
"matmul operands must agree in rank"
);
let split = left.rank() - 2;
assert_eq!(
&left.axes()[..split],
&right.axes()[..split],
"matmul batch axes do not agree"
);
let (rows, inner) = (left.axes()[split], left.axes()[split + 1]);
let (rhs_inner, columns) = (right.axes()[split], right.axes()[split + 1]);
assert_eq!(inner, rhs_inner, "matmul inner dimensions do not agree");
assert!(
rows > 0 && inner > 0 && columns > 0,
"matmul requires non-empty dimensions"
);
let batch_axes = &left.axes()[..split];
let batches: usize = batch_axes.iter().product();
let result_shape = Shape::new(batch_axes.iter().copied().chain([rows, columns]));
if batches == 0 {
return Self::dense(result_shape, Vec::new());
}
if let (Some((a, a_strides)), Some((b, b_strides))) =
(self.gemm_operand(), rhs.gemm_operand())
{
let mut elements = Vec::with_capacity(batches * rows * columns);
let mut index = vec![0usize; split];
loop {
let a_base: usize = index
.iter()
.zip(&a_strides[..split])
.map(|(&at, &stride)| at * stride)
.sum();
let b_base: usize = index
.iter()
.zip(&b_strides[..split])
.map(|(&at, &stride)| at * stride)
.sum();
let task = gemm::GemmTask::new(
&a[a_base..],
[a_strides[split], a_strides[split + 1]],
&b[b_base..],
[b_strides[split], b_strides[split + 1]],
rows,
inner,
columns,
);
match Element::gemm(&task) {
Some(product) => {
assert_eq!(
product.len(),
rows * columns,
"the `Elementary::gemm` contract requires `rows * columns` elements"
);
elements.extend(product);
}
None => elements.extend(gemm::multiply(&task)),
}
let mut axis = split;
loop {
if axis == 0 {
return Self::dense(result_shape, elements);
}
axis -= 1;
index[axis] += 1;
if index[axis] < batch_axes[axis] {
break;
}
index[axis] = 0;
}
}
}
let mut elements = Vec::with_capacity(batches * rows * columns);
for batch in 0..batches {
let a_base = batch * rows * inner;
let b_base = batch * inner * columns;
for row in 0..rows {
for column in 0..columns {
let mut total = self.get(a_base + row * inner).promote()
* rhs.get(b_base + column).promote();
for step in 1..inner {
total = total
+ self.get(a_base + row * inner + step).promote()
* rhs.get(b_base + step * columns + column).promote();
}
elements.push(Element::demote(total));
}
}
}
Self::dense(result_shape, elements)
}
pub fn transpose(&self) -> Self {
match &self.storage {
Storage::Dense { data, layout } => Self {
storage: Storage::Dense {
data: Arc::clone(data),
layout: layout.transpose(),
},
},
Storage::Constant { shape, value } => {
Self::constant(transpose_shape(shape), value.clone())
}
Storage::Selection { .. } => self.densify().transpose(),
}
}
pub fn sum(&self) -> Self {
let mut elements = self.iter();
let first = elements
.next()
.expect("sum requires a non-empty tensor")
.promote();
let total = elements.fold(first, |total, element| total + element.promote());
Self::constant(Shape::scalar(), Element::demote(total))
}
pub fn sum_along(&self, axis: usize) -> Self {
let shape = self.logical_shape();
let axes = shape.axes();
assert!(axis < axes.len(), "axis {axis} is out of rank for {shape}");
let outer: usize = axes[..axis].iter().product();
let extent = axes[axis];
let inner: usize = axes[axis + 1..].iter().product();
let mut elements = Vec::with_capacity(outer * inner);
for outer_index in 0..outer {
for inner_index in 0..inner {
let position = |step: usize| (outer_index * extent + step) * inner + inner_index;
let mut total = self.get(position(0)).promote();
for step in 1..extent {
total = total + self.get(position(step)).promote();
}
elements.push(Element::demote(total));
}
}
Self::dense(shape.without_axis(axis), elements)
}
pub fn logsumexp(&self, axis: usize) -> Self {
let peak = self.max_along(axis);
let shifted = self.clone() - peak.broadcast_along_like(axis, self);
peak + shifted.exp().sum_along(axis).ln()
}
pub fn log_softmax(&self, axis: usize) -> Self {
let peak = self.max_along(axis).broadcast_along_like(axis, self);
let shifted = self.clone() - peak;
let normalizer = shifted.exp().sum_along(axis).ln();
shifted.clone() - normalizer.broadcast_along_like(axis, &shifted)
}
pub fn max_along(&self, axis: usize) -> Self {
let shape = self.logical_shape();
let axes = shape.axes();
assert!(axis < axes.len(), "axis {axis} is out of rank for {shape}");
let outer: usize = axes[..axis].iter().product();
let extent = axes[axis];
let inner: usize = axes[axis + 1..].iter().product();
let mut elements = Vec::with_capacity(outer * inner);
for outer_index in 0..outer {
for inner_index in 0..inner {
let position = |step: usize| (outer_index * extent + step) * inner + inner_index;
let mut largest = self.get(position(0));
for step in 1..extent {
largest = largest.maximum(&self.get(position(step)));
}
elements.push(largest);
}
}
Self::dense(shape.without_axis(axis), elements)
}
pub fn broadcast(&self, shape: Shape) -> Self {
assert_eq!(
self.logical_shape().volume(),
1,
"broadcast requires a single-element tensor"
);
Self::constant(shape, self.get(0))
}
pub fn broadcast_like(&self, reference: &Self) -> Self {
self.broadcast(reference.logical_shape().clone())
}
pub fn broadcast_along(&self, axis: usize, extent: usize) -> Self {
let shape = self.logical_shape();
assert!(
axis <= shape.rank(),
"broadcast axis {axis} is out of rank for {shape}"
);
assert!(extent > 0, "broadcast extent must be positive");
let mut axes: Vec<usize> = shape.axes().to_vec();
axes.insert(axis, extent);
let widened = Shape::new(axes);
match &self.storage {
Storage::Constant { value, .. } => Self::constant(widened, value.clone()),
Storage::Dense { data, layout } => Self {
storage: Storage::Dense {
data: Arc::clone(data),
layout: layout.broadcast_along(axis, &widened),
},
},
Storage::Selection { .. } => self.densify().broadcast_along(axis, extent),
}
}
pub fn broadcast_along_like(&self, axis: usize, reference: &Self) -> Self {
let reference_shape = reference.logical_shape();
assert!(
axis < reference_shape.rank(),
"axis {axis} is out of rank for {reference_shape}"
);
assert_eq!(
self.logical_shape(),
&reference_shape.without_axis(axis),
"broadcast along axis {axis} of {reference_shape} requires the remaining shape"
);
self.broadcast_along(axis, reference_shape.axes()[axis])
}
pub fn reshape(&self, shape: Shape) -> Self {
assert_eq!(
self.logical_shape().volume(),
shape.volume(),
"reshape from {} to {shape} changes the number of elements",
self.logical_shape()
);
match &self.storage {
Storage::Constant { value, .. } => Self::constant(shape, value.clone()),
Storage::Dense { data, layout } => match layout.reshape(shape.clone()) {
Some(reshaped) => Self {
storage: Storage::Dense {
data: Arc::clone(data),
layout: reshaped,
},
},
None => Self::dense(shape, self.to_vec()),
},
Storage::Selection { .. } => Self::dense(shape, self.to_vec()),
}
}
pub fn permute(&self, order: &[usize]) -> Self {
let shape = self.logical_shape();
assert_eq!(
order.len(),
shape.rank(),
"permute order must cover every axis of {shape}"
);
let mut seen = vec![false; shape.rank()];
for &axis in order {
assert!(
axis < shape.rank(),
"permute axis {axis} is out of rank for {shape}"
);
assert!(
!std::mem::replace(&mut seen[axis], true),
"permute order repeats axis {axis}"
);
}
match &self.storage {
Storage::Constant { value, .. } => {
let axes = shape.axes();
let permuted = Shape::new(order.iter().map(|&axis| axes[axis]));
Self::constant(permuted, value.clone())
}
Storage::Dense { data, layout } => Self {
storage: Storage::Dense {
data: Arc::clone(data),
layout: layout.permute(order),
},
},
Storage::Selection { .. } => self.densify().permute(order),
}
}
pub fn narrow(&self, axis: usize, start: usize, len: usize) -> Self {
let shape = self.logical_shape();
assert!(
axis < shape.rank(),
"narrow axis {axis} is out of rank for {shape}"
);
assert!(len > 0, "narrow window must hold at least one element");
let extent = shape.axes()[axis];
let end = start
.checked_add(len)
.expect("narrow window end overflows `usize`");
assert!(
end <= extent,
"narrow window {start}..{end} exceeds axis {axis} extent {extent}"
);
match &self.storage {
Storage::Constant { value, .. } => {
let narrowed = Shape::new(
shape
.axes()
.iter()
.enumerate()
.map(|(index, &e)| if index == axis { len } else { e }),
);
Self::constant(narrowed, value.clone())
}
Storage::Dense { data, layout } => Self {
storage: Storage::Dense {
data: Arc::clone(data),
layout: layout.narrow(axis, start, len),
},
},
Storage::Selection { .. } => self.densify().narrow(axis, start, len),
}
}
pub fn pad(&self, axis: usize, start: usize, full_extent: usize) -> Self {
let shape = self.logical_shape();
assert!(
axis < shape.rank(),
"pad axis {axis} is out of rank for {shape}"
);
let axes = shape.axes();
let len = axes[axis];
let end = start
.checked_add(len)
.expect("pad window end overflows `usize`");
assert!(
end <= full_extent,
"pad window {start}..{end} exceeds the full extent {full_extent}"
);
let outer: usize = axes[..axis].iter().product();
let inner: usize = axes[axis + 1..].iter().product();
let zero = Element::zero();
let mut elements = Vec::with_capacity(outer * full_extent * inner);
for outer_index in 0..outer {
for position in 0..full_extent {
for inner_index in 0..inner {
if position >= start && position < end {
let source = (outer_index * len + (position - start)) * inner + inner_index;
elements.push(self.get(source));
} else {
elements.push(zero.clone());
}
}
}
}
let padded = Shape::new(
axes.iter()
.enumerate()
.map(|(index, &e)| if index == axis { full_extent } else { e }),
);
Self::dense(padded, elements)
}
pub fn unfold(&self, axis: usize, size: usize, step: usize, dilation: usize) -> Self {
let shape = self.logical_shape();
assert!(
axis < shape.rank(),
"unfold axis {axis} is out of rank for {shape}"
);
assert!(size > 0, "unfold windows must hold at least one element");
assert!(step > 0, "unfold step must be positive");
assert!(dilation > 0, "unfold dilation must be positive");
let extent = shape.axes()[axis];
let span = dilation
.checked_mul(size - 1)
.and_then(|reach| reach.checked_add(1))
.expect("unfold window span overflows `usize`");
assert!(
span <= extent,
"unfold window span {span} exceeds axis {axis} extent {extent}"
);
match &self.storage {
Storage::Constant { value, .. } => {
let count = (extent - span) / step + 1;
let mut unfolded: Vec<usize> = shape.axes().to_vec();
unfolded[axis] = count;
unfolded.insert(axis + 1, size);
Self::constant(Shape::new(unfolded), value.clone())
}
Storage::Dense { data, layout } => Self {
storage: Storage::Dense {
data: Arc::clone(data),
layout: layout.unfold(axis, size, step, dilation),
},
},
Storage::Selection { .. } => self.densify().unfold(axis, size, step, dilation),
}
}
pub fn fold(
&self,
axis: usize,
size: usize,
step: usize,
dilation: usize,
extent: usize,
) -> Self {
let shape = self.logical_shape();
let axes = shape.axes();
assert!(
axis + 1 < axes.len(),
"fold window axes {axis}, {} are out of rank for {shape}",
axis + 1
);
assert!(size > 0, "fold windows must hold at least one element");
assert!(step > 0, "fold step must be positive");
assert!(dilation > 0, "fold dilation must be positive");
let span = dilation
.checked_mul(size - 1)
.and_then(|reach| reach.checked_add(1))
.expect("fold window span overflows `usize`");
assert!(
span <= extent,
"fold window span {span} exceeds the extent {extent}"
);
let count = (extent - span) / step + 1;
assert_eq!(
axes[axis], count,
"fold window count {} disagrees with the {count} windows of extent {extent}",
axes[axis]
);
assert_eq!(
axes[axis + 1],
size,
"fold window size {} disagrees with {size}",
axes[axis + 1]
);
let outer: usize = axes[..axis].iter().product();
let inner: usize = axes[axis + 2..].iter().product();
let zero = Element::zero();
let mut elements = Vec::with_capacity(outer * extent * inner);
for outer_index in 0..outer {
for position in 0..extent {
for inner_index in 0..inner {
let mut total = zero.promote();
for k in 0..size {
let reach = k * dilation;
if reach > position {
break;
}
let rest = position - reach;
if !rest.is_multiple_of(step) {
continue;
}
let window = rest / step;
if window >= count {
continue;
}
let source =
((outer_index * count + window) * size + k) * inner + inner_index;
total = total + self.get(source).promote();
}
elements.push(Element::demote(total));
}
}
}
let folded = Shape::new(
axes[..axis]
.iter()
.copied()
.chain(std::iter::once(extent))
.chain(axes[axis + 2..].iter().copied()),
);
Self::dense(folded, elements)
}
pub fn windowed_patches(
&self,
kernel_height: usize,
kernel_width: usize,
stride: usize,
padding: usize,
) -> Self {
let shape = self.logical_shape();
let axes = shape.axes();
assert_eq!(
axes.len(),
4,
"windowed_product input must be rank 4 [batch, channels, height, width], got {shape}"
);
assert!(stride > 0, "windowed_product stride must be positive");
let (batch, channels, height, width) = (axes[0], axes[1], axes[2], axes[3]);
assert!(
kernel_height > 0
&& kernel_width > 0
&& kernel_height <= height + 2 * padding
&& kernel_width <= width + 2 * padding,
"windowed_product kernel {kernel_height}x{kernel_width} does not fit {shape} \
with padding {padding}"
);
let Some(elements) = self.as_slice() else {
return composed_windowed_patches(self, kernel_height, kernel_width, stride, padding);
};
let out_height = (height + 2 * padding - kernel_height) / stride + 1;
let out_width = (width + 2 * padding - kernel_width) / stride + 1;
let columns = channels * kernel_height * kernel_width;
let zero = Element::zero();
let mut patches = vec![zero; batch * out_height * out_width * columns];
for image in 0..batch {
for out_y in 0..out_height {
for out_x in 0..out_width {
let row = ((image * out_height + out_y) * out_width + out_x) * columns;
let source_x = (out_x * stride) as isize - padding as isize;
let clip_low = (-source_x).max(0) as usize;
let clip_high = kernel_width.min((width as isize - source_x).max(0) as usize);
if clip_low >= clip_high {
continue;
}
let run = clip_high - clip_low;
for channel in 0..channels {
for kernel_y in 0..kernel_height {
let source_y = (out_y * stride + kernel_y) as isize - padding as isize;
if source_y < 0 || source_y >= height as isize {
continue;
}
let source =
((image * channels + channel) * height + source_y as usize) * width
+ (source_x + clip_low as isize) as usize;
let target = row
+ (channel * kernel_height + kernel_y) * kernel_width
+ clip_low;
patches[target..target + run]
.clone_from_slice(&elements[source..source + run]);
}
}
}
}
}
Self::dense(
Shape::new([batch * out_height * out_width, columns]),
patches,
)
}
pub fn gather(&self, selection: &Self) -> Self {
let table = self.logical_shape();
let indices = selection.selection_indices();
assert!(table.rank() >= 1, "gather table needs at least one axis");
let vocabulary = selection.logical_shape().axes()[1];
assert_eq!(
vocabulary,
table.axes()[0],
"gather selection vocabulary {vocabulary} does not match table rows {}",
table.axes()[0]
);
let row_size: usize = table.axes()[1..].iter().product();
let mut elements = Vec::with_capacity(indices.len() * row_size);
for &row in indices {
for offset in 0..row_size {
elements.push(self.get(row * row_size + offset));
}
}
let result =
Shape::new(std::iter::once(indices.len()).chain(table.axes()[1..].iter().copied()));
Self::dense(result, elements)
}
pub fn scatter(&self, selection: &Self) -> Self {
let gradient = self.logical_shape();
assert!(
gradient.rank() >= 1,
"scatter needs a gradient with a leading selection axis"
);
let indices = selection.selection_indices();
assert_eq!(
gradient.axes()[0],
indices.len(),
"scatter gradient rows disagree with the selection count"
);
let rows = selection.logical_shape().axes()[1];
let row_size: usize = gradient.axes()[1..].iter().product();
let volume = rows
.checked_mul(row_size)
.expect("shape volume overflows `usize`");
let zero = Element::zero();
let mut accumulators = vec![zero.promote(); volume];
for (source, &target) in indices.iter().enumerate() {
for offset in 0..row_size {
let position = target * row_size + offset;
accumulators[position] =
accumulators[position].clone() + self.get(source * row_size + offset).promote();
}
}
let result = Shape::new(std::iter::once(rows).chain(gradient.axes()[1..].iter().copied()));
Self::dense(
result,
accumulators.into_iter().map(Element::demote).collect(),
)
}
}
impl<Element: Elementary> Tensor<Element> {
pub fn windowed_product(
&self,
kernel: &Self,
kernel_height: usize,
kernel_width: usize,
stride: usize,
padding: usize,
) -> Self {
self.windowed_patches(kernel_height, kernel_width, stride, padding)
.matmul(kernel)
}
}
impl<E: Element> Recordable for Tensor<E> {
fn shape(&self) -> Shape {
Tensor::shape(self)
}
fn zero_like(&self) -> Self {
Tensor::zero_like(self)
}
fn one_like(&self) -> Self {
Tensor::one_like(self)
}
fn exp(&self) -> Self {
Tensor::exp(self)
}
fn ln(&self) -> Self {
Tensor::ln(self)
}
fn sqrt(&self) -> Self {
Tensor::sqrt(self)
}
fn tanh(&self) -> Self {
Tensor::tanh(self)
}
fn sin(&self) -> Self {
Tensor::sin(self)
}
fn cos(&self) -> Self {
Tensor::cos(self)
}
fn log1p(&self) -> Self {
Tensor::log1p(self)
}
fn expm1(&self) -> Self {
Tensor::expm1(self)
}
fn erf(&self) -> Self {
Tensor::erf(self)
}
fn erf_derivative(&self) -> Self {
Tensor::erf_derivative(self)
}
fn powf(&self, exponent: Self) -> Self {
Tensor::powf(self, exponent)
}
fn maximum(&self, other: &Self) -> Self {
Tensor::maximum(self, other)
}
fn step(&self, threshold: &Self) -> Self {
Tensor::step(self, threshold)
}
fn matmul(&self, rhs: &Self) -> Self {
Tensor::matmul(self, rhs)
}
fn sum(&self) -> Self {
Tensor::sum(self)
}
fn sum_along(&self, axis: usize) -> Self {
Tensor::sum_along(self, axis)
}
fn logsumexp(&self, axis: usize) -> Self {
Tensor::logsumexp(self, axis)
}
fn log_softmax(&self, axis: usize) -> Self {
Tensor::log_softmax(self, axis)
}
fn broadcast(&self, shape: Shape) -> Self {
Tensor::broadcast(self, shape)
}
fn broadcast_along(&self, axis: usize, extent: usize) -> Self {
Tensor::broadcast_along(self, axis, extent)
}
fn reshape(&self, shape: Shape) -> Self {
Tensor::reshape(self, shape)
}
fn permute(&self, order: &[usize]) -> Self {
Tensor::permute(self, order)
}
fn narrow(&self, axis: usize, start: usize, len: usize) -> Self {
Tensor::narrow(self, axis, start, len)
}
fn pad(&self, axis: usize, start: usize, full_extent: usize) -> Self {
Tensor::pad(self, axis, start, full_extent)
}
fn unfold(&self, axis: usize, size: usize, step: usize, dilation: usize) -> Self {
Tensor::unfold(self, axis, size, step, dilation)
}
fn fold(&self, axis: usize, size: usize, step: usize, dilation: usize, extent: usize) -> Self {
Tensor::fold(self, axis, size, step, dilation, extent)
}
fn gather(&self, selection: &Self) -> Self {
Tensor::gather(self, selection)
}
fn scatter(&self, selection: &Self) -> Self {
Tensor::scatter(self, selection)
}
}
#[cfg(test)]
#[path = "tests/tensor_tests.rs"]
mod tests;