linrax: A JAX-Compatible, Simplex
Method Linear Program Solver

ThC21.1 @ ACC 2026

Brendan Gould
Akash Harapanahalli

Georgia Institute of Technology

Samuel Coogan

JAX Functional Paradigm

Function transformations:

  • jit: compile to fast XLA code (CPU or GPU)
  • grad / jacfwd / jacrev: automatic differentiation
  • vmap: automatic vectorization / batching

Combine for impressive computational abilities

  • Growing ecosystem of JAX-native control / optimization libraries: immrax [8], jaxopt [9]

Why Another LP Solver?

Many toolkits can solve general form:

\[ \min_{x \in [l, u]} c^\top x \quad \text{s.t.} \quad A_{\rm eq}\, x = b_{\rm eq}, \quad A_{\rm ub}\, x \leq b_{\rm ub}. \]

  1. Want transform compatibility for full control pipeline: rules out industrial solvers
  2. Need to solve general, automatically generated problems
    • May contain degenerate constraints: breaks first-order methods [9], [10]

Our Contribution

linrax: the first simplex-based LP solver in JAX.

  • Fully compatible with function transformations (jit, jacfwd, vmap)
  • Solves problems with degenerate constraints
  • Converges to an exact solution (up to floating-point error)

Control “Nudging” via Differentiable Reachability

Nominal optimal control trajectories often hug obstacles

  • Even small disturbances could cause a collision
  • Reachable-Set Analysis bounds effects of disturbances

Problem

We want to “nudge” the feedforward input \(u_{\rm ff}\) such that the entire reachable set robustly avoids the obstacle.

LP refinement greatly tightens reachable sets

JAX Control Pipeline

Optimal
Control
\(u_{\mathrm{ff}}\)
jit
Reachability
Interval
Arithmetic
solve degenerate LPs
Refine \(x_{1}\)
Refine \(x_{2}\)
Refine \(x_{n}\)
Increment time
Reachset
vmap \(\texttt{SafetyCheck}\):
Quantify
Intersection
jacfwd Compute
Gradient
\(u_{\mathrm{ff}} \leftarrow u_{\mathrm{ff}} - \alpha\,\nabla \texttt{SafetyCheck}(u_{\mathrm{ff}})\)
  • Solves \(2(n_x+n_a)T\) auto-generated LPs
  • Every LP has degenerate constraints

  • Left: unrefined reachable sets are too conservative for planning
  • Center: unmodified optimal control, refined reachable set intersects obstacle
  • Right: after 8 gradient steps on \(u_{\rm ff}\), reachable set is robustly safe

The Simplex Method: Pre-Solving

Convert to canonical form:

The Simplex Method: Pre-Solving

Convert to canonical form:

  • Remove unnecessary variables / constraints
  • Prune degenerate constraints

Problem

JAX does not allow the static shape of arrays to depend on traced values.

The Simplex Method: Pivoting

The Simplex Method: Two Phase Approach

Problem

Requires BFS (vertex of the feasible region) to initialize pivoting process.

Solution

Auxiliary problem finds a feasible vertex:

\[ \min_{x, a \ge 0} \begin{bmatrix} 0_{1 \times n} & 1_{1 \times m} \end{bmatrix} \begin{bmatrix} x \\ a\end{bmatrix} \text{ s.t. } \begin{bmatrix} A & \mathbf{I}_{m} \end{bmatrix} \begin{bmatrix} x \\ a \end{bmatrix} = b. \]

  • Phase 1 drives the auxiliary variables to zero \(\rightarrow\) feasible vertex of the real problem.
    • Easy to initialize: \(x=0\), \(a=b\)
  • Phase 2 then optimizes the original cost.

Problem

For degenerate LPs, auxiliary variables may remain in the active set after phase 1.

The Simplex Method: Two Phase Approach

Key Idea: Pivot Out Auxiliary Rows

Idea: Mark stuck auxiliary rows in ratio test so they exit the basis on the first opportunity in Phase 2.

Benchmark 1: small random LPs

20 input variables, 15 inequality constraints, \(N = 100\) samples.

Method jit (s) Solve time (s)
scipy \(1.08\cdot 10^{-3} \pm 7.96\cdot 10^{-4}\)
gurobi \(2.23\cdot 10^{-3} \pm 4.61\cdot 10^{-3}\)
cvxpy \(1.18\cdot 10^{-3} \pm 8.88\cdot 10^{-4}\)
jaxopt 0.95 \(3.48\cdot 10^{-2} \pm 3.50\cdot 10^{-3}\)
linrax 1.23 \(\mathbf{6.26\cdot 10^{-3} \pm 7.87\cdot 10^{-4}}\)

Results

  • Same order of magnitude as production solvers
  • ~5× faster than the other JAX option

Benchmark 2: the intended use case

Van der Pol reachable set computation, \(\mu = 1\), \(t_f = 0.628\), \(N = 10\).

Method jit (s) Total time (s) Bound size
scipy \(37.8 \pm 1.0\) \(7.01\cdot 10^{-2}\)
gurobi \(44.7 \pm 2.2\) \(6.86\cdot 10^{-2}\)
cvxpy crashed
jaxopt 4.13 \(56.6 \pm 0.6\) \(8.85\cdot 10^{-2}\) (loose)
linrax 6.25 \(\mathbf{2.13 \pm 0.06}\) \(6.87\cdot 10^{-2}\)

Results

  • ~17× faster than scipy, ~20× faster than gurobi
  • The whole pipeline JIT-compiles into a single XLA program
  • cvxpy crashes; jaxopt is slow and converges to a worse bound

Conclusions

JAX optimization for control has unique advantages and challenges

  • Function transformations enable fast, gradient-informed pipelines
  • Tracability requirements complicate classic algorithms

linrax: first simplex-based LP solver in JAX

  • Handles dependent constraints
  • Exact solution (up to floating point)
  • Fully composable with jit, jacfwd, vmap

Possible Future work:

  • smarter pivoting (beyond Bland’s rule)
  • GPU-side parallelism inside the solver
  • reverse-mode autodiff

pip install linrax

Personal website:

References

[1]
D. P. Bertsekas, Reinforcement learning and optimal control, 2nd printing. Belmont, Massachusetts: Athena Scientific, 2019.
[2]
A. D. Ames, S. Coogan, M. Egerstedt, G. Notomista, K. Sreenath, and P. Tabuada, “Control Barrier Functions: Theory and Applications,” in 2019 18th European Control Conference (ECC), June 2019, pp. 3420–3431. doi: 10.23919/ECC.2019.8796030.
[3]
S. M. Harwood and P. I. Barton, “Efficient polyhedral enclosures for the reachable set of nonlinear control systems,” Springer London, Feb. 2016.
[4]
K. Shen and J. Scott, “Rapid and Accurate Reachability Analysis for Nonlinear Dynamic Systems by Exploiting Model Redundancy,” Computers & Chemical Engineering, vol. 106, Aug. 2017, doi: 10.1016/j.compchemeng.2017.08.001.
[5]
P. Virtanen et al., SciPy 1.0: Fundamental algorithms for scientific computing in python,” Nature Methods, vol. 17, pp. 261–272, 2020, doi: 10.1038/s41592-019-0686-2.
[6]
Gurobi Optimization, LLC, “Gurobi optimizer reference manual.” 2024.
[7]
S. Diamond and S. Boyd, CVXPY: A Python-embedded modeling language for convex optimization,” Journal of Machine Learning Research, vol. 17, no. 83, pp. 1–5, 2016.
[8]
A. Harapanahalli, S. Jafarpour, and S. Coogan, “Immrax: A parallelizable and differentiable toolbox for interval analysis and mixed monotone reachability in JAX,” IFAC-PapersOnLine, vol. 58, no. 11, pp. 75–80, 2024, doi: 10.1016/j.ifacol.2024.07.428.
[9]
M. Blondel et al., “Efficient and modular implicit differentiation,” arXiv preprint arXiv:2105.15183, 2021, Available: https://arxiv.org/abs/2105.15183
[10]
H. Lu, Z. Peng, and J. Yang, MPAX: Mathematical programming in JAX,” arXiv preprint arXiv:2412.09734, 2024, Available: https://arxiv.org/abs/2412.09734