humancompatible-train 0.1.3__tar.gz → 0.1.4__tar.gz

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 (37) hide show
  1. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/PKG-INFO +1 -1
  2. humancompatible_train-0.1.4/experiments/run_dutch.py +439 -0
  3. humancompatible_train-0.1.4/experiments/run_experiment.py +196 -0
  4. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/experiments/run_folktables.py +145 -108
  5. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/algorithms/ssl_alm.py +1 -1
  6. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/fairness/utils/balanced_batch_sampler.py +5 -4
  7. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/fairness/utils/tests/test_balanced_batch_sampler.py +2 -0
  8. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible_train.egg-info/PKG-INFO +1 -1
  9. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible_train.egg-info/SOURCES.txt +2 -6
  10. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/pyproject.toml +1 -1
  11. humancompatible_train-0.1.3/experiments/run_folktables_torchalgs.py +0 -956
  12. humancompatible_train-0.1.3/humancompatible/train/fairness/constraints/__init__.py +0 -15
  13. humancompatible_train-0.1.3/humancompatible/train/fairness/constraints/constraint.py +0 -97
  14. humancompatible_train-0.1.3/humancompatible/train/fairness/constraints/constraint_fns.py +0 -244
  15. humancompatible_train-0.1.3/humancompatible/train/fairness/constraints/torch/__init__.py +0 -1
  16. humancompatible_train-0.1.3/humancompatible/train/fairness/constraints/torch/constraints.py +0 -36
  17. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/LICENCE.txt +0 -0
  18. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/README.md +0 -0
  19. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/experiments/__init__.py +0 -0
  20. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/experiments/calculate_iteration_values.py +0 -0
  21. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/__init__.py +0 -0
  22. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/__init__.py +0 -0
  23. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/algorithms/__init__.py +0 -0
  24. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/algorithms/ssl_alm_adam.py +0 -0
  25. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/algorithms/ssw.py +0 -0
  26. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/algorithms/test/__init__.py +0 -0
  27. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/algorithms/test/test_ssl_alm.py +0 -0
  28. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/algorithms/test/test_ssl_alm_adam.py +0 -0
  29. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/algorithms/test/test_ssw.py +0 -0
  30. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/fairness/__init__.py +0 -0
  31. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/fairness/utils/__init__.py +0 -0
  32. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/fairness/utils/tests/__init__.py +0 -0
  33. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible_train.egg-info/dependency_links.txt +0 -0
  34. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible_train.egg-info/requires.txt +0 -0
  35. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible_train.egg-info/top_level.txt +0 -0
  36. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/setup.cfg +0 -0
  37. {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/setup.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: humancompatible-train
3
- Version: 0.1.3
3
+ Version: 0.1.4
4
4
  Summary: PyTorch-based package for constrained training of neural networks
5
5
  Author: Gilles Bareilles, Jana Lepsova, Jakub Marecek
6
6
  Author-email: Andrii Kliachkin <kliachkin.andrii@gmail.com>
@@ -0,0 +1,439 @@
1
+ import importlib
2
+ from itertools import combinations
3
+ import os
4
+ import timeit
5
+ import warnings
6
+ import hydra
7
+ import numpy as np
8
+ import pandas as pd
9
+ import torch
10
+ from torch.utils.data import TensorDataset
11
+ from omegaconf import DictConfig, OmegaConf
12
+ from torch import nn, tensor
13
+ from utils.load_dutch import prepare_dutch
14
+ from utils.network import SimpleNet
15
+ from humancompatible.train.benchmark.algorithms.utils import net_grads_to_tensor
16
+ from itertools import combinations
17
+ from humancompatible.train.benchmark.constraints import FairnessConstraint
18
+
19
+
20
+
21
+ @hydra.main(version_base=None, config_path="conf", config_name="experiment")
22
+ def run(cfg: DictConfig) -> None:
23
+ warnings.filterwarnings("ignore", category=FutureWarning)
24
+
25
+ print(OmegaConf.to_yaml(cfg))
26
+ N_RUNS = cfg.n_runs
27
+ FT_STATE = 'dutch'
28
+ FT_TASK = 'dutch'
29
+ DOWNLOAD_DATA = cfg.data.download
30
+ DATA_PATH = cfg.data.path
31
+
32
+ if "constraint" in cfg.keys():
33
+ CONSTRAINT = cfg.constraint.import_name
34
+ LOSS_BOUND = cfg.constraint.bound
35
+ else:
36
+ CONSTRAINT = "unconstr"
37
+ LOSS_BOUND = 0
38
+
39
+ if cfg.device == "cpu":
40
+ device = "cpu"
41
+ elif cfg.alg == "ghost":
42
+ device = "cpu"
43
+ print("CUDA not supported for Stochastic Ghost")
44
+ elif torch.cuda.is_available():
45
+ device = "cuda"
46
+ print("CUDA found")
47
+ else:
48
+ device = "cpu"
49
+ print("CUDA not found")
50
+
51
+ print(f"{device = }")
52
+ torch.set_default_device(device)
53
+
54
+ DTYPE = torch.float32
55
+
56
+ ## load data ##
57
+
58
+ torch.set_default_dtype(DTYPE)
59
+ DATASET_NAME = FT_TASK + "_" + FT_STATE
60
+
61
+ (
62
+ (X_train, X_val, X_test),
63
+ (y_train, y_val, y_test),
64
+ (group_ind_train, group_ind_val, group_ind_test),
65
+ _,
66
+ _
67
+ ) = prepare_dutch(
68
+ random_state=None,
69
+ onehot=False,
70
+ stratify=True,
71
+ )
72
+ print(f'Train: {len(group_ind_train)} groups of size {[len(group) for group in group_ind_train]}')
73
+ print(f'Val: {len(group_ind_val)} groups of size {[len(group) for group in group_ind_val]}')
74
+ print(f'Test: {len(group_ind_test)} groups of size {[len(group) for group in group_ind_test]}')
75
+
76
+ X_train_tensor = tensor(X_train, dtype=DTYPE)
77
+ y_train_tensor = tensor(y_train, dtype=DTYPE)
78
+ train_ds = TensorDataset(X_train_tensor, y_train_tensor)
79
+
80
+ print(f"Train data loaded: {(FT_TASK, FT_STATE)}")
81
+ print(f"Data shape: {X_train_tensor.shape}")
82
+
83
+ ## prepare to save results ##
84
+
85
+ if "save_name" in cfg["alg"].keys():
86
+ alg_save_name = cfg.alg.save_name
87
+ else:
88
+ alg_save_name = cfg.alg.import_name
89
+
90
+ saved_models_path = os.path.abspath(
91
+ os.path.join(os.path.dirname(__file__), "utils", "saved_models")
92
+ )
93
+ directory = os.path.join(
94
+ saved_models_path, DATASET_NAME, CONSTRAINT, f"{LOSS_BOUND:.0E}"
95
+ )
96
+
97
+ model_name = os.path.join(directory, f"{alg_save_name}_{LOSS_BOUND}")
98
+
99
+ if not os.path.exists(directory):
100
+ os.makedirs(directory)
101
+
102
+ ## run experiments ##
103
+ histories = []
104
+ for EXP_IDX in range(1, N_RUNS+1):
105
+ print(f'Start run {EXP_IDX}\n')
106
+
107
+ ## define constraints ##
108
+ loss_fn = nn.BCEWithLogitsLoss()
109
+ constraint_fn_module = importlib.import_module("humancompatible.train.benchmark.constraints")
110
+ constraint_fn = getattr(constraint_fn_module, cfg.constraint.import_name)
111
+
112
+ if cfg.constraint.type == 'one_vs_mean':
113
+ c = [
114
+ FairnessConstraint(
115
+ train_ds,
116
+ [group_ind, np.concat(group_ind_train)],
117
+ fn=lambda net, inputs: constraint_fn(loss_fn, net, inputs) - cfg.constraint.bound,
118
+ batch_size=cfg.constraint.c_batch_size,
119
+ seed=EXP_IDX
120
+ )
121
+ for group_ind in group_ind_train
122
+ ]
123
+ if cfg.constraint.add_negative:
124
+ c.extend(
125
+ [
126
+ FairnessConstraint(
127
+ train_ds,
128
+ [group_ind, np.concat(group_ind_train)],
129
+ fn=lambda net, inputs: -constraint_fn(loss_fn, net, inputs) - cfg.constraint.bound,
130
+ batch_size=cfg.constraint.c_batch_size,
131
+ seed=EXP_IDX
132
+ )
133
+ for group_ind in group_ind_train
134
+ ]
135
+ )
136
+ elif cfg.constraint.type == 'one_vs_each':
137
+ c = [
138
+ FairnessConstraint(
139
+ train_ds,
140
+ group_idx,
141
+ fn=lambda net, inputs: constraint_fn(loss_fn, net, inputs) - cfg.constraint.bound,
142
+ batch_size=cfg.constraint.c_batch_size,
143
+ device=device,
144
+ seed=EXP_IDX,
145
+ )
146
+ for group_idx in combinations(group_ind_train, 2)
147
+ ]
148
+ if cfg.constraint.add_negative:
149
+ c.extend(
150
+ [
151
+ FairnessConstraint(
152
+ train_ds,
153
+ group_idx,
154
+ fn=lambda net, inputs: -constraint_fn(loss_fn, net, inputs) - cfg.constraint.bound,
155
+ batch_size=cfg.constraint.c_batch_size,
156
+ device=device,
157
+ seed=EXP_IDX,
158
+ )
159
+ for group_idx in combinations(group_ind_train, 2)
160
+ ]
161
+ )
162
+
163
+ torch.manual_seed(EXP_IDX)
164
+ net = SimpleNet(in_shape=X_test.shape[1], out_shape=1, dtype=DTYPE).to(device)
165
+ model_path = model_name + f"_trial{EXP_IDX}.pt"
166
+
167
+ optimizer_name = cfg.alg.import_name
168
+ module = importlib.import_module("humancompatible.train.benchmark.algorithms")
169
+ Optimizer = getattr(module, optimizer_name)
170
+ optimizer = Optimizer(net, train_ds, loss_fn, c)
171
+ # inconsequential backward pass cause first pass is very slow
172
+ x = net.forward(X_train_tensor[0])
173
+ x.backward()
174
+ net.zero_grad()
175
+ # train!
176
+ history = optimizer.optimize(
177
+ **cfg.alg.params,
178
+ max_iter=cfg.run_maxiter,
179
+ max_runtime=cfg.run_maxtime,
180
+ device=device,
181
+ # seed=EXP_IDX,
182
+ verbose=True,
183
+ )
184
+
185
+ ## SAVE RESULTS ##
186
+ params = pd.DataFrame(history["params"])
187
+ values = pd.DataFrame(history["values"])
188
+ t = pd.Series(history["time"], name="time")
189
+ histories.append(values.join(params, how="outer").join(t, how="outer"))
190
+
191
+ ## SAVE MODEL ##
192
+ print(f"Model saved to: {model_path}")
193
+ torch.save(net.state_dict(), model_path)
194
+ print("")
195
+
196
+ # Save DataFrames to CSV files
197
+ c_name = cfg.constraint.import_name
198
+ utils_path = os.path.abspath(
199
+ os.path.join(os.path.dirname(__file__), "utils", "exp_results", c_name)
200
+ )
201
+ if not os.path.exists(utils_path):
202
+ os.makedirs(utils_path)
203
+
204
+ if cfg.save_checkpoint_df:
205
+ fname = f"{alg_save_name}_{DATASET_NAME}_{LOSS_BOUND}.csv"
206
+ save_path = os.path.join(utils_path, fname)
207
+ print(f"Saving to: {save_path}")
208
+ histories = pd.concat(histories, keys=range(N_RUNS), names=["trial", "iteration"])
209
+ histories.to_pickle(save_path)
210
+ print("Saved!")
211
+
212
+ ####################################################
213
+ ### CALCULATE STATS ON EVERY ALGORITHM ITERATION ###
214
+ ####################################################
215
+
216
+ loss_fn = nn.BCEWithLogitsLoss()
217
+ constraint_fn_module = importlib.import_module("humancompatible.train.benchmark.constraints")
218
+ constraint_fn = getattr(constraint_fn_module, cfg.constraint.import_name)
219
+
220
+ print("----")
221
+ print("")
222
+
223
+ exp_iter_indices = [
224
+ histories.loc[exp_idx, :]
225
+ .index.get_level_values("iteration")[histories.loc[exp_idx]["w"].notna()]
226
+ .to_list()
227
+ for exp_idx in histories.index.get_level_values("trial").unique()
228
+ ]
229
+ exp_maxiter = np.argmax([ind[-1] for ind in exp_iter_indices])
230
+ longest_exp_indices = exp_iter_indices[exp_maxiter]
231
+ longest_exp_indices.extend(
232
+ [ei[-1] for ei in exp_iter_indices if ei[-1] not in longest_exp_indices]
233
+ )
234
+ longest_exp_indices = list(set(longest_exp_indices))
235
+ longest_exp_indices.sort()
236
+
237
+ index = pd.MultiIndex.from_product(
238
+ [longest_exp_indices, range(N_RUNS)],
239
+ names=("iteration", "trial"),
240
+ )
241
+ full_eval_train = pd.DataFrame(
242
+ index=index, columns=["G", "f", "fg", "c", "cg"]
243
+ ).sort_index()
244
+ full_eval_val = full_eval_train.copy()
245
+ full_eval_test = full_eval_train.copy()
246
+
247
+ loss_fn = nn.BCEWithLogitsLoss()
248
+ X_test_tensor = tensor(X_test, dtype=DTYPE).to(device)
249
+ y_test_tensor = tensor(y_test, dtype=DTYPE).to(device)
250
+
251
+ X_val_tensor = tensor(X_val, dtype=DTYPE).to(device)
252
+ y_val_tensor = tensor(y_val, dtype=DTYPE).to(device)
253
+
254
+ X_train_tensor = X_train_tensor.to(device=device)
255
+ y_train_tensor = y_train_tensor.to(device=device)
256
+
257
+ save_train = True
258
+ save_val = True
259
+ save_test = True
260
+ histories.dropna(subset=["w"], inplace=True)
261
+ with torch.inference_mode():
262
+ for exp_idx in range(N_RUNS):
263
+ for alg_iteration in histories.loc[exp_idx, :].index:
264
+ print(f"{exp_idx} | {alg_iteration}", end="\r")
265
+
266
+ w = histories["w"].loc[exp_idx, alg_iteration]
267
+ net.load_state_dict(w)
268
+ net = net.to(device)
269
+ if save_train:
270
+ if cfg.constraint.type=="one_vs_mean":
271
+ data_c = [
272
+ (
273
+ (X_train_tensor[g_idx], y_train_tensor[g_idx]),
274
+ (X_train_tensor, y_train_tensor)
275
+ )
276
+ for g_idx in group_ind_train
277
+ ]
278
+ elif cfg.constraint.type=="one_vs_each":
279
+ data_c = [
280
+ (
281
+ (X_train_tensor[g_idx_1], y_train_tensor[g_idx_1]),
282
+ (X_train_tensor[g_idx_2], y_train_tensor[g_idx_2]),
283
+ )
284
+ for g_idx_1, g_idx_2 in combinations(group_ind_train, 2)
285
+ ]
286
+ calculate_iteration_values(
287
+ alg=cfg.alg.import_name,
288
+ full_eval=full_eval_train,
289
+ index_to_save=[alg_iteration, exp_idx],
290
+ c=c,
291
+ loss_fn=loss_fn,
292
+ data_f=[X_train_tensor, y_train_tensor],
293
+ data_c=data_c,
294
+ net=net,
295
+ device=device,
296
+ add_negative=cfg.constraint.add_negative,
297
+ **params,
298
+ )
299
+
300
+ if save_val:
301
+ if cfg.constraint.type=="one_vs_mean":
302
+ data_c = [
303
+ (
304
+ (X_val_tensor[g_idx], y_val_tensor[g_idx]),
305
+ (X_val_tensor, y_val_tensor)
306
+ )
307
+ for g_idx in group_ind_val
308
+ ]
309
+ elif cfg.constraint.type=="one_vs_each":
310
+ data_c = [
311
+ (
312
+ (X_val_tensor[g_idx_1], y_val_tensor[g_idx_1]),
313
+ (X_val_tensor[g_idx_2], y_val_tensor[g_idx_2]),
314
+ )
315
+ for g_idx_1, g_idx_2 in combinations(group_ind_val, 2)
316
+ ]
317
+ calculate_iteration_values(
318
+ alg=cfg.alg.import_name,
319
+ full_eval=full_eval_val,
320
+ index_to_save=[alg_iteration, exp_idx],
321
+ c=c,
322
+ loss_fn=loss_fn,
323
+ data_f=[X_val_tensor, y_val_tensor],
324
+ data_c=data_c,
325
+ net=net,
326
+ device=device,
327
+ add_negative=cfg.constraint.add_negative,
328
+ **params,
329
+ )
330
+
331
+ if save_test:
332
+ if cfg.constraint.type=="one_vs_mean":
333
+ data_c = [
334
+ (
335
+ (X_test_tensor[g_idx], y_test_tensor[g_idx]),
336
+ (X_test_tensor, y_test_tensor)
337
+ )
338
+ for g_idx in group_ind_test
339
+ ]
340
+ elif cfg.constraint.type=="one_vs_each":
341
+ data_c = [
342
+ (
343
+ (X_test_tensor[g_idx_1], y_test_tensor[g_idx_1]),
344
+ (X_test_tensor[g_idx_2], y_test_tensor[g_idx_2]),
345
+ )
346
+ for g_idx_1, g_idx_2 in combinations(group_ind_test, 2)
347
+ ]
348
+ calculate_iteration_values(
349
+ alg=cfg.alg.import_name,
350
+ full_eval=full_eval_test,
351
+ index_to_save=[alg_iteration, exp_idx],
352
+ c=c,
353
+ loss_fn=loss_fn,
354
+ data_f=[X_test_tensor, y_test_tensor],
355
+ data_c=data_c,
356
+ net=net,
357
+ device=device,
358
+ add_negative=cfg.constraint.add_negative,
359
+ **params,
360
+ )
361
+
362
+ net.zero_grad()
363
+
364
+ fname = f"AFTER_{alg_save_name}_{DATASET_NAME}_{LOSS_BOUND}"
365
+ fext = ".csv"
366
+ if save_train:
367
+ fname_train = fname + "_train" + fext
368
+ save_path = os.path.join(utils_path, fname_train)
369
+ print(f"Saving to: {save_path}")
370
+ full_eval_train.to_pickle(save_path)
371
+
372
+ if save_val:
373
+ fname_val = fname + "_val" + fext
374
+ save_path = os.path.join(utils_path, fname_val)
375
+ print(f"Saving to: {save_path}")
376
+ full_eval_val.to_pickle(save_path)
377
+
378
+ if save_test:
379
+ fname_test = fname + "_test" + fext
380
+ save_path = os.path.join(utils_path, fname_test)
381
+ print(f"Saving to: {save_path}")
382
+ full_eval_test.to_pickle(save_path)
383
+
384
+
385
+ # helper function to calculate relevant values on full dataset (e.g. constraint gradient, AL function, etc)
386
+ # used to calculate those values at different points during algorithms run
387
+ def calculate_iteration_values(
388
+ alg,
389
+ full_eval,
390
+ index_to_save,
391
+ c,
392
+ loss_fn,
393
+ data_f,
394
+ data_c,
395
+ net,
396
+ device,
397
+ add_negative,
398
+ **params,
399
+ ):
400
+ c_val_vec, c_grads_mat = [], []
401
+
402
+ for i, c_i in enumerate(c):
403
+ cv = c_i.eval(net, data_c[i // 2 if add_negative else i])
404
+ c_val_vec.append(cv)
405
+ # cv.backward()
406
+ # cg = net_grads_to_tensor(net, flatten=True, device=device)
407
+ net.zero_grad()
408
+ # c_grads_mat.append(cg)
409
+ c_val_vec = torch.tensor(c_val_vec)
410
+ # c_grads_mat = torch.stack(c_grads_mat)
411
+ full_eval.loc[*index_to_save]["c"] = [c_val_vec.detach().cpu().numpy()]
412
+ # full_eval.loc[*index_to_save]["cg"] = [c_grads_mat.detach().cpu().numpy()]
413
+
414
+ X_tensor, y_tensor = data_f
415
+ outs = net(X_tensor)
416
+ if y_tensor.ndim < outs.ndim:
417
+ y_tensor = y_tensor.unsqueeze(1)
418
+ loss = loss_fn(outs, y_tensor)
419
+ # loss.backward()
420
+ # fg = net_grads_to_tensor(net, flatten=True, device=device)
421
+ # net.zero_grad()
422
+
423
+ full_eval.loc[*index_to_save]["f"] = loss.detach().cpu().numpy()
424
+ # full_eval.loc[*index_to_save]["fg"] = [fg.detach().cpu().numpy()]
425
+
426
+
427
+ def sample_or_restart_iterloader(loader):
428
+ try:
429
+ item = next(loader)
430
+ return item
431
+ except StopIteration:
432
+ loader._reset(loader)
433
+ # loader.gen
434
+ item = next(loader)
435
+ return item
436
+
437
+
438
+ if __name__ == "__main__":
439
+ run()
@@ -0,0 +1,196 @@
1
+ import importlib
2
+ import os
3
+ import warnings
4
+ import hydra
5
+ import pandas as pd
6
+ import torch
7
+ from torch.utils.data import TensorDataset
8
+ from omegaconf import DictConfig, OmegaConf
9
+ from torch import nn, tensor
10
+ from utils.load_folktables import prepare_folktables_multattr
11
+ from utils.network import SimpleNet
12
+ from utils.utils import create_constraint_from_cfg, run_summary_full_set
13
+ from humancompatible.train.benchmark.algorithms.utils import net_grads_to_tensor
14
+ from humancompatible.train.benchmark.constraints import FairnessConstraint
15
+
16
+
17
+ @hydra.main(version_base=None, config_path="conf", config_name="experiment")
18
+ def run(cfg: DictConfig) -> None:
19
+ warnings.filterwarnings("ignore", category=FutureWarning)
20
+
21
+ print(OmegaConf.to_yaml(cfg))
22
+ N_RUNS = cfg.n_runs
23
+ FT_STATE = cfg.data.state
24
+ FT_TASK = cfg.data.task
25
+ DOWNLOAD_DATA = cfg.data.download
26
+ DATA_PATH = cfg.data.path
27
+
28
+ if cfg.device == "cpu":
29
+ device = "cpu"
30
+ elif cfg.alg == "ghost":
31
+ device = "cpu"
32
+ print("CUDA not supported for Stochastic Ghost")
33
+ elif torch.cuda.is_available():
34
+ device = "cuda"
35
+ print("CUDA found")
36
+ else:
37
+ device = "cpu"
38
+ print("CUDA not found")
39
+
40
+ print(f"{device = }")
41
+ torch.set_default_device(device)
42
+
43
+ DTYPE = torch.float32
44
+
45
+ torch.set_default_dtype(DTYPE)
46
+ (
47
+ (X_train, X_val, X_test),
48
+ (y_train, y_val, y_test),
49
+ (group_ind_train, group_ind_val, group_ind_test),
50
+ _,
51
+ _
52
+ ) = prepare_folktables_multattr(
53
+ FT_TASK,
54
+ state=FT_STATE.upper(),
55
+ random_state=42,
56
+ onehot=False,
57
+ download=DOWNLOAD_DATA,
58
+ path=DATA_PATH,
59
+ sens_cols=cfg.data.sens_attr,
60
+ binarize=cfg.data.binarize,
61
+ stratify=cfg.data.stratify,
62
+ )
63
+
64
+ print(f'Train: {len(group_ind_train)} groups of size {[len(group) for group in group_ind_train]}')
65
+ print(f'Val: {len(group_ind_val)} groups of size {[len(group) for group in group_ind_val]}')
66
+ print(f'Test: {len(group_ind_test)} groups of size {[len(group) for group in group_ind_test]}')
67
+
68
+ X_train_tensor = tensor(X_train, dtype=DTYPE)
69
+ y_train_tensor = tensor(y_train, dtype=DTYPE)
70
+ train_ds = TensorDataset(X_train_tensor, y_train_tensor)
71
+
72
+ print(f"Train data loaded: {(FT_TASK, FT_STATE)}")
73
+ print(f"Data shape: {X_train_tensor.shape}")
74
+
75
+ #TODO change to actual current date time lol
76
+ EXPERIMENT_NAME = 'CURRENT_DATE_TIME'
77
+ exp_save_path = os.path.join(os.path.dirname(__file__), EXPERIMENT_NAME)
78
+ os.makedirs(exp_save_path, exist_ok=True)
79
+ OmegaConf.save(cfg, os.path.join(exp_save_path, 'config.yaml'))
80
+
81
+ for RUN_IDX in range(1, N_RUNS+1):
82
+
83
+ print(f'Start run {RUN_IDX}\n')
84
+
85
+ ## prepare files ##
86
+ run_save_path = os.path.abspath(
87
+ os.path.join(exp_save_path, str(RUN_IDX))
88
+ )
89
+ if not os.path.exists(run_save_path):
90
+ os.makedirs(run_save_path)
91
+
92
+ ## define constraints ##
93
+ loss_fn = nn.BCEWithLogitsLoss()
94
+ c = create_constraint_from_cfg(
95
+ cfg=cfg,
96
+ dataset=train_ds,
97
+ group_indices=group_ind_train,
98
+ loss_fn=loss_fn,
99
+ device=device,
100
+ seed=RUN_IDX)
101
+
102
+ ## define network ##
103
+ # TODO: add a choice of net to cfg
104
+ net = SimpleNet(in_shape=X_test.shape[1], out_shape=1, dtype=DTYPE).to(device)
105
+
106
+ ## define optimizer
107
+ optimizer_name = cfg.alg.import_name
108
+ module = importlib.import_module("humancompatible.train.benchmark.algorithms")
109
+ Optimizer = getattr(module, optimizer_name)
110
+ optimizer = Optimizer(net, train_ds, loss_fn, c)
111
+ # inconsequential backward pass cause first pass is very slow
112
+ x = net.forward(X_train_tensor[0])
113
+ x.backward()
114
+ net.zero_grad()
115
+ # train!
116
+ history = optimizer.optimize(
117
+ **cfg.alg.params,
118
+ max_iter=cfg.run_maxiter,
119
+ max_runtime=cfg.run_maxtime,
120
+ device=device,
121
+ verbose=True,
122
+ )
123
+
124
+ ## Process and save results ##
125
+ model_checkpoints_save_path = os.path.join(run_save_path, 'model_states')
126
+ os.makedirs(model_checkpoints_save_path, exist_ok=True)
127
+ for iteration, state_dict in history['params']['w'].items():
128
+ # save model
129
+ torch.save(state_dict, os.path.join(model_checkpoints_save_path, f'{iteration}.pt'))
130
+
131
+ # save optimizer states
132
+ params = pd.DataFrame(history["params"])
133
+ values = pd.DataFrame(history["values"])
134
+ t = pd.Series(history["time"], name="time")
135
+ alg_states_df = pd.concat([t, params, values], axis=1)
136
+ alg_states_df.to_csv(os.path.join(run_save_path, 'alg_states.csv'))
137
+
138
+ ## SAVE MODEL ##
139
+ torch.save(net.state_dict(), os.path.join(run_save_path, 'model.pt'))
140
+ print(f"Model saved to: {run_save_path} as model.pt")
141
+ print("")
142
+
143
+ ####################################################
144
+ ### CALCULATE STATS ON EVERY ALGORITHM ITERATION ###
145
+ ####################################################
146
+ save_train = True
147
+ save_val = True
148
+ save_test = True
149
+
150
+ save_states_index = alg_states_df.index % cfg.alg.params.save_state_interval == 0
151
+
152
+ if save_train:
153
+ # TODO: join with table of iters and times based on checkpoint save interval
154
+ table_train = run_summary_full_set(model=net, dataset=train_ds, path=model_checkpoints_save_path, constraints=c, loss_fn=torch.nn.BCEWithLogitsLoss())
155
+ table_train.index = alg_states_df.index[save_states_index]
156
+ table_train['time'] = alg_states_df['time'].loc[save_states_index]
157
+ if save_val:
158
+ X_val_tensor, y_val_tensor = tensor(X_val, dtype=DTYPE), tensor(y_val, dtype=DTYPE)
159
+ val_ds = TensorDataset(X_val_tensor, y_val_tensor)
160
+ c_val = create_constraint_from_cfg(cfg, dataset=val_ds, group_indices=group_ind_val, loss_fn=torch.nn.BCEWithLogitsLoss(), device=device)
161
+
162
+ table_val = run_summary_full_set(model=net, dataset=val_ds, path=model_checkpoints_save_path, constraints=c_val, loss_fn=torch.nn.BCEWithLogitsLoss())
163
+ table_val.index = alg_states_df.index[save_states_index]
164
+ table_val['time'] = alg_states_df['time'].loc[save_states_index]
165
+ if save_test:
166
+ X_test_tensor, y_test_tensor = tensor(X_test, dtype=DTYPE), tensor(y_test, dtype=DTYPE)
167
+ test_ds = TensorDataset(X_test_tensor, y_test_tensor)
168
+ c_test = create_constraint_from_cfg(cfg, dataset=test_ds, group_indices=group_ind_test, loss_fn=torch.nn.BCEWithLogitsLoss(), device=device)
169
+
170
+ table_test = run_summary_full_set(model=net, dataset=test_ds, path=model_checkpoints_save_path, constraints=c_test, loss_fn=torch.nn.BCEWithLogitsLoss())
171
+ table_test.index = alg_states_df.index[save_states_index]
172
+ table_test['time'] = alg_states_df['time'].loc[save_states_index]
173
+
174
+ fname = f"full_set_eval"
175
+ fext = ".csv"
176
+ if save_train:
177
+ fname_train = fname + "_train" + fext
178
+ save_path = os.path.join(run_save_path, fname_train)
179
+ print(f"Saving to: {save_path}")
180
+ table_train.to_pickle(save_path)
181
+
182
+ if save_val:
183
+ fname_val = fname + "_val" + fext
184
+ save_path = os.path.join(run_save_path, fname_val)
185
+ print(f"Saving to: {save_path}")
186
+ table_val.to_pickle(save_path)
187
+
188
+ if save_test:
189
+ fname_test = fname + "_test" + fext
190
+ save_path = os.path.join(run_save_path, fname_test)
191
+ print(f"Saving to: {save_path}")
192
+ table_test.to_pickle(save_path)
193
+
194
+
195
+ if __name__ == "__main__":
196
+ run()