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

Method __init__

distributed/ddp-tutorial-series/multinode.py:18–38  ·  view source on GitHub ↗
(
        self,
        model: torch.nn.Module,
        train_data: DataLoader,
        optimizer: torch.optim.Optimizer,
        save_every: int,
        snapshot_path: str,
    )

Source from the content-addressed store, hash-verified

16
17class 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}"

Callers

nothing calls this directly

Calls 1

_load_snapshotMethod · 0.95

Tested by

no test coverage detected