BoxOSQP does not work without equality constraints
- Dominant language
- Python
- Stars
- 1.1k
- Forks
- 76
- Avg merge
- 2d 21h
- Merged PRs (30d)
- 1
Description
Trying to run:
```
optimizer = jaxopt.BoxOSQP()
optimizer.run(
params_obj=(
jnp.eye(30, dtype=jnp.float32),
jnp.ones((30,), dtype=jnp.float32),
),
params_ineq=(-1, 1),
)
```
produces an error (some parts 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)
287 *args,
288 **kwargs) -> OptStep:
--> 289 state = self.init_state(init_params, *args, **kwargs)
290
291 # We unroll the very first iteration. This allows `init_val` and `body_fun`
[.../jaxopt/_src/osqp.py](...) in init_state(self, init_params, params_obj, params_eq, params_ineq)
456 A = self.matvec_A(params_eq)
457
--> 458 primal_residuals, dual_residuals = self._compute_residuals(Q, c, A, x, z, y)
459 solver_state = self._eq_qp_solve_impl.init_state(x, Q_params, params_eq,
460 self.sigma, self.rho_start)
[.../jaxopt/_src/osqp.py](...) in _compute_residuals(self, Q, c, A, x, z, y)
530 Ax, ATy = A.matvec_and_rmatvec(x, y)
531 primal_residuals = tree_sub(Ax, z)
--> 532 dual_residuals = tree_add(tree_add(Q(x), c), ATy)
533 return primal_residuals, dual_residuals
534
TypeError: unsupported operand type(s) for +: 'ArrayImpl' and 'NoneType'
```
Contributor guide
Assessment
This issue has not been assessed yet.