-
Notifications
You must be signed in to change notification settings - Fork 6
Expand file tree
/
Copy pathconstraint.py
More file actions
117 lines (105 loc) · 4.86 KB
/
Copy pathconstraint.py
File metadata and controls
117 lines (105 loc) · 4.86 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
import torch
from torch import nn
from torch.nn import functional as F
class ConstraintLoss(nn.Module):
def __init__(self, n_class=2, alpha=1, p_norm=2):
super(ConstraintLoss, self).__init__()
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
self.alpha = alpha
self.p_norm = p_norm
self.n_class = n_class
self.n_constraints = 2
self.dim_condition = self.n_class + 1
self.M = torch.zeros((self.n_constraints, self.dim_condition))
self.c = torch.zeros(self.n_constraints)
def mu_f(self, X=None, y=None, sensitive=None):
return torch.zeros(self.n_constraints)
def forward(self, X, out, sensitive, y=None):
sensitive = sensitive.view(out.shape)
if isinstance(y, torch.Tensor):
y = y.view(out.shape)
out = torch.sigmoid(out)
mu = self.mu_f(X=X, out=out, sensitive=sensitive, y=y)
gap_constraint = F.relu(
torch.mv(self.M.to(self.device), mu.to(self.device)) - self.c.to(self.device)
)
if self.p_norm == 2:
cons = self.alpha * torch.dot(gap_constraint, gap_constraint)
else:
cons = self.alpha * torch.dot(gap_constraint.detach(), gap_constraint)
return cons
class DemographicParityLoss(ConstraintLoss):
def __init__(self, sensitive_classes=[0, 1], alpha=1, p_norm=2):
"""loss of demograpfhic parity
Args:
sensitive_classes (list, optional): list of unique values of sensitive attribute. Defaults to [0, 1].
alpha (int, optional): [description]. Defaults to 1.
p_norm (int, optional): [description]. Defaults to 2.
"""
self.sensitive_classes = sensitive_classes
self.n_class = len(sensitive_classes)
super(DemographicParityLoss, self).__init__(
n_class=self.n_class, alpha=alpha, p_norm=p_norm
)
self.n_constraints = 2 * self.n_class
self.dim_condition = self.n_class + 1
self.M = torch.zeros((self.n_constraints, self.dim_condition))
for i in range(self.n_constraints):
j = i % 2
if j == 0:
self.M[i, j] = 1.0
self.M[i, -1] = -1.0
else:
self.M[i, j - 1] = -1.0
self.M[i, -1] = 1.0
self.c = torch.zeros(self.n_constraints)
def mu_f(self, X, out, sensitive, y=None):
expected_values_list = []
for v in self.sensitive_classes:
idx_true = sensitive == v # torch.bool
expected_values_list.append(out[idx_true].mean())
#("bismillah")
expected_values_list.append(out.mean())
return torch.stack(expected_values_list)
def forward(self, X, out, sensitive, y=None):
return super(DemographicParityLoss, self).forward(X, out, sensitive)
class AverageTreatmentEffectLoss(ConstraintLoss):
def __init__(self, sensitive_classes=[0, 1], alpha=1, p_norm=2):
self.sensitive_classes = sensitive_classes
self.y_classes = [1] # only consider positive outcome
self.n_class = len(sensitive_classes)
self.n_y_class = len(self.y_classes)
super(EqualOpportunityLoss, self).__init__(n_class=self.n_class, alpha=alpha, p_norm=p_norm)
self.n_constraints = self.n_class * self.n_y_class * 2
self.dim_condition = self.n_y_class * (self.n_class + 1)
self.M = torch.zeros((self.n_constraints, self.dim_condition))
self.c = torch.zeros(self.n_constraints)
element_K_A = self.sensitive_classes + [None]
for i_a, a_0 in enumerate(self.sensitive_classes):
for i_y, y_0 in enumerate(self.y_classes):
for i_s, s in enumerate([-1, 1]):
for j_y, y_1 in enumerate(self.y_classes):
for j_a, a_1 in enumerate(element_K_A):
i = i_a * (2 * self.n_y_class) + i_y * 2 + i_s
j = j_y + self.n_y_class * j_a
self.M[i, j] = self.__element_M(a_0, a_1, y_1, y_1, s)
def __element_M(self, a0, a1, y0, y1, s):
if a0 is None or a1 is None:
x = y0 == y1
return -1 * s * x
else:
x = (a0 == a1) & (y0 == y1)
return s * float(x)
def mu_f(self, X, out, sensitive, y):
expected_values_list = []
for u in self.sensitive_classes:
for v in self.y_classes:
idx_true = (y == v) * (sensitive == u)
expected_values_list.append(out[idx_true].mean())
# sensitive is star
for v in self.y_classes:
idx_true = y == v
expected_values_list.append(out[idx_true].mean())
return torch.stack(expected_values_list)
def forward(self, X, out, sensitive, y):
return super(EqualOpportunityLoss, self).forward(X, out, sensitive, y=y)