// tensor.mbt 鈥?Tensor struct + element-wise ops + broadcasting.
//
// Design:
// - Tensor wraps a flat `Array[Float]` plus a `shape` array.
// - Memory layout is row-major (C-order), same as NCHW conv /
// maxpool / linear / etc.
// - Reshape is a no-op data-wise (only updates shape) when total
// element count is preserved.
// - Element-wise ops use **broadcasting** following the standard
// right-align rule:
// align shapes from the right; for each axis dim is compatible if
// equal or one of them is 1; missing leading dims are treated as
// size-1 (implicit broadcast).
// - All ops return a fresh Tensor (no aliasing of inputs).
///|
/// Tensor struct: row-major flat data + shape.
pub struct Tensor {
data : Array[Float]
shape : Array[Int]
}
///|
fn total_size(shape : Array[Int]) -> Int {
let mut n = 1
for i in 0.. Tensor {
{ data: Array::make(total_size(shape), 0.0F), shape }
}
///|
/// Build a Tensor with the given shape, filled with ones.
pub fn Tensor::ones(shape : Array[Int]) -> Tensor {
{ data: Array::make(total_size(shape), 1.0F), shape }
}
///|
/// Build a Tensor from a flat data array + shape. Does NOT copy data.
pub fn Tensor::from(data : Array[Float], shape : Array[Int]) -> Tensor {
{ data, shape }
}
///|
/// Reshape to a new shape. Panics if total element count mismatches.
pub fn Tensor::reshape(self : Tensor, shape : Array[Int]) -> Tensor {
if total_size(shape) != self.data.length() {
abort("Tensor::reshape: total size mismatch")
}
{ data: self.data, shape }
}
///|
/// Element count.
pub fn Tensor::numel(self : Tensor) -> Int {
self.data.length()
}
///|
/// Compute the broadcast shape of two input shapes. Returns
/// `[-1]` if incompatible.
fn broadcast_shape(a : Array[Int], b : Array[Int]) -> Array[Int] {
let ndim = if a.length() > b.length() { a.length() } else { b.length() }
let out : Array[Int] = Array::make(ndim, 1)
let a_offset = ndim - a.length()
let b_offset = ndim - b.length()
let mut i = 0
while i < ndim {
let ai = if i < a_offset { 1 } else { a[i - a_offset] }
let bi = if i < b_offset { 1 } else { b[i - b_offset] }
if ai == bi {
out[i] = ai
} else if ai == 1 {
out[i] = bi
} else if bi == 1 {
out[i] = ai
} else {
return [-1]
}
i = i + 1
}
out
}
///|
/// Compute the row-major strides for a shape (right-most dim is stride 1).
fn compute_strides(shape : Array[Int]) -> Array[Int] {
let n = shape.length()
let strides : Array[Int] = Array::make(n, 1)
if n == 0 { return strides }
let mut j = n - 2
while j >= 0 {
strides[j] = strides[j + 1] * shape[j + 1]
j = j - 1
}
strides
}
///|
/// Apply a binary element-wise op on two broadcast-compatible tensors.
/// `f(a, b) -> Float` is the element-wise function.
fn elementwise_binary(
a : Tensor,
b : Tensor,
f : (Float, Float) -> Float,
) -> Tensor {
let out_shape = broadcast_shape(a.shape, b.shape)
if out_shape.length() == 1 && out_shape[0] == -1 {
abort("elementwise: incompatible broadcast shapes")
}
let n = total_size(out_shape)
let out_data : Array[Float] = Array::make(n, 0.0F)
let out_strides = compute_strides(out_shape)
let a_strides = compute_strides(a.shape)
let b_strides = compute_strides(b.shape)
let out_offset_a = out_shape.length() - a.shape.length()
let out_offset_b = out_shape.length() - b.shape.length()
let mut i = 0
while i < n {
let mut rem = i
let mut ai = 0
let mut bi = 0
let mut j = 0
while j < out_shape.length() {
let coord = rem / out_strides[j]
rem = rem - coord * out_strides[j]
let a_axis = j - out_offset_a
if a_axis >= 0 {
if a.shape[a_axis] != 1 {
ai = ai + coord * a_strides[a_axis]
}
}
let b_axis = j - out_offset_b
if b_axis >= 0 {
if b.shape[b_axis] != 1 {
bi = bi + coord * b_strides[b_axis]
}
}
j = j + 1
}
out_data[i] = f(a.data[ai], b.data[bi])
i = i + 1
}
{ data: out_data, shape: out_shape }
}
///|
/// Tensor add (with broadcasting).
pub fn tensor_add(a : Tensor, b : Tensor) -> Tensor {
elementwise_binary(a, b, fn(x, y) { x + y })
}
///|
/// Tensor subtract (with broadcasting).
pub fn tensor_sub(a : Tensor, b : Tensor) -> Tensor {
elementwise_binary(a, b, fn(x, y) { x - y })
}
///|
/// Tensor multiply (with broadcasting).
pub fn tensor_mul(a : Tensor, b : Tensor) -> Tensor {
elementwise_binary(a, b, fn(x, y) { x * y })
}
///|
/// Tensor divide (with broadcasting).
pub fn tensor_div(a : Tensor, b : Tensor) -> Tensor {
elementwise_binary(a, b, fn(x, y) { x / y })
}