diff --git a/AGENTS.md b/AGENTS.md index dd3dcd45..8121d855 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -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) @@ -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) @@ -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) diff --git a/core/src/main/scala/dimwit/autodiff/FloatTree.scala b/core/src/main/scala/dimwit/autodiff/FloatTree.scala index bf0b1cba..b51ee5f0 100644 --- a/core/src/main/scala/dimwit/autodiff/FloatTree.scala +++ b/core/src/main/scala/dimwit/autodiff/FloatTree.scala @@ -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]]) diff --git a/core/src/main/scala/dimwit/optimizer/GradientOptimizer.scala b/core/src/main/scala/dimwit/optimizer/GradientOptimizer.scala index 17a2d535..d02672bc 100644 --- a/core/src/main/scala/dimwit/optimizer/GradientOptimizer.scala +++ b/core/src/main/scala/dimwit/optimizer/GradientOptimizer.scala @@ -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.* @@ -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 @@ -71,11 +72,11 @@ 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. @@ -83,26 +84,26 @@ case class AdamState[P]( * @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 @@ -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 diff --git a/core/src/main/scala/dimwit/tensor/tensorops/StructuralOps.scala b/core/src/main/scala/dimwit/tensor/tensorops/StructuralOps.scala index afc7f0a1..a88d3054 100644 --- a/core/src/main/scala/dimwit/tensor/tensorops/StructuralOps.scala +++ b/core/src/main/scala/dimwit/tensor/tensorops/StructuralOps.scala @@ -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. diff --git a/core/src/test/scala/dimwit/optimizer/GradientOptimizerSuite.scala b/core/src/test/scala/dimwit/optimizer/GradientOptimizerSuite.scala new file mode 100644 index 00000000..931a2e84 --- /dev/null +++ b/core/src/test/scala/dimwit/optimizer/GradientOptimizerSuite.scala @@ -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 diff --git a/core/src/test/scala/dimwit/tensor/TensorOpsStructureSuite.scala b/core/src/test/scala/dimwit/tensor/TensorOpsStructureSuite.scala index 8bd7dd5a..338ac281 100644 --- a/core/src/test/scala/dimwit/tensor/TensorOpsStructureSuite.scala +++ b/core/src/test/scala/dimwit/tensor/TensorOpsStructureSuite.scala @@ -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)) diff --git a/docs/quickstart.md b/docs/quickstart.md index c75886c4..28569c8e 100644 --- a/docs/quickstart.md +++ b/docs/quickstart.md @@ -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) ``` diff --git a/examples/src/main/scala/basic/LogisticRegression.scala b/examples/src/main/scala/basic/LogisticRegression.scala index 1b35cf1c..aabda6d6 100644 --- a/examples/src/main/scala/basic/LogisticRegression.scala +++ b/examples/src/main/scala/basic/LogisticRegression.scala @@ -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 diff --git a/examples/src/main/scala/complex/VariationalAutoencoder.scala b/examples/src/main/scala/complex/VariationalAutoencoder.scala index a075f996..8fba39ef 100644 --- a/examples/src/main/scala/complex/VariationalAutoencoder.scala +++ b/examples/src/main/scala/complex/VariationalAutoencoder.scala @@ -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, ()) diff --git a/mdocs/AGENTS.md b/mdocs/AGENTS.md index 0cff8cf7..2a903bcc 100644 --- a/mdocs/AGENTS.md +++ b/mdocs/AGENTS.md @@ -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) @@ -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) @@ -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) diff --git a/mdocs/docs/quickstart.md b/mdocs/docs/quickstart.md index 231768e8..843046a1 100644 --- a/mdocs/docs/quickstart.md +++ b/mdocs/docs/quickstart.md @@ -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) ```