MCPcopy Create free account
hub / github.com/chengsen/PyTorch_TextGCN / train

Method train

trainer.py:160–181  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

158 self.val_lst = th.tensor(self.val_lst).long().to(self.device)
159
160 def train(self):
161 for epoch in range(self.max_epoch):
162 self.model.train()
163 self.optimizer.zero_grad()
164
165 logits = self.model.forward(self.features, self.adj)
166 loss = self.criterion(logits[self.train_lst],
167 self.target[self.train_lst])
168
169 loss.backward()
170 self.optimizer.step()
171
172 val_desc = self.val(self.val_lst)
173
174 desc = dict(**{"epoch" : epoch,
175 "train_loss": loss.item(),
176 }, **val_desc)
177
178 self.set_description(desc)
179
180 if self.earlystopping(val_desc["val_loss"]):
181 break
182
183 @th.no_grad()
184 def val(self, x, prefix="val"):

Callers 1

fitMethod · 0.95

Calls 3

valMethod · 0.95
set_descriptionMethod · 0.95
forwardMethod · 0.45

Tested by

no test coverage detected