Checkpoint engine manager to coordinate weight synchronization between trainer and rollout replicas. - ME: model engine, FSDP, MCore, VeOmni, export full tensor generator `get_per_tensor_param` - CE: checkpoint engine, NCCL, NIXL, etc In trainer, model engine and checkpoint engine are
| 283 | |
| 284 | |
| 285 | class CheckpointEngineManager: |
| 286 | """Checkpoint engine manager to coordinate weight synchronization between trainer and rollout replicas. |
| 287 | |
| 288 | - ME: model engine, FSDP, MCore, VeOmni, export full tensor generator `get_per_tensor_param` |
| 289 | - CE: checkpoint engine, NCCL, NIXL, etc |
| 290 | |
| 291 | In trainer, model engine and checkpoint engine are in same process. |
| 292 | In rollout, checkpoint engine and rollout worker are in separate process, update weights via cuda ipc. |
| 293 | |
| 294 | ``` |
| 295 | ┌────────┬────────┬─────┬────────┐ ┌───────────────────┬───────────────────┐ |
| 296 | │ ┌────┐ │ ┌────┐ │ │ ┌────┐ │ │ Replica 0 │ Replica 1 │ |
| 297 | │ │ ME0│ │ │ ME1│ │ │ │ MEn│ │ ├────┬────┬────┬────┼────┬────┬────┬────┤ |
| 298 | │ └──┬─┘ │ └────┘ │ ... │ └────┘ │ │ 0 │ 1 │ 2 │ 3 │ 0 │ 1 │ 2 │ 3 │ |
| 299 | │ v | | | | └──┬─┴──┬─┴──┬─┴──┬─┴──┬─┴──┬─┴──┬─┴──┬─┘ |
| 300 | | ┌──┴─┐ │ ┌────┐ │ │ ┌────┐ │ ^ ^ ^ cuda ipc ^ ^ ^ |
| 301 | │ │ CE │ │ │ CE │ │ │ │ CE │ │ ┌──┴─┬──┴─┬──┴─┬──┴─┬──┴─┬──┴─┬──┴─┬──┴─┐ |
| 302 | │ └──┬─┘ │ └────┘ │ │ └────┘ │ │ CE │ CE │ CE │ CE │ CE │ CE │ CE │ CE | |
| 303 | └────┼───┴────────┴─────┴────────┘ └──┬─┴──┬─┴──┬─┴──┬─┴──┬─┴──┬─┴──┬─┴──┬─┘ |
| 304 | v | | | | | | | | |
| 305 | └─────────────(nccl/nixl/..)─────────────┴────┴────┴────┴────┴────┴────┴────┘ |
| 306 | ``` |
| 307 | |
| 308 | Args: |
| 309 | backend: The checkpoint engine backend. |
| 310 | trainer: The trainer worker group. |
| 311 | replicas: The list of rollout replicas. |
| 312 | """ |
| 313 | |
| 314 | def __init__( |
| 315 | self, |
| 316 | backend: str, |
| 317 | trainer: RayWorkerGroup, |
| 318 | replicas: list[RolloutReplica], |
| 319 | ) -> None: |
| 320 | self.backend = backend |
| 321 | self.backend_cls = CheckpointEngineRegistry.get(backend) |
| 322 | self.trainer = trainer |
| 323 | self.replicas = replicas |
| 324 | |
| 325 | def build_process_group(self, rollout: RayWorkerGroup): |
| 326 | """Build process group for trainer and rollout replicas.""" |
| 327 | trainer = self.trainer |
| 328 | |
| 329 | # 1. prepare all workers |
| 330 | metadata = ray.get( |
| 331 | trainer.execute_checkpoint_engine(["prepare"] * trainer.world_size) |
| 332 | + rollout.execute_checkpoint_engine(["prepare"] * rollout.world_size) |
| 333 | ) |
| 334 | |
| 335 | # 2. build communication topology between all workers |
| 336 | trainer_kwargs, rollout_kwargs = self.backend_cls.build_topology( |
| 337 | trainer.world_size, rollout.world_size, metadata |
| 338 | ) |
| 339 | for k, v in trainer_kwargs.items(): |
| 340 | assert len(v) == trainer.world_size, f"trainer_kwargs[{k}] must have length of {trainer.world_size}" |
| 341 | for k, v in rollout_kwargs.items(): |
| 342 | assert len(v) == rollout.world_size, f"rollout_kwargs[{k}] must have length of {rollout.world_size}" |
no outgoing calls