| 12 | |
| 13 | |
| 14 | class CallBackVerification(object): |
| 15 | |
| 16 | def __init__(self, val_targets, rec_prefix, summary_writer=None, image_size=(112, 112), wandb_logger=None): |
| 17 | self.rank: int = distributed.get_rank() |
| 18 | self.highest_acc: float = 0.0 |
| 19 | self.highest_acc_list: List[float] = [0.0] * len(val_targets) |
| 20 | self.ver_list: List[object] = [] |
| 21 | self.ver_name_list: List[str] = [] |
| 22 | if self.rank is 0: |
| 23 | self.init_dataset(val_targets=val_targets, data_dir=rec_prefix, image_size=image_size) |
| 24 | |
| 25 | self.summary_writer = summary_writer |
| 26 | self.wandb_logger = wandb_logger |
| 27 | |
| 28 | def ver_test(self, backbone: torch.nn.Module, global_step: int): |
| 29 | results = [] |
| 30 | for i in range(len(self.ver_list)): |
| 31 | acc1, std1, acc2, std2, xnorm, embeddings_list = verification.test( |
| 32 | self.ver_list[i], backbone, 10, 10) |
| 33 | logging.info('[%s][%d]XNorm: %f' % (self.ver_name_list[i], global_step, xnorm)) |
| 34 | logging.info('[%s][%d]Accuracy-Flip: %1.5f+-%1.5f' % (self.ver_name_list[i], global_step, acc2, std2)) |
| 35 | |
| 36 | self.summary_writer: SummaryWriter |
| 37 | self.summary_writer.add_scalar(tag=self.ver_name_list[i], scalar_value=acc2, global_step=global_step, ) |
| 38 | if self.wandb_logger: |
| 39 | import wandb |
| 40 | self.wandb_logger.log({ |
| 41 | f'Acc/val-Acc1 {self.ver_name_list[i]}': acc1, |
| 42 | f'Acc/val-Acc2 {self.ver_name_list[i]}': acc2, |
| 43 | # f'Acc/val-std1 {self.ver_name_list[i]}': std1, |
| 44 | # f'Acc/val-std2 {self.ver_name_list[i]}': acc2, |
| 45 | }) |
| 46 | |
| 47 | if acc2 > self.highest_acc_list[i]: |
| 48 | self.highest_acc_list[i] = acc2 |
| 49 | logging.info( |
| 50 | '[%s][%d]Accuracy-Highest: %1.5f' % (self.ver_name_list[i], global_step, self.highest_acc_list[i])) |
| 51 | results.append(acc2) |
| 52 | |
| 53 | def init_dataset(self, val_targets, data_dir, image_size): |
| 54 | for name in val_targets: |
| 55 | path = os.path.join(data_dir, name + ".bin") |
| 56 | if os.path.exists(path): |
| 57 | data_set = verification.load_bin(path, image_size) |
| 58 | self.ver_list.append(data_set) |
| 59 | self.ver_name_list.append(name) |
| 60 | |
| 61 | def __call__(self, num_update, backbone: torch.nn.Module): |
| 62 | if self.rank is 0 and num_update > 0: |
| 63 | backbone.eval() |
| 64 | self.ver_test(backbone, num_update) |
| 65 | backbone.train() |
| 66 | |
| 67 | |
| 68 | class CallBackLogging(object): |