Hierarchical LSTM autoencoder - model not training

Viewed 38

I'm trying to reconstruct this paper about hierarchical autoencoder for paragraphs.

The idea is: Break a paragraph into sentences, then encode each sentence using an LSTM, and then using these encoding as an input for another LSTM that encode the entire paragraph.

Then, using a mirror decoder, decode the encoded paragraph using an LSTM into multiple sentences, and then use another LSTM to decode each word, with a linear layer on top and predicts the word.

The objective is to try to predict the original paragraph.

I've done some preprocessing, and right now I save each paragraph as a tensor of (maxSentence,maxWordsPerSentence,VocabSize), using one hot encoding.

My problem is, there model is not learning. The loss stays exactly the same and it doesn't seem as anything is happening.. I wasn't sure on how to calculate the loss (I've ran a batch all together and decoded it into multiple paragraphs, and then calculated the loss against the entire batch predictions, my train function is added below. I don't know if that is the problem (maybe I should calculate loss sentence by sentence instead the entire paragraph?) or maybe I have a problem in my model.

Encoder code:

class Encoder(nn.Module):
def __init__(self, input_dim, emb_dim, enc_hid_dim, dec_hid_dim, dropout):
    super().__init__()
    
    #self.embedding = nn.Embedding(input_dim, emb_dim)
    self.rnn_sent = nn.GRU(input_dim, enc_hid_dim, bidirectional = True)
    self.rnn_par = nn.GRU(enc_hid_dim*2, dec_hid_dim, bidirectional = True)

    
def forward(self, src):
    
    outputs, hidden = self.rnn_sent(src[:,0,0])
    total_out = outputs.unsqueeze(0).permute(1,0,2)
    for i in range(1,src.shape[1]):
      for j in range(src.shape[2]):
        outputs, hidden = self.rnn_sent(src[:,i,j],hidden)
      total_out = torch.cat((total_out,outputs.unsqueeze(0).permute(1,0,2)),dim=1)  

    outputs_par, hidden_par = self.rnn_par(total_out[:,0])
    
    for i in range(total_out.shape[1]):
        outputs_par, hidden_par = self.rnn_par(total_out[:,i],hidden_par)

    return outputs_par, hidden_par

Decoder code:

class Decoder(nn.Module):
def __init__(self, output_dim, emb_dim, enc_hid_dim, dec_hid_dim, dropout, attention):
    super().__init__()

    self.output_dim = output_dim
    self.attention = attention
    #self.embedding = nn.Embedding(output_dim, emb_dim)
    self.rnn_par = nn.GRU((enc_hid_dim * 2), dec_hid_dim*2)
    self.rnn_sen = nn.GRU(output_dim, dec_hid_dim*2)
    self.fc_out = nn.Linear(dec_hid_dim*2, output_dim)
    self.dropout = nn.Dropout(dropout)
    
def forward(self, input, hidden, encoder_outputs):
    output, hidden = self.rnn_par(encoder_outputs)
    all_par = output.unsqueeze(0).permute(1,0,2)
    for i in range(1,max_par_len):
      output,hidden = self.rnn_par(output,hidden)
      all_par = torch.cat((all_par,output.unsqueeze(0).permute(1,0,2)),dim=1)


    for i in range(max_par_len):
      output_arg = self.fc_out(all_par[:,i])
      #output_argmax = F.one_hot(output_arg.argmax(dim = 1), self.output_dim).to(torch.float)
      output_argmax = torch.softmax(output_arg,dim=1)
      output_sen, hidden_sen = self.rnn_sen(output_argmax)
      all_par_sen = output_argmax.unsqueeze(0).permute(1,0,2)
      for j in range(max_sen_len - 1):
        output_sen,hidden_sen = self.rnn_sen(output_argmax,hidden_sen)
        output_arg = self.fc_out(output_sen)
        output_argmax = torch.softmax(output_arg,dim=1)

        all_par_sen = torch.cat((all_par_sen,output_argmax.unsqueeze(0).permute(1,0,2)),dim=1)
      if i == 0:
        all_doc = all_par_sen.unsqueeze(0).permute(1,0,2,3)
      else:
        all_doc = torch.cat((all_doc,all_par_sen.unsqueeze(0).permute(1,0,2,3)),dim=1)

      i+=1
    return all_doc  ,hidden_sen

And my train function:

def train(model, iterator, optimizer, criterion, clip, epoch):

model.train()

epoch_loss = 0
data = tqdm(iterator)
for i, batch in enumerate(data):
    src = batch[0].to(device)#.to(torch.long)#.reshape(batch[0].shape[0],-1)
    trg = batch[0].to(device)#.to(torch.long)#.reshape(batch[0].shape[0],-1)
    target = torch.argmax(trg,dim=3).view(-1)
    print(target)
    optimizer.zero_grad()
    output = model(src, trg).view(-1,OUTPUT_DIM)
    loss = criterion(output, target)        
    loss.backward()
    torch.nn.utils.clip_grad_norm_(model.parameters(), clip)
    optimizer.step()
    epoch_loss += loss.item()


N_EPOCHS = 20
CLIP = 1
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)
criterion = nn.CrossEntropyLoss(ignore_index = vocabulary['<pad>'])
best_valid_loss = float('inf')

for epoch in range(N_EPOCHS):

    start_time = time.time()
    train_loader, valid_loader = data_loaders['train_loader'], data_loaders['test_loader']
    train_loss = train(model, train_loader, optimizer, criterion, CLIP,f'{epoch+1}/{N_EPOCHS}')
    #valid_loss = evaluate(model, valid_loader, criterion)

    end_time = time.time()

    epoch_mins, epoch_secs = epoch_time(start_time, end_time)

    print(f'Epoch: {epoch+1:02} | Time: {epoch_mins}m {epoch_secs}s')
    print(f'\tTrain Loss: {train_loss:.3f} | Train PPL: {math.exp(train_loss):7.3f}')
0 Answers
Related