From 71f1af9fa64989f83704bb5e7e60ac86836dbb3a Mon Sep 17 00:00:00 2001 From: Benjamin Meyer Date: Mon, 7 Sep 2026 08:47:37 +0200 Subject: [PATCH] Add after method for drop(n).next(); easier for non-Scala folks --- README.md | 4 +-- .../main/scala/deepwit/training/package.scala | 12 +++------ .../deepwit/training/TapEverySuite.scala | 26 ++++++++++++------- .../autoencoder/AutoEncoderTrain.scala | 5 ++-- .../scala/deepwit/examples/gpt/GPTTrain.scala | 5 ++-- .../mnistClassification/MNistCNNTrain.scala | 5 ++-- .../neuralImage/NeuralImageTrain.scala | 5 ++-- .../examples/regression/Regression.scala | 4 +-- .../examples/thinning/MoonsMLPTrain.scala | 5 ++-- .../VariationalAutoencoderTrain.scala | 5 ++-- mdocs/README.md | 5 ++-- 11 files changed, 39 insertions(+), 42 deletions(-) diff --git a/README.md b/README.md index 8ee7e72..3915d85 100644 --- a/README.md +++ b/README.md @@ -114,7 +114,7 @@ Training the model reduces to a termination condition on this iterator; here aft A model checkpointer serializes the final train state object. ```scala -val finalState = trainTrajectory.drop(numIterations).next() +val finalState = trainTrajectory.after(numIterations) TensorTreeCheckpointer.newIn(checkpointRoot).save(finalState, numIterations) ``` @@ -151,7 +151,7 @@ The user code composes these core modules into custom architectures given the us | `deepwit.init` | Xavier/Glorot normal and uniform, for matrices and vectors | | `deepwit.regularization` | `Perturbation` — thinning (dropout) as a mutation of the weights that *read* a feature | | `deepwit.optimizer` | `LearningRateSchedule` (constant, linear warmup, cosine decay), `LearningRateScheduler`, `clipGlobalNorm` | -| `deepwit.training` | `Monitor` (step, loss, throughput, learning rate), `tapEvery` | +| `deepwit.training` | `Monitor` (step, loss, throughput, learning rate), `tapEvery`, `after` | | `deepwit.checkpointing` | `TensorTreeCheckpointer` — save and load any `TensorTree` by iteration | ## Relationship to DimWit diff --git a/core/src/main/scala/deepwit/training/package.scala b/core/src/main/scala/deepwit/training/package.scala index 7784d30..4ec95fd 100644 --- a/core/src/main/scala/deepwit/training/package.scala +++ b/core/src/main/scala/deepwit/training/package.scala @@ -10,11 +10,7 @@ extension [T](it: Iterator[T]) if id > 0 && id % n == 0 then f(t, id) .map(_._1) -extension [T](it: LazyList[T]) - - def tapEvery(n: Int)(f: (T, Int) => Unit): LazyList[T] = - it - .zipWithIndex - .tapEach: (t, id) => - if id > 0 && id % n == 0 then f(t, id) - .map(_._1) + /** The state after n iterations: Advances the iterator n steps and returns the resulting element */ + def after(n: Int): T = + require(n >= 0, s"A number of steps must not be negative, but was $n.") + it.drop(n).next() diff --git a/core/src/test/scala/deepwit/training/TapEverySuite.scala b/core/src/test/scala/deepwit/training/TapEverySuite.scala index 0d1c704..5236196 100644 --- a/core/src/test/scala/deepwit/training/TapEverySuite.scala +++ b/core/src/test/scala/deepwit/training/TapEverySuite.scala @@ -22,15 +22,21 @@ class TapEverySuite extends AnyFunSpec with Matchers: Iterator.from(0).tapEvery(1)((_, id) => seen += id).take(3).toList seen.toList shouldBe List(1, 2) - describe("LazyList.tapEvery"): + describe("Iterator.after"): - it("fires at every n-th index but not at zero"): - val seen = ListBuffer.empty[(String, Int)] - LazyList.from(0).map(i => s"e$i").tapEvery(3)((t, id) => seen += ((t, id))).take(10).toList - seen.toList shouldBe List(("e3", 3), ("e6", 6), ("e9", 9)) + it("counts from zero, so the first element is the state after no steps"): + Iterator.from(0).after(0) shouldBe 0 - it("stays lazy until the elements are forced"): - val seen = ListBuffer.empty[Int] - val tapped = LazyList.from(0).tapEvery(1)((_, id) => seen += id) - seen.toList shouldBe empty - tapped.take(3).toList shouldBe List(0, 1, 2) + it("returns the element that many steps in"): + Iterator.from(0).after(3) shouldBe 3 + + it("advances the iterator past what it returns"): + val trajectory = Iterator.from(0) + trajectory.after(3) shouldBe 3 + trajectory.next() shouldBe 4 + + it("throws when the iterator ends first"): + a[NoSuchElementException] should be thrownBy Iterator(0, 1).after(5) + + it("rejects a negative number of steps"): + an[IllegalArgumentException] should be thrownBy Iterator.from(0).after(-1) diff --git a/examples/src/main/scala/deepwit/examples/autoencoder/AutoEncoderTrain.scala b/examples/src/main/scala/deepwit/examples/autoencoder/AutoEncoderTrain.scala index 8bbc9b2..537ee11 100644 --- a/examples/src/main/scala/deepwit/examples/autoencoder/AutoEncoderTrain.scala +++ b/examples/src/main/scala/deepwit/examples/autoencoder/AutoEncoderTrain.scala @@ -6,7 +6,7 @@ import dimwit.Conversions.given import deepwit.examples.dataset.MNISTLoader import MNISTLoader.TestSample -import deepwit.training.{Monitor, tapEvery} +import deepwit.training.{Monitor, after, tapEvery} import deepwit.checkpointing.TensorTreeCheckpointer import deepwit.loss.BinaryCrossEntropy import dimwit.optimizer.{Adam, AdamState} @@ -83,7 +83,6 @@ def train(): Unit = case (state, step) => checkpointer.save(state, step) println(s"Checkpoint saved at epoch $step") - .drop(numIterations) - .next() + .after(numIterations) println(s"Done. Wrote ${checkpointer.rootPath}.") diff --git a/examples/src/main/scala/deepwit/examples/gpt/GPTTrain.scala b/examples/src/main/scala/deepwit/examples/gpt/GPTTrain.scala index 1190ef2..f30d069 100644 --- a/examples/src/main/scala/deepwit/examples/gpt/GPTTrain.scala +++ b/examples/src/main/scala/deepwit/examples/gpt/GPTTrain.scala @@ -4,7 +4,7 @@ import deepwit.loss.CategoricalCrossEntropy import dimwit.* import dimwit.Conversions.given -import deepwit.training.{Monitor, tapEvery} +import deepwit.training.{Monitor, after, tapEvery} import deepwit.optimizer.* import dimwit.optimizer.{AdamW, Adam, AdamState} import dimwit.TreeOf.ops.* @@ -186,5 +186,4 @@ import Config.* logger.save(state, step) println(s"Checkpoint saved") println("-" * 30) - .drop(1_000_000_000) - .next() + .after(1_000_000_000) diff --git a/examples/src/main/scala/deepwit/examples/mnistClassification/MNistCNNTrain.scala b/examples/src/main/scala/deepwit/examples/mnistClassification/MNistCNNTrain.scala index 13eebec..794ed4b 100644 --- a/examples/src/main/scala/deepwit/examples/mnistClassification/MNistCNNTrain.scala +++ b/examples/src/main/scala/deepwit/examples/mnistClassification/MNistCNNTrain.scala @@ -8,7 +8,7 @@ import deepwit.loss.CategoricalCrossEntropy import deepwit.examples.dataset.{MNISTLoader, MNISTBatchSample} import dimwit.optimizer.GradientDescentState -import deepwit.training.{Monitor, tapEvery} +import deepwit.training.{Monitor, after, tapEvery} import deepwit.checkpointing.TensorTreeCheckpointer case class TrainState( @@ -81,7 +81,6 @@ def train(): Unit = case (state, step) => checkpointer.save(state, step) println(s"Checkpoint saved at epoch $step") - .drop(numIterations) - .next() + .after(numIterations) println(s"Done. Wrote ${checkpointer.rootPath}.") diff --git a/examples/src/main/scala/deepwit/examples/neuralImage/NeuralImageTrain.scala b/examples/src/main/scala/deepwit/examples/neuralImage/NeuralImageTrain.scala index a8f4578..cf98f9d 100644 --- a/examples/src/main/scala/deepwit/examples/neuralImage/NeuralImageTrain.scala +++ b/examples/src/main/scala/deepwit/examples/neuralImage/NeuralImageTrain.scala @@ -4,7 +4,7 @@ import dimwit.* import dimwit.Conversions.given import dimwit.optimizer.{Adam, AdamState} -import deepwit.training.{Monitor, tapEvery} +import deepwit.training.{Monitor, after, tapEvery} import deepwit.checkpointing.TensorTreeCheckpointer import deepwit.loss.SquaredError @@ -103,8 +103,7 @@ def train(): Unit = val finalState = trainTrajectory .tapEvery(100): case (state, step) => println(trainMonitor.report(step, state)) - .drop(numIterations) - .next() + .after(numIterations) // -- Save final state -- diff --git a/examples/src/main/scala/deepwit/examples/regression/Regression.scala b/examples/src/main/scala/deepwit/examples/regression/Regression.scala index 5342ec0..57c991a 100644 --- a/examples/src/main/scala/deepwit/examples/regression/Regression.scala +++ b/examples/src/main/scala/deepwit/examples/regression/Regression.scala @@ -9,6 +9,7 @@ import io.circe.Json import plotwit.* import plotwit.PlotTargets.desktopBrowser +import deepwit.training.after import deepwit.activation.gelu import deepwit.base.{AffineFormLayer, AffineLayer} import deepwit.checkpointing.TensorTreeCheckpointer @@ -104,8 +105,7 @@ def train(): Unit = // -- Run train trajectory -- val finalState = trainTrajectory - .drop(numIterations) - .next() + .after(numIterations) // -- Save the fitted state -- diff --git a/examples/src/main/scala/deepwit/examples/thinning/MoonsMLPTrain.scala b/examples/src/main/scala/deepwit/examples/thinning/MoonsMLPTrain.scala index fbdaf11..15bc5fe 100644 --- a/examples/src/main/scala/deepwit/examples/thinning/MoonsMLPTrain.scala +++ b/examples/src/main/scala/deepwit/examples/thinning/MoonsMLPTrain.scala @@ -8,7 +8,7 @@ import dimwit.Conversions.given import dimwit.optimizer.{Adam, AdamState} import deepwit.loss.CategoricalCrossEntropy -import deepwit.training.{Monitor, tapEvery} +import deepwit.training.{Monitor, after, tapEvery} import deepwit.checkpointing.TensorTreeCheckpointer case class TrainState( @@ -90,8 +90,7 @@ def train(): Unit = case (state, step) => checkpointer.save(state, step) println(s"Checkpoint saved at step $step") - .drop(numIterations) - .next() + .after(numIterations) println(f"Final cost: ${finalState.lastCost.item}%.6f") println(s"Done. Wrote ${checkpointer.rootPath}.") diff --git a/examples/src/main/scala/deepwit/examples/variationalAutoencoder/VariationalAutoencoderTrain.scala b/examples/src/main/scala/deepwit/examples/variationalAutoencoder/VariationalAutoencoderTrain.scala index 11f41a0..8b10c36 100644 --- a/examples/src/main/scala/deepwit/examples/variationalAutoencoder/VariationalAutoencoderTrain.scala +++ b/examples/src/main/scala/deepwit/examples/variationalAutoencoder/VariationalAutoencoderTrain.scala @@ -8,7 +8,7 @@ import deepwit.examples.dataset.MNISTLoader import deepwit.checkpointing.TensorTreeCheckpointer import deepwit.loss.BinaryCrossEntropy -import deepwit.training.{Monitor, tapEvery} +import deepwit.training.{Monitor, after, tapEvery} case class TrainState( params: VariationalAutoencoder.Params, @@ -101,8 +101,7 @@ def train(): Unit = case (state, step) => checkpointer.save(state, step) println(s"Checkpoint saved at step $step") - .drop(numIterations) - .next() + .after(numIterations) println(f"Final cost: ${finalState.lastCost.item}%.6f") println(s"Done. Wrote ${checkpointer.rootPath}.") diff --git a/mdocs/README.md b/mdocs/README.md index 3825016..fd5342c 100644 --- a/mdocs/README.md +++ b/mdocs/README.md @@ -7,6 +7,7 @@ import deepwit.activation.gelu import deepwit.base.{AffineFormLayer, AffineLayer} import deepwit.checkpointing.TensorTreeCheckpointer import deepwit.loss.SquaredError +import deepwit.training.after dimwit.initialize() @@ -142,7 +143,7 @@ Training the model reduces to a termination condition on this iterator; here aft A model checkpointer serializes the final train state object. ```scala mdoc:compile-only -val finalState = trainTrajectory.drop(numIterations).next() +val finalState = trainTrajectory.after(numIterations) TensorTreeCheckpointer.newIn(checkpointRoot).save(finalState, numIterations) ``` @@ -179,7 +180,7 @@ The user code composes these core modules into custom architectures given the us | `deepwit.init` | Xavier/Glorot normal and uniform, for matrices and vectors | | `deepwit.regularization` | `Perturbation` — thinning (dropout) as a mutation of the weights that *read* a feature | | `deepwit.optimizer` | `LearningRateSchedule` (constant, linear warmup, cosine decay), `LearningRateScheduler`, `clipGlobalNorm` | -| `deepwit.training` | `Monitor` (step, loss, throughput, learning rate), `tapEvery` | +| `deepwit.training` | `Monitor` (step, loss, throughput, learning rate), `tapEvery`, `after` | | `deepwit.checkpointing` | `TensorTreeCheckpointer` — save and load any `TensorTree` by iteration | ## Relationship to DimWit