humancompatible-train 0.3.2__tar.gz → 0.4.1__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (81) hide show
  1. {humancompatible_train-0.3.2/src/humancompatible_train.egg-info → humancompatible_train-0.4.1}/PKG-INFO +24 -5
  2. {humancompatible_train-0.3.2 → humancompatible_train-0.4.1}/README.md +2 -2
  3. humancompatible_train-0.4.1/pyproject.toml +56 -0
  4. {humancompatible_train-0.3.2 → humancompatible_train-0.4.1}/src/humancompatible/train/dual_optim/__init__.py +3 -0
  5. humancompatible_train-0.4.1/src/humancompatible/train/dual_optim/alm.py +348 -0
  6. humancompatible_train-0.4.1/src/humancompatible/train/dual_optim/base.py +502 -0
  7. humancompatible_train-0.4.1/src/humancompatible/train/dual_optim/ialm.py +248 -0
  8. {humancompatible_train-0.3.2 → humancompatible_train-0.4.1}/src/humancompatible/train/dual_optim/moreau.py +11 -3
  9. humancompatible_train-0.4.1/src/humancompatible/train/dual_optim/nupi.py +231 -0
  10. humancompatible_train-0.4.1/src/humancompatible/train/dual_optim/pbm.py +505 -0
  11. humancompatible_train-0.4.1/src/humancompatible/train/dual_optim/ssg.py +171 -0
  12. humancompatible_train-0.4.1/src/humancompatible/train/fairness/utils/balanced_batch_sampler.py +156 -0
  13. humancompatible_train-0.4.1/src/humancompatible/train/sqp/__init__.py +1 -0
  14. humancompatible_train-0.4.1/src/humancompatible/train/sqp/ghost/__init__.py +28 -0
  15. humancompatible_train-0.4.1/src/humancompatible/train/sqp/ghost/ghost_sqp.py +103 -0
  16. humancompatible_train-0.4.1/src/humancompatible/train/sqp/ghost/mlmc.py +74 -0
  17. humancompatible_train-0.4.1/src/humancompatible/train/sqp/ghost/oracle.py +75 -0
  18. humancompatible_train-0.4.1/src/humancompatible/train/sqp/ghost/sampler.py +54 -0
  19. humancompatible_train-0.4.1/src/humancompatible/train/sqp/ghost/subproblems.py +75 -0
  20. humancompatible_train-0.4.1/src/humancompatible/train/sqp/nonopt/__init__.py +12 -0
  21. humancompatible_train-0.4.1/src/humancompatible/train/sqp/nonopt/direction.py +421 -0
  22. humancompatible_train-0.4.1/src/humancompatible/train/sqp/nonopt/inverse_hessian.py +276 -0
  23. humancompatible_train-0.4.1/src/humancompatible/train/sqp/nonopt/line_search.py +156 -0
  24. humancompatible_train-0.4.1/src/humancompatible/train/sqp/nonopt/optimizer.py +442 -0
  25. humancompatible_train-0.4.1/src/humancompatible/train/sqp/nonopt/point_set.py +79 -0
  26. humancompatible_train-0.4.1/src/humancompatible/train/sqp/nonopt/qp.py +111 -0
  27. {humancompatible_train-0.3.2 → humancompatible_train-0.4.1/src/humancompatible_train.egg-info}/PKG-INFO +24 -5
  28. humancompatible_train-0.4.1/src/humancompatible_train.egg-info/SOURCES.txt +49 -0
  29. {humancompatible_train-0.3.2 → humancompatible_train-0.4.1}/src/humancompatible_train.egg-info/requires.txt +9 -1
  30. humancompatible_train-0.4.1/tests/test_alm.py +454 -0
  31. {humancompatible_train-0.3.2 → humancompatible_train-0.4.1}/tests/test_balanced_batch_sampler.py +114 -0
  32. humancompatible_train-0.4.1/tests/test_cooper_parity.py +159 -0
  33. humancompatible_train-0.4.1/tests/test_distributed.py +64 -0
  34. humancompatible_train-0.4.1/tests/test_dual_optimizer_base.py +331 -0
  35. humancompatible_train-0.4.1/tests/test_fairness_problems.py +212 -0
  36. humancompatible_train-0.4.1/tests/test_ghost.py +86 -0
  37. humancompatible_train-0.4.1/tests/test_golden.py +90 -0
  38. humancompatible_train-0.4.1/tests/test_ialm.py +95 -0
  39. humancompatible_train-0.4.1/tests/test_llm_gates.py +461 -0
  40. humancompatible_train-0.4.1/tests/test_nonopt.py +265 -0
  41. humancompatible_train-0.4.1/tests/test_nupi.py +146 -0
  42. {humancompatible_train-0.3.2 → humancompatible_train-0.4.1}/tests/test_pbm.py +47 -47
  43. humancompatible_train-0.4.1/tests/test_ssg.py +143 -0
  44. humancompatible_train-0.3.2/pyproject.toml +0 -29
  45. humancompatible_train-0.3.2/src/humancompatible/train/benchmark/algorithms/Algorithm.py +0 -25
  46. humancompatible_train-0.3.2/src/humancompatible/train/benchmark/algorithms/__init__.py +0 -8
  47. humancompatible_train-0.3.2/src/humancompatible/train/benchmark/algorithms/ghost.py +0 -254
  48. humancompatible_train-0.3.2/src/humancompatible/train/benchmark/algorithms/optim_wrapper.py +0 -184
  49. humancompatible_train-0.3.2/src/humancompatible/train/benchmark/algorithms/sgd.py +0 -114
  50. humancompatible_train-0.3.2/src/humancompatible/train/benchmark/algorithms/ssl_alm.py +0 -317
  51. humancompatible_train-0.3.2/src/humancompatible/train/benchmark/algorithms/switching_subgradient.py +0 -199
  52. humancompatible_train-0.3.2/src/humancompatible/train/benchmark/algorithms/utils.py +0 -61
  53. humancompatible_train-0.3.2/src/humancompatible/train/benchmark/constraints/__init__.py +0 -15
  54. humancompatible_train-0.3.2/src/humancompatible/train/benchmark/constraints/constraint.py +0 -114
  55. humancompatible_train-0.3.2/src/humancompatible/train/benchmark/constraints/constraint_fns.py +0 -251
  56. humancompatible_train-0.3.2/src/humancompatible/train/benchmark/constraints/torch/__init__.py +0 -1
  57. humancompatible_train-0.3.2/src/humancompatible/train/benchmark/constraints/torch/constraints.py +0 -48
  58. humancompatible_train-0.3.2/src/humancompatible/train/dual_optim/alm.py +0 -369
  59. humancompatible_train-0.3.2/src/humancompatible/train/dual_optim/ialm.py +0 -330
  60. humancompatible_train-0.3.2/src/humancompatible/train/dual_optim/nupi.py +0 -357
  61. humancompatible_train-0.3.2/src/humancompatible/train/dual_optim/pbm.py +0 -503
  62. humancompatible_train-0.3.2/src/humancompatible/train/fairness/__init__.py +0 -0
  63. humancompatible_train-0.3.2/src/humancompatible/train/fairness/utils/balanced_batch_sampler.py +0 -126
  64. humancompatible_train-0.3.2/src/humancompatible/train/optim/PBM.py +0 -571
  65. humancompatible_train-0.3.2/src/humancompatible/train/optim/__init__.py +0 -5
  66. humancompatible_train-0.3.2/src/humancompatible/train/optim/barrier.py +0 -158
  67. humancompatible_train-0.3.2/src/humancompatible/train/optim/ssl_alm.py +0 -236
  68. humancompatible_train-0.3.2/src/humancompatible/train/optim/ssl_alm_adam.py +0 -347
  69. humancompatible_train-0.3.2/src/humancompatible/train/optim/ssl_alm_adam_moment.py +0 -315
  70. humancompatible_train-0.3.2/src/humancompatible/train/optim/ssw.py +0 -156
  71. humancompatible_train-0.3.2/src/humancompatible/train/optim/ssw_barrier.py +0 -236
  72. humancompatible_train-0.3.2/src/humancompatible_train.egg-info/SOURCES.txt +0 -44
  73. humancompatible_train-0.3.2/tests/test_alm.py +0 -90
  74. {humancompatible_train-0.3.2 → humancompatible_train-0.4.1}/LICENCE.txt +0 -0
  75. {humancompatible_train-0.3.2 → humancompatible_train-0.4.1}/setup.cfg +0 -0
  76. {humancompatible_train-0.3.2 → humancompatible_train-0.4.1}/src/humancompatible/train/__init__.py +0 -0
  77. {humancompatible_train-0.3.2 → humancompatible_train-0.4.1}/src/humancompatible/train/dual_optim/barrier.py +0 -0
  78. {humancompatible_train-0.3.2/src/humancompatible/train/benchmark → humancompatible_train-0.4.1/src/humancompatible/train/fairness}/__init__.py +0 -0
  79. {humancompatible_train-0.3.2 → humancompatible_train-0.4.1}/src/humancompatible/train/fairness/utils/__init__.py +0 -0
  80. {humancompatible_train-0.3.2 → humancompatible_train-0.4.1}/src/humancompatible_train.egg-info/dependency_links.txt +0 -0
  81. {humancompatible_train-0.3.2 → humancompatible_train-0.4.1}/src/humancompatible_train.egg-info/top_level.txt +0 -0
@@ -1,10 +1,23 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: humancompatible-train
3
- Version: 0.3.2
3
+ Version: 0.4.1
4
4
  Summary: PyTorch-based package for constrained training of neural networks
5
- Author: Gilles Bareilles, Adam Bosak, Jana Lepsova, Jakub Marecek
5
+ Author: Adam Bosak, Gilles Bareilles, Jana Lepsova, Jakub Marecek
6
6
  Author-email: Andrii Kliachkin <kliacand@fel.cvut.cz>
7
7
  Maintainer-email: Andrii Kliachkin <kliacand@fel.cvut.cz>
8
+ License-Expression: Apache-2.0
9
+ Project-URL: Homepage, https://github.com/humancompatible/train
10
+ Project-URL: Repository, https://github.com/humancompatible/train
11
+ Project-URL: Issues, https://github.com/humancompatible/train/issues
12
+ Project-URL: Documentation, https://humancompatible-train.readthedocs.io
13
+ Keywords: constrained-optimization,pytorch,fairness,lagrangian,machine-learning
14
+ Classifier: Development Status :: 3 - Alpha
15
+ Classifier: Intended Audience :: Science/Research
16
+ Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
17
+ Classifier: Programming Language :: Python :: 3
18
+ Classifier: Programming Language :: Python :: 3.11
19
+ Classifier: Programming Language :: Python :: 3.12
20
+ Classifier: Operating System :: OS Independent
8
21
  Requires-Python: >=3.11
9
22
  Description-Content-Type: text/markdown
10
23
  License-File: LICENCE.txt
@@ -23,16 +36,22 @@ Requires-Dist: matplotlib; extra == "benchmark"
23
36
  Requires-Dist: pandas; extra == "benchmark"
24
37
  Requires-Dist: folktables; extra == "benchmark"
25
38
  Requires-Dist: pot; extra == "benchmark"
26
- Requires-Dist: hydra; extra == "benchmark"
39
+ Requires-Dist: hydra-core; extra == "benchmark"
27
40
  Requires-Dist: omegaconf; extra == "benchmark"
41
+ Provides-Extra: paper
42
+ Requires-Dist: hydra-core; extra == "paper"
43
+ Requires-Dist: omegaconf; extra == "paper"
44
+ Requires-Dist: hydra-submitit-launcher; extra == "paper"
28
45
  Provides-Extra: ghost
29
46
  Requires-Dist: qpsolvers; extra == "ghost"
30
47
  Requires-Dist: scipy; extra == "ghost"
48
+ Provides-Extra: compare
49
+ Requires-Dist: cooper-optim>=1.0.1; extra == "compare"
31
50
  Dynamic: license-file
32
51
 
33
52
  # humancompatible-train: a package for constrained machine learning
34
53
 
35
- [![License](https://img.shields.io/badge/License-Apache_2.0-blue.svg)](https://opensource.org/licenses/Apache-2.0) [![Setup](https://github.com/humancompatible/train/actions/workflows/setup.yml/badge.svg)](https://github.com/humancompatible/train/actions/workflows/setup.yml)
54
+ [![License](https://img.shields.io/badge/License-Apache_2.0-blue.svg)](https://opensource.org/licenses/Apache-2.0) [![Setup](https://github.com/humancompatible/train/actions/workflows/setup.yml/badge.svg)](https://github.com/humancompatible/train/actions/workflows/setup.yml) [![docs](https://app.readthedocs.org/projects/humancompatible-train/badge/?version=latest)](https://humancompatible-train.readthedocs.io/en/latest/?badge=latest)
36
55
 
37
56
  The toolkit implements algorithms for constrained training of neural networks based on PyTorch, and inspired by PyTorch's API.
38
57
  <!-- , as well as a tool to compare stochastic-constrained stochastic optimization algorithms on a _fair learning_ task in the `experiments` folder. -->
@@ -61,7 +80,7 @@ The only dependencies of this package are `numpy` and `torch`.
61
80
 
62
81
  ## Using the toolkit
63
82
 
64
- The toolkit implements algorithms for constrained training of neural networks based on PyTorch.
83
+ The toolkit implements algorithms for constrained training of neural networks based on PyTorch. For the documentation, please visit [our Read the Docs page!](https://humancompatible-train.readthedocs.io?version=latest)
65
84
 
66
85
  The algorithms are intended for use in tandem with classic PyTorch optimizers, calculating the Lagrangian and keeping track of the dual variables.
67
86
 
@@ -1,6 +1,6 @@
1
1
  # humancompatible-train: a package for constrained machine learning
2
2
 
3
- [![License](https://img.shields.io/badge/License-Apache_2.0-blue.svg)](https://opensource.org/licenses/Apache-2.0) [![Setup](https://github.com/humancompatible/train/actions/workflows/setup.yml/badge.svg)](https://github.com/humancompatible/train/actions/workflows/setup.yml)
3
+ [![License](https://img.shields.io/badge/License-Apache_2.0-blue.svg)](https://opensource.org/licenses/Apache-2.0) [![Setup](https://github.com/humancompatible/train/actions/workflows/setup.yml/badge.svg)](https://github.com/humancompatible/train/actions/workflows/setup.yml) [![docs](https://app.readthedocs.org/projects/humancompatible-train/badge/?version=latest)](https://humancompatible-train.readthedocs.io/en/latest/?badge=latest)
4
4
 
5
5
  The toolkit implements algorithms for constrained training of neural networks based on PyTorch, and inspired by PyTorch's API.
6
6
  <!-- , as well as a tool to compare stochastic-constrained stochastic optimization algorithms on a _fair learning_ task in the `experiments` folder. -->
@@ -29,7 +29,7 @@ The only dependencies of this package are `numpy` and `torch`.
29
29
 
30
30
  ## Using the toolkit
31
31
 
32
- The toolkit implements algorithms for constrained training of neural networks based on PyTorch.
32
+ The toolkit implements algorithms for constrained training of neural networks based on PyTorch. For the documentation, please visit [our Read the Docs page!](https://humancompatible-train.readthedocs.io?version=latest)
33
33
 
34
34
  The algorithms are intended for use in tandem with classic PyTorch optimizers, calculating the Lagrangian and keeping track of the dual variables.
35
35
 
@@ -0,0 +1,56 @@
1
+ [build-system]
2
+ requires = ["setuptools"]
3
+ build-backend = "setuptools.build_meta"
4
+
5
+ [project]
6
+ name = "humancompatible-train"
7
+ version = "0.4.1"
8
+ dependencies = [
9
+ "torch",
10
+ "numpy",
11
+ ]
12
+ requires-python = ">= 3.11"
13
+ authors = [
14
+ {name = "Andrii Kliachkin", email = "kliacand@fel.cvut.cz"},
15
+ {name = "Adam Bosak"},
16
+ {name = "Gilles Bareilles"},
17
+ {name = "Jana Lepsova"},
18
+ {name = "Jakub Marecek"},
19
+ ]
20
+ maintainers = [
21
+ {name = "Andrii Kliachkin", email = "kliacand@fel.cvut.cz"}
22
+ ]
23
+ description = "PyTorch-based package for constrained training of neural networks"
24
+ readme = "README.md"
25
+ license = "Apache-2.0"
26
+ keywords = ["constrained-optimization", "pytorch", "fairness", "lagrangian", "machine-learning"]
27
+ classifiers = [
28
+ "Development Status :: 3 - Alpha",
29
+ "Intended Audience :: Science/Research",
30
+ "Topic :: Scientific/Engineering :: Artificial Intelligence",
31
+ "Programming Language :: Python :: 3",
32
+ "Programming Language :: Python :: 3.11",
33
+ "Programming Language :: Python :: 3.12",
34
+ "Operating System :: OS Independent",
35
+ ]
36
+
37
+ [project.urls]
38
+ Homepage = "https://github.com/humancompatible/train"
39
+ Repository = "https://github.com/humancompatible/train"
40
+ Issues = "https://github.com/humancompatible/train/issues"
41
+ Documentation = "https://humancompatible-train.readthedocs.io"
42
+
43
+ [project.optional-dependencies]
44
+ examples = ["ipykernel", "ipympl", "fairret", "folktables", "scikit-learn", "matplotlib"]
45
+ # `hydra-core` is the distribution this repo imports. `hydra` on PyPI is an unrelated 2016
46
+ # project, which is what the lockfile had been resolving to (uv.lock pinned Hydra 2.5).
47
+ benchmark = ["fairret", "matplotlib", "pandas", "folktables", "pot", "hydra-core", "omegaconf"]
48
+ # The hyperparameter sweeps in paper/tune.py. The submitit launcher is what `-m
49
+ # hydra/launcher=slurm` needs; the basic sweeper covers the local case with no plugin.
50
+ # hydra-optuna-sweeper is deliberately absent: its 1.2.0 release requires optuna<3.0.0.
51
+ paper = ["hydra-core", "omegaconf", "hydra-submitit-launcher"]
52
+ ghost = ["qpsolvers", "scipy"]
53
+ # Cross-validation against another library (paper/e0/e_cooper.py,
54
+ # tests/test_cooper_parity.py). Pinned to a floor because E0e's claims are about
55
+ # specific update orderings in Cooper's optimizer classes.
56
+ compare = ["cooper-optim>=1.0.1"]
@@ -1,4 +1,7 @@
1
+ from .base import DualOptimizer
1
2
  from .alm import ALM
2
3
  from .ialm import iALM
3
4
  from .pbm import PBM
5
+ from .nupi import nuPI
6
+ from .ssg import SSG
4
7
  from .moreau import MoreauEnvelope
@@ -0,0 +1,348 @@
1
+ import torch
2
+ import torch.distributed as dist
3
+ from typing import Any, Optional, Tuple
4
+ from torch import Tensor
5
+
6
+ from .base import DualOptimizer
7
+
8
+ # cite: Stochastic Smoothed Primal-Dual Algorithms for Nonconvex Optimization with Linear Inequality Constraints
9
+ # https://arxiv.org/pdf/2504.07607
10
+
11
+ AUGMENTATIONS = ("quadratic", "hpr", None)
12
+
13
+
14
+ class ALM(DualOptimizer):
15
+ def __init__(
16
+ self,
17
+ m: int = None,
18
+ lr: float = 0.01,
19
+ init_duals: float | Tensor = None,
20
+ penalty: float = 0.0,
21
+ *,
22
+ dual_range: Tuple[float, float] = (-100.0, 100.0),
23
+ momentum: float = 0.0,
24
+ dampening: Optional[float] = None,
25
+ is_ineq: bool = False,
26
+ restart: bool = False,
27
+ augmentation: Optional[str] = None,
28
+ device=None,
29
+ process_group: Optional[dist.ProcessGroup] = None,
30
+ ) -> None:
31
+
32
+ # Unset + a nonzero penalty means "augment"; HPR is the better default surrogate.
33
+ if augmentation is None and penalty > 0:
34
+ augmentation = "hpr"
35
+
36
+ if augmentation not in AUGMENTATIONS:
37
+ raise ValueError(
38
+ f"Unknown augmentation: {augmentation!r}; expected one of "
39
+ f"{AUGMENTATIONS}"
40
+ )
41
+ if augmentation == "hpr" and penalty <= 0:
42
+ raise ValueError(
43
+ f"The 'hpr' augmentation divides by the penalty, so it requires "
44
+ f"penalty > 0; got {penalty}"
45
+ )
46
+
47
+ self.penalty = penalty
48
+ self.augmentation = augmentation
49
+ params, settings = self._make_group(
50
+ m, lr, momentum, dampening, init_duals, dual_range, is_ineq, restart, device
51
+ )
52
+ super().__init__(
53
+ [{"params": params, **settings}],
54
+ self._scalar_defaults(settings),
55
+ process_group=process_group,
56
+ )
57
+
58
+ @classmethod
59
+ def _make_group(
60
+ cls,
61
+ m: int = None,
62
+ lr: float = None,
63
+ momentum: float = None,
64
+ dampening: float = None,
65
+ init_duals: float | Tensor = None,
66
+ dual_range: Tuple[float, float] = None,
67
+ is_ineq: bool = None,
68
+ restart: bool = None,
69
+ device=None,
70
+ ):
71
+ if momentum is not None and (momentum < 0 or momentum > 1):
72
+ raise ValueError(f"momentum must be within [0,1]; got {momentum}")
73
+
74
+ # Default dampening to momentum (EMA) when unset and momentum > 0; else 0.
75
+ if dampening is None:
76
+ dampening = momentum if (momentum is not None and momentum > 0) else 0.0
77
+
78
+ if not isinstance(restart, bool):
79
+ raise ValueError(
80
+ f"Expected a Boolean value for restart, got {type(restart)}"
81
+ )
82
+
83
+ duals, settings = cls._base_group(
84
+ m, init_duals, dual_range, is_ineq, device
85
+ )
86
+ settings.update(
87
+ {
88
+ "lr": lr,
89
+ "momentum": momentum,
90
+ "dampening": dampening,
91
+ "momentum_buffer": torch.zeros_like(
92
+ duals.data, requires_grad=False, device=device
93
+ ),
94
+ "restart": restart,
95
+ }
96
+ )
97
+ return [duals], cls._drop_none(settings)
98
+
99
+ def add_constraint_group(
100
+ self,
101
+ m: int,
102
+ lr: float = None,
103
+ momentum: float = None,
104
+ dampening: Optional[float] = None,
105
+ init_duals: Tensor = None,
106
+ dual_range: tuple[float, float] = None,
107
+ is_ineq: bool = False,
108
+ restart: bool = False,
109
+ device = None,
110
+ *,
111
+ name: str = None,
112
+ bound: float = None,
113
+ ) -> None:
114
+ """
115
+ Allows to add a group of dual variables with separate initial values and learning rates.
116
+
117
+ :param m: Size of group (number of dual variables to add)
118
+ :type m: int
119
+ :param lr: Dual variable update rate.
120
+ :type lr: float
121
+ :param momentum: Momentum/Smoothing factor for dual variables. Equivalent to SGD momentum. Set to `0` to disable.
122
+ :type momentum: float
123
+ :param dampening: Dampening for momentum. Equivalent to SGD dampening. Set to `0` to disable.
124
+ :type dampening: float
125
+ :param init_duals: Initial values for the new dual variables. Defaults to the value set when creating the optimizer.
126
+ :type init_duals: Tensor
127
+ :param dual_range: After each dual update, the dual variables will be clamped to this range.
128
+ :type dual_range: Tuple[float, float]
129
+ :param is_ineq: Whether to treat the constraints as equality or inequality. If`True`, dual variables will be relaxed on strict satisfaction and lower-bounded by `max(dual_range[0], 0)`.
130
+ :type is_ineq: bool
131
+ :param restart: Whether to set the dual variables to zero immediately on strict satisfaction of corresponding constraints. Not recommended for stochastic constraints.
132
+ :type restart: bool
133
+ :param name: Name for this group, used when passing constraints as a mapping. Defaults to `group<k>`.
134
+ :type name: str
135
+ :param bound: Right-hand side of this group's constraints, if any. Only used by :meth:`violation`.
136
+ :type bound: float
137
+
138
+ .. note::
139
+ Parameters here will default to values set when initializing the dual optimizer.
140
+
141
+ """
142
+ params, settings = self._make_group(
143
+ m, lr, momentum, dampening, init_duals, dual_range, is_ineq, restart, device
144
+ )
145
+ if bound is not None:
146
+ settings["bound"] = bound
147
+ if name is not None:
148
+ settings["name"] = name
149
+ self.add_param_group({"params": params, **settings})
150
+
151
+ # ------------------------------------------------------------------ #
152
+ # hooks
153
+ # ------------------------------------------------------------------ #
154
+
155
+ def _ascent_direction(self, group: dict[str, Any], c: Tensor) -> Tensor:
156
+ """This group's dual ascent direction, i.e. the surrogate's gradient in the duals.
157
+ """
158
+ if self.augmentation in ("quadratic", None) or not group.get("is_ineq"):
159
+ return c
160
+ duals = group["params"][0]
161
+ return (torch.clamp(duals + self.penalty * c, min=0.0) - duals) / self.penalty
162
+
163
+ def _dual_update(self, group: dict[str, Any], c: Tensor) -> None:
164
+ momentum = group.get("momentum", 0.0)
165
+ buffer = group["momentum_buffer"]
166
+
167
+ d = self._ascent_direction(group, c)
168
+
169
+ if momentum > 0:
170
+ _update_c_buffers(d, momentum, group.get("dampening", 0.0), buffer)
171
+
172
+ _update_duals(
173
+ group["params"][0],
174
+ buffer if momentum > 0 else d,
175
+ group["lr"],
176
+ group.get("restart"),
177
+ raw_constraints=c,
178
+ )
179
+
180
+ def _add_constraint_contributions(
181
+ self, lagrangian: Tensor, group: dict[str, Any], snapshot: Any, c: Tensor
182
+ ) -> None:
183
+ if self.augmentation in ("quadratic", None):
184
+ lagrangian.add_(snapshot @ c)
185
+ return
186
+
187
+ rho = self.penalty
188
+ if group.get("is_ineq"):
189
+ with torch.no_grad():
190
+ lam_star = torch.clamp(snapshot + rho * c, min=0.0)
191
+ step = lam_star - snapshot
192
+ lagrangian.add_(lam_star @ c)
193
+ lagrangian.add_(step.dot(step), alpha=-1.0 / (2 * rho))
194
+ else:
195
+ lagrangian.add_(snapshot @ c)
196
+ lagrangian.add_(0.5 * rho * torch.dot(c, c))
197
+
198
+ def _add_global_terms(self, lagrangian: Tensor, constraints: Tensor) -> None:
199
+ # Under HPR the quadratic already sits inside each group's term.
200
+ if self.augmentation == "hpr" or self.penalty == 0:
201
+ return
202
+ # Violations only for inequality groups; see _penalty_constraints.
203
+ c = self._penalty_constraints(constraints)
204
+ lagrangian.add_(0.5 * self.penalty * torch.dot(c, c))
205
+
206
+ def _extra_state(self) -> dict[str, Any]:
207
+ return {"penalty": self.penalty, "augmentation": self.augmentation}
208
+
209
+ def _load_extra_state(self, state: dict[str, Any]) -> None:
210
+ self.penalty = state["penalty"]
211
+ # Checkpoints written before the augmentation option lack the key.
212
+ self.augmentation = state.get("augmentation", "quadratic")
213
+
214
+
215
+ def _update_c_buffers(
216
+ constraints: Tensor,
217
+ momentum: float,
218
+ dampening: float,
219
+ buffer: Tensor,
220
+ ) -> None:
221
+ """Update the constraint buffer with momentum."""
222
+ buffer.mul_(momentum).add_(constraints, alpha=1 - dampening)
223
+
224
+
225
+ def _update_duals(
226
+ duals: Tensor,
227
+ buffer: Tensor,
228
+ lr: float,
229
+ restart: bool,
230
+ raw_constraints: Tensor = None,
231
+ ) -> None:
232
+ """Update duals using the buffered constraint gradients."""
233
+ duals.add_(buffer, alpha=lr)
234
+ if restart:
235
+ # Use raw constraints (not the EMA buffer) to check satisfaction.
236
+ check = raw_constraints if raw_constraints is not None else buffer
237
+ duals.masked_fill_(check < 0, 0.0)
238
+
239
+
240
+
241
+ ALM.__doc__ = (
242
+
243
+ # \textbf{input}: \gamma \text{ (lr) }, \pmb{\lambda}_t \text{ (dual variables, created by method) }, \\
244
+ # \mathbf{c}(\theta) \text{ (constraints) }, f(\theta) \text{ (objective) }, \rho \text{ (penalty coefficient) } \\
245
+ r"""
246
+ A Dual Optimizer that works on the dual maximization tasks according to the (Augmented) Lagrangian rule. Creates and updates dual variables. Reference: https://doi.org/10.48550/arXiv.2504.07607
247
+
248
+ By default (``augmentation=None``, ``penalty=0``), this is dual ascent on the
249
+ plain Lagrangian:
250
+
251
+ .. math::
252
+
253
+ \mathcal{L}_{t+1} & \leftarrow f_t(\theta_{t}) + \pmb{\lambda}_{t+1}^T \mathbf{c}_t(\theta_{t})
254
+
255
+ \pmb{\lambda}_{t+1} & \leftarrow \pmb{\lambda}_t + \gamma \mathbf{c}_t(\theta_{t})
256
+
257
+ Setting ``penalty > 0`` augments the Lagrangian; ``augmentation`` picks the form
258
+ and defaults to ``"hpr"`` when left unset. Passing ``augmentation="quadratic"``
259
+ instead adds a quadratic penalty term; the dual update is unchanged:
260
+
261
+ .. math::
262
+
263
+ \mathcal{L}_{t+1} \leftarrow f_t(\theta_{t}) + \pmb{\lambda}_{t+1}^T \mathbf{c}_t(\theta_{t}) + \frac{\rho}{2} \| \mathbf{c}_t(\theta_{t}) \|^2_2
264
+
265
+ For constraint groups registered with ``is_ineq=True`` the quadratic term acts
266
+ on the violation, :math:`\frac{\rho}{2} \| [\mathbf{c}_t(\theta_t)]_+ \|^2_2`,
267
+ since penalising the raw value of an inequality constraint would also penalise
268
+ being strictly feasible. The linear term and the dual update always use the raw
269
+ values.
270
+
271
+ ``augmentation="hpr"`` (the implicit default whenever ``penalty > 0``) is the
272
+ Hestenes--Powell--Rockafellar augmentation, in which the linear and quadratic
273
+ terms are replaced by a single expression per group,
274
+
275
+ .. math::
276
+
277
+ \mathcal{L}_{t+1} \leftarrow f_t(\theta_t) + \frac{1}{2\rho} \left[
278
+ \big\| \big[ \pmb{\lambda}_{t+1} + \rho \, \mathbf{c}_t(\theta_t) \big]_+ \big\|^2_2
279
+ - \| \pmb{\lambda}_{t+1} \|^2_2 \right]
280
+
281
+ \pmb{\lambda}_{t+1} \leftarrow \left( 1 - \tfrac{\gamma}{\rho} \right) \pmb{\lambda}_t
282
+ + \tfrac{\gamma}{\rho} \big[ \pmb{\lambda}_t + \rho \, \mathbf{c}_t(\theta_t) \big]_+
283
+
284
+ for inequality groups, the dual update again being gradient ascent on the
285
+ surrogate. Writing :math:`\sigma = \pmb{\lambda} + \rho \mathbf{c}` for the trial
286
+ multiplier, the whole difference between the two augmented modes is where the
287
+ clamp sits: the weight the surrogate puts on :math:`\partial \mathbf{c} /
288
+ \partial \theta` is :math:`\max(\pmb{\lambda}, \sigma)` for ``"quadratic"`` but
289
+ :math:`\max(0, \sigma)` for ``"hpr"``. The quadratic penalty therefore never lets
290
+ a multiplier's pull on the primal step fall below :math:`\pmb{\lambda}` however
291
+ feasible the constraint is, whereas HPR switches that pull off entirely once
292
+ :math:`\sigma \le 0`. Three practical consequences:
293
+
294
+ * HPR is a **no-op for equality groups**, where
295
+ :math:`\frac{1}{2\rho}(\|\pmb{\mu} + \rho \mathbf{h}\|^2 - \|\pmb{\mu}\|^2)
296
+ = \pmb{\mu}^T \mathbf{h} + \frac{\rho}{2}\|\mathbf{h}\|^2` identically and the
297
+ dual update is unchanged.
298
+ * At ``lr == penalty`` both augmented modes perform the *same* dual update
299
+ :math:`\pmb{\lambda}_{t+1} = [\pmb{\lambda}_t + \rho \mathbf{c}_t]_+` --
300
+ ``"quadratic"`` reaches it via the non-negativity clamp -- so at that setting
301
+ the choice is purely a primal-side one. ``restart=True`` is also less necessary
302
+ under HPR, since an inactive multiplier stops influencing the primal step while
303
+ still nonzero.
304
+ * The clamp inside the dual update makes it a nonlinear function of the constraint
305
+ estimate, so with *stochastic* constraints HPR biases the multipliers upward
306
+ (Jensen), i.e. toward feasibility. The plain and quadratic dual updates are
307
+ linear in the estimate and carry no such bias.
308
+
309
+ :param m: Number of constraints (determines the number of dual variables to create)
310
+ :type m: int
311
+ :param lr: Dual variable update rate.
312
+ :type lr: float
313
+ :param init_duals: Initial values for the new dual variables. Defaults to 0 for all.
314
+ :type init_duals: float | Tensor
315
+ :param penalty: Augmented Lagrangian penalty parameter. Defaults to`0.`(no augmentation term). A nonzero value auto-selects the`"hpr"`augmentation unless`augmentation`is set explicitly.
316
+ :type penalty: float
317
+ :param dual_range: Safeguarding range for dual variables; they will be`clamp`-ed to this range.
318
+ :type dual_range: Tuple[float, float]
319
+ :param momentum: Momentum/Smoothing factor for dual variables. Equivalent to SGD momentum. Set to `0` to disable.
320
+ :type momentum: float
321
+ :param dampening: Dampening for momentum. Equivalent to SGD dampening. Set to `0` to disable.
322
+ :type dampening: float
323
+ :param is_ineq: Whether to treat the constraints as equality or inequality. If`True`, dual variables will be decreased on strict satisfaction and lower-bounded by `max(dual_range[0], 0)`.
324
+ :type is_ineq: bool
325
+ :param restart: Whether to set the dual variables to zero immediately on strict satisfaction of corresponding constraints. Not recommended for stochastic constraints.
326
+ :type restart: bool
327
+ :param ctol: Reserved for a constraint tolerance allowing tiny violations to account for noise. Accepted for API stability but **currently unused** by the dual update.
328
+ :type ctol: float
329
+ :param augmentation: Which augmentation to form: `None`(default -- no augmentation term if`penalty == 0`, else`"hpr"`), `"quadratic"`, or `"hpr"`. `"hpr"` requires `penalty > 0` and, on inequality groups, replaces the linear-plus-quadratic terms by the Hestenes-Powell-Rockafellar expression above.
330
+ :type augmentation: str, optional
331
+ :param process_group: Distributed process group for DDP. When set, constraint values are averaged across all workers via ``dist.all_reduce`` before each dual update, keeping dual variables consistent across replicas. Defaults to ``None`` (no synchronization).
332
+ :type process_group: dist.ProcessGroup, optional
333
+
334
+ .. note::
335
+ Constraint values may be passed to :meth:`forward` / :meth:`update` /
336
+ :meth:`forward_update` as a flat tensor, as one tensor per constraint
337
+ group, or as a mapping from group name to tensor. See
338
+ :meth:`~humancompatible.train.dual_optim.base.DualOptimizer._gather_constraints`.
339
+
340
+ .. note::
341
+ The HPR term also exists in this module as
342
+ :func:`~humancompatible.train.dual_optim.barrier.augmented_lagrangian`, which
343
+ :class:`~.pbm.PBM` reproduces with ``pbf="augmented_lagrangian"`` and
344
+ per-coordinate penalties :math:`p = \lambda / \rho`. That parametrisation is not
345
+ reused here because it divides by :math:`\lambda` and so breaks at
346
+ :math:`\lambda = 0`.
347
+ """
348
+ )