(data, group)
| 106 | |
| 107 | |
| 108 | def _serialize_to_tensor(data, group): |
| 109 | backend = dist.get_backend(group) |
| 110 | assert backend in ["gloo", "nccl"] |
| 111 | device = torch.device("cpu" if backend == "gloo" else "cuda") |
| 112 | |
| 113 | buffer = pickle.dumps(data) |
| 114 | if len(buffer) > 1024 ** 3: |
| 115 | logger = logging.getLogger(__name__) |
| 116 | logger.warning( |
| 117 | "Rank {} trying to all-gather {:.2f} GB of data on device {}".format( |
| 118 | get_rank(), len(buffer) / (1024 ** 3), device |
| 119 | ) |
| 120 | ) |
| 121 | storage = torch.ByteStorage.from_buffer(buffer) |
| 122 | tensor = torch.ByteTensor(storage).to(device=device) |
| 123 | return tensor |
| 124 | |
| 125 | |
| 126 | def _pad_to_largest_tensor(tensor, group): |
no test coverage detected