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

Class MixMatch

models/mixmatch/mixmatch.py:19–316  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

17
18
19class MixMatch:
20 def __init__(self, net_builder, num_classes, ema_m, T, lambda_u, \
21 t_fn=None, it=0, num_eval_iter=1000, tb_log=None, logger=None):
22 """
23 class Mixmatch contains setter of data_loader, optimizer, and model update methods.
24 Args:
25 net_builder: backbone network class (see net_builder in utils.py)
26 num_classes: # of label classes
27 ema_m: momentum of exponential moving average for eval_model
28 T: Temperature scaling parameter for output sharpening (only when hard_label = False)
29 p_cutoff: confidence cutoff parameters for loss masking
30 lambda_u: ratio of unsupervised loss to supervised loss
31 hard_label: If True, consistency regularization use a hard pseudo label.
32 it: initial iteration count
33 num_eval_iter: freqeuncy of iteration (after 500,000 iters)
34 tb_log: tensorboard writer (see train_utils.py)
35 logger: logger (see utils.py)
36 """
37
38 super(MixMatch, self).__init__()
39
40 # momentum update param
41 self.loader = {}
42 self.num_classes = num_classes
43 self.ema_m = ema_m
44
45 # create the encoders
46 # network is builded only by num_classes,
47 # other configs are covered in main.py
48 self.model = net_builder(num_classes=num_classes)
49 self.ema_model = deepcopy(self.model)
50
51 self.num_eval_iter = num_eval_iter
52 self.t_fn = Get_Scalar(T) # temperature params function
53 self.lambda_u = lambda_u
54 self.tb_log = tb_log
55
56 self.optimizer = None
57 self.scheduler = None
58
59 self.it = 0
60
61 self.logger = logger
62 self.print_fn = print if logger is None else logger.info
63 self.bn_controller = Bn_Controller()
64
65 def set_data_loader(self, loader_dict):
66 self.loader_dict = loader_dict
67 self.print_fn(f'[!] data loader keys: {self.loader_dict.keys()}')
68
69 def set_optimizer(self, optimizer, scheduler=None):
70 self.optimizer = optimizer
71 self.scheduler = scheduler
72
73 def train(self, args, logger=None):
74
75 ngpus_per_node = torch.cuda.device_count()
76

Callers 1

main_workerFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected