humancompatible-train 0.4.0__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 (51) hide show
  1. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/PKG-INFO +1 -1
  2. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/pyproject.toml +1 -1
  3. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/dual_optim/alm.py +40 -25
  4. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/fairness/utils/balanced_batch_sampler.py +14 -0
  5. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible_train.egg-info/PKG-INFO +1 -1
  6. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/tests/test_alm.py +1 -1
  7. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/tests/test_balanced_batch_sampler.py +6 -0
  8. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/LICENCE.txt +0 -0
  9. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/README.md +0 -0
  10. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/setup.cfg +0 -0
  11. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/__init__.py +0 -0
  12. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/dual_optim/__init__.py +0 -0
  13. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/dual_optim/barrier.py +0 -0
  14. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/dual_optim/base.py +0 -0
  15. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/dual_optim/ialm.py +0 -0
  16. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/dual_optim/moreau.py +0 -0
  17. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/dual_optim/nupi.py +0 -0
  18. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/dual_optim/pbm.py +0 -0
  19. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/dual_optim/ssg.py +0 -0
  20. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/fairness/__init__.py +0 -0
  21. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/fairness/utils/__init__.py +0 -0
  22. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/sqp/__init__.py +0 -0
  23. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/sqp/ghost/__init__.py +0 -0
  24. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/sqp/ghost/ghost_sqp.py +0 -0
  25. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/sqp/ghost/mlmc.py +0 -0
  26. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/sqp/ghost/oracle.py +0 -0
  27. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/sqp/ghost/sampler.py +0 -0
  28. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/sqp/ghost/subproblems.py +0 -0
  29. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/sqp/nonopt/__init__.py +0 -0
  30. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/sqp/nonopt/direction.py +0 -0
  31. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/sqp/nonopt/inverse_hessian.py +0 -0
  32. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/sqp/nonopt/line_search.py +0 -0
  33. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/sqp/nonopt/optimizer.py +0 -0
  34. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/sqp/nonopt/point_set.py +0 -0
  35. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible/train/sqp/nonopt/qp.py +0 -0
  36. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible_train.egg-info/SOURCES.txt +0 -0
  37. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible_train.egg-info/dependency_links.txt +0 -0
  38. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible_train.egg-info/requires.txt +0 -0
  39. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/src/humancompatible_train.egg-info/top_level.txt +0 -0
  40. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/tests/test_cooper_parity.py +0 -0
  41. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/tests/test_distributed.py +0 -0
  42. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/tests/test_dual_optimizer_base.py +0 -0
  43. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/tests/test_fairness_problems.py +0 -0
  44. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/tests/test_ghost.py +0 -0
  45. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/tests/test_golden.py +0 -0
  46. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/tests/test_ialm.py +0 -0
  47. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/tests/test_llm_gates.py +0 -0
  48. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/tests/test_nonopt.py +0 -0
  49. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/tests/test_nupi.py +0 -0
  50. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/tests/test_pbm.py +0 -0
  51. {humancompatible_train-0.4.0 → humancompatible_train-0.4.1}/tests/test_ssg.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: humancompatible-train
3
- Version: 0.4.0
3
+ Version: 0.4.1
4
4
  Summary: PyTorch-based package for constrained training of neural networks
5
5
  Author: Adam Bosak, Gilles Bareilles, Jana Lepsova, Jakub Marecek
6
6
  Author-email: Andrii Kliachkin <kliacand@fel.cvut.cz>
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "humancompatible-train"
7
- version = "0.4.0"
7
+ version = "0.4.1"
8
8
  dependencies = [
9
9
  "torch",
10
10
  "numpy",
@@ -8,7 +8,7 @@ from .base import DualOptimizer
8
8
  # cite: Stochastic Smoothed Primal-Dual Algorithms for Nonconvex Optimization with Linear Inequality Constraints
9
9
  # https://arxiv.org/pdf/2504.07607
10
10
 
11
- AUGMENTATIONS = ("quadratic", "hpr")
11
+ AUGMENTATIONS = ("quadratic", "hpr", None)
12
12
 
13
13
 
14
14
  class ALM(DualOptimizer):
@@ -17,18 +17,22 @@ class ALM(DualOptimizer):
17
17
  m: int = None,
18
18
  lr: float = 0.01,
19
19
  init_duals: float | Tensor = None,
20
- penalty: float = 1.0,
20
+ penalty: float = 0.0,
21
21
  *,
22
22
  dual_range: Tuple[float, float] = (-100.0, 100.0),
23
23
  momentum: float = 0.0,
24
24
  dampening: Optional[float] = None,
25
25
  is_ineq: bool = False,
26
26
  restart: bool = False,
27
- augmentation: str = "hpr",
27
+ augmentation: Optional[str] = None,
28
28
  device=None,
29
29
  process_group: Optional[dist.ProcessGroup] = None,
30
30
  ) -> None:
31
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
+
32
36
  if augmentation not in AUGMENTATIONS:
33
37
  raise ValueError(
34
38
  f"Unknown augmentation: {augmentation!r}; expected one of "
@@ -151,7 +155,7 @@ class ALM(DualOptimizer):
151
155
  def _ascent_direction(self, group: dict[str, Any], c: Tensor) -> Tensor:
152
156
  """This group's dual ascent direction, i.e. the surrogate's gradient in the duals.
153
157
  """
154
- if self.augmentation == "quadratic" or not group.get("is_ineq"):
158
+ if self.augmentation in ("quadratic", None) or not group.get("is_ineq"):
155
159
  return c
156
160
  duals = group["params"][0]
157
161
  return (torch.clamp(duals + self.penalty * c, min=0.0) - duals) / self.penalty
@@ -176,7 +180,7 @@ class ALM(DualOptimizer):
176
180
  def _add_constraint_contributions(
177
181
  self, lagrangian: Tensor, group: dict[str, Any], snapshot: Any, c: Tensor
178
182
  ) -> None:
179
- if self.augmentation == "quadratic":
183
+ if self.augmentation in ("quadratic", None):
180
184
  lagrangian.add_(snapshot @ c)
181
185
  return
182
186
 
@@ -241,11 +245,22 @@ ALM.__doc__ = (
241
245
  r"""
242
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
243
247
 
248
+ By default (``augmentation=None``, ``penalty=0``), this is dual ascent on the
249
+ plain Lagrangian:
250
+
244
251
  .. math::
245
252
 
253
+ \mathcal{L}_{t+1} & \leftarrow f_t(\theta_{t}) + \pmb{\lambda}_{t+1}^T \mathbf{c}_t(\theta_{t})
254
+
246
255
  \pmb{\lambda}_{t+1} & \leftarrow \pmb{\lambda}_t + \gamma \mathbf{c}_t(\theta_{t})
247
256
 
248
- \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
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
249
264
 
250
265
  For constraint groups registered with ``is_ineq=True`` the quadratic term acts
251
266
  on the violation, :math:`\frac{\rho}{2} \| [\mathbf{c}_t(\theta_t)]_+ \|^2_2`,
@@ -253,9 +268,9 @@ ALM.__doc__ = (
253
268
  being strictly feasible. The linear term and the dual update always use the raw
254
269
  values.
255
270
 
256
- Setting ``augmentation="hpr"`` switches to the Hestenes--Powell--Rockafellar
257
- augmentation, in which the linear and quadratic terms are replaced by a single
258
- expression per group,
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,
259
274
 
260
275
  .. math::
261
276
 
@@ -268,11 +283,11 @@ ALM.__doc__ = (
268
283
 
269
284
  for inequality groups, the dual update again being gradient ascent on the
270
285
  surrogate. Writing :math:`\sigma = \pmb{\lambda} + \rho \mathbf{c}` for the trial
271
- multiplier, the whole difference from the default is where the clamp sits: the
272
- weight the surrogate puts on :math:`\partial \mathbf{c} / \partial \theta` is
273
- :math:`\max(\pmb{\lambda}, \sigma)` for ``"quadratic"`` but
274
- :math:`\max(0, \sigma)` for ``"hpr"``. The default therefore never lets a
275
- multiplier's pull on the primal step fall below :math:`\pmb{\lambda}` however
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
276
291
  feasible the constraint is, whereas HPR switches that pull off entirely once
277
292
  :math:`\sigma \le 0`. Three practical consequences:
278
293
 
@@ -280,16 +295,16 @@ ALM.__doc__ = (
280
295
  :math:`\frac{1}{2\rho}(\|\pmb{\mu} + \rho \mathbf{h}\|^2 - \|\pmb{\mu}\|^2)
281
296
  = \pmb{\mu}^T \mathbf{h} + \frac{\rho}{2}\|\mathbf{h}\|^2` identically and the
282
297
  dual update is unchanged.
283
- * At ``lr == penalty`` both modes perform the *same* dual update
284
- :math:`\pmb{\lambda}_{t+1} = [\pmb{\lambda}_t + \rho \mathbf{c}_t]_+` -- the
285
- default reaches it via the non-negativity clamp -- so at that setting the choice
286
- is purely a primal-side one. ``restart=True`` is also less necessary under HPR,
287
- since an inactive multiplier stops influencing the primal step while still
288
- nonzero.
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.
289
304
  * The clamp inside the dual update makes it a nonlinear function of the constraint
290
305
  estimate, so with *stochastic* constraints HPR biases the multipliers upward
291
- (Jensen), i.e. toward feasibility. The default's dual update is linear in the
292
- estimate and carries no such bias.
306
+ (Jensen), i.e. toward feasibility. The plain and quadratic dual updates are
307
+ linear in the estimate and carry no such bias.
293
308
 
294
309
  :param m: Number of constraints (determines the number of dual variables to create)
295
310
  :type m: int
@@ -297,7 +312,7 @@ ALM.__doc__ = (
297
312
  :type lr: float
298
313
  :param init_duals: Initial values for the new dual variables. Defaults to 0 for all.
299
314
  :type init_duals: float | Tensor
300
- :param penalty: Augmented Lagrangian penalty parameter. Defaults to`1.`
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.
301
316
  :type penalty: float
302
317
  :param dual_range: Safeguarding range for dual variables; they will be`clamp`-ed to this range.
303
318
  :type dual_range: Tuple[float, float]
@@ -311,8 +326,8 @@ ALM.__doc__ = (
311
326
  :type restart: bool
312
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.
313
328
  :type ctol: float
314
- :param augmentation: Which augmentation to form, ``"quadratic"`` (default) or ``"hpr"``. The latter requires `penalty > 0` and, on inequality groups, replaces the linear-plus-quadratic terms by the Hestenes-Powell-Rockafellar expression above.
315
- :type augmentation: str
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
316
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).
317
332
  :type process_group: dist.ProcessGroup, optional
318
333
 
@@ -18,6 +18,12 @@ class BalancedBatchSampler(Sampler):
18
18
  Oversampling reuses a group's samples evenly: over one epoch, each sample of an extended group is drawn either
19
19
  `floor`or`ceil`of the per-sample average, and no sample is ever repeated within a single batch.
20
20
 
21
+ Because every batch draws the same number of samples per group regardless of true group size, the batch loss
22
+ over-represents small groups relative to their population share. The`group_weights`property gives the
23
+ inverse-propensity per-group weight (`n_groups * group_size / total_size`) that corrects for this: passing it
24
+ (indexed/broadcast per sample) as the`weight`argument of a mean-reduced loss recovers an unbiased estimate of
25
+ the population loss.
26
+
21
27
  :param group_indices: List of indices for each group. Defaults to`None`.
22
28
  :type group_indices: Iterable[Iterable[int]]
23
29
  :param group_onehot: Tensor of one-hot-encoded groups memberships of shape`(N, S)`, where`S`is the number of groups. Defaults to`None`.
@@ -140,3 +146,11 @@ class BalancedBatchSampler(Sampler):
140
146
  else self._group_sizes[group_idx]
141
147
  for group_idx in range(self._n_groups)
142
148
  ) // self._n_samples_per_group
149
+
150
+ @property
151
+ def group_weights(self) -> torch.Tensor:
152
+ """Inverse-propensity per-group weight (n_groups * population proportion), to
153
+ pass as `weight=` to a mean-reduced loss and correct for this sampler's equal
154
+ per-batch group representation."""
155
+ sizes = torch.tensor(self._group_sizes, dtype=torch.float32)
156
+ return self._n_groups * sizes / sizes.sum()
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: humancompatible-train
3
- Version: 0.4.0
3
+ Version: 0.4.1
4
4
  Summary: PyTorch-based package for constrained training of neural networks
5
5
  Author: Adam Bosak, Gilles Bareilles, Jana Lepsova, Jakub Marecek
6
6
  Author-email: Andrii Kliachkin <kliacand@fel.cvut.cz>
@@ -269,7 +269,7 @@ class TestALMHPR(unittest.TestCase):
269
269
  # non-negativity clamp; only the surrogate differs.
270
270
  rho = 1.0
271
271
  c = torch.tensor([0.5, -0.3, -2.0])
272
- quad = ALM(m=3, lr=rho, penalty=rho, init_duals=0.6, is_ineq=True)
272
+ quad = ALM(m=3, lr=rho, penalty=rho, init_duals=0.6, is_ineq=True, augmentation='quadratic')
273
273
  hpr = self._hpr(m=3, lr=rho, penalty=rho, init_duals=0.6, is_ineq=True)
274
274
 
275
275
  lag_q = quad.forward_update(self.loss, c)
@@ -66,6 +66,12 @@ class TestBalancedBatchSampler(unittest.TestCase):
66
66
  [i.tolist() for i in sampler._group_indices], self.subset_indices
67
67
  )
68
68
 
69
+ def test_group_weights(self):
70
+ # group sizes [2, 3, 5], n_groups=3 -> weight_g = 3 * size_g / 10
71
+ sampler = BalancedBatchSampler(group_indices=self.subset_indices, batch_size=3)
72
+ expected = torch.tensor([0.6, 0.9, 1.5])
73
+ torch.testing.assert_close(sampler.group_weights, expected)
74
+
69
75
  def test_iter(self):
70
76
  sampler = BalancedBatchSampler(
71
77
  group_indices=self.subset_indices, batch_size=6, drop_last=True