MCPcopy Create free account
hub / github.com/circuitnet/CircuitNet / train

Function train

net_delay_prediction/train.py:63–108  ·  view source on GitHub ↗
(model, args)

Source from the content-addressed store, hash-verified

61 test_dict(data_test)
62
63def 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
110if __name__ == '__main__':
111 args = parser.parse_args()

Callers 1

train.pyFile · 0.70

Calls 2

load_dataFunction · 0.90
test_netdelayFunction · 0.85

Tested by

no test coverage detected