CoolFace
Apppublic

RabbitRUI/ruispace

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
plot.py73 linesDownload Raw Back to utils
1# coding: utf-82 3import os4from pathlib import Path5 6import matplotlib.pyplot as plt7import numpy as np8import pandas as pd9from menpo.visualize.viewmatplotlib import sample_colours_from_colourmap10from prettytable import PrettyTable11from sklearn.metrics import roc_curve, auc12 13image_path = "/data/anxiang/IJB_release/IJBC"14files = [15        "./ms1mv3_arcface_r100/ms1mv3_arcface_r100/ijbc.npy"16]17 18 19def read_template_pair_list(path):20    pairs = pd.read_csv(path, sep=' ', header=None).values21    t1 = pairs[:, 0].astype(np.int)22    t2 = pairs[:, 1].astype(np.int)23    label = pairs[:, 2].astype(np.int)24    return t1, t2, label25 26 27p1, p2, label = read_template_pair_list(28    os.path.join('%s/meta' % image_path,29                 '%s_template_pair_label.txt' % 'ijbc'))30 31methods = []32scores = []33for file in files:34    methods.append(file.split('/')[-2])35    scores.append(np.load(file))36 37methods = np.array(methods)38scores = dict(zip(methods, scores))39colours = dict(40    zip(methods, sample_colours_from_colourmap(methods.shape[0], 'Set2')))41x_labels = [10 ** -6, 10 ** -5, 10 ** -4, 10 ** -3, 10 ** -2, 10 ** -1]42tpr_fpr_table = PrettyTable(['Methods'] + [str(x) for x in x_labels])43fig = plt.figure()44for method in methods:45    fpr, tpr, _ = roc_curve(label, scores[method])46    roc_auc = auc(fpr, tpr)47    fpr = np.flipud(fpr)48    tpr = np.flipud(tpr)  # select largest tpr at same fpr49    plt.plot(fpr,50             tpr,51             color=colours[method],52             lw=1,53             label=('[%s (AUC = %0.4f %%)]' %54                    (method.split('-')[-1], roc_auc * 100)))55    tpr_fpr_row = []56    tpr_fpr_row.append("%s-%s" % (method, "IJBC"))57    for fpr_iter in np.arange(len(x_labels)):58        _, min_index = min(59            list(zip(abs(fpr - x_labels[fpr_iter]), range(len(fpr)))))60        tpr_fpr_row.append('%.2f' % (tpr[min_index] * 100))61    tpr_fpr_table.add_row(tpr_fpr_row)62plt.xlim([10 ** -6, 0.1])63plt.ylim([0.3, 1.0])64plt.grid(linestyle='--', linewidth=1)65plt.xticks(x_labels)66plt.yticks(np.linspace(0.3, 1.0, 8, endpoint=True))67plt.xscale('log')68plt.xlabel('False Positive Rate')69plt.ylabel('True Positive Rate')70plt.title('ROC on IJB')71plt.legend(loc="lower right")72print(tpr_fpr_table)73