MCPcopy Create free account
hub / github.com/easy-graph/Easy-Graph / train

Method train

easygraph/functions/graph_embedding/sdne.py:196–258  ·  view source on GitHub ↗
(
        self,
        model,
        epochs=100,
        lr=0.006,
        bs=100,
        step_size=10,
        gamma=0.9,
        nu1=1e-5,
        nu2=1e-4,
        device="cpu",
        output="out.emb",
    )

Source from the content-addressed store, hash-verified

194 return L_1st, self.alpha * L_2nd, L_1st + self.alpha * L_2nd
195
196 def train(
197 self,
198 model,
199 epochs=100,
200 lr=0.006,
201 bs=100,
202 step_size=10,
203 gamma=0.9,
204 nu1=1e-5,
205 nu2=1e-4,
206 device="cpu",
207 output="out.emb",
208 ):
209 Adj, Node = get_adj(self.graph)
210 model = model.to(device)
211
212 opt = optim.Adam(model.parameters(), lr=lr)
213 scheduler = torch.optim.lr_scheduler.StepLR(
214 opt, step_size=step_size, gamma=gamma
215 )
216 Data = Dataload(Adj, Node)
217 Data = DataLoader(
218 Data,
219 batch_size=bs,
220 shuffle=True,
221 )
222
223 for epoch in range(1, epochs + 1):
224 loss_sum, loss_L1, loss_L2, loss_reg = 0, 0, 0, 0
225 for index in Data:
226 adj_batch = Adj[index]
227 adj_mat = adj_batch[:, index]
228 b_mat = torch.ones_like(adj_batch)
229 b_mat[adj_batch != 0] = self.beta
230
231 opt.zero_grad()
232 L_1st, L_2nd, L_all = model(adj_batch, adj_mat, b_mat)
233 L_reg = 0
234 for param in model.parameters():
235 L_reg += nu1 * torch.sum(torch.abs(param)) + nu2 * torch.sum(
236 param * param
237 )
238 Loss = L_all + L_reg
239 Loss.backward()
240 opt.step()
241 loss_sum += Loss
242 loss_L1 += L_1st
243 loss_L2 += L_2nd
244 loss_reg += L_reg
245 scheduler.step(epoch)
246 # print("The lr for epoch %d is %f" %(epoch, scheduler.get_lr()[0]))
247 print("loss for epoch %d is:" % epoch)
248 print("loss_sum is %f" % loss_sum)
249 print("loss_L1 is %f" % loss_L1)
250 print("loss_L2 is %f" % loss_L2)
251 print("loss_reg is %f" % loss_reg)
252
253 # model.eval()

Callers 2

line.pyFile · 0.80

Calls 4

get_adjFunction · 0.85
DataloadClass · 0.85
savectorMethod · 0.80
toMethod · 0.45

Tested by

no test coverage detected