CoolFace
Apppublic

IKMLab/MPTR_AutoT

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
utils.py139 linesDownload Raw Back to root
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