MCPcopy Create free account
hub / github.com/pytorch/examples / Trainer

Class Trainer

distributed/ddp-tutorial-series/single_gpu.py:7–47  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

5
6
7class Trainer:
8 def __init__(
9 self,
10 model: torch.nn.Module,
11 train_data: DataLoader,
12 optimizer: torch.optim.Optimizer,
13 gpu_id: int,
14 save_every: int,
15 ) -> None:
16 self.gpu_id = gpu_id
17 self.model = model.to(gpu_id)
18 self.train_data = train_data
19 self.optimizer = optimizer
20 self.save_every = save_every
21
22 def _run_batch(self, source, targets):
23 self.optimizer.zero_grad()
24 output = self.model(source)
25 loss = F.cross_entropy(output, targets)
26 loss.backward()
27 self.optimizer.step()
28
29 def _run_epoch(self, epoch):
30 b_sz = len(next(iter(self.train_data))[0])
31 print(f"[GPU{self.gpu_id}] Epoch {epoch} | Batchsize: {b_sz} | Steps: {len(self.train_data)}")
32 for source, targets in self.train_data:
33 source = source.to(self.gpu_id)
34 targets = targets.to(self.gpu_id)
35 self._run_batch(source, targets)
36
37 def _save_checkpoint(self, epoch):
38 ckp = self.model.state_dict()
39 PATH = "checkpoint.pt"
40 torch.save(ckp, PATH)
41 print(f"Epoch {epoch} | Training checkpoint saved at {PATH}")
42
43 def train(self, max_epochs: int):
44 for epoch in range(max_epochs):
45 self._run_epoch(epoch)
46 if epoch % self.save_every == 0:
47 self._save_checkpoint(epoch)
48
49
50def load_train_objs():

Callers 1

mainFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected