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,212 @@
|
|
|
1
|
+
from typing import Iterable, Optional, Union
|
|
2
|
+
|
|
3
|
+
import torch
|
|
4
|
+
from torch import Tensor
|
|
5
|
+
|
|
6
|
+
from torch.optim.optimizer import Optimizer, _use_grad_for_differentiable
|
|
7
|
+
|
|
8
|
+
class SSLALM(Optimizer):
|
|
9
|
+
def __init__(
|
|
10
|
+
self,
|
|
11
|
+
params,
|
|
12
|
+
m: int,
|
|
13
|
+
# tau in paper
|
|
14
|
+
lr: Union[float, Tensor] = 5e-2,
|
|
15
|
+
# eta in paper
|
|
16
|
+
dual_lr: Union[
|
|
17
|
+
float, Tensor
|
|
18
|
+
] = 5e-2, # keep as tensor for different learning rates for different constraints in the future? idk
|
|
19
|
+
dual_bound : Union[
|
|
20
|
+
float, Tensor
|
|
21
|
+
] = 100,
|
|
22
|
+
# penalty term multiplier
|
|
23
|
+
rho: float = 1.0,
|
|
24
|
+
# smoothing term multiplier
|
|
25
|
+
mu: float = 2.0,
|
|
26
|
+
# smoothing term update multiplier
|
|
27
|
+
beta: float = 0.5,
|
|
28
|
+
*,
|
|
29
|
+
init_dual_vars: Optional[Tensor] = None,
|
|
30
|
+
# whether some of the dual variables should not be updated
|
|
31
|
+
fix_dual_vars: Optional[Tensor] = None,
|
|
32
|
+
differentiable: bool = False,
|
|
33
|
+
# custom_project_fn: Optional[Callable] = project_fn
|
|
34
|
+
):
|
|
35
|
+
if isinstance(lr, torch.Tensor) and lr.numel() != 1:
|
|
36
|
+
raise ValueError("Tensor lr must be 1-element")
|
|
37
|
+
if isinstance(dual_lr, torch.Tensor) and lr.numel() != 1:
|
|
38
|
+
raise ValueError("Tensor dual_lr must be 1-element")
|
|
39
|
+
if lr < 0.0:
|
|
40
|
+
raise ValueError(f"Invalid learning rate: {lr}")
|
|
41
|
+
if dual_lr < 0.0:
|
|
42
|
+
raise ValueError(f"Invalid dual learning rate: {dual_lr}")
|
|
43
|
+
if init_dual_vars is not None and len(init_dual_vars) != m:
|
|
44
|
+
raise ValueError(
|
|
45
|
+
f"init_dual_vars should be of length m: expected {m}, got {len(init_dual_vars)}"
|
|
46
|
+
)
|
|
47
|
+
if fix_dual_vars is not None:
|
|
48
|
+
raise NotImplementedError()
|
|
49
|
+
if init_dual_vars is None and fix_dual_vars is not None:
|
|
50
|
+
raise ValueError(
|
|
51
|
+
f"if fix_dual_vars is not None, init_dual_vars should not be None."
|
|
52
|
+
)
|
|
53
|
+
|
|
54
|
+
if differentiable:
|
|
55
|
+
raise NotImplementedError("TorchSSLALM does not support differentiable")
|
|
56
|
+
|
|
57
|
+
defaults = dict(
|
|
58
|
+
lr=lr,
|
|
59
|
+
dual_lr=dual_lr,
|
|
60
|
+
rho=rho,
|
|
61
|
+
mu=mu,
|
|
62
|
+
beta=beta,
|
|
63
|
+
differentiable=differentiable,
|
|
64
|
+
# custom_project_fn=custom_project_fn
|
|
65
|
+
)
|
|
66
|
+
|
|
67
|
+
super().__init__(params, defaults)
|
|
68
|
+
|
|
69
|
+
# self.param_groups.append()
|
|
70
|
+
|
|
71
|
+
self.m = m
|
|
72
|
+
self.dual_lr = dual_lr
|
|
73
|
+
self.dual_bound = dual_bound
|
|
74
|
+
self.rho = rho
|
|
75
|
+
self.beta = beta
|
|
76
|
+
self.mu = mu
|
|
77
|
+
self.c_vals: list[Union[float, Tensor]] = []
|
|
78
|
+
self._c_val_average = [None]
|
|
79
|
+
# essentially, move everything here to self.state[param_group]
|
|
80
|
+
# self.state[param_group]['smoothing_avg'] <= z for that param_group;
|
|
81
|
+
# ...['grad'] <= grad w.r.t. that param_group
|
|
82
|
+
# ...['G'] <= G w.r.t. that param_group // idk if necessary
|
|
83
|
+
# ...['c_grad'][c_i] <= grad of ith constraint w.r.t. that group<w
|
|
84
|
+
if init_dual_vars is not None:
|
|
85
|
+
self._dual_vars = init_dual_vars
|
|
86
|
+
else:
|
|
87
|
+
self._dual_vars = torch.zeros(m, requires_grad=False)
|
|
88
|
+
|
|
89
|
+
def _init_group(self, group, params, grads, c_grads, smoothing):
|
|
90
|
+
# SHOULDN'T calculate values, only set them from the state of the respective param_group
|
|
91
|
+
# calculations only happen in step() (or rather in the func version of step)
|
|
92
|
+
has_sparse_grad = False
|
|
93
|
+
|
|
94
|
+
for p in group["params"]:
|
|
95
|
+
state = self.state[p]
|
|
96
|
+
|
|
97
|
+
params.append(p)
|
|
98
|
+
|
|
99
|
+
# load z (smoothing term)
|
|
100
|
+
# Lazy state initialization
|
|
101
|
+
if len(state) == 0:
|
|
102
|
+
state["smoothing"] = p.detach().clone()
|
|
103
|
+
state["c_grad"] = []
|
|
104
|
+
|
|
105
|
+
smoothing.append(state.get("smoothing"))
|
|
106
|
+
|
|
107
|
+
grads.append(p.grad)
|
|
108
|
+
c_grads.append(state.get("c_grad"))
|
|
109
|
+
|
|
110
|
+
return has_sparse_grad
|
|
111
|
+
|
|
112
|
+
def __setstate__(self, state):
|
|
113
|
+
super().__setstate__(state)
|
|
114
|
+
|
|
115
|
+
def dual_step(self, i: int, c_val: Tensor):
|
|
116
|
+
r"""Perform an update of the dual parameters.
|
|
117
|
+
Also saves constraint gradient for weight update. To be called BEFORE :func:`step` in an iteration!
|
|
118
|
+
|
|
119
|
+
Args:
|
|
120
|
+
i (int): index of the constraint
|
|
121
|
+
c_val (Tensor): an estimate of the value of the constraint at which the gradient was computed; used for dual parameter update
|
|
122
|
+
"""
|
|
123
|
+
|
|
124
|
+
# c_vals is cleaned in step()
|
|
125
|
+
self.c_vals.append(c_val.detach())
|
|
126
|
+
|
|
127
|
+
# update dual multipliers
|
|
128
|
+
dual_update_tensor = torch.zeros_like(self._dual_vars)
|
|
129
|
+
dual_update_tensor[i] = self.dual_lr * c_val
|
|
130
|
+
self._dual_vars.add_(dual_update_tensor)
|
|
131
|
+
for i in range(len(self._dual_vars)):
|
|
132
|
+
if self._dual_vars[i] >= self.dual_bound or self._dual_vars[i] < 0:
|
|
133
|
+
self._dual_vars[i].zero_()
|
|
134
|
+
|
|
135
|
+
# save constraint grad
|
|
136
|
+
for group in self.param_groups:
|
|
137
|
+
params: list[Tensor] = []
|
|
138
|
+
grads: list[Tensor] = []
|
|
139
|
+
c_grads: list[Tensor] = []
|
|
140
|
+
smoothing: list[Tensor] = []
|
|
141
|
+
_ = self._init_group(group, params, grads, c_grads, smoothing)
|
|
142
|
+
|
|
143
|
+
for p in group["params"]:
|
|
144
|
+
state = self.state[p]
|
|
145
|
+
# state['c_grad'] is cleaned in step()
|
|
146
|
+
# so it is always empty on dual_step()
|
|
147
|
+
state["c_grad"].append(p.grad)
|
|
148
|
+
|
|
149
|
+
@_use_grad_for_differentiable
|
|
150
|
+
def step(self, c_val: Union[Iterable | Tensor] = None):
|
|
151
|
+
r"""Perform an update of the primal parameters (network weights & slack variables). To be called AFTER :func:`dual_step` in an iteration!
|
|
152
|
+
|
|
153
|
+
Args:
|
|
154
|
+
c_val (Tensor): an Iterable of estimates of values of **ALL** constraints; used for primal parameter update.
|
|
155
|
+
Ideally, must be evaluated on an independent sample from the one used in :func:`dual_step`
|
|
156
|
+
"""
|
|
157
|
+
|
|
158
|
+
if c_val is None:
|
|
159
|
+
c_val = self.c_vals
|
|
160
|
+
if isinstance(c_val, Iterable) and not isinstance(c_val, torch.Tensor):
|
|
161
|
+
# if len(c_val) == 1 and isinstance(c_val[0], torch.Tensor):
|
|
162
|
+
# c_val = c_val[0]
|
|
163
|
+
# else:
|
|
164
|
+
c_val = torch.stack(c_val)
|
|
165
|
+
if c_val.ndim > 1:
|
|
166
|
+
c_val = c_val.squeeze(-1)
|
|
167
|
+
|
|
168
|
+
if c_val.numel() != self.m:
|
|
169
|
+
raise ValueError(f"Number of elements in c_val must be equal to m={self.m}, got {c_val.numel()}")
|
|
170
|
+
G = []
|
|
171
|
+
|
|
172
|
+
for group in self.param_groups:
|
|
173
|
+
params: list[Tensor] = []
|
|
174
|
+
grads: list[Tensor] = []
|
|
175
|
+
c_grads: list[Tensor] = []
|
|
176
|
+
smoothing: list[Tensor] = []
|
|
177
|
+
lr = group["lr"]
|
|
178
|
+
_ = self._init_group(group, params, grads, c_grads, smoothing)
|
|
179
|
+
|
|
180
|
+
for i, param in enumerate(params):
|
|
181
|
+
### calculate Lagrange f-n gradient (G) ###
|
|
182
|
+
|
|
183
|
+
# stack list of grads w.r.t. constraints to get
|
|
184
|
+
# tensor of shape (*param.shape, m)
|
|
185
|
+
l_term_grad = 0
|
|
186
|
+
aug_term_grad = 0
|
|
187
|
+
# if c_grads[i] is not None:
|
|
188
|
+
for j, c_grad in enumerate(c_grads[i]):
|
|
189
|
+
if c_grad is None:
|
|
190
|
+
continue
|
|
191
|
+
l_term_grad += c_grad * self._dual_vars[j]
|
|
192
|
+
aug_term_grad += c_grad * c_val[j]
|
|
193
|
+
|
|
194
|
+
G_i = (
|
|
195
|
+
grads[i]
|
|
196
|
+
+ l_term_grad
|
|
197
|
+
+ self.rho * aug_term_grad
|
|
198
|
+
+ self.mu * (param - smoothing[i])
|
|
199
|
+
)
|
|
200
|
+
G.append(G_i)
|
|
201
|
+
|
|
202
|
+
smoothing[i].add_(param - smoothing[i], alpha=self.beta)
|
|
203
|
+
|
|
204
|
+
param.add_(G_i, alpha=-lr)
|
|
205
|
+
|
|
206
|
+
## PROJECT (keep in mind we do layer by layer)
|
|
207
|
+
## add slack variables to params in constructor?
|
|
208
|
+
|
|
209
|
+
c_grads[i].clear()
|
|
210
|
+
|
|
211
|
+
self.c_vals.clear()
|
|
212
|
+
return G
|
|
@@ -0,0 +1,155 @@
|
|
|
1
|
+
from typing import Iterable, Optional, Union
|
|
2
|
+
|
|
3
|
+
import torch
|
|
4
|
+
from torch import Tensor
|
|
5
|
+
|
|
6
|
+
from torch.optim.optimizer import Optimizer, _use_grad_for_differentiable
|
|
7
|
+
|
|
8
|
+
# def project_fn(x, m):
|
|
9
|
+
# for i in range(1, m + 1):
|
|
10
|
+
# if x[-i] < 0:
|
|
11
|
+
# x[-i] = 0
|
|
12
|
+
# return x
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def _dual_step_func(dual_var, lr, cval):
|
|
16
|
+
return dual_var + lr * cval
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
# def step_fn(params, grads,)
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class SSG(Optimizer):
|
|
23
|
+
def __init__(
|
|
24
|
+
self,
|
|
25
|
+
params,
|
|
26
|
+
m: int = 1,
|
|
27
|
+
# constraint tolerance
|
|
28
|
+
# ctols: Union[
|
|
29
|
+
# float, Tensor
|
|
30
|
+
# ],
|
|
31
|
+
# ctols_rule = "const",
|
|
32
|
+
# learning rate
|
|
33
|
+
lr: Union[float, Tensor] = 5e-2,
|
|
34
|
+
# learning rate decrease rule
|
|
35
|
+
# lr_rule = "const",
|
|
36
|
+
# constraint learning rate
|
|
37
|
+
dual_lr: Union[
|
|
38
|
+
float, Tensor
|
|
39
|
+
] = 5e-2, # keep as tensor for different learning rates for different constraints in the future? idk
|
|
40
|
+
*,
|
|
41
|
+
differentiable: bool = False,
|
|
42
|
+
):
|
|
43
|
+
if isinstance(lr, torch.Tensor) and lr.numel() != 1:
|
|
44
|
+
raise ValueError("Tensor lr must be 1-element")
|
|
45
|
+
if isinstance(dual_lr, torch.Tensor) and lr.numel() != 1:
|
|
46
|
+
raise ValueError("Tensor dual_lr must be 1-element")
|
|
47
|
+
if lr < 0.0:
|
|
48
|
+
raise ValueError(f"Invalid learning rate: {lr}")
|
|
49
|
+
if dual_lr < 0.0:
|
|
50
|
+
raise ValueError(f"Invalid dual learning rate: {dual_lr}")
|
|
51
|
+
if not (m == 1):
|
|
52
|
+
raise ValueError(f"Switching Subgradient does not support multiple constraints."
|
|
53
|
+
"Consider taking the largest violation at each iteration.")
|
|
54
|
+
if differentiable:
|
|
55
|
+
raise NotImplementedError("SSw does not support differentiable")
|
|
56
|
+
|
|
57
|
+
defaults = dict(
|
|
58
|
+
lr=lr,
|
|
59
|
+
dual_lr=dual_lr,
|
|
60
|
+
differentiable=differentiable,
|
|
61
|
+
)
|
|
62
|
+
|
|
63
|
+
super().__init__(params, defaults)
|
|
64
|
+
|
|
65
|
+
# self.param_groups.append()
|
|
66
|
+
|
|
67
|
+
self.m = m
|
|
68
|
+
self.lr = lr
|
|
69
|
+
self.dual_lr = dual_lr
|
|
70
|
+
# self.lr_rule = lr_rule
|
|
71
|
+
# self.dual_lr_rule = dual_lr_rule
|
|
72
|
+
self.c_vals: list[Union[float, Tensor]] = []
|
|
73
|
+
# self.ctols = ctols
|
|
74
|
+
# essentially, move everything here to self.state[param_group]
|
|
75
|
+
# self.state[param_group]['smoothing_avg'] <= z for that param_group;
|
|
76
|
+
# ...['grad'] <= grad w.r.t. that param_group
|
|
77
|
+
# ...['G'] <= G w.r.t. that param_group // idk if necessary
|
|
78
|
+
# ...['c_grad'][c_i] <= grad of ith constraint w.r.t. that group<w
|
|
79
|
+
|
|
80
|
+
def _init_group(self, group, params, grads, c_grads):
|
|
81
|
+
# SHOULDN'T calculate values, only set them from the state of the respective param_group
|
|
82
|
+
# calculations only happen in step() (or rather in the func version of step)
|
|
83
|
+
has_sparse_grad = False
|
|
84
|
+
|
|
85
|
+
for p in group["params"]:
|
|
86
|
+
state = self.state[p]
|
|
87
|
+
|
|
88
|
+
params.append(p)
|
|
89
|
+
|
|
90
|
+
# Lazy state initialization
|
|
91
|
+
if len(state) == 0:
|
|
92
|
+
state["c_grad"] = []
|
|
93
|
+
|
|
94
|
+
grads.append(p.grad)
|
|
95
|
+
c_grads.append(state.get("c_grad"))
|
|
96
|
+
|
|
97
|
+
return has_sparse_grad
|
|
98
|
+
|
|
99
|
+
def __setstate__(self, state):
|
|
100
|
+
super().__setstate__(state)
|
|
101
|
+
|
|
102
|
+
def dual_step(self, i: int, c_val: Tensor = None):
|
|
103
|
+
r"""Save constraint gradient for weight update. To be called BEFORE :func:`step` in an iteration!
|
|
104
|
+
|
|
105
|
+
Args:
|
|
106
|
+
i (int): index of the constraint **(unused)**
|
|
107
|
+
c_val (Tensor): an estimate of the value of the constraint at which the gradient was computed **(unused)**
|
|
108
|
+
"""
|
|
109
|
+
|
|
110
|
+
if i > self.m:
|
|
111
|
+
raise ValueError("SSw does not support multiple constraints.")
|
|
112
|
+
|
|
113
|
+
# save constraint grad
|
|
114
|
+
for group in self.param_groups:
|
|
115
|
+
params: list[Tensor] = []
|
|
116
|
+
grads: list[Tensor] = []
|
|
117
|
+
c_grads: list[Tensor] = []
|
|
118
|
+
_ = self._init_group(group, params, grads, c_grads)
|
|
119
|
+
|
|
120
|
+
for p in group["params"]:
|
|
121
|
+
if p.grad is not None:
|
|
122
|
+
state = self.state[p]
|
|
123
|
+
# state['c_grad'] is cleaned in step()
|
|
124
|
+
# so it is always empty on dual_step()
|
|
125
|
+
state["c_grad"].append(p.grad)
|
|
126
|
+
|
|
127
|
+
@_use_grad_for_differentiable
|
|
128
|
+
def step(self, c_val: Union[Iterable | Tensor]):
|
|
129
|
+
r"""Perform an update of the primal parameters (network weights). To be called AFTER :func:`dual_step` in an iteration!
|
|
130
|
+
|
|
131
|
+
Args:
|
|
132
|
+
c_val (Tensor): an Iterable of estimates of values of **ALL** constraints; used for primal parameter update.
|
|
133
|
+
Ideally, must be evaluated on an independent sample from the one used in :func:`dual_step`
|
|
134
|
+
"""
|
|
135
|
+
|
|
136
|
+
# here assume c_val is a scalar
|
|
137
|
+
|
|
138
|
+
update_con = c_val > 0
|
|
139
|
+
|
|
140
|
+
for group in self.param_groups:
|
|
141
|
+
params: list[Tensor] = []
|
|
142
|
+
grads: list[Tensor] = []
|
|
143
|
+
c_grads: list[Tensor] = []
|
|
144
|
+
lr = group["lr"]
|
|
145
|
+
_ = self._init_group(group, params, grads, c_grads)
|
|
146
|
+
|
|
147
|
+
for i, param in enumerate(params):
|
|
148
|
+
|
|
149
|
+
if update_con:
|
|
150
|
+
param.add_(c_grads[i][0], alpha=-self.dual_lr)
|
|
151
|
+
else:
|
|
152
|
+
param.add_(param.grad, alpha=-lr)
|
|
153
|
+
|
|
154
|
+
if c_grads[i] is not None:
|
|
155
|
+
c_grads[i].clear()
|
|
@@ -0,0 +1,61 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
def net_params_to_tensor(
|
|
5
|
+
net: torch.nn.Module, flatten=False, copy=False
|
|
6
|
+
) -> torch.Tensor:
|
|
7
|
+
# flat_params = [ar.to_numpy(param) for param in net.parameters()]
|
|
8
|
+
if copy:
|
|
9
|
+
params = [param.detach().clone() for param in net.parameters()]
|
|
10
|
+
else:
|
|
11
|
+
params = [param for param in net.parameters()]
|
|
12
|
+
|
|
13
|
+
if flatten:
|
|
14
|
+
flat_params = [torch.flatten(param) for param in params]
|
|
15
|
+
return torch.concat(flat_params)
|
|
16
|
+
|
|
17
|
+
return params
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def check_same_sample(sample1, sample2):
|
|
21
|
+
s1w, s1nw = sample1[0], sample1[1]
|
|
22
|
+
s2w, s2nw = sample2[0], sample2[1]
|
|
23
|
+
|
|
24
|
+
s1wx, s1wy = s1w
|
|
25
|
+
s2wx, s2wy = s2w
|
|
26
|
+
|
|
27
|
+
s1nwx, s1nwy = s1nw
|
|
28
|
+
s2nwx, s2nwy = s2nw
|
|
29
|
+
|
|
30
|
+
return (
|
|
31
|
+
torch.all(s1wx == s2wx)
|
|
32
|
+
and torch.all(s1wy == s2wy)
|
|
33
|
+
and torch.all(s1nwx == s2nwx)
|
|
34
|
+
and torch.all(s1nwy == s2nwy)
|
|
35
|
+
)
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def net_grads_to_tensor(net, clip=False, flatten=True, device=None) -> torch.Tensor:
|
|
39
|
+
param_grads = []
|
|
40
|
+
if clip:
|
|
41
|
+
torch.nn.utils.clip_grad_norm_(net.parameters(), 0.5)
|
|
42
|
+
for param in net.parameters():
|
|
43
|
+
if param.grad is not None:
|
|
44
|
+
# Clone to avoid modifying the original tensor
|
|
45
|
+
device = param.grad.data.device if device is None else device
|
|
46
|
+
if flatten:
|
|
47
|
+
param_grads.append(param.grad.data.view(-1))
|
|
48
|
+
else:
|
|
49
|
+
param_grads.append(param.grad.data.to(device))
|
|
50
|
+
if flatten:
|
|
51
|
+
param_grads = torch.cat(param_grads)
|
|
52
|
+
return param_grads
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def _set_weights(net: torch.nn.Module, x):
|
|
56
|
+
start = 0
|
|
57
|
+
w = net_params_to_tensor(net, flatten=False, copy=False)
|
|
58
|
+
for i in range(len(w)):
|
|
59
|
+
end = start + w[i].numel()
|
|
60
|
+
w[i].set_(x[start:end].reshape(w[i].shape))
|
|
61
|
+
start = end
|
|
@@ -0,0 +1,87 @@
|
|
|
1
|
+
from typing import Callable, Iterable
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
import torch
|
|
5
|
+
from torch.utils.data import DataLoader, SubsetRandomSampler
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
def _make_dataloaders(dataset, group_indices, batch_size, device, drop_last, gen=None):
|
|
9
|
+
dataloaders = []
|
|
10
|
+
for idx in group_indices:
|
|
11
|
+
sampler = SubsetRandomSampler(idx, gen)
|
|
12
|
+
dataloaders.append(iter(DataLoader(dataset, batch_size, sampler=sampler, drop_last=True)))
|
|
13
|
+
return dataloaders
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class FairnessConstraint:
|
|
17
|
+
def __init__(
|
|
18
|
+
self,
|
|
19
|
+
dataset: torch.utils.data.Dataset,
|
|
20
|
+
group_indices: Iterable[Iterable[int]],
|
|
21
|
+
fn: Callable,
|
|
22
|
+
batch_size: int = None,
|
|
23
|
+
use_dataloaders=True,
|
|
24
|
+
device="cpu",
|
|
25
|
+
seed=None,
|
|
26
|
+
loader_drop_last=False,
|
|
27
|
+
):
|
|
28
|
+
self.dataset = dataset
|
|
29
|
+
self.group_sets = [
|
|
30
|
+
torch.utils.data.Subset(dataset, idx) for idx in group_indices
|
|
31
|
+
]
|
|
32
|
+
self._group_indices = group_indices
|
|
33
|
+
self.fn = fn
|
|
34
|
+
self._seed = seed
|
|
35
|
+
self._rng = np.random.default_rng(seed)
|
|
36
|
+
self._torch_rng = torch.manual_seed(seed) if seed is not None else torch.Generator(device=device)
|
|
37
|
+
self._device = device
|
|
38
|
+
self._drop_last = loader_drop_last
|
|
39
|
+
if batch_size is not None:
|
|
40
|
+
self._batch_size = batch_size
|
|
41
|
+
if use_dataloaders:
|
|
42
|
+
self.group_dataloaders = _make_dataloaders(
|
|
43
|
+
dataset, group_indices, batch_size, device, gen=self._torch_rng, drop_last=loader_drop_last
|
|
44
|
+
)
|
|
45
|
+
|
|
46
|
+
def group_sizes(self):
|
|
47
|
+
return [len(group) for group in self.group_sets]
|
|
48
|
+
|
|
49
|
+
def eval(self, net, sample, **kwargs):
|
|
50
|
+
return self.fn(net, sample, **kwargs)
|
|
51
|
+
|
|
52
|
+
def sample_loader(self):
|
|
53
|
+
samples = []
|
|
54
|
+
for i, l in enumerate(self.group_dataloaders):
|
|
55
|
+
try:
|
|
56
|
+
sample = next(l)
|
|
57
|
+
except StopIteration:
|
|
58
|
+
sampler = SubsetRandomSampler(self._group_indices[i], self._torch_rng)
|
|
59
|
+
l = iter(DataLoader(self.dataset, self._batch_size, sampler=sampler, drop_last=self._drop_last))
|
|
60
|
+
sample = next(l)
|
|
61
|
+
self.group_dataloaders[i] = l
|
|
62
|
+
|
|
63
|
+
samples.append(sample)
|
|
64
|
+
return samples
|
|
65
|
+
|
|
66
|
+
def sample_dataset(
|
|
67
|
+
self, N, rng: np.random.Generator = None, indices=None, return_indices=False
|
|
68
|
+
):
|
|
69
|
+
if rng is None:
|
|
70
|
+
rng = self._rng
|
|
71
|
+
|
|
72
|
+
if indices is None:
|
|
73
|
+
indices = []
|
|
74
|
+
# returns len(group) points if N > len(group)
|
|
75
|
+
for group in self.group_sets:
|
|
76
|
+
indices.append(
|
|
77
|
+
rng.choice(group.indices, N)
|
|
78
|
+
if N < len(group)
|
|
79
|
+
else rng.choice(group.indices, len(group))
|
|
80
|
+
)
|
|
81
|
+
|
|
82
|
+
sample = [self.dataset[indices[i]] for i, _ in enumerate(self.group_sets)]
|
|
83
|
+
|
|
84
|
+
if return_indices:
|
|
85
|
+
return sample, indices
|
|
86
|
+
else:
|
|
87
|
+
return sample
|
|
@@ -0,0 +1,118 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
from fairret.statistic import (
|
|
3
|
+
TruePositiveRate,
|
|
4
|
+
FalseNegativeFalsePositiveFraction,
|
|
5
|
+
Accuracy,
|
|
6
|
+
)
|
|
7
|
+
from fairret.loss import NormLoss
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def tpr_equality(_, net, c_data):
|
|
11
|
+
statistic = TruePositiveRate()
|
|
12
|
+
loss = NormLoss(statistic, p=2)
|
|
13
|
+
|
|
14
|
+
return fairret_stat_equality(net, c_data, loss)
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def ppv_equality(_, net, c_data):
|
|
18
|
+
statistic = FalseNegativeFalsePositiveFraction()
|
|
19
|
+
loss = NormLoss(statistic, p=2)
|
|
20
|
+
|
|
21
|
+
return fairret_stat_equality(net, c_data, loss)
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def acc_equality(_, net, c_data):
|
|
25
|
+
statistic = Accuracy()
|
|
26
|
+
loss = NormLoss(statistic, p=2)
|
|
27
|
+
|
|
28
|
+
return fairret_stat_equality(net, c_data, loss)
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def fairret_stat_equality(net, c_data, loss):
|
|
32
|
+
g1_inputs, g1_labels = c_data[0]
|
|
33
|
+
g2_inputs, g2_labels = c_data[1]
|
|
34
|
+
|
|
35
|
+
g1_outs = net(g1_inputs).squeeze()
|
|
36
|
+
g2_outs = net(g2_inputs).squeeze()
|
|
37
|
+
# if not (1 in a_labels or 1 in g2_labels):
|
|
38
|
+
# return torch.tensor(0)
|
|
39
|
+
|
|
40
|
+
group_codes = [0] * len(g1_labels) + [1] * len(g2_labels)
|
|
41
|
+
group_codes = torch.tensor(
|
|
42
|
+
[[0.0, 1.0] if x == 1 else [1.0, 0.0] for x in group_codes]
|
|
43
|
+
)
|
|
44
|
+
|
|
45
|
+
return loss(
|
|
46
|
+
torch.concat([g1_outs, g2_outs]).unsqueeze(1),
|
|
47
|
+
group_codes,
|
|
48
|
+
torch.concat([g1_labels, g2_labels]).unsqueeze(1),
|
|
49
|
+
)
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def dummy(_, net, c_data):
|
|
53
|
+
r = torch.zeros(1)
|
|
54
|
+
r.grad = 0
|
|
55
|
+
return r
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
def loss_equality(loss, net, c_data):
|
|
59
|
+
g1_inputs, g1_labels = c_data[0]
|
|
60
|
+
g2_inputs, g2_labels = c_data[1]
|
|
61
|
+
g1_outs = net(g1_inputs)
|
|
62
|
+
if g1_labels.ndim == 0:
|
|
63
|
+
g1_labels = g1_labels.reshape(1)
|
|
64
|
+
g2_labels = g2_labels.reshape(1)
|
|
65
|
+
if g1_labels.ndim < g1_outs.ndim:
|
|
66
|
+
g1_labels = g1_labels.unsqueeze(1)
|
|
67
|
+
g2_labels = g2_labels.unsqueeze(1)
|
|
68
|
+
g1_loss = loss(g1_outs, g1_labels)
|
|
69
|
+
g2_outs = net(g2_inputs)
|
|
70
|
+
g2_loss = loss(g2_outs, g2_labels)
|
|
71
|
+
|
|
72
|
+
val = g1_loss - g2_loss
|
|
73
|
+
return val
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
def abs_loss_equality(loss, net, c_data):
|
|
77
|
+
g1_inputs, g1_labels = c_data[0]
|
|
78
|
+
g2_inputs, g2_labels = c_data[1]
|
|
79
|
+
g1_outs = net(g1_inputs)
|
|
80
|
+
if g1_labels.ndim == 0:
|
|
81
|
+
g1_labels = g1_labels.reshape(1)
|
|
82
|
+
g2_labels = g2_labels.reshape(1)
|
|
83
|
+
if g1_labels.ndim < g1_outs.ndim:
|
|
84
|
+
g1_labels = g1_labels.unsqueeze(1)
|
|
85
|
+
g2_labels = g2_labels.unsqueeze(1)
|
|
86
|
+
g1_loss = loss(g1_outs, g1_labels)
|
|
87
|
+
g2_outs = net(g2_inputs)
|
|
88
|
+
g2_loss = loss(g2_outs, g2_labels)
|
|
89
|
+
|
|
90
|
+
val = g1_loss - g2_loss
|
|
91
|
+
return torch.abs(val)
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
def fairret_constr(loss, net, c_data):
|
|
95
|
+
g1_inputs, g1_labels = c_data[0]
|
|
96
|
+
g2_inputs, g2_labels = c_data[1]
|
|
97
|
+
g1_logits = net(g1_inputs)
|
|
98
|
+
g2_logits = net(g2_inputs)
|
|
99
|
+
g1_onehot = torch.tensor([[0.0, 1.0]] * len(g1_inputs))
|
|
100
|
+
g2_onehot = torch.tensor([[1.0, 0.0]] * len(g2_inputs))
|
|
101
|
+
logits = torch.concat([g1_logits, g2_logits])
|
|
102
|
+
sens = torch.vstack([g1_onehot, g2_onehot])
|
|
103
|
+
labels = torch.hstack([g1_labels, g2_labels]).unsqueeze(1)
|
|
104
|
+
|
|
105
|
+
return loss(logits, sens, label=labels)
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def fairret_pr_constr(loss, net, c_data):
|
|
109
|
+
g1_inputs, _ = c_data[0]
|
|
110
|
+
g2_inputs, _ = c_data[1]
|
|
111
|
+
g1_logits = net(g1_inputs)
|
|
112
|
+
g2_logits = net(g2_inputs)
|
|
113
|
+
g1_onehot = torch.tensor([[0.0, 1.0]] * len(g1_inputs))
|
|
114
|
+
g2_onehot = torch.tensor([[1.0, 0.0]] * len(g2_inputs))
|
|
115
|
+
logits = torch.concat([g1_logits, g2_logits])
|
|
116
|
+
sens = torch.vstack([g1_onehot, g2_onehot])
|
|
117
|
+
|
|
118
|
+
return loss(logits, sens)
|
|
File without changes
|
|
@@ -0,0 +1,15 @@
|
|
|
1
|
+
from .constraint import FairnessConstraint
|
|
2
|
+
from .constraint_fns import (
|
|
3
|
+
fairret_stat_equality,
|
|
4
|
+
ppv_equality,
|
|
5
|
+
acc_equality,
|
|
6
|
+
tpr_equality,
|
|
7
|
+
abs_loss_equality,
|
|
8
|
+
loss_equality,
|
|
9
|
+
abs_diff_tpr,
|
|
10
|
+
abs_diff_fpr,
|
|
11
|
+
abs_max_dev_from_overall_tpr,
|
|
12
|
+
abs_diff_pr
|
|
13
|
+
)
|
|
14
|
+
|
|
15
|
+
__all__ = ["FairnessConstraint, loss_equality"]
|