diff --git a/cosmos_rl/policy/model/dino_cradio_v3/cosmos_config.toml b/cosmos_rl/policy/model/dino_cradio_v3/cosmos_config.toml index 7be237bb5..af522d1ec 100644 --- a/cosmos_rl/policy/model/dino_cradio_v3/cosmos_config.toml +++ b/cosmos_rl/policy/model/dino_cradio_v3/cosmos_config.toml @@ -8,9 +8,9 @@ model_name_or_path = "nvidia/C-RADIOv3-g" tp_size = 1 cp_size = 1 ep_size = 1 -dp_shard_size = -1 # FSDP-only, autodetect shard size based on WORLD_SIZE +dp_shard_size = 1 pp_size = 1 -dp_replicate_size = 1 +dp_replicate_size = 8 [train] resume = false diff --git a/cosmos_rl/policy/model/dino_cradio_v3/parallelize.py b/cosmos_rl/policy/model/dino_cradio_v3/parallelize.py index 38303954b..10fd8a019 100644 --- a/cosmos_rl/policy/model/dino_cradio_v3/parallelize.py +++ b/cosmos_rl/policy/model/dino_cradio_v3/parallelize.py @@ -1,6 +1,10 @@ from typing import Optional, Callable import torch.nn as nn +from torch.distributed.device_mesh import DeviceMesh +from torch.distributed._composable.replicate import replicate + +from cosmos_rl.utils.logging import logger from cosmos_rl.utils.parallelism import ParallelDims from cosmos_rl.policy.config import Config as CosmosConfig @@ -11,7 +15,18 @@ def parallelize_model( config: CosmosConfig, pp_loss_fn: Optional[Callable], ) -> nn.Module: - if parallel_dims.world_size > 1: - raise NotImplementedError("Only single GPU runs supported!") + if parallel_dims.world_size == 1: + return None, None + + world_mesh = parallel_dims.mesh + # DDP + if parallel_dims.dp_replicate_enabled: + assert world_mesh.ndim == 1, "DDP does not support > 1D parallelism" + _apply_ddp(model, world_mesh) return None, None + + +def _apply_ddp(model: nn.Module, dp_mesh: DeviceMesh): + replicate(model, device_mesh=dp_mesh) + logger.info("Applied DDP to the model")