diff --git a/AGENTS.md b/AGENTS.md index 2c82bdd2..2bbce1fd 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -1240,8 +1240,8 @@ val wrong = intTensor.exp // exp requires IsFloating constraint // dimwit.tensor.Labels.consTuple[MdocApp12.this.A, EmptyTuple.type]( // this.A.derived$Label, dimwit.tensor.Labels.emptyTuple), // /* missing */ -// summon[dimwit.tensor.TensorOps.IsFloating[dimwit.tensor.DType.Int32]] -// ) +// summon[dimwit.tensor.ValueTypeClasses.IsFloating[dimwit.tensor.DType.Int32]] +// ) // // failed with: // @@ -1261,7 +1261,7 @@ val wrong = boolTensor.mean // dimwit.tensor.Labels.consTuple[MdocApp12.this.A, EmptyTuple.type]( // this.A.derived$Label, dimwit.tensor.Labels.emptyTuple), // /* missing */ -// summon[dimwit.tensor.TensorOps.IsFloating[dimwit.tensor.DType.Bool]] +// summon[dimwit.tensor.ValueTypeClasses.IsFloating[dimwit.tensor.DType.Bool]] // ) // // failed with: @@ -1332,9 +1332,8 @@ val wrong = t1 +! t2 // Cannot broadcast tensors of shapes Tuple1[MdocApp12.this.A] and Tuple1[MdocApp12.this.A]. If same shape no broadcasting allowed!. // I found: // -// dimwit.tensor.tensorops.TensorOpsUtil.Broadcast.broadcastLeft[ -// Tuple1[MdocApp12.this.A], Tuple1[MdocApp12.this.A], -// dimwit.tensor.DType.Float32]( +// dimwit.tensor.Broadcast.broadcastLeft[Tuple1[MdocApp12.this.A], +// Tuple1[MdocApp12.this.A], dimwit.tensor.DType.Float32]( // dimwit.tensor.Labels.consTuple[MdocApp12.this.A, EmptyTuple.type]( // this.A.derived$Label, dimwit.tensor.Labels.emptyTuple), // dimwit.tensor.Labels.consTuple[MdocApp12.this.A, EmptyTuple.type]( @@ -1422,13 +1421,13 @@ val wrong = Autodiff.grad(nonScalar) // Use jacobian instead // None of the overloaded alternatives of method grad in object Autodiff with types // [Input, V] // (f: Input => dimwit.tensor.Tensor0[V]) -// (using evidence$1: dimwit.tensor.TensorOps.IsFloating[V], +// (using evidence$1: dimwit.tensor.ValueTypeClasses.IsFloating[V], // 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²], +// (using evidence$1²: dimwit.tensor.ValueTypeClasses.IsFloating[V²], // t1Tree: dimwit.tensortree.TensorTree[T1], // t2Tree: dimwit.tensortree.TensorTree[T2], // t3Tree: dimwit.tensortree.TensorTree[T3], outTree²: @@ -1436,7 +1435,7 @@ val wrong = Autodiff.grad(nonScalar) // Use jacobian instead // dimwit.autodiff.Grad[(T1, T2, T3)] // [T1², T2², V³] // (f: (T1², T2²) => dimwit.tensor.Tensor0[V³]) -// (using evidence$1³: dimwit.tensor.TensorOps.IsFloating[V³], +// (using evidence$1³: dimwit.tensor.ValueTypeClasses.IsFloating[V³], // t1Tree²: dimwit.tensortree.TensorTree[T1²], // t2Tree²: dimwit.tensortree.TensorTree[T2²], outTree³: // dimwit.tensortree.TensorTree[dimwit.tensor.Tensor0[V³]]): (T1², T2²) => diff --git a/core/src/main/scala/dimwit/autodiff/Autodiff.scala b/core/src/main/scala/dimwit/autodiff/Autodiff.scala index 28d79462..6e2ed243 100644 --- a/core/src/main/scala/dimwit/autodiff/Autodiff.scala +++ b/core/src/main/scala/dimwit/autodiff/Autodiff.scala @@ -6,7 +6,7 @@ import dimwit.jax.Jax import dimwit.prime.PrimeConcat import dimwit.tensor.Tensor import dimwit.tensor.Tensor0 -import dimwit.tensor.TensorOps.IsFloating +import dimwit.tensor.ValueTypeClasses.IsFloating import dimwit.tensortree.TensorTree import me.shadaj.scalapy.py diff --git a/core/src/main/scala/dimwit/linalg/LinearAlgebra.scala b/core/src/main/scala/dimwit/linalg/LinearAlgebra.scala index 1afc6295..cfd457c0 100644 --- a/core/src/main/scala/dimwit/linalg/LinearAlgebra.scala +++ b/core/src/main/scala/dimwit/linalg/LinearAlgebra.scala @@ -9,8 +9,8 @@ import dimwit.tensor.Tensor import dimwit.tensor.Tensor0 import dimwit.tensor.Tensor1 import dimwit.tensor.Tensor2 -import dimwit.tensor.TensorOps.IsFloating -import dimwit.tensor.TensorOps.IsNumber +import dimwit.tensor.ValueTypeClasses.IsFloating +import dimwit.tensor.ValueTypeClasses.IsNumber import me.shadaj.scalapy.py /** Common linear algebra operations. diff --git a/core/src/main/scala/dimwit/package.scala b/core/src/main/scala/dimwit/package.scala index 555ecde4..7d0e7bb5 100644 --- a/core/src/main/scala/dimwit/package.scala +++ b/core/src/main/scala/dimwit/package.scala @@ -62,10 +62,34 @@ package object dimwit: // Export the Prime axis marker and the type classes that manipulate it export dimwit.prime.{Prime, PrimeRemover, PrimeRest, PrimeConcat} - // Export operations - export dimwit.tensor.TensorOps.* + // Export the type classes on value types, e.g. IsFloating + export dimwit.tensor.ValueTypeClasses.* + + // Export the extension methods on tensors, e.g. `t.relu` or `t.sum(Axis[A])`. + // The functions they call are exported on the Tensor companion objects, e.g. `Tensor.relu(t)`. + export dimwit.tensor.tensorops.ElementWiseExtensions.* + export dimwit.tensor.tensorops.ReductionExtensions.* + export dimwit.tensor.tensorops.AlongAxisExtensions.* + export dimwit.tensor.tensorops.ContractionExtensions.* + export dimwit.tensor.tensorops.ConvolutionExtensions.* + export dimwit.tensor.tensorops.LinearAlgebraExtensions.* + export dimwit.tensor.tensorops.StructuralExtensions.* + export dimwit.tensor.tensorops.FunctionalExtensions.* + export dimwit.tensor.tensorops.Tensor0Extensions.* + export dimwit.tensor.tensorops.Tensor1Extensions.* + export dimwit.tensor.tensorops.Tensor2Extensions.* + export dimwit.tensor.tensorops.Tensor3Extensions.* + export dimwit.tensor.ValueExtensions.* + + // Export the functions that have no extension method or operator doing the same, e.g. `maximum(t1, t2)` or `stack(tensors, Axis[A])`. + export dimwit.tensor.tensorops.ElementWiseOps.{maximum, minimum, maximum_!, minimum_!} + export dimwit.tensor.tensorops.StructuralOps.{where, where_!, triu, tril, stack, concatenate} + export dimwit.tensor.tensorops.FunctionalOps.zipvmap + + // Export convolution options + export dimwit.tensor.{Padding, Stride1, Stride2, Stride3} + export dimwit.linalg.LinearAlgebra.{VectorNormType, MatrixNormType, QRMode} - export dimwit.tensor.ValueOps.* // Export devices export dimwit.hardware.Device diff --git a/core/src/main/scala/dimwit/random/Random.scala b/core/src/main/scala/dimwit/random/Random.scala index bdd89bc0..d0f8ce28 100644 --- a/core/src/main/scala/dimwit/random/Random.scala +++ b/core/src/main/scala/dimwit/random/Random.scala @@ -4,7 +4,7 @@ import dimwit.tensortree.TensorTree import dimwit.jax.Jax import dimwit.python.PyBridge.liftPyTensor import dimwit.tensor.DType.Int32 -import dimwit.tensor.TensorOps.* +import dimwit.* import dimwit.tensor.TupleHelpers.TupleNOf import dimwit.tensor.* diff --git a/core/src/main/scala/dimwit/stats/Distributions.scala b/core/src/main/scala/dimwit/stats/Distributions.scala index 8efb99d1..72ff5929 100644 --- a/core/src/main/scala/dimwit/stats/Distributions.scala +++ b/core/src/main/scala/dimwit/stats/Distributions.scala @@ -2,7 +2,6 @@ package dimwit.stats import dimwit.* import dimwit.random.Random -import dimwit.tensor.TensorOps opaque type LogProb = Float32 opaque type Prob = Float32 @@ -15,8 +14,8 @@ object LogProb: extension [T <: Tuple: Labels](t: Tensor[T, LogProb]) - def exp: Tensor[T, Prob] = TensorOps.exp(t) - def log: Tensor[T, Float32] = TensorOps.log(t) // Lose LogProb if we log again + def exp: Tensor[T, Prob] = Tensor.exp(t) + def log: Tensor[T, Float32] = Tensor.log(t) // Lose LogProb if we log again def asFloat: Tensor[T, Float32] = t object Prob: @@ -27,8 +26,8 @@ object Prob: extension [T <: Tuple: Labels](t: Tensor[T, Prob]) - def exp: Tensor[T, Float32] = TensorOps.exp(t) // Lose Prob if we exp again - def log: Tensor[T, LogProb] = TensorOps.log(t) + def exp: Tensor[T, Float32] = Tensor.exp(t) // Lose Prob if we exp again + def log: Tensor[T, LogProb] = Tensor.log(t) def asFloat: Tensor[T, Float32] = t trait Distribution[EventShape <: Tuple: Labels, V]: diff --git a/core/src/main/scala/dimwit/tensor/ArrayWriter.scala b/core/src/main/scala/dimwit/tensor/ArrayWriter.scala index bb453fcb..d9bfea1d 100644 --- a/core/src/main/scala/dimwit/tensor/ArrayWriter.scala +++ b/core/src/main/scala/dimwit/tensor/ArrayWriter.scala @@ -1,9 +1,9 @@ package dimwit.tensor import dimwit.jax.Jax -import dimwit.tensor.TensorOps.IsBoolean -import dimwit.tensor.TensorOps.IsFloating -import dimwit.tensor.TensorOps.IsInteger +import dimwit.tensor.ValueTypeClasses.IsBoolean +import dimwit.tensor.ValueTypeClasses.IsFloating +import dimwit.tensor.ValueTypeClasses.IsInteger import me.shadaj.scalapy.py import me.shadaj.scalapy.py.SeqConverters import me.shadaj.scalapy.readwrite.Writer diff --git a/core/src/main/scala/dimwit/tensor/Broadcast.scala b/core/src/main/scala/dimwit/tensor/Broadcast.scala new file mode 100644 index 00000000..eb09148f --- /dev/null +++ b/core/src/main/scala/dimwit/tensor/Broadcast.scala @@ -0,0 +1,90 @@ +package dimwit.tensor + +import dimwit.tensor.TupleHelpers.StrictSubset +import dimwit.tensor.tensorops.StructuralExtensions.broadcastTo + +import scala.annotation.implicitNotFound + +@implicitNotFound("Cannot broadcast tensors of shapes ${T1} and ${T2}. If same shape no broadcasting allowed!") +sealed trait Broadcast[T1 <: Tuple, T2 <: Tuple, V]: + type Out <: Tuple + given labelsOut: Labels[Out] + def broadcast(t1: Tensor[T1, V], t2: Tensor[T2, V]): (Tensor[Out, V], Tensor[Out, V]) + def applyTo[V2](t1: Tensor[T1, V], t2: Tensor[T2, V])(f: (Tensor[Out, V], Tensor[Out, V]) => Tensor[Out, V2]): Tensor[Out, V2] = + val (bt1, bt2) = broadcast(t1, t2) + f(bt1, bt2) + +object Broadcast extends BroadcastLowPriority: + + given broadcastLeft[T1 <: Tuple: Labels, T2 <: Tuple: Labels, V](using + StrictSubset[T2, T1] + ): Broadcast[T1, T2, V] with + type Out = T1 + val labelsOut = summon[Labels[T1]] + def broadcast(t1: Tensor[T1, V], t2: Tensor[T2, V]) = + (t1, t2.broadcastTo[T1](t1.shape)) + +trait BroadcastLowPriority: + given broadcastRight[T1 <: Tuple: Labels, T2 <: Tuple: Labels, V](using + StrictSubset[T1, T2] + ): Broadcast[T1, T2, V] with + type Out = T2 + val labelsOut = summon[Labels[T2]] + def broadcast(t1: Tensor[T1, V], t2: Tensor[T2, V]) = + (t1.broadcastTo[T2](t2.shape), t2) + +/** Broadcasts three tensors to their common shape, which is the one of the three shapes containing all + * axes of the other two. As for [[Broadcast]] at least one of the tensors has to be broadcast. + */ +@implicitNotFound( + "Cannot broadcast tensors of shapes ${T1}, ${T2} and ${T3}. One of them must contain all axes of the other two. If all same shape no broadcasting allowed!" +) +sealed trait Broadcast3[T1 <: Tuple, T2 <: Tuple, T3 <: Tuple, V]: + type Out <: Tuple + given labelsOut: Labels[Out] + def broadcast[V1](t1: Tensor[T1, V1], t2: Tensor[T2, V], t3: Tensor[T3, V]): (Tensor[Out, V1], Tensor[Out, V], Tensor[Out, V]) + +object Broadcast3 extends Broadcast3LowPriority: + + /** `t2` and `t3` broadcast against each other, `t1` already has their common shape. */ + given valuesBroadcast[O <: Tuple, T2 <: Tuple, T3 <: Tuple, V](using + bc: Broadcast[T2, T3, V] { type Out = O } + ): Broadcast3[O, T2, T3, V] with + type Out = O + val labelsOut = bc.labelsOut + def broadcast[V1](t1: Tensor[O, V1], t2: Tensor[T2, V], t3: Tensor[T3, V]) = + val (bt2, bt3) = bc.broadcast(t2, t3) + (t1, bt2, bt3) + + /** `t2` and `t3` broadcast against each other, `t1` is broadcast to their common shape. */ + given conditionAndValuesBroadcast[T1 <: Tuple: Labels, T2 <: Tuple, T3 <: Tuple, O <: Tuple, V](using + bc: Broadcast[T2, T3, V] { type Out = O }, + ev: StrictSubset[T1, O] + ): Broadcast3[T1, T2, T3, V] with + type Out = O + val labelsOut = bc.labelsOut + def broadcast[V1](t1: Tensor[T1, V1], t2: Tensor[T2, V], t3: Tensor[T3, V]) = + given Labels[O] = bc.labelsOut + val (bt2, bt3) = bc.broadcast(t2, t3) + (t1.broadcastTo[O](bt2.shape), bt2, bt3) + + /** `t2` and `t3` have the same shape, only `t1` is broadcast to it. */ + given conditionBroadcast[T1 <: Tuple: Labels, T <: Tuple: Labels, V](using + ev: StrictSubset[T1, T] + ): Broadcast3[T1, T, T, V] with + type Out = T + val labelsOut = summon[Labels[T]] + def broadcast[V1](t1: Tensor[T1, V1], t2: Tensor[T, V], t3: Tensor[T, V]) = + (t1.broadcastTo[T](t2.shape), t2, t3) + +trait Broadcast3LowPriority: + + /** `t2` and `t3` are both broadcast to the shape of `t1`. */ + given valuesBroadcastToFirst[T1 <: Tuple: Labels, T2 <: Tuple: Labels, T3 <: Tuple: Labels, V](using + ev2: StrictSubset[T2, T1], + ev3: StrictSubset[T3, T1] + ): Broadcast3[T1, T2, T3, V] with + type Out = T1 + val labelsOut = summon[Labels[T1]] + def broadcast[V1](t1: Tensor[T1, V1], t2: Tensor[T2, V], t3: Tensor[T3, V]) = + (t1, t2.broadcastTo[T1](t1.shape), t3.broadcastTo[T1](t1.shape)) diff --git a/core/src/main/scala/dimwit/tensor/Convolution.scala b/core/src/main/scala/dimwit/tensor/Convolution.scala new file mode 100644 index 00000000..f1c3a312 --- /dev/null +++ b/core/src/main/scala/dimwit/tensor/Convolution.scala @@ -0,0 +1,20 @@ +package dimwit.tensor + +/** Padding options for convolution operations. + * SAME: Output size is the same as input size (with appropriate padding). + * VALID: No padding, output size is reduced based on kernel size. + * + * Refer to JAX documentation for more details on padding behavior. + * https://jax.readthedocs.io/en/latest/_autosummary/jax.lax.conv_general_dilated.html + */ +enum Padding: + case SAME, VALID + +/** Stride of a 1D convolution. */ +type Stride1[S1] = AxisExtent[S1] + +/** Stride of a 2D convolution. */ +type Stride2[S1, S2] = (AxisExtent[S1], AxisExtent[S2]) + +/** Stride of a 3D convolution. */ +type Stride3[S1, S2, S3] = (AxisExtent[S1], AxisExtent[S2], AxisExtent[S3]) diff --git a/core/src/main/scala/dimwit/tensor/DType.scala b/core/src/main/scala/dimwit/tensor/DType.scala index a59c39fc..77eb17c4 100644 --- a/core/src/main/scala/dimwit/tensor/DType.scala +++ b/core/src/main/scala/dimwit/tensor/DType.scala @@ -1,9 +1,9 @@ package dimwit.tensor import dimwit.jax.JaxDType import dimwit.tensor.HasScalar -import dimwit.tensor.TensorOps.IsBoolean -import dimwit.tensor.TensorOps.IsFloating -import dimwit.tensor.TensorOps.IsInteger +import dimwit.tensor.ValueTypeClasses.IsBoolean +import dimwit.tensor.ValueTypeClasses.IsFloating +import dimwit.tensor.ValueTypeClasses.IsInteger import me.shadaj.scalapy.py import java.nio.ByteBuffer diff --git a/core/src/main/scala/dimwit/tensor/Tensor.scala b/core/src/main/scala/dimwit/tensor/Tensor.scala index 9f57b344..69e0e0e5 100644 --- a/core/src/main/scala/dimwit/tensor/Tensor.scala +++ b/core/src/main/scala/dimwit/tensor/Tensor.scala @@ -7,10 +7,10 @@ import dimwit.jax.Jax.PyDynamic import dimwit.jax.JaxDType import dimwit.tensor.Label import dimwit.tensor.Labels -import dimwit.tensor.TensorOps.IsBoolean -import dimwit.tensor.TensorOps.IsFloating -import dimwit.tensor.TensorOps.IsInteger -import dimwit.tensor.TensorOps.IsNumber +import dimwit.tensor.ValueTypeClasses.IsBoolean +import dimwit.tensor.ValueTypeClasses.IsFloating +import dimwit.tensor.ValueTypeClasses.IsInteger +import dimwit.tensor.ValueTypeClasses.IsNumber import dimwit.tensor.TypedIndex import dimwit.tensor.VType import me.shadaj.scalapy.py @@ -201,6 +201,17 @@ object Tensor: /** Use the [[LikeFactory]] to create a tensor */ def like[T <: Tuple: Labels, V](template: Tensor[T, V]): LikeFactory[T, V] = LikeFactory(template) + // The functions on tensors, e.g. `Tensor.relu(t)` or `Tensor.sum(t, Axis[A])`. + // Each one has an extension method with the same name, e.g. `t.relu` or `t.sum(Axis[A])`. + export tensorops.ElementWiseOps.* + export tensorops.ReductionOps.* + export tensorops.AlongAxisOps.* + export tensorops.ContractionOps.* + export tensorops.ConvolutionOps.* + // StructuralOps and FunctionalOps also hold the type classes of these functions, which are not exported. + export tensorops.StructuralOps.{where, where_!, triu, tril, stack, concatenate} + export tensorops.FunctionalOps.zipvmap + /** Type aliases for tensors of different ranks. */ type Tensor0[V] = Tensor[EmptyTuple, V] type Tensor1[L, V] = Tensor[Tuple1[L], V] @@ -321,54 +332,9 @@ object Tensor1: def apply[L: Label](axisExtent: AxisExtent[L]): Tensor.ShapedFactory[Tuple1[L]] = Tensor.ShapedFactory(Shape(axisExtent)) def apply[L: Label, V](axisExtent: AxisExtent[L], vtype: VType[V]): Tensor.ShapedTypedFactory[Tuple1[L], V] = Tensor.ShapedTypedFactory(Shape(axisExtent), vtype) - // --------------------------------------------------------- - // Functions from a Tensor1 to a Tensor1. + // Functions from a Tensor1 to a Tensor1, e.g. `Tensor1.softmax(t)`. // They can be lifted to any tensor shape with `vapply`, e.g. `t.vapply(Axis[A])(Tensor1.softmax)`. - // --------------------------------------------------------- - - /** sorts the Tensor1 `t`. */ - def sort[L: Label, V: IsNumber](t: Tensor1[L, V]): Tensor1[L, V] = - Tensor(Jax.jnp.sort(t.jaxValue, axis = 0)) - - /** returns the indices that would sort the Tensor1 `t`. */ - def argsort[L: Label, V: IsNumber](t: Tensor1[L, V]): Tensor1[L, Int32] = - Tensor(Jax.jnp.argsort(t.jaxValue, axis = 0)) - - /** computes the cumulative sum of the Tensor1 `t`. */ - def cumsum[L: Label, V: IsNumber](t: Tensor1[L, V]): Tensor1[L, V] = - Tensor(Jax.jnp.cumsum(t.jaxValue, axis = 0)) - - /** computes the cumulative product of the Tensor1 `t`. */ - def cumprod[L: Label, V: IsNumber](t: Tensor1[L, V]): Tensor1[L, V] = - Tensor(Jax.jnp.cumprod(t.jaxValue, axis = 0)) - - /** computes the cumulative maximum of the Tensor1 `t`. */ - def cummax[L: Label, V: IsNumber](t: Tensor1[L, V]): Tensor1[L, V] = - Tensor(Jax.lax.cummax(t.jaxValue, axis = 0)) - - /** computes the cumulative minimum of the Tensor1 `t`. */ - def cummin[L: Label, V: IsNumber](t: Tensor1[L, V]): Tensor1[L, V] = - Tensor(Jax.lax.cummin(t.jaxValue, axis = 0)) - - /** computes the cumulative log-sum-exp of the Tensor1 `t`, i.e. a numerically stable `log(cumsum(exp(t)))`. */ - def logcumsumexp[L: Label, V: IsFloating](t: Tensor1[L, V]): Tensor1[L, V] = - Tensor(Jax.lax.cumlogsumexp(t.jaxValue, axis = 0)) - - /** computes the discrete difference of the Tensor1 `t`, reducing its size by one. */ - def diff[L: Label, V: IsNumber](t: Tensor1[L, V]): Tensor1[L, V] = - Tensor(Jax.jnp.diff(t.jaxValue, axis = 0)) - - /** rolls the elements of the Tensor1 `t` by `shift` positions; elements shifted beyond the end re-appear at the start. */ - def roll[L: Label, V](shift: Int)(t: Tensor1[L, V]): Tensor1[L, V] = - Tensor(Jax.jnp.roll(t.jaxValue, shift = shift, axis = 0)) - - /** computes the softmax of the Tensor1 `t`. */ - def softmax[L: Label, V: IsFloating](t: Tensor1[L, V]): Tensor1[L, V] = - Tensor(Jax.jnn.softmax(t.jaxValue, axis = 0)) - - /** computes the log of the softmax of the Tensor1 `t`, more stable than `softmax(t).log`. */ - def logSoftmax[L: Label, V: IsFloating](t: Tensor1[L, V]): Tensor1[L, V] = - Tensor(Jax.jnn.log_softmax(t.jaxValue, axis = 0)) + export tensorops.AlongAxisTensor1Ops.* /* Companion object for Tensors of rank 2 (matrices). * Provides factory methods for creating tensors of rank 2 with various value types. diff --git a/core/src/main/scala/dimwit/tensor/VType.scala b/core/src/main/scala/dimwit/tensor/VType.scala index 80ec6fc2..3a6c22df 100644 --- a/core/src/main/scala/dimwit/tensor/VType.scala +++ b/core/src/main/scala/dimwit/tensor/VType.scala @@ -1,6 +1,6 @@ package dimwit.tensor -import dimwit.tensor.TensorOps.HasDType +import dimwit.tensor.ValueTypeClasses.HasDType object VType: def apply[V](tensor: Tensor[?, V]): VType[V] = VTypeImpl[V](tensor.dtype) diff --git a/core/src/main/scala/dimwit/tensor/ValueExtensions.scala b/core/src/main/scala/dimwit/tensor/ValueExtensions.scala new file mode 100644 index 00000000..a24851a7 --- /dev/null +++ b/core/src/main/scala/dimwit/tensor/ValueExtensions.scala @@ -0,0 +1,202 @@ +package dimwit.tensor + +import dimwit.tensor.DType.Bool +import dimwit.tensor.ValueTypeClasses.IsFloating +import dimwit.tensor.ValueTypeClasses.IsNumber +import dimwit.tensor.tensorops.ElementWiseOps + +/** Operators with a Scala scalar on the left, e.g. `2.0f *! t` or `3 < t0`. + * + * `t op scalar` works through the implicit conversions in [[dimwit.Conversions]], which convert the scalar to a + * `Tensor0` of the tensor's value type. `scalar op t` cannot use them, because Scala does not convert the receiver of + * a method call. These extension methods mirror every operator `t op scalar` by asking for the very same conversion, + * so both orders compile for exactly the same scalar and value types, and in both the scalar takes the precision of + * the tensor. Named methods like `elementEquals` or `and` are deliberately not mirrored: `2.0f.approxEquals(t)` reads + * as a method of the scalar. + * + * There is one extension block per Scala scalar type. A single generic `extension [S](scalar: S)` would also match + * e.g. `"-" * 3` and hide the `*` that `String` gets through `StringOps`. + */ +private[dimwit] object ValueExtensions: + + extension (scalar: Boolean) + + // operators with a Tensor0 + def +[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Boolean, Tensor0[V]]): Tensor0[V] = ElementWiseOps.add(toTensor0(scalar), t) + def -[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Boolean, Tensor0[V]]): Tensor0[V] = ElementWiseOps.subtract(toTensor0(scalar), t) + def *[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Boolean, Tensor0[V]]): Tensor0[V] = ElementWiseOps.multiply(toTensor0(scalar), t) + def %[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Boolean, Tensor0[V]]): Tensor0[V] = ElementWiseOps.mod(toTensor0(scalar), t) + def /[V: IsFloating](t: Tensor0[V])(using toTensor0: Conversion[Boolean, Tensor0[V]]): Tensor0[V] = ElementWiseOps.divide(toTensor0(scalar), t) + def <[V](t: Tensor0[V])(using toTensor0: Conversion[Boolean, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.less(toTensor0(scalar), t) + def <=[V](t: Tensor0[V])(using toTensor0: Conversion[Boolean, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.lessEqual(toTensor0(scalar), t) + def >[V](t: Tensor0[V])(using toTensor0: Conversion[Boolean, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.greater(toTensor0(scalar), t) + def >=[V](t: Tensor0[V])(using toTensor0: Conversion[Boolean, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.greaterEqual(toTensor0(scalar), t) + def ===[V](t: Tensor0[V])(using toTensor0: Conversion[Boolean, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.arrayEqual(toTensor0(scalar), t) + + // broadcasting operators + def +![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Boolean, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.add_!(toTensor0(scalar), t) + def -![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Boolean, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.subtract_!(toTensor0(scalar), t) + def *![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Boolean, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.multiply_!(toTensor0(scalar), t) + def %![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Boolean, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.mod_!(toTensor0(scalar), t) + def /![T <: Tuple, V: IsFloating](t: Tensor[T, V])(using toTensor0: Conversion[Boolean, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.divide_!(toTensor0(scalar), t) + def `![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Boolean, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, Bool] = ElementWiseOps.greater_!(toTensor0(scalar), t) + def >=![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Boolean, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, Bool] = ElementWiseOps.greaterEqual_!(toTensor0(scalar), t) + def ===![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Boolean, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor0[Bool] = ElementWiseOps.arrayEqual_!(toTensor0(scalar), t) + + extension (scalar: Byte) + + // operators with a Tensor0 + def +[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Byte, Tensor0[V]]): Tensor0[V] = ElementWiseOps.add(toTensor0(scalar), t) + def -[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Byte, Tensor0[V]]): Tensor0[V] = ElementWiseOps.subtract(toTensor0(scalar), t) + def *[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Byte, Tensor0[V]]): Tensor0[V] = ElementWiseOps.multiply(toTensor0(scalar), t) + def %[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Byte, Tensor0[V]]): Tensor0[V] = ElementWiseOps.mod(toTensor0(scalar), t) + def /[V: IsFloating](t: Tensor0[V])(using toTensor0: Conversion[Byte, Tensor0[V]]): Tensor0[V] = ElementWiseOps.divide(toTensor0(scalar), t) + def <[V](t: Tensor0[V])(using toTensor0: Conversion[Byte, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.less(toTensor0(scalar), t) + def <=[V](t: Tensor0[V])(using toTensor0: Conversion[Byte, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.lessEqual(toTensor0(scalar), t) + def >[V](t: Tensor0[V])(using toTensor0: Conversion[Byte, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.greater(toTensor0(scalar), t) + def >=[V](t: Tensor0[V])(using toTensor0: Conversion[Byte, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.greaterEqual(toTensor0(scalar), t) + def ===[V](t: Tensor0[V])(using toTensor0: Conversion[Byte, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.arrayEqual(toTensor0(scalar), t) + + // broadcasting operators + def +![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Byte, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.add_!(toTensor0(scalar), t) + def -![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Byte, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.subtract_!(toTensor0(scalar), t) + def *![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Byte, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.multiply_!(toTensor0(scalar), t) + def %![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Byte, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.mod_!(toTensor0(scalar), t) + def /![T <: Tuple, V: IsFloating](t: Tensor[T, V])(using toTensor0: Conversion[Byte, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.divide_!(toTensor0(scalar), t) + def `![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Byte, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, Bool] = ElementWiseOps.greater_!(toTensor0(scalar), t) + def >=![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Byte, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, Bool] = ElementWiseOps.greaterEqual_!(toTensor0(scalar), t) + def ===![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Byte, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor0[Bool] = ElementWiseOps.arrayEqual_!(toTensor0(scalar), t) + + extension (scalar: Short) + + // operators with a Tensor0 + def +[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Short, Tensor0[V]]): Tensor0[V] = ElementWiseOps.add(toTensor0(scalar), t) + def -[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Short, Tensor0[V]]): Tensor0[V] = ElementWiseOps.subtract(toTensor0(scalar), t) + def *[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Short, Tensor0[V]]): Tensor0[V] = ElementWiseOps.multiply(toTensor0(scalar), t) + def %[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Short, Tensor0[V]]): Tensor0[V] = ElementWiseOps.mod(toTensor0(scalar), t) + def /[V: IsFloating](t: Tensor0[V])(using toTensor0: Conversion[Short, Tensor0[V]]): Tensor0[V] = ElementWiseOps.divide(toTensor0(scalar), t) + def <[V](t: Tensor0[V])(using toTensor0: Conversion[Short, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.less(toTensor0(scalar), t) + def <=[V](t: Tensor0[V])(using toTensor0: Conversion[Short, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.lessEqual(toTensor0(scalar), t) + def >[V](t: Tensor0[V])(using toTensor0: Conversion[Short, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.greater(toTensor0(scalar), t) + def >=[V](t: Tensor0[V])(using toTensor0: Conversion[Short, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.greaterEqual(toTensor0(scalar), t) + def ===[V](t: Tensor0[V])(using toTensor0: Conversion[Short, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.arrayEqual(toTensor0(scalar), t) + + // broadcasting operators + def +![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Short, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.add_!(toTensor0(scalar), t) + def -![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Short, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.subtract_!(toTensor0(scalar), t) + def *![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Short, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.multiply_!(toTensor0(scalar), t) + def %![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Short, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.mod_!(toTensor0(scalar), t) + def /![T <: Tuple, V: IsFloating](t: Tensor[T, V])(using toTensor0: Conversion[Short, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.divide_!(toTensor0(scalar), t) + def `![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Short, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, Bool] = ElementWiseOps.greater_!(toTensor0(scalar), t) + def >=![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Short, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, Bool] = ElementWiseOps.greaterEqual_!(toTensor0(scalar), t) + def ===![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Short, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor0[Bool] = ElementWiseOps.arrayEqual_!(toTensor0(scalar), t) + + extension (scalar: Int) + + // operators with a Tensor0 + def +[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Int, Tensor0[V]]): Tensor0[V] = ElementWiseOps.add(toTensor0(scalar), t) + def -[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Int, Tensor0[V]]): Tensor0[V] = ElementWiseOps.subtract(toTensor0(scalar), t) + def *[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Int, Tensor0[V]]): Tensor0[V] = ElementWiseOps.multiply(toTensor0(scalar), t) + def %[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Int, Tensor0[V]]): Tensor0[V] = ElementWiseOps.mod(toTensor0(scalar), t) + def /[V: IsFloating](t: Tensor0[V])(using toTensor0: Conversion[Int, Tensor0[V]]): Tensor0[V] = ElementWiseOps.divide(toTensor0(scalar), t) + def <[V](t: Tensor0[V])(using toTensor0: Conversion[Int, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.less(toTensor0(scalar), t) + def <=[V](t: Tensor0[V])(using toTensor0: Conversion[Int, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.lessEqual(toTensor0(scalar), t) + def >[V](t: Tensor0[V])(using toTensor0: Conversion[Int, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.greater(toTensor0(scalar), t) + def >=[V](t: Tensor0[V])(using toTensor0: Conversion[Int, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.greaterEqual(toTensor0(scalar), t) + def ===[V](t: Tensor0[V])(using toTensor0: Conversion[Int, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.arrayEqual(toTensor0(scalar), t) + + // broadcasting operators + def +![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Int, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.add_!(toTensor0(scalar), t) + def -![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Int, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.subtract_!(toTensor0(scalar), t) + def *![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Int, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.multiply_!(toTensor0(scalar), t) + def %![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Int, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.mod_!(toTensor0(scalar), t) + def /![T <: Tuple, V: IsFloating](t: Tensor[T, V])(using toTensor0: Conversion[Int, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.divide_!(toTensor0(scalar), t) + def `![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Int, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, Bool] = ElementWiseOps.greater_!(toTensor0(scalar), t) + def >=![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Int, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, Bool] = ElementWiseOps.greaterEqual_!(toTensor0(scalar), t) + def ===![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Int, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor0[Bool] = ElementWiseOps.arrayEqual_!(toTensor0(scalar), t) + + extension (scalar: Long) + + // operators with a Tensor0 + def +[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Long, Tensor0[V]]): Tensor0[V] = ElementWiseOps.add(toTensor0(scalar), t) + def -[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Long, Tensor0[V]]): Tensor0[V] = ElementWiseOps.subtract(toTensor0(scalar), t) + def *[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Long, Tensor0[V]]): Tensor0[V] = ElementWiseOps.multiply(toTensor0(scalar), t) + def %[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Long, Tensor0[V]]): Tensor0[V] = ElementWiseOps.mod(toTensor0(scalar), t) + def /[V: IsFloating](t: Tensor0[V])(using toTensor0: Conversion[Long, Tensor0[V]]): Tensor0[V] = ElementWiseOps.divide(toTensor0(scalar), t) + def <[V](t: Tensor0[V])(using toTensor0: Conversion[Long, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.less(toTensor0(scalar), t) + def <=[V](t: Tensor0[V])(using toTensor0: Conversion[Long, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.lessEqual(toTensor0(scalar), t) + def >[V](t: Tensor0[V])(using toTensor0: Conversion[Long, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.greater(toTensor0(scalar), t) + def >=[V](t: Tensor0[V])(using toTensor0: Conversion[Long, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.greaterEqual(toTensor0(scalar), t) + def ===[V](t: Tensor0[V])(using toTensor0: Conversion[Long, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.arrayEqual(toTensor0(scalar), t) + + // broadcasting operators + def +![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Long, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.add_!(toTensor0(scalar), t) + def -![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Long, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.subtract_!(toTensor0(scalar), t) + def *![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Long, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.multiply_!(toTensor0(scalar), t) + def %![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Long, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.mod_!(toTensor0(scalar), t) + def /![T <: Tuple, V: IsFloating](t: Tensor[T, V])(using toTensor0: Conversion[Long, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.divide_!(toTensor0(scalar), t) + def `![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Long, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, Bool] = ElementWiseOps.greater_!(toTensor0(scalar), t) + def >=![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Long, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, Bool] = ElementWiseOps.greaterEqual_!(toTensor0(scalar), t) + def ===![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Long, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor0[Bool] = ElementWiseOps.arrayEqual_!(toTensor0(scalar), t) + + extension (scalar: Float) + + // operators with a Tensor0 + def +[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Float, Tensor0[V]]): Tensor0[V] = ElementWiseOps.add(toTensor0(scalar), t) + def -[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Float, Tensor0[V]]): Tensor0[V] = ElementWiseOps.subtract(toTensor0(scalar), t) + def *[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Float, Tensor0[V]]): Tensor0[V] = ElementWiseOps.multiply(toTensor0(scalar), t) + def %[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Float, Tensor0[V]]): Tensor0[V] = ElementWiseOps.mod(toTensor0(scalar), t) + def /[V: IsFloating](t: Tensor0[V])(using toTensor0: Conversion[Float, Tensor0[V]]): Tensor0[V] = ElementWiseOps.divide(toTensor0(scalar), t) + def <[V](t: Tensor0[V])(using toTensor0: Conversion[Float, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.less(toTensor0(scalar), t) + def <=[V](t: Tensor0[V])(using toTensor0: Conversion[Float, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.lessEqual(toTensor0(scalar), t) + def >[V](t: Tensor0[V])(using toTensor0: Conversion[Float, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.greater(toTensor0(scalar), t) + def >=[V](t: Tensor0[V])(using toTensor0: Conversion[Float, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.greaterEqual(toTensor0(scalar), t) + def ===[V](t: Tensor0[V])(using toTensor0: Conversion[Float, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.arrayEqual(toTensor0(scalar), t) + + // broadcasting operators + def +![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Float, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.add_!(toTensor0(scalar), t) + def -![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Float, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.subtract_!(toTensor0(scalar), t) + def *![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Float, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.multiply_!(toTensor0(scalar), t) + def %![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Float, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.mod_!(toTensor0(scalar), t) + def /![T <: Tuple, V: IsFloating](t: Tensor[T, V])(using toTensor0: Conversion[Float, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.divide_!(toTensor0(scalar), t) + def `![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Float, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, Bool] = ElementWiseOps.greater_!(toTensor0(scalar), t) + def >=![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Float, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, Bool] = ElementWiseOps.greaterEqual_!(toTensor0(scalar), t) + def ===![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Float, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor0[Bool] = ElementWiseOps.arrayEqual_!(toTensor0(scalar), t) + + extension (scalar: Double) + + // operators with a Tensor0 + def +[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Double, Tensor0[V]]): Tensor0[V] = ElementWiseOps.add(toTensor0(scalar), t) + def -[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Double, Tensor0[V]]): Tensor0[V] = ElementWiseOps.subtract(toTensor0(scalar), t) + def *[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Double, Tensor0[V]]): Tensor0[V] = ElementWiseOps.multiply(toTensor0(scalar), t) + def %[V: IsNumber](t: Tensor0[V])(using toTensor0: Conversion[Double, Tensor0[V]]): Tensor0[V] = ElementWiseOps.mod(toTensor0(scalar), t) + def /[V: IsFloating](t: Tensor0[V])(using toTensor0: Conversion[Double, Tensor0[V]]): Tensor0[V] = ElementWiseOps.divide(toTensor0(scalar), t) + def <[V](t: Tensor0[V])(using toTensor0: Conversion[Double, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.less(toTensor0(scalar), t) + def <=[V](t: Tensor0[V])(using toTensor0: Conversion[Double, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.lessEqual(toTensor0(scalar), t) + def >[V](t: Tensor0[V])(using toTensor0: Conversion[Double, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.greater(toTensor0(scalar), t) + def >=[V](t: Tensor0[V])(using toTensor0: Conversion[Double, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.greaterEqual(toTensor0(scalar), t) + def ===[V](t: Tensor0[V])(using toTensor0: Conversion[Double, Tensor0[V]]): Tensor0[Bool] = ElementWiseOps.arrayEqual(toTensor0(scalar), t) + + // broadcasting operators + def +![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Double, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.add_!(toTensor0(scalar), t) + def -![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Double, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.subtract_!(toTensor0(scalar), t) + def *![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Double, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.multiply_!(toTensor0(scalar), t) + def %![T <: Tuple, V: IsNumber](t: Tensor[T, V])(using toTensor0: Conversion[Double, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.mod_!(toTensor0(scalar), t) + def /![T <: Tuple, V: IsFloating](t: Tensor[T, V])(using toTensor0: Conversion[Double, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = ElementWiseOps.divide_!(toTensor0(scalar), t) + def `![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Double, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, Bool] = ElementWiseOps.greater_!(toTensor0(scalar), t) + def >=![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Double, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, Bool] = ElementWiseOps.greaterEqual_!(toTensor0(scalar), t) + def ===![T <: Tuple, V](t: Tensor[T, V])(using toTensor0: Conversion[Double, Tensor0[V]], bc: Broadcast[EmptyTuple, T, V]): Tensor0[Bool] = ElementWiseOps.arrayEqual_!(toTensor0(scalar), t) diff --git a/core/src/main/scala/dimwit/tensor/ValueOps.scala b/core/src/main/scala/dimwit/tensor/ValueOps.scala deleted file mode 100644 index f15a5992..00000000 --- a/core/src/main/scala/dimwit/tensor/ValueOps.scala +++ /dev/null @@ -1,52 +0,0 @@ -package dimwit.tensor - -import dimwit.tensor.TensorOps.IsFloating -import dimwit.tensor.TensorOps.IsNumber -import dimwit.tensor.DType.Bool -import dimwit.tensor.tensorops.ElementWiseOps.add -import dimwit.tensor.tensorops.ElementWiseOps.divide -import dimwit.tensor.tensorops.ElementWiseOps.equal -import dimwit.tensor.tensorops.ElementWiseOps.greater -import dimwit.tensor.tensorops.ElementWiseOps.greaterEqual -import dimwit.tensor.tensorops.ElementWiseOps.less -import dimwit.tensor.tensorops.ElementWiseOps.lessEqual -import dimwit.tensor.tensorops.ElementWiseOps.multiply -import dimwit.tensor.tensorops.ElementWiseOps.subtract -import dimwit.tensor.tensorops.TensorOpsUtil.Broadcast - -object ValueOps: - - extension [V: IsNumber](t: Tensor0[V]) - - def +(t2: Tensor0[V]): Tensor0[V] = TensorOps.add(t, t2) - def -(t2: Tensor0[V]): Tensor0[V] = TensorOps.subtract(t, t2) - def *(t2: Tensor0[V]): Tensor0[V] = TensorOps.multiply(t, t2) - - extension [V: IsFloating](t: Tensor0[V]) - - def /(scalar: Tensor0[V]): Tensor0[V] = TensorOps.divide(t, scalar) - - extension (scalar: Float) - - def +[V: IsNumber](t: Tensor0[V]): Tensor0[V] = add(Tensor0.likeDType(t)(scalar), t) - def +![T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V])(using bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = bc.applyTo(Tensor0.likeDType(t)(scalar), t)(add) - - def -[V: IsNumber](t: Tensor0[V]): Tensor0[V] = subtract(Tensor0.likeDType(t)(scalar), t) - def -![T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V])(using bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = bc.applyTo(Tensor0.likeDType(t)(scalar), t)(subtract) - - def *[V: IsNumber](t: Tensor0[V]): Tensor0[V] = multiply(Tensor0.likeDType(t)(scalar), t) - def *![T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V])(using bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = bc.applyTo(Tensor0.likeDType(t)(scalar), t)(multiply) - - // Comparing a scalar against a tensor. `![T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V])(using bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, Bool] = bc.applyTo(Tensor0.likeDType(t)(scalar), t)(greater) - def >=![T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V])(using bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, Bool] = bc.applyTo(Tensor0.likeDType(t)(scalar), t)(greaterEqual) - def elementEquals_![T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V])(using bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, Bool] = bc.applyTo(Tensor0.likeDType(t)(scalar), t)(equal) - - extension (scalar: Float) - - def /[V: IsFloating](t: Tensor0[V]): Tensor0[V] = divide(Tensor0.likeDType(t)(scalar), t) - def /![T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V])(using bc: Broadcast[EmptyTuple, T, V]): Tensor[bc.Out, V] = bc.applyTo(Tensor0.likeDType(t)(scalar), t)(divide) diff --git a/core/src/main/scala/dimwit/tensor/TensorOps.scala b/core/src/main/scala/dimwit/tensor/ValueTypeClasses.scala similarity index 72% rename from core/src/main/scala/dimwit/tensor/TensorOps.scala rename to core/src/main/scala/dimwit/tensor/ValueTypeClasses.scala index e05fa51c..023f411a 100644 --- a/core/src/main/scala/dimwit/tensor/TensorOps.scala +++ b/core/src/main/scala/dimwit/tensor/ValueTypeClasses.scala @@ -1,18 +1,11 @@ package dimwit.tensor import dimwit.jax.Jax -import dimwit.tensor.HasScalar -import dimwit.tensor.Label -import dimwit.tensor.Labels -import dimwit.tensor.ShapeTypeHelpers.* -import dimwit.tensor.TupleHelpers.* import scala.annotation.implicitNotFound -import scala.annotation.targetName -object TensorOps: - - import dimwit.tensor.tensorops.TensorOpsUtil.* +/** Type classes on the value type `V` of a `Tensor[T, V]`, e.g. `IsFloating[V]` for floating point tensors. */ +object ValueTypeClasses: /** Typeclass to map a type V to its corresponding DType. */ @@ -58,17 +51,3 @@ object TensorOps: object IsBoolean: def apply[V](using ev: IsBoolean[V]): IsBoolean[V] = ev - - export tensorops.ElementWiseOps.* - export tensorops.ReductionOps.* - export tensorops.AlongAxisOps.* - export tensorops.ContractionOps.* - export tensorops.ConvolutionOps.* - export tensorops.LinearAlgebraOps.* - export tensorops.StructuralOps.* - export tensorops.FunctionalOps.* - - export tensorops.Tensor0Ops.* - export tensorops.Tensor1Ops.* - export tensorops.Tensor2Ops.* - export tensorops.Tensor3Ops.* diff --git a/core/src/main/scala/dimwit/tensor/tensorops/AlongAxisOps.scala b/core/src/main/scala/dimwit/tensor/tensorops/AlongAxisOps.scala index 3736bc53..2fe8bac8 100644 --- a/core/src/main/scala/dimwit/tensor/tensorops/AlongAxisOps.scala +++ b/core/src/main/scala/dimwit/tensor/tensorops/AlongAxisOps.scala @@ -6,65 +6,160 @@ import dimwit.tensor.DType.Int32 import dimwit.tensor.Label import dimwit.tensor.Labels import dimwit.tensor.ShapeTypeHelpers.AxisIndex -import dimwit.tensor.ShapeTypeHelpers.AxisIndices -import dimwit.tensor.ShapeTypeHelpers.UnwrapAxes import dimwit.tensor.Tensor import dimwit.tensor.Tensor1 -import dimwit.tensor.TensorOps.IsFloating -import dimwit.tensor.TensorOps.IsNumber -import me.shadaj.scalapy.py.SeqConverters -import dimwit.tensor.tensorops.FunctionalOps.vapply +import dimwit.tensor.ValueTypeClasses.IsFloating +import dimwit.tensor.ValueTypeClasses.IsNumber +import dimwit.tensor.tensorops.FunctionalExtensions.vapply -/** Operations along an axis that keep all axes of the tensor (unlike reductions, which remove them). */ -object AlongAxisOps: +/** Functions on a Tensor1 that keep its axis, exported as e.g. `Tensor1.softmax`. */ +private[dimwit] object AlongAxisTensor1Ops: + + // --------------------------------------------------------- + // Functions from a Tensor1 to a Tensor1. + // They are lifted to any tensor shape with `vapply` by the functions in [[AlongAxisOps]]. + // --------------------------------------------------------- + + /** sorts the Tensor1 `t`. */ + def sort[L: Label, V: IsNumber](t: Tensor1[L, V]): Tensor1[L, V] = + Tensor(Jax.jnp.sort(t.jaxValue, axis = 0)) + + /** returns the indices that would sort the Tensor1 `t`. */ + def argsort[L: Label, V: IsNumber](t: Tensor1[L, V]): Tensor1[L, Int32] = + Tensor(Jax.jnp.argsort(t.jaxValue, axis = 0)) + + /** computes the cumulative sum of the Tensor1 `t`. */ + def cumsum[L: Label, V: IsNumber](t: Tensor1[L, V]): Tensor1[L, V] = + Tensor(Jax.jnp.cumsum(t.jaxValue, axis = 0)) + + /** computes the cumulative product of the Tensor1 `t`. */ + def cumprod[L: Label, V: IsNumber](t: Tensor1[L, V]): Tensor1[L, V] = + Tensor(Jax.jnp.cumprod(t.jaxValue, axis = 0)) + + /** computes the cumulative maximum of the Tensor1 `t`. */ + def cummax[L: Label, V: IsNumber](t: Tensor1[L, V]): Tensor1[L, V] = + Tensor(Jax.lax.cummax(t.jaxValue, axis = 0)) + + /** computes the cumulative minimum of the Tensor1 `t`. */ + def cummin[L: Label, V: IsNumber](t: Tensor1[L, V]): Tensor1[L, V] = + Tensor(Jax.lax.cummin(t.jaxValue, axis = 0)) + + /** computes the cumulative log-sum-exp of the Tensor1 `t`, i.e. a numerically stable `log(cumsum(exp(t)))`. */ + def logcumsumexp[L: Label, V: IsFloating](t: Tensor1[L, V]): Tensor1[L, V] = + Tensor(Jax.lax.cumlogsumexp(t.jaxValue, axis = 0)) + + /** computes the discrete difference of the Tensor1 `t`, reducing its size by one. */ + def diff[L: Label, V: IsNumber](t: Tensor1[L, V]): Tensor1[L, V] = + Tensor(Jax.jnp.diff(t.jaxValue, axis = 0)) + + /** rolls the elements of the Tensor1 `t` by `shift` positions; elements shifted beyond the end re-appear at the start. */ + def roll[L: Label, V](shift: Int)(t: Tensor1[L, V]): Tensor1[L, V] = + Tensor(Jax.jnp.roll(t.jaxValue, shift = shift, axis = 0)) + + /** computes the softmax of the Tensor1 `t`. */ + def softmax[L: Label, V: IsFloating](t: Tensor1[L, V]): Tensor1[L, V] = + Tensor(Jax.jnn.softmax(t.jaxValue, axis = 0)) + + /** computes the log of the softmax of the Tensor1 `t`, more stable than `softmax(t).log`. */ + def logSoftmax[L: Label, V: IsFloating](t: Tensor1[L, V]): Tensor1[L, V] = + Tensor(Jax.jnn.log_softmax(t.jaxValue, axis = 0)) + +/** Operations along an axis that keep all axes of the tensor (unlike reductions, which remove them), exported as e.g. `Tensor.softmax(t, Axis[A])`. */ +private[dimwit] object AlongAxisOps: + + /** rolls the elements of `t` along the specified axis by `shift` positions; elements shifted beyond the end re-appear at the start. */ + def roll[T <: Tuple: Labels, V, L: Label](t: Tensor[T, V], axis: Axis[L], shift: Int)(using AxisIndex[T, L]): Tensor[T, V] = + t.vapply(axis)(AlongAxisTensor1Ops.roll(shift)) + + /** returns the indices that would sort `t` along the specified axis. */ + def argsort[T <: Tuple: Labels, V: IsNumber, L: Label](t: Tensor[T, V], axis: Axis[L])(using AxisIndex[T, L]): Tensor[T, Int32] = + t.vapply(axis)(AlongAxisTensor1Ops.argsort) + + /** sorts `t` along the specified axis. */ + def sort[T <: Tuple: Labels, V: IsNumber, L: Label](t: Tensor[T, V], axis: Axis[L])(using AxisIndex[T, L]): Tensor[T, V] = + t.vapply(axis)(AlongAxisTensor1Ops.sort) + + /** computes the cumulative sum of `t` along the specified axis. */ + def cumsum[T <: Tuple: Labels, V: IsNumber, L: Label](t: Tensor[T, V], axis: Axis[L])(using AxisIndex[T, L]): Tensor[T, V] = + t.vapply(axis)(AlongAxisTensor1Ops.cumsum) + + /** computes the cumulative product of `t` along the specified axis. */ + def cumprod[T <: Tuple: Labels, V: IsNumber, L: Label](t: Tensor[T, V], axis: Axis[L])(using AxisIndex[T, L]): Tensor[T, V] = + t.vapply(axis)(AlongAxisTensor1Ops.cumprod) + + /** computes the cumulative maximum of `t` along the specified axis. */ + def cummax[T <: Tuple: Labels, V: IsNumber, L: Label](t: Tensor[T, V], axis: Axis[L])(using AxisIndex[T, L]): Tensor[T, V] = + t.vapply(axis)(AlongAxisTensor1Ops.cummax) + + /** computes the cumulative minimum of `t` along the specified axis. */ + def cummin[T <: Tuple: Labels, V: IsNumber, L: Label](t: Tensor[T, V], axis: Axis[L])(using AxisIndex[T, L]): Tensor[T, V] = + t.vapply(axis)(AlongAxisTensor1Ops.cummin) + + /** computes the discrete difference of `t` along the specified axis, reducing that axis' size by one. */ + def diff[T <: Tuple: Labels, V: IsNumber, L: Label](t: Tensor[T, V], axis: Axis[L])(using AxisIndex[T, L]): Tensor[T, V] = + t.vapply(axis)(AlongAxisTensor1Ops.diff) + + /** computes the cumulative log-sum-exp of `t` along the specified axis. */ + def logcumsumexp[T <: Tuple: Labels, V: IsFloating, L: Label](t: Tensor[T, V], axis: Axis[L])(using AxisIndex[T, L]): Tensor[T, V] = + t.vapply(axis)(AlongAxisTensor1Ops.logcumsumexp) + + /** computes the softmax of `t` along the specified axis. */ + def softmax[T <: Tuple: Labels, V: IsFloating, L: Label](t: Tensor[T, V], axis: Axis[L])(using AxisIndex[T, L]): Tensor[T, V] = + t.vapply(axis)(AlongAxisTensor1Ops.softmax) + + /** computes the log-softmax of `t` along the specified axis. */ + def logSoftmax[T <: Tuple: Labels, V: IsFloating, L: Label](t: Tensor[T, V], axis: Axis[L])(using AxisIndex[T, L]): Tensor[T, V] = + t.vapply(axis)(AlongAxisTensor1Ops.logSoftmax) + +/** Extension methods for operations along an axis, e.g. `t.softmax(Axis[A])`. */ +private[dimwit] object AlongAxisExtensions: extension [T <: Tuple: Labels, V](t: Tensor[T, V]) /** rolls the elements of `t` along the specified axis by `shift` positions. */ def roll[L: Label](axis: Axis[L], shift: Int)(using AxisIndex[T, L]): Tensor[T, V] = - t.vapply(axis)(Tensor1.roll(shift)) + AlongAxisOps.roll(t, axis, shift) extension [T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V]) - /** Returns a tensor of indices that would sort `t` along the specified axes */ - def argsort[Inputs <: Tuple](axes: Inputs)(using ev: AxisIndices[T, UnwrapAxes[Inputs]]): Tensor[T, Int32] = Tensor(Jax.jnp.argsort(t.jaxValue, axis = ev.indices.toPythonProxy)) + /** Returns a tensor of indices that would sort `t` along the specified axis */ def argsort[L: Label](axis: Axis[L])(using AxisIndex[T, L]): Tensor[T, Int32] = - t.vapply(axis)(Tensor1.argsort) + AlongAxisOps.argsort(t, axis) /** sorts the tensor `t` along the specified axis */ def sort[L: Label](axis: Axis[L])(using AxisIndex[T, L]): Tensor[T, V] = - t.vapply(axis)(Tensor1.sort) + AlongAxisOps.sort(t, axis) /** computes the cumulative sum of the tensor `t` along the specified axis. */ def cumsum[L: Label](axis: Axis[L])(using AxisIndex[T, L]): Tensor[T, V] = - t.vapply(axis)(Tensor1.cumsum) + AlongAxisOps.cumsum(t, axis) /** computes the cumulative product of the tensor `t` along the specified axis. */ def cumprod[L: Label](axis: Axis[L])(using AxisIndex[T, L]): Tensor[T, V] = - t.vapply(axis)(Tensor1.cumprod) + AlongAxisOps.cumprod(t, axis) /** computes the cumulative maximum of the tensor `t` along the specified axis. */ def cummax[L: Label](axis: Axis[L])(using AxisIndex[T, L]): Tensor[T, V] = - t.vapply(axis)(Tensor1.cummax) + AlongAxisOps.cummax(t, axis) /** computes the cumulative minimum of the tensor `t` along the specified axis. */ def cummin[L: Label](axis: Axis[L])(using AxisIndex[T, L]): Tensor[T, V] = - t.vapply(axis)(Tensor1.cummin) + AlongAxisOps.cummin(t, axis) /** computes the discrete difference of the tensor `t` along the specified axis, reducing that axis' size by one. */ def diff[L: Label](axis: Axis[L])(using AxisIndex[T, L]): Tensor[T, V] = - t.vapply(axis)(Tensor1.diff) + AlongAxisOps.diff(t, axis) extension [T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]) /** computes the cumulative log-sum-exp of the tensor `t` along the specified axis. */ def logcumsumexp[L: Label](axis: Axis[L])(using AxisIndex[T, L]): Tensor[T, V] = - t.vapply(axis)(Tensor1.logcumsumexp) + AlongAxisOps.logcumsumexp(t, axis) /** computes the softmax of `t` along the specified axis. */ def softmax[L: Label](axis: Axis[L])(using AxisIndex[T, L]): Tensor[T, V] = - t.vapply(axis)(Tensor1.softmax) + AlongAxisOps.softmax(t, axis) /** computes the log-softmax of `t` along the specified axis. */ def logSoftmax[L: Label](axis: Axis[L])(using AxisIndex[T, L]): Tensor[T, V] = - t.vapply(axis)(Tensor1.logSoftmax) + AlongAxisOps.logSoftmax(t, axis) diff --git a/core/src/main/scala/dimwit/tensor/tensorops/ContractionOps.scala b/core/src/main/scala/dimwit/tensor/tensorops/ContractionOps.scala index 33bce6e2..25c9c407 100644 --- a/core/src/main/scala/dimwit/tensor/tensorops/ContractionOps.scala +++ b/core/src/main/scala/dimwit/tensor/tensorops/ContractionOps.scala @@ -12,10 +12,60 @@ import me.shadaj.scalapy.readwrite.Writer import scala.annotation.targetName -/** Provides extension methods for tensor contraction operations, - * including outer products and dot products. - */ -object ContractionOps: +/** Tensor contraction operations: outer products and dot products. */ +private[dimwit] object ContractionOps: + + /** Computes the outer product of `t1` and `t2`. + * Automatically primes the labels of the resulting tensor to avoid label collisions. + */ + def outerProduct[T <: Tuple: Labels, OtherShape <: Tuple: Labels, V](t1: Tensor[T, V], t2: Tensor[OtherShape, V])(using + primeConcat: PrimeConcat[T, OtherShape], + labels: Labels[primeConcat.Out] + ): Tensor[primeConcat.Out, V] = Tensor( + Jax.jnp.tensordot(t1.jaxValue, t2.jaxValue, axes = 0) // generalized outer product + ) + + /** Computes the dot product of `t1` and `t2` along the specified axis. + * The axis must be present in both tensors and will be contracted (removed) from the resulting tensor. + * + * @param axis The axis along which to contract. Must be present in both tensors. + */ + def dot[T <: Tuple, ContractAxis, OtherShape <: Tuple, V](t1: Tensor[T, V], axis: Axis[ContractAxis], t2: Tensor[OtherShape, V])(using + ev: AxisRemover[T, ContractAxis], + evOther: AxisRemover[OtherShape, ContractAxis] + )(using + primeConcat: PrimeConcat[ev.RemainingAxes, evOther.RemainingAxes], + labelsOut: Labels[primeConcat.Out] + ): Tensor[primeConcat.Out, V] = + tensordot(t1, ev.index, t2, evOther.index) + + /** Computes the dot product of `t1` and `t2` along the specified pair of axes. + * The axes must be present in their respective tensors and will be contracted (removed) from the resulting tensor. + * + * @param axisPair The pair of axes along which to contract. Each axis must be present in its respective tensor. + */ + @targetName("dotOn") + def dot[T <: Tuple, ContractAxisA, ContractAxisB, OtherShape <: Tuple, V]( + t1: Tensor[T, V], + axisPair: (Axis[ContractAxisA], Axis[ContractAxisB]), + t2: Tensor[OtherShape, V] + )(using + ev: AxisRemover[T, ContractAxisA], + evOther: AxisRemover[OtherShape, ContractAxisB] + )(using + primeConcat: PrimeConcat[ev.RemainingAxes, evOther.RemainingAxes], + outLabels: Labels[primeConcat.Out] + ): Tensor[primeConcat.Out, V] = + tensordot(t1, ev.index, t2, evOther.index) + + private def tensordot[Out <: Tuple: Labels, V](t1: Tensor[?, V], index1: Int, t2: Tensor[?, V], index2: Int): Tensor[Out, V] = + val axesTuple1 = Jax.Dynamic.global.tuple(Seq(index1).toPythonProxy) + val axesTuple2 = Jax.Dynamic.global.tuple(Seq(index2).toPythonProxy) + val axesPair = Jax.Dynamic.global.tuple(Seq(axesTuple1, axesTuple2).toPythonProxy) + Tensor(Jax.jnp.tensordot(t1.jaxValue, t2.jaxValue, axes = axesPair)) + +/** Extension methods for tensor contraction operations, e.g. `t1.dot(Axis[A])(t2)`. */ +private[dimwit] object ContractionExtensions: extension [T <: Tuple: Labels, V](tensor: Tensor[T, V]) @@ -25,9 +75,7 @@ object ContractionOps: def outerProduct[OtherShape <: Tuple: Labels](other: Tensor[OtherShape, V])(using primeConcat: PrimeConcat[T, OtherShape], labels: Labels[primeConcat.Out] - ): Tensor[primeConcat.Out, V] = Tensor( - Jax.jnp.tensordot(tensor.jaxValue, other.jaxValue, axes = 0) // generalized outer product - ) + ): Tensor[primeConcat.Out, V] = ContractionOps.outerProduct(tensor, other) /** Computes the dot product of this tensor with another tensor along the specified axis. * The axis must be present in both tensors and will be contracted (removed) from the resulting tensor. @@ -44,12 +92,7 @@ object ContractionOps: )(using primeConcat: PrimeConcat[ev.RemainingAxes, evOther.RemainingAxes], labelsOut: Labels[primeConcat.Out] - ): Tensor[primeConcat.Out, V] = - val axesTuple1 = Jax.Dynamic.global.tuple(Seq(ev.index).toPythonProxy) - val axesTuple2 = Jax.Dynamic.global.tuple(Seq(evOther.index).toPythonProxy) - val axesPair = Jax.Dynamic.global.tuple(Seq(axesTuple1, axesTuple2).toPythonProxy) - - Tensor(Jax.jnp.tensordot(tensor.jaxValue, other.jaxValue, axes = axesPair)) + ): Tensor[primeConcat.Out, V] = ContractionOps.dot(tensor, axis, other) /** Computes the dot product of this tensor with another tensor along the specified pair of axes. * The axes must be present in their respective tensors and will be contracted (removed) from the resulting tensor. @@ -68,9 +111,4 @@ object ContractionOps: )(using primeConcat: PrimeConcat[ev.RemainingAxes, evOther.RemainingAxes], outLabels: Labels[primeConcat.Out] - ): Tensor[primeConcat.Out, V] = - val axesTuple1 = Jax.Dynamic.global.tuple(Seq(ev.index).toPythonProxy) - val axesTuple2 = Jax.Dynamic.global.tuple(Seq(evOther.index).toPythonProxy) - val axesPair = Jax.Dynamic.global.tuple(Seq(axesTuple1, axesTuple2).toPythonProxy) - - Tensor(Jax.jnp.tensordot(tensor.jaxValue, other.jaxValue, axes = axesPair)) + ): Tensor[primeConcat.Out, V] = ContractionOps.dot(tensor, axisPair, other) diff --git a/core/src/main/scala/dimwit/tensor/tensorops/ConvolutionOps.scala b/core/src/main/scala/dimwit/tensor/tensorops/ConvolutionOps.scala index d98b7549..d97ab338 100644 --- a/core/src/main/scala/dimwit/tensor/tensorops/ConvolutionOps.scala +++ b/core/src/main/scala/dimwit/tensor/tensorops/ConvolutionOps.scala @@ -8,29 +8,236 @@ import dimwit.tensor.Labels import dimwit.tensor.ShapeTypeHelpers.AxisIndex import dimwit.tensor.Tensor import dimwit.tensor.Tensor3 -import dimwit.tensor.TensorOps.IsFloating -import dimwit.tensor.TensorOps.swap +import dimwit.tensor.ValueTypeClasses.IsFloating +import dimwit.tensor.Padding +import dimwit.tensor.Stride1 +import dimwit.tensor.Stride2 +import dimwit.tensor.Stride3 +import dimwit.tensor.tensorops.StructuralExtensions.swap import me.shadaj.scalapy.py import me.shadaj.scalapy.py.SeqConverters import me.shadaj.scalapy.readwrite.Writer -/** Provides extension methods for convolution operations on tensors. +/** Convolution operations on tensors. * Convolution operations are restricted to 1D, 2D and 3D convolutions, * and support both standard and transposed convolutions. */ -object ConvolutionOps: +private[dimwit] object ConvolutionOps: - /** Padding options for convolution operations. - * SAME: Output size is the same as input size (with appropriate padding). - * VALID: No padding, output size is reduced based on kernel size. + /** Computes the 1D convolution of `input` with the specified kernel tensor. * - * Refer to JAX documentation for more details on padding behavior. - * https://jax.readthedocs.io/en/latest/_autosummary/jax.lax.conv_general_dilated.html + * @param kernel - The convolution kernel + * @param stride - Stride for the convolution. + * @param padding - Padding mode for the convolution. + * @return A new tensor representing the result of the convolution operation. */ - enum Padding: - case SAME, VALID + def conv1d[S1: Label, InChannel: Label, V: IsFloating, OutChannel: Label]( + input: Tensor[S1 *: InChannel *: EmptyTuple, V], + kernel: Tensor3[S1, InChannel, OutChannel, V], + stride: Stride1[S1] | Int = 1, + padding: Padding = Padding.SAME + ): Tensor[S1 *: OutChannel *: EmptyTuple, V] = + require( + input.shape(Axis[InChannel]) == kernel.shape(Axis[InChannel]), + s"Input channels mismatch: input has ${input.shape(Axis[InChannel])} channels, kernel expects ${kernel.shape(Axis[InChannel])} channels" + ) + val strides = stride match + case s: Int => Seq(s) + case ae: AxisExtent[S1] => Seq(ae.size) + // JAX requires input and kernel to have same rank, so we must add (and remove) dummy dim to input. + val batchInput = Jax.jnp.expand_dims(input.jaxValue, axis = 0) // add dummy dim + val convResult = Jax.lax.conv_general_dilated( + lhs = batchInput, + rhs = kernel.jaxValue, + window_strides = strides.toPythonProxy, + padding = padding.toString, + dimension_numbers = py.Dynamic.global.tuple(Seq("NHC", "HIO", "NHC").toPythonProxy) + ) + val unbatchedRes = Jax.jnp.squeeze(convResult, axis = 0) // remove dummy dim + Tensor(unbatchedRes) - type Stride1[S1] = AxisExtent[S1] + /** Computes the transposed 1D convolution of + * `input` with the specified kernel tensor. + * @param kernel - The convolution kernel + * @param stride - Stride for the convolution. + * @param padding - Padding mode for the convolution. + * @return A new tensor representing the result of the transposed convolution operation. + */ + def transposeConv1d[S1: Label, OutChannel: Label, V: IsFloating, InChannel: Label]( + input: Tensor[S1 *: OutChannel *: EmptyTuple, V], + kernel: Tensor[S1 *: InChannel *: OutChannel *: EmptyTuple, V], + stride: Stride1[S1] | Int = 1, + padding: Padding = Padding.SAME + ): Tensor[S1 *: InChannel *: EmptyTuple, V] = + require( + input.shape(Axis[OutChannel]) == kernel.shape(Axis[OutChannel]), + s"Input channels mismatch: input has ${input.shape(Axis[OutChannel])} channels (OutChannel), kernel expects ${kernel.shape(Axis[OutChannel])}" + ) + val strides = stride match + case s: Int => Seq(s) + case ex: AxisExtent[S1] => Seq(ex.size) + + // kernel -> kernal adjoint: swap in/out channels and flip spatial dims + var kernelAdjoint = kernel.swap(Axis[InChannel], Axis[OutChannel]).jaxValue + kernelAdjoint = Jax.jnp.flip(kernelAdjoint, axis = 0) // flip S1 + + val batchInput = Jax.jnp.expand_dims(input.jaxValue, axis = 0) // add dummy dim + val convResult = Jax.lax.conv_transpose( + lhs = batchInput, + rhs = kernelAdjoint, + strides = strides.toPythonProxy, + padding = padding.toString, + dimension_numbers = py.Dynamic.global.tuple(Seq("NHC", "HIO", "NHC").toPythonProxy) + ) + val unbatchedRes = Jax.jnp.squeeze(convResult, axis = 0) // remove dummy dim + Tensor(unbatchedRes) + + /** Computes the 2D convolution of `input` with the specified kernel tensor. + * + * @param kernel - The convolution kernel tensor with shape (S1, S2, InChannel, OutChannel). + * @param stride - Stride for the convolution. + * @param padding - Padding mode for the convolution. + * @return A new tensor representing the result of the convolution operation. + */ + def conv2d[S1: Label, S2: Label, InChannel: Label, V: IsFloating, OutChannel: Label]( + input: Tensor[S1 *: S2 *: InChannel *: EmptyTuple, V], + kernel: Tensor[S1 *: S2 *: InChannel *: OutChannel *: EmptyTuple, V], + stride: Stride2[S1, S2] | Int = 1, + padding: Padding = Padding.SAME + ): Tensor[S1 *: S2 *: OutChannel *: EmptyTuple, V] = + require( + input.shape(Axis[InChannel]) == kernel.shape(Axis[InChannel]), + s"Input channels mismatch: input has ${input.shape(Axis[InChannel])} channels, kernel expects ${kernel.shape(Axis[InChannel])} channels" + ) + val strides = stride match + case s: Int => Seq(s, s) + case (ae1, ae2) => Seq(ae1.size, ae2.size) + // JAX requires input and kernel to have same rank, so we must add (and remove) dummy dim to input. + val batchInput = Jax.jnp.expand_dims(input.jaxValue, axis = 0) // add dummy dim + val convResult = Jax.lax.conv_general_dilated( + lhs = batchInput, + rhs = kernel.jaxValue, + window_strides = strides.toPythonProxy, + padding = padding.toString, + dimension_numbers = py.Dynamic.global.tuple(Seq("NHWC", "HWIO", "NHWC").toPythonProxy) + ) + val unbatchedRes = Jax.jnp.squeeze(convResult, axis = 0) // remove dummy dim + Tensor(unbatchedRes) + + /** Computes the transposed 2D convolution of `input` with the specified kernel tensor. + * + * @param kernel - The convolution kernel tensor with shape (S1, S2, InChannel, OutChannel). + * @param stride - Stride for the convolution. + * @param padding - Padding mode for the convolution. + * @return A new tensor representing the result of the transposed convolution operation. + */ + def transposeConv2d[S1: Label, S2: Label, OutChannel: Label, V: IsFloating, InChannel: Label]( + input: Tensor[S1 *: S2 *: OutChannel *: EmptyTuple, V], + kernel: Tensor[S1 *: S2 *: InChannel *: OutChannel *: EmptyTuple, V], + stride: Stride2[S1, S2] | Int = 1, + padding: Padding = Padding.SAME + ): Tensor[S1 *: S2 *: InChannel *: EmptyTuple, V] = + require( + input.shape(Axis[OutChannel]) == kernel.shape(Axis[OutChannel]), + s"Input channels mismatch: input has ${input.shape(Axis[OutChannel])} channels (OutChannel), kernel expects ${kernel.shape(Axis[OutChannel])}" + ) + + // JAX requires input and kernel to have same rank. Add dummy batch dim if needed. + val strides = stride match + case s: Int => Seq(s, s) + case (ae1, ae2) => Seq(ae1.size, ae2.size) + + // kernel -> kernal adjoint: swap in/out channels and flip spatial dims + var kernelAdjoint = kernel.swap(Axis[InChannel], Axis[OutChannel]).jaxValue + kernelAdjoint = Jax.jnp.flip(kernelAdjoint, axis = 0) // flip S1 + kernelAdjoint = Jax.jnp.flip(kernelAdjoint, axis = 1) // flip S2 + + val batchInput = Jax.jnp.expand_dims(input.jaxValue, axis = 0) // add dummy dim + val convResult = Jax.lax.conv_transpose( + lhs = batchInput, + rhs = kernelAdjoint, + strides = strides.toPythonProxy, + padding = padding.toString, + dimension_numbers = py.Dynamic.global.tuple(Seq("NHWC", "HWIO", "NHWC").toPythonProxy) + ) + val unbatchedRes = Jax.jnp.squeeze(convResult, axis = 0) // remove dummy dim + Tensor(unbatchedRes) + + /** Computes the 3D convolution of `input` with the specified kernel tensor. + * + * @param kernel - The convolution kernel tensor + * @param stride - Stride for the convolution. + * @param padding - Padding mode for the convolution. + * @return A new tensor representing the result of the convolution operation. + */ + def conv3d[S1: Label, S2: Label, S3: Label, InChannel: Label, V: IsFloating, OutChannel: Label]( + input: Tensor[S1 *: S2 *: S3 *: InChannel *: EmptyTuple, V], + kernel: Tensor[S1 *: S2 *: S3 *: InChannel *: OutChannel *: EmptyTuple, V], + stride: Stride3[S1, S2, S3] | Int = 1, + padding: Padding = Padding.SAME + ): Tensor[S1 *: S2 *: S3 *: OutChannel *: EmptyTuple, V] = + require( + input.shape(Axis[InChannel]) == kernel.shape(Axis[InChannel]), + s"Input channels mismatch: input has ${input.shape(Axis[InChannel])} channels, kernel expects ${kernel.shape(Axis[InChannel])} channels" + ) + val strides = stride match + case s: Int => Seq(s, s, s) + case (dim1, dim2, dim3) => Seq(dim1.size, dim2.size, dim3.size) + + // JAX requires input and kernel to have same rank, so we must add (and remove) dummy dim to input. + // 3D Layout: NDHWC (Batch, Depth, Height, Width, Channel) + val batchInput = Jax.jnp.expand_dims(input.jaxValue, axis = 0) // add dummy dim + val convResult = Jax.lax.conv_general_dilated( + lhs = batchInput, + rhs = kernel.jaxValue, + window_strides = strides.toPythonProxy, + padding = padding.toString, + dimension_numbers = py.Dynamic.global.tuple(Seq("NDHWC", "DHWIO", "NDHWC").toPythonProxy) + ) + val unbatchedRes = Jax.jnp.squeeze(convResult, axis = 0) // remove dummy dim + Tensor(unbatchedRes) + + /** Computes the transposed 3D convolution of `input` with the specified kernel tensor. + * + * @param kernel - The convolution kernel tensor + * @param stride - Stride for the convolution. + * @param padding - Padding mode for the convolution. + * @return A new tensor representing the result of the transposed convolution operation. + */ + def transposeConv3d[S1: Label, S2: Label, S3: Label, OutChannel: Label, V: IsFloating, InChannel: Label]( + input: Tensor[S1 *: S2 *: S3 *: OutChannel *: EmptyTuple, V], + kernel: Tensor[S1 *: S2 *: S3 *: InChannel *: OutChannel *: EmptyTuple, V], + stride: Stride3[S1, S2, S3] | Int = 1, + padding: Padding = Padding.SAME + ): Tensor[S1 *: S2 *: S3 *: InChannel *: EmptyTuple, V] = + require( + input.shape(Axis[OutChannel]) == kernel.shape(Axis[OutChannel]), + s"Input channels mismatch: input has ${input.shape(Axis[OutChannel])} channels (OutChannel), kernel expects ${kernel.shape(Axis[OutChannel])}" + ) + + val strides = stride match + case s: Int => Seq(s, s, s) + case (ae1, ae2, ae3) => Seq(ae1.size, ae2.size, ae3.size) + + // kernel -> kernel adjoint: swap in/out channels and flip all spatial dims + var kernelAdjoint = kernel.swap(Axis[InChannel], Axis[OutChannel]).jaxValue + kernelAdjoint = Jax.jnp.flip(kernelAdjoint, axis = 0) // flip S1 (Depth) + kernelAdjoint = Jax.jnp.flip(kernelAdjoint, axis = 1) // flip S2 (Height) + kernelAdjoint = Jax.jnp.flip(kernelAdjoint, axis = 2) // flip S3 (Width) + + val batchInput = Jax.jnp.expand_dims(input.jaxValue, axis = 0) // add dummy dim + val convResult = Jax.lax.conv_transpose( + lhs = batchInput, + rhs = kernelAdjoint, + strides = strides.toPythonProxy, + padding = padding.toString, + dimension_numbers = py.Dynamic.global.tuple(Seq("NDHWC", "DHWIO", "NDHWC").toPythonProxy) + ) + val unbatchedRes = Jax.jnp.squeeze(convResult, axis = 0) // remove dummy dim + Tensor(unbatchedRes) + +/** Extension methods for convolution operations, e.g. `input.conv1d(kernel)`. */ +private[dimwit] object ConvolutionExtensions: extension [S1: Label, InChannel: Label, V: IsFloating](input: Tensor[S1 *: InChannel *: EmptyTuple, V]) @@ -45,25 +252,7 @@ object ConvolutionOps: kernel: Tensor3[S1, InChannel, OutChannel, V], stride: Stride1[S1] | Int = 1, padding: Padding = Padding.SAME - ): Tensor[S1 *: OutChannel *: EmptyTuple, V] = - require( - input.shape(Axis[InChannel]) == kernel.shape(Axis[InChannel]), - s"Input channels mismatch: input has ${input.shape(Axis[InChannel])} channels, kernel expects ${kernel.shape(Axis[InChannel])} channels" - ) - val strides = stride match - case s: Int => Seq(s) - case ae: AxisExtent[S1] => Seq(ae.size) - // JAX requires input and kernel to have same rank, so we must add (and remove) dummy dim to input. - val batchInput = Jax.jnp.expand_dims(input.jaxValue, axis = 0) // add dummy dim - val convResult = Jax.lax.conv_general_dilated( - lhs = batchInput, - rhs = kernel.jaxValue, - window_strides = strides.toPythonProxy, - padding = padding.toString, - dimension_numbers = py.Dynamic.global.tuple(Seq("NHC", "HIO", "NHC").toPythonProxy) - ) - val unbatchedRes = Jax.jnp.squeeze(convResult, axis = 0) // remove dummy dim - Tensor(unbatchedRes) + ): Tensor[S1 *: OutChannel *: EmptyTuple, V] = ConvolutionOps.conv1d(input, kernel, stride, padding) extension [S1: Label, OutChannel: Label, V: IsFloating](input: Tensor[S1 *: OutChannel *: EmptyTuple, V]) @@ -78,31 +267,7 @@ object ConvolutionOps: kernel: Tensor[S1 *: InChannel *: OutChannel *: EmptyTuple, V], stride: Stride1[S1] | Int = 1, padding: Padding = Padding.SAME - ): Tensor[S1 *: InChannel *: EmptyTuple, V] = - require( - input.shape(Axis[OutChannel]) == kernel.shape(Axis[OutChannel]), - s"Input channels mismatch: input has ${input.shape(Axis[OutChannel])} channels (OutChannel), kernel expects ${kernel.shape(Axis[OutChannel])}" - ) - val strides = stride match - case s: Int => Seq(s) - case ex: AxisExtent[S1] => Seq(ex.size) - - // kernel -> kernal adjoint: swap in/out channels and flip spatial dims - var kernelAdjoint = kernel.swap(Axis[InChannel], Axis[OutChannel]).jaxValue - kernelAdjoint = Jax.jnp.flip(kernelAdjoint, axis = 0) // flip S1 - - val batchInput = Jax.jnp.expand_dims(input.jaxValue, axis = 0) // add dummy dim - val convResult = Jax.lax.conv_transpose( - lhs = batchInput, - rhs = kernelAdjoint, - strides = strides.toPythonProxy, - padding = padding.toString, - dimension_numbers = py.Dynamic.global.tuple(Seq("NHC", "HIO", "NHC").toPythonProxy) - ) - val unbatchedRes = Jax.jnp.squeeze(convResult, axis = 0) // remove dummy dim - Tensor(unbatchedRes) - - type Stride2[S1, S2] = (AxisExtent[S1], AxisExtent[S2]) + ): Tensor[S1 *: InChannel *: EmptyTuple, V] = ConvolutionOps.transposeConv1d(input, kernel, stride, padding) extension [S1: Label, S2: Label, InChannel: Label, V: IsFloating](input: Tensor[S1 *: S2 *: InChannel *: EmptyTuple, V]) @@ -117,25 +282,7 @@ object ConvolutionOps: kernel: Tensor[S1 *: S2 *: InChannel *: OutChannel *: EmptyTuple, V], stride: Stride2[S1, S2] | Int = 1, padding: Padding = Padding.SAME - ): Tensor[S1 *: S2 *: OutChannel *: EmptyTuple, V] = - require( - input.shape(Axis[InChannel]) == kernel.shape(Axis[InChannel]), - s"Input channels mismatch: input has ${input.shape(Axis[InChannel])} channels, kernel expects ${kernel.shape(Axis[InChannel])} channels" - ) - val strides = stride match - case s: Int => Seq(s, s) - case (ae1, ae2) => Seq(ae1.size, ae2.size) - // JAX requires input and kernel to have same rank, so we must add (and remove) dummy dim to input. - val batchInput = Jax.jnp.expand_dims(input.jaxValue, axis = 0) // add dummy dim - val convResult = Jax.lax.conv_general_dilated( - lhs = batchInput, - rhs = kernel.jaxValue, - window_strides = strides.toPythonProxy, - padding = padding.toString, - dimension_numbers = py.Dynamic.global.tuple(Seq("NHWC", "HWIO", "NHWC").toPythonProxy) - ) - val unbatchedRes = Jax.jnp.squeeze(convResult, axis = 0) // remove dummy dim - Tensor(unbatchedRes) + ): Tensor[S1 *: S2 *: OutChannel *: EmptyTuple, V] = ConvolutionOps.conv2d(input, kernel, stride, padding) extension [S1: Label, S2: Label, OutChannel: Label, V: IsFloating](input: Tensor[S1 *: S2 *: OutChannel *: EmptyTuple, V]) @@ -150,34 +297,7 @@ object ConvolutionOps: kernel: Tensor[S1 *: S2 *: InChannel *: OutChannel *: EmptyTuple, V], stride: Stride2[S1, S2] | Int = 1, padding: Padding = Padding.SAME - ): Tensor[S1 *: S2 *: InChannel *: EmptyTuple, V] = - require( - input.shape(Axis[OutChannel]) == kernel.shape(Axis[OutChannel]), - s"Input channels mismatch: input has ${input.shape(Axis[OutChannel])} channels (OutChannel), kernel expects ${kernel.shape(Axis[OutChannel])}" - ) - - // JAX requires input and kernel to have same rank. Add dummy batch dim if needed. - val strides = stride match - case s: Int => Seq(s, s) - case (ae1, ae2) => Seq(ae1.size, ae2.size) - - // kernel -> kernal adjoint: swap in/out channels and flip spatial dims - var kernelAdjoint = kernel.swap(Axis[InChannel], Axis[OutChannel]).jaxValue - kernelAdjoint = Jax.jnp.flip(kernelAdjoint, axis = 0) // flip S1 - kernelAdjoint = Jax.jnp.flip(kernelAdjoint, axis = 1) // flip S2 - - val batchInput = Jax.jnp.expand_dims(input.jaxValue, axis = 0) // add dummy dim - val convResult = Jax.lax.conv_transpose( - lhs = batchInput, - rhs = kernelAdjoint, - strides = strides.toPythonProxy, - padding = padding.toString, - dimension_numbers = py.Dynamic.global.tuple(Seq("NHWC", "HWIO", "NHWC").toPythonProxy) - ) - val unbatchedRes = Jax.jnp.squeeze(convResult, axis = 0) // remove dummy dim - Tensor(unbatchedRes) - - type Stride3[S1, S2, S3] = (AxisExtent[S1], AxisExtent[S2], AxisExtent[S3]) + ): Tensor[S1 *: S2 *: InChannel *: EmptyTuple, V] = ConvolutionOps.transposeConv2d(input, kernel, stride, padding) extension [S1: Label, S2: Label, S3: Label, InChannel: Label, V: IsFloating](input: Tensor[S1 *: S2 *: S3 *: InChannel *: EmptyTuple, V]) @@ -192,27 +312,7 @@ object ConvolutionOps: kernel: Tensor[S1 *: S2 *: S3 *: InChannel *: OutChannel *: EmptyTuple, V], stride: Stride3[S1, S2, S3] | Int = 1, padding: Padding = Padding.SAME - ): Tensor[S1 *: S2 *: S3 *: OutChannel *: EmptyTuple, V] = - require( - input.shape(Axis[InChannel]) == kernel.shape(Axis[InChannel]), - s"Input channels mismatch: input has ${input.shape(Axis[InChannel])} channels, kernel expects ${kernel.shape(Axis[InChannel])} channels" - ) - val strides = stride match - case s: Int => Seq(s, s, s) - case (dim1, dim2, dim3) => Seq(dim1.size, dim2.size, dim3.size) - - // JAX requires input and kernel to have same rank, so we must add (and remove) dummy dim to input. - // 3D Layout: NDHWC (Batch, Depth, Height, Width, Channel) - val batchInput = Jax.jnp.expand_dims(input.jaxValue, axis = 0) // add dummy dim - val convResult = Jax.lax.conv_general_dilated( - lhs = batchInput, - rhs = kernel.jaxValue, - window_strides = strides.toPythonProxy, - padding = padding.toString, - dimension_numbers = py.Dynamic.global.tuple(Seq("NDHWC", "DHWIO", "NDHWC").toPythonProxy) - ) - val unbatchedRes = Jax.jnp.squeeze(convResult, axis = 0) // remove dummy dim - Tensor(unbatchedRes) + ): Tensor[S1 *: S2 *: S3 *: OutChannel *: EmptyTuple, V] = ConvolutionOps.conv3d(input, kernel, stride, padding) extension [S1: Label, S2: Label, S3: Label, OutChannel: Label, V: IsFloating](input: Tensor[S1 *: S2 *: S3 *: OutChannel *: EmptyTuple, V]) @@ -227,29 +327,4 @@ object ConvolutionOps: kernel: Tensor[S1 *: S2 *: S3 *: InChannel *: OutChannel *: EmptyTuple, V], stride: Stride3[S1, S2, S3] | Int = 1, padding: Padding = Padding.SAME - ): Tensor[S1 *: S2 *: S3 *: InChannel *: EmptyTuple, V] = - require( - input.shape(Axis[OutChannel]) == kernel.shape(Axis[OutChannel]), - s"Input channels mismatch: input has ${input.shape(Axis[OutChannel])} channels (OutChannel), kernel expects ${kernel.shape(Axis[OutChannel])}" - ) - - val strides = stride match - case s: Int => Seq(s, s, s) - case (ae1, ae2, ae3) => Seq(ae1.size, ae2.size, ae3.size) - - // kernel -> kernel adjoint: swap in/out channels and flip all spatial dims - var kernelAdjoint = kernel.swap(Axis[InChannel], Axis[OutChannel]).jaxValue - kernelAdjoint = Jax.jnp.flip(kernelAdjoint, axis = 0) // flip S1 (Depth) - kernelAdjoint = Jax.jnp.flip(kernelAdjoint, axis = 1) // flip S2 (Height) - kernelAdjoint = Jax.jnp.flip(kernelAdjoint, axis = 2) // flip S3 (Width) - - val batchInput = Jax.jnp.expand_dims(input.jaxValue, axis = 0) // add dummy dim - val convResult = Jax.lax.conv_transpose( - lhs = batchInput, - rhs = kernelAdjoint, - strides = strides.toPythonProxy, - padding = padding.toString, - dimension_numbers = py.Dynamic.global.tuple(Seq("NDHWC", "DHWIO", "NDHWC").toPythonProxy) - ) - val unbatchedRes = Jax.jnp.squeeze(convResult, axis = 0) // remove dummy dim - Tensor(unbatchedRes) + ): Tensor[S1 *: S2 *: S3 *: InChannel *: EmptyTuple, V] = ConvolutionOps.transposeConv3d(input, kernel, stride, padding) diff --git a/core/src/main/scala/dimwit/tensor/tensorops/ElementWiseOps.scala b/core/src/main/scala/dimwit/tensor/tensorops/ElementWiseOps.scala index 8a15aaaa..6dc81de6 100644 --- a/core/src/main/scala/dimwit/tensor/tensorops/ElementWiseOps.scala +++ b/core/src/main/scala/dimwit/tensor/tensorops/ElementWiseOps.scala @@ -7,39 +7,293 @@ import dimwit.tensor.DType.Int32 import dimwit.tensor.Labels import dimwit.tensor.Tensor import dimwit.tensor.Tensor0 -import dimwit.tensor.TensorOps.IsBoolean -import dimwit.tensor.TensorOps.IsFloating -import dimwit.tensor.TensorOps.IsInteger -import dimwit.tensor.TensorOps.IsNumber +import dimwit.tensor.ValueTypeClasses.IsBoolean +import dimwit.tensor.ValueTypeClasses.IsFloating +import dimwit.tensor.ValueTypeClasses.IsInteger +import dimwit.tensor.ValueTypeClasses.IsNumber import dimwit.tensor.VType -import dimwit.tensor.tensorops.TensorOpsUtil.Broadcast +import dimwit.tensor.Broadcast -object ElementWiseOps: +private[dimwit] object ElementWiseOps: + + /** Tensors of the same type have the same labels, but their extents can still differ. + * JAX would then silently broadcast axes of extent 1, so operations without `!` check that the extents match. + */ + private[dimwit] def requireSameShape(t1: Tensor[?, ?], t2: Tensor[?, ?]): Unit = + require(t1.shape.dimensions == t2.shape.dimensions, s"Shape mismatch: ${t1.shape} vs ${t2.shape}") // --------------------------------------------------------- // General operations on any tensor type // --------------------------------------------------------- - /** Elementwise maximum of two tensors. */ - def maximum[T <: Tuple: Labels, V](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.maximum(t1.jaxValue, t2.jaxValue)) + def maximum[T <: Tuple: Labels, V](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, V] = + requireSameShape(t1, t2) + Tensor(Jax.jnp.maximum(t1.jaxValue, t2.jaxValue)) /** Elementwise minimum of two tensors. */ - def minimum[T <: Tuple: Labels, V](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.minimum(t1.jaxValue, t2.jaxValue)) + def minimum[T <: Tuple: Labels, V](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, V] = + requireSameShape(t1, t2) + Tensor(Jax.jnp.minimum(t1.jaxValue, t2.jaxValue)) /** Elementwise `<` of two tensors of the same shape. */ - def less[T <: Tuple: Labels, V](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, Bool] = Tensor(Jax.jnp.less(t1.jaxValue, t2.jaxValue)) + def less[T <: Tuple: Labels, V](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, Bool] = + requireSameShape(t1, t2) + Tensor(Jax.jnp.less(t1.jaxValue, t2.jaxValue)) /** Elementwise `<=` of two tensors of the same shape. */ - def lessEqual[T <: Tuple: Labels, V](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, Bool] = Tensor(Jax.jnp.less_equal(t1.jaxValue, t2.jaxValue)) + def lessEqual[T <: Tuple: Labels, V](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, Bool] = + requireSameShape(t1, t2) + Tensor(Jax.jnp.less_equal(t1.jaxValue, t2.jaxValue)) /** Elementwise `>` of two tensors of the same shape. */ - def greater[T <: Tuple: Labels, V](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, Bool] = Tensor(Jax.jnp.greater(t1.jaxValue, t2.jaxValue)) + def greater[T <: Tuple: Labels, V](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, Bool] = + requireSameShape(t1, t2) + Tensor(Jax.jnp.greater(t1.jaxValue, t2.jaxValue)) /** Elementwise `>=` of two tensors of the same shape. */ - def greaterEqual[T <: Tuple: Labels, V](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, Bool] = Tensor(Jax.jnp.greater_equal(t1.jaxValue, t2.jaxValue)) + def greaterEqual[T <: Tuple: Labels, V](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, Bool] = + requireSameShape(t1, t2) + Tensor(Jax.jnp.greater_equal(t1.jaxValue, t2.jaxValue)) + + /** Checks full array equality, returns true if all elements are equal. */ + def arrayEqual[T <: Tuple: Labels, V](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor0[Bool] = + requireSameShape(t1, t2) + Tensor0(Jax.jnp.array_equal(t1.jaxValue, t2.jaxValue)) /** Elementwise equality of two tensors of the same shape. */ - def equal[T <: Tuple: Labels, V](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, Bool] = Tensor(Jax.jnp.equal(t1.jaxValue, t2.jaxValue)) + def equal[T <: Tuple: Labels, V](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, Bool] = + requireSameShape(t1, t2) + Tensor(Jax.jnp.equal(t1.jaxValue, t2.jaxValue)) + + /** Performs element-wise addition of two tensors of the same shape and type. */ + def add[T <: Tuple: Labels, V: IsNumber](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, V] = + requireSameShape(t1, t2) + Tensor(Jax.jnp.add(t1.jaxValue, t2.jaxValue)) + + /** Returns a new tensor with each element negated. */ + def negate[T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.negative(t.jaxValue)) + + /** Subtracts one tensor from another of the same shape and type, returning a new tensor. */ + def subtract[T <: Tuple: Labels, V: IsNumber](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, V] = + requireSameShape(t1, t2) + Tensor(Jax.jnp.subtract(t1.jaxValue, t2.jaxValue)) + + /** Multiplies two tensors of the same shape and type element-wise, returning a new tensor. */ + def multiply[T <: Tuple: Labels, V: IsNumber]( + t1: Tensor[T, V], + t2: Tensor[T, V] + ): Tensor[T, V] = + requireSameShape(t1, t2) + Tensor(Jax.jnp.multiply(t1.jaxValue, t2.jaxValue)) + + /** Multiplies each element of `t` by the scalar `s`. */ + def scale[T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V], s: Tensor0[V]): Tensor[T, V] = Tensor(Jax.jnp.multiply(t.jaxValue, s.jaxValue)) + + /** Computes the element-wise remainder of `t1 / t2`, matching Python's `%` operator (the result takes the sign of the divisor). */ + def mod[T <: Tuple: Labels, V: IsNumber](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, V] = + requireSameShape(t1, t2) + Tensor(Jax.jnp.mod(t1.jaxValue, t2.jaxValue)) + + // --------------------------------------------------------- + // Operations on Floating tensors + // --------------------------------------------------------- + /** Divides two tensors of the same shape and type element-wise, returning a new tensor. */ + def divide[T <: Tuple: Labels, V: IsFloating](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, V] = + requireSameShape(t1, t2) + Tensor(Jax.jnp.divide(t1.jaxValue, t2.jaxValue)) + + // --------------------------------------------------------- + // IsBoolean operations + // --------------------------------------------------------- + /** Elementwise logical AND of two tensors of the same shape and type. */ + def logicalAnd[T <: Tuple: Labels, V: IsBoolean](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, V] = + requireSameShape(t1, t2) + Tensor(Jax.jnp.logical_and(t1.jaxValue, t2.jaxValue)) + + /** Elementwise logical OR of two tensors of the same shape and type. */ + def logicalOr[T <: Tuple: Labels, V: IsBoolean](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, V] = + requireSameShape(t1, t2) + Tensor(Jax.jnp.logical_or(t1.jaxValue, t2.jaxValue)) + + /** Elementwise logical XOR of two tensors of the same shape and type. */ + def logicalXor[T <: Tuple: Labels, V: IsBoolean](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, V] = + requireSameShape(t1, t2) + Tensor(Jax.jnp.logical_xor(t1.jaxValue, t2.jaxValue)) + + /** Elementwise logical NOT of a tensor. */ + def logicalNot[T <: Tuple: Labels, V: IsBoolean](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.logical_not(t.jaxValue)) + + // --------------------------------------------------------- + // Broadcasting variants: like the operations above, but `t1` and `t2` are first broadcast to their common shape, + // which is the one of the two shapes containing all axes of the other. + // --------------------------------------------------------- + + /** Like [[maximum]], but broadcasts `t1` and `t2` to their common shape first. */ + def maximum_![T1 <: Tuple, T2 <: Tuple, V](t1: Tensor[T1, V], t2: Tensor[T2, V])(using bc: Broadcast[T1, T2, V]): Tensor[bc.Out, V] = bc.applyTo(t1, t2)(maximum) + + /** Like [[minimum]], but broadcasts `t1` and `t2` to their common shape first. */ + def minimum_![T1 <: Tuple, T2 <: Tuple, V](t1: Tensor[T1, V], t2: Tensor[T2, V])(using bc: Broadcast[T1, T2, V]): Tensor[bc.Out, V] = bc.applyTo(t1, t2)(minimum) + + /** Like [[less]], but broadcasts `t1` and `t2` to their common shape first. */ + def less_![T1 <: Tuple, T2 <: Tuple, V](t1: Tensor[T1, V], t2: Tensor[T2, V])(using bc: Broadcast[T1, T2, V]): Tensor[bc.Out, Bool] = bc.applyTo(t1, t2)(less) + + /** Like [[lessEqual]], but broadcasts `t1` and `t2` to their common shape first. */ + def lessEqual_![T1 <: Tuple, T2 <: Tuple, V](t1: Tensor[T1, V], t2: Tensor[T2, V])(using bc: Broadcast[T1, T2, V]): Tensor[bc.Out, Bool] = bc.applyTo(t1, t2)(lessEqual) + + /** Like [[greater]], but broadcasts `t1` and `t2` to their common shape first. */ + def greater_![T1 <: Tuple, T2 <: Tuple, V](t1: Tensor[T1, V], t2: Tensor[T2, V])(using bc: Broadcast[T1, T2, V]): Tensor[bc.Out, Bool] = bc.applyTo(t1, t2)(greater) + + /** Like [[greaterEqual]], but broadcasts `t1` and `t2` to their common shape first. */ + def greaterEqual_![T1 <: Tuple, T2 <: Tuple, V](t1: Tensor[T1, V], t2: Tensor[T2, V])(using bc: Broadcast[T1, T2, V]): Tensor[bc.Out, Bool] = bc.applyTo(t1, t2)(greaterEqual) + + /** Like [[arrayEqual]], but broadcasts `t1` and `t2` to their common shape first. */ + def arrayEqual_![T1 <: Tuple, T2 <: Tuple, V](t1: Tensor[T1, V], t2: Tensor[T2, V])(using bc: Broadcast[T1, T2, V]): Tensor0[Bool] = + val (bt1, bt2) = bc.broadcast(t1, t2) + arrayEqual(bt1, bt2)(using bc.labelsOut) + + /** Like [[equal]], but broadcasts `t1` and `t2` to their common shape first. */ + def equal_![T1 <: Tuple, T2 <: Tuple, V](t1: Tensor[T1, V], t2: Tensor[T2, V])(using bc: Broadcast[T1, T2, V]): Tensor[bc.Out, Bool] = bc.applyTo(t1, t2)(equal) + + /** Like [[add]], but broadcasts `t1` and `t2` to their common shape first. */ + def add_![T1 <: Tuple, T2 <: Tuple, V: IsNumber](t1: Tensor[T1, V], t2: Tensor[T2, V])(using bc: Broadcast[T1, T2, V]): Tensor[bc.Out, V] = bc.applyTo(t1, t2)(add) + + /** Like [[subtract]], but broadcasts `t1` and `t2` to their common shape first. */ + def subtract_![T1 <: Tuple, T2 <: Tuple, V: IsNumber](t1: Tensor[T1, V], t2: Tensor[T2, V])(using bc: Broadcast[T1, T2, V]): Tensor[bc.Out, V] = bc.applyTo(t1, t2)(subtract) + + /** Like [[multiply]], but broadcasts `t1` and `t2` to their common shape first. */ + def multiply_![T1 <: Tuple, T2 <: Tuple, V: IsNumber](t1: Tensor[T1, V], t2: Tensor[T2, V])(using bc: Broadcast[T1, T2, V]): Tensor[bc.Out, V] = bc.applyTo(t1, t2)(multiply) + + /** Like [[mod]], but broadcasts `t1` and `t2` to their common shape first. */ + def mod_![T1 <: Tuple, T2 <: Tuple, V: IsNumber](t1: Tensor[T1, V], t2: Tensor[T2, V])(using bc: Broadcast[T1, T2, V]): Tensor[bc.Out, V] = bc.applyTo(t1, t2)(mod) + + /** Like [[divide]], but broadcasts `t1` and `t2` to their common shape first. */ + def divide_![T1 <: Tuple, T2 <: Tuple, V: IsFloating](t1: Tensor[T1, V], t2: Tensor[T2, V])(using bc: Broadcast[T1, T2, V]): Tensor[bc.Out, V] = bc.applyTo(t1, t2)(divide) + + /** Like [[logicalAnd]], but broadcasts `t1` and `t2` to their common shape first. */ + def logicalAnd_![T1 <: Tuple, T2 <: Tuple, V: IsBoolean](t1: Tensor[T1, V], t2: Tensor[T2, V])(using bc: Broadcast[T1, T2, V]): Tensor[bc.Out, V] = bc.applyTo(t1, t2)(logicalAnd) + + /** Like [[logicalOr]], but broadcasts `t1` and `t2` to their common shape first. */ + def logicalOr_![T1 <: Tuple, T2 <: Tuple, V: IsBoolean](t1: Tensor[T1, V], t2: Tensor[T2, V])(using bc: Broadcast[T1, T2, V]): Tensor[bc.Out, V] = bc.applyTo(t1, t2)(logicalOr) + + /** Like [[logicalXor]], but broadcasts `t1` and `t2` to their common shape first. */ + def logicalXor_![T1 <: Tuple, T2 <: Tuple, V: IsBoolean](t1: Tensor[T1, V], t2: Tensor[T2, V])(using bc: Broadcast[T1, T2, V]): Tensor[bc.Out, V] = bc.applyTo(t1, t2)(logicalXor) + + // --------------------------------------------------------- + // Unary operations + // --------------------------------------------------------- + + /** elementwise absolute value of `t`. */ + def abs[T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.abs(t.jaxValue)) + + /** elementwise sign (-1, 0 or 1) of `t`. */ + def sign[T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.sign(t.jaxValue)) + + /** clips the elements of `t` to the range [`min`, `max`]. */ + def clip[T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V], min: Tensor0[V], max: Tensor0[V]): Tensor[T, V] = Tensor(Jax.jnp.clip(t.jaxValue, min.jaxValue, max.jaxValue)) + + /** raises each element of `t` to the power `n`. */ + def pow[T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V], n: Tensor0[V]): Tensor[T, V] = Tensor(Jax.jnp.power(t.jaxValue, n.jaxValue)) + + // Elementwise operations on floating tensors + + /** elementwise square root of `t`. */ + def sqrt[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.sqrt(t.jaxValue)) + + /** elementwise exponential of `t`. */ + def exp[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.exp(t.jaxValue)) + + /** elementwise natural logarithm of `t`. */ + def log[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.log(t.jaxValue)) + + /** elementwise sine of `t`. */ + def sin[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.sin(t.jaxValue)) + + /** elementwise cosine of `t`. */ + def cos[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.cos(t.jaxValue)) + + /** elementwise hyperbolic tangent of `t`. */ + def tanh[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.tanh(t.jaxValue)) + + /** elementwise inverse sine of `t`. */ + def arcsin[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.arcsin(t.jaxValue)) + + /** elementwise inverse cosine of `t`. */ + def arccos[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.arccos(t.jaxValue)) + + /** elementwise inverse tangent of `t`. */ + def arctan[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.arctan(t.jaxValue)) + + /** elementwise floor of `t`. */ + def floor[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.floor(t.jaxValue)) + + /** elementwise ceiling of `t`. */ + def ceil[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.ceil(t.jaxValue)) + + /** elementwise rounding to the nearest integer of `t`. */ + def round[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.round(t.jaxValue)) + + /** elementwise test for NaN. */ + def isnan[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor[T, Bool] = Tensor(Jax.jnp.isnan(t.jaxValue)) + + /** elementwise test for finiteness (not NaN and not ±inf). */ + def isfinite[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor[T, Bool] = Tensor(Jax.jnp.isfinite(t.jaxValue)) + + /** replaces NaN by `nan`, +inf by `posInf` and -inf by `negInf`. + * By default, ±inf become the largest/smallest finite value of the dtype. + */ + def nanToNum[T <: Tuple: Labels, V](t: Tensor[T, V])(using + IsFloating[V] + )( + nan: Tensor0[V], + posInf: Tensor0[V] = IsFloating[V].maxFinite, + negInf: Tensor0[V] = IsFloating[V].minFinite + ): Tensor[T, V] = + Tensor(Jax.jnp.nan_to_num(t.jaxValue, nan = nan.jaxValue, posinf = posInf.jaxValue, neginf = negInf.jaxValue)) + + /** elementwise sigmoid activation of `t`. */ + def sigmoid[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnn.sigmoid(t.jaxValue)) + + /** elementwise ReLU activation of `t`. */ + def relu[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnn.relu(t.jaxValue)) + + /** elementwise GELU activation of `t`. */ + def gelu[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnn.gelu(t.jaxValue)) + + /** returns true if all elements of `t` and `other` are equal within `tolerance`. */ + def approxEquals[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V], other: Tensor[T, V], tolerance: Float = 1e-6f): Tensor0[Bool] = + all(approxElementEquals(t, other, tolerance)) + + /** compares `t` and `other` elementwise within `tolerance`. */ + def approxElementEquals[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V], other: Tensor[T, V], tolerance: Float = 1e-6f): Tensor[T, Bool] = + requireSameShape(t, other) + Tensor( + Jax.jnp.isclose( + t.jaxValue, + other.jaxValue, + atol = tolerance, + rtol = tolerance + ) + ) + + /** Like [[approxEquals]], but broadcasts `t1` and `t2` to their common shape first. */ + def approxEquals_![T1 <: Tuple, T2 <: Tuple, V: IsFloating](t1: Tensor[T1, V], t2: Tensor[T2, V], tolerance: Float = 1e-6f)(using bc: Broadcast[T1, T2, V]): Tensor0[Bool] = + all(approxElementEquals_!(t1, t2, tolerance)) + + /** Like [[approxElementEquals]], but broadcasts `t1` and `t2` to their common shape first. */ + def approxElementEquals_![T1 <: Tuple, T2 <: Tuple, V: IsFloating](t1: Tensor[T1, V], t2: Tensor[T2, V], tolerance: Float = 1e-6f)(using bc: Broadcast[T1, T2, V]): Tensor[bc.Out, Bool] = + bc.applyTo(t1, t2)((a, b) => approxElementEquals(a, b, tolerance)) + + // Operations on boolean tensors + + /** returns true if all elements of `t` are true, false otherwise */ + def all[T <: Tuple: Labels, V: IsBoolean](t: Tensor[T, V]): Tensor0[V] = Tensor0(Jax.jnp.all(t.jaxValue)) + + /** returns true if any element of `t` is true, false otherwise */ + def any[T <: Tuple: Labels, V: IsBoolean](t: Tensor[T, V]): Tensor0[V] = Tensor0(Jax.jnp.any(t.jaxValue)) + +private[dimwit] object ElementWiseExtensions: + + import ElementWiseOps.* // extension methods for comparisons extension [T <: Tuple: Labels, V](t: Tensor[T, V]) @@ -54,27 +308,28 @@ object ElementWiseOps: * Must be written backticked (``a `]], but broadcasts both sides to their common shape first. */ - def >![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor[bc.Out, Bool] = bc.applyTo(t, other)(greater) + def >![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor[bc.Out, Bool] = greater_!(t, other) /** Like [[>=]], but broadcasts both sides to their common shape first. */ - def >=![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor[bc.Out, Bool] = bc.applyTo(t, other)(greaterEqual) + def >=![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor[bc.Out, Bool] = greaterEqual_!(t, other) /** Checks full array equality, returns true if all elements are equal */ - def ===(other: Tensor[T, V]): Tensor0[Bool] = Tensor0(Jax.jnp.array_equal(t.jaxValue, other.jaxValue)) + def ===(other: Tensor[T, V]): Tensor0[Bool] = ElementWiseOps.arrayEqual(t, other) + + /** Like [[===]], but broadcasts both sides to their common shape first. */ + def ===![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor0[Bool] = ElementWiseOps.arrayEqual_!(t, other) /** Elementwise equality, returns a tensor of bools indicating which elements are equal */ - def elementEquals(other: Tensor[T, V]): Tensor[T, Bool] = - require(t.shape.dimensions == other.shape.dimensions, s"Shape mismatch: ${t.shape.dimensions} vs ${other.shape.dimensions}") - equal(t, other) + def elementEquals(other: Tensor[T, V]): Tensor[T, Bool] = equal(t, other) /** Like [[elementEquals]], but broadcasts both sides to their common shape first. */ - def elementEquals_![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor[bc.Out, Bool] = bc.applyTo(t, other)(equal) + def elementEquals_![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor[bc.Out, Bool] = equal_!(t, other) /** Casts the elements of this tensor to a tensor of type Bool. */ def asBool: Tensor[T, Bool] = t.asType(VType[Bool]) @@ -102,36 +357,6 @@ object ElementWiseOps: */ def asFloat[NewV: IsFloating](vtype: VType[NewV]): Tensor[T, NewV] = t.asType(vtype) - /** Performs element-wise addition of two tensors of the same shape and type. */ - def add[T <: Tuple: Labels, V: IsNumber](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.add(t1.jaxValue, t2.jaxValue)) - - /** Adds a scalar tensor to each element of a tensor. */ - def addScalar[T <: Tuple: Labels, V: IsNumber](t1: Tensor[T, V], s: Tensor0[V]): Tensor[T, V] = Tensor(Jax.jnp.add(t1.jaxValue, s.jaxValue)) - - /** Returns a new tensor with each element negated. */ - def negate[T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.negative(t.jaxValue)) - - /** Subtracts one tensor from another of the same shape and type, returning a new tensor. */ - def subtract[T <: Tuple: Labels, V: IsNumber](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.subtract(t1.jaxValue, t2.jaxValue)) - - /** Subtracts a scalar tensor from each element of a tensor. */ - def subtractScalar[T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V], s: Tensor0[V]): Tensor[T, V] = Tensor(Jax.jnp.subtract(t.jaxValue, s.jaxValue)) - - /** Multiplies two tensors of the same shape and type element-wise, returning a new tensor. */ - def multiply[T <: Tuple: Labels, V: IsNumber]( - t1: Tensor[T, V], - t2: Tensor[T, V] - ): Tensor[T, V] = Tensor(Jax.jnp.multiply(t1.jaxValue, t2.jaxValue)) - - /** Multiplies each element of a tensor by a scalar tensor, returning a new tensor. */ - def multiplyScalar[T <: Tuple: Labels, V: IsNumber](t1: Tensor[T, V], s: Tensor0[V]): Tensor[T, V] = Tensor(Jax.jnp.multiply(t1.jaxValue, s.jaxValue)) - - /** Computes the element-wise remainder of `t1 / t2`, matching Python's `%` operator (the result takes the sign of the divisor). */ - def mod[T <: Tuple: Labels, V: IsNumber](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.mod(t1.jaxValue, t2.jaxValue)) - - /** Computes the remainder of dividing each element of a tensor by a scalar tensor. */ - def modScalar[T <: Tuple: Labels, V: IsNumber](t1: Tensor[T, V], s: Tensor0[V]): Tensor[T, V] = Tensor(Jax.jnp.mod(t1.jaxValue, s.jaxValue)) - // extension methods for the binary operations on two tensors extension [T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V]) @@ -143,104 +368,72 @@ object ElementWiseOps: // extension methods for the scalar operations. extension [T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V]) - def +![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor[bc.Out, V] = bc.applyTo(t, other)(add) + def +![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor[bc.Out, V] = add_!(t, other) def unary_- : Tensor[T, V] = negate(t) - def -![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor[bc.Out, V] = bc.applyTo(t, other)(subtract) + def -![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor[bc.Out, V] = subtract_!(t, other) - def *![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor[bc.Out, V] = bc.applyTo(t, other)(multiply) - def scale(other: Tensor0[V]): Tensor[T, V] = multiplyScalar(t, other) - def %![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor[bc.Out, V] = bc.applyTo(t, other)(mod) + def *![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor[bc.Out, V] = multiply_!(t, other) + def scale(other: Tensor0[V]): Tensor[T, V] = ElementWiseOps.scale(t, other) + def %![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor[bc.Out, V] = mod_!(t, other) // extension methods extension [T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V]) - def abs: Tensor[T, V] = Tensor(Jax.jnp.abs(t.jaxValue)) - def sign: Tensor[T, V] = Tensor(Jax.jnp.sign(t.jaxValue)) - def clip(min: Tensor0[V], max: Tensor0[V]): Tensor[T, V] = Tensor(Jax.jnp.clip(t.jaxValue, min.jaxValue, max.jaxValue)) - def pow(n: Tensor0[V]): Tensor[T, V] = Tensor(Jax.jnp.power(t.jaxValue, n.jaxValue)) - - // --------------------------------------------------------- - // Operations on Floating tensors - // --------------------------------------------------------- - - /** Divides two tensors of the same shape and type element-wise, returning a new tensor. */ - def divide[T <: Tuple: Labels, V: IsFloating](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.divide(t1.jaxValue, t2.jaxValue)) - - /** Divides each element of a tensor by a scalar tensor, returning a new tensor. */ - def divideScalar[T <: Tuple: Labels, V: IsFloating](t1: Tensor[T, V], t2: Tensor0[V]): Tensor[T, V] = Tensor(Jax.jnp.divide(t1.jaxValue, t2.jaxValue)) + def abs: Tensor[T, V] = ElementWiseOps.abs(t) + def sign: Tensor[T, V] = ElementWiseOps.sign(t) + def clip(min: Tensor0[V], max: Tensor0[V]): Tensor[T, V] = ElementWiseOps.clip(t, min, max) + def pow(n: Tensor0[V]): Tensor[T, V] = ElementWiseOps.pow(t, n) // extension methods on floating tensors extension [T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]) def /(other: Tensor[T, V]): Tensor[T, V] = divide(t, other) - def /![O <: Tuple](other: Tensor[O, V])(using join: Broadcast[T, O, V]): Tensor[join.Out, V] = join.applyTo(t, other)(divide) - - def sqrt: Tensor[T, V] = Tensor(Jax.jnp.sqrt(t.jaxValue)) - def exp: Tensor[T, V] = Tensor(Jax.jnp.exp(t.jaxValue)) - def log: Tensor[T, V] = Tensor(Jax.jnp.log(t.jaxValue)) - def sin: Tensor[T, V] = Tensor(Jax.jnp.sin(t.jaxValue)) - def cos: Tensor[T, V] = Tensor(Jax.jnp.cos(t.jaxValue)) - def tanh: Tensor[T, V] = Tensor(Jax.jnp.tanh(t.jaxValue)) - def arcsin: Tensor[T, V] = Tensor(Jax.jnp.arcsin(t.jaxValue)) - def arccos: Tensor[T, V] = Tensor(Jax.jnp.arccos(t.jaxValue)) - def arctan: Tensor[T, V] = Tensor(Jax.jnp.arctan(t.jaxValue)) - def floor: Tensor[T, V] = Tensor(Jax.jnp.floor(t.jaxValue)) - def ceil: Tensor[T, V] = Tensor(Jax.jnp.ceil(t.jaxValue)) - def round: Tensor[T, V] = Tensor(Jax.jnp.round(t.jaxValue)) - def isnan: Tensor[T, Bool] = Tensor(Jax.jnp.isnan(t.jaxValue)) - def isfinite: Tensor[T, Bool] = Tensor(Jax.jnp.isfinite(t.jaxValue)) + def /![O <: Tuple](other: Tensor[O, V])(using join: Broadcast[T, O, V]): Tensor[join.Out, V] = divide_!(t, other) + + def sqrt: Tensor[T, V] = ElementWiseOps.sqrt(t) + def exp: Tensor[T, V] = ElementWiseOps.exp(t) + def log: Tensor[T, V] = ElementWiseOps.log(t) + def sin: Tensor[T, V] = ElementWiseOps.sin(t) + def cos: Tensor[T, V] = ElementWiseOps.cos(t) + def tanh: Tensor[T, V] = ElementWiseOps.tanh(t) + def arcsin: Tensor[T, V] = ElementWiseOps.arcsin(t) + def arccos: Tensor[T, V] = ElementWiseOps.arccos(t) + def arctan: Tensor[T, V] = ElementWiseOps.arctan(t) + def floor: Tensor[T, V] = ElementWiseOps.floor(t) + def ceil: Tensor[T, V] = ElementWiseOps.ceil(t) + def round: Tensor[T, V] = ElementWiseOps.round(t) + def isnan: Tensor[T, Bool] = ElementWiseOps.isnan(t) + def isfinite: Tensor[T, Bool] = ElementWiseOps.isfinite(t) /** replaces NaN by `nan`, +inf by `posInf` and -inf by `negInf`. * By default, ±inf become the largest/smallest finite value of the dtype. */ def nanToNum(using - IsFloating[V] + ev: IsFloating[V] )( nan: Tensor0[V], posInf: Tensor0[V] = IsFloating[V].maxFinite, negInf: Tensor0[V] = IsFloating[V].minFinite ): Tensor[T, V] = - Tensor(Jax.jnp.nan_to_num(t.jaxValue, nan = nan.jaxValue, posinf = posInf.jaxValue, neginf = negInf.jaxValue)) + ElementWiseOps.nanToNum(t)(using ev)(nan, posInf, negInf) // activation functions - def sigmoid: Tensor[T, V] = Tensor(Jax.jnn.sigmoid(t.jaxValue)) - def relu: Tensor[T, V] = Tensor(Jax.jnn.relu(t.jaxValue)) - def gelu: Tensor[T, V] = Tensor(Jax.jnn.gelu(t.jaxValue)) - - def approxEquals(other: Tensor[T, V], tolerance: Float = 1e-6f): Tensor0[Bool] = approxElementEquals(other, tolerance).all - def approxElementEquals(other: Tensor[T, V], tolerance: Float = 1e-6f): Tensor[T, Bool] = - Tensor( - Jax.jnp.allclose( - t.jaxValue, - other.jaxValue, - atol = tolerance, - rtol = tolerance - ) - ) - - // --------------------------------------------------------- - // IsBoolean operations - // --------------------------------------------------------- + def sigmoid: Tensor[T, V] = ElementWiseOps.sigmoid(t) + def relu: Tensor[T, V] = ElementWiseOps.relu(t) + def gelu: Tensor[T, V] = ElementWiseOps.gelu(t) - /** Elementwise logical AND of two tensors of the same shape and type. */ - def logicalAnd[T <: Tuple: Labels, V: IsBoolean](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.logical_and(t1.jaxValue, t2.jaxValue)) - - /** Elementwise logical OR of two tensors of the same shape and type. */ - def logicalOr[T <: Tuple: Labels, V: IsBoolean](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.logical_or(t1.jaxValue, t2.jaxValue)) - - /** Elementwise logical XOR of two tensors of the same shape and type. */ - def logicalXor[T <: Tuple: Labels, V: IsBoolean](t1: Tensor[T, V], t2: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.logical_xor(t1.jaxValue, t2.jaxValue)) - - /** Elementwise logical NOT of a tensor. */ - def logicalNot[T <: Tuple: Labels, V: IsBoolean](t: Tensor[T, V]): Tensor[T, V] = Tensor(Jax.jnp.logical_not(t.jaxValue)) + def approxEquals(other: Tensor[T, V], tolerance: Float = 1e-6f): Tensor0[Bool] = ElementWiseOps.approxEquals(t, other, tolerance) + def approxElementEquals(other: Tensor[T, V], tolerance: Float = 1e-6f): Tensor[T, Bool] = ElementWiseOps.approxElementEquals(t, other, tolerance) + def approxEquals_![O <: Tuple](other: Tensor[O, V], tolerance: Float = 1e-6f)(using bc: Broadcast[T, O, V]): Tensor0[Bool] = ElementWiseOps.approxEquals_!(t, other, tolerance) + def approxElementEquals_![O <: Tuple](other: Tensor[O, V], tolerance: Float = 1e-6f)(using bc: Broadcast[T, O, V]): Tensor[bc.Out, Bool] = ElementWiseOps.approxElementEquals_!(t, other, tolerance) extension [T <: Tuple: Labels, V: IsBoolean](t: Tensor[T, V]) /** returns true if all elements of the tensor are true, false otherwise */ - def all: Tensor0[V] = Tensor0(Jax.jnp.all(t.jaxValue)) + def all: Tensor0[V] = ElementWiseOps.all(t) /** return true if any element of the tensor is true, false otherwise */ - def any: Tensor0[V] = Tensor0(Jax.jnp.any(t.jaxValue)) + def any: Tensor0[V] = ElementWiseOps.any(t) /** returns a tensor of the same shape with each element negated (logical NOT) */ def unary_! : Tensor[T, V] = logicalNot(t) @@ -255,10 +448,10 @@ object ElementWiseOps: infix def xor(other: Tensor[T, V]): Tensor[T, V] = logicalXor(t, other) /** elementwise logical AND with a broadcastable tensor */ - infix def and_![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor[bc.Out, V] = bc.applyTo(t, other)(logicalAnd) + infix def and_![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor[bc.Out, V] = logicalAnd_!(t, other) /** elementwise logical OR with a broadcastable tensor */ - infix def or_![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor[bc.Out, V] = bc.applyTo(t, other)(logicalOr) + infix def or_![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor[bc.Out, V] = logicalOr_!(t, other) /** elementwise logical XOR with a broadcastable tensor */ - infix def xor_![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor[bc.Out, V] = bc.applyTo(t, other)(logicalXor) + infix def xor_![O <: Tuple](other: Tensor[O, V])(using bc: Broadcast[T, O, V]): Tensor[bc.Out, V] = logicalXor_!(t, other) diff --git a/core/src/main/scala/dimwit/tensor/tensorops/FunctionalOps.scala b/core/src/main/scala/dimwit/tensor/tensorops/FunctionalOps.scala index 49c60584..83ac9212 100644 --- a/core/src/main/scala/dimwit/tensor/tensorops/FunctionalOps.scala +++ b/core/src/main/scala/dimwit/tensor/tensorops/FunctionalOps.scala @@ -21,7 +21,8 @@ import me.shadaj.scalapy.readwrite.Writer import scala.NamedTuple.NamedTuple import scala.annotation.implicitNotFound -object FunctionalOps: +/** Functional operations on tensors, e.g. `Tensor.zipvmap`, and the type classes they need. */ +private[dimwit] object FunctionalOps: /** Prepends the axis `L` to every tensor of the tensor tree `FOut`: the result * type of a `vmap`/`zipvmap` whose body returned `FOut`. @@ -120,6 +121,12 @@ object FunctionalOps: fromPyTree.fromPyTree(jaxResult) export ZipVmap.zipvmap +/** Extension methods for functional operations, e.g. `t.vmap(Axis[A])(f)`. */ +private[dimwit] object FunctionalExtensions: + + import FunctionalOps.PrependAxis + import FunctionalOps.ZipVmap + extension [T <: Tuple: Labels, V](t: Tensor[T, V]) /** Zips the current tensor with another tensor along the specified axis diff --git a/core/src/main/scala/dimwit/tensor/tensorops/LinearAlgebraOps.scala b/core/src/main/scala/dimwit/tensor/tensorops/LinearAlgebraExtensions.scala similarity index 94% rename from core/src/main/scala/dimwit/tensor/tensorops/LinearAlgebraOps.scala rename to core/src/main/scala/dimwit/tensor/tensorops/LinearAlgebraExtensions.scala index cde95f15..ce1b3d25 100644 --- a/core/src/main/scala/dimwit/tensor/tensorops/LinearAlgebraOps.scala +++ b/core/src/main/scala/dimwit/tensor/tensorops/LinearAlgebraExtensions.scala @@ -8,10 +8,10 @@ import dimwit.tensor.Tensor import dimwit.tensor.Tensor0 import dimwit.tensor.Tensor1 import dimwit.tensor.Tensor2 -import dimwit.tensor.TensorOps.IsFloating -import dimwit.tensor.TensorOps.IsNumber +import dimwit.tensor.ValueTypeClasses.IsFloating +import dimwit.tensor.ValueTypeClasses.IsNumber -object LinearAlgebraOps: +private[dimwit] object LinearAlgebraExtensions: extension [L1: Label, L2: Label, V](t: Tensor2[L1, L2, V]) diff --git a/core/src/main/scala/dimwit/tensor/tensorops/ReductionOps.scala b/core/src/main/scala/dimwit/tensor/tensorops/ReductionOps.scala index 4aec2e49..c0787fa2 100644 --- a/core/src/main/scala/dimwit/tensor/tensorops/ReductionOps.scala +++ b/core/src/main/scala/dimwit/tensor/tensorops/ReductionOps.scala @@ -10,38 +10,95 @@ import dimwit.tensor.ShapeTypeHelpers.AxisRemover import dimwit.tensor.ShapeTypeHelpers.UnwrapAxes import dimwit.tensor.Tensor import dimwit.tensor.Tensor0 -import dimwit.tensor.TensorOps.IsFloating -import dimwit.tensor.TensorOps.IsNumber +import dimwit.tensor.ValueTypeClasses.IsFloating +import dimwit.tensor.ValueTypeClasses.IsNumber import me.shadaj.scalapy.py import me.shadaj.scalapy.py.SeqConverters import me.shadaj.scalapy.readwrite.Writer -object ReductionOps: +private[dimwit] object ReductionOps: + + // the axes reduced over are removed from the result, without an axis `t` is reduced to a scalar. + + /** sums `t` (over all elements, or) along the specified axis or axes. */ + def sum[T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V]): Tensor0[V] = Tensor0(Jax.jnp.sum(t.jaxValue)) + def sum[T <: Tuple: Labels, V: IsNumber, L: Label](t: Tensor[T, V], axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.sum(t.jaxValue, axis = ev.index)) + def sum[T <: Tuple: Labels, V: IsNumber, Inputs <: Tuple](t: Tensor[T, V], axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.sum(t.jaxValue, axis = ev.indices.toPythonProxy)) + + /** takes the maximum of `t` (over all elements, or) along the specified axis or axes. */ + def max[T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V]): Tensor0[V] = Tensor0(Jax.jnp.max(t.jaxValue)) + def max[T <: Tuple: Labels, V: IsNumber, L: Label](t: Tensor[T, V], axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.max(t.jaxValue, axis = ev.index)) + def max[T <: Tuple: Labels, V: IsNumber, Inputs <: Tuple](t: Tensor[T, V], axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.max(t.jaxValue, axis = ev.indices.toPythonProxy)) + + /** takes the minimum of `t` (over all elements, or) along the specified axis or axes. */ + def min[T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V]): Tensor0[V] = Tensor0(Jax.jnp.min(t.jaxValue)) + def min[T <: Tuple: Labels, V: IsNumber, L: Label](t: Tensor[T, V], axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.min(t.jaxValue, axis = ev.index)) + def min[T <: Tuple: Labels, V: IsNumber, Inputs <: Tuple](t: Tensor[T, V], axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.min(t.jaxValue, axis = ev.indices.toPythonProxy)) + + /** returns the index of the maximum of `t` along the specified axis or axes. */ + def argmax[T <: Tuple: Labels, V: IsNumber, L: Label](t: Tensor[T, V], axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, Int32] = Tensor(Jax.jnp.argmax(t.jaxValue, axis = ev.index)) + def argmax[T <: Tuple: Labels, V: IsNumber, Inputs <: Tuple](t: Tensor[T, V], axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, Int32] = Tensor(Jax.jnp.argmax(t.jaxValue, axis = ev.indices.toPythonProxy)) + + /** returns the index of the minimum of `t` along the specified axis or axes. */ + def argmin[T <: Tuple: Labels, V: IsNumber, L: Label](t: Tensor[T, V], axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, Int32] = Tensor(Jax.jnp.argmin(t.jaxValue, axis = ev.index)) + def argmin[T <: Tuple: Labels, V: IsNumber, Inputs <: Tuple](t: Tensor[T, V], axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, Int32] = Tensor(Jax.jnp.argmin(t.jaxValue, axis = ev.indices.toPythonProxy)) + + /** computes the mean of `t` (over all elements, or) along the specified axis or axes. */ + def mean[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor0[V] = Tensor0(Jax.jnp.mean(t.jaxValue)) + def mean[T <: Tuple: Labels, V: IsFloating, L: Label](t: Tensor[T, V], axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.mean(t.jaxValue, axis = ev.index)) + def mean[T <: Tuple: Labels, V: IsFloating, Inputs <: Tuple](t: Tensor[T, V], axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.mean(t.jaxValue, axis = ev.indices.toPythonProxy)) + + /** computes the standard deviation of `t` (over all elements, or) along the specified axis or axes. */ + def std[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor0[V] = Tensor0(Jax.jnp.std(t.jaxValue)) + def std[T <: Tuple: Labels, V: IsFloating, L: Label](t: Tensor[T, V], axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.std(t.jaxValue, axis = ev.index)) + def std[T <: Tuple: Labels, V: IsFloating, Inputs <: Tuple](t: Tensor[T, V], axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.std(t.jaxValue, axis = ev.indices.toPythonProxy)) + + /** computes the median of `t` (over all elements, or) along the specified axis or axes. */ + def median[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor0[V] = Tensor0(Jax.jnp.median(t.jaxValue)) + def median[T <: Tuple: Labels, V: IsFloating, L: Label](t: Tensor[T, V], axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.median(t.jaxValue, axis = ev.index)) + def median[T <: Tuple: Labels, V: IsFloating, Inputs <: Tuple](t: Tensor[T, V], axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.median(t.jaxValue, axis = ev.indices.toPythonProxy)) + + /** computes the mean of `t`, ignoring NaN values, (over all elements, or) along the specified axis or axes. */ + def nanmean[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor0[V] = Tensor0(Jax.jnp.nanmean(t.jaxValue)) + def nanmean[T <: Tuple: Labels, V: IsFloating, L: Label](t: Tensor[T, V], axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.nanmean(t.jaxValue, axis = ev.index)) + def nanmean[T <: Tuple: Labels, V: IsFloating, Inputs <: Tuple](t: Tensor[T, V], axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.nanmean(t.jaxValue, axis = ev.indices.toPythonProxy)) + + /** computes the median of `t`, ignoring NaN values, (over all elements, or) along the specified axis or axes. */ + def nanmedian[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]): Tensor0[V] = Tensor0(Jax.jnp.nanmedian(t.jaxValue)) + def nanmedian[T <: Tuple: Labels, V: IsFloating, L: Label](t: Tensor[T, V], axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.nanmedian(t.jaxValue, axis = ev.index)) + def nanmedian[T <: Tuple: Labels, V: IsFloating, Inputs <: Tuple](t: Tensor[T, V], axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.nanmedian(t.jaxValue, axis = ev.indices.toPythonProxy)) + + /** computes the `q`th quantile of `t` (over all elements, or) along the specified axis or axes. */ + def quantile[T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V], q: Float): Tensor0[V] = Tensor0(Jax.jnp.quantile(t.jaxValue, q)) + def quantile[T <: Tuple: Labels, V: IsFloating, L: Label](t: Tensor[T, V], q: Float, axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.quantile(t.jaxValue, q, axis = ev.index)) + def quantile[T <: Tuple: Labels, V: IsFloating, Inputs <: Tuple](t: Tensor[T, V], q: Float, axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.quantile(t.jaxValue, q, axis = ev.indices.toPythonProxy)) + +private[dimwit] object ReductionExtensions: extension [T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V]) /** sums the tensor `t` along the specified axes, returning a new tensor with those axes removed. */ - def sum[Inputs <: Tuple](axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.sum(t.jaxValue, axis = ev.indices.toPythonProxy)) - def sum[L: Label](axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.sum(t.jaxValue, axis = ev.index)) - def sum: Tensor0[V] = Tensor0(Jax.jnp.sum(t.jaxValue)) + def sum[Inputs <: Tuple](axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = ReductionOps.sum(t, axes) + def sum[L: Label](axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = ReductionOps.sum(t, axis) + def sum: Tensor0[V] = ReductionOps.sum(t) /** takes the maximum of the tensor `t` along the specified axes, returning a new tensor with those axes removed. */ - def max[Inputs <: Tuple](axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.max(t.jaxValue, axis = ev.indices.toPythonProxy)) - def max[L: Label](axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.max(t.jaxValue, axis = ev.index)) - def max: Tensor0[V] = Tensor0(Jax.jnp.max(t.jaxValue)) + def max[Inputs <: Tuple](axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = ReductionOps.max(t, axes) + def max[L: Label](axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = ReductionOps.max(t, axis) + def max: Tensor0[V] = ReductionOps.max(t) /** takes the minimum of the tensor `t` along the specified axes, returning a new tensor with those axes removed. */ - def min[Inputs <: Tuple](axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.min(t.jaxValue, axis = ev.indices.toPythonProxy)) - def min[L: Label](axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.min(t.jaxValue, axis = ev.index)) - def min: Tensor0[V] = Tensor0(Jax.jnp.min(t.jaxValue)) + def min[Inputs <: Tuple](axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = ReductionOps.min(t, axes) + def min[L: Label](axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = ReductionOps.min(t, axis) + def min: Tensor0[V] = ReductionOps.min(t) /** argument of the maximum of the tensor `t` along the specified axes, returning a new tensor with those axes removed. */ - def argmax[Inputs <: Tuple](axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, Int32] = Tensor(Jax.jnp.argmax(t.jaxValue, axis = ev.indices.toPythonProxy)) - def argmax[L: Label](axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, Int32] = Tensor(Jax.jnp.argmax(t.jaxValue, axis = ev.index)) + def argmax[Inputs <: Tuple](axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, Int32] = ReductionOps.argmax(t, axes) + def argmax[L: Label](axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, Int32] = ReductionOps.argmax(t, axis) /** argument of the minimum of the tensor `t` along the specified axes, returning a new tensor with those axes removed. */ - def argmin[Inputs <: Tuple](axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, Int32] = Tensor(Jax.jnp.argmin(t.jaxValue, axis = ev.indices.toPythonProxy)) - def argmin[L: Label](axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, Int32] = Tensor(Jax.jnp.argmin(t.jaxValue, axis = ev.index)) + def argmin[Inputs <: Tuple](axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, Int32] = ReductionOps.argmin(t, axes) + def argmin[L: Label](axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, Int32] = ReductionOps.argmin(t, axis) // --------------------------------------------------------- // IsFloat operations (IsFloat or IsInt) @@ -50,31 +107,31 @@ object ReductionOps: extension [T <: Tuple: Labels, V: IsFloating](t: Tensor[T, V]) /** computes the mean of the tensor `t` along the specified axes, returning a new tensor with those axes removed. */ - def mean[Inputs <: Tuple](axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.mean(t.jaxValue, axis = ev.indices.toPythonProxy)) - def mean[L: Label](axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.mean(t.jaxValue, axis = ev.index)) - def mean: Tensor0[V] = Tensor0(Jax.jnp.mean(t.jaxValue)) + def mean[Inputs <: Tuple](axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = ReductionOps.mean(t, axes) + def mean[L: Label](axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = ReductionOps.mean(t, axis) + def mean: Tensor0[V] = ReductionOps.mean(t) /** computes the mean of the tensor `t` along the specified axes, returning a new tensor with those axes removed. */ - def std[Inputs <: Tuple](axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.std(t.jaxValue, axis = ev.indices.toPythonProxy)) - def std[L: Label](axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.std(t.jaxValue, axis = ev.index)) - def std: Tensor0[V] = Tensor0(Jax.jnp.std(t.jaxValue)) + def std[Inputs <: Tuple](axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = ReductionOps.std(t, axes) + def std[L: Label](axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = ReductionOps.std(t, axis) + def std: Tensor0[V] = ReductionOps.std(t) /** computes the qth quantile of the tensor `t` along the specified axes, returning a new tensor with those axes removed. */ - def quantile[Inputs <: Tuple](q: Float, axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.quantile(t.jaxValue, q, axis = ev.indices.toPythonProxy)) - def quantile[L: Label](q: Float, axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.quantile(t.jaxValue, q, axis = ev.index)) - def quantile(q: Float): Tensor0[V] = Tensor0(Jax.jnp.quantile(t.jaxValue, q)) + def quantile[Inputs <: Tuple](q: Float, axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = ReductionOps.quantile(t, q, axes) + def quantile[L: Label](q: Float, axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = ReductionOps.quantile(t, q, axis) + def quantile(q: Float): Tensor0[V] = ReductionOps.quantile(t, q) /** computes the median of the tensor `t` along the specified axes, returning a new tensor with those axes removed. */ - def median: Tensor0[V] = Tensor0(Jax.jnp.median(t.jaxValue)) - def median[L: Label](axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.median(t.jaxValue, axis = ev.index)) - def median[Inputs <: Tuple](axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.median(t.jaxValue, axis = ev.indices.toPythonProxy)) + def median: Tensor0[V] = ReductionOps.median(t) + def median[L: Label](axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = ReductionOps.median(t, axis) + def median[Inputs <: Tuple](axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = ReductionOps.median(t, axes) /** computes the mean of the tensor `t` along the specified axes, ignoring na values and returning a new tensor with those axes removed. */ - def nanmean: Tensor0[V] = Tensor0(Jax.jnp.nanmean(t.jaxValue)) - def nanmean[L: Label](axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.nanmean(t.jaxValue, axis = ev.index)) - def nanmean[Inputs <: Tuple](axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.nanmean(t.jaxValue, axis = ev.indices.toPythonProxy)) + def nanmean: Tensor0[V] = ReductionOps.nanmean(t) + def nanmean[L: Label](axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = ReductionOps.nanmean(t, axis) + def nanmean[Inputs <: Tuple](axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = ReductionOps.nanmean(t, axes) /** computes the median of the tensor `t` along the specified axes, ignoring na values and returning a new tensor with those axes removed. */ - def nanmedian[Inputs <: Tuple](axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.nanmedian(t.jaxValue, axis = ev.indices.toPythonProxy)) - def nanmedian[L: Label](axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = Tensor(Jax.jnp.nanmedian(t.jaxValue, axis = ev.index)) - def nanmedian: Tensor0[V] = Tensor0(Jax.jnp.nanmedian(t.jaxValue)) + def nanmedian[Inputs <: Tuple](axes: Inputs)(using ev: AxesRemover[T, UnwrapAxes[Inputs]], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = ReductionOps.nanmedian(t, axes) + def nanmedian[L: Label](axis: Axis[L])(using ev: AxisRemover[T, L], l: Labels[ev.RemainingAxes]): Tensor[ev.RemainingAxes, V] = ReductionOps.nanmedian(t, axis) + def nanmedian: Tensor0[V] = ReductionOps.nanmedian(t) diff --git a/core/src/main/scala/dimwit/tensor/tensorops/StructuralOps.scala b/core/src/main/scala/dimwit/tensor/tensorops/StructuralOps.scala index bdc7b1c3..0c9d405c 100644 --- a/core/src/main/scala/dimwit/tensor/tensorops/StructuralOps.scala +++ b/core/src/main/scala/dimwit/tensor/tensorops/StructuralOps.scala @@ -34,7 +34,7 @@ import dimwit.tensor.TensorEvidence.CheckValid import dimwit.tensor.TensorEvidence.ComputeMissing import dimwit.tensor.TensorEvidence.IsPermutation import dimwit.tensor.TensorEvidence.ValidationResult -import dimwit.tensor.tensorops.TensorOpsUtil.Broadcast3 +import dimwit.tensor.Broadcast3 import dimwit.|+| import me.shadaj.scalapy.py import me.shadaj.scalapy.py.SeqConverters @@ -44,7 +44,8 @@ import me.shadaj.scalapy.readwrite.Writer import scala.annotation.implicitNotFound import scala.util.NotGiven -object StructuralOps: +/** Structural operations on tensors, e.g. `Tensor.stack` or `Tensor.concatenate`, and the type classes they need. */ +private[dimwit] object StructuralOps: /** Inserts axis `New` directly after axis `Anchor`. */ trait AxisInserter[T <: Tuple, Anchor, New]: @@ -93,7 +94,7 @@ object StructuralOps: tail: AxisSwapper.Aux[T, L1, L2, O] ): AxisSwapper.Aux[H *: T, L1, L2, H *: O] = AxisSwapper.instance - private object Util: + private[dimwit] object Util: type ExtractLabel[X] = X match case AxisAtIndex[l] => l @@ -152,6 +153,8 @@ object StructuralOps: ifTrue: Tensor[T, V], ifFalse: Tensor[T, V] ): Tensor[T, V] = + ElementWiseOps.requireSameShape(condition, ifTrue) + ElementWiseOps.requireSameShape(ifTrue, ifFalse) Tensor(Jax.jnp.where(condition.jaxValue, ifTrue.jaxValue, ifFalse.jaxValue)) /** Like [[where]], but broadcasts condition, `ifTrue` and `ifFalse` to their common shape, @@ -359,6 +362,12 @@ object StructuralOps: val headTensor = Tensor[NewShape, V](currentArr)(using newLabelsWitness) headTensor *: tailMaker(arrays.tail, compLabels.tail, originalLabels, splitIndex) +/** Extension methods for structural operations, e.g. `t.transpose`, `t.slice(...)` or `t.rearrange(...)`. */ +private[dimwit] object StructuralExtensions: + + import StructuralOps.* + import StructuralOps.Util.* + extension [T <: Tuple, V](tensor: Tensor[T, V]) /** takes a concatenated tensor and splits it into a tuple of tensors along the specified axis, @@ -862,9 +871,6 @@ object StructuralOps: newLabels: Labels[ev.NewShape] ): Tensor[ev.NewShape, V] = Tensor(tensor.jaxValue) - def retag[newT <: Tuple](using newLabels: Labels[newT]): Tensor[newT, V] = - Tensor(tensor.jaxValue)(using newLabels) - def relabelAll[newT <: Tuple]( newAxes: newT )(using diff --git a/core/src/main/scala/dimwit/tensor/tensorops/Tensor0Ops.scala b/core/src/main/scala/dimwit/tensor/tensorops/Tensor0Extensions.scala similarity index 98% rename from core/src/main/scala/dimwit/tensor/tensorops/Tensor0Ops.scala rename to core/src/main/scala/dimwit/tensor/tensorops/Tensor0Extensions.scala index 8ea24de4..ec743505 100644 --- a/core/src/main/scala/dimwit/tensor/tensorops/Tensor0Ops.scala +++ b/core/src/main/scala/dimwit/tensor/tensorops/Tensor0Extensions.scala @@ -3,7 +3,7 @@ package dimwit.tensor.tensorops import dimwit.tensor.DType.* import dimwit.tensor.Tensor0 -object Tensor0Ops: +private[dimwit] object Tensor0Extensions: private inline def checkTracer[V, R](scalar: Tensor0[V]): Unit = require( diff --git a/core/src/main/scala/dimwit/tensor/tensorops/Tensor1Ops.scala b/core/src/main/scala/dimwit/tensor/tensorops/Tensor1Extensions.scala similarity index 97% rename from core/src/main/scala/dimwit/tensor/tensorops/Tensor1Ops.scala rename to core/src/main/scala/dimwit/tensor/tensorops/Tensor1Extensions.scala index 93eb1b49..95114bd5 100644 --- a/core/src/main/scala/dimwit/tensor/tensorops/Tensor1Ops.scala +++ b/core/src/main/scala/dimwit/tensor/tensorops/Tensor1Extensions.scala @@ -13,7 +13,7 @@ import me.shadaj.scalapy.py import me.shadaj.scalapy.py.SeqConverters import me.shadaj.scalapy.readwrite.Writer -object Tensor1Ops: +private[dimwit] object Tensor1Extensions: extension [L, V](t: Tensor1[L, V]) diff --git a/core/src/main/scala/dimwit/tensor/tensorops/Tensor2Ops.scala b/core/src/main/scala/dimwit/tensor/tensorops/Tensor2Extensions.scala similarity index 92% rename from core/src/main/scala/dimwit/tensor/tensorops/Tensor2Ops.scala rename to core/src/main/scala/dimwit/tensor/tensorops/Tensor2Extensions.scala index c376c822..b3c08df8 100644 --- a/core/src/main/scala/dimwit/tensor/tensorops/Tensor2Ops.scala +++ b/core/src/main/scala/dimwit/tensor/tensorops/Tensor2Extensions.scala @@ -6,7 +6,7 @@ import dimwit.tensor.Label import dimwit.tensor.Labels import dimwit.tensor.Tensor2 -object Tensor2Ops: +private[dimwit] object Tensor2Extensions: extension [L1: Label, L2: Label, V](t: Tensor2[L1, L2, V]) @@ -18,7 +18,7 @@ object Tensor2Ops: * @param axis1 the second axis to swap * @return a new Tensor2 with the specified axes transposed */ - def transpose(axis2: Axis[L2], axis1: Axis[L1]): Tensor2[L2, L1, V] = StructuralOps.transpose(t)(axis2, axis1) + def transpose(axis2: Axis[L2], axis1: Axis[L1]): Tensor2[L2, L1, V] = StructuralExtensions.transpose(t)(axis2, axis1) extension [L1, L2, V, X](t: Tensor2[L1, L2, V])(using ev: HasScalar[V, X]) /** Converts a Tensor2 to a nested Scala Array (Array of Arrays). diff --git a/core/src/main/scala/dimwit/tensor/tensorops/Tensor3Ops.scala b/core/src/main/scala/dimwit/tensor/tensorops/Tensor3Extensions.scala similarity index 94% rename from core/src/main/scala/dimwit/tensor/tensorops/Tensor3Ops.scala rename to core/src/main/scala/dimwit/tensor/tensorops/Tensor3Extensions.scala index 8b705168..468309cb 100644 --- a/core/src/main/scala/dimwit/tensor/tensorops/Tensor3Ops.scala +++ b/core/src/main/scala/dimwit/tensor/tensorops/Tensor3Extensions.scala @@ -3,7 +3,7 @@ package dimwit.tensor.tensorops import dimwit.tensor.HasScalar import dimwit.tensor.Tensor3 -object Tensor3Ops: +private[dimwit] object Tensor3Extensions: extension [L1, L2, L3, V, X](t: Tensor3[L1, L2, L3, V])(using ev: HasScalar[V, X]) /** Converts a Tensor3 to a nested Scala Array (Array of Arrays of Arrays). diff --git a/core/src/main/scala/dimwit/tensor/tensorops/TensorOpsUtils.scala b/core/src/main/scala/dimwit/tensor/tensorops/TensorOpsUtils.scala deleted file mode 100644 index 7f4e43bc..00000000 --- a/core/src/main/scala/dimwit/tensor/tensorops/TensorOpsUtils.scala +++ /dev/null @@ -1,97 +0,0 @@ -package dimwit.tensor.tensorops - -import dimwit.tensor.Labels -import dimwit.tensor.Tensor -import dimwit.tensor.TupleHelpers.StrictSubset - -import scala.annotation.implicitNotFound - -object TensorOpsUtil: - - import dimwit.tensor.TensorOps.broadcastTo - - @implicitNotFound("Cannot broadcast tensors of shapes ${T1} and ${T2}. If same shape no broadcasting allowed!") - sealed trait Broadcast[T1 <: Tuple, T2 <: Tuple, V]: - type Out <: Tuple - given labelsOut: Labels[Out] - def broadcast(t1: Tensor[T1, V], t2: Tensor[T2, V]): (Tensor[Out, V], Tensor[Out, V]) - def applyTo[V2](t1: Tensor[T1, V], t2: Tensor[T2, V])(f: (Tensor[Out, V], Tensor[Out, V]) => Tensor[Out, V2]): Tensor[Out, V2] = - val (bt1, bt2) = broadcast(t1, t2) - f(bt1, bt2) - - object Broadcast extends BroadcastLowPriority: - - given broadcastLeft[T1 <: Tuple: Labels, T2 <: Tuple: Labels, V](using - StrictSubset[T2, T1] - ): Broadcast[T1, T2, V] with - type Out = T1 - val labelsOut = summon[Labels[T1]] - def broadcast(t1: Tensor[T1, V], t2: Tensor[T2, V]) = - (t1, t2.broadcastTo[T1](t1.shape)) - - trait BroadcastLowPriority: - given broadcastRight[T1 <: Tuple: Labels, T2 <: Tuple: Labels, V](using - StrictSubset[T1, T2] - ): Broadcast[T1, T2, V] with - type Out = T2 - val labelsOut = summon[Labels[T2]] - def broadcast(t1: Tensor[T1, V], t2: Tensor[T2, V]) = - (t1.broadcastTo[T2](t2.shape), t2) - - /** Broadcasts three tensors to their common shape, which is the one of the three shapes containing all - * axes of the other two. As for [[Broadcast]] at least one of the tensors has to be broadcast. - */ - @implicitNotFound( - "Cannot broadcast tensors of shapes ${T1}, ${T2} and ${T3}. One of them must contain all axes of the other two. If all same shape no broadcasting allowed!" - ) - sealed trait Broadcast3[T1 <: Tuple, T2 <: Tuple, T3 <: Tuple, V]: - type Out <: Tuple - given labelsOut: Labels[Out] - def broadcast[V1](t1: Tensor[T1, V1], t2: Tensor[T2, V], t3: Tensor[T3, V]): (Tensor[Out, V1], Tensor[Out, V], Tensor[Out, V]) - - object Broadcast3 extends Broadcast3LowPriority: - - /** `t2` and `t3` broadcast against each other, `t1` already has their common shape. */ - given valuesBroadcast[O <: Tuple, T2 <: Tuple, T3 <: Tuple, V](using - bc: Broadcast[T2, T3, V] { type Out = O } - ): Broadcast3[O, T2, T3, V] with - type Out = O - val labelsOut = bc.labelsOut - def broadcast[V1](t1: Tensor[O, V1], t2: Tensor[T2, V], t3: Tensor[T3, V]) = - val (bt2, bt3) = bc.broadcast(t2, t3) - (t1, bt2, bt3) - - /** `t2` and `t3` broadcast against each other, `t1` is broadcast to their common shape. */ - given conditionAndValuesBroadcast[T1 <: Tuple: Labels, T2 <: Tuple, T3 <: Tuple, O <: Tuple, V](using - bc: Broadcast[T2, T3, V] { type Out = O }, - ev: StrictSubset[T1, O] - ): Broadcast3[T1, T2, T3, V] with - type Out = O - val labelsOut = bc.labelsOut - def broadcast[V1](t1: Tensor[T1, V1], t2: Tensor[T2, V], t3: Tensor[T3, V]) = - given Labels[O] = bc.labelsOut - val (bt2, bt3) = bc.broadcast(t2, t3) - (t1.broadcastTo[O](bt2.shape), bt2, bt3) - - /** `t2` and `t3` have the same shape, only `t1` is broadcast to it. */ - given conditionBroadcast[T1 <: Tuple: Labels, T <: Tuple: Labels, V](using - ev: StrictSubset[T1, T] - ): Broadcast3[T1, T, T, V] with - type Out = T - val labelsOut = summon[Labels[T]] - def broadcast[V1](t1: Tensor[T1, V1], t2: Tensor[T, V], t3: Tensor[T, V]) = - (t1.broadcastTo[T](t2.shape), t2, t3) - - trait Broadcast3LowPriority: - - /** `t2` and `t3` are both broadcast to the shape of `t1`. */ - given valuesBroadcastToFirst[T1 <: Tuple: Labels, T2 <: Tuple: Labels, T3 <: Tuple: Labels, V](using - ev2: StrictSubset[T2, T1], - ev3: StrictSubset[T3, T1] - ): Broadcast3[T1, T2, T3, V] with - type Out = T1 - val labelsOut = summon[Labels[T1]] - def broadcast[V1](t1: Tensor[T1, V1], t2: Tensor[T2, V], t3: Tensor[T3, V]) = - (t1, t2.broadcastTo[T1](t1.shape), t3.broadcastTo[T1](t1.shape)) - -end TensorOpsUtil diff --git a/core/src/main/scala/dimwit/tensortree/TensorTree.scala b/core/src/main/scala/dimwit/tensortree/TensorTree.scala index c23fb675..8109dd08 100644 --- a/core/src/main/scala/dimwit/tensortree/TensorTree.scala +++ b/core/src/main/scala/dimwit/tensortree/TensorTree.scala @@ -106,28 +106,22 @@ object TensorTree: // extends TensorTreeLowPriority: */ given tensor[Q <: Tuple, V](using n: Labels[Q]): TensorTree[Tensor[Q, V]] with def map(t: Tensor[Q, V], f: [T <: Tuple, V2] => (Labels[T]) ?=> (Tensor[T, V2] => Tensor[T, V2])): Tensor[Q, V] = - import TensorOps.retag - f[Q, V](using n)(t.retag[Q](using n)) + f[Q, V](using n)(t) def mapWithName(t: Tensor[Q, V], f: [T <: Tuple, V2] => (Labels[T]) ?=> ((String, Tensor[T, V2]) => Tensor[T, V2]), path: String = ""): Tensor[Q, V] = - import TensorOps.retag - f[Q, V](using n)(path, t.retag[Q](using n)) + f[Q, V](using n)(path, t) def mapLeaves[A](t: Tensor[Q, V], f: [T <: Tuple, V2] => (Labels[T]) ?=> (Tensor[T, V2] => A)): Iterator[A] = - import TensorOps.retag - Iterator(f[Q, V](using n)(t.retag[Q](using n))) + Iterator(f[Q, V](using n)(t)) def foreach(t: Tensor[Q, V], f: [T <: Tuple, V2] => (Labels[T]) ?=> (Tensor[T, V2] => Unit)): Unit = - import TensorOps.retag - f[Q, V](using n)(t.retag[Q](using n)) + f[Q, V](using n)(t) def foreachWithName(t: Tensor[Q, V], f: [T <: Tuple, V2] => (Labels[T]) ?=> ((String, Tensor[T, V2]) => Unit), path: String = ""): Unit = - import TensorOps.retag - f[Q, V](using n)(path, t.retag[Q](using n)) + f[Q, V](using n)(path, t) def zipMap(p1: Tensor[Q, V], p2: Tensor[Q, V], f: [T <: Tuple, V2] => (Labels[T]) ?=> ((Tensor[T, V2], Tensor[T, V2]) => Tensor[T, V2])): Tensor[Q, V] = - import TensorOps.retag - f[Q, V](using n)(p1.retag[Q](using n), p2.retag[Q](using n)) + f[Q, V](using n)(p1, p2) def toPyTree(p: Tensor[Q, V]): Jax.PyAny = p.jaxValue def fromPyTree(pyVal: Jax.PyAny): Tensor[Q, V] = Tensor(pyVal.as[Jax.PyDynamic]) diff --git a/core/src/main/scala/dimwit/tensortree/TreeOf.scala b/core/src/main/scala/dimwit/tensortree/TreeOf.scala index 7b13145a..2dc07cad 100644 --- a/core/src/main/scala/dimwit/tensortree/TreeOf.scala +++ b/core/src/main/scala/dimwit/tensortree/TreeOf.scala @@ -1,6 +1,6 @@ package dimwit.tensortree -import dimwit.tensor.TensorOps.* +import dimwit.* import dimwit.tensor.* import scala.NamedTuple.NamedTuple @@ -129,12 +129,12 @@ object TreeOf: def `//!`(p2: Tensor0[V]): P = p1.map([T <: Tuple] => (n: Labels[T]) ?=> (a: Tensor[T, V]) => a /! p2) extension [P: TensorTree, V](p1: P)(using TreeOf[P, V], NotGiven[P <:< Tensor[?, ?]])(using IsFloating[V]) - def sqrt: P = p1.map([T <: Tuple] => (n: Labels[T]) ?=> (a: Tensor[T, V]) => TensorOps.sqrt(a)) + def sqrt: P = p1.map([T <: Tuple] => (n: Labels[T]) ?=> (a: Tensor[T, V]) => Tensor.sqrt(a)) def pow(exponent: Float): P = pow(Tensor0(VType[V])(exponent)) - def pow(exponent: Tensor0[V]): P = p1.map([T <: Tuple] => (n: Labels[T]) ?=> (a: Tensor[T, V]) => TensorOps.pow(a)(exponent)) - def scale(scalar: Tensor0[V]): P = p1.map([T <: Tuple] => (n: Labels[T]) ?=> (a: Tensor[T, V]) => TensorOps.scale(a)(scalar)) - def sign: P = p1.map([T <: Tuple] => (n: Labels[T]) ?=> (a: Tensor[T, V]) => TensorOps.sign(a)) + def pow(exponent: Tensor0[V]): P = p1.map([T <: Tuple] => (n: Labels[T]) ?=> (a: Tensor[T, V]) => Tensor.pow(a, exponent)) + def scale(scalar: Tensor0[V]): P = p1.map([T <: Tuple] => (n: Labels[T]) ?=> (a: Tensor[T, V]) => Tensor.scale(a, scalar)) + def sign: P = p1.map([T <: Tuple] => (n: Labels[T]) ?=> (a: Tensor[T, V]) => Tensor.sign(a)) def fillCopy(value: Float): P = p1.map([T <: Tuple] => (n: Labels[T]) ?=> (a: Tensor[T, V]) => Tensor(a.shape, VType[V]).fill(value)) diff --git a/core/src/test/scala/dimwit/package.scala b/core/src/test/scala/dimwit/package.scala index 4bc580bb..904c66ed 100644 --- a/core/src/test/scala/dimwit/package.scala +++ b/core/src/test/scala/dimwit/package.scala @@ -25,7 +25,7 @@ def approxEqual[T <: Tuple: Labels](right: Tensor[T, Float32], tolerance: Float new Matcher[Tensor[T, Float32]]: def apply(left: Tensor[T, Float32]): MatchResult = - val areEqual = (left `approxEquals` (right, tolerance)).item + val areEqual = left.approxEquals(right, tolerance).item lazy val diffMsg = if areEqual then "" else s"Max diff: ${(left - right).abs.max}" MatchResult( diff --git a/core/src/test/scala/dimwit/tensor/TensorCovarianceSuite.scala b/core/src/test/scala/dimwit/tensor/TensorCovarianceSuite.scala index 398518f3..8ce44e70 100644 --- a/core/src/test/scala/dimwit/tensor/TensorCovarianceSuite.scala +++ b/core/src/test/scala/dimwit/tensor/TensorCovarianceSuite.scala @@ -1,6 +1,7 @@ package dimwit.tensor import dimwit.* +import dimwit.Conversions.given import scala.collection.View.Empty class TensorCovarianceSuite extends DimwitTest: diff --git a/core/src/test/scala/dimwit/tensor/TensorOpsAlongAxisSuite.scala b/core/src/test/scala/dimwit/tensor/TensorOpsAlongAxisSuite.scala index 7dac4f00..a9e97ba7 100644 --- a/core/src/test/scala/dimwit/tensor/TensorOpsAlongAxisSuite.scala +++ b/core/src/test/scala/dimwit/tensor/TensorOpsAlongAxisSuite.scala @@ -194,3 +194,19 @@ class TensorOpsAlongAxisSuite extends DimwitTest: it("roll with Tensor1.roll through vapply"): unsorted.vapply(Axis[B])(Tensor1.roll(1)) shouldEqual unsorted.roll(Axis[B], shift = 1) + + describe("Function forms (Tensor.op(t, axis) is t.op(axis))"): + it("numeric ops"): + Tensor.argsort(unsorted, Axis[B]) shouldEqual unsorted.argsort(Axis[B]) + Tensor.sort(unsorted, Axis[A]) shouldEqual unsorted.sort(Axis[A]) + Tensor.cumsum(unsorted, Axis[B]) shouldEqual unsorted.cumsum(Axis[B]) + Tensor.cumprod(unsorted, Axis[B]) shouldEqual unsorted.cumprod(Axis[B]) + Tensor.cummax(unsorted, Axis[B]) shouldEqual unsorted.cummax(Axis[B]) + Tensor.cummin(unsorted, Axis[A]) shouldEqual unsorted.cummin(Axis[A]) + Tensor.diff(unsorted, Axis[B]) shouldEqual unsorted.diff(Axis[B]) + Tensor.roll(unsorted, Axis[B], shift = 1) shouldEqual unsorted.roll(Axis[B], shift = 1) + + it("floating ops"): + Tensor.logcumsumexp(unsorted, Axis[B]) shouldEqual unsorted.logcumsumexp(Axis[B]) + Tensor.softmax(unsorted, Axis[B]) shouldEqual unsorted.softmax(Axis[B]) + Tensor.logSoftmax(unsorted, Axis[A]) shouldEqual unsorted.logSoftmax(Axis[A]) diff --git a/core/src/test/scala/dimwit/tensor/TensorOpsBroadcastSuite.scala b/core/src/test/scala/dimwit/tensor/TensorOpsBroadcastSuite.scala index 4909e27a..43d16682 100644 --- a/core/src/test/scala/dimwit/tensor/TensorOpsBroadcastSuite.scala +++ b/core/src/test/scala/dimwit/tensor/TensorOpsBroadcastSuite.scala @@ -222,3 +222,75 @@ class TensorOpsBroadcastSuite extends DimwitTest: val ab = Tensor2(Axis[A], Axis[B]).fromArray(Array(Array(1.0f, 2.0f))) val bc = Tensor2(Axis[B], Axis[C]).fromArray(Array(Array(10.0f), Array(20.0f))) "ab +! bc" shouldNot compile // TODO add support for this + + describe("Function forms (Tensor.op_!(t1, t2) is t1 op! t2)"): + val bAB = Tensor2(Axis[A], Axis[B]).fromArray(Array(Array(true, false), Array(false, true))) + val bA = Tensor1(Axis[A]).fromArray(Array(true, false)) + + it("arithmetic"): + Tensor.add_!(tAB, tA) shouldEqual (tAB +! tA) + Tensor.subtract_!(tA, tAB) shouldEqual (tA -! tAB) + Tensor.multiply_!(tAB, tA) shouldEqual (tAB *! tA) + Tensor.divide_!(tAB, tA) shouldEqual (tAB /! tA) + Tensor.mod_!(iAB, iA) shouldEqual (iAB %! iA) + + it("comparisons"): + Tensor.less_!(tAB, tA) shouldEqual (tAB `! tA) + Tensor.greaterEqual_!(tAB, tA) shouldEqual (tAB >=! tA) + Tensor.equal_!(iAB, iA) shouldEqual iAB.elementEquals_!(iA) + + it("logical"): + Tensor.logicalAnd_!(bAB, bA) shouldEqual (bAB and_! bA) + Tensor.logicalOr_!(bAB, bA) shouldEqual (bAB or_! bA) + Tensor.logicalXor_!(bAB, bA) shouldEqual (bAB xor_! bA) + + it("maximum and minimum"): + val tA25 = Tensor1(Axis[A]).fromArray(Array(15.0f, 35.0f)) + maximum_!(tAB, tA25) shouldEqual Tensor.like(tAB).fromArray(Array(15.0f, 20.0f, 35.0f, 40.0f)) + minimum_!(tA25, tAB) shouldEqual Tensor.like(tAB).fromArray(Array(10.0f, 15.0f, 30.0f, 35.0f)) + + it("scale"): + Tensor.scale(tAB, Tensor0(2.0f)) shouldEqual tAB.scale(Tensor0(2.0f)) + + it("approxEquals and approxElementEquals"): + val nearA = Tensor1(Axis[A]).fromArray(Array(10.0000001f, 40.0000001f)) + tAB.approxElementEquals_!(nearA) shouldEqual Tensor(tAB.shape).fromArray(Array(true, false, false, true)) + Tensor.approxElementEquals_!(tAB, nearA) shouldEqual tAB.approxElementEquals_!(nearA) + tAB.approxEquals_!(nearA).item shouldBe false + Tensor.like(tAB).fill(3.0f).approxEquals_!(Tensor0(3.0000001f)).item shouldBe true + Tensor.approxEquals_!(tAB, Tensor0(10.0f)) shouldEqual tAB.approxEquals_!(Tensor0(10.0f)) + + describe("No implicit broadcasting without !"): + val tAB22 = Tensor2(Axis[A], Axis[B]).fromArray(Array(Array(1.0f, 2.0f), Array(3.0f, 4.0f))) + val tAB12 = Tensor2(Axis[A], Axis[B]).fromArray(Array(Array(10.0f, 20.0f))) + + it("same labels but different extents fail fast"): + an[IllegalArgumentException] should be thrownBy (tAB22 + tAB12) + an[IllegalArgumentException] should be thrownBy (tAB22 / tAB12) + an[IllegalArgumentException] should be thrownBy (tAB22 < tAB12) + an[IllegalArgumentException] should be thrownBy (tAB22 === tAB12) + an[IllegalArgumentException] should be thrownBy maximum(tAB22, tAB12) + an[IllegalArgumentException] should be thrownBy where(tAB22 > tAB22, tAB22, tAB12) + an[IllegalArgumentException] should be thrownBy tAB22.approxElementEquals(tAB12) + + it("=== and ===!"): + (tAB === tAB).item shouldBe true + Tensor.arrayEqual(tAB, tAB) shouldEqual (tAB === tAB) + (Tensor.like(tAB).fill(3.0f) ===! Tensor0(3.0f)).item shouldBe true + (tAB ===! tA).item shouldBe false + Tensor.arrayEqual_!(tAB, tA) shouldEqual (tAB ===! tA) + + describe("Scalar first"): + it("computes scalar op tensor, not tensor op scalar"): + (10 -! iAB) shouldEqual (Tensor0(10) -! iAB) + (10 % Tensor0(3)) shouldEqual Tensor0(1) + (10 %! iAB) shouldEqual (Tensor0(10) %! iAB) + (2.0 /! tAB) shouldEqual (Tensor0(2.0f) /! tAB) + (25.0 ` 5, Axis[In] -> 2)).fill(1.0f) + val kernel = Tensor(Shape(Axis[S1] -> 3, Axis[In] -> 2, Axis[Out] -> 4)).fill(0.5f) + Tensor.conv1d(input, kernel, stride = 2, padding = Padding.VALID) shouldEqual input.conv1d(kernel, stride = 2, padding = Padding.VALID) + val output = input.conv1d(kernel) + Tensor.transposeConv1d(output, kernel) shouldEqual output.transposeConv1d(kernel) + + it("conv2d and transposeConv2d"): + val input = Tensor(Shape(Axis[S1] -> 4, Axis[S2] -> 4, Axis[In] -> 2)).fill(1.0f) + val kernel = Tensor(Shape(Axis[S1] -> 2, Axis[S2] -> 2, Axis[In] -> 2, Axis[Out] -> 3)).fill(0.5f) + Tensor.conv2d(input, kernel) shouldEqual input.conv2d(kernel) + val output = input.conv2d(kernel) + Tensor.transposeConv2d(output, kernel) shouldEqual output.transposeConv2d(kernel) + + it("conv3d and transposeConv3d"): + val input = Tensor(Shape(Axis[S1] -> 3, Axis[S2] -> 3, Axis[S3] -> 3, Axis[In] -> 2)).fill(1.0f) + val kernel = Tensor(Shape(Axis[S1] -> 2, Axis[S2] -> 2, Axis[S3] -> 2, Axis[In] -> 2, Axis[Out] -> 3)).fill(0.5f) + Tensor.conv3d(input, kernel) shouldEqual input.conv3d(kernel) + val output = input.conv3d(kernel) + Tensor.transposeConv3d(output, kernel) shouldEqual output.transposeConv3d(kernel) diff --git a/core/src/test/scala/dimwit/tensor/TensorOpsElementwiseSuite.scala b/core/src/test/scala/dimwit/tensor/TensorOpsElementwiseSuite.scala index e69f05cc..2df7d639 100644 --- a/core/src/test/scala/dimwit/tensor/TensorOpsElementwiseSuite.scala +++ b/core/src/test/scala/dimwit/tensor/TensorOpsElementwiseSuite.scala @@ -116,6 +116,10 @@ class TensorOpsElementwiseSuite extends DimwitTest: t2.approxEquals(t2Near).item shouldBe true t2.approxElementEquals(t2Near).all.item shouldBe true + it("approxElementEquals is elementwise"): + val t2Partial = Tensor.like(t2).fromArray(Array(-1.0f, 0.5f, 1.0f, 4.0f)) + t2.approxElementEquals(t2Partial) shouldEqual Tensor.like(b2).fromArray(Array(true, false, true, true)) + describe("Int ops (Tensor2)"): it("abs"): @@ -205,3 +209,41 @@ class TensorOpsElementwiseSuite extends DimwitTest: f2.asBool shouldEqual Tensor(f2.shape).fromArray(Array(true, false, true, true)) f2.asInt32 shouldEqual Tensor(f2.shape).fromArray(Array(-1, 0, 0, 2)) f2.asFloat32 shouldEqual f2 + + describe("Function forms (Tensor.op(t) is t.op)"): + + // values in (0, 1), so that log, sqrt, arcsin and arccos are defined + val p2 = Tensor.like(t2).fromArray(Array(0.1f, 0.2f, 0.3f, 0.4f)) + + it("numeric ops"): + Tensor.abs(t2) shouldEqual t2.abs + Tensor.sign(t2) shouldEqual t2.sign + Tensor.clip(t2, Tensor0(0.0f), Tensor0(1.0f)) shouldEqual t2.clip(Tensor0(0.0f), Tensor0(1.0f)) + Tensor.pow(t2, Tensor0(2.0f)) shouldEqual t2.pow(Tensor0(2.0f)) + Tensor.abs(i2) shouldEqual i2.abs + + it("floating ops"): + Tensor.sqrt(p2) shouldEqual p2.sqrt + Tensor.exp(p2) shouldEqual p2.exp + Tensor.log(p2) shouldEqual p2.log + Tensor.sin(p2) shouldEqual p2.sin + Tensor.cos(p2) shouldEqual p2.cos + Tensor.tanh(p2) shouldEqual p2.tanh + Tensor.arcsin(p2) shouldEqual p2.arcsin + Tensor.arccos(p2) shouldEqual p2.arccos + Tensor.arctan(p2) shouldEqual p2.arctan + Tensor.floor(t2) shouldEqual t2.floor + Tensor.ceil(t2) shouldEqual t2.ceil + Tensor.round(t2) shouldEqual t2.round + Tensor.isnan(t2) shouldEqual t2.isnan + Tensor.isfinite(t2) shouldEqual t2.isfinite + Tensor.nanToNum(t2)(Tensor0(0.0f)) shouldEqual t2.nanToNum(Tensor0(0.0f)) + + it("activation functions"): + Tensor.sigmoid(t2) shouldEqual t2.sigmoid + Tensor.relu(t2) shouldEqual t2.relu + Tensor.gelu(t2) shouldEqual t2.gelu + + it("boolean ops"): + Tensor.all(b2) shouldEqual b2.all + Tensor.any(b2) shouldEqual b2.any diff --git a/core/src/test/scala/dimwit/tensor/TensorOpsReductionSuite.scala b/core/src/test/scala/dimwit/tensor/TensorOpsReductionSuite.scala index 6db8d583..b04eccac 100644 --- a/core/src/test/scala/dimwit/tensor/TensorOpsReductionSuite.scala +++ b/core/src/test/scala/dimwit/tensor/TensorOpsReductionSuite.scala @@ -1,6 +1,7 @@ package dimwit.tensor import dimwit.* +import dimwit.Conversions.given class TensorOpsReductionSuite extends DimwitTest: @@ -32,7 +33,7 @@ class TensorOpsReductionSuite extends DimwitTest: it("=== (Tensor0[Boolean])"): (t2 === t2).item shouldBe true - (t2 === (t2 *! Tensor0(0.0f))).item shouldBe false + (t2 === (t2 *! 0.0f)).item shouldBe false describe("Reduction Ops"): it("sum"): @@ -152,3 +153,26 @@ class TensorOpsReductionSuite extends DimwitTest: t2.approxEquals(t2Near).item shouldBe true val t2Far = t2 *! Tensor0(1.1f) t2.approxEquals(t2Far).item shouldBe false + Tensor.approxEquals(t2, t2Near) shouldEqual t2.approxEquals(t2Near) + + describe("Function forms (Tensor.op(t, ...) is t.op(...))"): + it("numeric reductions"): + Tensor.sum(t2) shouldEqual t2.sum + Tensor.sum(t2, Axis[A]) shouldEqual t2.sum(Axis[A]) + Tensor.sum(t2, (Axis[A], Axis[B])) shouldEqual t2.sum((Axis[A], Axis[B])) + Tensor.max(t2) shouldEqual t2.max + Tensor.max(t2, Axis[B]) shouldEqual t2.max(Axis[B]) + Tensor.min(t2) shouldEqual t2.min + Tensor.min(t2, Axis[B]) shouldEqual t2.min(Axis[B]) + Tensor.argmax(t2, Axis[B]) shouldEqual t2.argmax(Axis[B]) + Tensor.argmin(t2, Axis[B]) shouldEqual t2.argmin(Axis[B]) + + it("floating reductions"): + Tensor.mean(t2) shouldEqual t2.mean + Tensor.mean(t2, Axis[A]) shouldEqual t2.mean(Axis[A]) + Tensor.std(t2, Axis[B]) shouldEqual t2.std(Axis[B]) + Tensor.quantile(t2, 0.5f) shouldEqual t2.quantile(0.5f) + Tensor.quantile(t2, 0.5f, Axis[B]) shouldEqual t2.quantile(0.5f, Axis[B]) + Tensor.median(t2, Axis[B]) shouldEqual t2.median(Axis[B]) + Tensor.nanmean(t2, Axis[B]) shouldEqual t2.nanmean(Axis[B]) + Tensor.nanmedian(t2, Axis[B]) shouldEqual t2.nanmedian(Axis[B])