MCPcopy Create free account
hub / github.com/TrustAGI-Lab/MTGODE / train

Method train

trainer.py:22–50  ·  view source on GitHub ↗
(self, input, real_val, idx=None)

Source from the content-addressed store, hash-verified

20 self.cl = cl
21
22 def train(self, input, real_val, idx=None):
23 self.model.train()
24 self.optimizer.zero_grad()
25 output = self.model(input, idx=idx).transpose(1, 3)
26 nfe_1 = self.model.ODE.odefunc.nfe # get CTA nfe
27 nfe_2 = self.model.ODE.odefunc.stnet.gconv_1.CGPODE.odefunc.nfe // nfe_1 # get CPG nfe
28 self.model.ODE.odefunc.nfe = 0 # reset CTA nfe
29 self.model.ODE.odefunc.stnet.gconv_1.CGPODE.odefunc.nfe = 0 # reset CGP 1 nfe
30 self.model.ODE.odefunc.stnet.gconv_2.CGPODE.odefunc.nfe = 0 # reset CGP 2 nfe
31 real = torch.unsqueeze(real_val, dim=1)
32 predict = self.scaler.inverse_transform(output)
33 if self.iter % self.step == 0 and self.task_level <= self.seq_out_len:
34 self.task_level += 1
35 if self.cl:
36 loss = self.loss(predict[:, :, :, :self.task_level], real[:, :, :, :self.task_level], 0.0)
37 else:
38 loss = self.loss(predict, real, 0.0)
39
40 loss.backward()
41
42 if self.clip is not None:
43 torch.nn.utils.clip_grad_norm_(self.model.parameters(), self.clip)
44
45 self.optimizer.step()
46 mape = util.masked_mape(predict, real, 0.0).item()
47 rmse = util.masked_rmse(predict, real, 0.0).item()
48 self.iter += 1
49
50 return loss.item(), mape, rmse, nfe_1, nfe_2
51
52 def eval(self, input, real_val):
53 self.model.eval()

Callers 2

mainFunction · 0.95
trainFunction · 0.80

Calls 1

inverse_transformMethod · 0.80

Tested by

no test coverage detected