oookiku/route-explainer
1
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)