diff --git a/core/src/main/scala/dimwit/tensortree/FloatTree.scala b/core/src/main/scala/dimwit/tensortree/FloatTree.scala index 4f8b2cb..4b3cd99 100644 --- a/core/src/main/scala/dimwit/tensortree/FloatTree.scala +++ b/core/src/main/scala/dimwit/tensortree/FloatTree.scala @@ -3,6 +3,7 @@ package dimwit.tensortree import dimwit.tensor.TensorOps.* import dimwit.tensor.* +import scala.NamedTuple.NamedTuple import scala.deriving.* import scala.util.NotGiven @@ -31,6 +32,9 @@ object FloatTree: given mapInstance[K, A, V](using FloatTree[A, V]): FloatTree[Map[K, A], V] with {} + // 4. Named tuples, delegating to the FloatTree instance of the underlying value tuple + given namedTupleInstance[N <: Tuple, V <: Tuple, Fl](using FloatTree[V, Fl]): FloatTree[NamedTuple[N, V], Fl] with {} + inline given derived[P <: Product, V](using evNotTuple: NotGiven[P <:< Tuple], m: Mirror.ProductOf[P], diff --git a/core/src/main/scala/dimwit/tensortree/TensorTree.scala b/core/src/main/scala/dimwit/tensortree/TensorTree.scala index 1e742bb..9467e45 100644 --- a/core/src/main/scala/dimwit/tensortree/TensorTree.scala +++ b/core/src/main/scala/dimwit/tensortree/TensorTree.scala @@ -6,6 +6,7 @@ import dimwit.tensor.DType.Float32 import me.shadaj.scalapy.py import me.shadaj.scalapy.py.SeqConverters +import scala.NamedTuple.NamedTuple import scala.compiletime.* import scala.deriving.* @@ -200,6 +201,31 @@ object TensorTree: // extends TensorTreeLowPriority: val len = py.Dynamic.global.len(pyList).as[Int] List.tabulate(len)(i => tp.fromPyTree(pyList.bracketAccess(i))) + given namedTupleInstance[N <: Tuple, V <: Tuple](using tt: TensorTree[V]): TensorTree[NamedTuple[N, V]] with + def map(p: NamedTuple[N, V], f: [T <: Tuple, V2] => (Labels[T]) ?=> (Tensor[T, V2] => Tensor[T, V2])): NamedTuple[N, V] = + tt.map(p.toTuple, f) + + def mapWithName(p: NamedTuple[N, V], f: [T <: Tuple, V2] => (Labels[T]) ?=> ((String, Tensor[T, V2]) => Tensor[T, V2]), path: String = ""): NamedTuple[N, V] = + tt.mapWithName(p.toTuple, f, path) + + def mapLeaves[A](p: NamedTuple[N, V], f: [T <: Tuple, V2] => (Labels[T]) ?=> (Tensor[T, V2] => A)): Iterator[A] = + tt.mapLeaves(p.toTuple, f) + + def foreach(p: NamedTuple[N, V], f: [T <: Tuple, V2] => (Labels[T]) ?=> (Tensor[T, V2] => Unit)): Unit = + tt.foreach(p.toTuple, f) + + def foreachWithName(p: NamedTuple[N, V], f: [T <: Tuple, V2] => (Labels[T]) ?=> ((String, Tensor[T, V2]) => Unit), path: String = ""): Unit = + tt.foreachWithName(p.toTuple, f, path) + + def zipMap(p1: NamedTuple[N, V], p2: NamedTuple[N, V], f: [T <: Tuple, V2] => (Labels[T]) ?=> ((Tensor[T, V2], Tensor[T, V2]) => Tensor[T, V2])): NamedTuple[N, V] = + tt.zipMap(p1.toTuple, p2.toTuple, f) + + def toPyTree(p: NamedTuple[N, V]): Jax.PyAny = + tt.toPyTree(p.toTuple) + + def fromPyTree(pyVal: Jax.PyAny): NamedTuple[N, V] = + tt.fromPyTree(pyVal) + /** automatically derive a TensorTree instance for any case class (or product type) * whose fields all have TensorTree instances. */ diff --git a/core/src/test/scala/dimwit/tensortree/PyTreeSuite.scala b/core/src/test/scala/dimwit/tensortree/PyTreeSuite.scala deleted file mode 100644 index d1407a3..0000000 --- a/core/src/test/scala/dimwit/tensortree/PyTreeSuite.scala +++ /dev/null @@ -1,83 +0,0 @@ -package dimwit.tensortree - -import dimwit.* -import dimwit.jax.Jax -import me.shadaj.scalapy.py -class ToPyTreeSuite extends DimwitTest: - - describe("TensorTree Identity (fromPyTree(toPyTree(x)) == x)"): - - it("1-level case class"): - case class Params( - val w: Tensor1[A, Float32], - val b: Tensor0[Float32] - ) - val params = Params( - Tensor1(Axis[A]).fromArray(Array(0.1f, 0.2f, 0.3f)), - Tensor0(0.5f) - ) - - val tc = TensorTree[Params] - val reconstructed = tc.fromPyTree(tc.toPyTree(params)) - - reconstructed.w should approxEqual(params.w) - reconstructed.b should approxEqual(params.b) - - it("2-level case class"): - case class LayerParams( - val w: Tensor2[A, B, Float32], - val b: Tensor0[Float32] - ) - case class ModelParams( - val layer1: LayerParams, - val layer2: LayerParams - ) - - val params = ModelParams( - LayerParams( - Tensor2(Axis[A], Axis[B]).fromArray(Array(Array(0.1f, 0.2f), Array(0.3f, 0.4f))), - Tensor0(0.25f) - ), - LayerParams( - Tensor2(Axis[A], Axis[B]).fromArray(Array(Array(0.5f, 0.6f), Array(0.7f, 0.8f))), - Tensor0(0.75f) - ) - ) - - val tc = TensorTree[ModelParams] - val reconstructed = tc.fromPyTree(tc.toPyTree(params)) - - reconstructed.layer1.w should approxEqual(params.layer1.w) - reconstructed.layer1.b should approxEqual(params.layer1.b) - reconstructed.layer2.w should approxEqual(params.layer2.w) - reconstructed.layer2.b should approxEqual(params.layer2.b) - - it("tuple"): - val myTuple = ( - Tensor1(Axis[A]).fromArray(Array(0.1f, 0.2f, 0.3f)), - Tensor0(0.5f) - ) - - val tc = TensorTree[(Tensor1[A, Float32], Tensor0[Float32])] - val reconstructed = tc.fromPyTree(tc.toPyTree(myTuple)) - - reconstructed._1 should approxEqual(myTuple._1) - reconstructed._2 should approxEqual(myTuple._2) - - it("case class with list"): - case class Params( - val layerWeights: List[Tensor2[A, B, Float32]] - ) - val params = Params( - List( - Tensor2(Axis[A], Axis[B]).fromArray(Array(Array(0.1f, 0.2f), Array(0.3f, 0.4f))), - Tensor2(Axis[A], Axis[B]).fromArray(Array(Array(1.1f, 1.2f), Array(1.3f, 1.4f))) - ) - ) - - val tc = TensorTree[Params] - val reconstructed = tc.fromPyTree(tc.toPyTree(params)) - - reconstructed.layerWeights.size shouldBe params.layerWeights.size - reconstructed.layerWeights(0) should approxEqual(params.layerWeights(0)) - reconstructed.layerWeights(1) should approxEqual(params.layerWeights(1)) diff --git a/core/src/test/scala/dimwit/tensortree/TensorTreeSuite.scala b/core/src/test/scala/dimwit/tensortree/TensorTreeSuite.scala index d04347f..98c8880 100644 --- a/core/src/test/scala/dimwit/tensortree/TensorTreeSuite.scala +++ b/core/src/test/scala/dimwit/tensortree/TensorTreeSuite.scala @@ -22,6 +22,17 @@ class TensorTreeSuite extends DimwitTest: tree2.counts should equal(params.counts) tree2.flags should equal(params.flags) + it("named tuple"): + type Data = (numbers: Tensor1[A, Float32], counts: Tensor1[A, Int32]) + val params: Data = ( + numbers = Tensor1(Axis[A]).fromArray(Array(0.1f, 0.2f, 0.3f)), + counts = Tensor1(Axis[A]).fromArray(Array(1, 2, 3)) + ) + val tree = TensorTree[Data] + val tree2 = tree.map(params, [T <: Tuple, V] => (labels: Labels[T]) ?=> (x: Tensor[T, V]) => x) + tree2.numbers should approxEqual(params.numbers) + tree2.counts should equal(params.counts) + describe("zipmap"): it("1-level case class"): case class Params( @@ -41,6 +52,15 @@ class TensorTreeSuite extends DimwitTest: res.w1 should approxEqual(maximum(params1.w1, params2.w1)) res.b1 should equal(maximum(params1.b1, params2.b1)) + it("named tuple"): + type Params = (w1: Tensor1[A, Float32], b1: Tensor0[Int32]) + val params1: Params = (w1 = Tensor1(Axis[A]).fromArray(Array(0.1f, 0.2f, 0.3f)), b1 = Tensor0(0)) + val params2: Params = (w1 = Tensor1(Axis[A]).fromArray(Array(0.4f, 0.5f, 0.6f)), b1 = Tensor0(1)) + val ftTree = TensorTree[Params] + val res = ftTree.zipMap(params1, params2, [T <: Tuple, V] => (labels: Labels[T]) ?=> (x1: Tensor[T, V], x2: Tensor[T, V]) => maximum(x1, x2)) + res.w1 should approxEqual(maximum(params1.w1, params2.w1)) + res.b1 should equal(maximum(params1.b1, params2.b1)) + describe("mapLeaves"): it("1-level case class"): @@ -78,6 +98,17 @@ class TensorTreeSuite extends DimwitTest: val leavesCount = tree.mapLeaves(paramsList, [T <: Tuple, V] => (labels: Labels[T]) ?=> (x: Tensor[T, V]) => 1).sum leavesCount should equal(3) + it("named tuple"): + type Data = (numbers: Tensor1[A, Float32], counts: Tensor1[A, Int32], flags: Tensor1[A, Bool]) + val params: Data = ( + numbers = Tensor1(Axis[A]).fromArray(Array(0.1f, 0.2f, 0.3f)), + counts = Tensor1(Axis[A]).fromArray(Array(1, 2, 3)), + flags = Tensor1(Axis[A]).fromArray(Array(true, false, true)) + ) + val tree = TensorTree[Data] + val leavesCount = tree.mapLeaves(params, [T <: Tuple, V] => (labels: Labels[T]) ?=> (x: Tensor[T, V]) => 1).sum + leavesCount should equal(3) + describe("foreach"): it("1-level case class"): case class Data( @@ -144,6 +175,21 @@ class TensorTreeSuite extends DimwitTest: ) paths.toList should equal(List("layers[0]", "layers[1]", "extra._1", "extra._2")) + it("named tuple (paths fall back to the underlying tuple positions)"): + type Params = (w1: Tensor1[A, Float32], b1: Tensor0[Int32]) + val params: Params = (w1 = Tensor1(Axis[A]).fromArray(Array(0.1f, 0.2f, 0.3f)), b1 = Tensor0(0)) + val tree = TensorTree[Params] + var paths = Vector.empty[String] + tree.mapWithName( + params, + [T <: Tuple, V] => + (labels: Labels[T]) ?=> + (path: String, x: Tensor[T, V]) => + paths = paths :+ path + x + ) + paths.toList should equal(List("_1", "_2")) + describe("foreachWithName"): it("1-level case class"): case class Data( @@ -182,3 +228,89 @@ class TensorTreeSuite extends DimwitTest: paths = paths :+ path ) paths.toList should equal(List("inner1.w", "inner2.w")) + + describe("toPyTree / fromPyTree (fromPyTree(toPyTree(x)) == x)"): + it("1-level case class"): + case class Params( + val w: Tensor1[A, Float32], + val b: Tensor0[Float32] + ) + val params = Params( + Tensor1(Axis[A]).fromArray(Array(0.1f, 0.2f, 0.3f)), + Tensor0(0.5f) + ) + + val tc = TensorTree[Params] + val reconstructed = tc.fromPyTree(tc.toPyTree(params)) + + reconstructed.w should approxEqual(params.w) + reconstructed.b should approxEqual(params.b) + + it("2-level case class"): + case class LayerParams( + val w: Tensor2[A, B, Float32], + val b: Tensor0[Float32] + ) + case class ModelParams( + val layer1: LayerParams, + val layer2: LayerParams + ) + + val params = ModelParams( + LayerParams( + Tensor2(Axis[A], Axis[B]).fromArray(Array(Array(0.1f, 0.2f), Array(0.3f, 0.4f))), + Tensor0(0.25f) + ), + LayerParams( + Tensor2(Axis[A], Axis[B]).fromArray(Array(Array(0.5f, 0.6f), Array(0.7f, 0.8f))), + Tensor0(0.75f) + ) + ) + + val tc = TensorTree[ModelParams] + val reconstructed = tc.fromPyTree(tc.toPyTree(params)) + + reconstructed.layer1.w should approxEqual(params.layer1.w) + reconstructed.layer1.b should approxEqual(params.layer1.b) + reconstructed.layer2.w should approxEqual(params.layer2.w) + reconstructed.layer2.b should approxEqual(params.layer2.b) + + it("tuple"): + val myTuple = ( + Tensor1(Axis[A]).fromArray(Array(0.1f, 0.2f, 0.3f)), + Tensor0(0.5f) + ) + + val tc = TensorTree[(Tensor1[A, Float32], Tensor0[Float32])] + val reconstructed = tc.fromPyTree(tc.toPyTree(myTuple)) + + reconstructed._1 should approxEqual(myTuple._1) + reconstructed._2 should approxEqual(myTuple._2) + + it("case class with list"): + case class Params( + val layerWeights: List[Tensor2[A, B, Float32]] + ) + val params = Params( + List( + Tensor2(Axis[A], Axis[B]).fromArray(Array(Array(0.1f, 0.2f), Array(0.3f, 0.4f))), + Tensor2(Axis[A], Axis[B]).fromArray(Array(Array(1.1f, 1.2f), Array(1.3f, 1.4f))) + ) + ) + + val tc = TensorTree[Params] + val reconstructed = tc.fromPyTree(tc.toPyTree(params)) + + reconstructed.layerWeights.size shouldBe params.layerWeights.size + reconstructed.layerWeights(0) should approxEqual(params.layerWeights(0)) + reconstructed.layerWeights(1) should approxEqual(params.layerWeights(1)) + + it("named tuple"): + type Params = (w: Tensor1[A, Float32], b: Tensor0[Float32]) + val params: Params = (w = Tensor1(Axis[A]).fromArray(Array(0.1f, 0.2f, 0.3f)), b = Tensor0(0.5f)) + + val tc = TensorTree[Params] + val reconstructed = tc.fromPyTree(tc.toPyTree(params)) + + reconstructed.w should approxEqual(params.w) + reconstructed.b should approxEqual(params.b)