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.
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/PKG-INFO +1 -1
- humancompatible_train-0.1.4/experiments/run_dutch.py +439 -0
- humancompatible_train-0.1.4/experiments/run_experiment.py +196 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/experiments/run_folktables.py +145 -108
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/algorithms/ssl_alm.py +1 -1
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/fairness/utils/balanced_batch_sampler.py +5 -4
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/fairness/utils/tests/test_balanced_batch_sampler.py +2 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible_train.egg-info/PKG-INFO +1 -1
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible_train.egg-info/SOURCES.txt +2 -6
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/pyproject.toml +1 -1
- humancompatible_train-0.1.3/experiments/run_folktables_torchalgs.py +0 -956
- humancompatible_train-0.1.3/humancompatible/train/fairness/constraints/__init__.py +0 -15
- humancompatible_train-0.1.3/humancompatible/train/fairness/constraints/constraint.py +0 -97
- humancompatible_train-0.1.3/humancompatible/train/fairness/constraints/constraint_fns.py +0 -244
- humancompatible_train-0.1.3/humancompatible/train/fairness/constraints/torch/__init__.py +0 -1
- humancompatible_train-0.1.3/humancompatible/train/fairness/constraints/torch/constraints.py +0 -36
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/LICENCE.txt +0 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/README.md +0 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/experiments/__init__.py +0 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/experiments/calculate_iteration_values.py +0 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/__init__.py +0 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/__init__.py +0 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/algorithms/__init__.py +0 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/algorithms/ssl_alm_adam.py +0 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/algorithms/ssw.py +0 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/algorithms/test/__init__.py +0 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/algorithms/test/test_ssl_alm.py +0 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/algorithms/test/test_ssl_alm_adam.py +0 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/algorithms/test/test_ssw.py +0 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/fairness/__init__.py +0 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/fairness/utils/__init__.py +0 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible/train/fairness/utils/tests/__init__.py +0 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible_train.egg-info/dependency_links.txt +0 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible_train.egg-info/requires.txt +0 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/humancompatible_train.egg-info/top_level.txt +0 -0
- {humancompatible_train-0.1.3 → humancompatible_train-0.1.4}/setup.cfg +0 -0
- {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
|
+
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()
|