Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 8 additions & 9 deletions AGENTS.md
Original file line number Diff line number Diff line change
Expand Up @@ -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:
//
Expand All @@ -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:
Expand Down Expand Up @@ -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](
Expand Down Expand Up @@ -1422,21 +1421,21 @@ 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²:
// dimwit.tensortree.TensorTree[dimwit.tensor.Tensor0[V²]]): (T1, T2, T3) =>
// dimwit.autodiff.Grad[(T1, T2, T3)]
// [T1², T2², V³]
// (f: (T1², T2²) => dimwit.tensor.Tensor0[V³])
// (using evidence$1³: dimwit.tensor.TensorOps.IsFloating[V³],
// (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²) =>
Expand Down
2 changes: 1 addition & 1 deletion core/src/main/scala/dimwit/autodiff/Autodiff.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
4 changes: 2 additions & 2 deletions core/src/main/scala/dimwit/linalg/LinearAlgebra.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
30 changes: 27 additions & 3 deletions core/src/main/scala/dimwit/package.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion core/src/main/scala/dimwit/random/Random.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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.*

Expand Down
9 changes: 4 additions & 5 deletions core/src/main/scala/dimwit/stats/Distributions.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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:
Expand All @@ -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]:
Expand Down
6 changes: 3 additions & 3 deletions core/src/main/scala/dimwit/tensor/ArrayWriter.scala
Original file line number Diff line number Diff line change
@@ -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
Expand Down
90 changes: 90 additions & 0 deletions core/src/main/scala/dimwit/tensor/Broadcast.scala
Original file line number Diff line number Diff line change
@@ -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))
20 changes: 20 additions & 0 deletions core/src/main/scala/dimwit/tensor/Convolution.scala
Original file line number Diff line number Diff line change
@@ -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])
6 changes: 3 additions & 3 deletions core/src/main/scala/dimwit/tensor/DType.scala
Original file line number Diff line number Diff line change
@@ -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
Expand Down
68 changes: 17 additions & 51 deletions core/src/main/scala/dimwit/tensor/Tensor.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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]
Expand Down Expand Up @@ -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.
Expand Down
2 changes: 1 addition & 1 deletion core/src/main/scala/dimwit/tensor/VType.scala
Original file line number Diff line number Diff line change
@@ -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)
Expand Down
Loading
Loading