Skip to content
Open
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
6 changes: 3 additions & 3 deletions AGENTS.md
Original file line number Diff line number Diff line change
Expand Up @@ -902,7 +902,7 @@ val lossFunc = mse(trainData, trainLabels)
val gradFunc = Autodiff.grad(lossFunc)

// Create optimizer
val optimizer = GradientDescent(learningRate = Tensor0(0.01f))
val optimizer = GradientDescent(learningRate = 0.01f)

// Training loop with iterator
val trained = optimizer.iterate(initModelParams)(gradFunc)
Expand All @@ -919,7 +919,7 @@ val trained = optimizer.iterate(initModelParams)(gradFunc)
import dimwit.optimizer.Lion

// Lion optimizer with momentum
val lionOptimizer = Lion(learningRate = Tensor0(1e-3f), beta1 = Tensor0(0.9f), beta2 = Tensor0(0.99f), weightDecay = Tensor0(0.0f))
val lionOptimizer = Lion(learningRate = 1e-3f, beta1 = 0.9f, beta2 = 0.99f, weightDecay = 0.0f)

// Training with Lion
val trainedLion = lionOptimizer.iterate(initModelParams)(gradFunc)
Expand Down Expand Up @@ -962,7 +962,7 @@ val initRegressionParams = RegressionParams(initSlope, initIntercept)

// Train
val regressionGrad = Autodiff.grad(regressionLoss(xData, yData))
val gdOptimizer = GradientDescent(learningRate = Tensor0(0.1f))
val gdOptimizer = GradientDescent(learningRate = 0.1f)

val finalParams = gdOptimizer.iterate(initRegressionParams)(regressionGrad)
.take(100)
Expand Down
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
Original file line number Diff line number Diff line change
Expand Up @@ -553,9 +553,9 @@ object StructuralOps:
ev: AxisRemover[T, L],
labelR: Labels[ev.RemainingAxes]
): Seq[Tensor[ev.RemainingAxes, V]] =
val axisIdx = ev.index
val unstacked = Jax.jnp.split(tensor.jaxValue, tensor.shape.dimensions(axisIdx), axis = axisIdx).as[Seq[Jax.PyDynamic]]
unstacked.map(x => Tensor[ev.RemainingAxes, V](x))
(0 until tensor.shape.dimensions(ev.index)).map: i =>
val slicedJax = Jax.jnp.take(tensor.jaxValue, Jax.jnp.array(i), axis = ev.index)
Tensor[ev.RemainingAxes, V](slicedJax)

/** splits the tensor into chunks of the specified size along the given axis
* returning a sequence of tensors corresponding to the chunks.
Expand Down
69 changes: 69 additions & 0 deletions core/src/test/scala/dimwit/optimizer/GradientOptimizerSuite.scala
Original file line number Diff line number Diff line change
@@ -0,0 +1,69 @@
package dimwit.optimizer

import dimwit.*
import dimwit.Conversions.given
import dimwit.autodiff.FloatTree.*
import dimwit.autodiff.FloatTree.ops.*
import dimwit.autodiff.*

class GradientOptimizerSuite extends DimwitTest:

describe("GradientDescent"):
it("should converge towards the minimum of f(x) = (x+1)^2 at x = -1"):
val optimizer = GradientDescent(learningRate = 0.1)
val minX = optimizer.iterate(Tensor0(2.0f))(x => Grad(2 * (x + 1))).drop(1000).next()
minX.item shouldBe -1.0f +- 0.1f

describe("Adam"):
it("should converge towards the minimum of f(x) = (x+1)^2 at x = -1"):
val optimizer = Adam(learningRate = 0.1)
val minX = optimizer.iterate(Tensor0(2.0f))(x => Grad(2 * (x + 1))).drop(1000).next()
minX.item shouldBe -1.0f +- 0.1f

it("should compute the exact momentum and velocity updates (single step)"):
val optimizer = Adam(learningRate = 0.1, b1 = 0.9, b2 = 0.999)
val initParams = Tensor0(2.0f)
val initState = optimizer.init(initParams)

val grad = Grad(Tensor0(6.0f))
val (nextParams, nextState) = optimizer.update(grad, initParams, initState)

nextParams.item shouldBe 1.9f +- 1e-5f
nextState.momentums.item shouldBe 0.6f +- 1e-5f
nextState.velocities.item shouldBe 0.036f +- 1e-5f
nextState.b1.item shouldBe 0.9f +- 1e-5f
nextState.b2.item shouldBe 0.999f +- 1e-5f

describe("AdamW"):
it("should converge towards the minimum of f(x) = (x+1)^2 at x = -1"):
val optimizer = AdamW(Adam(learningRate = 0.1), weightDecayFactor = 0.1)
val minX = optimizer.iterate(Tensor0(2.0f))(x => Grad(2 * (x + 1))).drop(1000).next()
minX.item shouldBe -1.0f +- 0.1f

it("should apply decoupled weight decay (single step)"):
val adam = Adam(learningRate = 0.1)
val adamW = AdamW(adam, weightDecayFactor = 0.1)

val initParams = Tensor0(2.0)
val grad = Grad(Tensor0(6.0))
val (adamParams, _) = adam.update(grad, initParams, adam.init(initParams))
val (adamWParams, _) = adamW.update(grad, initParams, adamW.init(initParams))

adamWParams.item shouldBe (adamParams.item - 0.02) +- 1e-5

describe("Lion"):
it("should converge towards the minimum of f(x) = (x+1)^2"):
val optimizer = Lion(learningRate = 0.1)
val minX = optimizer.iterate(Tensor0(2.0f))(x => Grad(2 * (x + 1))).drop(1000).next()
minX.item shouldBe -1.0f +- 0.1f

it("should compute the exact sign-based update and momentum (single step)"):
val optimizer = Lion(learningRate = 0.1, beta1 = 0.9, beta2 = 0.99)
val initParams = Tensor0(2.0)
val initMomentum = optimizer.init(initParams)

val grad = Grad(Tensor0(6.0))
val (nextParams, nextMomentum) = optimizer.update(grad, initParams, initMomentum)

nextParams.item shouldBe 1.9d +- 1e-5d
nextMomentum.item shouldBe 0.06d +- 1e-5d
25 changes: 25 additions & 0 deletions core/src/test/scala/dimwit/tensor/TensorOpsStructureSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -485,6 +485,31 @@ class TensorOpsStructureSuite extends DimwitTest:
unstacked(1) should approxEqual(Tensor1(Axis[B]).fromArray(Array(3.0f, 4.0f)))
unstacked(2) should approxEqual(Tensor1(Axis[B]).fromArray(Array(5.0f, 6.0f)))

it("should correctly unstack a 3D tensor along the first axis"):
val data = Tensor3(Axis[A], Axis[B], Axis[C]).fromArray(
Array(
Array(
Array(1.0f, 2.0f),
Array(3.0f, 4.0f)
),
Array(
Array(2.0f, 5.0f),
Array(0.0f, 3.0f)
),
Array(
Array(99.0f, 13.0f),
Array(0.0f, 22.0f)
)
)
)

data.shape.dimensions shouldBe List(3, 2, 2)

val unstacked = data.unstack(Axis[A])
unstacked.length shouldBe 3
unstacked.foreach: slice =>
slice.shape.dimensions shouldBe List(2, 2)

it("unstack ∘ stack is identity"):
val t1 = Tensor1(Axis[B]).fromArray(Array(1.0f, 2.0f))
val t2 = Tensor1(Axis[B]).fromArray(Array(3.0f, 4.0f))
Expand Down
2 changes: 1 addition & 1 deletion docs/quickstart.md
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,7 @@ def fit(x: Tensor2[Batch, Feature, Float32], y: Tensor1[Batch, Float32]): Iterat
val gradFn = grad(loss(x, y))

// gradient based optimization
val gd = GradientDescent(learningRate = Tensor0(0.1f)) // this is wrong, should be 0.1f not Tensor0
val gd = GradientDescent(learningRate = 0.1f) // this is wrong, should be 0.1f not Tensor0
gd.iterate(p0)(gradFn)
```

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
6 changes: 3 additions & 3 deletions mdocs/AGENTS.md
Original file line number Diff line number Diff line change
Expand Up @@ -709,7 +709,7 @@ val lossFunc = mse(trainData, trainLabels)
val gradFunc = Autodiff.grad(lossFunc)

// Create optimizer
val optimizer = GradientDescent(learningRate = Tensor0(0.01f))
val optimizer = GradientDescent(learningRate = 0.01f)

// Training loop with iterator
val trained = optimizer.iterate(initModelParams)(gradFunc)
Expand All @@ -726,7 +726,7 @@ val trained = optimizer.iterate(initModelParams)(gradFunc)
import dimwit.optimizer.Lion

// Lion optimizer with momentum
val lionOptimizer = Lion(learningRate = Tensor0(1e-3f), beta1 = Tensor0(0.9f), beta2 = Tensor0(0.99f), weightDecay = Tensor0(0.0f))
val lionOptimizer = Lion(learningRate = 1e-3f, beta1 = 0.9f, beta2 = 0.99f, weightDecay = 0.0f)

// Training with Lion
val trainedLion = lionOptimizer.iterate(initModelParams)(gradFunc)
Expand Down Expand Up @@ -769,7 +769,7 @@ val initRegressionParams = RegressionParams(initSlope, initIntercept)

// Train
val regressionGrad = Autodiff.grad(regressionLoss(xData, yData))
val gdOptimizer = GradientDescent(learningRate = Tensor0(0.1f))
val gdOptimizer = GradientDescent(learningRate = 0.1f)

val finalParams = gdOptimizer.iterate(initRegressionParams)(regressionGrad)
.take(100)
Expand Down
2 changes: 1 addition & 1 deletion mdocs/docs/quickstart.md
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,7 @@ def fit(x: Tensor2[Batch, Feature, Float32], y: Tensor1[Batch, Float32]): Iterat
val gradFn = grad(loss(x, y))

// gradient based optimization
val gd = GradientDescent(learningRate = Tensor0(0.1f)) // this is wrong, should be 0.1f not Tensor0
val gd = GradientDescent(learningRate = 0.1f) // this is wrong, should be 0.1f not Tensor0
gd.iterate(p0)(gradFn)
```

Expand Down
Loading