diff --git a/core/src/main/scala/dimwit/autodiff/FloatTree.scala b/core/src/main/scala/dimwit/autodiff/FloatTree.scala index bf0b1cb..b51ee5f 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 17a2d53..d02672b 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/examples/src/main/scala/basic/LogisticRegression.scala b/examples/src/main/scala/basic/LogisticRegression.scala index 1b35cf1..aabda6d 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 a075f99..8fba39e 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, ())