| 1109 | """ |
| 1110 | |
| 1111 | def __init__(self, input_shape, num_classes, num_domains, hparams): |
| 1112 | super(SANDMask, self).__init__(input_shape, num_classes, num_domains, hparams) |
| 1113 | |
| 1114 | self.tau = hparams["tau"] |
| 1115 | self.k = hparams["k"] |
| 1116 | betas = (0.9, 0.999) |
| 1117 | self.optimizer = torch.optim.Adam( |
| 1118 | self.network.parameters(), |
| 1119 | lr=self.hparams["lr"], |
| 1120 | weight_decay=self.hparams['weight_decay'], |
| 1121 | betas=betas |
| 1122 | ) |
| 1123 | |
| 1124 | self.register_buffer('update_count', torch.tensor([0])) |
| 1125 | |
| 1126 | def update(self, minibatches, unlabeled=None): |
| 1127 | |