MCPcopy Create free account
hub / github.com/kwuking/TimeMixer / train

Method train

exp/exp_classification.py:81–157  ·  view source on GitHub ↗
(self, setting)

Source from the content-addressed store, hash-verified

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

Callers 1

valiMethod · 0.45

Calls 8

_get_dataMethod · 0.95
_select_optimizerMethod · 0.95
_select_criterionMethod · 0.95
valiMethod · 0.95
EarlyStoppingClass · 0.90
adjust_learning_rateFunction · 0.90
loadMethod · 0.80
backwardMethod · 0.45

Tested by

no test coverage detected