import random import numpy as np def train(rnn, training_data, n_epoch = 10, n_batch_size = 64, report_every = 50, learning_rate = 0.2, criterion = [nn.NLLLoss](https://docs.pytorch.org/docs/stable/generated/torch.nn.NLLLoss.html#torch.nn.NLLLoss "torch.nn.NLLLoss")()): """ Learn on a batch of training_data for a specified number of iterations and reporting thresholds """ # Keep track of losses for plotting current_loss = 0 all_losses = [] [rnn.train](https://docs.pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module.train "torch.nn.Module.train")() optimizer = [torch.optim.SGD](https://docs.pytorch.org/docs/stable/generated/torch.optim.SGD.html#torch.optim.SGD "torch.optim.SGD")([rnn.parameters](https://docs.pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module.parameters "torch.nn.Module.parameters")(), lr=learning_rate) start = time.time() print(f"training on data set with n = {len(training_data)}") for iter in range(1, n_epoch + 1): [rnn.zero_grad](https://docs.pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module.zero_grad "torch.nn.Module.zero_grad")() # clear the gradients # create some minibatches # we cannot use dataloaders because each of our names is a different length batches = list(range(len(training_data))) random.shuffle(batches) batches = np.array_split(batches, len(batches) //n_batch_size ) for idx, batch in enumerate(batches): batch_loss = 0 for i in batch: #for each example in this batch (label_tensor, text_tensor, label, text) = training_data[i] [output](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = [rnn.forward](https://docs.pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module.forward "torch.nn.Module.forward")(text_tensor) loss = criterion([output](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor"), label_tensor) batch_loss += loss # optimize parameters batch_loss.backward() [nn.utils.clip_grad_norm_](https://docs.pytorch.org/docs/stable/generated/torch.nn.utils.clip_grad_norm_.html#torch.nn.utils.clip_grad_norm_ "torch.nn.utils.clip_grad_norm_")([rnn.parameters](https://docs.pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module.parameters "torch.nn.Module.parameters")(), 3) optimizer.step() optimizer.zero_grad() current_loss += batch_loss.item() / len(batch) all_losses.append(current_loss / len(batches) ) if iter % report_every == 0: print(f"{iter} ({iter / n_epoch:.0%}): \t average batch loss = {all_losses[-1]}") current_loss = 0 return all_losses