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

Class Trainer

distributed/ddp-tutorial-series/multigpu.py:24–66  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

22 init_process_group(backend="nccl", rank=rank, world_size=world_size)
23
24class Trainer:
25 def __init__(
26 self,
27 model: torch.nn.Module,
28 train_data: DataLoader,
29 optimizer: torch.optim.Optimizer,
30 gpu_id: int,
31 save_every: int,
32 ) -> None:
33 self.gpu_id = gpu_id
34 self.model = model.to(gpu_id)
35 self.train_data = train_data
36 self.optimizer = optimizer
37 self.save_every = save_every
38 self.model = DDP(model, device_ids=[gpu_id])
39
40 def _run_batch(self, source, targets):
41 self.optimizer.zero_grad()
42 output = self.model(source)
43 loss = F.cross_entropy(output, targets)
44 loss.backward()
45 self.optimizer.step()
46
47 def _run_epoch(self, epoch):
48 b_sz = len(next(iter(self.train_data))[0])
49 print(f"[GPU{self.gpu_id}] Epoch {epoch} | Batchsize: {b_sz} | Steps: {len(self.train_data)}")
50 self.train_data.sampler.set_epoch(epoch)
51 for source, targets in self.train_data:
52 source = source.to(self.gpu_id)
53 targets = targets.to(self.gpu_id)
54 self._run_batch(source, targets)
55
56 def _save_checkpoint(self, epoch):
57 ckp = self.model.module.state_dict()
58 PATH = "checkpoint.pt"
59 torch.save(ckp, PATH)
60 print(f"Epoch {epoch} | Training checkpoint saved at {PATH}")
61
62 def train(self, max_epochs: int):
63 for epoch in range(max_epochs):
64 self._run_epoch(epoch)
65 if self.gpu_id == 0 and epoch % self.save_every == 0:
66 self._save_checkpoint(epoch)
67
68
69def load_train_objs():

Callers 1

mainFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected