google / google/jaxopt

BoxOSQP does not work without equality constraints

Open
#546 5 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

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

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.