022/Deep-ECG
/Deep-ECG/src/DeepECG/utils.py
import os
import torch
from torch import nn
import numpy as np
import pandas as pd
from sklearn.metrics import average_precision_score
from sklearn.metrics import roc_auc_score
def compute_auc(y_true, y_score):
return roc_auc_score(y_true, y_score)
def compute_average_precision(y_true, y_score):
return average_precision_score(y_true, y_score)
def save_checkpoint(state, is_best, filename="checkpoint.pth.tar", best_model_file='model_best.pth.tar'):
torch.save(state, filename)
if is_best:
torch.save(state, best_model_file)
def save_checkpoint(state, is_best, filename="checkpoint.pth.tar", best_model_file='model_best.pth.tar'):
torch.save(state, filename)
if is_best:
torch.save(state, best_model_file)
def get_best_model_file(best_model_file='model_best.pth.tar'):
if os.path.exists(best_model_file):
return best_model_file
else:
return 'checkpoint.pth.tar'
def get_best_model(best_model_file='model_best.pth.tar'):
if os.path.exists(best_model_file):
return torch.load(best_model_file)
else:
return torch.load(get_best_model_file(best_model_file))
def get_best_metrics(best_model_file='model_best.pth.tar'):
best_model = get_best_model(best_model_file)
return best_model['metrics']
def get_checkpoint(checkpoint_file="checkpoint.pth.tar"):
if os.path.exists(checkpoint_file):
return torch.load(checkpoint_file)
else:
return None
def get_best_metrics(checkpoint_file="checkpoint.pth.tar"):
checkpoint = get_checkpoint(checkpoint_file)
return checkpoint['metrics']
def get_best_model(checkpoint_file="checkpoint.pth.tar"):
checkpoint = get_checkpoint(checkpoint_file)
return checkpoint['best_model']
def get_best_metrics_from_path(path):
if os.path.isfile(path):
return get_best_metrics(path)
else:
file_name = os.listdir(path)
file_name.sort()
best_model_file = os.path.join(path, file_name[-1])
return get_best_metrics(best_model_file)
def get_checkpoint_from_path(path):
if os.path.isfile(path):
return get_checkpoint(path)
else:
file_name = os.listdir(path)
file_name.sort()
checkpoint_file = os.path.join(path, file_name[-1])
return get_checkpoint(checkpoint_file)
def get_best_model_from_path(path):
checkpoint = get_checkpoint_from_path(path)
return checkpoint['best_model']
def get_best_model_from_path(path):
checkpoint = get_checkpoint_from_path(path)
return checkpoint['best_model']
def load_model(model, model_file, optimizer=None, epoch=-1):
checkpoint = get_checkpoint(model_file)
state_dict = checkpoint['state_dict']
model.load_state_dict(state_dict)
if optimizer is not None:
if epoch == -1:
optimizer.load_state_dict(checkpoint['optimizer'])
else:
for i in range(epoch + 1):
optimizer.load_state_dict(checkpoint['optimizer'][i])
def load_model(model, model_file, optimizer=None, epoch=-1):
checkpoint = get_checkpoint(model_file)
state_dict = checkpoint['state_dict']
model.load_state_dict(state_dict)
if optimizer is not None:
if epoch == -1:
optimizer.load_state_dict(checkpoint['optimizer'])
else:
for i in range(epoch + 1):
optimizer.load_state_dict(checkpoint['optimizer'][i])
def load_model(model, model_file, optimizer=None, epoch=-1):
checkpoint = get_checkpoint(model_file)
state_dict = checkpoint['state_dict']
model.load_state_dict(state_dict)
if optimizer is not None:
if epoch == -1:
optimizer.load_state_dict(checkpoint['optimizer'])
else:
for i in range(epoch + 1):
optimizer.load_state_dict(checkpoint['optimizer'][i])
def load_model_with_best_metrics(model, best_model_file):
checkpoint = get_best_model(best_model_file)
state_dict = checkpoint['state_dict']
model.load_state_dict(state_dict)
optimizer = checkpoint['optimizer']
return model, optimizer
def load_model(model, model_file, optimizer=None, epoch=-1):
checkpoint = get_checkpoint(model_file)
state_dict = checkpoint['state_dict']
model.load_state_dict(state_dict)
if optimizer is not None:
if epoch == -1:
optimizer.load_state_dict(checkpoint['optimizer'])
else:
for i in range(epoch + 1):
optimizer.load_state_dict(checkpoint['optimizer'][i])
def save_model(model, model_file, optimizer=None, epoch=-1):
checkpoint = {}
state_dict = model.state_dict()
checkpoint['state_dict'] = state_dict
if optimizer is not None:
if epoch == -1:
optimizer.load_state_dict(checkpoint['optimizer'])
else:
for i in range(epoch + 1):
optimizer.load_state_dict(checkpoint['optimizer'][i])
torch.save(checkpoint, model_file)
return model_file
def save_model_with_best_metrics(model, best_model_file):
checkpoint = {}
state_dict = model.state_dict()
checkpoint['state_dict'] = state_dict
optimizer = model.optimizer
checkpoint['optimizer'] = optimizer
torch.save(checkpoint, best_model_file)
return best_model_file
def save_checkpoint(state, is_best, filename="checkpoint.pth.tar"):
torch.save(state, filename)
if is_best:
torch.save(state, "model_best.pth.tar")
def get_best_metrics_from_path(path):
if os.path.isfile(path):
return get_best_metrics(path)
else:
file_name = os