Add commonly-used methods on tensors - #164
marcelluethi wants to merge 9 commits into
Conversation
|
|
||
| /** 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)) |
There was a problem hiding this comment.
I would not support the default to the last axis.
There was a problem hiding this comment.
I am not sure about it. On one hand it is a confusing default behavior. On the other hand it gracefully handles the case of Tensor1. As far as I know, we cannot have separate extension methods with the same name on Tensor1 and generic Tensor. Argsort, argmin, etc all have this kind of Axis less version.
There was a problem hiding this comment.
Sorry, I did not express my review clearly before (too early in the morning :D)
I would NOT support the default case. The user should always be explicit about which axis the sort is applied to:
val t: Tensor2[Batch, TimeStep, Int32] = ???
t.sort(Axis[TimeStep])Even if Feature is the last dimension. DimWit works completely without positional assumptions; suddenly having a default here is wrong.
Actually, we should be more extreme and define sort only on Tensor1. Then the above statement must be:
t.vapply(Axis[TimeStep])(_.sort)
// or
t.vmap(Axis[Batch])(_.sort)
t.sort // compile-error => axis param missingThis would be identical to how linear layers or softmax work now.
val t: Tensor2[Batch, Feature, Float32] = ???
t.vmap(Axis[Batch])(linearLayer)
t.vapply(Axis[Feature])(softmax)def softmax[L: Label, V: IsFloating](t: Tensor1[L, V]): Tensor1[L, V] =
liftPyTensor(Jax.jnn.softmax(toPyTensor(t), axis = 0)).sort is a function on Tensor1: Taking a vector and sorting that vector. It does not know anything about higher-dimensional tensors. This is the strict and minimal conceptual scope of .sort.
Note that some functions are more general than their minimal scope, like .dot and relu. So your version of sort wouldn't be the only one, but I think we should keep scopes very strict, especially for less common methods. With application to higher tensors with vapply and vmap.
There was a problem hiding this comment.
Additionally, sort is not a ReductionOps. Same for argsort actually...
|
|
||
| /** 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)) |
There was a problem hiding this comment.
I would not support the default to the last axis.
There was a problem hiding this comment.
Good catch. The default should not be axis=-1 but rather just `Tensor(Jax.jnp.cumsum(t.jaxValue) if we want to be consistent. Fixed in the last commit
| 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)) |
There was a problem hiding this comment.
This should provide arguments for nan=0.0, posinf=None, neginf=None that it passes to JAX. I would also consider removing the default for nan, so the user must explicitly specify 0.0.
https://docs.jax.dev/en/latest/_autosummary/jax.numpy.nan_to_num.html
There was a problem hiding this comment.
Maybe like this:
def nanToNum(valueForNan: Double, valueForPosInf: Double, valueForPosNegInf: Double): Tensor[T, V]
def nanToNum(valueForNan: Double): Tensor[T, V] = nanToNum(valueForNan, valueForNan, valueForNan)
```|
|
||
| /** 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)) |
There was a problem hiding this comment.
This would return a flattened vector, and the type Tensor[T, V] is incorrect. See the axis comment at:
https://docs.jax.dev/en/latest/_autosummary/jax.numpy.cumsum.html
We could change the return value to Tensor1[R, V] with merger: AxesMerger.Aux; see flatten to fix this, or not support this (for now).
Actually, should cumsum just be an operation on Tensor1? Similar to sort.
t.vapply(Axis[A])(_.cumsum)This would allow:
t.flatten.cumsumFor the default flatten case.
There was a problem hiding this comment.
Motivation is similar: cumsum is an operation over a list of values, which is a Vector / Tensor1 in tensorland.
There was a problem hiding this comment.
If we decide what to do here, do the same for cumprod. And diff(I think).
|
@marcelluethi I made some targeted comments for specific lines of code. Overall, my view is this: JAX has many functions that are conceptually functions on Tensor1, but are defined for higher-order tensors, taking an An illustrative example for this is https://docs.jax.dev/en/latest/_autosummary/jax.numpy.cross.html |
|
Maybe this approach is the best of both worlds (let's discuss tomorrow 👍 ): Operations, which allow no or multiple axes in JAX, are define on extension (t: Tensor[...])
def sum(...)
t.sum
t.sum(Axis[A])Operations, which require exactly one axis in JAX, are functions on vectors and should be defined on // main branch
extension (t: Tensor1[...])
def softmax: Tensor1[...] = ... // correct scoping
// not yet on main branch
extension (t: Tensor[...])
def softmax(axis: Axis[A]): Tensor[...] = t.vapply(axis)(softmax) // syntax sugar
t.vapply(Axis[A])(softmax) // current
t.softmax(Axis[A]) // new option to do thisThe implementation might be difficult due to This consideration applys to |
|
So I understand that there are operations like object Tensor1Transform:
def softmax[V : Floating](t : Tensor1[A, V]) : Tensor1[A, V] = ???
def diff[V](t : Tensor1[A, V]) : Tensor1[A, V] = ???and in extension t : Tensor[T, V : Floating]
def softmax[L](axis : Axis[L]) : Tensor[T, V] = vapply(axis)(Tensor1Transform.softmax)To use the methods, the user has two options:
|
|
@benikm91 I started doing the refactoring discussed above and incorporating your comments. |
| def isfinite: Tensor[T, Bool] = Tensor(Jax.jnp.isfinite(t.jaxValue)) | ||
|
|
||
| /** replaces NaN by `nan`, +inf by `posInf` and -inf by `negInf`. | ||
| * By default, NaN becomes 0 and ±inf become the largest/smallest finite value of the dtype. |
| /** 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) | ||
| def sort: Tensor[T, V] = Tensor(Jax.jnp.sort(t.jaxValue)) |
There was a problem hiding this comment.
Remove default sort (last axis by convention).
| /** replaces NaN by `nan`, +inf by `posInf` and -inf by `negInf`. | ||
| * By default, NaN becomes 0 and ±inf become the largest/smallest finite value of the dtype. | ||
| */ | ||
| def nanToNum(using |
There was a problem hiding this comment.
This IsFloating[V] can be removed. IsFloating evidence already in extention method.
| // 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)) |
There was a problem hiding this comment.
This changed syntax to t.relu. We need relu(t), as this is more natural for activation functions. I propose to put this in DimWit, here, but I am fine to move this also to DeepWit only.
DimWits wrapping of the basic tensor methods supported in jax is rather patchy and was introduced on a per need basis.
This PR proposes to add some more, frequently used tensor operations. Each of them is just a one-liner, wrapping the underlying jax function.
The methods added are:
sortcumsum,cumproddifffloor,ceil,roundarcsin,arccos,arctanisnan,isfinite,nanToNummod(with a%operator)