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,311 @@
1
+ import timeit
2
+ from copy import deepcopy
3
+ from typing import Callable
4
+
5
+ import numpy as np
6
+ import torch
7
+
8
+ from humancompatible.train.algorithms.Algorithm import Algorithm
9
+ from humancompatible.train.algorithms.utils import _set_weights, net_params_to_tensor
10
+
11
+
12
+ class SSLALM(Algorithm):
13
+ def __init__(
14
+ self, net, data, loss, constraints, custom_project_fn: Callable = None
15
+ ):
16
+ super().__init__(net, data, loss, constraints)
17
+ self.project = custom_project_fn if custom_project_fn else self.project_fn
18
+
19
+ @staticmethod
20
+ def project_fn(x, m):
21
+ for i in range(1, m + 1):
22
+ if x[-i] < 0:
23
+ x[-i] = 0
24
+ return x
25
+
26
+ def optimize(
27
+ self,
28
+ tau=0.01,
29
+ eta=0.05,
30
+ lambda_bound=25.,
31
+ rho=1.,
32
+ mu=2.,
33
+ beta=0.5,
34
+ tau_mult=1.,
35
+ eta_mult=1.,
36
+ batch_size=16,
37
+ epochs=None,
38
+ start_lambda=None,
39
+ max_runtime=None,
40
+ max_iter=None,
41
+ seed=None,
42
+ device="cpu",
43
+ verbose=True,
44
+ use_unbiased_penalty_grad=True,
45
+ save_state_interval=1
46
+ ):
47
+ self.state_history = {}
48
+ self.state_history["params"] = {"w": {}, "dual_ms": {}, "z": {}, "slack": {}}
49
+ # self.history['vars_full'] = {'G': {}, 'f': {}, 'fg': {}, 'c': {}, 'cg': {}}
50
+ self.state_history["values"] = {"G": {}, "f": {}, "fg": {}, "c": {}, "cg": {}}
51
+ self.state_history["time"] = {}
52
+
53
+ m = len(self.constraints)
54
+ slack_vars = torch.zeros(m, requires_grad=True)
55
+ _lambda = (
56
+ torch.zeros(m, requires_grad=True) if start_lambda is None else start_lambda
57
+ )
58
+
59
+ z = torch.concat(
60
+ [net_params_to_tensor(self.net, flatten=True, copy=True), slack_vars]
61
+ )
62
+ z_par = torch.narrow(z, 0, 0, z.shape[-1] - m)
63
+
64
+ c = self.constraints
65
+
66
+ run_start = timeit.default_timer()
67
+
68
+ if epochs is None:
69
+ epochs = np.inf
70
+ if max_iter is None:
71
+ max_iter = np.inf
72
+ if max_runtime is None:
73
+ max_runtime = np.inf
74
+
75
+ gen = torch.Generator(device=device)
76
+ if seed is not None:
77
+ gen = gen.manual_seed(seed)
78
+ loss_loader = torch.utils.data.DataLoader(
79
+ self.dataset, batch_size, shuffle=(gen.device == 'cpu'), generator=gen
80
+ )
81
+ loss_iter = iter(loss_loader)
82
+
83
+ epoch = 0
84
+ iteration = 0
85
+ total_iters = 0
86
+
87
+ ### initial f and f_grad estimate ###
88
+ f_grad_estimate = 0
89
+ pre_loader = torch.utils.data.DataLoader(
90
+ self.dataset, batch_size, shuffle=(gen.device == 'cpu'), generator=gen
91
+ )
92
+ pre_iter = iter(pre_loader)
93
+ (f_inputs, f_labels) = next(pre_iter)
94
+ _, f_grad_estimate = self._objective_estimate(f_inputs, f_labels)
95
+ self.net.zero_grad()
96
+
97
+ ### initial c_val and c_grad estimate ###
98
+ c_sample = [ci.sample_loader() for ci in c]
99
+ _c_val_estimate = self._c_value_estimate(slack_vars, c, c_sample)
100
+ c_val_estimate = torch.concat(_c_val_estimate)
101
+ c_grad_estimate = self._constraint_grad_estimate(slack_vars, _c_val_estimate)
102
+
103
+ ### c_val estimate ###
104
+ if use_unbiased_penalty_grad:
105
+ c_sample = [ci.sample_loader() for ci in c]
106
+ c_val_estimate_2 = torch.concat(self._c_value_estimate(slack_vars, c, c_sample))
107
+ else:
108
+ c_val_estimate_2 = c_val_estimate
109
+
110
+ n_iters_c_satisfied = 0
111
+ percent_iters_c_satisfied = 0
112
+
113
+ while True:
114
+ elapsed = timeit.default_timer() - run_start
115
+ iteration += 1
116
+ total_iters += 1
117
+ if epoch >= epochs or total_iters >= max_iter or elapsed > max_runtime:
118
+ break
119
+
120
+ self.state_history["time"][total_iters] = elapsed
121
+ if total_iters % save_state_interval == 0:
122
+ self.state_history["params"]["w"][total_iters] = deepcopy(
123
+ self.net.state_dict()
124
+ )
125
+ self.state_history["params"]["dual_ms"][total_iters] = (
126
+ _lambda.detach().cpu().numpy()
127
+ )
128
+ self.state_history["params"]["z"][total_iters] = (
129
+ z_par.detach().cpu().numpy()
130
+ )
131
+ self.state_history["params"]["slack"][total_iters] = (
132
+ slack_vars.detach().cpu().numpy()
133
+ )
134
+
135
+ percent_iters_c_satisfied = n_iters_c_satisfied / total_iters
136
+
137
+ try:
138
+ (f_inputs, f_labels) = next(loss_iter)
139
+ except StopIteration:
140
+ epoch += 1
141
+ iteration = 0
142
+ gen = gen
143
+ loss_loader = torch.utils.data.DataLoader(
144
+ self.dataset, batch_size, shuffle=(gen.device == 'cpu'), generator=gen
145
+ )
146
+ loss_iter = iter(loss_loader)
147
+ (f_inputs, f_labels) = next(loss_iter)
148
+ tau *= tau_mult
149
+ eta *= eta_mult
150
+ # rho *= rho_mult
151
+
152
+ ########################
153
+ ## UPDATE MULTIPLIERS ##
154
+ ########################
155
+ self.net.zero_grad()
156
+ slack_vars.grad = None
157
+
158
+ # sample for and calculate self.constraints (lines 2, 3)
159
+ # update multipliers (line 3)
160
+ with torch.no_grad():
161
+ _lambda = _lambda + eta * c_val_estimate
162
+ # dual safeguard (lines 4,5)
163
+ for i, l in enumerate(_lambda):
164
+ if l >= lambda_bound: #or l < 0:
165
+ _lambda[i] = 0
166
+ # if torch.norm(_lambda) >= lambda_bound:
167
+ # _lambda = torch.zeros_like(_lambda, requires_grad=True)
168
+
169
+ x_t = torch.concat(
170
+ [
171
+ net_params_to_tensor(self.net, flatten=True, copy=True),
172
+ slack_vars,
173
+ ]
174
+ )
175
+
176
+ G = (
177
+ f_grad_estimate
178
+ + c_grad_estimate.T @ _lambda
179
+ + rho * (c_grad_estimate.T @ c_val_estimate_2)
180
+ )
181
+
182
+ if mu > 0:
183
+ smoothing = mu * (x_t - z)
184
+ G += smoothing
185
+
186
+ x_t1 = self.project(x_t - tau * G, m)
187
+
188
+ if mu > 0:
189
+ z += beta * (x_t - z)
190
+
191
+ ###################
192
+ ## UPDATE PARAMS ##
193
+ ###################
194
+
195
+ with torch.no_grad():
196
+ _set_weights(self.net, x_t1)
197
+ for i in range(len(slack_vars)):
198
+ slack_vars[i] = x_t1[i - len(slack_vars)]
199
+ # objective gradient
200
+ loss_eval, f_grad_1 = self._objective_estimate(f_inputs, f_labels)
201
+ self.net.zero_grad()
202
+
203
+ # constraint value abd grad (1)
204
+ c_sample = [ci.sample_loader() for ci in c]
205
+ _c_val_1 = self._c_value_estimate(slack_vars, c, c_sample)
206
+ c_val_1 = torch.concat(_c_val_1)
207
+ c_grad_1 = self._constraint_grad_estimate(slack_vars, _c_val_1)
208
+
209
+ # constraint value (2) (independent)
210
+ if use_unbiased_penalty_grad:
211
+ c_sample = [ci.sample_loader() for ci in c]
212
+ c_val_2 = torch.concat(self._c_value_estimate(slack_vars, c, c_sample))
213
+ else:
214
+ c_val_2 = c_val_1
215
+
216
+ f_grad_estimate = f_grad_1
217
+ c_val_estimate = c_val_1
218
+ c_val_estimate_2 = c_val_2
219
+ c_grad_estimate = c_grad_1
220
+
221
+ if total_iters % save_state_interval == 0:
222
+ with torch.no_grad():
223
+ f_grad_par = torch.narrow(
224
+ f_grad_estimate, 0, 0, f_grad_estimate.shape[-1] - m
225
+ )
226
+ c_grad_par = torch.narrow(
227
+ c_grad_estimate, 1, 0, c_grad_estimate.shape[-1] - m
228
+ )
229
+ G_par = torch.narrow(G, 0, 0, G.shape[-1] - m)
230
+ z_par = torch.narrow(z, 0, 0, z.shape[-1] - m)
231
+
232
+ self.state_history["values"]["G"][total_iters] = (
233
+ torch.norm(G_par).detach().cpu().numpy()
234
+ )
235
+ self.state_history["values"]["f"][total_iters] = (
236
+ loss_eval.detach().cpu().numpy()
237
+ )
238
+ self.state_history["values"]["fg"][total_iters] = (
239
+ torch.norm(f_grad_par).detach().cpu().numpy()
240
+ )
241
+ self.state_history["values"]["c"][total_iters] = (
242
+ c_val_2.detach().cpu().numpy()
243
+ )
244
+ self.state_history["values"]["cg"][total_iters] = (
245
+ torch.norm(c_grad_par, dim=1).detach().cpu().numpy()
246
+ )
247
+
248
+ if torch.all(c_val_1 <= 0):
249
+ n_iters_c_satisfied += 1
250
+
251
+ if verbose:
252
+ with np.printoptions(
253
+ precision=3,
254
+ suppress=True,
255
+ floatmode="fixed",
256
+ sign=" ",
257
+ linewidth=200,
258
+ ):
259
+ print(
260
+ f"{epoch:2}|{iteration:5}|{tau:.3f}|"
261
+ # f"{loss_eval.detach().cpu().numpy():1.3f}|"
262
+ f"{_lambda.detach().cpu().numpy()}|"
263
+ f"{c_val_estimate.detach().cpu().numpy() - slack_vars.detach().cpu().numpy()}|",
264
+ # f"{slack_vars.detach().cpu().numpy()} | {100*percent_iters_c_satisfied:2.1f}%",
265
+ end="\r",
266
+ )
267
+
268
+ return self.state_history
269
+
270
+
271
+
272
+ def _c_value_estimate(self, slack_vars, c, c_sample):
273
+ c_val = [
274
+ ci.eval(self.net, c_sample[i]).reshape(1) + slack_vars[i]
275
+ for i, ci in enumerate(c)
276
+ ]
277
+
278
+ return c_val
279
+
280
+ def _objective_estimate(self, f_inputs, f_labels):
281
+ m = len(self.constraints)
282
+ # breakpoint()
283
+ outputs = self.net(f_inputs)
284
+ # if f_labels.dim() < outputs.dim():
285
+ # f_labels = f_labels.unsqueeze(1)
286
+ loss_eval = self.loss_fn(outputs.squeeze(), f_labels)
287
+ f_grad = torch.autograd.grad(loss_eval, self.net.parameters())
288
+ f_grad = torch.concat([*[g.flatten() for g in f_grad], torch.zeros(m)])
289
+
290
+ return loss_eval, f_grad
291
+
292
+ def _constraint_grad_estimate(self, slack_vars, c):
293
+ c_grad = []
294
+ # breakpoint()
295
+ for ci in c:
296
+ ci_grad = torch.autograd.grad(ci, self.net.parameters())
297
+ if slack_vars is None:
298
+ c_grad.append(torch.concat([g.flatten() for g in ci_grad]))
299
+ else:
300
+ slack_grad = torch.autograd.grad(ci, slack_vars, materialize_grads=True)
301
+ # if torch.sum(slack_grad[0]) != 1:
302
+ # breakpoint()
303
+ c_grad.append(
304
+ torch.concat([*[g.flatten() for g in ci_grad], *slack_grad])
305
+ )
306
+ slack_vars.grad = None
307
+ # slack_vars.zero_grad_
308
+
309
+ self.net.zero_grad()
310
+ c_grad = torch.stack(c_grad)
311
+ return c_grad
@@ -0,0 +1,192 @@
1
+ import timeit
2
+ from copy import deepcopy
3
+ from typing import Callable
4
+
5
+ import numpy as np
6
+ import torch
7
+
8
+ from humancompatible.train.algorithms.Algorithm import Algorithm
9
+ from humancompatible.train.algorithms.utils import net_params_to_tensor
10
+
11
+
12
+ class SSG(Algorithm):
13
+ def __init__(
14
+ self, net, data, loss, constraints, custom_project_fn: Callable = None
15
+ ):
16
+ super().__init__(net, data, loss, constraints)
17
+ self.project = custom_project_fn if custom_project_fn else self.project_fn
18
+
19
+ @staticmethod
20
+ def project_fn(x, m):
21
+ return x
22
+
23
+ def optimize(
24
+ self,
25
+ ctol_rule,
26
+ ctol,
27
+ f_stepsize_rule,
28
+ f_stepsize,
29
+ c_stepsize_rule,
30
+ c_stepsize,
31
+ batch_size,
32
+ epochs=None,
33
+ save_iter=None,
34
+ device="cpu",
35
+ seed=None,
36
+ verbose=True,
37
+ max_runtime=None,
38
+ max_iter=None,
39
+ save_state_interval=100
40
+ ):
41
+ self.state_history = {}
42
+ self.state_history["params"] = {"w": {}}
43
+ self.state_history["values"] = {"G": {}, "f": {}, "c": {}}
44
+ self.state_history["time"] = {}
45
+
46
+ f_eta_t = f_stepsize
47
+ c_eta_t = c_stepsize
48
+
49
+ loss_eval = None
50
+ c_t = None
51
+ eta_f_sum = total_iters = iteration = f_iters = c_iters = epoch = 0
52
+ eta_f_list = []
53
+ _ctol = ctol
54
+ if epochs is None:
55
+ epochs = np.inf
56
+ if max_iter is None:
57
+ max_iter = np.inf
58
+
59
+ gen = torch.Generator(device=device)
60
+ loss_loader = torch.utils.data.DataLoader(
61
+ self.dataset, batch_size, shuffle=True, generator=gen
62
+ )
63
+ loss_iter = iter(loss_loader)
64
+
65
+ run_start = timeit.default_timer()
66
+ while True:
67
+ elapsed = timeit.default_timer() - run_start
68
+ iteration += 1
69
+ total_iters += 1
70
+ if epoch >= epochs or total_iters >= max_iter or elapsed > max_runtime:
71
+ break
72
+
73
+ self.state_history["time"][total_iters] = elapsed
74
+ if total_iters % save_state_interval == 0:
75
+ self.state_history["params"]["w"][total_iters] = deepcopy(
76
+ self.net.state_dict()
77
+ )
78
+
79
+ try:
80
+ f_sample = next(loss_iter)
81
+ except StopIteration:
82
+ epoch += 1
83
+ iteration = 0
84
+ loss_loader = torch.utils.data.DataLoader(
85
+ self.dataset, batch_size, shuffle=True, generator=gen
86
+ )
87
+ loss_iter = iter(loss_loader)
88
+ f_sample = next(loss_iter)
89
+
90
+ self.net.zero_grad()
91
+ if ctol_rule == 'dimin':
92
+ _ctol = ctol / np.sqrt(total_iters)
93
+
94
+ if save_iter is not None and total_iters >= save_iter:
95
+ eta_f_list.append(f_eta_t)
96
+ eta_f_sum += f_eta_t
97
+
98
+ # generate sample of constraints
99
+ c_sample = [ci.sample_loader() for ci in self.constraints]
100
+ # calc constraints and update multipliers (line 3)
101
+ with torch.no_grad():
102
+ c_t = np.array(
103
+ [
104
+ ci.eval(self.net, c_sample[i]).reshape(1)
105
+ for i, ci in enumerate(self.constraints)
106
+ ]
107
+ ).flatten()
108
+ c_argmax = np.argmax(c_t)
109
+ c_max = c_t[c_argmax]
110
+
111
+ x_t = net_params_to_tensor(self.net, flatten=True, copy=True)
112
+
113
+ if c_max >= _ctol:
114
+ c_iters += 1
115
+ c_max2 = self.constraints[c_argmax].eval(self.net, c_sample[c_argmax]).reshape(1)
116
+
117
+ c_grad = torch.autograd.grad(c_max2, self.net.parameters())
118
+ c_grad = torch.concat([cg.flatten() for cg in c_grad])
119
+
120
+ if c_stepsize_rule == "adaptive":
121
+ c_eta_t = c_max / (1e-6 + torch.norm(c_grad) ** 2)
122
+ elif c_stepsize_rule == "const":
123
+ c_eta_t = c_stepsize
124
+ elif c_stepsize_rule == "dimin":
125
+ c_eta_t = c_stepsize / np.sqrt(total_iters)
126
+
127
+ x_t1 = self.project(x_t - c_eta_t * c_grad, m=len(self.constraints))
128
+
129
+ else:
130
+ f_iters += 1
131
+ f_inputs, f_labels = f_sample
132
+ outputs = self.net(f_inputs)
133
+ if f_labels.dim() < outputs.dim():
134
+ f_labels = f_labels.unsqueeze(1)
135
+ loss_eval = self.loss_fn(outputs, f_labels)
136
+
137
+ f_grad = torch.autograd.grad(loss_eval, self.net.parameters())
138
+ f_grad = torch.concat([fg.flatten() for fg in f_grad])
139
+
140
+ if f_stepsize_rule == "dimin":
141
+ f_eta_t = f_stepsize / np.sqrt(total_iters)
142
+ elif f_stepsize_rule == "const":
143
+ f_eta_t = f_stepsize
144
+ x_t1 = self.project(x_t - f_eta_t * f_grad, m=len(self.constraints))
145
+
146
+ start = 0
147
+ with torch.no_grad():
148
+ w = net_params_to_tensor(self.net, flatten=False, copy=False)
149
+ for i in range(len(w)):
150
+ end = start + w[i].numel()
151
+ w[i].set_(x_t1[start:end].reshape(w[i].shape))
152
+ start = end
153
+
154
+ if total_iters % save_state_interval == 0:
155
+ if c_max is not None:
156
+ self.state_history["values"]["c"][total_iters] = (
157
+ c_t
158
+ )
159
+ if loss_eval is not None:
160
+ self.state_history["values"]["f"][total_iters] = (
161
+ loss_eval.cpu().detach().numpy()
162
+ )
163
+
164
+ if verbose and loss_eval is not None and c_t is not None:
165
+ with np.printoptions(
166
+ precision=3,
167
+ suppress=True,
168
+ floatmode="fixed",
169
+ sign=" ",
170
+ linewidth=100,
171
+ ):
172
+ print(
173
+ f"{epoch:2}|"
174
+ f"{_ctol:.3}|"
175
+ f"{iteration:5}|"
176
+ f"{loss_eval.detach().cpu().numpy():.5f}|"
177
+ f"{c_t}",
178
+ end="\r",
179
+ )
180
+
181
+ ######################
182
+ ### POSTPROCESSING ###
183
+ ######################
184
+
185
+ if save_iter is not None:
186
+ model_ind = np.random.default_rng(seed=seed).choice(
187
+ np.arange(start=save_iter, stop=total_iters),
188
+ p=np.array(eta_f_list) / np.sum(eta_f_list),
189
+ )
190
+ self.net.load_state_dict(self.state_history["params"]["w"].iloc[model_ind])
191
+
192
+ return self.state_history
@@ -0,0 +1,4 @@
1
+ from .ssl_alm import SSLALM
2
+ from .ssw import SSG
3
+
4
+ __all__ = ["SSLALM", "SSG"]