MCPcopy Create free account
hub / github.com/alexrame/fishr / __init__

Method __init__

domainbed/algorithms.py:203–232  ·  view source on GitHub ↗
(self, input_shape, num_classes, num_domains, hparams, conditional, class_balance)

Source from the content-addressed store, hash-verified

201 """Domain-Adversarial Neural Networks (abstract class)"""
202
203 def __init__(self, input_shape, num_classes, num_domains, hparams, conditional, class_balance):
204
205 super(AbstractDANN, self).__init__(input_shape, num_classes, num_domains, hparams)
206
207 self.register_buffer('update_count', torch.tensor([0]))
208 self.conditional = conditional
209 self.class_balance = class_balance
210
211 # Algorithms
212 self.featurizer = networks.Featurizer(input_shape, self.hparams)
213 self.classifier = networks.Classifier(
214 self.featurizer.n_outputs, num_classes, self.hparams['nonlinear_classifier']
215 )
216 self.discriminator = networks.MLP(self.featurizer.n_outputs, num_domains, self.hparams)
217 self.class_embeddings = nn.Embedding(num_classes, self.featurizer.n_outputs)
218
219 # Optimizers
220 self.disc_opt = torch.optim.Adam(
221 (list(self.discriminator.parameters()) + list(self.class_embeddings.parameters())),
222 lr=self.hparams["lr_d"],
223 weight_decay=self.hparams['weight_decay_d'],
224 betas=(self.hparams['beta1'], 0.9)
225 )
226
227 self.gen_opt = torch.optim.Adam(
228 (list(self.featurizer.parameters()) + list(self.classifier.parameters())),
229 lr=self.hparams["lr_g"],
230 weight_decay=self.hparams['weight_decay_g'],
231 betas=(self.hparams['beta1'], 0.9)
232 )
233
234 def update(self, minibatches, unlabeled=None):
235 device = "cuda" if minibatches[0][0].is_cuda else "cpu"

Callers

nothing calls this directly

Calls 1

__init__Method · 0.45

Tested by

no test coverage detected