Commit 7db1e3e2 authored by opitz's avatar opitz
Browse files

add metrics

parent db29fb79
Loading
Loading
Loading
Loading

src/metrics.py

0 → 100644
+96 −0
Changes for src/metrics.py: 96 added lines, 0 removed lines.
Original line number Diff line number Diff line
import numpy as np

from sklearn.metrics import roc_auc_score
from sklearn.metrics import accuracy_score
from sklearn.metrics import classification_report
from sklearn.metrics import f1_score


"""
def predict(model,X):
    _, p_embedding,n_embedding,a_embedding = model.predict(X)
    #scores_p = np.dot(p_embedding,a_embedding.T)
    #scores_n = np.dot(p_embedding,a_embedding.T)
    scores_p = np.einsum("ij,ij -> i",p_embedding,a_embedding)
    scores_n = np.einsum("ij,ij -> i",n_embedding,a_embedding)
    return np.vstack((scores_p,scores_n))

"""
def predict_nli(model,X):
    _, p_embedding,n_embedding= model.predict(X)
    
    scores_p = np.einsum("ij -> i",p_embedding)
    scores_n = np.einsum("ij -> i",n_embedding)
    return np.vstack((scores_p,scores_n))
"""
def predict_nli_naive(model,X):
    scores = model.predict(X) 
    s = []
    for i in range(int(len(scores)/2)):
        s.append([scores[i],scores[i+int(len(scores)/2.00)]])
    return np.array(s)

def get_accuracy_naive(model,X,y):
    scores = predict_nli_naive(model,X)
    nodecision = [i for i in range(len(scores)) if abs(scores[i][0] - scores[i][1]) < 0.00001]
    correct = [i for i in range(len(scores)) if scores[i][0] > scores[i][1] and i not in nodecision]
    incorrect = [i for i in range(len(scores)) if scores[i][0] < scores[i][1] and i not in nodecision]
    n_c = len(correct)+len(nodecision)/2.00
    n_i = len(incorrect) +len(nodecision)/2.00
    return n_c/(n_c+n_i)

def get_accuracy(model,X):
    scores = predict_nli(model,X).T
    nodecision = [i for i in range(len(scores)) if abs(scores[i][0] - scores[i][1]) < 0.00001]
    correct = [i for i in range(len(scores)) if scores[i][0] > scores[i][1] and i not in nodecision]
    incorrect = [i for i in range(len(scores)) if scores[i][0] < scores[i][1] and i not in nodecision]
    n_c = len(correct)+len(nodecision)/2.00
    n_i = len(incorrect) +len(nodecision)/2.00
    return n_c/(n_c+n_i)
"""
def get_classif_report(model,X,labels,bin_labels=["support","attack"]):
    if type(model) != list:
        scores = predict_nli(model,X).T
    else:
        scores = predict_nli(model[0],X).T

        for m in model[1:]:
            print(scores[:10])
            scores+=predict_nli(m,X).T
    nodecision = [i for i in range(len(scores)) if abs(scores[i][0] - scores[i][1]) < 0.00001]
    print("nodec",len(nodecision),set(labels))
    correct = [i for i in range(len(scores)) if scores[i][0] > scores[i][1] and i not in nodecision]
    incorrect = [i for i in range(len(scores)) if scores[i][0] < scores[i][1] and i not in nodecision]
    preds = []
    for i,la in enumerate(labels):
        if i in correct:
            preds.append(labels[i])
        else:
            preds.append([ l for l in bin_labels if l != la][0])
    print(classification_report(labels,preds))
    return f1_score(labels,preds,average="macro")


#this is the functio we use
def get_classif_report_scores(scores,X=None,labels=None,bin_labels=["support","attack"],printreport=False,evalfun=f1_score,pos_label=None):
    nodecision = [i for i in range(len(scores)) if abs(scores[i][0] - scores[i][1]) < 0.00001]
    
    #correct are examples where scores[0] > scores[1] 
    #(two instances processed by same simaese network, 
    #the correct one is the first, goal: first achieves higher score than second)
    correct = [i for i in range(len(scores)) if scores[i][0] > scores[i][1] and i not in nodecision]
    incorrect = [i for i in range(len(scores)) if scores[i][0] < scores[i][1] and i not in nodecision]
    preds = []
    for i,la in enumerate(labels):
        # if correct label achieves higher score take true label
        if i in correct:
            preds.append(labels[i])
        # if false label achieves higher score take opposite label
        else:
            preds.append([ l for l in bin_labels if l != la][0])
    if printreport:
        print(classification_report(labels,preds))
    if pos_label:
        return evalfun(labels,preds,pos_label=pos_label)

    return evalfun(labels,preds,average="macro")