(self, args=None)
| 8 | # This is the baseline encoder-decoder we used in the ablation study |
| 9 | class NNET(nn.Module): |
| 10 | def __init__(self, args=None): |
| 11 | super(NNET, self).__init__() |
| 12 | self.encoder = Encoder() |
| 13 | self.decoder = Decoder(num_classes=4) |
| 14 | |
| 15 | def forward(self, x, **kwargs): |
| 16 | out = self.decoder(self.encoder(x), **kwargs) |