Replies: 20 comments
|
Hi Nardi, great questions!
asdex's sparsity detection and coloring are currently not implemented in JAX and therefore can't be jitted. The assumption I made in my design (so far) is that users would "prepare" a Jacobian (or Hessian) function using jac_fn = asdex.jacobian(f, x_sample) # <-- preparation step containing non-JAX operations
jac_fn = jax.jit(jac_fn) # <-- prepared function can be jitted and transformed
for x in inputs:
J = jac_fn(x) # <-- reuse for different inputs in your solverHowever, it would indeed be good to support both approaches.
Thanks for the pointer, I was not aware of this! |
|
Okay, that makes sense! Indeed, it is fine that the coloring logic lives outside of the JIT context, but it would be nice if the result can be passed into JITted functions. I tried to just register the classes as PyTrees, but I had some issues with the cached properties on the
I'll leave it to you to decide what makes sense here :) |
|
For some context, I have been setting up a sort of "extension library" to lineax that wraps a number of JAX-compatible sparse solver routines: https://github.com/nardi/splineax I think it would be nice to add some sparse autodiff integration, as lineax has a However, the API is as such that the operator is only instantiated quite last-minute (when the point of Jacobian evaluation For now I'll try to do a first integration with my |
|
Option 2 of precomputing the indices makes sense, as they stay static after preparation. The construction of the
Yes, this is one of the use-cases we had in mind for asdex! I've already talked a bit with @jpbrodrick89 about this, who might be interested in this discussion. |
|
Hi guys, yes the elegant design would be a solver that calls bcoo/csr_jacobian in solver.init. The key hurdle is that in most cases you do NOT want the indices calculation to be jittable as you often want to use non-jit-friendly operations and most importantly store the indices as STATIC arrays. There are two ways around this, either you trick jax into thinking the arrays are hashable or (my preference) you convert from numpy arrays to jax arrays at the final step and XLA should just constant fold all the non jit operations and you only end up doing the index calculation once and they are stored as constant jax arrays in the compiled jaxpr. So perhaps just extracting the arrays from the pattern objects or declaring them as pytrees/equinox modules will just work. The main requirement for this to work is that at no point should index calculation depend on any concrete VALUES and only abstract shape/dtype (e.g. of example x) so that it can truly be computed at compile time. |
|
@jpbrodrick89 I can see that the index calculation itself should be performed at compile time, but is it important that the resulting index arrays themselves are static? They just get passed as inputs to The fact that the coloring and index calculation needs to happen outside JIT is in fact my motivation for wanting to pass the result into a JITted function. So that the coloring can be done outside of the solver context, and the solver only has to consume it (or preferably, the operator, since it is a property of the function and solver-independent). The design I have now is to add a However in many cases, the Jacobian evaluation point is only known within a JITted context (e.g. an iterative optimization algorithm), so if we do not want to recalculate the coloring more often than necessary, we need to calculate the sparsity and coloring beforehand, then use that information to create the operator object later. The API for this exists of course, but then if the result of I guess a different approach that would make this less important would be to have a kind of "coloring cache" for each function, similar to JAX's |
This is already a general requirement for asdex's pattern types: |
Sorry I dont understand why you need to create a SparseJacobianLinearOperator. I would just create new solver that calls bcoo_jacobian(operator.mv). Alternatively you could define a to_bcoo singledispatch that calls bcoo_jacobian(operator.mv) by default but allows a direct path for a BCOOLinearOperator if you find that helpful. For a reference, please look at lineax's TridiagonalSolver, this never converts an operator to a TridiagonalOperator but just calls tridiagonal which extracts the diagonals (either directly or by coloring rules based on singledispatch). |
No its not important, just helpful to ensure things are done at compile time, allowing them to be dynamic jax arrays that XLA "realises" are constant and folds accordingly should be fine. |
As @adrhill said you don't need to know the jacobian evaluation point to evaluate the sparsify pattern and construct the coloring, you just need its shape and dtype, so your concern is not warranted I believe.
I don't think this is necessary JAX's XLA constant folding should handle this automatically, but if compile time becomes a bottleneck you can always try AOT. |
Why are you needing to differentiate the solver routine itself, lineax's implicit differentiation rules should sidestep this. |
Sorry, what is |
Yes you're right, it's not so much about the solving step itself, but rather that a it is part of a larger function which is differentiated and compiled, and I do not want to recalculate the coloring. For example, I may have two functions:
My situation is that if I naively call Another example with the same behavior would be if I have only one function, but change a separate unrelated static argument, triggering recompilation and redoing the coloring. |
|
@jpbrodrick89 Ah, I see you are referring to |
|
Sorry I'm not so familiar with the API here, I was just getting mixed up with my lineaxpr API. Lineaxpr is still experimental and in development and i have many commits that I have not yet pushed to the repo, would be interesting to know if it works more out the box for what you're trying to achieve. I think both approaches should have advantages and disadvantages. |
|
asdex already supports multiple decompression targets (e.g. JAX's BCOO and dense arrays, NumPy and sparse SciPy formats). If lineax has a favored format, we should add it. |
|
I've taken the freedom to convert this issue into a discussion to coordinate the interface between asdex/lineaxpr and lineax/splineax. I opened #171 to track the pytree compatibility feature request. |
|
@adrhill I added support for sparse Jacobian operators via asdex to splineax! https://github.com/nardi/splineax For now I went with the static coloring approach, I wrapped it in an object that is identified by the combination of function and arguments, which should indicate a unique coloring. That seems good enough for now, if a change is made to make the coloring object a pytree the wrapper object can be removed without changing the API :) |
|
Sorry, just catching up on this discussion (co-developer of SparseMatrixColorings.jl here). Can someone explain to me why you folks need to manipulate the patterns directly, as opposed to the functions that |
|
The original question has been resolved in #176 and |
Uh oh!
There was an error while loading. Please reload this page.
Hey, not sure if an issue is appropriate for this, but didn't see any better option :)
I think the library looks great! I have been wanting something like this for a long time. I've tried to do some sparsity detection myself in https://github.com/nardi/jax-nansparse, and then use https://github.com/mfschubert/sparsejac for the coloring and compressed AD, but this library seems to be a lot more feature complete.
I've found it very useful to be able to calculate a Jacobian coloring in advance and cache it, which this library of course supports. However, it is not so clear to me why the
SparsityPatternandColoredPatternobjects are not PyTrees, while most of their fields are just arrays. Now, it is quite difficult to pass them into a JITted function, because they are neither properly hashable nor traceable by JAX.I have some functions where I do nested differentiation, e.g. I define a sparse solver routine which internally calls a (sparse)
jacfwdwith a coloring, then that whole solver routine is itself differentiated (densely). I've been able to achieve this with a fork ofsparsejacwhere I've split the coloring and the AD code, and there the coloring objects are simple PyTrees: https://github.com/nardi/sparsejac/blob/refactor_coloring/src/sparsejac/sparsejac.pyI'm assuming there is a reason for this decision, but I can't see it yet when reading the code. Is there something in the (de)compression code that requires the sparsity/coloring objects to be static perhaps?
Thank you!
All reactions