From 1c928693deed5d0a8d2c81e3179df229e7575fe1 Mon Sep 17 00:00:00 2001 From: Marcel Luethi Date: Sat, 1 Aug 2026 16:51:02 +0200 Subject: [PATCH 1/4] move tensor tree related methods to separate package --- core/src/main/scala/dimwit/MemoryHelper.scala | 2 +- core/src/main/scala/dimwit/autodiff/Autodiff.scala | 1 + core/src/main/scala/dimwit/jax/Jit.scala | 2 +- .../src/main/scala/dimwit/optimizer/GradientOptimizer.scala | 5 +++-- core/src/main/scala/dimwit/package.scala | 4 +++- core/src/main/scala/dimwit/python/PyBridge.scala | 2 +- core/src/main/scala/dimwit/random/Random.scala | 2 +- .../scala/dimwit/{autodiff => tensortree}/FloatTree.scala | 2 +- .../scala/dimwit/{autodiff => tensortree}/TensorTree.scala | 2 +- core/src/test/scala/dimwit/python/PyWrapSuite.scala | 4 ++-- .../{autodiff => tensortree}/FloatTensorTreeSuite.scala | 6 +++--- .../scala/dimwit/{autodiff => tensortree}/PyTreeSuite.scala | 2 +- .../dimwit/{autodiff => tensortree}/TensorTreeSuite.scala | 2 +- .../src/main/scala/complex/VariationalAutoencoder.scala | 2 +- 14 files changed, 21 insertions(+), 17 deletions(-) rename core/src/main/scala/dimwit/{autodiff => tensortree}/FloatTree.scala (99%) rename core/src/main/scala/dimwit/{autodiff => tensortree}/TensorTree.scala (99%) rename core/src/test/scala/dimwit/{autodiff => tensortree}/FloatTensorTreeSuite.scala (98%) rename core/src/test/scala/dimwit/{autodiff => tensortree}/PyTreeSuite.scala (99%) rename core/src/test/scala/dimwit/{autodiff => tensortree}/TensorTreeSuite.scala (99%) diff --git a/core/src/main/scala/dimwit/MemoryHelper.scala b/core/src/main/scala/dimwit/MemoryHelper.scala index ac7604a9..a5a67095 100644 --- a/core/src/main/scala/dimwit/MemoryHelper.scala +++ b/core/src/main/scala/dimwit/MemoryHelper.scala @@ -1,6 +1,6 @@ package dimwit -import dimwit.autodiff.TensorTree +import dimwit.tensortree.TensorTree import me.shadaj.scalapy.py private[dimwit] object MemoryHelper: diff --git a/core/src/main/scala/dimwit/autodiff/Autodiff.scala b/core/src/main/scala/dimwit/autodiff/Autodiff.scala index 03f8ee15..613e471c 100644 --- a/core/src/main/scala/dimwit/autodiff/Autodiff.scala +++ b/core/src/main/scala/dimwit/autodiff/Autodiff.scala @@ -6,6 +6,7 @@ import dimwit.tensor.Tensor import dimwit.tensor.Tensor0 import dimwit.tensor.TensorOps.IsFloating import dimwit.tensor.TupleHelpers.PrimeConcatType +import dimwit.tensortree.TensorTree import me.shadaj.scalapy.py object Autodiff: diff --git a/core/src/main/scala/dimwit/jax/Jit.scala b/core/src/main/scala/dimwit/jax/Jit.scala index c10c4293..5f39dbc0 100644 --- a/core/src/main/scala/dimwit/jax/Jit.scala +++ b/core/src/main/scala/dimwit/jax/Jit.scala @@ -1,7 +1,7 @@ package dimwit.jax import dimwit.OnError -import dimwit.autodiff.TensorTree +import dimwit.tensortree.TensorTree import dimwit.jax.Jax import dimwit.jax.Jax.PyDynamic import me.shadaj.scalapy.py diff --git a/core/src/main/scala/dimwit/optimizer/GradientOptimizer.scala b/core/src/main/scala/dimwit/optimizer/GradientOptimizer.scala index 17a2d535..c1cfdba5 100644 --- a/core/src/main/scala/dimwit/optimizer/GradientOptimizer.scala +++ b/core/src/main/scala/dimwit/optimizer/GradientOptimizer.scala @@ -1,8 +1,9 @@ package dimwit.optimizer import dimwit.* -import dimwit.autodiff.FloatTree.* -import dimwit.autodiff.FloatTree.ops.* +import dimwit.tensortree.* +import dimwit.tensortree.FloatTree.* +import dimwit.tensortree.FloatTree.ops.* import dimwit.autodiff.* /** Gradient optimizer interface with functional state management. diff --git a/core/src/main/scala/dimwit/package.scala b/core/src/main/scala/dimwit/package.scala index 5242a5a1..56a807ff 100644 --- a/core/src/main/scala/dimwit/package.scala +++ b/core/src/main/scala/dimwit/package.scala @@ -84,7 +84,9 @@ package object dimwit: // Export devices export dimwit.hardware.Device // Export automatic differentiation - export dimwit.autodiff.{Autodiff, TensorTree, FloatTree, Grad} + export dimwit.autodiff.{Autodiff, Grad} + // Export tensor trees + export dimwit.tensortree.{TensorTree, TensorTreeIO, TensorTreeFormat, FloatTree} // Export Just-in-Time compilation export dimwit.jax.Jit.{jit, jitDonating, jitDonatingUnsafe} export dimwit.jax.EagerCleanup.eagerCleanup diff --git a/core/src/main/scala/dimwit/python/PyBridge.scala b/core/src/main/scala/dimwit/python/PyBridge.scala index d3da1a7a..e8edb203 100644 --- a/core/src/main/scala/dimwit/python/PyBridge.scala +++ b/core/src/main/scala/dimwit/python/PyBridge.scala @@ -1,7 +1,7 @@ package dimwit.python import dimwit.OnError -import dimwit.autodiff.TensorTree +import dimwit.tensortree.TensorTree import dimwit.jax.Jax import dimwit.tensor.* import me.shadaj.scalapy.py diff --git a/core/src/main/scala/dimwit/random/Random.scala b/core/src/main/scala/dimwit/random/Random.scala index d6acbf0a..3e93222b 100644 --- a/core/src/main/scala/dimwit/random/Random.scala +++ b/core/src/main/scala/dimwit/random/Random.scala @@ -1,6 +1,6 @@ package dimwit.random -import dimwit.autodiff.TensorTree +import dimwit.tensortree.TensorTree import dimwit.jax.Jax import dimwit.python.PyBridge.liftPyTensor import dimwit.tensor.DType.Int32 diff --git a/core/src/main/scala/dimwit/autodiff/FloatTree.scala b/core/src/main/scala/dimwit/tensortree/FloatTree.scala similarity index 99% rename from core/src/main/scala/dimwit/autodiff/FloatTree.scala rename to core/src/main/scala/dimwit/tensortree/FloatTree.scala index bf0b1cba..bdf131ae 100644 --- a/core/src/main/scala/dimwit/autodiff/FloatTree.scala +++ b/core/src/main/scala/dimwit/tensortree/FloatTree.scala @@ -1,4 +1,4 @@ -package dimwit.autodiff +package dimwit.tensortree import dimwit.tensor.TensorOps.* import dimwit.tensor.* diff --git a/core/src/main/scala/dimwit/autodiff/TensorTree.scala b/core/src/main/scala/dimwit/tensortree/TensorTree.scala similarity index 99% rename from core/src/main/scala/dimwit/autodiff/TensorTree.scala rename to core/src/main/scala/dimwit/tensortree/TensorTree.scala index ec7eb0c4..1e742bbe 100644 --- a/core/src/main/scala/dimwit/autodiff/TensorTree.scala +++ b/core/src/main/scala/dimwit/tensortree/TensorTree.scala @@ -1,4 +1,4 @@ -package dimwit.autodiff +package dimwit.tensortree import dimwit.jax.Jax import dimwit.tensor.* diff --git a/core/src/test/scala/dimwit/python/PyWrapSuite.scala b/core/src/test/scala/dimwit/python/PyWrapSuite.scala index 432584a8..5ab4b191 100644 --- a/core/src/test/scala/dimwit/python/PyWrapSuite.scala +++ b/core/src/test/scala/dimwit/python/PyWrapSuite.scala @@ -77,7 +77,7 @@ class PyWrapSuite extends DimwitTest: // Call through a Python lambda that just forwards to pyFn val caller = py.eval("lambda fn, x: fn(x)") val input = Tensor1(Axis[A]).fromArray(Array(2f, 3f, 4f)) - val pyResult = caller(pyFn, dimwit.autodiff.TensorTree[Tensor1[A, Float32]].toPyTree(input)) - val result = dimwit.autodiff.TensorTree[Tensor1[A, Float32]].fromPyTree(pyResult) + val pyResult = caller(pyFn, dimwit.tensortree.TensorTree[Tensor1[A, Float32]].toPyTree(input)) + val result = dimwit.tensortree.TensorTree[Tensor1[A, Float32]].fromPyTree(pyResult) result should approxEqual(Tensor1(Axis[A]).fromArray(Array(6f, 9f, 12f))) diff --git a/core/src/test/scala/dimwit/autodiff/FloatTensorTreeSuite.scala b/core/src/test/scala/dimwit/tensortree/FloatTensorTreeSuite.scala similarity index 98% rename from core/src/test/scala/dimwit/autodiff/FloatTensorTreeSuite.scala rename to core/src/test/scala/dimwit/tensortree/FloatTensorTreeSuite.scala index ad4fd8f5..418baa48 100644 --- a/core/src/test/scala/dimwit/autodiff/FloatTensorTreeSuite.scala +++ b/core/src/test/scala/dimwit/tensortree/FloatTensorTreeSuite.scala @@ -1,9 +1,9 @@ -package dimwit.autodiff +package dimwit.tensortree import dimwit.* import dimwit.Conversions.given -import dimwit.autodiff.FloatTree.* -import dimwit.autodiff.FloatTree.ops.* +import dimwit.tensortree.FloatTree.* +import dimwit.tensortree.FloatTree.ops.* class FloatTensorTreeSuite extends DimwitTest: diff --git a/core/src/test/scala/dimwit/autodiff/PyTreeSuite.scala b/core/src/test/scala/dimwit/tensortree/PyTreeSuite.scala similarity index 99% rename from core/src/test/scala/dimwit/autodiff/PyTreeSuite.scala rename to core/src/test/scala/dimwit/tensortree/PyTreeSuite.scala index b8c67259..d1407a3c 100644 --- a/core/src/test/scala/dimwit/autodiff/PyTreeSuite.scala +++ b/core/src/test/scala/dimwit/tensortree/PyTreeSuite.scala @@ -1,4 +1,4 @@ -package dimwit.autodiff +package dimwit.tensortree import dimwit.* import dimwit.jax.Jax diff --git a/core/src/test/scala/dimwit/autodiff/TensorTreeSuite.scala b/core/src/test/scala/dimwit/tensortree/TensorTreeSuite.scala similarity index 99% rename from core/src/test/scala/dimwit/autodiff/TensorTreeSuite.scala rename to core/src/test/scala/dimwit/tensortree/TensorTreeSuite.scala index 7e8587cf..d04347f8 100644 --- a/core/src/test/scala/dimwit/autodiff/TensorTreeSuite.scala +++ b/core/src/test/scala/dimwit/tensortree/TensorTreeSuite.scala @@ -1,4 +1,4 @@ -package dimwit.autodiff +package dimwit.tensortree import dimwit.* diff --git a/examples/src/main/scala/complex/VariationalAutoencoder.scala b/examples/src/main/scala/complex/VariationalAutoencoder.scala index a075f996..3bc5e224 100644 --- a/examples/src/main/scala/complex/VariationalAutoencoder.scala +++ b/examples/src/main/scala/complex/VariationalAutoencoder.scala @@ -2,7 +2,7 @@ package examples.complex.vae import dimwit.Conversions.given import dimwit.* -import dimwit.autodiff.FloatTree.* +import dimwit.tensortree.FloatTree.* import dimwit.autodiff.* import dimwit.nn.ActivationFunctions.relu import dimwit.nn.ActivationFunctions.sigmoid From 5bdaaf829793140d5233d51a275a511154a9cef5 Mon Sep 17 00:00:00 2001 From: Marcel Luethi Date: Sat, 1 Aug 2026 16:51:47 +0200 Subject: [PATCH 2/4] add tensor tree IO Introduced mechanism to serialize a tensor tree. The default implementation is pkl. --- .../dimwit/tensortree/TensorTreeFormat.scala | 40 +++++++++++++++++ .../dimwit/tensortree/TensorTreeIO.scala | 27 ++++++++++++ .../dimwit/tensortree/TensorTreeIOSuite.scala | 44 +++++++++++++++++++ 3 files changed, 111 insertions(+) create mode 100644 core/src/main/scala/dimwit/tensortree/TensorTreeFormat.scala create mode 100644 core/src/main/scala/dimwit/tensortree/TensorTreeIO.scala create mode 100644 core/src/test/scala/dimwit/tensortree/TensorTreeIOSuite.scala diff --git a/core/src/main/scala/dimwit/tensortree/TensorTreeFormat.scala b/core/src/main/scala/dimwit/tensortree/TensorTreeFormat.scala new file mode 100644 index 00000000..62ceee91 --- /dev/null +++ b/core/src/main/scala/dimwit/tensortree/TensorTreeFormat.scala @@ -0,0 +1,40 @@ +package dimwit.tensortree + +import dimwit.jax.Jax +import me.shadaj.scalapy.py + +import java.nio.file.Path + +/** Provides an interface for reading and writing tensor trees to and from disk in various formats. + */ +trait TensorTreeFormat: + def write[P](p: P, path: Path)(using tt: TensorTree[P]): Unit + def read[P](path: Path)(using tt: TensorTree[P]): P + +object TensorTreeFormat: + + /** Pickle format for saving and loading tensor trees. + * + * Arrays are moved off-device to numpy before pickling, and re-materialized + * as JAX arrays on load, so files remain portable across host/GPU/TPU. + * `jax.tree_util.tree_map` handles the tuple/list/None nesting natively, so + * no recursive walk is needed here. + */ + object Pickle extends TensorTreeFormat: + private lazy val pickle = py.module("pickle") + private lazy val builtins = py.module("builtins") + + def write[P](p: P, path: Path)(using tt: TensorTree[P]): Unit = + val toHost = (x: Jax.PyDynamic) => Jax.np.asarray(Jax.jax.device_get(x)) + val numpyTree = Jax.jax.tree_util.tree_map(toHost, tt.toPyTree(p)) + val file = builtins.open(path.toAbsolutePath().toString(), "wb").as[py.Dynamic] + try pickle.dump(numpyTree, file) + finally file.close() + + def read[P](path: Path)(using tt: TensorTree[P]): P = + val file = builtins.open(path.toAbsolutePath().toString(), "rb").as[py.Dynamic] + val numpyTree = + try pickle.load(file).as[py.Dynamic] + finally file.close() + val toDevice = (x: Jax.PyDynamic) => Jax.jnp.asarray(x) + tt.fromPyTree(Jax.jax.tree_util.tree_map(toDevice, numpyTree)) diff --git a/core/src/main/scala/dimwit/tensortree/TensorTreeIO.scala b/core/src/main/scala/dimwit/tensortree/TensorTreeIO.scala new file mode 100644 index 00000000..54eede30 --- /dev/null +++ b/core/src/main/scala/dimwit/tensortree/TensorTreeIO.scala @@ -0,0 +1,27 @@ +package dimwit.tensortree + +import java.nio.file.Path + +/** Provides methods to save and load tensor trees to and from disk. + * The default format is pickle, but other formats can be specified if needed. + */ +object TensorTreeIO: + + /** Saves a tensor tree to disk in the specified format. + * + * @param p The tensor tree structure (e.g., model parameters) to be saved. + * @param path The file path where the tensor tree will be saved. + * @param format The format in which to save the tensor tree (default is pickle). + */ + def save[P](p: P, path: Path, format: TensorTreeFormat = TensorTreeFormat.Pickle)(using tt: TensorTree[P]): Unit = + format.write(p, path) + + /** Loads a tensor tree from disk in the specified format. + * + * @tparam P The type of the tensor tree structure to be loaded. + * @param path The file path from which to load the tensor tree. + * @param format The format in which to load the tensor tree (default is pickle). + * @return The loaded tensor tree structure. + */ + def load[P](path: Path, format: TensorTreeFormat = TensorTreeFormat.Pickle)(using tt: TensorTree[P]): P = + format.read(path) diff --git a/core/src/test/scala/dimwit/tensortree/TensorTreeIOSuite.scala b/core/src/test/scala/dimwit/tensortree/TensorTreeIOSuite.scala new file mode 100644 index 00000000..5a8da924 --- /dev/null +++ b/core/src/test/scala/dimwit/tensortree/TensorTreeIOSuite.scala @@ -0,0 +1,44 @@ +package dimwit.tensortree + +import dimwit.* + +import java.nio.file.Files + +class TensorTreeIOSuite extends DimwitTest: + + describe("save and load"): + it("round-trips a 1-level case class"): + case class Params(w1: Tensor1[A, Float32], b1: Tensor0[Int32]) + val params = Params( + Tensor1(Axis[A]).fromArray(Array(0.1f, 0.2f, 0.3f)), + Tensor0(42) + ) + val path = Files.createTempFile("tensortree-io", ".pkl") + try + TensorTreeIO.save(params, path) + val restored = TensorTreeIO.load[Params](path) + restored.w1 should approxEqual(params.w1) + restored.b1 should equal(params.b1) + finally Files.deleteIfExists(path) + + it("round-trips nested structures (lists and tuples)"): + case class Model(layers: List[Tensor0[Float32]], extra: (Tensor0[Int32], Tensor0[Int32])) + val params = Model( + List(Tensor0(1.0f), Tensor0(2.0f), Tensor0(3.0f)), + (Tensor0(3), Tensor0(4)) + ) + val path = Files.createTempFile("tensortree-io-nested", ".pkl") + try + TensorTreeIO.save(params, path) + val restored = TensorTreeIO.load[Model](path) + restored.layers should equal(params.layers) + restored.extra should equal(params.extra) + finally Files.deleteIfExists(path) + + it("round-trips an empty (Unit) tree"): + val path = Files.createTempFile("tensortree-io-unit", ".pkl") + try + TensorTreeIO.save[Unit]((), path) + val restored = TensorTreeIO.load[Unit](path) + restored should equal(()) + finally Files.deleteIfExists(path) From 594f5d7db8cc691263809e1061566ee80a6a9862 Mon Sep 17 00:00:00 2001 From: Marcel Luethi Date: Sat, 1 Aug 2026 17:08:36 +0200 Subject: [PATCH 3/4] fix agent doc --- AGENTS.md | 21 +++++++++++---------- mdocs/AGENTS.md | 3 ++- 2 files changed, 13 insertions(+), 11 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index dd3dcd45..e4fbd269 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -758,7 +758,8 @@ Use **case classes** to group parameters. DimWit automatically derives `TensorTr ```scala import dimwit.* -import dimwit.autodiff.{TensorTree, FloatTree, Autodiff} +import dimwit.autodiff.{Autodiff} +import dimwit.tensortree.{TensorTree, FloatTree} trait Feature derives Label trait Hidden derives Label @@ -1283,23 +1284,23 @@ val wrong = Autodiff.grad(nonScalar) // Use jacobian instead // [Input, V] // (f: Input => dimwit.tensor.Tensor0[V]) // (using evidence$1: dimwit.tensor.TensorOps.IsFloating[V], -// inTree: dimwit.autodiff.TensorTree[Input], outTree: -// dimwit.autodiff.TensorTree[dimwit.tensor.Tensor0[V]]): Input => +// inTree: dimwit.tensortree.TensorTree[Input], outTree: +// dimwit.tensortree.TensorTree[dimwit.tensor.Tensor0[V]]): Input => // dimwit.autodiff.Grad[Input] // [T1, T2, T3, V²] // (f: (T1, T2, T3) => dimwit.tensor.Tensor0[V²]) // (using evidence$1²: dimwit.tensor.TensorOps.IsFloating[V²], -// t1Tree: dimwit.autodiff.TensorTree[T1], -// t2Tree: dimwit.autodiff.TensorTree[T2], -// t3Tree: dimwit.autodiff.TensorTree[T3], outTree²: -// dimwit.autodiff.TensorTree[dimwit.tensor.Tensor0[V²]]): (T1, T2, T3) => +// t1Tree: dimwit.tensortree.TensorTree[T1], +// t2Tree: dimwit.tensortree.TensorTree[T2], +// t3Tree: dimwit.tensortree.TensorTree[T3], outTree²: +// dimwit.tensortree.TensorTree[dimwit.tensor.Tensor0[V²]]): (T1, T2, T3) => // dimwit.autodiff.Grad[(T1, T2, T3)] // [T1², T2², V³] // (f: (T1², T2²) => dimwit.tensor.Tensor0[V³]) // (using evidence$1³: dimwit.tensor.TensorOps.IsFloating[V³], -// t1Tree²: dimwit.autodiff.TensorTree[T1²], -// t2Tree²: dimwit.autodiff.TensorTree[T2²], outTree³: -// dimwit.autodiff.TensorTree[dimwit.tensor.Tensor0[V³]]): (T1², T2²) => +// t1Tree²: dimwit.tensortree.TensorTree[T1²], +// t2Tree²: dimwit.tensortree.TensorTree[T2²], outTree³: +// dimwit.tensortree.TensorTree[dimwit.tensor.Tensor0[V³]]): (T1², T2²) => // dimwit.autodiff.Grad[(T1², T2²)] // match arguments (dimwit.tensor.Tensor1[MdocApp12.this.A, dimwit.Float32] => // dimwit.tensor.Tensor1[MdocApp12.this.A, dimwit.Float32]) diff --git a/mdocs/AGENTS.md b/mdocs/AGENTS.md index 0cff8cf7..dfda177a 100644 --- a/mdocs/AGENTS.md +++ b/mdocs/AGENTS.md @@ -573,7 +573,8 @@ Use **case classes** to group parameters. DimWit automatically derives `TensorTr ```scala mdoc:reset:silent import dimwit.* -import dimwit.autodiff.{TensorTree, FloatTree, Autodiff} +import dimwit.autodiff.{Autodiff} +import dimwit.tensortree.{TensorTree, FloatTree} trait Feature derives Label trait Hidden derives Label From 2ad4dd2903bcb567db5dbab2eaeca1358a25455a Mon Sep 17 00:00:00 2001 From: Marcel Luethi Date: Mon, 3 Aug 2026 15:36:48 +0200 Subject: [PATCH 4/4] fix tensortree import path in test --- .../test/scala/dimwit/optimizer/GradientOptimizerSuite.scala | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/core/src/test/scala/dimwit/optimizer/GradientOptimizerSuite.scala b/core/src/test/scala/dimwit/optimizer/GradientOptimizerSuite.scala index 931a2e84..dc8e2fb9 100644 --- a/core/src/test/scala/dimwit/optimizer/GradientOptimizerSuite.scala +++ b/core/src/test/scala/dimwit/optimizer/GradientOptimizerSuite.scala @@ -2,8 +2,8 @@ package dimwit.optimizer import dimwit.* import dimwit.Conversions.given -import dimwit.autodiff.FloatTree.* -import dimwit.autodiff.FloatTree.ops.* +import dimwit.tensortree.FloatTree.* +import dimwit.tensortree.FloatTree.ops.* import dimwit.autodiff.* class GradientOptimizerSuite extends DimwitTest: