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.
- experiments/__init__.py +0 -0
- experiments/calculate_iteration_values.py +310 -0
- experiments/run_folktables.py +408 -0
- experiments/run_folktables_torchalgs.py +956 -0
- humancompatible/__init__.py +0 -0
- humancompatible/train/__init__.py +3 -0
- humancompatible/train/algorithms/Algorithm.py +25 -0
- humancompatible/train/algorithms/__init__.py +8 -0
- humancompatible/train/algorithms/ghost.py +250 -0
- humancompatible/train/algorithms/sgd.py +107 -0
- humancompatible/train/algorithms/ssl_alm.py +311 -0
- humancompatible/train/algorithms/switching_subgradient.py +192 -0
- humancompatible/train/algorithms/torch/__init__.py +4 -0
- humancompatible/train/algorithms/torch/ssl_alm.py +212 -0
- humancompatible/train/algorithms/torch/ssw.py +155 -0
- humancompatible/train/algorithms/utils.py +61 -0
- humancompatible/train/constraints/__init__.py +11 -0
- humancompatible/train/constraints/constraint.py +87 -0
- humancompatible/train/constraints/constraint_fns.py +118 -0
- humancompatible/train/fairness/__init__.py +0 -0
- humancompatible/train/fairness/constraints/__init__.py +15 -0
- humancompatible/train/fairness/constraints/constraint.py +97 -0
- humancompatible/train/fairness/constraints/constraint_fns.py +244 -0
- humancompatible/train/fairness/constraints/torch/__init__.py +1 -0
- humancompatible/train/fairness/constraints/torch/constraints.py +36 -0
- humancompatible/train/fairness/utils/__init__.py +1 -0
- humancompatible/train/fairness/utils/balanced_batch_sampler.py +65 -0
- humancompatible_train-0.1.0.dist-info/METADATA +188 -0
- humancompatible_train-0.1.0.dist-info/RECORD +32 -0
- humancompatible_train-0.1.0.dist-info/WHEEL +5 -0
- humancompatible_train-0.1.0.dist-info/licenses/LICENCE.txt +201 -0
- humancompatible_train-0.1.0.dist-info/top_level.txt +2 -0
|
File without changes
|
|
@@ -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,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
|