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

Method __init__

domainbed/algorithms.py:1172–1193  ·  view source on GitHub ↗
(self, input_shape, num_classes, num_domains, hparams)

Source from the content-addressed store, hash-verified

1170 "Invariant Gradients variances for Out-of-distribution Generalization"
1171
1172 def __init__(self, input_shape, num_classes, num_domains, hparams):
1173 assert backpack is not None, "Install backpack with: 'pip install backpack-for-pytorch==1.3.0'"
1174 super(Fishr, self).__init__(input_shape, num_classes, num_domains, hparams)
1175 self.num_domains = num_domains
1176
1177 self.featurizer = networks.Featurizer(input_shape, self.hparams)
1178 self.classifier = extend(
1179 networks.Classifier(
1180 self.featurizer.n_outputs,
1181 num_classes,
1182 self.hparams['nonlinear_classifier'],
1183 )
1184 )
1185 self.network = nn.Sequential(self.featurizer, self.classifier)
1186
1187 self.register_buffer("update_count", torch.tensor([0]))
1188 self.bce_extended = extend(nn.CrossEntropyLoss(reduction='none'))
1189 self.ema_per_domain = [
1190 MovingAverage(ema=self.hparams["ema"], oneminusema_correction=True)
1191 for _ in range(self.num_domains)
1192 ]
1193 self._init_optimizer()
1194
1195 def _init_optimizer(self):
1196 self.optimizer = torch.optim.Adam(

Callers

nothing calls this directly

Calls 3

_init_optimizerMethod · 0.95
MovingAverageClass · 0.90
__init__Method · 0.45

Tested by

no test coverage detected