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,956 @@
1
+ from copy import deepcopy
2
+ import importlib
3
+ from itertools import combinations
4
+ import os
5
+ import timeit
6
+ import warnings
7
+ import hydra
8
+ import numpy as np
9
+ import pandas as pd
10
+ import torch
11
+ from omegaconf import DictConfig, OmegaConf
12
+ from torch import nn, tensor
13
+ from torch.utils.data import TensorDataset, DataLoader, SubsetRandomSampler
14
+ from humancompatible.train.fairness.constraints.constraint_fns import fairret_stat_equality
15
+ from utils.load_folktables import prepare_folktables_multattr
16
+ from utils.network import SimpleNet
17
+ from humancompatible.train.algorithms.utils import net_grads_to_tensor, net_params_to_tensor
18
+ from itertools import combinations
19
+ from humancompatible.train.fairness.constraints import FairnessConstraint
20
+
21
+
22
+
23
+ @hydra.main(version_base=None, config_path="conf", config_name="experiment")
24
+ def run(cfg: DictConfig) -> None:
25
+ warnings.filterwarnings("ignore", category=FutureWarning)
26
+
27
+ print(OmegaConf.to_yaml(cfg))
28
+ N_RUNS = cfg.n_runs
29
+ FT_STATE = cfg.data.state
30
+ FT_TASK = cfg.data.task
31
+ DOWNLOAD_DATA = cfg.data.download
32
+ DATA_PATH = cfg.data.path
33
+
34
+ if "constraint" in cfg.keys():
35
+ CONSTRAINT = cfg.constraint.import_name
36
+ LOSS_BOUND = cfg.constraint.bound
37
+ else:
38
+ CONSTRAINT = "unconstr"
39
+ LOSS_BOUND = 0
40
+
41
+ if cfg.device == "cpu":
42
+ device = "cpu"
43
+ elif cfg.alg == "ghost":
44
+ device = "cpu"
45
+ print("CUDA not supported for Stochastic Ghost")
46
+ elif torch.cuda.is_available():
47
+ device = "cuda"
48
+ print("CUDA found")
49
+ else:
50
+ device = "cpu"
51
+ print("CUDA not found")
52
+
53
+ print(f"{device = }")
54
+ torch.set_default_device(device)
55
+
56
+ DTYPE = torch.float32
57
+
58
+ ## load data ##
59
+
60
+ torch.set_default_dtype(DTYPE)
61
+ DATASET_NAME = FT_TASK + "_" + FT_STATE
62
+
63
+ (
64
+ X_train,
65
+ y_train,
66
+ group_ind_train,
67
+ group_onehot_train,
68
+ sep_group_ind_train,
69
+ X_test,
70
+ y_test,
71
+ group_ind_test,
72
+ sep_group_ind_test,
73
+ group_onehot_test,
74
+ _
75
+ ) = prepare_folktables_multattr(
76
+ FT_TASK,
77
+ state=FT_STATE.upper(),
78
+ random_state=42,
79
+ onehot=False,
80
+ download=DOWNLOAD_DATA,
81
+ path=DATA_PATH,
82
+ sens_cols=cfg.data.sens_attr,
83
+ binarize=cfg.data.binarize,
84
+ stratify=False,
85
+ )
86
+ print('Groups:')
87
+ print(len(group_ind_train))
88
+ X_train_tensor = tensor(X_train, dtype=DTYPE)
89
+ y_train_tensor = tensor(y_train, dtype=DTYPE)
90
+ train_ds = TensorDataset(X_train_tensor, y_train_tensor)
91
+
92
+ print(f"Train data loaded: {(FT_TASK, FT_STATE)}")
93
+ print(f"Data shape: {X_train_tensor.shape}")
94
+
95
+ ## prepare to save results ##
96
+
97
+ if "save_name" in cfg["alg"].keys():
98
+ alg_save_name = cfg.alg.save_name
99
+ else:
100
+ alg_save_name = cfg.alg.import_name
101
+
102
+ saved_models_path = os.path.abspath(
103
+ os.path.join(os.path.dirname(__file__), "utils", "saved_models")
104
+ )
105
+ directory = os.path.join(
106
+ saved_models_path, DATASET_NAME, CONSTRAINT, f"{LOSS_BOUND:.0E}"
107
+ )
108
+
109
+ model_name = os.path.join(directory, f"{alg_save_name}_{LOSS_BOUND}")
110
+
111
+ if not os.path.exists(directory):
112
+ os.makedirs(directory)
113
+
114
+ ## run experiments ##
115
+
116
+ histories = []
117
+
118
+ # experiment loop
119
+ for EXP_IDX in range(N_RUNS):
120
+
121
+ net = SimpleNet(in_shape=X_test.shape[1], out_shape=1, dtype=DTYPE).to(device)
122
+
123
+ ## define constraints ##
124
+
125
+ criterion = nn.BCEWithLogitsLoss()
126
+ constraint_fn_module = importlib.import_module("humancompatible.train.fairness.constraints")
127
+ try:
128
+ constraint_fn = getattr(constraint_fn_module, cfg.constraint.import_name)
129
+ except:
130
+ constraint_fn = getattr(importlib.import_module("humancompatible.train.fairness.constraints.torch"), "loss_equality")
131
+
132
+ if cfg.constraint.import_name == 'abs_max_dev_from_overall_tpr':
133
+ c = [FairnessConstraint(
134
+ train_ds,
135
+ group_ind_train,
136
+ fn=lambda net, inputs: constraint_fn(criterion, net, inputs) - cfg.constraint.bound,
137
+ batch_size=cfg.constraint.c_batch_size,
138
+ seed=EXP_IDX
139
+ )]
140
+ elif cfg.constraint.import_name in ['abs_diff_tpr', 'abs_diff_pr']:
141
+ c = [
142
+ FairnessConstraint(
143
+ train_ds,
144
+ [group_ind, np.concat(group_ind_train)],
145
+ fn=lambda net, inputs: constraint_fn(criterion, net, inputs) - cfg.constraint.bound,
146
+ batch_size=cfg.constraint.c_batch_size,
147
+ seed=EXP_IDX
148
+ )
149
+ for group_ind in group_ind_train
150
+ ]
151
+ elif cfg.constraint.import_name in ['abs_diff_fpr', 'abs_diff_pr']:
152
+ c = [
153
+ FairnessConstraint(
154
+ train_ds,
155
+ [group_ind, np.concat(group_ind_train)],
156
+ fn=lambda net, inputs: constraint_fn(criterion, net, inputs) - cfg.constraint.bound,
157
+ batch_size=cfg.constraint.c_batch_size,
158
+ seed=EXP_IDX
159
+ )
160
+ for group_ind in group_ind_train
161
+ ]
162
+ constraint_fn1 = getattr(constraint_fn_module, 'abs_diff_tpr')
163
+ c += [
164
+ FairnessConstraint(
165
+ train_ds,
166
+ [group_ind, np.concat(group_ind_train)],
167
+ fn=lambda net, inputs: constraint_fn1(criterion, net, inputs) - cfg.constraint.bound,
168
+ batch_size=cfg.constraint.c_batch_size,
169
+ seed=EXP_IDX
170
+ )
171
+ for group_ind in group_ind_train
172
+ ]
173
+ else:
174
+ c = construct_constraints(
175
+ bound=cfg.constraint.bound,
176
+ add_negative=cfg.constraint.add_negative,
177
+ batch_size=cfg.constraint.c_batch_size,
178
+ device=device,
179
+ constraint_groups=[group_ind_train],
180
+ dataset=train_ds,
181
+ seed=EXP_IDX,
182
+ constraint_fn=lambda net, inputs: constraint_fn(criterion, net, inputs),
183
+ max_0 = False
184
+ )
185
+ # breakpoint()
186
+
187
+ torch.manual_seed(EXP_IDX)
188
+ model_path = model_name + f"_trial{EXP_IDX}.pt"
189
+
190
+ net = SimpleNet(in_shape=X_test.shape[1], out_shape=1, dtype=DTYPE).to(device)
191
+
192
+ history = {
193
+ "params": {"w": {}, "slack": {}, "dual_ms": {}},
194
+ "values": {"f": {}, "c": {}, "G": {}},
195
+ "time": {},
196
+ }
197
+
198
+ ## train ##
199
+
200
+ if cfg.alg.import_name.startswith("fairret"):
201
+ m = len(group_ind_train)
202
+
203
+ _fairret_loss_module = importlib.import_module("fairret.loss")
204
+ _fairret_statistic_module = importlib.import_module("fairret.statistic")
205
+ fstat = getattr(_fairret_statistic_module, cfg.alg.params.statistic)()
206
+ floss = getattr(_fairret_loss_module, cfg.alg.params.loss)(fstat, p=1)
207
+
208
+ run_start = timeit.default_timer()
209
+
210
+ criterion = torch.nn.BCEWithLogitsLoss()
211
+ optimizer = torch.optim.SGD(net.parameters(), lr=cfg.alg.params.lr)
212
+ c_batch_size = cfg.alg.params.c_batch_size
213
+ obj_batch_size = cfg.alg.params.obj_batch_size
214
+ mult = cfg.alg.params.pmult
215
+
216
+ total_iters = 0
217
+ gen = torch.Generator(device=device)
218
+ gen.manual_seed(EXP_IDX)
219
+ obj_loader = iter(
220
+ torch.utils.data.DataLoader(
221
+ train_ds,
222
+ obj_batch_size,
223
+ shuffle=True,
224
+ generator=gen,
225
+ drop_last=True,
226
+ )
227
+ )
228
+
229
+ constr_dataloaders = []
230
+ for group_indices in group_ind_train:
231
+ sampler = SubsetRandomSampler(group_indices, gen)
232
+ constr_dataloaders.append(
233
+ iter(
234
+ DataLoader(
235
+ train_ds, c_batch_size, sampler=sampler, drop_last=True
236
+ )
237
+ )
238
+ )
239
+
240
+ epoch = 0
241
+ iteration = 0
242
+ total_iters = 0
243
+
244
+ group_ind_onehot = torch.empty(m, (c_batch_size * m))
245
+ for j in range(1, m):
246
+ group_ind_onehot[j][c_batch_size * (j - 1) : c_batch_size * j] = (
247
+ torch.ones(c_batch_size)
248
+ )
249
+ group_ind_onehot = group_ind_onehot.T
250
+
251
+ while True:
252
+ elapsed = timeit.default_timer() - run_start
253
+ iteration += 1
254
+ total_iters += 1
255
+ if (
256
+ cfg.run_maxiter is not None and total_iters >= cfg.run_maxiter
257
+ ) or elapsed > cfg.run_maxtime:
258
+ break
259
+ if total_iters % cfg.alg.params.save_state_interval == 0:
260
+ history["params"]["w"][total_iters] = deepcopy(net.state_dict())
261
+ history["time"][total_iters] = elapsed
262
+
263
+ net.zero_grad()
264
+
265
+ inputs, labels = sample_or_restart_iterloader(obj_loader)
266
+ outputs = net(inputs)
267
+ loss_obj = criterion(outputs.squeeze(), labels)
268
+
269
+ c_inputs, c_labels = [], []
270
+ for j in range(m):
271
+ group_inputs, group_labels = sample_or_restart_iterloader(
272
+ constr_dataloaders[j]
273
+ )
274
+ c_inputs.append(group_inputs)
275
+ c_labels.append(group_labels)
276
+
277
+ c_inputs = torch.concat(c_inputs)
278
+ c_labels = torch.concat(c_labels)
279
+
280
+ outputs_c = net(c_inputs).squeeze()
281
+ loss_c = floss(
282
+ outputs_c.unsqueeze(1), group_ind_onehot, c_labels.unsqueeze(1)
283
+ )
284
+
285
+ loss = loss_obj + mult * loss_c
286
+
287
+ loss.backward()
288
+ optimizer.step()
289
+
290
+ with np.printoptions(precision=6, suppress=True):
291
+ print(
292
+ f"{epoch:2} | {iteration:5} | {loss_obj.detach().cpu().numpy():.4} | {loss_c.detach().cpu().numpy():.4}",
293
+ end="\r",
294
+ )
295
+ elif cfg.alg.import_name.startswith("TorchSSLALM"):
296
+ epochs = 1000
297
+ avg_epoch_c_log = []
298
+ avg_epoch_loss_log = []
299
+ m = len(list(combinations(group_ind_train, 2)))*2
300
+
301
+ from fairret.statistic import TruePositiveRate, PositiveRate
302
+ from fairret.loss import NormLoss
303
+
304
+ # breakpoint()
305
+ train_ds_sens = TensorDataset(
306
+ X_train_tensor,
307
+ group_onehot_train,
308
+ y_train_tensor
309
+ )
310
+
311
+
312
+ slack_vars = torch.zeros(m, requires_grad=True)
313
+ obj_batch_size = 16
314
+ c_batch_size = cfg.constraint.c_batch_size
315
+
316
+ from humancompatible.train.algorithms.torch import SSLALM
317
+ optimizer = SSLALM(
318
+ net.parameters(),
319
+ lr=0.01,
320
+ dual_lr=0.1,
321
+ rho=1.0,
322
+ mu=2.0,
323
+ beta=0.5,
324
+ m=m,
325
+ )
326
+ c_bound = torch.tensor([cfg.constraint.bound]*m)
327
+ optimizer.add_param_group(param_group={"params": slack_vars, "name": "slack"})
328
+ # constr = FalseNegativeFalsePositiveFraction()
329
+ # constr = PositiveRate()
330
+ constr = constraint_fn
331
+ # fair_loss = NormLoss(constr)
332
+
333
+ time = timeit.default_timer()
334
+ total_iters = 0
335
+ c_criterion = torch.nn.BCEWithLogitsLoss(reduction='none')
336
+
337
+ for epoch in range(epochs):
338
+ elapsed = timeit.default_timer()
339
+ if elapsed - time > cfg.run_maxtime:
340
+ break
341
+ loss_log = []
342
+ c_log = []
343
+ gen = torch.Generator(device=device)
344
+ gen.manual_seed(EXP_IDX + epoch)
345
+ from humancompatible.train.fairness.utils import BalancedBatchSampler
346
+
347
+ sampler = BalancedBatchSampler(
348
+ # subgroup_indices=group_ind_train,
349
+ subgroup_onehot=group_onehot_train,
350
+ batch_size=c_batch_size,
351
+ drop_last=True
352
+ )
353
+ dataloader = DataLoader(
354
+ train_ds_sens,
355
+ batch_sampler=sampler
356
+ )
357
+ # c_loader = iter(dataloader)
358
+ for batch_input, batch_sens, batch_label in dataloader:
359
+ elapsed = timeit.default_timer()
360
+ if elapsed - time > cfg.run_maxtime:
361
+ break
362
+ history["time"][total_iters] = elapsed - time
363
+
364
+ # evaluate constraints and constraint grads
365
+ c_vals = []
366
+ c_vals_raw = []
367
+
368
+ c_inputs = batch_input[::2]
369
+ c_labels = batch_label[::2]
370
+ c_sens = batch_sens[::2]
371
+ c_sens_norm = c_sens.div(torch.sum(c_sens, axis=0))
372
+
373
+ # calculate loss for each group
374
+ c_loss = c_criterion(net(c_inputs).squeeze(), c_labels) @ c_sens_norm
375
+ c_val_raw_vec = []
376
+ for (l1, l2) in combinations(c_loss, 2):
377
+ c_val_raw_vec.append(l1-l2)
378
+ c_val_raw_vec.append(l2-l1)
379
+
380
+
381
+ # c_outputs = torch.nn.functional.sigmoid(net(c_inputs))
382
+ # c_outputs_pos_idx = (c_outputs >= 0).squeeze()
383
+ # c_overall = constr(c_outputs[c_outputs_pos_idx], None
384
+ # , c_labels[c_outputs_pos_idx].unsqueeze(1)
385
+ # )
386
+ # c_val_raw_vec = constr(c_outputs[c_outputs_pos_idx], c_sens[c_outputs_pos_idx]
387
+ # , c_labels[c_outputs_pos_idx].unsqueeze(1)
388
+ # )
389
+ # c_val_raw_vec = torch.abs(c_val_raw_vec - c_overall)
390
+
391
+ for i in range(m):
392
+ optimizer.zero_grad()
393
+ c_val = c_val_raw_vec[i] + slack_vars[i] - c_bound[i]
394
+ # retain_graph in all but last iteration to calculate grads
395
+ c_val.backward(retain_graph = i < m-1)
396
+ optimizer.dual_step(i, c_val)
397
+
398
+ c_vals.append(c_val.detach())
399
+ c_vals_raw.append(c_val_raw_vec[i].detach())
400
+
401
+
402
+ optimizer.zero_grad()
403
+ # evaluate loss and loss grad
404
+ logits = net(batch_input)
405
+ loss = criterion(logits.squeeze(), batch_label) + torch.zeros_like(slack_vars) @ slack_vars # SLACK
406
+ loss.backward()
407
+
408
+ if cfg.alg.params.use_unbiased_penalty_grad:
409
+ with torch.no_grad():
410
+ c_inputs = batch_input[1::2]
411
+ c_labels = batch_label[1::2]
412
+ c_sens = batch_sens[1::2]
413
+ c_sens_norm = c_sens.div(torch.sum(c_sens, axis=0))
414
+ c_loss = c_criterion(net(c_inputs).squeeze(), c_labels) @ c_sens_norm
415
+ c_val_raw_vec = []
416
+ for (l1, l2) in combinations(c_loss, 2):
417
+ c_val_raw_vec.append(l1-l2)
418
+ c_val_raw_vec.append(l2-l1)
419
+
420
+ # c_vals = []
421
+ # c_vals_raw = []
422
+ # c_inputs = batch_input[1::2]
423
+ # c_labels = batch_label[1::2]
424
+ # c_sens = batch_sens[1::2]
425
+ # c_outputs = torch.nn.functional.sigmoid(net(c_inputs))
426
+ # c_outputs_pos_idx = (c_outputs >= 0).squeeze()
427
+ # c_overall = constr(c_outputs[c_outputs_pos_idx], None
428
+ # ,c_labels[c_outputs_pos_idx].unsqueeze(1)
429
+ # )
430
+ # c_val_raw_vec = constr(c_outputs[c_outputs_pos_idx], c_sens[c_outputs_pos_idx]
431
+ # , c_labels[c_outputs_pos_idx].unsqueeze(1)
432
+ # )
433
+ # c_val_raw_vec = torch.abs(c_val_raw_vec - c_overall)
434
+ # c_val = c_val_raw_vec + slack_vars - c_bound
435
+ # c_vals.append(c_val.detach())
436
+
437
+ # c_vals_raw = c_val_raw_vec.detach()
438
+
439
+ optimizer.step(c_vals)
440
+ optimizer.zero_grad()
441
+ if total_iters % cfg.alg.params.save_state_interval == 0:
442
+ history["params"]["w"][total_iters] = deepcopy(net.state_dict())
443
+ history["time"][total_iters] = elapsed - time
444
+
445
+ total_iters += 1
446
+
447
+ with torch.no_grad():
448
+ for s in slack_vars:
449
+ if s < 0:
450
+ s.zero_()
451
+
452
+ loss_log.append(loss.item())
453
+ c_log.append([c.item() for c in c_vals_raw])
454
+
455
+ # print(optimizer._dual_vars)
456
+ avg_epoch_loss_log.append(np.mean(loss_log))
457
+ avg_epoch_c_log.append(np.mean(c_log, axis=0))
458
+ with np.printoptions(precision=4):
459
+ print(
460
+ f"Epoch: {epoch}, loss: {avg_epoch_loss_log[-1]}, constraints: {avg_epoch_c_log[-1]}, dual: {optimizer._dual_vars.detach().numpy()}"
461
+ )
462
+ elif cfg.alg.import_name.startswith("TorchSSG"):
463
+ epochs = 100000
464
+ avg_epoch_c_log = []
465
+ avg_epoch_loss_log = []
466
+ m = 1
467
+
468
+ from fairret.statistic import PositiveRate
469
+
470
+ # breakpoint()
471
+ train_ds_sens = TensorDataset(X_train_tensor, group_onehot_train, y_train_tensor)
472
+
473
+ obj_batch_size = 16
474
+ c_batch_size = cfg.constraint.c_batch_size
475
+
476
+ from humancompatible.train.algorithms.torch import SSG
477
+ optimizer = SSG(
478
+ net.parameters(),
479
+ lr=0.05,
480
+ dual_lr=0.05,
481
+ m=m,
482
+ )
483
+ c_bound = torch.tensor([cfg.constraint.bound]*5)
484
+ c_tol = torch.tensor([cfg.constraint.bound*2]*5)
485
+
486
+ time = timeit.default_timer()
487
+ total_iters = 0
488
+
489
+ for epoch in range(epochs):
490
+ elapsed = timeit.default_timer()
491
+ if elapsed - time > cfg.run_maxtime:
492
+ break
493
+ loss_log = []
494
+ c_log = []
495
+ gen = torch.Generator(device=device)
496
+ gen.manual_seed(EXP_IDX + epoch)
497
+ obj_loader = iter(
498
+ torch.utils.data.DataLoader(
499
+ train_ds,
500
+ obj_batch_size,
501
+ shuffle=True,
502
+ generator=gen,
503
+ )
504
+ )
505
+
506
+ gen = torch.Generator(device=device)
507
+ gen.manual_seed(EXP_IDX + epoch)
508
+ from humancompatible.train.fairness.utils import BalancedBatchSampler
509
+
510
+ sampler = BalancedBatchSampler(
511
+ subgroup_indices=group_ind_train,
512
+ batch_size=c_batch_size,
513
+ drop_last=True)
514
+ dataloader = DataLoader(train_ds_sens, batch_sampler=sampler)
515
+ max_idx_log = []
516
+ for batch_input, batch_sens, batch_label in dataloader:
517
+ elapsed = timeit.default_timer()
518
+ if elapsed - time > cfg.run_maxtime:
519
+ break
520
+ history["time"][total_iters] = elapsed - time
521
+
522
+ # evaluate constraints and largest constraint grad
523
+ c_vals = []
524
+ c_vals_raw = []
525
+
526
+ c_inputs = batch_input
527
+ c_sens = batch_sens
528
+ c_labels = batch_label
529
+ c_outputs = torch.nn.functional.sigmoid(net(c_inputs))
530
+
531
+ constr = PositiveRate()
532
+ c_overall = constr(c_outputs, None)
533
+ c_val_raw_vec = constr(c_outputs, c_sens)
534
+ # breakpoint()
535
+ c_val_raw_vec = torch.abs(c_val_raw_vec - c_overall)
536
+
537
+ c_val = c_val_raw_vec - c_bound
538
+ c_max_viol_idx = torch.argmax(c_val - c_tol)
539
+ c_max_viol = c_val[c_max_viol_idx]
540
+ c_max_viol.backward()
541
+ max_idx_log.append(c_max_viol_idx)
542
+
543
+ optimizer.dual_step(i=0)
544
+
545
+ c_vals = c_max_viol
546
+ c_vals_raw.append(c_val_raw_vec.detach())
547
+
548
+ optimizer.zero_grad()
549
+ # evaluate loss and loss grad
550
+ logits = net(batch_input)
551
+ loss = criterion(logits.squeeze(), batch_label)
552
+ loss.backward()
553
+
554
+ optimizer.step(c_vals)
555
+ optimizer.zero_grad()
556
+
557
+ total_iters += 1
558
+ c_tol /= np.sqrt(total_iters)
559
+
560
+ if total_iters % cfg.alg.params.save_state_interval == 0:
561
+ history["params"]["w"][total_iters] = deepcopy(net.state_dict())
562
+
563
+
564
+ loss_log.append(loss.item())
565
+ c_log.append([
566
+ c_val.detach()
567
+ ])
568
+
569
+ # print(optimizer._dual_vars)
570
+ avg_epoch_loss_log.append(np.mean(loss_log))
571
+ avg_epoch_c_log.append(np.mean(c_log, axis=0))
572
+ print(
573
+ f"Epoch: {epoch}, loss: {avg_epoch_loss_log[-1]}, constraints: {avg_epoch_c_log[-1]}"
574
+ )
575
+ elif cfg.alg.import_name.startswith("SGD"):
576
+
577
+ optimizer = torch.optim.Adam(params=net.parameters())
578
+
579
+ time = timeit.default_timer()
580
+ total_iters = 0
581
+ train_ds_sens = TensorDataset(X_train_tensor, group_onehot_train, y_train_tensor)
582
+ epochs = 1000000
583
+ avg_epoch_loss_log = []
584
+
585
+ for epoch in range(epochs):
586
+ elapsed = timeit.default_timer()
587
+ if elapsed - time > cfg.run_maxtime:
588
+ break
589
+ loss_log = []
590
+ gen = torch.Generator(device=device)
591
+ gen.manual_seed(EXP_IDX + epoch)
592
+
593
+ gen = torch.Generator(device=device)
594
+ gen.manual_seed(EXP_IDX + epoch)
595
+ from humancompatible.train.fairness.utils import BalancedBatchSampler
596
+
597
+ sampler = BalancedBatchSampler(
598
+ subgroup_indices=group_ind_train,
599
+ batch_size=cfg.constraint.c_batch_size,
600
+ drop_last=True)
601
+ dataloader = DataLoader(train_ds_sens, batch_sampler=sampler)
602
+
603
+ for batch_input, batch_sens, batch_label in dataloader:
604
+ elapsed = timeit.default_timer()
605
+ if elapsed - time > cfg.run_maxtime:
606
+ break
607
+ history["time"][total_iters] = elapsed - time
608
+ logits = net(batch_input)
609
+ loss = criterion(logits.squeeze(), batch_label)
610
+ loss.backward()
611
+
612
+ optimizer.step()
613
+ optimizer.zero_grad()
614
+
615
+ if total_iters % cfg.alg.params.save_state_interval == 0:
616
+ history["params"]["w"][total_iters] = deepcopy(net.state_dict())
617
+
618
+ total_iters += 1
619
+
620
+ loss_log.append(loss.item())
621
+
622
+ avg_epoch_loss_log.append(np.mean(loss_log))
623
+
624
+ print(
625
+ f"Epoch: {epoch}, loss: {avg_epoch_loss_log[-1]}"
626
+ )
627
+ else:
628
+ optimizer_name = cfg.alg.import_name
629
+ module = importlib.import_module("humancompatible.train.algorithms")
630
+ Optimizer = getattr(module, optimizer_name)
631
+
632
+ optimizer = Optimizer(net, train_ds, criterion, c)
633
+ history = optimizer.optimize(
634
+ **cfg.alg.params,
635
+ max_iter=cfg.run_maxiter,
636
+ max_runtime=cfg.run_maxtime,
637
+ device=device,
638
+ seed=EXP_IDX,
639
+ verbose=True,
640
+ )
641
+
642
+ ## SAVE RESULTS ##
643
+ params = pd.DataFrame(history["params"])
644
+ values = pd.DataFrame(history["values"])
645
+ t = pd.Series(history["time"], name="time")
646
+ histories.append(values.join(params, how="outer").join(t, how="outer"))
647
+
648
+ ## SAVE MODEL ##
649
+ print(f"Model saved to: {model_path}")
650
+ torch.save(net.state_dict(), model_path)
651
+ print("")
652
+
653
+ # Save DataFrames to CSV files
654
+ if cfg.alg.import_name.lower() == "sgd":
655
+ c_name = "unconstrained"
656
+ else:
657
+ c_name = cfg.constraint.import_name
658
+ utils_path = os.path.abspath(
659
+ os.path.join(os.path.dirname(__file__), "utils", "exp_results", c_name)
660
+ )
661
+ if not os.path.exists(utils_path):
662
+ os.makedirs(utils_path)
663
+ fname = f"{alg_save_name}_{DATASET_NAME}_{LOSS_BOUND}.csv"
664
+ save_path = os.path.join(utils_path, fname)
665
+ print(f"Saving to: {save_path}")
666
+ histories = pd.concat(histories, keys=range(N_RUNS), names=["trial", "iteration"])
667
+ histories.to_pickle(save_path)
668
+ print("Saved!")
669
+
670
+ ####################################################
671
+ ### CALCULATE STATS ON EVERY ALGORITHM ITERATION ###
672
+ ####################################################
673
+
674
+ criterion = nn.BCEWithLogitsLoss()
675
+ # constraint_fn_module = importlib.import_module("humancompatible.train.fairness.constraints")
676
+ # constraint_fn = getattr(constraint_fn_module, cfg.constraint.import_name)
677
+ # if cfg.constraint.import_name != 'abs_max_dev_from_overall_tpr':
678
+ # c = construct_constraints(
679
+ # bound=cfg.constraint.bound,
680
+ # add_negative=cfg.constraint.add_negative,
681
+ # batch_size=cfg.constraint.c_batch_size,
682
+ # device=device,
683
+ # constraint_groups=[group_ind_train],
684
+ # dataset=train_ds,
685
+ # seed=EXP_IDX,
686
+ # constraint_fn=lambda net, inputs: constraint_fn(loss_fn, net, inputs),
687
+ # max_0 = False
688
+ # )
689
+ # else:
690
+ # c = [FairnessConstraint(
691
+ # train_ds,
692
+ # group_ind_train,
693
+ # fn=lambda net, inputs: constraint_fn(loss_fn, net, inputs) - cfg.constraint.bound,
694
+ # batch_size=cfg.constraint.c_batch_size // len(group_ind_train),
695
+ # seed=EXP_IDX
696
+ # )]
697
+
698
+ print("----")
699
+ print("")
700
+
701
+ exp_iter_indices = [
702
+ histories.loc[exp_idx, :]
703
+ .index.get_level_values("iteration")[histories.loc[exp_idx]["w"].notna()]
704
+ .to_list()
705
+ for exp_idx in histories.index.get_level_values("trial").unique()
706
+ ]
707
+ exp_maxiter = np.argmax([ind[-1] for ind in exp_iter_indices])
708
+ longest_exp_indices = exp_iter_indices[exp_maxiter]
709
+ longest_exp_indices.extend(
710
+ [ei[-1] for ei in exp_iter_indices if ei[-1] not in longest_exp_indices]
711
+ )
712
+ longest_exp_indices = list(set(longest_exp_indices))
713
+ longest_exp_indices.sort()
714
+
715
+ index = pd.MultiIndex.from_product(
716
+ [longest_exp_indices, range(N_RUNS)],
717
+ names=("iteration", "trial"),
718
+ )
719
+ full_eval_train = pd.DataFrame(
720
+ index=index, columns=["G", "f", "fg", "c", "cg"]
721
+ ).sort_index()
722
+ full_eval_test = pd.DataFrame(
723
+ index=index, columns=["G", "f", "fg", "c", "cg"]
724
+ ).sort_index()
725
+
726
+ criterion = nn.BCEWithLogitsLoss()
727
+ X_test_tensor = tensor(X_test, dtype=DTYPE).to(device)
728
+ y_test_tensor = tensor(y_test, dtype=DTYPE).to(device)
729
+ X_train_tensor = X_train_tensor.to(device=device)
730
+ y_train_tensor = y_train_tensor.to(device=device)
731
+
732
+ save_train = True
733
+ save_test = True
734
+ histories.dropna(subset=["w"], inplace=True)
735
+
736
+ for exp_idx in range(N_RUNS):
737
+ for alg_iteration in histories.loc[exp_idx, :].index:
738
+ print(f"{exp_idx} | {alg_iteration}", end="\r")
739
+
740
+ w = histories["w"].loc[exp_idx, alg_iteration]
741
+ net.load_state_dict(w)
742
+ net = net.to(device)
743
+ if cfg.alg.import_name.lower() == "sslalm":
744
+ x_t = net_params_to_tensor(net, flatten=True, copy=True)
745
+ lambdas = histories["dual_ms"].loc[exp_idx, alg_iteration]
746
+ z = histories["z"].loc[exp_idx, alg_iteration]
747
+ params = {
748
+ "x_t": x_t,
749
+ "lambdas": lambdas,
750
+ "z": z,
751
+ "rho": cfg.alg.params.rho,
752
+ "mu": cfg.alg.params.mu,
753
+ }
754
+
755
+ if save_train:
756
+ if cfg.constraint.import_name == 'abs_max_dev_from_overall_tpr':
757
+ data_c = [[
758
+ (X_train_tensor[g_idx], y_train_tensor[g_idx]) for g_idx in group_ind_train
759
+ ]]
760
+ elif cfg.constraint.import_name in ['abs_diff_pr', 'abs_diff_tpr']:
761
+ data_c = [
762
+ (
763
+ (X_train_tensor[g_idx], y_train_tensor[g_idx]),
764
+ (X_train_tensor, y_train_tensor)
765
+ )
766
+ for g_idx in group_ind_train
767
+ ]
768
+ else:
769
+ data_c = [
770
+ (
771
+ (X_train_tensor[g_idx_1], y_train_tensor[g_idx_1]),
772
+ (X_train_tensor[g_idx_2], y_train_tensor[g_idx_2]),
773
+ )
774
+ for g_idx_1, g_idx_2 in combinations(group_ind_train, 2)
775
+ ]
776
+ calculate_iteration_values(
777
+ alg=cfg.alg.import_name,
778
+ full_eval=full_eval_train,
779
+ index_to_save=[alg_iteration, exp_idx],
780
+ c=c,
781
+ loss_fn=criterion,
782
+ data_f=[X_train_tensor, y_train_tensor],
783
+ data_c=data_c,
784
+ net=net,
785
+ device=device,
786
+ add_negative=cfg.constraint.add_negative,
787
+ **params,
788
+ )
789
+
790
+
791
+ if save_test:
792
+ if cfg.constraint.import_name == 'abs_max_dev_from_overall_tpr':
793
+ data_c = [[
794
+ (X_test_tensor[g_idx], y_test_tensor[g_idx]) for g_idx in group_ind_test
795
+ ]]
796
+ elif cfg.constraint.import_name in ['abs_diff_tpr', 'abs_diff_pr']:
797
+ data_c = [
798
+ (
799
+ (X_test_tensor[g_idx], y_test_tensor[g_idx]),
800
+ (X_test_tensor, y_test_tensor)
801
+ )
802
+ for g_idx in group_ind_test
803
+ ]
804
+ else:
805
+ data_c = [
806
+ (
807
+ (X_test_tensor[g_idx_1], y_test_tensor[g_idx_1]),
808
+ (X_test_tensor[g_idx_2], y_test_tensor[g_idx_2]),
809
+ )
810
+ for g_idx_1, g_idx_2 in combinations(group_ind_test, 2)
811
+ ]
812
+ calculate_iteration_values(
813
+ alg=cfg.alg.import_name,
814
+ full_eval=full_eval_test,
815
+ index_to_save=[alg_iteration, exp_idx],
816
+ c=c,
817
+ loss_fn=criterion,
818
+ data_f=[X_test_tensor, y_test_tensor],
819
+ data_c=data_c,
820
+ net=net,
821
+ device=device,
822
+ add_negative=cfg.constraint.add_negative,
823
+ **params,
824
+ )
825
+
826
+ net.zero_grad()
827
+
828
+ fname = f"AFTER_{alg_save_name}_{DATASET_NAME}_{LOSS_BOUND}"
829
+ fext = ".csv"
830
+ if save_train:
831
+ fname_train = fname + "_train" + fext
832
+ save_path = os.path.join(utils_path, fname_train)
833
+ print(f"Saving to: {save_path}")
834
+ full_eval_train.to_pickle(save_path)
835
+
836
+ if save_test:
837
+ fname_test = fname + "_test" + fext
838
+ save_path = os.path.join(utils_path, fname_test)
839
+ print(f"Saving to: {save_path}")
840
+ full_eval_test.to_pickle(save_path)
841
+
842
+
843
+ # helper function to construct pairwise constraints for every combination of provided groups
844
+ def construct_constraints(
845
+ constraint_fn,
846
+ bound,
847
+ dataset,
848
+ constraint_groups,
849
+ batch_size,
850
+ add_negative,
851
+ device,
852
+ seed,
853
+ max_0 = False
854
+ ):
855
+ c = []
856
+
857
+ for group_indices in constraint_groups:
858
+ for group_idx in combinations(group_indices, 2):
859
+ c1 = FairnessConstraint(
860
+ dataset,
861
+ group_idx,
862
+ fn=lambda net, d: torch.max(constraint_fn(net, d) - bound, torch.zeros(1)) if max_0 else constraint_fn(net, d) - bound,
863
+ batch_size=batch_size // 2,
864
+ device=device,
865
+ seed=seed,
866
+ )
867
+ c.append(c1)
868
+
869
+ if add_negative:
870
+ c2 = FairnessConstraint(
871
+ dataset,
872
+ group_idx,
873
+ fn=lambda net, d: torch.max(-constraint_fn(net, d) - bound, torch.zeros(1)) if max_0 else -constraint_fn(net, d) - bound,
874
+ batch_size=batch_size // 2,
875
+ device=device,
876
+ seed=seed,
877
+ )
878
+ c.append(c2)
879
+
880
+ return c
881
+
882
+ # helper function to calculate relevant values on full dataset (e.g. constraint gradient, AL function, etc)
883
+ # used to calculate those values at different points during algorithms run
884
+ def calculate_iteration_values(
885
+ alg,
886
+ full_eval,
887
+ index_to_save,
888
+ c,
889
+ loss_fn,
890
+ data_f,
891
+ data_c,
892
+ net,
893
+ device,
894
+ add_negative,
895
+ **params,
896
+ ):
897
+ c_val_vec, c_grads_mat = [], []
898
+
899
+ for i, c_i in enumerate(c):
900
+ cv = c_i.eval(net, data_c[i // 2 if add_negative else i])
901
+ c_val_vec.append(cv)
902
+ cv.backward()
903
+ cg = net_grads_to_tensor(net, flatten=True, device=device)
904
+ net.zero_grad()
905
+ c_grads_mat.append(cg)
906
+ c_val_vec = torch.tensor(c_val_vec)
907
+ c_grads_mat = torch.stack(c_grads_mat)
908
+ full_eval.loc[*index_to_save]["c"] = [c_val_vec.detach().cpu().numpy()]
909
+ full_eval.loc[*index_to_save]["cg"] = [c_grads_mat.detach().cpu().numpy()]
910
+
911
+ X_tensor, y_tensor = data_f
912
+ outs = net(X_tensor)
913
+ if y_tensor.ndim < outs.ndim:
914
+ y_tensor = y_tensor.unsqueeze(1)
915
+ loss = loss_fn(outs, y_tensor)
916
+ loss.backward()
917
+ fg = net_grads_to_tensor(net, flatten=True, device=device)
918
+ net.zero_grad()
919
+
920
+ full_eval.loc[*index_to_save]["f"] = loss.detach().cpu().numpy()
921
+ full_eval.loc[*index_to_save]["fg"] = [fg.detach().cpu().numpy()]
922
+
923
+ # if alg.lower() == "sgd":
924
+ # return
925
+
926
+ # elif alg.lower() == "sslalm":
927
+ # x_t, z, rho, mu, lambdas = (
928
+ # params["x_t"],
929
+ # params["z"],
930
+ # params["rho"],
931
+ # params["mu"],
932
+ # params["lambdas"],
933
+ # )
934
+ # G = (
935
+ # fg
936
+ # + c_grads_mat.T @ lambdas
937
+ # + rho * (c_grads_mat.T @ c_val_vec)
938
+ # + mu * (x_t - z)
939
+ # )
940
+
941
+ # full_eval.loc[*index_to_save]["G"] = [G.detach().cpu().numpy()]
942
+
943
+
944
+ def sample_or_restart_iterloader(loader):
945
+ try:
946
+ item = next(loader)
947
+ return item
948
+ except StopIteration:
949
+ loader._reset(loader)
950
+ # loader.gen
951
+ item = next(loader)
952
+ return item
953
+
954
+
955
+ if __name__ == "__main__":
956
+ run()