Save policy, optimizer, and scheduler state to disk.
(self, output_dir: Optional[str] = None, metrics: Optional[Dict] = None)
| 413 | }, output_path) |
| 414 | |
| 415 | def save(self, output_dir: Optional[str] = None, metrics: Optional[Dict] = None): |
| 416 | """Save policy, optimizer, and scheduler state to disk.""" |
| 417 | |
| 418 | policy_state_dict = self.policy.state_dict() |
| 419 | self.write_state_dict(self.example_counter, policy_state_dict, metrics, 'policy.pt', output_dir) |
| 420 | del policy_state_dict |
| 421 | |
| 422 | optimizer_state_dict = self.optimizer.state_dict() |
| 423 | self.write_state_dict(self.example_counter, optimizer_state_dict, metrics, 'optimizer.pt', output_dir) |
| 424 | del optimizer_state_dict |
| 425 | |
| 426 | scheduler_state_dict = self.scheduler.state_dict() |
| 427 | self.write_state_dict(self.example_counter, scheduler_state_dict, metrics, 'scheduler.pt', output_dir) |
| 428 | |
| 429 | |
| 430 | class FSDPTrainer(BasicTrainer): |