(model: torch.nn.Module, config: DistributedConfig, model_type: ModelType, device_mesh: DeviceMesh)
| 35 | |
| 36 | |
| 37 | def prepare_model_for_distributed(model: torch.nn.Module, config: DistributedConfig, model_type: ModelType, device_mesh: DeviceMesh) -> torch.nn.Module: |
| 38 | if config.use_ddp: |
| 39 | return DDP(model, device_ids=[device_mesh.get_local_rank()], output_device=device_mesh.get_local_rank(), find_unused_parameters=False) |
| 40 | if config.use_fsdp: |
| 41 | mp_policy = config.get_mixed_precision_policy() |
| 42 | fsdp_kwargs = { |
| 43 | "reshard_after_forward": config.reshard_after_forward, |
| 44 | "mp_policy": mp_policy, |
| 45 | "offload_policy": config.offload_policy, |
| 46 | } |
| 47 | |
| 48 | def shard_layers(layers: Iterable[torch.nn.Module]) -> None: |
| 49 | for layer in layers: |
| 50 | fully_shard(layer, mesh=device_mesh, **fsdp_kwargs) |
| 51 | |
| 52 | if model_type in {ModelType.VideoTokenizer, ModelType.LatentActionModel}: |
| 53 | shard_layers(model.encoder.transformer.blocks) |
| 54 | |
| 55 | if model_type == ModelType.VideoTokenizer: |
| 56 | shard_layers(model.encoder.latent_head) |
| 57 | |
| 58 | if model_type == ModelType.LatentActionModel: |
| 59 | shard_layers(model.encoder.action_head) |
| 60 | shard_layers(model.decoder.transformer.blocks) |
| 61 | shard_layers(model.decoder.frame_head) |
| 62 | |
| 63 | fully_shard(model.encoder, mesh=device_mesh, **fsdp_kwargs) |
| 64 | fully_shard(model.decoder, mesh=device_mesh, **fsdp_kwargs) |
| 65 | fully_shard(model.quantizer, mesh=device_mesh, **fsdp_kwargs) |
| 66 | |
| 67 | elif model_type == ModelType.DynamicsModel: |
| 68 | shard_layers(model.transformer.blocks) |
| 69 | fully_shard(model.latent_embed, mesh=device_mesh, **fsdp_kwargs) |
| 70 | fully_shard(model.output_mlp, mesh=device_mesh, **fsdp_kwargs) |
| 71 | else: |
| 72 | raise ValueError('Unknown model type') |
| 73 | fully_shard( |
| 74 | model, |
| 75 | mesh=device_mesh, |
| 76 | **fsdp_kwargs, |
| 77 | ) |
| 78 | |
| 79 | return model |
| 80 | |
| 81 | |
| 82 | def unwrap_model(model: torch.nn.Module) -> torch.nn.Module: |
no test coverage detected