humancompatible-train 0.1.0__py3-none-any.whl

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 (32) hide show
  1. experiments/__init__.py +0 -0
  2. experiments/calculate_iteration_values.py +310 -0
  3. experiments/run_folktables.py +408 -0
  4. experiments/run_folktables_torchalgs.py +956 -0
  5. humancompatible/__init__.py +0 -0
  6. humancompatible/train/__init__.py +3 -0
  7. humancompatible/train/algorithms/Algorithm.py +25 -0
  8. humancompatible/train/algorithms/__init__.py +8 -0
  9. humancompatible/train/algorithms/ghost.py +250 -0
  10. humancompatible/train/algorithms/sgd.py +107 -0
  11. humancompatible/train/algorithms/ssl_alm.py +311 -0
  12. humancompatible/train/algorithms/switching_subgradient.py +192 -0
  13. humancompatible/train/algorithms/torch/__init__.py +4 -0
  14. humancompatible/train/algorithms/torch/ssl_alm.py +212 -0
  15. humancompatible/train/algorithms/torch/ssw.py +155 -0
  16. humancompatible/train/algorithms/utils.py +61 -0
  17. humancompatible/train/constraints/__init__.py +11 -0
  18. humancompatible/train/constraints/constraint.py +87 -0
  19. humancompatible/train/constraints/constraint_fns.py +118 -0
  20. humancompatible/train/fairness/__init__.py +0 -0
  21. humancompatible/train/fairness/constraints/__init__.py +15 -0
  22. humancompatible/train/fairness/constraints/constraint.py +97 -0
  23. humancompatible/train/fairness/constraints/constraint_fns.py +244 -0
  24. humancompatible/train/fairness/constraints/torch/__init__.py +1 -0
  25. humancompatible/train/fairness/constraints/torch/constraints.py +36 -0
  26. humancompatible/train/fairness/utils/__init__.py +1 -0
  27. humancompatible/train/fairness/utils/balanced_batch_sampler.py +65 -0
  28. humancompatible_train-0.1.0.dist-info/METADATA +188 -0
  29. humancompatible_train-0.1.0.dist-info/RECORD +32 -0
  30. humancompatible_train-0.1.0.dist-info/WHEEL +5 -0
  31. humancompatible_train-0.1.0.dist-info/licenses/LICENCE.txt +201 -0
  32. humancompatible_train-0.1.0.dist-info/top_level.txt +2 -0
@@ -0,0 +1,97 @@
1
+ from typing import Callable, Iterable
2
+
3
+ import numpy as np
4
+ import fairret
5
+ import torch
6
+ from torch.utils.data import DataLoader, SubsetRandomSampler
7
+
8
+
9
+ def _make_dataloaders(dataset, group_indices, batch_size, device, drop_last, gen=None):
10
+ dataloaders = []
11
+ for idx in group_indices:
12
+ sampler = SubsetRandomSampler(idx, gen)
13
+ dataloaders.append(iter(DataLoader(dataset, batch_size, sampler=sampler, drop_last=True)))
14
+ return dataloaders
15
+
16
+
17
+ class FairnessConstraint:
18
+ def __init__(
19
+ self,
20
+ dataset: torch.utils.data.Dataset,
21
+ group_indices: Iterable[Iterable[int]],
22
+ fn: Callable,
23
+ batch_size: int = None,
24
+ use_dataloaders=True,
25
+ device="cpu",
26
+ seed=None,
27
+ loader_drop_last=False,
28
+ ):
29
+ self.dataset = dataset
30
+ self.group_sets = [
31
+ torch.utils.data.Subset(dataset, idx) for idx in group_indices
32
+ ]
33
+ self._group_indices = group_indices
34
+ self.fn = fn
35
+ self._seed = seed
36
+ self._rng = np.random.default_rng(seed)
37
+ self._torch_rng = torch.manual_seed(seed) if seed is not None else torch.Generator(device=device)
38
+ self._device = device
39
+ self._drop_last = loader_drop_last
40
+ if batch_size is not None:
41
+ self._batch_size = batch_size
42
+ if use_dataloaders:
43
+ self.group_dataloaders = _make_dataloaders(
44
+ dataset, group_indices, batch_size, device, gen=self._torch_rng, drop_last=loader_drop_last
45
+ )
46
+
47
+ def group_sizes(self):
48
+ return [len(group) for group in self.group_sets]
49
+
50
+ def eval(self, net, sample, **kwargs):
51
+ return self.fn(net, sample, **kwargs)
52
+
53
+ # def eval_fairret(self, net, sample, group_id, **kwargs):
54
+ # if self._fairret_results is None:
55
+ # statistic = fairret.statistic.TruePositiveRate()
56
+ # loss = fairret.loss.NormLoss(statistic)
57
+
58
+
59
+ # return self._fairret_results[group_id]
60
+
61
+ def sample_loader(self):
62
+ self._fairret_results = None
63
+ samples = []
64
+ for i, l in enumerate(self.group_dataloaders):
65
+ try:
66
+ sample = next(l)
67
+ except StopIteration:
68
+ sampler = SubsetRandomSampler(self._group_indices[i], self._torch_rng)
69
+ l = iter(DataLoader(self.dataset, self._batch_size, sampler=sampler, drop_last=self._drop_last))
70
+ sample = next(l)
71
+ self.group_dataloaders[i] = l
72
+
73
+ samples.append(sample)
74
+ return samples
75
+
76
+ def sample_dataset(
77
+ self, N, rng: np.random.Generator = None, indices=None, return_indices=False
78
+ ):
79
+ if rng is None:
80
+ rng = self._rng
81
+
82
+ if indices is None:
83
+ indices = []
84
+ # returns len(group) points if N > len(group)
85
+ for group in self.group_sets:
86
+ indices.append(
87
+ rng.choice(group.indices, N)
88
+ if N < len(group)
89
+ else rng.choice(group.indices, len(group))
90
+ )
91
+
92
+ sample = [self.dataset[indices[i]] for i, _ in enumerate(self.group_sets)]
93
+
94
+ if return_indices:
95
+ return sample, indices
96
+ else:
97
+ return sample
@@ -0,0 +1,244 @@
1
+ import torch
2
+ from fairret.statistic import (
3
+ TruePositiveRate,
4
+ FalseNegativeFalsePositiveFraction,
5
+ FalsePositiveRate,
6
+ PositiveRate,
7
+ Accuracy,
8
+ )
9
+ from fairret.loss import NormLoss
10
+
11
+
12
+ def tpr_equality(_, net, c_data):
13
+ statistic = TruePositiveRate()
14
+ loss = NormLoss(statistic, p=1)
15
+
16
+ return fairret_stat_equality(net, c_data, loss)
17
+
18
+
19
+ def ppv_equality(_, net, c_data):
20
+ statistic = FalseNegativeFalsePositiveFraction()
21
+ loss = NormLoss(statistic, p=1)
22
+
23
+ return fairret_stat_equality(net, c_data, loss)
24
+
25
+
26
+ def acc_equality(_, net, c_data):
27
+ statistic = Accuracy()
28
+ loss = NormLoss(statistic, p=1)
29
+
30
+ return fairret_stat_equality(net, c_data, loss)
31
+
32
+
33
+ def fairret_stat_equality(net, c_data, loss):
34
+ g1_inputs, g1_labels = c_data[0]
35
+ g2_inputs, g2_labels = c_data[1]
36
+
37
+ g1_outs = net(g1_inputs).squeeze()
38
+ g2_outs = net(g2_inputs).squeeze()
39
+ # if not (1 in a_labels or 1 in g2_labels):
40
+ # return torch.tensor(0)
41
+
42
+ group_codes = [0] * len(g1_labels) + [1] * len(g2_labels)
43
+ group_codes = torch.tensor(
44
+ [[0.0, 1.0] if x == 1 else [1.0, 0.0] for x in group_codes]
45
+ )
46
+
47
+ return loss(
48
+ torch.concat([g1_outs, g2_outs]).unsqueeze(1),
49
+ group_codes,
50
+ torch.concat([g1_labels, g2_labels]).unsqueeze(1),
51
+ )
52
+
53
+
54
+ def dummy(_, net, c_data):
55
+ r = torch.zeros(1)
56
+ r.grad = 0
57
+ return r
58
+
59
+
60
+ def loss_equality(loss, net, c_data):
61
+ g1_inputs, g1_labels = c_data[0]
62
+ g2_inputs, g2_labels = c_data[1]
63
+ g1_outs = net(g1_inputs)
64
+ if g1_labels.ndim == 0:
65
+ g1_labels = g1_labels.reshape(1)
66
+ g2_labels = g2_labels.reshape(1)
67
+ if g1_labels.ndim < g1_outs.ndim:
68
+ g1_labels = g1_labels.unsqueeze(1)
69
+ g2_labels = g2_labels.unsqueeze(1)
70
+ g1_loss = loss(g1_outs, g1_labels)
71
+ g2_outs = net(g2_inputs)
72
+ g2_loss = loss(g2_outs, g2_labels)
73
+
74
+ val = g1_loss - g2_loss
75
+ return val
76
+
77
+
78
+
79
+ def abs_diff_tpr(_, net, c_data):#, stat):
80
+ tpr = TruePositiveRate()
81
+ g1_inputs, g1_labels = c_data[0]
82
+ g2_inputs, g2_labels = c_data[1]
83
+ g1_outs = torch.nn.functional.sigmoid(net(g1_inputs))
84
+ g2_outs = torch.nn.functional.sigmoid(net(g2_inputs))
85
+
86
+ g1_pos_pred_mask = (g1_outs >= 0).squeeze()
87
+ g2_pos_pred_mask = (g2_outs >= 0).squeeze()
88
+
89
+ if g1_labels.ndim == 0:
90
+ g1_labels = g1_labels.reshape(1)
91
+ g2_labels = g2_labels.reshape(1)
92
+ if g1_labels.ndim < g1_outs.ndim:
93
+ g1_labels = g1_labels.unsqueeze(1)
94
+ g2_labels = g2_labels.unsqueeze(1)
95
+
96
+ g1_loss = tpr(g1_outs[g1_pos_pred_mask], None, g1_labels[g1_pos_pred_mask])
97
+ g2_loss = tpr(g2_outs[g2_pos_pred_mask], None, g2_labels[g2_pos_pred_mask])
98
+
99
+ val = abs(g1_loss - g2_loss)
100
+
101
+ return val
102
+
103
+ def abs_diff_pr(_, net, c_data):#, stat):
104
+ pr = PositiveRate()
105
+ g1_inputs, g1_labels = c_data[0]
106
+ g2_inputs, g2_labels = c_data[1]
107
+ g1_outs = torch.nn.functional.sigmoid(net(g1_inputs))
108
+ g2_outs = torch.nn.functional.sigmoid(net(g2_inputs))
109
+
110
+ if g1_labels.ndim == 0:
111
+ g1_labels = g1_labels.reshape(1)
112
+ g2_labels = g2_labels.reshape(1)
113
+ if g1_labels.ndim < g1_outs.ndim:
114
+ g1_labels = g1_labels.unsqueeze(1)
115
+ g2_labels = g2_labels.unsqueeze(1)
116
+
117
+ g1_loss = pr(g1_outs, None)
118
+ g2_loss = pr(g2_outs, None)
119
+
120
+ val = abs(g1_loss - g2_loss)
121
+
122
+ return val
123
+
124
+ def abs_diff_fpr(_, net, c_data):#, stat):
125
+ tpr = FalsePositiveRate()
126
+ g1_inputs, g1_labels = c_data[0]
127
+ g2_inputs, g2_labels = c_data[1]
128
+ g1_outs = torch.nn.functional.sigmoid(net(g1_inputs))
129
+ g2_outs = torch.nn.functional.sigmoid(net(g2_inputs))
130
+
131
+ if g1_labels.ndim == 0:
132
+ g1_labels = g1_labels.reshape(1)
133
+ g2_labels = g2_labels.reshape(1)
134
+ if g1_labels.ndim < g1_outs.ndim:
135
+ g1_labels = g1_labels.unsqueeze(1)
136
+ g2_labels = g2_labels.unsqueeze(1)
137
+
138
+ g1_loss = tpr(g1_outs, None, g1_labels)
139
+ g2_loss = tpr(g2_outs, None, g2_labels)
140
+
141
+ val = abs(g1_loss - g2_loss)
142
+
143
+ return val
144
+
145
+ def abs_max_dev_from_overall_tpr(_, net, c_data):
146
+ stats = []
147
+ st = TruePositiveRate()
148
+ for input, label in c_data:
149
+ out = net(input)
150
+ pred_sigm = torch.nn.functional.sigmoid(out)
151
+ pos_pred_mask = (pred_sigm >= 0).squeeze()
152
+ tpr = st(pred_sigm[pos_pred_mask], None, label[pos_pred_mask].unsqueeze(1))
153
+ stats.append(tpr)
154
+
155
+ stats = torch.cat(stats)
156
+ all_inp = torch.cat([x[0] for x in c_data])
157
+ all_lab = torch.cat([x[1] for x in c_data])
158
+ all_out = net(all_inp)
159
+ all_pred_sigm = torch.nn.functional.sigmoid(all_out)
160
+ pos_pred_mask = (all_pred_sigm >= 0).squeeze()
161
+ all_tpr = st(all_pred_sigm[pos_pred_mask], None, all_lab[pos_pred_mask].unsqueeze(1))
162
+
163
+ val = torch.max(
164
+ torch.abs(stats - all_tpr)
165
+ )
166
+
167
+ # val = torch.max(
168
+ # torch.abs(stats/all_tpr - 1)
169
+ # )
170
+
171
+ return val
172
+
173
+
174
+
175
+ def abs_max_dev_from_overall_fpr(_, net, c_data):
176
+ stats = []
177
+ st = FalsePositiveRate()
178
+ for input, label in c_data:
179
+ out = net(input)
180
+ pred_sigm = torch.nn.functional.sigmoid(out)
181
+ tpr = st(pred_sigm, None, label.unsqueeze(1))
182
+ stats.append(tpr)
183
+
184
+ stats = torch.cat(stats)
185
+ all_inp = torch.cat([x[0] for x in c_data])
186
+ all_lab = torch.cat([x[1] for x in c_data])
187
+ all_out = net(all_inp)
188
+ all_pred_sigm = torch.nn.functional.sigmoid(all_out)
189
+ all_tpr = st(all_pred_sigm, None, all_lab.unsqueeze(1))
190
+
191
+ val = torch.max(
192
+ torch.abs(stats - all_tpr)
193
+ )
194
+
195
+ # val = torch.max(
196
+ # torch.abs(stats/all_tpr - 1)
197
+ # )
198
+
199
+ return val
200
+
201
+
202
+ def abs_loss_equality(loss, net, c_data):
203
+ g1_inputs, g1_labels = c_data[0]
204
+ g2_inputs, g2_labels = c_data[1]
205
+ g1_outs = net(g1_inputs)
206
+ if g1_labels.ndim == 0:
207
+ g1_labels = g1_labels.reshape(1)
208
+ g2_labels = g2_labels.reshape(1)
209
+ if g1_labels.ndim < g1_outs.ndim:
210
+ g1_labels = g1_labels.unsqueeze(1)
211
+ g2_labels = g2_labels.unsqueeze(1)
212
+ g1_loss = loss(g1_outs, g1_labels)
213
+ g2_outs = net(g2_inputs)
214
+ g2_loss = loss(g2_outs, g2_labels)
215
+
216
+ val = g1_loss - g2_loss
217
+ return torch.abs(val)
218
+
219
+
220
+ def fairret_constr(loss, net, c_data):
221
+ g1_inputs, g1_labels = c_data[0]
222
+ g2_inputs, g2_labels = c_data[1]
223
+ g1_logits = net(g1_inputs)
224
+ g2_logits = net(g2_inputs)
225
+ g1_onehot = torch.tensor([[0.0, 1.0]] * len(g1_inputs))
226
+ g2_onehot = torch.tensor([[1.0, 0.0]] * len(g2_inputs))
227
+ logits = torch.concat([g1_logits, g2_logits])
228
+ sens = torch.vstack([g1_onehot, g2_onehot])
229
+ labels = torch.hstack([g1_labels, g2_labels]).unsqueeze(1)
230
+
231
+ return loss(logits, sens, label=labels)
232
+
233
+
234
+ def fairret_pr_constr(loss, net, c_data):
235
+ g1_inputs, _ = c_data[0]
236
+ g2_inputs, _ = c_data[1]
237
+ g1_logits = net(g1_inputs)
238
+ g2_logits = net(g2_inputs)
239
+ g1_onehot = torch.tensor([[0.0, 1.0]] * len(g1_inputs))
240
+ g2_onehot = torch.tensor([[1.0, 0.0]] * len(g2_inputs))
241
+ logits = torch.concat([g1_logits, g2_logits])
242
+ sens = torch.vstack([g1_onehot, g2_onehot])
243
+
244
+ return loss(logits, sens)
@@ -0,0 +1 @@
1
+ from .constraints import loss_equality
@@ -0,0 +1,36 @@
1
+ import torch
2
+ from torch import Tensor
3
+ from torch.nn import Module
4
+
5
+ def loss_equality(preds: Tensor, sens: Tensor, labels: Tensor, criterion: Module = None, diff_to_overall: bool = False):
6
+ """
7
+ A constraint that penalizes the sum of difference between the loss of each group and the overall loss if`diff_to_overall`is`True`,
8
+ and the absolute difference in loss between the two groups if`diff_to_overall`is`False`.
9
+
10
+ Args:
11
+ logits (torch.Tensor): Predictions of shape :math:`(N)`, as we assume to be performing binary
12
+ classification or regression.
13
+ sens (torch.Tensor): One-hot encoding of group membership of shape`(N, S)`with`S`the number of sensitive features.
14
+ `S`must be 2 if`diff_to_overall`is`True`.
15
+ labels (torch.Tensor): Predictions of shape`(N)`.
16
+ loss: (torch.nn.Module): The loss function to calculate.
17
+ diff_to_overall: (bool): Determines whether to penalize the sum of the absolute difference between
18
+ each group's loss and the overall loss if`True`, or the absolute difference in losses of two groups otherwise.
19
+ """
20
+ if not diff_to_overall and sens.shape[-1] != 2:
21
+ raise ValueError(f"If`diff_to_overall` is`False`, expected`sens.shape[-1]` to be 2, got {sens.shape[-1]}")
22
+
23
+ if criterion is None:
24
+ criterion = torch.nn.BCEWithLogitsLoss()
25
+
26
+ sens_t = sens.T
27
+ group_losses = torch.empty(sens.shape[-1])
28
+ for group in range(sens.shape[-1]):
29
+ group_preds, group_labels = preds[sens_t[group] == 1], labels[sens_t[group] == 1]
30
+ group_losses[group] = criterion(group_preds.squeeze(), group_labels)
31
+
32
+ if not diff_to_overall:
33
+ return torch.abs(group_losses[0] - group_losses[1])
34
+
35
+ overall_loss = criterion(preds.squeeze(), labels)
36
+ return torch.sum(torch.abs(overall_loss-group_losses))
@@ -0,0 +1 @@
1
+ from .balanced_batch_sampler import BalancedBatchSampler
@@ -0,0 +1,65 @@
1
+ import numpy as np
2
+ import torch
3
+ from torch.utils.data import Sampler
4
+
5
+ class BalancedBatchSampler(Sampler):
6
+ def __init__(self, subgroup_onehot=None, subgroup_indices=None, batch_size=1, drop_last=True):
7
+ """
8
+ A Sampler that yields an equal number of samples from each group specified with either one-hot encoding or indices.
9
+
10
+ Args:
11
+ subset_indices (list of list): List of indices for each subset. Defaults to None.
12
+ subgroup_onehot (tensor): Tensor of one-hot-encoded group memberships of shape `(N, S)`, where`S`is the number of subgroups. Defaults to None.
13
+ batch_size (int): Number of samples per batch.
14
+ drop_last (bool): If True, drop the last incomplete batch.
15
+ """
16
+
17
+ if subgroup_indices is None and subgroup_onehot is None:
18
+ raise ValueError(f"Exactly one of`subgroup_indices`,`subgroup_onehot`must be`None`")
19
+
20
+ if subgroup_onehot is not None:
21
+ subgroup_onehot = subgroup_onehot.numpy()
22
+ subgroup_indices = [
23
+ np.argwhere(subgroup_onehot[:, gr] == 1).squeeze() for gr in range(subgroup_onehot.shape[-1])
24
+ ]
25
+
26
+ self.subset_indices = subgroup_indices
27
+ self.batch_size = batch_size
28
+ if drop_last is False:
29
+ raise NotImplementedError('drop_last=True not supported yet!')
30
+ self.drop_last = drop_last
31
+ self.n_subsets = len(subgroup_indices)
32
+ self.subset_sizes = [len(indices) for indices in subgroup_indices]
33
+ self.n_samples_per_subset = batch_size // self.n_subsets
34
+ # Check if batch_size is divisible by the number of subsets
35
+ assert batch_size % self.n_subsets == 0, (
36
+ f"Batch size ({batch_size}) must be divisible by the number of subsets ({self.n_subsets})."
37
+ )
38
+
39
+ def __iter__(self):
40
+ # Shuffle indices within each subset
41
+ shuffled_subset_indices = [torch.randperm(len(indices)).tolist() for indices in self.subset_indices]
42
+
43
+ # Calculate the maximum number of batches per subset
44
+ max_batches = min(len(indices) // self.n_samples_per_subset for indices in self.subset_indices)
45
+ if not self.drop_last and any(len(indices) % self.n_samples_per_subset != 0 for indices in self.subset_indices):
46
+ max_batches += 1 # Include partial batches if drop_last is False
47
+ # TODO: randomly permute the batch as well
48
+ # Yield balanced batches
49
+ for batch_idx in range(max_batches):
50
+ batch = []
51
+ for subset_idx in range(self.n_subsets):
52
+ start = batch_idx * self.n_samples_per_subset
53
+ end = start + self.n_samples_per_subset
54
+ subset_batch_indices = shuffled_subset_indices[subset_idx][start:end]
55
+ batch.extend([self.subset_indices[subset_idx][i] for i in subset_batch_indices])
56
+
57
+ # Yield the global indices for the batch
58
+ yield batch
59
+
60
+ def __len__(self):
61
+ if self.drop_last:
62
+ return min(len(indices) // self.n_samples_per_subset for indices in self.subset_indices)
63
+ else:
64
+ return max((len(indices) + self.n_samples_per_subset - 1) // self.n_samples_per_subset
65
+ for indices in self.subset_indices)
@@ -0,0 +1,188 @@
1
+ Metadata-Version: 2.4
2
+ Name: humancompatible-train
3
+ Version: 0.1.0
4
+ Summary: PyTorch-based package for constrained training of neural networks
5
+ Author: Gilles Bareilles, Jana Lepsova, Jakub Marecek
6
+ Author-email: Andrii Kliachkin <kliachkin.andrii@gmail.com>
7
+ Maintainer-email: Andrii Kliachkin <kliachkin.andrii@gmail.com>
8
+ Requires-Python: >=3.11
9
+ Description-Content-Type: text/markdown
10
+ License-File: LICENCE.txt
11
+ Requires-Dist: torch
12
+ Requires-Dist: numpy
13
+ Provides-Extra: ghost
14
+ Requires-Dist: qpsolvers; extra == "ghost"
15
+ Requires-Dist: scipy; extra == "ghost"
16
+ Dynamic: license-file
17
+
18
+ # Benchmarking Stochastic Approximation Algorithms for Fairness-Constrained Training of Deep Neural Networks
19
+
20
+ [![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)
21
+
22
+ The toolkit implements algorithms for constrained training of neural networks based on PyTorch, and inspired by PyTorch's API, as well as a tool to compare stochastic-constrained stochastic optimization algorithms on a _fair learning_ task in the `experiments` folder.
23
+
24
+ ## Table of Contents
25
+ 1. [Basic installation instructions](#basic-installation-instructions)
26
+ 2. [Using the toolkit](#using-the-toolkit)
27
+ 3. [Reproducing the Benchmark](#reproducing-the-benchmark)
28
+ 4. [Extending the toolkit](#extending-the-toolkit) <!-- 6. [Citing humancompatible/train](#Citing-humancompatible/train) -->
29
+ 5. [License and terms of use](#license-and-terms-of-use)
30
+ 6. [References](#references)
31
+
32
+ Humancompatible/train is still under active development! If you find bugs or have feature
33
+ requests, please file a
34
+ [Github issue](https://github.com/humancompatible/train/issues).
35
+
36
+ ## Basic installation instructions
37
+ The code requires Python version ```3.11```.
38
+
39
+ 1. Create a virtual environment
40
+
41
+ **bash** (Linux)
42
+ ```
43
+ python3.11 -m venv fairbenchenv
44
+ source fairbenchenv/bin/activate
45
+ ```
46
+ **cmd** (Windows)
47
+ ```
48
+ python -m venv fairbenchenv
49
+ fairbenchenv\Scripts\activate.bat
50
+ ```
51
+ 2. Install from source.
52
+ ```
53
+ git clone https://github.com/humancompatible/train.git
54
+ cd train
55
+ pip install -r requirements.txt
56
+ pip install .
57
+ ```
58
+
59
+ If you wish to edit the code of the algorithms, install as an editable package:
60
+ ```
61
+ pip install -e .
62
+ ```
63
+
64
+ __Warning__: it is recommended to use Stochastic Ghost with the mkl-accelerated version of the scipy package with Stochastic Ghost; to install it, run
65
+
66
+ ```pip install --force-reinstall -i https://software.repos.intel.com/python/pypi scipy```
67
+
68
+ after installing requirements.txt; otherwise, the algorithm will run slower. However, this is not supported on MacOS and may fail on some Windows devices.
69
+
70
+ <!-- Install via pip -->
71
+ <!-- ``` -->
72
+ <!-- pip install folktables -->
73
+ <!-- ``` -->
74
+
75
+ ## Using the toolkit
76
+
77
+ The toolkit implements algorithms for constrained training of neural networks based on PyTorch, and inspired by PyTorch's API.
78
+
79
+ ### Code examples
80
+
81
+ You are invited to check out the new API presented in notebooks in the `examples` folder.
82
+
83
+ The algorithms follow the `dual_step()` - `step()` framework: taking inspiration from PyTorch, the `double_step` does updates related to the dual parameters and prepares for the primal update (by, e.g., saving constraint gradients), and `step()` updates the primal parameters.
84
+
85
+ The idea is to make different algorithms nearly interchangable in the code.
86
+
87
+ The legacy API used for the benchmark is presented in `examples/_old_/algorithm_demo.ipynb` and `examples/_old_/constraint_demo.ipynb`.
88
+
89
+ ## Reproducing the Benchmark
90
+
91
+ ### Running the algorithms
92
+
93
+ The benchmark comprises the following algorithms:
94
+ - Stochastic Ghost [[2]](#2),
95
+ - SSL-ALM [[3]](#3),
96
+ - Stochastic Switching Subgradient [[4]](#4).
97
+
98
+ To reproduce the experiments of the paper, run the following:
99
+ ``` python
100
+ cd experiments
101
+ python run_folktables.py data=folktables alg=sslalm
102
+ python run_folktables.py data=folktables alg=alm
103
+ python run_folktables.py data=folktables alg=ghost
104
+ python run_folktables.py data=folktables alg=ssg
105
+ python run_folktables.py data=folktables alg=sgd # baseline, no fairness
106
+ python run_folktables.py data=folktables alg=fairret # baseline, fairness with regularizer
107
+ ```
108
+ Each command will start 10 runs of the `alg`, 30 seconds each.
109
+ The results will be saved to `experiments/utils/saved_models` and `experiments/utils/exp_results`.
110
+ <!-- In the repository, we include the configuration needed to reproduce the experiments in the paper. To do so, go to `experiments` and run `python run_folktables.py data=folktables alg=sslalm`. -->
111
+ <!-- Repeat for the other algorithms by changing the `alg` parameter. -->
112
+
113
+ This repository uses [Hydra](https://hydra.cc/) to manage parameters; see `experiments/conf` for configuration files.
114
+ * To change the parameters of the experiment, such as the number of runs for each algorithm, run time, the dataset used (*note: for now supports only Folktables*) - use `experiment.yaml`.
115
+ * To change the dataset settings - such as file location - or do dataset-specific adjustments - such as the configuration of the protected attributes - use `data/{dataset_name}.yaml`
116
+ * To change algorithm hyperparameters, use `alg/{algorithm_name}.yaml`.
117
+ * To change constraint hyperparameters, use `constraint/{constraint_name}.yaml`
118
+
119
+ <!-- ; it is installed as one of the dependencies. -->
120
+ <!-- To learn more about using Hydra, please check out the [official tutorial](https://hydra.cc/docs/tutorials/basic/your_first_app). -->
121
+
122
+ ### Producing plots
123
+ The plots and tables like the ones in the paper can be produced using the two notebooks. `experiments/algo_plots.ipynb` houses the convergence plots, and `experiments/model_plots.ipynb` - all the others.
124
+
125
+ ## Extending the toolkit
126
+
127
+ ### Adding new code
128
+
129
+ **To add a new algorithm**, you can subclass the PyTorch ```Optimizer``` class and proceed following the API guideline presented above.
130
+
131
+ ## License and terms of use
132
+
133
+ humancompatible/train is provided under the Apache 2.0 Licence.
134
+
135
+ The benchmark part of the package relies on the Folktables package, provided under MIT Licence.
136
+ It provides code to download data from the American Community Survey
137
+ (ACS) Public Use Microdata Sample (PUMS) files managed by the US Census Bureau.
138
+ The data itself is governed by the terms of use provided by the Census Bureau.
139
+ For more information, see https://www.census.gov/data/developers/about/terms-of-service.html
140
+
141
+ <!-- ## Cite this work -->
142
+
143
+ <!-- If you use this work, we encourage you to cite our paper, and the folktables dataset [[1]](#1). -->
144
+
145
+ <!-- ``` -->
146
+ <!-- @article{ding2021retiring, -->
147
+ <!-- title={Retiring Adult: New Datasets for Fair Machine Learning}, -->
148
+ <!-- author={Ding, Frances and Hardt, Moritz and Miller, John and Schmidt, Ludwig}, -->
149
+ <!-- journal={Advances in Neural Information Processing Systems}, -->
150
+ <!-- volume={34}, -->
151
+ <!-- year={2021} -->
152
+ <!-- } -->
153
+ <!-- ``` -->
154
+
155
+ ## Future work
156
+
157
+ - Add more algorithms with PyTorch-like API
158
+ - Add more examples from different fields where constrained training of DNNs is employed
159
+ - Migrate the benchmark to the new API
160
+
161
+ ## References
162
+
163
+ If you use this work, we encourage you to cite [our paper](https://arxiv.org/abs/2507.04033),
164
+
165
+ ```
166
+ @misc{kliachkin2025benchmarkingstochasticapproximationalgorithms,
167
+ title={Benchmarking Stochastic Approximation Algorithms for Fairness-Constrained Training of Deep Neural Networks},
168
+ author={Andrii Kliachkin and Jana Lepšová and Gilles Bareilles and Jakub Mareček},
169
+ year={2025},
170
+ eprint={2507.04033},
171
+ archivePrefix={arXiv},
172
+ primaryClass={cs.LG},
173
+ url={https://arxiv.org/abs/2507.04033},
174
+ }
175
+ ```
176
+
177
+ <a id="1">[1]</a>
178
+ Ding, Hardt & Miller et al. (2021) Retiring Adult: New Datasets for Fair Machine Learning, Curran Associates, Inc..
179
+
180
+ <a id="2">[2]</a>
181
+ Facchinei & Kungurtsev (2023) Stochastic Approximation for Expectation Objective and Expectation Inequality-Constrained Nonconvex Optimization, arXiv.
182
+
183
+ <a id="3">[3]</a>
184
+ Huang, Zhang & Alacaoglu (2025) Stochastic Smoothed Primal-Dual Algorithms for Nonconvex Optimization with Linear Inequality Constraints, arXiv.
185
+
186
+ <a id="4">[4]</a>
187
+ Huang & Lin (2023) Oracle Complexity of Single-Loop Switching Subgradient Methods for Non-Smooth Weakly Convex Functional Constrained Optimization, Curran Associates Inc..
188
+
@@ -0,0 +1,32 @@
1
+ experiments/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
2
+ experiments/calculate_iteration_values.py,sha256=ZkqrARRW7b_hvayoliNN2TlA-KYHteDL2EbJWXlt4tY,10509
3
+ experiments/run_folktables.py,sha256=8pxav2pGAVQP_fzvL0r0q0H5FJWh9OjT-NqFjWnoYWc,13942
4
+ experiments/run_folktables_torchalgs.py,sha256=5SoDXx9e2Wm1An--nWDLIsbA2S2eO20WfvwvxsslToo,36506
5
+ humancompatible/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
6
+ humancompatible/train/__init__.py,sha256=5OfDdJBkjlvz4VJUsgn3TGT0u0qyAdNfThRREzMy2d4,74
7
+ humancompatible/train/algorithms/Algorithm.py,sha256=mt39i5hcr1eaIfGPFwuEVyyEmlhqKFZJRLMINUqK9mE,650
8
+ humancompatible/train/algorithms/__init__.py,sha256=1Euo6eQo362PKuX54EihBlvSWnFW7vHF7TtD6DEsKy0,243
9
+ humancompatible/train/algorithms/ghost.py,sha256=_S6tDWT-98G7O11gp2r0efFSynL8T0MH1xtQqMc6hiU,8247
10
+ humancompatible/train/algorithms/sgd.py,sha256=aqaug5XJqdFTf6z9LrXTnZq_HO7zDcRaa63pJUa6P34,3360
11
+ humancompatible/train/algorithms/ssl_alm.py,sha256=tsGr2IRAHaUm7Vmc-y5eOn9rOi4Hodavdtw8FqtHht4,11190
12
+ humancompatible/train/algorithms/switching_subgradient.py,sha256=qZr_su-nNn_5kOq28w-LBBKwJUf4mZpGgnLzP7eTUOg,6723
13
+ humancompatible/train/algorithms/utils.py,sha256=h32UCmjT-KlbW2LAC_AWkyA6uvHUsl45fNj7n_iXeu4,1742
14
+ humancompatible/train/algorithms/torch/__init__.py,sha256=T4kthURISj0xgLfG_7VVywrJi-XChkSh7Lop-eCSOFI,78
15
+ humancompatible/train/algorithms/torch/ssl_alm.py,sha256=2ufFtVKoYZEzMFMmti6vYjnp414wIe83HQkfLpjaMlw,7907
16
+ humancompatible/train/algorithms/torch/ssw.py,sha256=9HRUBp3RsG1S-Ph39vIKJDTBygv36gsMgg9-AXjrfCY,5301
17
+ humancompatible/train/constraints/__init__.py,sha256=1rhhYnXKFhkrdJr1Cse24M3kw3hrIv5-P_qDxaqMVuY,247
18
+ humancompatible/train/constraints/constraint.py,sha256=T_xw_JU7Kpr7IY02mzw6LvMGdfLmO_rir2q6JIqmvYg,2948
19
+ humancompatible/train/constraints/constraint_fns.py,sha256=wtAx_OOUeq9Xv7hwNiGvoMTiGooNaisXjbWNtaABecc,3323
20
+ humancompatible/train/fairness/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
21
+ humancompatible/train/fairness/constraints/__init__.py,sha256=00l89E-p9pF6C7zMEewLRA7jwmjYuulaJOQPdOuujDA,333
22
+ humancompatible/train/fairness/constraints/constraint.py,sha256=u736vUoFHDyB1ao3TOtAQ3hw18GkRA5jg2AccLUVHAk,3304
23
+ humancompatible/train/fairness/constraints/constraint_fns.py,sha256=--vuLrNYL_NWiZ8iAUhm1qMIdKlBnTJ6NVxmIaArgYQ,6994
24
+ humancompatible/train/fairness/constraints/torch/__init__.py,sha256=y4NZpfk0OnggOItrrwwP57_LM42txDIYB5Cy3NpZUv8,38
25
+ humancompatible/train/fairness/constraints/torch/constraints.py,sha256=38uUXfvkcwPaMUqR6CaUUG2ij0aAPmfgXdiGnPBqMH4,1897
26
+ humancompatible/train/fairness/utils/__init__.py,sha256=ij77OdmUzgMKnx0L4VHaENsprZn3rLy7HcgUCpLBhGE,56
27
+ humancompatible/train/fairness/utils/balanced_batch_sampler.py,sha256=J6eHeUIqZG3qGCgCZSBuT710vuBEcudNK1q4pkhWT8Q,3296
28
+ humancompatible_train-0.1.0.dist-info/licenses/LICENCE.txt,sha256=vjeFbYz-1t5Gs7Nj2kXpLKRiLtRrhlx4drCybRNjNm4,11338
29
+ humancompatible_train-0.1.0.dist-info/METADATA,sha256=-VtMN0aZt-C9OenUSl26diKriYJeVw1_7grM6f8UxMk,8357
30
+ humancompatible_train-0.1.0.dist-info/WHEEL,sha256=_zCd3N1l69ArxyTb8rzEoP9TpbYXkqRFSNOD5OuxnTs,91
31
+ humancompatible_train-0.1.0.dist-info/top_level.txt,sha256=hQ41HOSAFFSm6ksFDFX7HlOfle9UI72wCdD_sbWd4Yk,28
32
+ humancompatible_train-0.1.0.dist-info/RECORD,,
@@ -0,0 +1,5 @@
1
+ Wheel-Version: 1.0
2
+ Generator: setuptools (80.9.0)
3
+ Root-Is-Purelib: true
4
+ Tag: py3-none-any
5
+