Is there a way to replace the 'allreduce_hook' used for DDP(DistributedDataParallel) in Pytorch?

Viewed 125

I know that Pytorch DDP uses 'allreduce_hook' as the default communication hook. Is there a way to replace this default hook with 'quantization_pertensor_hook' or 'powerSGD_hook'. There is an official Pytorch documentation introducing the DDP communication hooks, but I still got confused about how to do this in practice.

This is how I initiate the process group and create the DDP model.

import torch.distributed as dist
import torch.nn as nn

dist.init_process_group(backend='nccl', init_method='env://', world_size=args.world_size, rank=rank)
model = nn.parallel.DistributedDataParallel(model, device_ids=[0])

Is there any way to declare the hook that I want based on this code?

1 Answers

This could do the job


dist.init_process_group(backend='nccl', init_method='env://', world_size=args.world_size, rank=rank)
model = nn.parallel.DistributedDataParallel(model, device_ids=[0])

state = powerSGD.PowerSGDState(process_group=None, matrix_approximation_rank=1, start_powerSGD_iter=10, min_compression_rate=0.5)
model.register_comm_hook(state, powerSGD.powerSGD_hook)
...
Related