import torch import torch.nn as nn import torch.optim as optim from tqdm import tqdm model = BERT().to(device) token_criterion = nn.CrossEntropyLoss() # Expect indices, not one-hot vectors classification_criterion = nn.BCEWithLogitsLoss() # For logits directly optimizer = optim.Adam(model.parameters(), lr=0.001) #keeping track of the losses across all epochs losses = [[]] #these epochs can take a while, keeping it at a fairly small number for epoch in range(4): for sequence_batch, location_batch, classtarg_batch in tqdm(zip(sequence_tokens_batches, sentence_location_batches, is_positives_batches)): # Zeroing out gradients from last iteration optimizer.zero_grad() # Masking the tokens in the input sequence masked_tokens, masked_token_locations = mask_batch(sequence_batch) # Generating class and masked token predictions clsf_logits, token_logits = model(masked_tokens, location_batch, masked_token_locations) # Setting up target for masked token prediction masked_token_targets = sequence_batch[masked_token_locations.bool()] # Calculating loss for next sentence classification loss_clsf = classification_criterion(clsf_logits.squeeze(), classtarg_batch.float()) # Calculating loss for masked language modeling loss_mlm = token_criterion(token_logits, masked_token_targets) # Combining losses loss = loss_mlm + loss_clsf #keeping track of loss across the current epoch losses[-1].append(float(loss)) # Backpropagation loss.backward() optimizer.step() print(f'=======Epoch {epoch} Completed=======') print(f'average loss in epoch: {np.mean(losses[-1])}') losses.append([])