Transformer Autoencoder with low loss and good accuracy can't reconstruct song without teacher forcing

Viewed 270

I am working on an Adversarial Autoencoder with Compressive Transformer for music generation and interpolation. The input for the decoder is a sequence of 8 bars, where each bars is made by 200 tokens. Each bar has 4 tracks which are respectively: drums, bass, guitar and strings. Each note is represented with a sequence <starting time> <note pitch> <note duration>. The compressed representation is the first 16 latent of the encoder output on the last bar (the 8th one), which should capture information about all the 8 provided bars thanks to the memories of the Compressive Transformer. The training is done with the WGAN-GP loss.

The encoder output distribution matches the prior distribution (a mixture of 4 gaussians), the accuracy reaches the value of 0.6 and the loss decrases well:

enter image description here

enter image description here

When I try to sample from the distribution, the song generated is kinda weird:

generated song

But the problem is when I try to reconstruct a song:

original

reconstructed

I think that the problem is releated to what is discussed here: at training time, when the decoder needs to reconstruct the token in position k, the previous <k token are the right ones because of the teacher forcing, while at testing time the previous token are feed in an autoregressive way and, because not all of them are the right ones, this confuses the Decoder. The code for Greedy Decoding is the following:

    def greedy_decode(self, latent, n_bars, desc):
        _, _, d_mems, d_cmems = get_memories(n_batch=1)
        outs = []
        for _ in tqdm(range(n_bars), position=0, leave=True, desc=desc):
            trg = np.full((4, 1, 1), config["tokens"]["sos"])
            trg = torch.LongTensor(trg).to(config["train"]["device"])
            for _ in range(config["model"]["seq_len"] - 1):  # for each token of each bar
                trg_mask = create_trg_mask(trg.cpu().numpy())
                out, _, _, _, _, _ = self.decoder(trg, trg_mask, None, latent, d_mems, d_cmems)
                out = torch.max(out, dim=-2).indices
                out = out.permute(2, 0, 1)
                trg = torch.cat((trg, out[..., -1:]), dim=-1)
            trg_mask = create_trg_mask(trg.cpu().numpy())
            out, _, _, d_mems, d_cmems, _ = self.decoder(trg, trg_mask, None, latent, d_mems, d_cmems)
            out = torch.max(out, dim=-2).indices
            out = out.permute(2, 0, 1)
            outs.append(out)
        return outs

I also tried Beam Search, but it does not seem to solve the problem. Now I am trying to use the previous generated token instead to the true one with a probability of 1/2, but the training is very slow. Does anyone know a way to teach the transformer to be not too dependent on teacher forcing?

0 Answers
Related