MCPcopy Create free account
hub / github.com/DirectMolecularConfGen/DMCG / train

Function train

evaluate.py:29–59  ·  view source on GitHub ↗
(model, device, loader, optimizer, scheduler, args)

Source from the content-addressed store, hash-verified

27
28
29def train(model, device, loader, optimizer, scheduler, args):
30 model.train()
31 loss_accum_dict = defaultdict(float)
32 pbar = tqdm(loader, desc="Iteration")
33 for step, batch in enumerate(pbar):
34 batch = batch.to(device)
35
36 if batch.x.shape[0] == 1 or batch.batch[-1] == 0:
37 pass
38 else:
39 atom_pred_list, extra_output = model(batch)
40 optimizer.zero_grad()
41
42 loss, loss_dict = model.compute_loss(atom_pred_list, extra_output, batch, args)
43 loss.backward()
44 optimizer.step()
45 scheduler.step()
46
47 for k, v in loss_dict.items():
48 loss_accum_dict[k] += v.detach().item()
49
50 if step % args.log_interval == 0:
51 description = f"Iteration loss: {loss_accum_dict['loss'] / (step + 1):6.4f} lr: {scheduler.get_last_lr()[0]:.5e}"
52 # for k in loss_accum_dict.keys():
53 # description += f" {k}: {loss_accum_dict[k]/(step+1):6.4f}"
54
55 pbar.set_description(description)
56
57 for k in loss_accum_dict.keys():
58 loss_accum_dict[k] /= step + 1
59 return loss_accum_dict
60
61
62def get_rmsd_min(inputargs):

Callers 1

mainFunction · 0.70

Calls 2

compute_lossMethod · 0.80
stepMethod · 0.45

Tested by

no test coverage detected