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
2 changes: 1 addition & 1 deletion .scalafmt.conf
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
version = "3.9.8"
version = "3.10.3"
runner.dialect = scala3
maxColumn = 999
rewrite.scala3.convertToNewSyntax = true
Expand Down
24 changes: 9 additions & 15 deletions AGENTS.md
Original file line number Diff line number Diff line change
Expand Up @@ -266,9 +266,7 @@ val descendingSpaced = Tensor1(Axis[A]).linspace(Tensor0(1.0f), Tensor0(0.0f), 3
val samples = Tensor1(Axis[B]).fromArray(Array(4.0f, 2.0f, 8.0f))
val binEdges = Tensor1(Axis[A]).linspace(samples.min, samples.max, 4)

// The typed factory fixes the value type; with dimwit.Conversions.given
// plain literals are converted to Tensor0 of that type
import dimwit.Conversions.given
// The typed factory fixes the value type; plain literals are converted to Tensor0 of that type
val halfSpaced = Tensor1(Axis[A], VType[Float16]).linspace(0.0f, 1.0f, 5)
```

Expand Down Expand Up @@ -474,10 +472,10 @@ val wrong = t.sum(Axis[C])
// Conflicting definitions:
// val t:
// dimwit.tensor.Tensor2[MdocApp0.this.A, MdocApp0.this.B,
// dimwit.tensor.DType.Float32] in class MdocApp0 at line 71 and
// dimwit.tensor.DType.Float32] in class MdocApp0 at line 70 and
// val t:
// dimwit.tensor.Tensor2[MdocApp0.this.A, MdocApp0.this.B,
// dimwit.tensor.DType.Float32] in class MdocApp0 at line 117
// dimwit.tensor.DType.Float32] in class MdocApp0 at line 116
//
```

Expand Down Expand Up @@ -519,10 +517,10 @@ val wrong = t + 5.0f // Use +! instead
// Conflicting definitions:
// val t:
// dimwit.tensor.Tensor2[MdocApp0.this.A, MdocApp0.this.B,
// dimwit.tensor.DType.Float32] in class MdocApp0 at line 71 and
// dimwit.tensor.DType.Float32] in class MdocApp0 at line 70 and
// val t:
// dimwit.tensor.Tensor2[MdocApp0.this.A, MdocApp0.this.B,
// dimwit.tensor.DType.Float32] in class MdocApp0 at line 126
// dimwit.tensor.DType.Float32] in class MdocApp0 at line 125
//
```

Expand Down Expand Up @@ -612,19 +610,19 @@ val wrong = m1.dot(Axis[B])(m2)
// Conflicting definitions:
// val m1:
// dimwit.tensor.Tensor2[MdocApp1.this.A, MdocApp1.this.B,
// dimwit.tensor.DType.Float32] in class MdocApp1 at line 148 and
// dimwit.tensor.DType.Float32] in class MdocApp1 at line 147 and
// val m1:
// dimwit.tensor.Tensor2[MdocApp1.this.A, MdocApp1.this.B,
// dimwit.tensor.DType.Float32] in class MdocApp1 at line 151
// dimwit.tensor.DType.Float32] in class MdocApp1 at line 150
//
// error:
// Conflicting definitions:
// val m2:
// dimwit.tensor.Tensor2[MdocApp1.this.B, MdocApp1.this.C,
// dimwit.tensor.DType.Float32] in class MdocApp1 at line 149 and
// dimwit.tensor.DType.Float32] in class MdocApp1 at line 148 and
// val m2:
// dimwit.tensor.Tensor2[MdocApp1.this.C, MdocApp1.this.D,
// dimwit.tensor.DType.Float32] in class MdocApp1 at line 152
// dimwit.tensor.DType.Float32] in class MdocApp1 at line 151
//
```

Expand Down Expand Up @@ -1007,7 +1005,6 @@ println(s"Block Hessian shapes: ${h_x1x1.shape}, ${h_x1x2.shape}, ${h_x2x1.shape

```scala
import dimwit.*
import dimwit.Conversions.given
import dimwit.optimizer.{GradientDescent, GradientOptimizer}
import dimwit.random.Random

Expand Down Expand Up @@ -1056,7 +1053,6 @@ val trained = optimizer.iterate(initModelParams)(gradFunc)

```scala
import dimwit.optimizer.Lion
import dimwit.Conversions.given // enables implicit conversion from Float to Tensor[V]

// Lion optimizer with momentum
val lionOptimizer = Lion(learningRate = 1e-3f, beta1 = 0.9f, beta2 = 0.99f, weightDecay = 0.0f)
Expand All @@ -1071,8 +1067,6 @@ val trainedLion = lionOptimizer.iterate(initModelParams)(gradFunc)
### Complete Training Example: Linear Regression

```scala
import dimwit.Conversions.given // enables implicit conversion from Float to Tensor[V]

// Define problem dimensions
trait Sample derives Label
trait InputDim derives Label
Expand Down
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
package dimwit.optimizer

import dimwit.*
import dimwit.Conversions.given
import dimwit.autodiff.*
import dimwit.autodiff.Grad
import dimwit.tensortree.TreeOf
Expand Down
3 changes: 0 additions & 3 deletions core/src/main/scala/dimwit/package.scala
Original file line number Diff line number Diff line change
Expand Up @@ -101,9 +101,6 @@ package object dimwit:
export dimwit.jax.Jit.{jit, jitDonating, jitDonatingUnsafe}
export dimwit.jax.EagerCleanup.eagerCleanup

object Conversions:
export dimwit.tensor.Tensor0.{boolean2BooleanTensor, byte2IntegerTensor, short2IntegerTensor, int2IntegerTensor, long2IntegerTensor, float2FloatingTensor, int2FloatingTensor, double2FloatingTensor}

// Export random object
export dimwit.random.Random
export dimwit.random.Random.Key
Expand Down
30 changes: 5 additions & 25 deletions core/src/main/scala/dimwit/tensor/Tensor.scala
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@ import DType.*
* @param T The shape of the tensor, represented as a tuple of axis labels.
* @param V The data type of the tensor elements.
*/
class Tensor[T <: Tuple: Labels, V] private[dimwit] (
into class Tensor[T <: Tuple: Labels, V] private[dimwit] (

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I did not know that into also works on the class level. This is much more elegant than the solution I found, which would have used into in every signature that has a Tensor[T, V]. Nice.

private[dimwit] val jaxValue: Jax.PyDynamic
):

Expand Down Expand Up @@ -201,6 +201,10 @@ 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 conversions from Scala scalars to a Tensor0, e.g. `t *! 2.0f`. Being members of the companion object, they
// are found without an import, and `into class Tensor` allows them without a feature warning.
export Tensor0Conversions.given

// 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.*
Expand All @@ -224,30 +228,6 @@ type Tensor4[L1, L2, L3, L4, V] = Tensor[(L1, L2, L3, L4), V]
*/
object Tensor0:

given boolean2BooleanTensor[V: IsBoolean]: Conversion[Boolean, Tensor0[V]] with
def apply(value: Boolean): Tensor0[V] = Tensor0(VType[V])(value)

given byte2IntegerTensor[V: IsInteger]: Conversion[Byte, Tensor0[V]] with
def apply(value: Byte): Tensor0[V] = Tensor0(VType[V])(value)

given short2IntegerTensor[V: IsInteger]: Conversion[Short, Tensor0[V]] with
def apply(value: Short): Tensor0[V] = Tensor0(VType[V])(value)

given int2IntegerTensor[V: IsInteger]: Conversion[Int, Tensor0[V]] with
def apply(value: Int): Tensor0[V] = Tensor0(VType[V])(value)

given int2FloatingTensor[V: IsFloating]: Conversion[Int, Tensor0[V]] with
def apply(value: Int): Tensor0[V] = Tensor0(VType[V])(value.toFloat)

given long2IntegerTensor[V: IsInteger]: Conversion[Long, Tensor0[V]] with
def apply(value: Long): Tensor0[V] = Tensor0(VType[V])(value)

given float2FloatingTensor[V: IsFloating]: Conversion[Float, Tensor0[V]] with
def apply(value: Float): Tensor0[V] = Tensor0(VType[V])(value)

given double2FloatingTensor[V: IsFloating]: Conversion[Double, Tensor0[V]] with
def apply(value: Double): Tensor0[V] = Tensor0(VType[V])(value)

object Value0Factory:

def apply(value: Boolean): Tensor0[Bool] = Tensor0(VType[Bool])(value)
Expand Down
35 changes: 35 additions & 0 deletions core/src/main/scala/dimwit/tensor/Tensor0Conversions.scala
Original file line number Diff line number Diff line change
@@ -0,0 +1,35 @@
package dimwit.tensor

import dimwit.tensor.ValueTypeClasses.IsBoolean
import dimwit.tensor.ValueTypeClasses.IsFloating
import dimwit.tensor.ValueTypeClasses.IsInteger

/** Conversions from Scala scalars to a `Tensor0` of the tensor's value type, e.g. for `t *! 2.0f`.
*
* Users never import them: the `Tensor` companion object exports them, so the compiler finds them automatically.
*/
private[dimwit] object Tensor0Conversions:

given boolean2BooleanTensor[V: IsBoolean]: Conversion[Boolean, Tensor0[V]] with
def apply(value: Boolean): Tensor0[V] = Tensor0(VType[V])(value)

given byte2IntegerTensor[V: IsInteger]: Conversion[Byte, Tensor0[V]] with
def apply(value: Byte): Tensor0[V] = Tensor0(VType[V])(value)

given short2IntegerTensor[V: IsInteger]: Conversion[Short, Tensor0[V]] with
def apply(value: Short): Tensor0[V] = Tensor0(VType[V])(value)

given int2IntegerTensor[V: IsInteger]: Conversion[Int, Tensor0[V]] with
def apply(value: Int): Tensor0[V] = Tensor0(VType[V])(value)

given int2FloatingTensor[V: IsFloating]: Conversion[Int, Tensor0[V]] with
def apply(value: Int): Tensor0[V] = Tensor0(VType[V])(value.toFloat)

given long2IntegerTensor[V: IsInteger]: Conversion[Long, Tensor0[V]] with
def apply(value: Long): Tensor0[V] = Tensor0(VType[V])(value)

given float2FloatingTensor[V: IsFloating]: Conversion[Float, Tensor0[V]] with
def apply(value: Float): Tensor0[V] = Tensor0(VType[V])(value)

given double2FloatingTensor[V: IsFloating]: Conversion[Double, Tensor0[V]] with
def apply(value: Double): Tensor0[V] = Tensor0(VType[V])(value)
4 changes: 2 additions & 2 deletions core/src/main/scala/dimwit/tensor/ValueExtensions.scala
Original file line number Diff line number Diff line change
Expand Up @@ -7,8 +7,8 @@ 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
* `t op scalar` works through the implicit conversions exported on the [[Tensor]] companion object, 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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -141,8 +141,6 @@ private[dimwit] object StructuralOps:
object AxisAbsent:
given notContained[T <: Tuple, L](using NotGiven[Tuple.Contains[T, L] =:= true]): AxisAbsent[T, L] = new AxisAbsent[T, L] {}

import Util.*

@benikm91 benikm91 Sep 27, 2026 •

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Note this change is a clean up from the previous PR


object TensorWhere:

/** Returns a new tensor selecting elements from `ifTrue` where the condition is true
Expand Down
1 change: 0 additions & 1 deletion core/src/test/scala/dimwit/autodiff/AutodiffSuite.scala
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
package dimwit.autodiff

import dimwit.*
import dimwit.Conversions.given

/** A parameter tree, declared top level so its Mirror is available. */
case class JacParams(w: Tensor1[A, Float32], b: Tensor1[B, Float32]) derives TensorTree
Expand Down
1 change: 0 additions & 1 deletion core/src/test/scala/dimwit/jax/JitSuite.scala
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
package dimwit.jax

import dimwit.*
import dimwit.Conversions.given
import me.shadaj.scalapy.py

class JitSuite extends DimwitTest:
Expand Down
1 change: 0 additions & 1 deletion core/src/test/scala/dimwit/memory/DimWitMemorySuite.scala
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
package dimwit.memory

import dimwit.*
import dimwit.Conversions.given
import org.scalatest.DoNotDiscover
import scala.compiletime.testing.typeCheckErrors
import scala.compiletime.ops.double
Expand Down
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
package dimwit.optimizer

import dimwit.*
import dimwit.Conversions.given

class GradientOptimizerSuite extends DimwitTest:

Expand Down
1 change: 0 additions & 1 deletion core/src/test/scala/dimwit/python/PyWrapSuite.scala
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
package dimwit.python

import dimwit.*
import dimwit.Conversions.given
import dimwit.python.PyBridge
import dimwit.jax.Jax
import me.shadaj.scalapy.py
Expand Down
1 change: 0 additions & 1 deletion core/src/test/scala/dimwit/stats/DistributionSuite.scala
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
package dimwit.stats

import dimwit.*
import dimwit.Conversions.given
import dimwit.jax.Jax
import dimwit.random.Random
import dimwit.jax.Jax.scipy_stats as jstats
Expand Down
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
package dimwit.tensor

import dimwit.*
import dimwit.Conversions.given
import scala.collection.View.Empty

class TensorCovarianceSuite extends DimwitTest:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -193,7 +193,6 @@ class TensorCreationSuite extends DimwitTest:
Tensor1(Axis[A]).linspace(Tensor0(0.0), Tensor0(1.0), 3).dtype shouldBe DType.Float64

it("typed factory fixes the value type and accepts converted literals"):
import dimwit.Conversions.given
Tensor1(Axis[A], VType[Float16]).linspace(0.0f, 1.0f, 3).dtype shouldBe DType.Float16
Tensor1(Axis[A], VType[Float32]).linspace(0.0f, 1.0f, 3, endpoint = false) shouldEqual Tensor1(Axis[A]).fromArray(Array(0.0f, 1.0f / 3.0f, 2.0f / 3.0f))

Expand Down
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
package dimwit.tensor

import dimwit.*
import dimwit.Conversions.given

class TensorOpsBroadcastSuite extends DimwitTest:

Expand Down
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
package dimwit.tensor

import dimwit.*
import dimwit.Conversions.given

class TensorOpsElementwiseSuite extends DimwitTest:

Expand Down
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
package dimwit.tensor

import dimwit.*
import dimwit.Conversions.given

class TensorOpsFunctionalSuite extends DimwitTest:

Expand Down
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
package dimwit.tensor

import dimwit.*
import dimwit.Conversions.given

class TensorOpsReductionSuite extends DimwitTest:

Expand Down
1 change: 0 additions & 1 deletion core/src/test/scala/dimwit/tensortree/TreeOfSuite.scala
Original file line number Diff line number Diff line change
@@ -1,7 +1,6 @@
package dimwit.tensortree

import dimwit.*
import dimwit.Conversions.given
import dimwit.tensortree.TreeOf.*
import dimwit.tensortree.TreeOf.given
import dimwit.tensortree.TreeOf.ops.*
Expand Down
5 changes: 2 additions & 3 deletions docs/quickstart.md
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,6 @@ Before we start exploring the features of DimWit, let's look at a simple example
```scala
// main imports for basic tensor operations and automatic differentiation
import dimwit.*
import dimwit.Conversions.given
import dimwit.Autodiff.grad // TODO replace with cleaner import after PR is merged
import dimwit.optimizer.GradientDescent // TODO replace with cleaner import after refactoring

Expand Down Expand Up @@ -220,10 +219,10 @@ tensor1 + tensor3
// Conflicting definitions:
// val tensor1:
// dimwit.tensor.Tensor[(MdocApp1.this.A, MdocApp1.this.B),
// dimwit.tensor.DType.Float32] in class MdocApp1 at line 64 and
// dimwit.tensor.DType.Float32] in class MdocApp1 at line 63 and
// val tensor1:
// dimwit.tensor.Tensor[(MdocApp1.this.A, MdocApp1.this.B),
// dimwit.tensor.DType.Float32] in class MdocApp1 at line 68
// dimwit.tensor.DType.Float32] in class MdocApp1 at line 67
//
// val tensor1 = Tensor(Shape(Axis[A] -> 3, Axis[B] -> 2)).fill(1.0f)
// ^
Expand Down
1 change: 0 additions & 1 deletion examples/src/main/scala/dimwit/basic/KMeans.scala
Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
package dimwit.examples.basic.kmeans

import dimwit.Conversions.given
import dimwit.*
import dimwit.random.Random
import dimwit.stats.Normal
Expand Down
Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
package dimwit.examples.basic

import dimwit.Conversions.given
import dimwit.*
import dimwit.autodiff.*
import dimwit.optimizer.GradientDescent
Expand Down
Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
package dimwit.examples.complex.vae

import dimwit.Conversions.given
import dimwit.*
import dimwit.tensortree.TreeOf.*
import dimwit.autodiff.*
Expand Down
1 change: 0 additions & 1 deletion examples/src/main/scala/dimwit/dataset/MNISTLoader.scala
Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
package dimwit.examples.dataset

import dimwit.Conversions.given
import dimwit.*
import me.shadaj.scalapy.py

Expand Down
Loading
Loading