(epoch, net, optimizer, dataloader)
| 262 | |
| 263 | |
| 264 | def warmup(epoch, net, optimizer, dataloader): |
| 265 | current_lr = adjust_learning_rate(optimizer, epoch) |
| 266 | net.train() |
| 267 | |
| 268 | num_iter = (len(dataloader.dataset.cIndex) // dataloader.batch_size) + 1 |
| 269 | for batch_idx, (input10, input11, input2, label1, label2) in enumerate(dataloader): |
| 270 | labels = torch.cat((label1, label1, label2), 0) |
| 271 | input1 = torch.cat((input10, input11,), 0) |
| 272 | |
| 273 | input1 = input1.cuda() |
| 274 | input2 = input2.cuda() |
| 275 | labels = labels.cuda() |
| 276 | |
| 277 | _, out0, = net(input1, input2) |
| 278 | loss_id = criterion_id(out0, labels) |
| 279 | |
| 280 | optimizer.zero_grad() |
| 281 | loss_id.backward() |
| 282 | optimizer.step() |
| 283 | |
| 284 | if batch_idx % 50 == 0: |
| 285 | print('%s:%.1f-%s | Epoch [%3d/%3d] Iter[%3d/%3d]\t CE-loss: %.4f\t Current-lr: %.4f' |
| 286 | % (args.dataset, args.noise_rate, args.noise_mode, epoch, 80, batch_idx + 1, |
| 287 | num_iter, loss_id.item(), current_lr)) |
| 288 | |
| 289 | |
| 290 | def eval_train(net, dataloader, type): |
no test coverage detected