| 116 | |
| 117 | |
| 118 | def train_one_epoch(model, optimizer, data_loader, device, epoch): |
| 119 | model.train() |
| 120 | loss_function = torch.nn.CrossEntropyLoss() |
| 121 | mean_loss = torch.zeros(1).to(device) |
| 122 | optimizer.zero_grad() |
| 123 | |
| 124 | data_loader = tqdm(data_loader, file=sys.stdout) |
| 125 | |
| 126 | for step, data in enumerate(data_loader): |
| 127 | images, labels = data |
| 128 | |
| 129 | pred = model(images.to(device)) |
| 130 | |
| 131 | loss = loss_function(pred, labels.to(device)) |
| 132 | loss.backward() |
| 133 | mean_loss = (mean_loss * step + loss.detach()) / (step + 1) # update mean losses |
| 134 | |
| 135 | data_loader.desc = "[epoch {}] mean loss {}".format(epoch, round(mean_loss.item(), 3)) |
| 136 | |
| 137 | if not torch.isfinite(loss): |
| 138 | print('WARNING: non-finite loss, ending training ', loss) |
| 139 | sys.exit(1) |
| 140 | |
| 141 | optimizer.step() |
| 142 | optimizer.zero_grad() |
| 143 | |
| 144 | return mean_loss.item() |
| 145 | |
| 146 | |
| 147 | @torch.no_grad() |