-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathplot_by_attr_non_uniform.py
More file actions
206 lines (182 loc) · 9.53 KB
/
Copy pathplot_by_attr_non_uniform.py
File metadata and controls
206 lines (182 loc) · 9.53 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
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
import cmd
import sys
import os
import csv
import math
import numpy as np
import matplotlib
import matplotlib.pyplot as plt
matplotlib.style.use('ggplot')
import pandas as pd
from pandas import DataFrame, Series
class CSVReaderConst(object):
HIGHEST_RISK = 10
LOWEST_RISK = 1
RACES_TO_CORRECT = ["African American", "White"]
RECIDIVISM_COL_NAME = "two_year_recid"
class DataAnalyzer(object):
def __init__(self, filepath_in):
self.plot_filepath = filepath_in
self.df = pd.read_csv(filepath_in)
def trait_breakdown(self, col_name):
breakdown = self.df[col_name.lower()].value_counts(sort=True, ascending=False)
print(breakdown)
return breakdown
def get_trait_key(self, trait, score):
return "{0!s}_{1!s}".format(trait, score)
def correct_for(self, col_name, recid_dec_col_name, traits=[], rms=False):
'''
Across same col_name, correct the attribute for each trait to remove bias
'''
if traits == []:
trait_dict = self.trait_breakdown(col_name=col_name)
for trait, percentage in trait_dict.iteritems():
traits.append(trait)
print(traits)
baseline_error_dict = {}
baseline_abs_error_dict = {}
residual_sq_error_dict = {}
rms_error_dict = {}
people_per_trait = {}
baseline_bias_dict = {}
for trait in traits:
group = self.df[self.df[col_name] == trait]
# print(group)
num_members = len(group)
for index, person in group.iterrows():
trait_key = self.get_trait_key(trait, int(person[recid_dec_col_name]))
if int(person[CSVReaderConst.RECIDIVISM_COL_NAME]) == 1:
person_error = float(person[recid_dec_col_name]) - CSVReaderConst.HIGHEST_RISK
else:
person_error = float(person[recid_dec_col_name]) - CSVReaderConst.LOWEST_RISK
if baseline_error_dict.get(trait_key, None) is None:
baseline_error_dict[trait_key] = person_error
baseline_abs_error_dict[trait_key] = abs(person_error)
residual_sq_error_dict[trait_key] = pow(person_error, 2)
people_per_trait[trait_key] = 1
else:
baseline_error_dict[trait_key] += person_error
baseline_abs_error_dict[trait_key] += abs(person_error)
residual_sq_error_dict[trait_key] += pow(person_error, 2)
people_per_trait[trait_key] += 1
if num_members == 0:
raise ValueError("No members found in group {0!s}".format(trait))
for trait_key, error in baseline_error_dict.items():
baseline_bias_dict[trait_key] = float(error)/float(people_per_trait[trait_key])
rms_error_dict[trait_key] = math.sqrt(float(residual_sq_error_dict[trait_key])/float(people_per_trait[trait_key]))
# print(total_error)
if rms:
print("For group {0!s}, root mean squared error: {1:.3f}, baseline bias: {2:.3f}".format(trait_key,
rms_error_dict[trait_key],
baseline_bias_dict[trait_key]))
else:
print("For group {0!s}, baseline absolute error: {1:.3f}, baseline bias: {2:.3f}".format(trait_key,
baseline_abs_error_dict[trait_key],
baseline_bias_dict[trait_key]))
print("=========================================")
new_error_dict = {}
new_abs_error_dict = {}
new_rms_error_dict = {}
new_residual_sq_error_dict = {}
new_bias_dict = {}
t_err = 0
for trait in traits:
group = self.df[self.df[col_name] == trait]
# print(group)
num_members = len(group)
for index, person in group.iterrows():
trait_key = self.get_trait_key(trait, int(person[recid_dec_col_name]))
corrected_decile = float(person[recid_dec_col_name]) - float(baseline_bias_dict[trait_key])
# print("corrected_decile {0!s}".format(corrected_decile))
if int(person[CSVReaderConst.RECIDIVISM_COL_NAME]) == 1:
person_error = float(corrected_decile) - CSVReaderConst.HIGHEST_RISK
else:
person_error = float(corrected_decile) - CSVReaderConst.LOWEST_RISK
t_err += person_error
if new_error_dict.get(trait_key, None) is None:
new_error_dict[trait_key] = person_error
new_abs_error_dict[trait_key] = abs(person_error)
new_residual_sq_error_dict[trait_key] = pow(person_error, 2)
else:
new_error_dict[trait_key] += person_error
new_abs_error_dict[trait_key] += abs(person_error)
new_residual_sq_error_dict[trait_key] += pow(person_error, 2)
if num_members == 0:
raise ValueError("No members found in group {0!s}".format(trait))
print("t_err: {0!s}".format(t_err))
for trait_key, error in new_error_dict.items():
new_bias_dict[trait_key] = float(error)/float(people_per_trait[trait_key])
new_rms_error_dict[trait_key] = math.sqrt(float(new_residual_sq_error_dict[trait_key])/float(people_per_trait[trait_key]))
if rms:
print("For group {0!s}, new root mean squared error: {1:.3f}, new baseline bias: {2:.3f}".format(trait_key,
new_rms_error_dict[trait_key],
new_bias_dict[trait_key]))
else:
print("For group {0!s}, new absolute error: {1:.3f}, new baseline bias: {2:.3f}".format(trait_key,
new_abs_error_dict[trait_key],
new_bias_dict[trait_key]))
baseline_errors = []
new_errors = []
trait_labels = []
for trait_key in sorted(baseline_abs_error_dict.iterkeys()):
if rms:
baseline_errors.append(rms_error_dict[trait_key])
new_errors.append(new_rms_error_dict[trait_key])
else:
baseline_errors.append(baseline_abs_error_dict[trait_key])
new_errors.append(new_abs_error_dict[trait_key])
trait_labels.append(trait_key)
num_items = len(baseline_error_dict)
ind = np.arange(num_items)
width = 0.35
fig, ax = plt.subplots()
rects1 = ax.bar(ind, baseline_errors, width, color='r')
rects2 = ax.bar(ind + width, new_errors, width, color='y')
if rms:
ax.set_ylabel("Root Mean Squared Errors")
else:
ax.set_ylabel("Absolute Errors")
ax.set_xticks(ind + width)
ax.set_xticklabels(trait_labels)
ax.legend( (rects1[0], rects2[0]), ('Baseline', 'Corrected') )
plt.draw()
plt.pause(0.001)
return baseline_error_dict, baseline_bias_dict, new_error_dict
class AnalyzerShell(cmd.Cmd):
intro = 'Welcome to the analyzer shell. Type help or ? to list commands.\n'
prompt = '(analyzer) '
def setup(self, file_path):
self.data_analyzer = DataAnalyzer(file_path)
def do_trait_breakdown(self, arg):
'Get a percentage and count breakdown of specified argument'
self.data_analyzer.trait_breakdown(col_name=arg)
def do_plot_recid(self, arg):
'Plot the value of recidivism decile, with a stacked chart of how many actually had recidivism for this attribute.\n \
i.e. plot_recid race Caucasian decile_score'
split_up = arg.split(" ")
self.data_analyzer.plot_recid(*split_up)
def do_correct_for(self, arg, calc_rms=False):
'Correct a particular decile score attribute based on a specific column.\nSpecificy traits in the column to correct, or "ALL" for an analysis of all. \
\ni.e. correct_for decile_score race African-American, Caucasian OR correct_for decile_score race ALL'
split_up = arg.split(" ", 2)
# print(split_up)
dec_name = split_up[0]
col_name = split_up[1]
traits = split_up[2].split(", ")
if (traits[0] == "ALL"):
traits = []
self.data_analyzer.correct_for(col_name=col_name, recid_dec_col_name=dec_name, traits=traits, rms=calc_rms)
def do_correct_for_rms(self, arg):
'Same as correct_for, but using root mean squared error instead of linear error calculation'
self.do_correct_for(arg=arg, calc_rms=True)
def do_quit(self, arg):
'Quit'
print('Thank you for using analyzer')
return True
if __name__ == '__main__':
if len(sys.argv) < 2:
print("Need filepath")
else:
shell = AnalyzerShell()
shell.setup(os.path.normpath(sys.argv[1]))
shell.cmdloop()