MCPcopy Create free account
hub / github.com/TorchSSL/TorchSSL / PiModel

Class PiModel

models/pimodel/pimodel.py:16–274  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

14
15
16class 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)

Callers 1

main_workerFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected