import torch.distributed.algorithms.ddp_comm_hooks.default_hooks as ddp_hks def configure_ddp(model, rank): dist.init_process_group("nccl") model = DDP(model, device_ids=[rank], static_graph=True, gradient_as_bucket_view=True, bucket_cap_mb=2000) model.register_comm_hook(state=None, hook=ddp_hks.bf16_compress_hook) return model