| 60 | return ZippedDataset(self.enc_preproc.dataset(section), self.dec_preproc.dataset(section)) |
| 61 | |
| 62 | def __init__(self, preproc, device, encoder, decoder): |
| 63 | super().__init__() |
| 64 | self.preproc = preproc |
| 65 | self.encoder = registry.construct( |
| 66 | 'encoder', encoder, device=device, preproc=preproc.enc_preproc) |
| 67 | self.decoder = registry.construct( |
| 68 | 'decoder', decoder, device=device, preproc=preproc.dec_preproc) |
| 69 | self.decoder.visualize_flag = False |
| 70 | |
| 71 | if getattr(self.encoder, 'batched'): |
| 72 | self.compute_loss = self._compute_loss_enc_batched |
| 73 | else: |
| 74 | self.compute_loss = self._compute_loss_unbatched |
| 75 | |
| 76 | def _compute_loss_enc_batched(self, batch, debug=False): |
| 77 | losses = [] |