MCPcopy Create free account
hub / github.com/InternLM/InternBootcamp / CheckpointEngineManager

Class CheckpointEngineManager

verl/verl/checkpoint_engine/base.py:285–409  ·  view source on GitHub ↗

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

Source from the content-addressed store, hash-verified

283
284
285class 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}"

Callers 8

test_server_adapterFunction · 0.90
init_workersMethod · 0.90
init_workersMethod · 0.90

Calls

no outgoing calls

Tested by 6

test_server_adapterFunction · 0.72