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,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,11 @@
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
+ )
10
+
11
+ __all__ = ["FairnessConstraint, loss_equality"]
@@ -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"]