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

Class PseudoLabel

models/pseudolabel/pseudolabel.py:17–291  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

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

Callers 1

main_workerFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected