(self, setting)
| 79 | return total_loss, accuracy |
| 80 | |
| 81 | def train(self, setting): |
| 82 | train_data, train_loader = self._get_data(flag='TRAIN') |
| 83 | vali_data, vali_loader = self._get_data(flag='TEST') |
| 84 | test_data, test_loader = self._get_data(flag='TEST') |
| 85 | |
| 86 | path = os.path.join(self.args.checkpoints, setting) |
| 87 | if not os.path.exists(path): |
| 88 | os.makedirs(path) |
| 89 | |
| 90 | time_now = time.time() |
| 91 | |
| 92 | train_steps = len(train_loader) |
| 93 | early_stopping = EarlyStopping(patience=self.args.patience, verbose=True) |
| 94 | |
| 95 | model_optim = self._select_optimizer() |
| 96 | criterion = self._select_criterion() |
| 97 | |
| 98 | scheduler = lr_scheduler.OneCycleLR(optimizer=model_optim, |
| 99 | steps_per_epoch=train_steps, |
| 100 | pct_start=self.args.pct_start, |
| 101 | epochs=self.args.train_epochs, |
| 102 | max_lr=self.args.learning_rate) |
| 103 | |
| 104 | for epoch in range(self.args.train_epochs): |
| 105 | iter_count = 0 |
| 106 | train_loss = [] |
| 107 | |
| 108 | self.model.train() |
| 109 | epoch_time = time.time() |
| 110 | |
| 111 | for i, (batch_x, label, padding_mask) in enumerate(train_loader): |
| 112 | iter_count += 1 |
| 113 | model_optim.zero_grad() |
| 114 | |
| 115 | batch_x = batch_x.float().to(self.device) |
| 116 | padding_mask = padding_mask.float().to(self.device) |
| 117 | label = label.to(self.device) |
| 118 | |
| 119 | outputs = self.model(batch_x, padding_mask, None, None) |
| 120 | loss = criterion(outputs, label.long().squeeze(-1)) |
| 121 | train_loss.append(loss.item()) |
| 122 | |
| 123 | if (i + 1) % 100 == 0: |
| 124 | print("\titers: {0}, epoch: {1} | loss: {2:.7f}".format(i + 1, epoch + 1, loss.item())) |
| 125 | speed = (time.time() - time_now) / iter_count |
| 126 | left_time = speed * ((self.args.train_epochs - epoch) * train_steps - i) |
| 127 | print('\tspeed: {:.4f}s/iter; left time: {:.4f}s'.format(speed, left_time)) |
| 128 | iter_count = 0 |
| 129 | time_now = time.time() |
| 130 | |
| 131 | loss.backward() |
| 132 | nn.utils.clip_grad_norm_(self.model.parameters(), max_norm=4.0) |
| 133 | model_optim.step() |
| 134 | |
| 135 | # if self.args.lradj == 'TST': |
| 136 | # adjust_learning_rate(model_optim, scheduler, epoch + 1, self.args, printout=False) |
| 137 | # scheduler.step() |
| 138 |
no test coverage detected