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:
When I try to sample from the distribution, the song generated is kinda weird:
But the problem is when I try to reconstruct a song:
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?

