| 137 | |
| 138 | @staticmethod |
| 139 | def load_optimizer(model, optimizer, path): |
| 140 | if os.path.isdir(path): |
| 141 | # dcp |
| 142 | state_dict = get_optimizer_state_dict( |
| 143 | model, |
| 144 | optimizer, |
| 145 | options=StateDictOptions( |
| 146 | full_state_dict=False, |
| 147 | cpu_offload=True, |
| 148 | ), |
| 149 | ) |
| 150 | DCP.load(state_dict=state_dict, checkpoint_id=path) |
| 151 | else: |
| 152 | state_dict = torch.load(path, map_location='cpu') |
| 153 | set_optimizer_state_dict( |
| 154 | model, |
| 155 | optimizer, |
| 156 | optim_state_dict=state_dict, |
| 157 | options=StateDictOptions(full_state_dict=False, strict=True), |
| 158 | ) |
| 159 | if dist.is_initialized(): |
| 160 | dist.barrier() |
| 161 | |
| 162 | @staticmethod |
| 163 | def save_model(rank, model, out_path: str, dcp=False, lora=False): |