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

Method fit

trainer.py:109–127  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

107 np.random.seed(self.args.seed)
108
109 def fit(self):
110 self.prepare_data()
111 self.model = self.model(nfeat=self.nfeat_dim,
112 nhid=self.args.nhid,
113 nclass=self.nclass,
114 dropout=self.args.dropout)
115 print(self.model.parameters)
116 self.model = self.model.to(self.device)
117
118 self.optimizer = th.optim.Adam(self.model.parameters(), lr=self.args.lr)
119 self.criterion = th.nn.CrossEntropyLoss()
120
121 self.model_param = sum(param.numel() for param in self.model.parameters())
122 print('# model parameters:', self.model_param)
123 self.convert_tensor()
124
125 start = time()
126 self.train()
127 self.train_time = time() - start
128
129 @classmethod
130 def set_description(cls, desc):

Callers 1

mainFunction · 0.95

Calls 3

prepare_dataMethod · 0.95
convert_tensorMethod · 0.95
trainMethod · 0.95

Tested by

no test coverage detected