| 4 | |
| 5 | |
| 6 | class 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() |