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

Class Trainer

trainer.py:6–65  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

4
5
6class Trainer():
7 def __init__(self, model, lrate, wdecay, clip, step_size, seq_out_len, scaler, device, cl=True):
8 self.scaler = scaler
9 self.model = model
10 self.model.to(device)
11 self.optimizer = optim.Adam(self.model.parameters(), lr=lrate, weight_decay=wdecay)
12 self.loss = util.masked_mae
13 # self.loss = util.masked_mse
14 # self.loss = util.masked_rmses
15 self.clip = clip
16 self.step = step_size
17 self.iter = 1
18 self.task_level = 1
19 self.seq_out_len = seq_out_len
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()
54 output = self.model(input)
55 self.model.ODE.odefunc.nfe = 0 # reset CTA nfe
56 self.model.ODE.odefunc.stnet.gconv_1.CGPODE.odefunc.nfe = 0 # reset CGP 1 nfe
57 self.model.ODE.odefunc.stnet.gconv_2.CGPODE.odefunc.nfe = 0 # reset CGP 2 nfe
58 output = output.transpose(1,3)
59 real = torch.unsqueeze(real_val, dim=1)
60 predict = self.scaler.inverse_transform(output)
61 loss = self.loss(predict, real, 0.0)
62 mape = util.masked_mape(predict, real, 0.0).item()
63 rmse = util.masked_rmse(predict, real, 0.0).item()

Callers 1

mainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected