Team Ai
Apppublic

oookiku/route-explainer

sourceHugging Faceotherupdated 3y agoView on Hugging Face
1likes
eval_classifier.py241 linesDownload Raw Back to root
1import os2import argparse3import json4import multiprocessing5import torch6import time7from tqdm import tqdm8from torch.utils.data import DataLoader9from torchmetrics.classification import MulticlassAccuracy, MulticlassF1Score10from utils.util_calc import TemporalConfusionMatrix11from models.classifiers.nn_classifiers.nn_classifier import NNClassifier12from models.classifiers.ground_truth.ground_truth import GroundTruth13from models.classifiers.ground_truth.ground_truth_base import FAIL_FLAG14from utils.data_utils.tsptw_dataset import TSPTWDataloader15from utils.data_utils.pctsp_dataset import PCTSPDataloader16from utils.data_utils.pctsptw_dataset import PCTSPTWDataloader17from utils.data_utils.cvrp_dataset import CVRPDataloader18from utils.utils import set_device19from utils.utils import load_dataset20 21def load_eval_dataset(dataset_path, problem, model_type, batch_size, num_workers, parallel, num_cpus):22    if model_type == "nn":23        if problem == "tsptw":24            eval_dataset = TSPTWDataloader(dataset_path, sequential=True, parallel=parallel, num_cpus=num_cpus)25        elif problem == "pctsp":26            eval_dataset = PCTSPDataloader(dataset_path, sequential=True, parallel=parallel, num_cpus=num_cpus)27        elif problem == "pctsptw":28            eval_dataset = PCTSPTWDataloader(dataset_path, sequential=True, parallel=parallel, num_cpus=num_cpus)29        elif problem == "cvrp":30            eval_dataset = CVRPDataloader(dataset_path, sequential=True, parallel=parallel, num_cpus=num_cpus)31        else:32            raise NotImplementedError33 34        #------------35        # dataloader36        #------------37        def pad_seq_length(batch):38            data = {}39            for key in batch[0].keys():40                padding_value = True if key == "mask" else 0.041                # post-padding42                data[key] = torch.nn.utils.rnn.pad_sequence([d[key] for d in batch], batch_first=True, padding_value=padding_value)43            pad_mask = torch.nn.utils.rnn.pad_sequence([torch.full((d["mask"].size(0), ), True) for d in batch], batch_first=True, padding_value=False)44            data.update({"pad_mask": pad_mask})45            return data46        eval_dataloader = DataLoader(eval_dataset,47                                     batch_size=batch_size,48                                     shuffle=False,49                                     collate_fn=pad_seq_length,50                                     num_workers=num_workers)51        return eval_dataloader52    else:53        eval_dataset = load_dataset(dataset_path)54        return eval_dataset55 56def eval_classifier(problem: str, 57                    dataset, 58                    model_type: str, 59                    model_dir: str = None, 60                    gpu: int = -1, 61                    num_workers: int = 4, 62                    batch_size: int = 128, 63                    parallel: bool = True,64                    solver: str = "ortools",65                    num_cpus: int = 1):66    #--------------67    # gpu settings68    #--------------69    use_cuda, device = set_device(gpu)70 71    #-------72    # model73    #-------74    num_classes = 3 if problem == "pctsptw" else 275    if model_type == "nn":76        assert model_dir is not None, "please specify model_path when model_type is nn."77        params = argparse.ArgumentParser()78        # model_dir = os.path.split(args.model_path)[0]79        with open(f"{model_dir}/cmd_args.dat", "r") as f:80            params.__dict__ = json.load(f)81        assert params.problem == problem, "problem of the trained model should match that of the dataset"82        model = NNClassifier(problem=params.problem,83                             node_enc_type=params.node_enc_type,84                             edge_enc_type=params.edge_enc_type,85                             dec_type=params.dec_type,86                             emb_dim=params.emb_dim,87                             num_enc_mlp_layers=params.num_enc_mlp_layers,88                             num_dec_mlp_layers=params.num_dec_mlp_layers,89                             num_classes=num_classes,90                             dropout=params.dropout,91                             pos_encoder=params.pos_encoder)92        # load trained weights (the best epoch)93        with open(f"{model_dir}/best_epoch.dat", "r") as f:94            best_epoch = int(f.read())95        print(f"loaded {model_dir}/model_epoch{best_epoch}.pth.")96        model.load_state_dict(torch.load(f"{model_dir}/model_epoch{best_epoch}.pth"))97        if use_cuda:98            model.to(device)99        is_sequential = model.is_sequential100    elif model_type == "ground_truth":101        model = GroundTruth(problem=problem, solver_type=solver)102        is_sequential = False103    else:104        assert False, f"Invalid model type: {model_type}"105 106    #---------107    # Metrics108    #---------109    overall_accuracy = MulticlassF1Score(num_classes=num_classes, average="macro").to(device)110    eval_accuracy_dict = {} # MulticlassAccuracy(num_classes=num_classes, average="macro")111    temp_confmat_dict  = {} # TemporalConfusionMatrix(num_classes=num_classes, seq_length=50, device=device)112    temporal_accuracy_dict  = {}113    num_nodes_dist_dict = {}114                        115    #------------116    # Evaluation117    #------------118    if model_type == "nn":119        model.eval()120        eval_time = 0.0121        print("Evaluating models ...", end="")122        start_time = time.perf_counter()123        for data in dataset:124            if use_cuda:125                data = {key: value.to(device) for key, value in data.items()}126            if not is_sequential:127                shp = data["curr_node_id"].size()128                data = {key: value.flatten(0, 1) for key, value in data.items()}129            probs = model(data) # [batch_size x num_classes] or [batch_size x max_seq_length x num_classes]130            if not is_sequential:131                probs = probs.view(*shp, -1) # [batch_size x max_seq_length x num_classes]132                data["labels"] = data["labels"].view(*shp)133                data["pad_mask"] = data["pad_mask"].view(*shp)134            #------------135            # evaluation136            #------------137            start_eval_time = time.perf_counter()138            # accuracy139            seq_length_list = torch.unique(data["pad_mask"].sum(-1)) 140            for seq_length_tensor in seq_length_list:141                seq_length = seq_length_tensor.item()142                if seq_length not in eval_accuracy_dict.keys():143                    eval_accuracy_dict[seq_length] = MulticlassF1Score(num_classes=num_classes, average="macro").to(device)144                    temp_confmat_dict[seq_length]  = TemporalConfusionMatrix(num_classes=num_classes, seq_length=seq_length, device=device)145                    temporal_accuracy_dict[seq_length] = [MulticlassF1Score(num_classes=num_classes, average="macro").to(device) for _ in range(seq_length)]146                    num_nodes_dist_dict[seq_length] = 0147                seq_length_mask = (data["pad_mask"].sum(-1) == seq_length) # [batch_size]148                extracted_labels = data["labels"][seq_length_mask]149                extracted_probs  = probs[seq_length_mask]150                extracted_mask   = data["pad_mask"][seq_length_mask].view(-1) # [batch_size x max_seq_length] -> [(batch_size*max_seq_length)]151                eval_accuracy_dict[seq_length](extracted_probs.argmax(-1).view(-1)[extracted_mask], extracted_labels.view(-1)[extracted_mask])152                mask = data["pad_mask"].view(-1)153                overall_accuracy(probs.argmax(-1).view(-1)[mask], data["labels"].view(-1)[mask])154                # confusion matrix155                temp_confmat_dict[seq_length].update(probs.argmax(-1), data["labels"], data["pad_mask"])         156                # temporal accuracy157                for step in range(seq_length):158                    temporal_accuracy_dict[seq_length][step](extracted_probs[:, step, :], extracted_labels[:, step])159                # number of samples whose sequence length is seq_length160                num_nodes_dist_dict[seq_length] += len(extracted_labels)161            eval_time += time.perf_counter() - start_eval_time162        calc_time = time.perf_counter() - start_time - eval_time163        total_eval_accuracy = {key: value.compute().item() for key, value in eval_accuracy_dict.items()}164        overall_accuracy = overall_accuracy.compute() #.item()165        temporal_confmat = {key: value.compute() for key, value in temp_confmat_dict.items()}166        temporal_accuracy = {key: [value.compute().item() for value in values] for key, values in temporal_accuracy_dict.items()}167        print("done")168        return overall_accuracy, total_eval_accuracy, temporal_accuracy, calc_time, temporal_confmat, num_nodes_dist_dict169    else:170        eval_accuracy = MulticlassF1Score(num_classes=num_classes, average="macro").to(device)171        print("Loading data ...", end=" ")172        with multiprocessing.Pool(num_cpus) as pool:173            input_list = list(pool.starmap(model.get_inputs, [(instance["tour"], 0, instance) for instance in dataset]))174        print("done")175 176        print("Infering labels ...", end="")177        pool = multiprocessing.Pool(num_cpus)178        start_time = time.perf_counter()179        prob_list = list(pool.starmap(model, tqdm([(inputs, False, False) for inputs in input_list])))180        calc_time = time.perf_counter() - start_time181        pool.close()182        print("done")183 184        print("Evaluating models ...", end="")185        for i, instance in enumerate(dataset):186            labels = instance["labels"]187            for vehicle_id in range(len(labels)):188                for step, label in labels[vehicle_id]:189                    pred_label = prob_list[i][vehicle_id][step-1] # [num_classes]190                    if pred_label == FAIL_FLAG:191                        pred_label = label - 1 if label != 0 else label + 1192                    eval_accuracy(torch.LongTensor([pred_label]).view(1, -1), torch.LongTensor([label]).view(1, -1))193        total_eval_accuracy = eval_accuracy.compute()194        print("done")195        return total_eval_accuracy.item(), calc_time196 197if __name__ == "__main__":198    parser = argparse.ArgumentParser()199    #-----------------200    # general settings201    #-----------------202    parser.add_argument("--gpu", default=-1, type=int, help="Used GPU Number: gpu=-1 indicates using cpu")203    parser.add_argument("--num_workers", default=4, type=int, help="Number of workers in dataloader")204    parser.add_argument("--parallel", )205 206    #-------------207    # data setting208    #-------------209    parser.add_argument("--dataset_path", type=str, help="Path to a dataset", required=True)210    211    #------------------212    # Metrics settings213    #------------------214 215 216    #----------------217    # model settings218    #----------------219    parser.add_argument("--model_type", type=str, default="nn", help="Select from [nn, ground_truth]")220    # nn classifier221    parser.add_argument("--model_dir", type=str, default=None)222    parser.add_argument("--batch_size", type=int, default=256)223    parser.add_argument("--parallel", action="store_true")224    # ground truth225    parser.add_argument("--solver", type=str, default="ortools")226    parser.add_argument("--num_cpus", type=int, default=os.cpu_count())227    args = parser.parse_args()228 229    problem   = str(os.path.basename(os.path.dirname(args.dataset_path)))230 231    dataset = load_eval_dataset(args.dataset_path, problem, args.model_type, args.batch_size, args.num_workers, args.parallel, args.num_cpus)232    eval_classifier(problem=problem, 233         dataset=dataset,234         model_type=args.model_type,235         model_dir=args.model_dir,236         gpu=args.gpu,237         num_workers=args.num_workers,238         batch_size=args.batch_size,239         parallel=args.parallel,240         solver=args.solver,241         num_cpus=args.num_cpus)