Split tensor ops into functions - #168
Conversation
|
Note: Coding was done by Claude; however, it took several iterations, with me reviewing/ @marcelluethi, so maybe go over the PR description, not the code, for now. The changes are extensive but mostly internal. |
marcelluethi
left a comment
There was a problem hiding this comment.
I haven´t carefully checked all code but went through the most crucial places, like Tensor TensorNExtension, ValueExtension, AlongAxisOps. The new structure makes a lot of sense and I actually like that all functions on Tensor are now only available using the Tensor object and not globally. This should make both the internal organisation as well as the user experience better.
…ions) Every tensor operation now also exists as a function, e.g. Tensor.relu(t), Tensor.sum(t, Axis[A]), Tensor.softmax(t, Axis[A]), Tensor.dot(t1, Axis[A], t2). Each tensorops file has an XOps object (functions, exported on Tensor/Tensor1) and an XExtensions object (extension methods calling XOps, exported from package.scala). All tensorops objects are private[dimwit]. - TensorOps -> ValueTypeClasses, TensorOpsUtil -> dimwit.tensor.Broadcast, ValueOps -> ValueExtensions; Padding/Stride moved to dimwit.tensor - Add broadcasting functions (add_!, less_!, maximum_!, arrayEqual_!, ...), ===!, approxEquals_! and approxElementEquals_! - Scalar-first operators mirror scalar-last for all Scala scalar types - Non-! binary ops fail fast on differing extents instead of JAX silently broadcasting axes of extent 1 - Fix approxElementEquals returning a single boolean (allclose -> isclose) - Remove broken argsort(axes: Tuple), unused retag and the *Scalar functions; quantile takes a Tensor0; add/less/logicalAnd/... are no longer top-level (use operators or Tensor.add etc.)
1b9b9ae to
8b75f9a
Compare
Split tensor ops into functions (
*Ops) and extension methods (*Extensions)Until now most tensor operations existed only as extension methods (
t.relu,t.sum(Axis[A])), with no function you could pass around (e.g. toTreeOf.map). This PR adds a function for every operation and restructurestensoropsconsistently.Structure
tensorops/X.scalanow hasXOps(the functions) andXExtensions(extension methods that only callXOps). They can't share one object: a function and an extension method with the same name have the same erased signature.Tensor.relu(t),Tensor.sum(t, Axis[A]),Tensor.softmax(t, Axis[A]),Tensor.dot(t1, Axis[A], t2),Tensor.conv2d(input, kernel). TheTensor1 => Tensor1functions stay onTensor1(Tensor1.softmax). Users access operations as functions on Tensor/Tensor1 (e.g. Tensor.relu(t)) or as extension methods via import dimwit.* (e.g. t.relu). XOps and XExtensions are private[dimwit] and only structure the code internally.import dimwit.*:maximum,minimum,maximum_!,minimum_!,where,where_!,triu,tril,stack,concatenate,zipvmap.TensorOpsrenamed toValueTypeClasses: it only holdsIsFloating,IsNumber, … now.TensorOpsUtilrenamed todimwit.tensor.Broadcast(public, top-level).PaddingandStride1/2/3moved todimwit.tensor.ValueOpsrenamed toValueExtensions.New
add_!,subtract_!,multiply_!,divide_!,mod_!,less_!, …equal_!,logicalAnd_!/Or_!/Xor_!,maximum_!,minimum_!,arrayEqual_!. The!operators call them.===!,approxEquals_!,approxElementEquals_!.3 %! tand2.0 *! tnow work, and the scalar takes the tensor's precision.Fixes
!two-tensor ops (+,<,maximum,where,===, …) now fail fast when the extents differ. Before, JAX silently broadcast axes of extent 1 (e.g. 2×2+1×2). This didn't align with DimWit's intent of explicit broadcasting and no extent broadcasting.approxElementEqualsreally compares elementwise (jnp.isclose). Before, it returned a single boolean (allclose) typed as a tensor.argsort(axes: Tuple). It always failed at runtime, becausejnp.argsortonly takes one axis.Breaking changes
add,subtract,multiply,divide,mod,negate,less,lessEqual,greater,greaterEqual,equalandlogicalAnd/Or/Xor/Notare no longer top-level: use the operators, orTensor.add(...)etc.addScalar,subtractScalar,divideScalar,modScalar: use+!,-!,/!,%!with aTensor0.multiplyScalar→Tensor.scale.2.0f *! t) now needimport dimwit.Conversions.given, liket *! 2.0falready did.quantile(q)takes aTensor0[V]instead of aFloat.retag(unused).dimwit.tensor.TensorOps.*,dimwit.tensor.tensorops.*,ValueOps. Code usingimport dimwit.*only sees the changes above.