Type precision issue in BoxOSQP
- 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
Assessment
This issue has not been assessed yet.