// autograd.mbt
///|
/// Enum representing operators in the computation graph.
pub enum Op {
Add(Tensor, Tensor)
Sub(Tensor, Tensor)
Mul(Tensor, Tensor)
Div(Tensor, Tensor)
MatMul(Tensor, Tensor)
Reshape(Tensor, Array[Int])
Transpose(Tensor, Int, Int)
Slice(Tensor, Int, Int, Int)
SliceMulti(Tensor, Array[(Int, Int)])
ReLU(Tensor)
Sigmoid(Tensor)
Tanh(Tensor)
MSELoss(Tensor, Tensor)
ReduceSum(Tensor, Int?, Bool, Bool)
Softmax(Tensor, Int)
LogSoftmax(Tensor, Int)
CrossEntropy(Tensor, Tensor)
Abs(Tensor)
Sqrt(Tensor)
Exp(Tensor)
Log(Tensor)
Pow(Tensor, Double)
Clamp(Tensor, Double, Double)
}
///|
/// Helper to get the inputs of an operation.
fn get_op_inputs(op : Op) -> Array[Tensor] {
match op {
Add(a, b) => [a, b]
Sub(a, b) => [a, b]
Mul(a, b) => [a, b]
Div(a, b) => [a, b]
MatMul(a, b) => [a, b]
Reshape(a, _) => [a]
Transpose(a, _, _) => [a]
Slice(a, _, _, _) => [a]
SliceMulti(a, _) => [a]
ReLU(a) => [a]
Sigmoid(a) => [a]
Tanh(a) => [a]
MSELoss(a, b) => [a, b]
ReduceSum(a, _, _, _) => [a]
Softmax(a, _) => [a]
LogSoftmax(a, _) => [a]
CrossEntropy(a, b) => [a, b]
Abs(a) => [a]
Sqrt(a) => [a]
Exp(a) => [a]
Log(a) => [a]
Pow(a, _) => [a]
Clamp(a, _, _) => [a]
}
}
///|
/// Performs topological sort starting from the root tensor.
fn topological_sort(root : Tensor) -> Array[Tensor] {
let visited : Array[Tensor] = Array::new()
let result : Array[Tensor] = Array::new()
fn visit(node : Tensor) {
let mut found = false
for v in visited {
if v.id == node.id {
found = true
break
}
}
if !found {
visited.push(node)
match node.creator {
Some(op) => {
let inputs = get_op_inputs(op)
for input in inputs {
visit(input)
}
}
None => ()
}
result.push(node)
}
}
visit(root)
result
}
///|
/// Sum-reduce gradient delta along broadcasted dimensions.
fn reduce_grad(
delta : Array[Double],
out_shape : Array[Int],
target_shape : Array[Int],
) -> Array[Double] {
let target_len = target_shape.length()
let mut target_size = 1
for dim in target_shape {
target_size = target_size * dim
}
if target_len == 0 {
target_size = 1
}
let target_grad = Array::make(target_size, 0.0)
let out_strides = shape_to_strides(out_shape)
let target_strides = shape_to_strides(target_shape)
let size = delta.length()
for i in 0.. Unit {
if !t.requires_grad {
return
}
match t.grad {
None => {
let g = Array::make(t.data.length(), 0.0)
for i in 0..
for i in 0.. Unit {
if !t.requires_grad {
return
}
let reduced = if t.shape != out_shape {
reduce_grad(delta, out_shape, t.shape)
} else {
delta
}
match t.grad {
None => {
let g = Array::make(t.data.length(), 0.0)
for i in 0..
for i in 0.. Unit {
match op {
Add(a, b) => {
accumulate_grad(a, grad, out_shape)
accumulate_grad(b, grad, out_shape)
}
Sub(a, b) => {
accumulate_grad(a, grad, out_shape)
let neg_grad = Array::make(grad.length(), 0.0)
for i in 0.. {
let out_strides = shape_to_strides(out_shape)
if a.requires_grad {
let grad_a = Array::make(grad.length(), 0.0)
for i in 0.. {
let out_strides = shape_to_strides(out_shape)
if a.requires_grad {
let grad_a = Array::make(grad.length(), 0.0)
for i in 0.. {
let m = a.shape[0]
let k = a.shape[1]
let n = b.shape[1]
if a.requires_grad {
let grad_a = Array::make(m * k, 0.0)
for r in 0.. accumulate_grad_direct(a, grad)
Transpose(a, dim0, dim1) =>
if a.requires_grad {
let size = grad.length()
let grad_a = Array::make(a.data.length(), 0.0)
let len = a.shape.length()
let out_strides = shape_to_strides(out_shape)
for i in 0..
if a.requires_grad {
let size = grad.length()
let grad_a = Array::make(a.data.length(), 0.0)
let len = a.shape.length()
let out_strides = shape_to_strides(out_shape)
for i in 0..
if a.requires_grad {
let size = grad.length()
let grad_a = Array::make(a.data.length(), 0.0)
let len = a.shape.length()
let out_strides = shape_to_strides(out_shape)
for i in 0..
if a.requires_grad {
let grad_a = Array::make(grad.length(), 0.0)
for i in 0.. 0.0 { grad[i] } else { 0.0 }
}
accumulate_grad(a, grad_a, out_shape)
}
Sigmoid(a) =>
if a.requires_grad {
let grad_a = Array::make(grad.length(), 0.0)
for i in 0..
if a.requires_grad {
let grad_a = Array::make(grad.length(), 0.0)
for i in 0..
if a.requires_grad {
let grad_a = Array::make(grad.length(), 0.0)
for i in 0.. 0.0 {
1.0
} else if a.data[i] < 0.0 {
-1.0
} else {
0.0
}
grad_a[i] = grad[i] * sign
}
accumulate_grad_direct(a, grad_a)
}
Sqrt(a) =>
if a.requires_grad {
let grad_a = Array::make(grad.length(), 0.0)
for i in 0..
if a.requires_grad {
let grad_a = Array::make(grad.length(), 0.0)
for i in 0..
if a.requires_grad {
let grad_a = Array::make(grad.length(), 0.0)
for i in 0..
if a.requires_grad {
let grad_a = Array::make(grad.length(), 0.0)
for i in 0..= 0 {
for _ in 0..
if a.requires_grad {
let grad_a = Array::make(grad.length(), 0.0)
for i in 0.. lower && value < upper { grad[i] } else { 0.0 }
}
accumulate_grad_direct(a, grad_a)
}
ReduceSum(a, dim, keepdim, mean) =>
if a.requires_grad {
let grad_a = Array::make(a.data.length(), 0.0)
match dim {
None => {
let factor = if mean {
1.0 / a.data.length().to_double()
} else {
1.0
}
for i in 0.. {
let count = a.shape[d].to_double()
let factor = if mean { 1.0 / count } else { 1.0 }
let output_strides = shape_to_strides(out_shape)
for i in 0..
if a.requires_grad {
let values = softmax_data(a, dim)
let width = a.shape[dim]
let stride = a.strides[dim]
let grad_a = Array::make(a.data.length(), 0.0)
for i in 0..
if a.requires_grad {
let values = log_softmax_data(a, dim)
let width = a.shape[dim]
let stride = a.strides[dim]
let grad_a = Array::make(a.data.length(), 0.0)
for i in 0..
if logits.requires_grad {
let batch = logits.shape[0]
let classes = logits.shape[1]
let grad_logits = Array::make(logits.data.length(), 0.0)
for row in 0.. maximum {
maximum = logits.data[offset + col]
}
}
let mut denominator = 0.0
for col in 0.. {
let size = pred.data.length()
let factor = 2.0 / size.to_double() * grad[0]
if pred.requires_grad {
let grad_pred = Array::make(size, 0.0)
for i in 0.. Unit {
if self.grad is None {
self.grad = Some(Array::make(self.data.length(), 1.0))
}
let sorted = topological_sort(self)
let len = sorted.length()
for i in 0.. ()
Some(op) =>
match t.grad {
None => ()
Some(g) => propagate_op_grad(op, g, t.shape)
}
}
}
}