ft_losses = [[]*1] for epoch in range(5): for i in tqdm(range(0, train_pos.shape[0], batch_size)): if i+batch_size>=train_pos.shape[0]: break #getting batch train_pos_batch = train_pos[i:i+batch_size] train_tok_batch = train_tok[i:i+batch_size] train_targ_batch = train_targ[i:i+batch_size] # Zeroing out gradients from last iteration optimizer.zero_grad() # Masking the tokens in the input sequence masked_tokens, masked_token_locations = mask_batch(train_tok_batch) # Generating class and masked token predictions clsf_logits, token_logits = model(train_tok_batch, train_pos_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(), train_targ_batch.float()) # Combining losses loss = loss_clsf ft_losses[-1].append(float(loss)) # Backpropagation loss.backward() optimizer.step() print(f'=======Epoch {epoch} Completed=======') print(f'average loss in epoch: {np.mean(ft_losses[-1])}') losses.append([])