| 22 | init_process_group(backend="nccl", rank=rank, world_size=world_size) |
| 23 | |
| 24 | class 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 | |
| 69 | def load_train_objs(): |