IKMLab/MPTR_AutoT
0
1from pathlib import Path2import torch3import os4import random5import argparse6import json7import pandas as pd8import numpy as np9from sklearn.metrics import precision_recall_fscore_support10from ast import literal_eval11 12 13def pred_by_threshold(14 t: float,15 y_true: np.array,16 similarities: np.array,17 classes: dict,18):19 preds = (similarities >= t) * 120 sk_results = precision_recall_fscore_support(21 y_true,22 preds,23 # average="samples", # For calculating sample-wise P and R scores.24 )25 outputs = {26 "f1": np.average(sk_results[2]),27 "P": np.average(sk_results[0]),28 "R": np.average(sk_results[1]),29 }30 for label_name, idx in classes.items():31 outputs[f"{label_name}_f1"] = sk_results[2][idx]32 return outputs33 34 35def get_avg_length(dataset: torch.utils.data.Dataset):36 all_lengths = 037 data_size = len(dataset)38 for i in range(data_size):39 all_lengths += len(dataset[i]["input_ids"])40 return all_lengths / data_size41 42 43def load_csv_multi_label(filename: str, col_name: str = "labels") -> pd.DataFrame:44 """Prevent Pandas from converting lists of int into lists of strings.45 46 Args:47 filename (str): path of a csv file48 col_name (str, optional): column name of lists of int. Defaults to 'labels'.49 50 Returns:51 pd.DataFrame: a Pandas dataframe52 """53 return pd.read_csv(filename, converters={col_name: literal_eval})54 55 56def save_logged_results(filename: str, results: dict):57 try:58 old_df = pd.read_csv(filename)59 df = pd.concat([old_df, pd.DataFrame(results)], ignore_index=True)60 except FileNotFoundError:61 df = pd.DataFrame(results)62 63 df.to_csv(filename, index=None)64 65 66def set_seed(seed):67 """68 Args:69 seed: an integer number to initialize a pseudorandom number generator70 """71 os.environ["PYTHONHASHSEED"] = str(seed)72 random.seed(seed)73 np.random.seed(seed)74 torch.manual_seed(seed)75 76 if torch.cuda.is_available():77 torch.cuda.manual_seed(seed)78 # torch.cuda.manual_seed_all(seed) # if using more than one GPUs79 torch.backends.cudnn.deterministic = True80 torch.backends.cudnn.benchmark = False81 82 83def save_baseline_table(84 y_preds: list,85 baseline_name: str,86 baseline_result_file: str = "results/baselines.pkl",87 all_doc_idx: list = None,88) -> None:89 if Path(baseline_result_file).exists():90 df = pd.read_pickle(baseline_result_file)91 else:92 assert all_doc_idx is not None93 df = pd.DataFrame({"doc_idx": all_doc_idx})94 95 df[baseline_name] = y_preds96 df.to_pickle(baseline_result_file)97 98 99def load_params(path_of_params):100 with open(path_of_params, "r") as f:101 params = json.load(f)102 return argparse.Namespace(**params)103 104 105def get_label_words(classes: list, use_multi_label_words=False) -> list:106 mapping = {107 "cyst": "cyst",108 "HCC": "hcc", # hepatoma109 "cirrhosis": "cirrhosis",110 "post-treatment": "posttreatment",111 "steatosis": "steatosis",112 "metastasis": "metastasis",113 "hemangioma": "hemangioma",114 }115 if use_multi_label_words:116 mapping = {117 "cyst": ["cyst"],118 "HCC": ["hcc", "hepatoma"], # hepatoma119 "cirrhosis": ["cirrhosis"],120 "post-treatment": ["posttreatment"],121 "steatosis": ["steatosis", "steatohepatitis"],122 "metastasis": ["metastasis"],123 "hemangioma": ["hemangioma"],124 }125 126 label_words = [mapping[c] for c in classes]127 return label_words128 129 130def seed_mapper(data_type: str) -> list:131 mapping = {132 "train_8": [2, 4, 7, 11, 21, 23, 24, 36, 44, 128],133 "train_32": [0, 1, 3, 7, 10],134 }135 if data_type in mapping:136 return mapping[data_type]137 else:138 raise NotImplementedError139 