Skip to content
Open
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: 17 additions & 0 deletions core/src/main/scala/dimwit/tensor/tensorops/ElementWiseOps.scala
Original file line number Diff line number Diff line change
Expand Up @@ -126,12 +126,19 @@ object ElementWiseOps:
/** 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])

def +(other: Tensor[T, V]): Tensor[T, V] = add(t, other)
def -(other: Tensor[T, V]): Tensor[T, V] = subtract(t, other)
def *(other: Tensor[T, V]): Tensor[T, V] = multiply(t, other)
def %(other: Tensor[T, V]): Tensor[T, V] = mod(t, other)

// extension methods for the scalar operations.
extension [T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V])
Expand All @@ -143,6 +150,7 @@ object ElementWiseOps:

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)

// extension methods
extension [T <: Tuple: Labels, V: IsNumber](t: Tensor[T, V])
Expand Down Expand Up @@ -173,6 +181,15 @@ object ElementWiseOps:
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 nanToNum: Tensor[T, V] = Tensor(Jax.jnp.nan_to_num(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] =
Expand Down
16 changes: 16 additions & 0 deletions core/src/main/scala/dimwit/tensor/tensorops/ReductionOps.scala
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,22 @@ object ReductionOps:
def argsort[L: Label](axis: Axis[L])(using ev: AxisIndex[T, L]): Tensor[T, Int32] = Tensor(Jax.jnp.argsort(t.jaxValue, axis = ev.index))
def argsort: Tensor[T, Int32] = Tensor(Jax.jnp.argsort(t.jaxValue))

/** sorts the tensor `t` along the specified axis */
def sort[L: Label](axis: Axis[L])(using ev: AxisIndex[T, L]): Tensor[T, V] = Tensor(Jax.jnp.sort(t.jaxValue, axis = ev.index))
def sort: Tensor[T, V] = Tensor(Jax.jnp.sort(t.jaxValue))

/** computes the cumulative sum of the tensor `t` along the specified axis. */
def cumsum[L: Label](axis: Axis[L])(using ev: AxisIndex[T, L]): Tensor[T, V] = Tensor(Jax.jnp.cumsum(t.jaxValue, axis = ev.index))
def cumsum: Tensor[T, V] = Tensor(Jax.jnp.cumsum(t.jaxValue, axis = -1))

/** computes the cumulative product of the tensor `t` along the specified axis. */
def cumprod[L: Label](axis: Axis[L])(using ev: AxisIndex[T, L]): Tensor[T, V] = Tensor(Jax.jnp.cumprod(t.jaxValue, axis = ev.index))
def cumprod: Tensor[T, V] = Tensor(Jax.jnp.cumprod(t.jaxValue, axis = -1))

/** 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 ev: AxisIndex[T, L]): Tensor[T, V] = Tensor(Jax.jnp.diff(t.jaxValue, axis = ev.index))
def diff: Tensor[T, V] = Tensor(Jax.jnp.diff(t.jaxValue))

// ---------------------------------------------------------
// IsFloat operations (IsFloat or IsInt)
// ---------------------------------------------------------
Expand Down
34 changes: 34 additions & 0 deletions core/src/test/scala/dimwit/tensor/TensorOpsElementwiseSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,35 @@ class TensorOpsElementwiseSuite extends DimwitTest:
tZero.cos should approxEqual(Tensor.like(t2).fill(1f))
tZero.tanh should approxEqual(tZero)

it("arcsin/arccos/arctan"):
Tensor.like(t2).fill(0.5f).arcsin should approxEqual(Tensor.like(t2).fill((math.Pi / 6).toFloat), tolerance = 1e-5f)
Tensor.like(t2).fill(0.5f).arccos should approxEqual(Tensor.like(t2).fill((math.Pi / 3).toFloat), tolerance = 1e-5f)
Tensor.like(t2).fill(1.0f).arctan should approxEqual(Tensor.like(t2).fill((math.Pi / 4).toFloat), tolerance = 1e-5f)

it("floor/ceil/round"):
val t = Tensor.like(t2).fromArray(Array(-1.5f, 0.4f, 1.5f, 2.6f))
t.floor should approxEqual(Tensor.like(t2).fromArray(Array(-2.0f, 0.0f, 1.0f, 2.0f)))
t.ceil should approxEqual(Tensor.like(t2).fromArray(Array(-1.0f, 1.0f, 2.0f, 3.0f)))
t.round should approxEqual(Tensor.like(t2).fromArray(Array(-2.0f, 0.0f, 2.0f, 3.0f)))

it("isnan/isfinite"):
val t = Tensor.like(t2).fromArray(Array(Float.NaN, Float.PositiveInfinity, 1.0f, 0.0f))
t.isnan shouldEqual Tensor.like(b2).fromArray(Array(true, false, false, false))
t.isfinite shouldEqual Tensor.like(b2).fromArray(Array(false, false, true, true))

it("nanToNum"):
val t = Tensor.like(t2).fromArray(Array(Float.NaN, Float.PositiveInfinity, Float.NegativeInfinity, 1.0f))
t.nanToNum shouldEqual Tensor.like(t2).fromArray(Array(0.0f, Float.MaxValue, -Float.MaxValue, 1.0f))

it("mod"):
val t = Tensor.like(t2).fromArray(Array(-7.0f, 7.0f, -7.0f, 7.0f))
val divisor = Tensor.like(t2).fromArray(Array(3.0f, 3.0f, -3.0f, -3.0f))
(t % divisor) should approxEqual(Tensor.like(t2).fromArray(Array(2.0f, 1.0f, -1.0f, -2.0f)))

it("mod broadcasting (%!)"):
val t = Tensor1(Axis[A]).fromArray(Array(-7.0f, 7.0f))
(t %! Tensor0(3.0f)) should approxEqual(Tensor1(Axis[A]).fromArray(Array(2.0f, 1.0f)))

it("clip"):
t2.clip(0.0f, 2.0f) should approxEqual(Tensor.like(t2).fromArray(Array(0.0f, 0.0f, 1.0f, 2.0f)))

Expand All @@ -75,6 +104,11 @@ class TensorOpsElementwiseSuite extends DimwitTest:
it("pow"):
i2.pow(Tensor0(3)) shouldEqual Tensor.like(i2).fromArray(Array(-1, 0, 1, 8))

it("mod"):
val t = Tensor.like(i2).fromArray(Array(-7, 7, -7, 7))
val divisor = Tensor.like(i2).fromArray(Array(3, 3, -3, -3))
(t % divisor) shouldEqual Tensor.like(i2).fromArray(Array(2, 1, -1, -2))

it("clip"):
i2.clip(0, 1) shouldEqual Tensor.like(i2).fromArray(Array(0, 0, 1, 1))

Expand Down
82 changes: 82 additions & 0 deletions core/src/test/scala/dimwit/tensor/TensorOpsReductionSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -168,6 +168,88 @@ class TensorOpsReductionSuite extends DimwitTest:
)
)

it("sort"):
val descendingAlongB = Tensor2(Axis[A], Axis[B]).fromArray(
Array(
Array(3.0f, 2.0f, 1.0f),
Array(6.0f, 5.0f, 4.0f)
)
)
descendingAlongB.sort shouldEqual Tensor2(
Axis[A],
Axis[B]
).fromArray(
Array(
Array(1.0f, 2.0f, 3.0f),
Array(4.0f, 5.0f, 6.0f)
)
)

it("sort axis A"):
val descendingAlongA = Tensor2(Axis[A], Axis[B]).fromArray(
Array(
Array(4.0f, 5.0f, 6.0f),
Array(1.0f, 2.0f, 3.0f)
)
)
val res = descendingAlongA.sort(axis = Axis[A])
res shouldEqual Tensor2(
Axis[A],
Axis[B]
).fromArray(
Array(
Array(1.0f, 2.0f, 3.0f),
Array(4.0f, 5.0f, 6.0f)
)
)

it("sort axis B"):
val descendingAlongB = Tensor2(Axis[A], Axis[B]).fromArray(
Array(
Array(3.0f, 2.0f, 1.0f),
Array(6.0f, 5.0f, 4.0f)
)
)
val res = descendingAlongB.sort(axis = Axis[B])
res shouldEqual Tensor2(
Axis[A],
Axis[B]
).fromArray(
Array(
Array(1.0f, 2.0f, 3.0f),
Array(4.0f, 5.0f, 6.0f)
)
)

it("cumsum"):
val res = t2.cumsum(axis = Axis[B])
res shouldEqual Tensor.like(res).fromArray(Array(1.0f, 3.0f, 6.0f, 4.0f, 9.0f, 15.0f))

it("cumsum default axis"):
t2.cumsum shouldEqual t2.cumsum(axis = Axis[B])

it("cumsum axis A"):
val res = t2.cumsum(axis = Axis[A])
res shouldEqual Tensor.like(res).fromArray(Array(1.0f, 2.0f, 3.0f, 5.0f, 7.0f, 9.0f))

it("cumprod"):
val res = t2.cumprod(axis = Axis[B])
res shouldEqual Tensor.like(res).fromArray(Array(1.0f, 2.0f, 6.0f, 4.0f, 20.0f, 120.0f))

it("cumprod default axis"):
t2.cumprod shouldEqual t2.cumprod(axis = Axis[B])

it("diff axis B"):
val res = t2.diff(axis = Axis[B])
res shouldEqual Tensor.like(res).fromArray(Array(1.0f, 1.0f, 1.0f, 1.0f))

it("diff default axis"):
t2.diff shouldEqual t2.diff(axis = Axis[B])

it("diff axis A"):
val res = t2.diff(axis = Axis[A])
res shouldEqual Tensor.like(res).fromArray(Array(3.0f, 3.0f, 3.0f))

describe("Boolean Reductions"):
it("all"):
b2.all shouldEqual Tensor0(false)
Expand Down
Loading