use oxmera_core::{DType, Device, Error, Result, Shape};
use crate::autograd::{GradFn, is_recording};
use crate::backend::{Backend, BinaryOp, ReduceOp, UnaryOp, backend_for};
use crate::tensor::{Tensor, ViewKind};
use std::sync::Arc;
fn same_device(a: &Tensor, b: &Tensor, op: &'static str) -> Result<Device> {
if a.device() != b.device() {
return Err(Error::DeviceMismatch {
lhs: a.device(),
rhs: b.device(),
op,
});
}
Ok(a.device())
}
fn record(
out: Tensor,
inputs: Vec<Tensor>,
vjp: impl Fn(&Tensor) -> Result<Vec<Option<Tensor>>> + Send + Sync + 'static,
) -> Tensor {
if is_recording() && inputs.iter().any(Tensor::is_tracked) {
out.with_grad_fn(GradFn {
inputs,
vjp: Box::new(vjp),
})
} else {
out
}
}
pub(crate) fn reduce_to_shape(grad: &Tensor, shape: &Shape) -> Result<Tensor> {
if grad.shape() == shape {
return Ok(grad.clone());
}
let gdims = grad.dims().to_vec();
let tdims = shape.dims();
let lead = gdims.len() - tdims.len();
let mut axes: Vec<usize> = (0..lead).collect();
for (i, &td) in tdims.iter().enumerate() {
if td == 1 && gdims[lead + i] != 1 {
axes.push(lead + i);
}
}
let reduced = if axes.is_empty() {
grad.clone()
} else {
grad.sum_keepdim(&axes, true)?
};
reduced.reshape(shape.clone())
}
pub(crate) fn record_view(input: &Tensor, out: Tensor, kind: ViewKind) -> Tensor {
let in_shape = input.shape().clone();
record(out, vec![input.clone()], move |g| {
let gi = match &kind {
ViewKind::Reshape | ViewKind::Contiguous => g.reshape(in_shape.clone())?,
ViewKind::Permute(perm) => {
let mut inverse = vec![0usize; perm.len()];
for (i, &p) in perm.iter().enumerate() {
inverse[p] = i;
}
g.permute(&inverse)?
}
ViewKind::Narrow { dim, start, len } => {
let indices: Vec<i64> = (*start..start + len).map(|i| i as i64).collect();
let indices = Tensor::from_vec_i64(indices, Shape::from([*len]))?;
Tensor::zeros(in_shape.clone())
.to_dtype(g.dtype())?
.to_device(g.device())?
.index_add(*dim, &indices, g)?
}
ViewKind::Broadcast => reduce_to_shape(g, &in_shape)?,
};
Ok(vec![Some(gi)])
})
}
impl Tensor {
fn backend(&self) -> Result<Arc<dyn Backend>> {
backend_for(self.device())
}
fn unary_op(&self, op: UnaryOp) -> Result<Tensor> {
let out = self.backend()?.unary(op, self)?;
let a = self.clone();
let o = out.clone();
Ok(record(out, vec![self.clone()], move |g| {
let gi = match op {
UnaryOp::Neg => g.neg()?,
UnaryOp::Exp => g.mul(&o)?,
UnaryOp::Ln => g.div(&a)?,
UnaryOp::Abs => {
let sign = a
.gt_mask(&Tensor::scalar_on(&a, 0.0)?)?
.sub(&Tensor::scalar_on(&a, 0.0)?.gt_mask(&a)?)?;
g.mul(&sign)?
}
UnaryOp::Sqrt => g.mul(&Tensor::scalar_on(&a, 0.5)?)?.div(&o)?,
UnaryOp::Sin => g.mul(&a.cos()?)?,
UnaryOp::Cos => g.mul(&a.sin()?.neg()?)?,
UnaryOp::Tanh => {
let one = Tensor::scalar_on(&a, 1.0)?;
g.mul(&one.sub(&o.mul(&o)?)?)?
}
UnaryOp::Relu => g.mul(&a.gt_mask(&Tensor::scalar_on(&a, 0.0)?)?)?,
UnaryOp::Gelu => {
let c = Tensor::scalar_on(&a, 0.797_884_6)?;
let k = Tensor::scalar_on(&a, 0.044_715)?;
let one = Tensor::scalar_on(&a, 1.0)?;
let half = Tensor::scalar_on(&a, 0.5)?;
let three_k = Tensor::scalar_on(&a, 3.0 * 0.044_715)?;
let x2 = a.mul(&a)?;
let u = c.mul(&a.add(&k.mul(&x2.mul(&a)?)?)?)?;
let t = u.tanh()?;
let sech2 = one.sub(&t.mul(&t)?)?;
let du = c.mul(&one.add(&three_k.mul(&x2)?)?)?;
let d = half
.mul(&one.add(&t)?)?
.add(&half.mul(&a)?.mul(&sech2)?.mul(&du)?)?;
g.mul(&d)?
}
UnaryOp::Sigmoid => {
let one = Tensor::scalar_on(&a, 1.0)?;
g.mul(&o)?.mul(&one.sub(&o)?)?
}
};
Ok(vec![Some(gi)])
}))
}
pub fn neg(&self) -> Result<Tensor> {
self.unary_op(UnaryOp::Neg)
}
pub fn exp(&self) -> Result<Tensor> {
self.unary_op(UnaryOp::Exp)
}
pub fn ln(&self) -> Result<Tensor> {
self.unary_op(UnaryOp::Ln)
}
pub fn abs(&self) -> Result<Tensor> {
self.unary_op(UnaryOp::Abs)
}
pub fn sqrt(&self) -> Result<Tensor> {
self.unary_op(UnaryOp::Sqrt)
}
pub fn sin(&self) -> Result<Tensor> {
self.unary_op(UnaryOp::Sin)
}
pub fn cos(&self) -> Result<Tensor> {
self.unary_op(UnaryOp::Cos)
}
pub fn tanh(&self) -> Result<Tensor> {
self.unary_op(UnaryOp::Tanh)
}
pub fn relu(&self) -> Result<Tensor> {
self.unary_op(UnaryOp::Relu)
}
pub fn gelu(&self) -> Result<Tensor> {
self.unary_op(UnaryOp::Gelu)
}
pub fn sigmoid(&self) -> Result<Tensor> {
self.unary_op(UnaryOp::Sigmoid)
}
pub fn scalar_on(like: &Tensor, value: f32) -> Result<Tensor> {
let s = match like.dtype() {
DType::F64 => Tensor::from_vec_f64(vec![f64::from(value)], Shape::from([]))?,
_ => Tensor::scalar(value),
};
s.to_device(like.device())
}
pub fn to_dtype(&self, dtype: DType) -> Result<Tensor> {
if self.dtype() == dtype {
return Ok(self.clone());
}
if self.device() != Device::Cpu {
return Err(Error::UnsupportedDType {
dtype,
op: "to_dtype (device tensors are f32; convert on the CPU)",
});
}
let shape = self.shape().clone();
let out = match (self.dtype(), dtype) {
(DType::F32, DType::F64) => Tensor::from_vec_f64(
self.to_vec_f32()?.into_iter().map(f64::from).collect(),
shape,
)?,
(DType::F64, DType::F32) => Tensor::from_vec_f32(
self.to_vec_f64()?.into_iter().map(|x| x as f32).collect(),
shape,
)?,
(DType::I64, DType::F32) => Tensor::from_vec_f32(
self.to_vec_i64()?.into_iter().map(|x| x as f32).collect(),
shape,
)?,
(DType::I64, DType::F64) => Tensor::from_vec_f64(
self.to_vec_i64()?.into_iter().map(|x| x as f64).collect(),
shape,
)?,
(_, to) => {
return Err(Error::UnsupportedDType {
dtype: to,
op: "to_dtype",
});
}
};
let from = self.dtype();
Ok(record(out, vec![self.clone()], move |g| {
Ok(vec![Some(g.to_dtype(from)?)])
}))
}
fn binary_op(&self, op: BinaryOp, rhs: &Tensor) -> Result<Tensor> {
let device = same_device(self, rhs, "binary")?;
let out = backend_for(device)?.binary(op, self, rhs)?;
let (a, b) = (self.clone(), rhs.clone());
let o = out.clone();
Ok(record(out, vec![self.clone(), rhs.clone()], move |g| {
let (ga, gb): (Option<Tensor>, Option<Tensor>) = match op {
BinaryOp::Add => (Some(g.clone()), Some(g.clone())),
BinaryOp::Sub => (Some(g.clone()), Some(g.neg()?)),
BinaryOp::Mul => (Some(g.mul_raw(&b)?), Some(g.mul_raw(&a)?)),
BinaryOp::Div => {
let ga = g.div_raw(&b)?;
let gb = g.mul_raw(&o)?.div_raw(&b)?.neg()?;
(Some(ga), Some(gb))
}
BinaryOp::Pow => {
let one = Tensor::scalar_on(&a, 1.0)?;
let ga = g.mul_raw(&b)?.mul_raw(&a.pow(&b.sub(&one)?)?)?;
let gb = g.mul_raw(&o)?.mul_raw(&a.ln()?)?;
(Some(ga), Some(gb))
}
BinaryOp::Maximum => {
let mask = a.gt_mask(&b)?;
let one = Tensor::scalar_on(&a, 1.0)?;
let ga = g.mul_raw(&mask)?;
let gb = g.mul_raw(&one.sub(&mask)?)?;
(Some(ga), Some(gb))
}
BinaryOp::Minimum => {
let mask = b.gt_mask(&a)?;
let one = Tensor::scalar_on(&a, 1.0)?;
let ga = g.mul_raw(&mask)?;
let gb = g.mul_raw(&one.sub(&mask)?)?;
(Some(ga), Some(gb))
}
BinaryOp::Gt | BinaryOp::Eq => (None, None),
};
let ga = match ga {
Some(t) => Some(reduce_to_shape(&t, a.shape())?),
None => None,
};
let gb = match gb {
Some(t) => Some(reduce_to_shape(&t, b.shape())?),
None => None,
};
Ok(vec![ga, gb])
}))
}
fn mul_raw(&self, rhs: &Tensor) -> Result<Tensor> {
let device = same_device(self, rhs, "mul")?;
backend_for(device)?.binary(BinaryOp::Mul, self, rhs)
}
fn div_raw(&self, rhs: &Tensor) -> Result<Tensor> {
let device = same_device(self, rhs, "div")?;
backend_for(device)?.binary(BinaryOp::Div, self, rhs)
}
pub fn add(&self, rhs: &Tensor) -> Result<Tensor> {
self.binary_op(BinaryOp::Add, rhs)
}
pub fn sub(&self, rhs: &Tensor) -> Result<Tensor> {
self.binary_op(BinaryOp::Sub, rhs)
}
pub fn mul(&self, rhs: &Tensor) -> Result<Tensor> {
self.binary_op(BinaryOp::Mul, rhs)
}
pub fn div(&self, rhs: &Tensor) -> Result<Tensor> {
self.binary_op(BinaryOp::Div, rhs)
}
pub fn pow(&self, rhs: &Tensor) -> Result<Tensor> {
self.binary_op(BinaryOp::Pow, rhs)
}
pub fn maximum(&self, rhs: &Tensor) -> Result<Tensor> {
self.binary_op(BinaryOp::Maximum, rhs)
}
pub fn minimum(&self, rhs: &Tensor) -> Result<Tensor> {
self.binary_op(BinaryOp::Minimum, rhs)
}
pub fn gt_mask(&self, rhs: &Tensor) -> Result<Tensor> {
self.binary_op(BinaryOp::Gt, rhs)
}
pub fn eq_mask(&self, rhs: &Tensor) -> Result<Tensor> {
self.binary_op(BinaryOp::Eq, rhs)
}
pub fn add_scalar(&self, s: f32) -> Result<Tensor> {
self.add(&Tensor::scalar_on(self, s)?)
}
pub fn mul_scalar(&self, s: f32) -> Result<Tensor> {
self.mul(&Tensor::scalar_on(self, s)?)
}
pub fn matmul(&self, rhs: &Tensor) -> Result<Tensor> {
let device = same_device(self, rhs, "matmul")?;
if self.ndim() > 3 || rhs.ndim() > 3 {
return self.matmul_lowered(rhs);
}
let out = backend_for(device)?.matmul(self, rhs)?;
let (a, b) = (self.clone(), rhs.clone());
Ok(record(out, vec![self.clone(), rhs.clone()], move |g| {
let ga = backend_for(g.device())?.matmul(g, &b.t()?)?;
let gb = backend_for(g.device())?.matmul(&a.t()?, g)?;
Ok(vec![
Some(reduce_to_shape(&ga, a.shape())?),
Some(reduce_to_shape(&gb, b.shape())?),
])
}))
}
fn matmul_lowered(&self, rhs: &Tensor) -> Result<Tensor> {
let (ad, bd) = (self.dims(), rhs.dims());
if ad.len() < 2 || bd.len() < 2 {
return Err(Error::InvalidArgument {
op: "matmul",
detail: format!("operands need rank >= 2; got {}x{}", ad.len(), bd.len()),
});
}
let (m, k) = (ad[ad.len() - 2], ad[ad.len() - 1]);
let (kb, n) = (bd[bd.len() - 2], bd[bd.len() - 1]);
if k != kb {
return Err(Error::ShapeMismatch {
expected: Shape::new(bd[..bd.len() - 2].iter().copied().chain([k, n]).collect()),
got: rhs.shape().clone(),
op: "matmul",
});
}
let a_batch = Shape::new(ad[..ad.len() - 2].to_vec());
let b_batch = Shape::new(bd[..bd.len() - 2].to_vec());
let batch = oxmera_core::shape::broadcast_shapes(&a_batch, &b_batch).map_err(|_| {
Error::BroadcastIncompatible {
lhs: self.shape().clone(),
rhs: rhs.shape().clone(),
}
})?;
let batch_numel = batch.numel();
let lower = |t: &Tensor, own: &Shape, rows: usize, cols: usize| -> Result<Tensor> {
if own.numel() == 1 {
return t.reshape(Shape::from([rows, cols]));
}
let full: Vec<usize> = batch.dims().iter().copied().chain([rows, cols]).collect();
let expanded = if own.dims() == batch.dims() {
t.clone()
} else {
let lead = batch.ndim() - own.ndim();
let padded: Vec<usize> = std::iter::repeat_n(1usize, lead)
.chain(own.dims().iter().copied())
.chain([rows, cols])
.collect();
t.reshape(Shape::new(padded))?
.broadcast_to(Shape::new(full.clone()))?
.contiguous()?
};
expanded.reshape(Shape::from([batch_numel, rows, cols]))
};
let a3 = lower(self, &a_batch, m, k)?;
let b3 = lower(rhs, &b_batch, k, n)?;
let out = a3.matmul(&b3)?;
let out_shape: Vec<usize> = batch.dims().iter().copied().chain([m, n]).collect();
out.reshape(Shape::new(out_shape))
}
fn reduce_op(&self, op: ReduceOp, axes: &[usize], keepdim: bool) -> Result<Tensor> {
let axes = normalize_axes(axes, self.ndim(), "reduce")?;
let out = self.backend()?.reduce(op, self, &axes, keepdim)?;
let a = self.clone();
let o = out.clone();
let axes_c = axes.clone();
Ok(record(out, vec![self.clone()], move |g| {
let g_keep = if keepdim {
g.clone()
} else {
unsqueeze_axes(g, &axes_c)?
};
let gi = match op {
ReduceOp::Sum => g_keep.broadcast_to(a.shape().clone())?.contiguous()?,
ReduceOp::Max | ReduceOp::Min => {
let o_keep = if keepdim {
o.clone()
} else {
unsqueeze_axes(&o, &axes_c)?
};
let mask = a.eq_mask(&o_keep.broadcast_to(a.shape().clone())?)?;
let count = mask.sum_keepdim(&axes_c, true)?;
g_keep
.broadcast_to(a.shape().clone())?
.mul_raw(&mask)?
.div_raw(&count.broadcast_to(a.shape().clone())?.contiguous()?)?
}
};
Ok(vec![Some(gi)])
}))
}
pub fn sum(&self, axes: &[usize]) -> Result<Tensor> {
self.reduce_op(ReduceOp::Sum, axes, false)
}
pub fn sum_keepdim(&self, axes: &[usize], keepdim: bool) -> Result<Tensor> {
self.reduce_op(ReduceOp::Sum, axes, keepdim)
}
pub fn max(&self, axes: &[usize]) -> Result<Tensor> {
self.reduce_op(ReduceOp::Max, axes, false)
}
pub fn max_keepdim(&self, axes: &[usize], keepdim: bool) -> Result<Tensor> {
self.reduce_op(ReduceOp::Max, axes, keepdim)
}
pub fn min(&self, axes: &[usize]) -> Result<Tensor> {
self.reduce_op(ReduceOp::Min, axes, false)
}
pub fn mean(&self, axes: &[usize]) -> Result<Tensor> {
self.mean_keepdim(axes, false)
}
pub fn mean_keepdim(&self, axes: &[usize], keepdim: bool) -> Result<Tensor> {
let axes_n = normalize_axes(axes, self.ndim(), "mean")?;
let n: usize = axes_n.iter().map(|&ax| self.dims()[ax]).product();
self.sum_keepdim(&axes_n, keepdim)?
.mul_scalar(1.0 / n as f32)
}
pub fn argmax(&self, dim: usize, keepdim: bool) -> Result<Tensor> {
if dim >= self.ndim() {
return Err(Error::InvalidArgument {
op: "argmax",
detail: format!("dim {dim} out of range for rank {}", self.ndim()),
});
}
self.backend()?.argmax(self, dim, keepdim)
}
pub fn softmax(&self, dim: usize) -> Result<Tensor> {
let shifted = self.sub(&self.max_keepdim(&[dim], true)?.detach())?;
let e = shifted.exp()?;
let denom = e.sum_keepdim(&[dim], true)?;
e.div(&denom)
}
pub fn log_softmax(&self, dim: usize) -> Result<Tensor> {
let shifted = self.sub(&self.max_keepdim(&[dim], true)?.detach())?;
let lse = shifted.exp()?.sum_keepdim(&[dim], true)?.ln()?;
shifted.sub(&lse)
}
pub fn index_select(&self, dim: usize, indices: &Tensor) -> Result<Tensor> {
let out = dispatch_index(self, &[indices], |be, t, extra| {
be.index_select(t, dim, &extra[0])
})?;
let in_shape = self.shape().clone();
let idx = indices.clone();
Ok(record(out, vec![self.clone()], move |g| {
let zeros = Tensor::zeros(in_shape.clone())
.to_dtype(g.dtype())?
.to_device(g.device())?;
Ok(vec![Some(zeros.index_add(dim, &idx, g)?)])
}))
}
pub fn index_add(&self, dim: usize, indices: &Tensor, src: &Tensor) -> Result<Tensor> {
same_device(self, src, "index_add")?;
let out = dispatch_index(self, &[indices, src], |be, t, extra| {
be.index_add(t, dim, &extra[0], &extra[1])
})?;
let idx = indices.clone();
Ok(record(out, vec![self.clone(), src.clone()], move |g| {
Ok(vec![Some(g.clone()), Some(g.index_select(dim, &idx)?)])
}))
}
pub fn to_device(&self, device: Device) -> Result<Tensor> {
if self.device() == device {
return Ok(self.clone());
}
if self.dtype() == DType::F64 {
return Err(Error::UnsupportedDType {
dtype: DType::F64,
op: "to_device (f64 tensors live on the CPU; to_dtype(F32) first)",
});
}
let out = match (self.device(), device) {
(Device::Cpu, target) => backend_for(target)?.upload(&self.contiguous_data()?)?,
(_, Device::Cpu) => self.backend()?.download(self)?,
(_, target) => {
let host = self.backend()?.download(self)?;
backend_for(target)?.upload(&host)?
}
};
let source = self.device();
Ok(record(out, vec![self.clone()], move |g| {
Ok(vec![Some(g.to_device(source)?)])
}))
}
}
fn unsqueeze_axes(t: &Tensor, axes: &[usize]) -> Result<Tensor> {
let mut out = t.clone();
let mut sorted = axes.to_vec();
sorted.sort_unstable();
for &ax in &sorted {
out = out.unsqueeze(ax)?;
}
Ok(out)
}
fn normalize_axes(axes: &[usize], ndim: usize, op: &'static str) -> Result<Vec<usize>> {
let mut axes: Vec<usize> = if axes.is_empty() {
(0..ndim).collect()
} else {
axes.to_vec()
};
axes.sort_unstable();
axes.dedup();
if let Some(&bad) = axes.iter().find(|&&a| a >= ndim) {
return Err(Error::InvalidArgument {
op,
detail: format!("axis {bad} out of range for rank {ndim}"),
});
}
Ok(axes)
}
fn dispatch_index(
t: &Tensor,
extra: &[&Tensor],
f: impl Fn(&dyn Backend, &Tensor, &[Tensor]) -> Result<Tensor>,
) -> Result<Tensor> {
let backend = backend_for(t.device())?;
let on_device: Vec<Tensor> = extra.iter().map(|e| (*e).clone()).collect();
match f(backend.as_ref(), t, &on_device) {
Err(Error::NotImplemented { .. }) if t.device() != Device::Cpu => {
let cpu = backend.download(t)?;
let cpu_extra: Vec<Tensor> = extra
.iter()
.map(|e| e.to_device(Device::Cpu))
.collect::<Result<_>>()?;
let cpu_backend = backend_for(Device::Cpu)?;
let out = f(cpu_backend.as_ref(), &cpu, &cpu_extra)?;
backend_for(t.device())?.upload(&out)
}
other => other,
}
}