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
|
@@ -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
|