| 14 | |
| 15 | |
| 16 | class PiModel: |
| 17 | def __init__(self, net_builder, num_classes, lambda_u, |
| 18 | num_eval_iter=1000, tb_log=None, ema_m=0.999, logger=None): |
| 19 | """ |
| 20 | class PiModel contains setter of data_loader, optimizer, and model update methods. |
| 21 | Args: |
| 22 | net_builder: backbone network class (see net_builder in utils.py) |
| 23 | num_classes: # of label classes |
| 24 | lambda_u: ratio of unsupervised loss to supervised loss |
| 25 | it: initial iteration count |
| 26 | num_eval_iter: frequency of evaluation. |
| 27 | tb_log: tensorboard writer (see train_utils.py) |
| 28 | logger: logger (see utils.py) |
| 29 | """ |
| 30 | |
| 31 | super(PiModel, self).__init__() |
| 32 | |
| 33 | # momentum update param |
| 34 | self.loader = {} |
| 35 | self.num_classes = num_classes |
| 36 | |
| 37 | # create the encoders |
| 38 | # network is builded only by num_classes, |
| 39 | # other configs are covered in main.py |
| 40 | |
| 41 | self.model = net_builder(num_classes=num_classes) |
| 42 | self.num_eval_iter = num_eval_iter |
| 43 | self.lambda_u = lambda_u |
| 44 | self.tb_log = tb_log |
| 45 | |
| 46 | self.optimizer = None |
| 47 | self.scheduler = None |
| 48 | |
| 49 | self.it = 0 |
| 50 | |
| 51 | self.logger = logger |
| 52 | self.print_fn = print if logger is None else logger.info |
| 53 | |
| 54 | self.ema_m = ema_m |
| 55 | self.ema_model = deepcopy(self.model) |
| 56 | |
| 57 | self.bn_controller = Bn_Controller() |
| 58 | |
| 59 | def set_data_loader(self, loader_dict): |
| 60 | self.loader_dict = loader_dict |
| 61 | self.print_fn(f'[!] data loader keys: {self.loader_dict.keys()}') |
| 62 | |
| 63 | def set_optimizer(self, optimizer, scheduler=None): |
| 64 | self.optimizer = optimizer |
| 65 | self.scheduler = scheduler |
| 66 | |
| 67 | def train(self, args): |
| 68 | |
| 69 | ngpus_per_node = torch.cuda.device_count() |
| 70 | |
| 71 | # lb: labeled, ulb: unlabeled |
| 72 | self.model.train() |
| 73 | self.ema = EMA(self.model, self.ema_m) |