This function load state dict to model and optimizer on cpu, and move it back to device. This avoids a GPU memory surge issue. NOTE: assume checkpoint is loaded in cpu
(self, checkpoint, device, trainable=True)
| 22 | ) |
| 23 | |
| 24 | def resume_from_cpu(self, checkpoint, device, trainable=True): |
| 25 | """ |
| 26 | This function load state dict to model and optimizer on cpu, and move it back to device. |
| 27 | This avoids a GPU memory surge issue. |
| 28 | NOTE: assume checkpoint is loaded in cpu |
| 29 | """ |
| 30 | # handles model |
| 31 | self.model = self.model.cpu() |
| 32 | self.model.load_state_dict(checkpoint["model_state_dict"]) |
| 33 | self.model = self.model.to(device) |
| 34 | if trainable: |
| 35 | # handles optimizer |
| 36 | self.optimizer.load_state_dict(checkpoint["optimizer_state_dict"]) |
| 37 | optimizer_to(self.optimizer, device) |
| 38 | # possible extension: reinitialize scheduler based on this new optimizer |
| 39 | self.scheduler = self.scheduler = torch.optim.lr_scheduler.MultiStepLR( |
| 40 | self.optimizer, milestones=[50, 100, 150, 200], gamma=0.5 |
| 41 | ) |
| 42 | |
| 43 | # used by scene completion task |
| 44 | def step_completion(self, data, batch_size, loss_fn='ce', trainable=False): |