google / google/jaxopt

Type precision issue in BoxOSQP

Open
#547 10 comments 0 reactions 1 assignee Claimed by @Algue-Rythme View on GitHub
Dominant language
Python
Stars
1.1k
Forks
76
Avg merge
2d 21h
Merged PRs (30d)
1

Description

When I have `float64` support enabled in JAX and try to run:

```
optimizer = jaxopt.BoxOSQP()
optimizer.run(
params_obj=(
jnp.eye(30, dtype=jax.numpy.float32),
jnp.ones((30,), dtype=jax.numpy.float32),
),
params_eq=jnp.ones((1, 30), dtype=jax.numpy.float32),
params_ineq=(-1, 1),
)
```

I get an internal error from the implementation (partially redacted):

```
[.../jaxopt/_src/osqp.py](...) in run(self, init_params, params_obj, params_eq, params_ineq)
763 init_params = self.init_params(None, params_obj, params_eq, params_ineq)
764
--> 765 return super().run(init_params, params_obj, params_eq, params_ineq)
766
767 def l2_optimality_error(

[.../jaxopt/_src/base.py](...) in run(self, init_params, *args, **kwargs)
345 run = decorator(run)
346
--> 347 return run(init_params, *args, **kwargs)
348
349 def __post_init__(self):

[.../jaxopt/_src/implicit_diff.py](...) in wrapped_solver_fun(*args, **kwargs)
249 args, kwargs = _signature_bind(solver_fun_signature, *args, **kwargs)
250 keys, vals = list(kwargs.keys()), list(kwargs.values())
--> 251 return make_custom_vjp_solver_fun(solver_fun, keys)(*args, *vals)
252
253 return wrapped_solver_fun

[.../jaxopt/_src/implicit_diff.py](...) in solver_fun_flat(*flat_args)
205 def solver_fun_flat(*flat_args):
206 args, kwargs = _extract_kwargs(kwarg_keys, flat_args)
--> 207 return solver_fun(*args, **kwargs)
208
209 def solver_fun_fwd(*flat_args):

[.../jaxopt/_src/base.py](...) in _run(self, init_params, *args, **kwargs)
307 zero_step = self._make_zero_step(init_params, state)
308
--> 309 opt_step = self.update(init_params, state, *args, **kwargs)
310 init_val = (opt_step, (args, kwargs))
311

[.../jaxopt/_src/osqp.py](...) in update(self, params, state, params_obj, params_eq, params_ineq)
703 # We need our own ifelse cond because automatic jitting of jax.lax.cond branches
704 # could pose problems with non jittable matvecs, or prevent printing when verbose > 0.
--> 705 rho_bar, solver_state = cond(
706 jnp.mod(state.iter_num, self.stepsize_updates_frequency) == 0,
707 lambda _: self._update_stepsize(rho_bar, solver_state, primal_residuals, dual_residuals, Q, c, A, x, y),

[.../jaxopt/_src/cond.py](...) in cond(cond, if_fun, else_fun, jit, *operands)
22 with jax.disable_jit():
23 return jax.lax.cond(cond, if_fun, else_fun, *operands)
---> 24 return jax.lax.cond(cond, if_fun, else_fun, *operands)

TypeError: true_fun and false_fun output must have identical types, got
('DIFFERENT ShapedArray(float32[]) vs. ShapedArray(float64[], weak_type=True)', ('ShapedArray(float32[30])', ('ShapedArray(float32[30,30])', 'ShapedArray(float32[1,30])', 'ShapedArray(float64[], weak_type=True)', 'DIFFERENT ShapedArray(float32[]) vs. ShapedArray(float64[], weak_type=True)'), None)).
```

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.