This repository contains code for training and evaluating optimal control policies using implicit neural networks with Jacobian-Free Backpropagation (JFB) and Jacobian-Based Backpropagation (JBB/CVX) for the three examples presented in our ICML paper.
ImplicitNets.py: Implicit neural network architectures (JFB method)ImplicitOC.py: Implicit optimal control layer with HJB optimality conditionsCVXPolicy.py: CVXPY-based policies (JBB method)DirectControlNets.py: Direct transcription baseline policies (for comparison)OptimalControlTrainer.py: Unified training framework for all policy typesutils.py: Utility functions
MultiBicycle.py: Multi-agent bicycle optimal control problemQuadcopter.py: Single and multi-agent quadcopter optimal controlConsumption.py: Multi-agent consumption-savings optimal control
example_multibicycle.py: Train JFB policy on multi-bicycle problem
example_multi_quadcopter.py: Train JFB and JBB policies on quadrotor problem
example_multiConsumption.py: Train JFB policy on multi-agent consumption problem
torch>=2.0.0
numpy>=1.20.0
pandas>=1.3.0
matplotlib>=3.4.0
cvxpy>=1.2.0
cvxpylayers>=0.1.5pip install torch numpy pandas matplotlib cvxpy cvxpylayersTrain with JFB:
python example_multibicycle.pyConfiguration parameters in the script:
batch_size: Batch size for training (default: 100)nt: Number of time steps (default: 60)t_final: Final time horizon (default: 4.0)n_b: Number of bicycles (fixed at 100)alphaG: Terminal cost weight (default: 500.0)epochs: Number of training epochs (default: 500)
Train with JFB (and optionally JBB):
python example_multi_quadcopter.pyThe script supports training with different numbers of quadrotors and training methods via command-line arguments:
Basic usage (default: 100 quadrotors, JFB training, CPU device):
python example_multi_quadcopter.pyUsing GPU:
python example_multi_quadcopter.py --device cuda
# or specify GPU device
python example_multi_quadcopter.py --device cuda:0Single Quadrotor (1 agent):
python example_multi_quadcopter.py --num_quadcopters 16 Quadrotors:
python example_multi_quadcopter.py --num_quadcopters 6Train with JBB (CVXPyLayers):
python example_multi_quadcopter.py --train_jbbTrain with both JFB and JBB:
python example_multi_quadcopter.py --train_jfb --train_jbbDisable JFB training:
python example_multi_quadcopter.py --no_train_jfb --train_jbbOther useful arguments:
python example_multi_quadcopter.py --num_quadcopters 100 --epochs 1000 --lr 0.005 --device cuda:0Configuration parameters:
batch_size: Batch size for training (default: 50)nt: Number of time steps (default: 160)t_final: Final time horizon (default: 4.5)num_quadcopters: Number of quadrotors - 1, 6, or 100 (default: 100)alphaG: Terminal cost weight (default: 1000.0)epochs: Number of training epochs (default: 500)
Train with JFB:
python example_multiConsumption.pyConfiguration parameters:
batch_size: Batch size for training (default: 128)nt: Number of time steps (default: 100)t_final: Final time horizon (default: 2.0)m: Number of agents (fixed at 100)epochs: Number of training epochs (default: 500)
Numerical verification is performed automatically during training through the compute_loss() function when the save_history=True flag is set in the training configuration. This generates CSV files containing:
- Total loss (control objective = running_cost + alphaG * terminal_cost)
- Running cost and terminal cost (separate)
- Optimality condition violations (cHJB, cHJBfin, cadj, cadjfin)
- Gradient metrics (max_grad_H, avg_grad_H)
- Contractivity verification: max_grad_T_u (maximum gradient of T_theta operator)
- M_theta conditioning: smallest_M_sdval, largest_M_sdval (singular values of M_theta matrix)
- Descent direction: angle between expected gradients
The training scripts save these history CSV files (e.g., history_best_policy_JFB_*.csv) which contain all numerical verification metrics reported in the paper. These CSVs can be loaded and analyzed to reproduce the numerical verification plots shown in the paper.
Training scripts generate:
best_policy_*.pth: Saved model weightshistory_*.csv: Training history with all numerical verification metrics (loss, costs, optimality violations, gradient norms, contractivity checks, M_theta singular values, descent angles)*_run.log: Training logs- Trajectory plots in
results_*/directories
- GPU Usage: Set
device='cuda'ordevice='cuda:X'in config for GPU acceleration - Hyperparameters: The provided hyperparameters are tuned for each problem
- Convergence: Monitor training logs for:
- Loss convergence
- Optimality condition violations (cHJB, cadj)
- Gradient norms
- Memory: Large batch sizes may require significant GPU memory
- Multiple Trials: Run multiple trials (set
n_trialsin main()) for statistical significance
batch_size: Number of initial conditions per training batchnt: Discretization time stepst_final: Time horizonalphaG: Weight on terminal costalphaHJB: Weights for optimality condition penalties [cHJB_weight, cHJBfin_weight]lr: Learning rate (default: 1e-3 for JFB, 5e-4 for Direct Control)epochs: Number of training epochs
- JFB:
max_iters,tol,tracked_iters,alphafor fixed-point solver - JBB/CVX:
tolfor CVXPY solver tolerance - Direct Transcription:
weight_decayfor explicit regularization (required due to lack of optimality constraints)
The comparison scripts demonstrate a key insight: Direct transcription policies (which directly optimize control sequences without enforcing optimality conditions) require:
- 10x smaller learning rate (5e-4 vs 5e-3)
- Explicit weight decay for regularization
- 100x smaller weight initialization
These differences arise because JFB enforces optimality conditions (∇_u H = 0), providing implicit regularization. Direct transcription lacks these constraints and requires explicit regularization to match performance.
If you use this code, please cite our paper:
@misc{gelphman2025jfb,
title={End-to-End Training of High-Dimensional Optimal Control with Implicit Hamiltonians via Jacobian-Free Backpropagation},
author={Eric Gelphman and Deepanshu Verma and Nicole Tianjiao Yang and Stanley Osher and Samy Wu Fung},
year={2025},
eprint={2510.00359},
archivePrefix={arXiv},
primaryClass={math.OC},
url={https://arxiv.org/abs/2510.00359},
}
@misc{gelphman2026convergence,
title={On the Convergence of Jacobian-Free Backpropagation for Optimal Control Problems with Implicit Hamiltonians},
author={Eric Gelphman and Deepanshu Verma and Nicole Tianjiao Yang and Stanley Osher and Samy Wu Fung},
year={2026},
eprint={2602.00921},
archivePrefix={arXiv},
primaryClass={math.OC},
url={https://arxiv.org/abs/2602.00921},
}