patrick-kidger / patrick-kidger/optimistix

Interactively step through solve with jax.lax.scan

Open
#103 2 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

question
Dominant language
Python
Stars
623
Forks
54
PR merge metrics
No merged PRs in 30d

Description

Hi,

I am trying to use jax.lax.scan together with the interactive solving approach (example in documentation) for some minimization.
However I found that for some solvers (least_squares) using lax.scan results in a type error. (Everything works with the standard for-loop.)

This is the example im working with:

import jax.numpy as jnp
import optimistix
from jax.tree_util import Partial
import jax
import numpy as np

### work with lax.scan
# solver = optimistix.BFGS(rtol=1e-3, atol=1e-3)
# solver = optimistix.NonlinearCG(rtol=1e-3, atol=1e-3)

### do NOT work with lax.scan
# solver = optimistix.GaussNewton(rtol=1e-3, atol=1e-3)
# solver = optimistix.LevenbergMarquardt(rtol=1e-3, atol=1e-3)
# solver = optimistix.IndirectLevenbergMarquardt(rtol=1e-3, atol=1e-3)
# solver = optimistix.NelderMead(rtol=1e-3, atol=1e-3)
# solver = optimistix.Dogleg(rtol=1e-3, atol=1e-3)



def test_func(x, *args):
    return jnp.sum(x**2) + 1.0, None


fn = test_func
y = jnp.array(np.random.uniform(-1,1,size=(5,5)))

args = None
options = dict(lower=-1.0, upper=1.0)
f_struct = jax.ShapeDtypeStruct((), jnp.float32)
aux_struct = None
tags = frozenset()

state = solver.init(fn, y, args, options, f_struct, aux_struct, tags)
step = Partial(solver.step, fn=fn, args=args, options=options, tags=tags)

def step_helper(carry, xs):
    y, state, _ = carry
    return step(y=y, state=state), None

carry = (y, state, None)

carry, _ = jax.lax.scan(step_helper, carry, length=10)

#for _ in range(10):
#    carry, _ = step_helper(carry, None)

y, state, aux = carry

The error which is raised is:

TypeError: Value { lambda a:f32[5,5]; b:f32[5,5]. let
c:f32[5,5] = mul b a
d:f32[] = reduce_sum[axes=(0, 1)] c
in (d,) } with type <class 'jax._src.core.Jaxpr'> is not a valid JAX type

Am I doing something wrong or is there something else going on?

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start with the interactive-solving example in the documentation and reproduce the failure using jax.lax.scan with the listed least-squares solvers. Compare it with the working standard for-loop and verify that the affected solvers can be stepped through scan without the Jaxpr type error.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
machine-learning
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Stale
Clarity
Mostly clear
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.