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
File without changes
@@ -0,0 +1,3 @@
1
+ # from .algorithms, .constraints
2
+
3
+ __all__ = ["algorithms", "constraints"]
@@ -0,0 +1,25 @@
1
+ from typing import Callable, Iterable
2
+
3
+ from torch.nn import Module
4
+ from torch.utils.data import Dataset
5
+
6
+ from humancompatible.train.fairness.constraints import FairnessConstraint
7
+
8
+
9
+ class Algorithm:
10
+ def __init__(
11
+ self,
12
+ net: Module,
13
+ data: Dataset,
14
+ loss: Callable,
15
+ constraints: Iterable[FairnessConstraint],
16
+ ):
17
+ self.net = net
18
+ self.constraints = constraints
19
+ self.loss_fn = loss
20
+ self.dataset = data
21
+
22
+ self.history = {"loss": [], "constr": [], "w": [], "time": [], "n_samples": []}
23
+
24
+ def optimize(self, max_runtime: float = None, max_iter: int = None):
25
+ pass
@@ -0,0 +1,8 @@
1
+ from .ghost import StochasticGhost
2
+ from .ssl_alm import SSLALM
3
+ from .switching_subgradient import SSG
4
+ from .sgd import SGD
5
+ # from .torch.ssl_alm import SSLALM
6
+ # from .torch.ssw import SSG
7
+
8
+ __all__ = ["SSLALM", "StochasticGhost", "SSG", "SGD"]
@@ -0,0 +1,250 @@
1
+ # import autoray as ar
2
+ import timeit
3
+ from copy import deepcopy
4
+
5
+ import numpy as np
6
+ import scipy as sp
7
+ import torch
8
+ from qpsolvers import solve_qp
9
+ from scipy.optimize import linprog
10
+
11
+ from humancompatible.train.algorithms.Algorithm import Algorithm
12
+ from humancompatible.train.algorithms.utils import net_params_to_tensor
13
+
14
+
15
+ class StochasticGhost(Algorithm):
16
+ def __init__(self, net, data, loss, constraints):
17
+ super().__init__(net, data, loss, constraints)
18
+
19
+ @staticmethod
20
+ def solvesubp(
21
+ fgrad,
22
+ cval,
23
+ cgrad,
24
+ kap_val,
25
+ beta,
26
+ tau,
27
+ hesstype,
28
+ mc,
29
+ n,
30
+ qp_solver="osqp",
31
+ solver_params={},
32
+ ):
33
+ if hesstype == "diag":
34
+ P = tau * sp.sparse.identity(n, format="csc")
35
+ kap = kap_val * np.ones(mc)
36
+ cval = np.array(cval)
37
+ return solve_qp(
38
+ P,
39
+ fgrad.reshape((n,)),
40
+ cgrad.reshape((mc, n)),
41
+ kap - cval,
42
+ np.zeros((0, n)),
43
+ np.zeros((0,)),
44
+ -beta * np.ones((n,)),
45
+ beta * np.ones((n,)),
46
+ qp_solver,
47
+ )
48
+
49
+ @staticmethod
50
+ def compute_kappa(cval, cgrad, lamb, rho, mc, n):
51
+ term1 = (1 - lamb) * np.maximum(cval, 0).max()
52
+ obj = np.zeros(n + 1)
53
+ obj[0] = 1.0
54
+ A_ub = np.hstack([-np.ones((mc, 1)), cgrad])
55
+ b_ub = -cval
56
+ bounds = [(0, None)] + [(-rho, rho) for _ in range(n)]
57
+
58
+ try:
59
+ res = linprog(c=obj, A_ub=A_ub, b_ub=b_ub, bounds=bounds, method="highs")
60
+ if res.success:
61
+ term2 = lamb * res.fun
62
+ else:
63
+ term2 = lamb * rho
64
+ except:
65
+ term2 = lamb * rho
66
+
67
+ return term1 + term2
68
+
69
+ def optimize(
70
+ self,
71
+ alpha,
72
+ stepsize_rule="inv_iter",
73
+ zeta=0.5,
74
+ gamma0=0.1,
75
+ rho=0.8,
76
+ lamb=0.5,
77
+ beta=10.0,
78
+ tau=1.0,
79
+ device="cpu",
80
+ seed=None,
81
+ verbose=True,
82
+ max_runtime=None,
83
+ max_iter=None,
84
+ save_state_interval=1
85
+ ):
86
+ self.state_history = {}
87
+ self.state_history["params"] = {"w": {}}
88
+ self.state_history["values"] = {"f": {}, "d": {}, "c": {}, "n_samples": {}}
89
+ self.state_history["time"] = {}
90
+
91
+ max_sample_size = np.max([c.group_sizes() for c in self.constraints])
92
+ n = sum(p.numel() for p in self.net.parameters())
93
+
94
+ rng = np.random.default_rng(seed=seed)
95
+ run_start = timeit.default_timer()
96
+
97
+ total_iters = 0
98
+ while True:
99
+ total_iters += 1
100
+ if max_iter is not None and total_iters >= max_iter:
101
+ break
102
+ current_time = timeit.default_timer()
103
+ if total_iters % save_state_interval == 0:
104
+ self.state_history["time"][total_iters] = current_time - run_start
105
+
106
+ if max_runtime > 0 and current_time - run_start >= max_runtime:
107
+ print(current_time - run_start)
108
+ # self.history["constr"] = pd.DataFrame(self.history["constr"])
109
+ return self.state_history
110
+
111
+ if stepsize_rule == "inv_iter":
112
+ gamma = gamma0 / (total_iters + 1) ** zeta
113
+ elif stepsize_rule == "dimin":
114
+ if total_iters == 1:
115
+ gamma = gamma0
116
+ else:
117
+ gamma *= 1 - zeta * gamma
118
+
119
+ Nsamp = rng.geometric(p=alpha) - 1
120
+ while (2 ** (Nsamp + 1)) > max_sample_size:
121
+ Nsamp = rng.geometric(p=alpha) - 1
122
+
123
+ n_samples_used = 3 * (
124
+ 1 + 2 ** (Nsamp + 1)
125
+ )
126
+
127
+ dsols = np.zeros((4, n))
128
+
129
+ ################
130
+ ### sampling ###
131
+ ################
132
+ indices_f = []
133
+ samples_c = []
134
+
135
+ subp_batch_size = 2 ** (Nsamp + 1)
136
+
137
+ indices_f.append(rng.choice(len(self.dataset), size=1))
138
+ samples_c.append([c.sample_dataset(1) for c in self.constraints])
139
+
140
+ idx_f = rng.choice(len(self.dataset), size=subp_batch_size)
141
+ indices_f.extend([idx_f[::2], idx_f[1::2], idx_f])
142
+ s_c = [c.sample_dataset(subp_batch_size) for c in self.constraints]
143
+ samples_c.extend(
144
+ [
145
+ [[(x[::2], y[::2]) for x, y in c_sample] for c_sample in s_c],
146
+ [[(x[1::2], y[1::2]) for x, y in c_sample] for c_sample in s_c],
147
+ s_c,
148
+ ]
149
+ )
150
+
151
+ ##############
152
+ ### update ###
153
+ ##############
154
+ for j, samples in enumerate(zip(indices_f, samples_c)):
155
+ self.net.zero_grad()
156
+
157
+ idx = samples[0]
158
+ obj_batch = self.dataset[idx]
159
+ c_batch = samples[1]
160
+
161
+ # calculate autograd jacobian of obj fun w.r.t. params
162
+ outs = self.net(obj_batch[0])
163
+ if obj_batch[1].ndim < outs.ndim:
164
+ outs = outs.squeeze(1)
165
+ feval = self.loss_fn(outs, obj_batch[1])
166
+
167
+ dfdw = torch.autograd.grad(feval, self.net.parameters())
168
+ dfdw = torch.concat([dfdwi.flatten() for dfdwi in dfdw])
169
+
170
+ # calculate autograd jacobian of self.constraints fun w.r.t. params
171
+
172
+ constraint_eval = []
173
+ dcdw = []
174
+ for i, c in enumerate(self.constraints):
175
+ self.net.zero_grad()
176
+ # print(j, i)
177
+ c_val = c.eval(self.net, c_batch[i])
178
+
179
+ c_grad = torch.autograd.grad(c_val, self.net.parameters())
180
+ c_grad = (
181
+ torch.concat([cg.flatten() for cg in c_grad]).detach().numpy()
182
+ )
183
+
184
+ constraint_eval.append(c_val.detach())
185
+ dcdw.append(c_grad)
186
+
187
+ constraint_eval = np.array(constraint_eval).flatten()
188
+ dcdw = np.array(dcdw)
189
+
190
+ kappa = self.compute_kappa(
191
+ constraint_eval,
192
+ dcdw,
193
+ rho,
194
+ lamb,
195
+ mc=len(self.constraints),
196
+ n=len(dfdw),
197
+ )
198
+
199
+ # solve subproblem
200
+ feval = feval.detach().numpy()
201
+ dfdw = dfdw.detach().numpy()
202
+ dsol = self.solvesubp(
203
+ dfdw,
204
+ constraint_eval,
205
+ dcdw,
206
+ kappa,
207
+ beta,
208
+ tau,
209
+ hesstype="diag",
210
+ mc=len(self.constraints),
211
+ n=len(dfdw),
212
+ qp_solver="osqp",
213
+ )
214
+
215
+ dsols[j, :] = dsol
216
+
217
+ # aggregate solutions to the subproblem according to Eq. 23
218
+ dsol = dsols[0, :] + (
219
+ dsols[3, :] - 0.5 * dsols[1, :] - 0.5 * dsols[2, :]
220
+ ) / (alpha * ((1 - alpha) ** (Nsamp)))
221
+
222
+ start = 0
223
+ print(f"{total_iters}", end="\r")
224
+ with torch.no_grad():
225
+ w = net_params_to_tensor(self.net)
226
+ if any([torch.any(torch.isnan(lw)) for lw in w]):
227
+ print("NaNs!")
228
+ return self.state_history
229
+ for i in range(len(w)):
230
+ end = start + w[i].numel()
231
+ w[i].add_(
232
+ torch.tensor(
233
+ gamma * np.reshape(dsol[start:end], np.shape(w[i]))
234
+ )
235
+ )
236
+ start = end
237
+
238
+ if total_iters % save_state_interval == 0:
239
+ self.state_history["params"]["w"][total_iters] = deepcopy(
240
+ self.net.state_dict()
241
+ )
242
+ self.state_history["values"]["n_samples"][total_iters] = n_samples_used
243
+
244
+ self.state_history["values"]["d"][total_iters] = dsol
245
+ # self.history["w"].append(deepcopy(self.net.state_dict()))
246
+
247
+ feval = self.loss_fn(outs, obj_batch[1])
248
+
249
+ # self.history["constr"] = pd.DataFrame(self.history["constr"])
250
+ return self.state_history
@@ -0,0 +1,107 @@
1
+ import timeit
2
+ from copy import deepcopy
3
+
4
+ import numpy as np
5
+ import torch
6
+
7
+ from humancompatible.train.algorithms.Algorithm import Algorithm
8
+
9
+
10
+ class SGD(Algorithm):
11
+ def __init__(self, net, data, loss, constraints):
12
+ super().__init__(net, data, loss, constraints)
13
+
14
+ def optimize(
15
+ self,
16
+ lr,
17
+ batch_size,
18
+ epochs=None,
19
+ max_runtime=None,
20
+ max_iter=None,
21
+ seed=None,
22
+ device="cpu",
23
+ verbose=True,
24
+ save_state_interval=1000,
25
+ ):
26
+ self.state_history = {}
27
+ self.state_history["params"] = {"w": {}}
28
+ self.state_history["values"] = {"f": {}, "fg": {}}
29
+ self.state_history["time"] = {}
30
+
31
+ run_start = timeit.default_timer()
32
+
33
+ if epochs is None:
34
+ epochs = np.inf
35
+ if max_iter is None:
36
+ max_iter = np.inf
37
+ if max_runtime is None:
38
+ max_runtime = np.inf
39
+
40
+ gen = torch.Generator(device=device)
41
+ if seed is not None:
42
+ gen.manual_seed(seed)
43
+ loss_loader = torch.utils.data.DataLoader(
44
+ self.dataset, batch_size, shuffle=True, generator=gen
45
+ )
46
+ loss_iter = iter(loss_loader)
47
+
48
+ epoch = 0
49
+ iteration = 0
50
+ total_iters = 0
51
+
52
+ optimizer = torch.optim.SGD(self.net.parameters(), lr=lr)
53
+
54
+ while True:
55
+ elapsed = timeit.default_timer() - run_start
56
+ iteration += 1
57
+ total_iters += 1
58
+ if epoch >= epochs or total_iters >= max_iter or elapsed > max_runtime:
59
+ break
60
+
61
+ if total_iters % save_state_interval == 0:
62
+ self.state_history["params"]["w"][total_iters] = deepcopy(
63
+ self.net.state_dict()
64
+ )
65
+ self.state_history["time"][total_iters] = elapsed
66
+
67
+ try:
68
+ (f_inputs, f_labels) = next(loss_iter)
69
+ except StopIteration:
70
+ epoch += 1
71
+ iteration = 0
72
+ gen = gen
73
+ loss_loader = torch.utils.data.DataLoader(
74
+ self.dataset, batch_size, shuffle=True, generator=gen
75
+ )
76
+ loss_iter = iter(loss_loader)
77
+ (f_inputs, f_labels) = next(loss_iter)
78
+ # lr *= 0.8
79
+
80
+ ########################
81
+ ## UPDATE MULTIPLIERS ##
82
+ ########################
83
+ self.net.zero_grad()
84
+ outputs = self.net(f_inputs)
85
+ loss = self.loss_fn(outputs.squeeze(), f_labels)
86
+ loss.backward()
87
+ optimizer.step()
88
+
89
+ # f_grad_estimate =
90
+
91
+ with torch.no_grad():
92
+ if total_iters % save_state_interval == 0:
93
+ self.state_history["values"]["f"][total_iters] = (
94
+ loss.detach().cpu().numpy()
95
+ )
96
+ # self.state_history['values']['fg'][total_iters] = torch.norm(f_grad_estimate).detach().cpu().numpy()
97
+
98
+ if verbose:
99
+ with np.printoptions(
100
+ precision=8, suppress=True, floatmode="fixed", sign=" "
101
+ ):
102
+ print(
103
+ f"""{epoch:2}|{iteration:5} | {lr} | {loss.detach().cpu().numpy():1.5f}""",
104
+ end="\r",
105
+ )
106
+
107
+ return self.state_history