| 203 | return step_num # type: ignore |
| 204 | |
| 205 | async def _pull_latest_weights(self): |
| 206 | self.logger.info("Start to pull latest model weights.") |
| 207 | new_version = await self.synchronizer.wait_new_model_state_dict.remote( |
| 208 | current_version=self.model_version, |
| 209 | ) |
| 210 | if new_version > self.model_version: |
| 211 | if self.model_version != -1 or new_version > 0: |
| 212 | self.logger.info(f"New model weights version: {new_version}") |
| 213 | await asyncio.gather( |
| 214 | *[ |
| 215 | model.sync_model_weights( |
| 216 | new_version, |
| 217 | self.config.synchronizer.sync_method, |
| 218 | timeout=self.config.synchronizer.sync_timeout, |
| 219 | ) |
| 220 | for model in self.models |
| 221 | ] |
| 222 | ) |
| 223 | self.model_version = new_version |
| 224 | else: |
| 225 | self.logger.warning( |
| 226 | f"No new model weights found, current version: {self.model_version}" |
| 227 | ) |
| 228 | |
| 229 | async def _nccl_weights_update(self): |
| 230 | new_version = await self.synchronizer.ready_to_nccl_sync.remote( |