MCPcopy Create free account
hub / github.com/DSL-Lab/GRBM / train

Function train

main.py:31–49  ·  view source on GitHub ↗
(model,
          train_loader,
          optimizer,
          config)

Source from the content-addressed store, hash-verified

29
30
31def train(model,
32 train_loader,
33 optimizer,
34 config):
35 model.train()
36 for ii, (data, _) in enumerate(tqdm(train_loader)):
37 if config['cuda']:
38 data = data.cuda()
39
40 optimizer.zero_grad()
41 model.CD_grad(data)
42 if config['clip_norm'] > 0:
43 nn.utils.clip_grad_norm_(model.parameters(), config['clip_norm'])
44 optimizer.step()
45
46 if ii == len(train_loader) - 1:
47 recon_loss = model.reconstruction(data).item()
48
49 return recon_loss
50
51
52def create_dataset(config):

Callers 1

train_modelFunction · 0.85

Calls 2

CD_gradMethod · 0.80
reconstructionMethod · 0.80

Tested by

no test coverage detected