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

Method __init__

models/pimodel/pimodel.py:17–57  ·  view source on GitHub ↗

class PiModel contains setter of data_loader, optimizer, and model update methods. Args: net_builder: backbone network class (see net_builder in utils.py) num_classes: # of label classes lambda_u: ratio of unsupervised loss to supervised loss

(self, net_builder, num_classes, lambda_u,
                 num_eval_iter=1000, tb_log=None, ema_m=0.999, logger=None)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 2

Bn_ControllerClass · 0.90
net_builderFunction · 0.85

Tested by

no test coverage detected