Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions core/src/main/scala/dimwit/autodiff/FloatTree.scala
Original file line number Diff line number Diff line change
Expand Up @@ -104,6 +104,13 @@ object FloatTree:
def **![P](p1: P)(using TensorTree[P], FloatTree[P, V]): P = p1.map([T <: Tuple] => (n: Labels[T]) ?=> (a: Tensor[T, V]) => a *! p2)
def `//!`[P](p1: P)(using TensorTree[P], FloatTree[P, V]): P = p1.map([T <: Tuple] => (n: Labels[T]) ?=> (a: Tensor[T, V]) => a /! p2)

// Scalar broadcast extensions (Tensor0 op Tree)
extension [V: IsFloating](p2: Double)
def ++![P](p1: P)(using TensorTree[P], FloatTree[P, V]): P = Tensor0(VType[V])(p2) ++! p1
def --![P](p1: P)(using TensorTree[P], FloatTree[P, V]): P = Tensor0(VType[V])(p2) --! p1
def **![P](p1: P)(using TensorTree[P], FloatTree[P, V]): P = Tensor0(VType[V])(p2) **! p1
def `//!`[P](p1: P)(using TensorTree[P], FloatTree[P, V]): P = Tensor0(VType[V])(p2) `//!` p1

// Tree extensions (Tree op Tree, Tree op Scalar, and math ops)
// Excluded for bare Tensor[T, V] to avoid conflicts with tensor's own operators
extension [P, V](p1: P)(using tt: TensorTree[P], ft: FloatTree[P, V], isF: IsFloating[V], ev: NotGiven[IsFloatingTensor[P, V]])
Expand Down
67 changes: 34 additions & 33 deletions core/src/main/scala/dimwit/optimizer/GradientOptimizer.scala
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package dimwit.optimizer

import dimwit.*
import dimwit.Conversions.given
import dimwit.autodiff.FloatTree.*
import dimwit.autodiff.FloatTree.ops.*
import dimwit.autodiff.*
Expand All @@ -25,43 +26,43 @@ import dimwit.autodiff.*
* }}}
*/
trait GradientOptimizer:
type State[_]
type State[_, V]

// Core API
def init[Params: TensorTree: FloatTreeFor[Float32]](params: Params): State[Params]
def update[Params: TensorTree: FloatTreeFor[Float32]](gradients: Grad[Params], params: Params, state: State[Params]): (Params, State[Params])
def init[V, Params: TensorTree: FloatTreeFor[V]](params: Params)(using IsFloating[V]): State[Params, V]
def update[V, Params: TensorTree: FloatTreeFor[V]](gradients: Grad[Params], params: Params, state: State[Params, V])(using IsFloating[V]): (Params, State[Params, V])

// Convenience: iterator with fixed gradient function
def iterateWithState[Params: TensorTree: FloatTreeFor[Float32]](init: Params)(df: Params => Grad[Params]): Iterator[(Params, State[Params])] =
def iterateWithState[V, Params: TensorTree: FloatTreeFor[V]](init: Params)(df: Params => Grad[Params])(using IsFloating[V]): Iterator[(Params, State[Params, V])] =
Iterator.iterate((init, this.init(init))): (params, state) =>
val grads = df(params)
update(grads, params, state)

def iterate[Params: TensorTree: FloatTreeFor[Float32]](init: Params)(df: Params => Grad[Params]): Iterator[Params] =
def iterate[V, Params: TensorTree: FloatTreeFor[V]](init: Params)(df: Params => Grad[Params])(using IsFloating[V]): Iterator[Params] =
iterateWithState(init)(df).map(_._1)

case class GradientDescent(learningRate: Tensor0[Float32]) extends GradientOptimizer:
case class GradientDescent(learningRate: Double) extends GradientOptimizer:

type State[P] = Unit // Stateless optimizer
type State[P, V] = Unit // Stateless optimizer

def init[Params: TensorTree: FloatTreeFor[Float32]](params: Params): Unit = ()
def init[V, Params: TensorTree: FloatTreeFor[V]](params: Params)(using IsFloating[V]): Unit = ()

def update[Params: TensorTree: FloatTreeFor[Float32]](gradients: Grad[Params], params: Params, state: Unit): (Params, Unit) =
def update[V, Params: TensorTree: FloatTreeFor[V]](gradients: Grad[Params], params: Params, state: Unit)(using IsFloating[V]): (Params, Unit) =
val newParams = params -- gradients.value.scale(learningRate)
(newParams, ())

case class Lion(learningRate: Tensor0[Float32], weightDecay: Tensor0[Float32] = Tensor0(0.0f), beta1: Tensor0[Float32] = Tensor0(0.9f), beta2: Tensor0[Float32] = Tensor0(0.99f)) extends GradientOptimizer:
case class Lion(learningRate: Double, weightDecay: Double = 0.0f, beta1: Double = 0.9f, beta2: Double = 0.99f) extends GradientOptimizer:

type State[P] = P // momentum state has same structure as params
type State[P, V] = P // momentum state has same structure as params

def init[Params: TensorTree: FloatTreeFor[Float32]](params: Params): Params =
def init[V, Params: TensorTree: FloatTreeFor[V]](params: Params)(using IsFloating[V]): Params =
params.map([T <: Tuple] =>
(n: Labels[T]) ?=>
(t: Tensor[T, Float32]) =>
(t: Tensor[T, V]) =>
Tensor(t.shape).fill(0f)
)

def update[Params: TensorTree: FloatTreeFor[Float32]](gradients: Grad[Params], params: Params, momentums: Params): (Params, Params) =
def update[V, Params: TensorTree: FloatTreeFor[V]](gradients: Grad[Params], params: Params, momentums: Params)(using IsFloating[V]): (Params, Params) =
// the direction (1 or -1)
// is determined by the sign of the momentum + gradient
val updateDirection = (momentums **! beta1 ++ gradients.value **! (1f - beta1)).sign
Expand All @@ -71,38 +72,38 @@ case class Lion(learningRate: Tensor0[Float32], weightDecay: Tensor0[Float32] =

(updatedParams, newMomentums)

case class AdamState[P](
case class AdamState[P, V: IsFloating](
momentums: P, // momentums
velocities: P, // velocities
b1: Tensor0[Float32], // decay rate for momentums mᵗ
b2: Tensor0[Float32] // decay rate for velocities vᵗ
b1: Tensor0[V], // decay rate for momentums mᵗ
b2: Tensor0[V] // decay rate for velocities vᵗ
)

/** Implements the Adam optimization algorithm.
*
* @see [[https://arxiv.org/abs/1412.6980 Adam: A Method for Stochastic Optimization]]
*/
case class Adam(
learningRate: Tensor0[Float32], // step size (learning rate)
b1: Tensor0[Float32] = Tensor0(0.9f), // decay rate for momentums mᵗ
b2: Tensor0[Float32] = Tensor0(0.999f), // decay rate for velocities vᵗ
epsilon: Tensor0[Float32] = Tensor0(1e-8f) // small constant to prevent division by zero
learningRate: Double, // step size (learning rate)
b1: Double = 0.9, // decay rate for momentums mᵗ
b2: Double = 0.999, // decay rate for velocities vᵗ
epsilon: Double = 1e-8 // small constant to prevent division by zero
) extends GradientOptimizer:

private val β1 = b1
private val β2 = b2

type State[P] = AdamState[P]
type State[P, V] = AdamState[P, V]

def init[Params: TensorTree: FloatTreeFor[Float32]](params: Params): State[Params] =
def init[V, Params: TensorTree: FloatTreeFor[V]](params: Params)(using IsFloating[V]): State[Params, V] =
def zeros = params.fillCopy(0f)
AdamState(zeros, zeros, b1 = Tensor0(1f), b2 = Tensor0(1f))
AdamState[Params, V](zeros, zeros, b1 = Tensor0(VType[V])(1f), b2 = Tensor0(VType[V])(1f))

def update[Params: TensorTree: FloatTreeFor[Float32]](
def update[V, Params: TensorTree: FloatTreeFor[V]](
gradients: Grad[Params],
params: Params,
state: State[Params]
): (Params, State[Params]) =
state: State[Params, V]
)(using IsFloating[V]): (Params, State[Params, V]) =
// rename state variables to last time step for clarity
val `mₜ₋₁` = state.momentums
val `vₜ₋₁` = state.velocities
Expand Down Expand Up @@ -140,18 +141,18 @@ case class Adam(
*/
case class AdamW(
val adam: Adam,
val weightDecayFactor: Tensor0[Float32]
val weightDecayFactor: Double
) extends GradientOptimizer:

type State[P] = adam.State[P]
type State[P, V] = adam.State[P, V]

def init[Params: TensorTree: FloatTreeFor[Float32]](params: Params): State[Params] = adam.init(params)
def init[V, Params: TensorTree: FloatTreeFor[V]](params: Params)(using IsFloating[V]): State[Params, V] = adam.init(params)

def update[Params: TensorTree: FloatTreeFor[Float32]](
def update[V, Params: TensorTree: FloatTreeFor[V]](
gradients: Grad[Params],
params: Params,
state: State[Params]
): (Params, State[Params]) =
state: State[Params, V]
)(using IsFloating[V]): (Params, State[Params, V]) =
val α = adam.learningRate
val `θₜ₋₁` = params
val `λ'` = weightDecayFactor
Expand Down
2 changes: 1 addition & 1 deletion examples/src/main/scala/basic/LogisticRegression.scala
Original file line number Diff line number Diff line change
Expand Up @@ -125,7 +125,7 @@ object LogisticRegression:
val trainLoss = jit(BinaryLogisticRegression.loss(trainingData, trainLabels))
val valLoss = jit(BinaryLogisticRegression.loss(valData, valLabels))
val learningRate = 5e-1f
val gd = GradientDescent(Tensor0(learningRate))
val gd = GradientDescent(learningRate)

// Training loop
val numiterations = 1000
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -209,7 +209,7 @@ object VariationalAutoencoderExample:
losses.sum / batchSize.toFloat

val batches = trainImages.chunk(Axis[TrainSample], numSamples / batchSize)
val optimizer = GradientDescent(learningRate = Tensor0(learningRate))
val optimizer = GradientDescent(learningRate = learningRate)
def trainBatch(trainKey: Random.Key, batch: Tensor3[TrainSample, Height, Width, Float32], params: Params): Params =
val grads = grad(batchLoss(trainKey, batch))(params)
val (newParams, _) = optimizer.update(grads, params, ())
Expand Down
Loading