MCPcopy Create free account
hub / github.com/AtlasAnalyticsLab/AdaFisher / train

Function train

Language_Model/run_exp.py:23–58  ·  view source on GitHub ↗
(epoch, model, dataloader, optimizer, args, gm = None)

Source from the content-addressed store, hash-verified

21from optimizers.sgd import SGD
22
23def train(epoch, model, dataloader, optimizer, args, gm = None):
24 model.train()
25 losses = []
26 total_iters = 0
27 start_time = time.time()
28 for idx, batch in enumerate(
29 tqdm(
30 dataloader, desc="Epoch {0}".format(epoch), disable=(not args.progress_bar)
31 )
32 ):
33 batch = to_device(batch, args.device)
34 if args.optimizer in ["Shampoo", "kfac"]:
35 dummy_y = gm.setup_model_call(model, batch["source"])
36 gm.setup_loss_call(model.loss, dummy_y, batch["target"], batch["mask"])
37 outputs, loss = gm.forward_and_backward()
38 torch.nn.utils.clip_grad_norm_(model.parameters(),
39 args.clip_norm)
40 else:
41 optimizer.zero_grad()
42 log_probas = model(batch["source"])
43 loss = model.loss(log_probas, batch["target"], batch["mask"])
44 losses.append(loss.item() * batch["mask"].sum().item())
45 if args.optimizer == 'AdaHessian':
46 loss.backward(create_graph=True)
47 else:
48 loss.backward()
49 optimizer.step()
50 total_iters += 1
51 if idx % args.print_every == 0:
52 tqdm.write(f"[TRAIN] Epoch: {epoch}, Iter : {idx} / {len(dataloader)}, Loss: {loss.item():.5f}")
53
54 mean_loss = np.mean(losses)
55 mean_loss /= args.batch_size * dataloader.dataset.max_length
56 perplexity = math.exp(mean_loss)
57 tqdm.write(f"== [TRAIN] Epoch: {epoch}, Perplexity: {perplexity:.3f} ==>")
58 return mean_loss, perplexity, time.time() - start_time
59
60
61def evaluate(epoch, model, dataloader, args, mode="val"):

Callers 1

mainFunction · 0.85

Calls 4

to_deviceFunction · 0.90
trainMethod · 0.80
lossMethod · 0.80
stepMethod · 0.45

Tested by

no test coverage detected