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

Method __init__

models/remixmatch/remixmatch.py:21–67  ·  view source on GitHub ↗

class Fixmatch 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 ema_m: momentum of exponential moving average for eval_mod

(self, net_builder, num_classes, ema_m, T, lambda_u, \
                 w_match,
                 t_fn=None, it=0, num_eval_iter=1000, tb_log=None, logger=None)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 3

Bn_ControllerClass · 0.90
net_builderFunction · 0.85
Get_ScalarClass · 0.70

Tested by

no test coverage detected