(
self,
model: torch.nn.Module,
train_data: DataLoader,
optimizer: torch.optim.Optimizer,
save_every: int,
snapshot_path: str,
)
| 16 | |
| 17 | class Trainer: |
| 18 | def __init__( |
| 19 | self, |
| 20 | model: torch.nn.Module, |
| 21 | train_data: DataLoader, |
| 22 | optimizer: torch.optim.Optimizer, |
| 23 | save_every: int, |
| 24 | snapshot_path: str, |
| 25 | ) -> None: |
| 26 | self.local_rank = int(os.environ["LOCAL_RANK"]) |
| 27 | self.global_rank = int(os.environ["RANK"]) |
| 28 | self.model = model.to(self.local_rank) |
| 29 | self.train_data = train_data |
| 30 | self.optimizer = optimizer |
| 31 | self.save_every = save_every |
| 32 | self.epochs_run = 0 |
| 33 | self.snapshot_path = snapshot_path |
| 34 | if os.path.exists(snapshot_path): |
| 35 | print("Loading snapshot") |
| 36 | self._load_snapshot(snapshot_path) |
| 37 | |
| 38 | self.model = DDP(self.model, device_ids=[self.local_rank]) |
| 39 | |
| 40 | def _load_snapshot(self, snapshot_path): |
| 41 | loc = f"cuda:{self.local_rank}" |
nothing calls this directly
no test coverage detected