(model, args)
| 61 | test_dict(data_test) |
| 62 | |
| 63 | def train(model, args): |
| 64 | data_train, data_test = load_data(args) |
| 65 | |
| 66 | optimizer = torch.optim.Adam(model.parameters(), lr=args.lr) |
| 67 | |
| 68 | for e in range(args.iteration): |
| 69 | model.train() |
| 70 | train_loss_tot_net_delays = 0 |
| 71 | optimizer.zero_grad() |
| 72 | |
| 73 | for k, g in random.sample(data_train.items(), args.batch_size): |
| 74 | pred_net_delays= model(g) |
| 75 | loss_net_delays = 0 |
| 76 | |
| 77 | loss_net_delays = F.mse_loss(pred_net_delays, g.edges['net_out'].data['net_delays_log']) |
| 78 | train_loss_tot_net_delays += loss_net_delays.item() |
| 79 | loss_net_delays.backward() |
| 80 | |
| 81 | optimizer.step() |
| 82 | |
| 83 | if e == 0 or e % 20 == 19: |
| 84 | with torch.no_grad(): |
| 85 | model.eval() |
| 86 | test_loss_tot_net_delays= 0 |
| 87 | for k, g in data_test.items(): |
| 88 | pred_net_delays= model(g) |
| 89 | |
| 90 | test_loss_tot_net_delays += F.mse_loss(pred_net_delays, g.edges['net_out'].data['net_delays_log']).item() |
| 91 | |
| 92 | print('Epoch {}, net delay {:.6f}/{:.6f})'.format( |
| 93 | e, |
| 94 | train_loss_tot_net_delays / args.batch_size, |
| 95 | test_loss_tot_net_delays / len(data_test) |
| 96 | ) |
| 97 | ) |
| 98 | |
| 99 | if e == 0 or e % 200 == 199 or (e > 6000 and test_loss_tot_net_delays / len(data_test) < 6): |
| 100 | if args.checkpoint: |
| 101 | save_path = './checkpoints/{}/{}.pth'.format(args.checkpoint, e) |
| 102 | torch.save(model.state_dict(), save_path) |
| 103 | print('saved model to', save_path) |
| 104 | try: |
| 105 | test_netdelay(model) |
| 106 | except ValueError as e: |
| 107 | print(e) |
| 108 | print('Error testing, but ignored') |
| 109 | |
| 110 | if __name__ == '__main__': |
| 111 | args = parser.parse_args() |
no test coverage detected