File size: 9,826 Bytes
178d33b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
import csv
import os

import numpy as np


def make_args_list(benchmarks, methods, metrics, benchmark_dict):
    args_list = []
    for metric in metrics:
        for benchmark in set(benchmarks) & set(benchmark_dict[metric]):
            for method in methods:
                args_list.append([benchmark, method, metric])
    return args_list


def write_metric(args, folder_list, save_line_dict, benchmark_dict):

    metric_list = [
        'fpr95', 'auroc', 'aupr_in', 'aupr_out', 'ccr_4', 'ccr_3', 'ccr_2',
        'ccr_1', 'acc'
    ]
    save_list = []
    for metric in args.metric2save:
        save_list.append(metric_list.index(metric) + 1)

    for metric in args.metrics:
        if metric == 'ood':
            for benchmark in set(args.benchmarks) & set(
                    benchmark_dict[metric]):
                args_list = make_args_list([benchmark], args.methods, ['ood'],
                                           benchmark_dict)
                sub_form_content = []
                for key_param in args_list:
                    for folder in folder_list:
                        key_folder = folder.split('_')
                        if all(key in key_folder for key in key_param):
                            target_folder = folder
                            break
                    else:
                        print("No respective folder path, something's wrong.")
                        raise FileNotFoundError
                        # quit()

                    with open(
                            os.path.join(args.output_dir, target_folder,
                                         'ood.csv'), 'r') as f:
                        lines = f.readlines()[save_line_dict[key_param[-1]]:]
                    sub_line_content = {}
                    sub_line_content['method/{}'.format(
                        args.metric2save)] = key_param[1]
                    for line in lines:
                        split = line.split(',')
                        content = ''
                        for metric in save_list:
                            content = content + '{:.2f}'.format(
                                float(split[metric])) + ' / '
                        else:
                            content = content[:-3]
                        # use method name as key
                        sub_line_content[split[0]] = content
                    sub_form_content.append(sub_line_content)
                csv_path = os.path.join(args.output_dir,
                                        '{}_ood.csv'.format(key_param[0]))
                with open(csv_path, 'w', newline='') as csvfile:
                    fieldnames = order_fieldnames(
                        list(sub_form_content[0].keys()), args)
                    writer = csv.DictWriter(csvfile, fieldnames=fieldnames)
                    writer.writeheader()
                    for sub_line_content in sub_form_content:
                        writer.writerow(sub_line_content)

        elif metric == 'osr':
            sub_form_content = []
            for method in args.methods:
                args_list = make_args_list(args.benchmarks, [method], ['osr'],
                                           benchmark_dict)
                sub_line_content = {}

                for key_param in args_list:
                    sub_line_content['method/{}'.format(
                        args.metric2save)] = key_param[1]
                    target_folder = []
                    seeds = ['seed1', 'seed2', 'seed3', 'seed4', 'seed5']
                    for seed in seeds:
                        key_param.append(seed)
                        for folder in folder_list:
                            key_folder = folder.split('_')
                            if all(key in key_folder for key in key_param):
                                target_folder.append(folder)
                                break
                        else:
                            print(
                                "No respective folder path, something's wrong."
                            )
                            raise FileNotFoundError
                            # quit()
                        key_param.pop(-1)

                    temp = np.ndarray(shape=(len(seeds), len(save_list)))
                    for i, folder in enumerate(target_folder):
                        with open(
                                os.path.join(args.output_dir, folder,
                                             'ood.csv'), 'r') as f:
                            lines = f.readlines(
                            )[save_line_dict[key_param[-1]]:]
                        for line in lines:
                            split = line.split(',')
                            for j, metric_index in enumerate(save_list):
                                temp[i][j] = split[metric_index]
                    content = ''
                    for item in np.mean(temp, axis=0):
                        content = content + '{:.2f}'.format(item) + ' / '
                    else:
                        content = content[:-3]

                    sub_line_content[key_param[0]] = content
                sub_form_content.append(sub_line_content)

            csv_path = os.path.join(args.output_dir, 'total_osr.csv')
            with open(csv_path, 'w', newline='') as csvfile:
                fieldnames = order_fieldnames(list(sub_form_content[0].keys()),
                                              args)
                writer = csv.DictWriter(csvfile, fieldnames=fieldnames)
                writer.writeheader()
                for sub_line_content in sub_form_content:
                    writer.writerow(sub_line_content)


def write_total(args, folder_list, save_line_dict, benchmark_dict,
                main_content_extract_dict):
    main_form_content = []
    for method in args.methods:
        main_line_content = {}
        for metric in args.metrics:
            args_list = make_args_list(args.benchmarks, [method], [metric],
                                       benchmark_dict)
            for key_param in args_list:
                main_line_content['method --> auroc'] = key_param[1]

                if metric == 'ood':
                    for folder in folder_list:
                        key_folder = folder.split('_')
                        if all(key in key_folder for key in key_param):
                            target_folder = folder
                            break
                    else:
                        print("No respective folder path, something's wrong.")
                        # quit()

                    with open(
                            os.path.join(args.output_dir, target_folder,
                                         'ood.csv'), 'r') as f:
                        lines = f.readlines()[save_line_dict[key_param[-1]]:]

                    content = ''
                    for line in lines:
                        if line.split(',')[0] in main_content_extract_dict[
                                key_param[-1]]:

                            # take auroc only
                            content = content + '{:.2f}'.format(
                                float(line.split(',')[2])) + ' / '
                    else:
                        content = content[:-3]
                    # use benchmark name as key
                    main_line_content[key_param[0]] = content

                if metric == 'osr':
                    target_folder = []
                    seeds = ['seed1', 'seed2', 'seed3', 'seed4', 'seed5']
                    for seed in seeds:
                        key_param.append(seed)
                        for folder in folder_list:
                            key_folder = folder.split('_')
                            if all(key in key_folder for key in key_param):
                                target_folder.append(folder)
                                break
                        else:
                            print(
                                "No respective folder path, something's wrong."
                            )
                            # quit()
                        key_param.pop(-1)

                    temp = np.ndarray(shape=(len(seeds), 1))
                    for i, folder in enumerate(target_folder):
                        with open(
                                os.path.join(args.output_dir, folder,
                                             'ood.csv'), 'r') as f:
                            lines = f.readlines(
                            )[save_line_dict[key_param[-1]]:]
                        for line in lines:
                            split = line.split(',')
                            temp[i] = split[2]
                    content = '{:.2f}'.format(np.mean(temp, axis=0).item())
                    main_line_content[key_param[0]] = content

        main_form_content.append(main_line_content)

    csv_path = os.path.join(args.output_dir, 'total_result.csv')
    with open(csv_path, 'w', newline='') as csvfile:
        fieldnames = order_fieldnames(list(main_form_content[0].keys()), args)
        writer = csv.DictWriter(csvfile, fieldnames=fieldnames)
        writer.writeheader()
        for main_line_content in main_form_content:
            writer.writerow(main_line_content)


verify_dir = './results/total'
for folder in os.listdir(verify_dir):
    if os.path.isdir(os.path.join(verify_dir, folder)):
        if 'ood.csv' not in os.listdir(os.path.join(verify_dir, folder)):
            # if 'seed1' in folder.split('_'):
            print(folder)


def order_fieldnames(keys, args):

    ordered_keys = []
    ordered_keys.append(keys[0])
    for item in args.benchmarks:
        if item in keys:
            ordered_keys.append(item)

    return ordered_keys