Save (unsharded) policy state to disk.
(self, output_dir=None, metrics=None)
| 536 | self.reference_model = tp.tensor_parallel(reference_model, sharded=False) |
| 537 | |
| 538 | def save(self, output_dir=None, metrics=None): |
| 539 | """Save (unsharded) policy state to disk.""" |
| 540 | with tp.save_tensor_parallel(self.policy): |
| 541 | policy_state_dict = self.policy.state_dict() |
| 542 | |
| 543 | self.write_state_dict(self.example_counter, policy_state_dict, metrics, 'policy.pt', output_dir) |
| 544 | del policy_state_dict |
| 545 |
nothing calls this directly
no test coverage detected