| 267 | return codes.transpose(0, 1) |
| 268 | |
| 269 | def decode_latent(self, codes, y_mask, refer, refer_mask, ge): |
| 270 | quantized = self.quantizer.decode(codes) |
| 271 | |
| 272 | y = self.vq_proj(quantized) * y_mask |
| 273 | y = self.encoder_ssl(y * y_mask, y_mask) |
| 274 | |
| 275 | y = self.mrte(y, y_mask, refer, refer_mask, ge) |
| 276 | |
| 277 | y = self.encoder2(y * y_mask, y_mask) |
| 278 | |
| 279 | stats = self.proj(y) * y_mask |
| 280 | m, logs = torch.split(stats, self.out_channels, dim=1) |
| 281 | return y, m, logs, y_mask, quantized |
| 282 | |
| 283 | |
| 284 | class ResidualCouplingBlock(nn.Module): |