chore: import upstream snapshot with attribution
This commit is contained in:
@@ -0,0 +1,5 @@
|
||||
# Copyright (c) Microsoft Corporation.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# DeepSpeed Team
|
||||
'''Copyright The Microsoft DeepSpeed Team'''
|
||||
@@ -0,0 +1,218 @@
|
||||
# Copyright (c) Microsoft Corporation.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# DeepSpeed Team
|
||||
"""
|
||||
batched collective operations for overhead amortization and better
|
||||
bandwidth utilization
|
||||
"""
|
||||
|
||||
import math
|
||||
from typing import List, Any
|
||||
import torch
|
||||
from torch import Tensor
|
||||
from deepspeed import comm as dist
|
||||
from deepspeed.comm import ProcessGroup, all_to_all_single
|
||||
from deepspeed.accelerator import get_accelerator
|
||||
from deepspeed.utils import instrument_w_nvtx
|
||||
from deepspeed.ops import op_builder
|
||||
from deepspeed.utils import logger
|
||||
|
||||
|
||||
def _torch_reduce_scatter_fn(input_tensor: Tensor, output_tensor: Tensor, group=None, async_op=False, prof=False):
|
||||
return instrument_w_nvtx(dist.reduce_scatter_fn)(output_tensor, input_tensor, group=group, async_op=False)
|
||||
|
||||
|
||||
quantizer_module = None
|
||||
|
||||
|
||||
@instrument_w_nvtx
|
||||
@torch.no_grad()
|
||||
def all_to_all_quant_reduce(tensors: List[Tensor], groups: {}) -> List[Tensor]:
|
||||
global quantizer_module
|
||||
if quantizer_module is None:
|
||||
quantizer_module = op_builder.QuantizerBuilder().load()
|
||||
local_world_size = get_accelerator().device_count()
|
||||
global_world_size = dist.get_world_size()
|
||||
num_nodes = global_world_size // local_world_size
|
||||
this_rank = dist.get_rank()
|
||||
intra_idx = int(this_rank / local_world_size)
|
||||
inter_idx = this_rank % local_world_size
|
||||
output_lst: List[Tensor] = [None] * len(tensors)
|
||||
for idx, tensor in enumerate(tensors):
|
||||
if tensor.dim() == 1:
|
||||
output_lst[idx] = reduce_scatter_coalesced([tensor])[0]
|
||||
elif tensor.numel() % (2 * global_world_size) != 0:
|
||||
# Due to the constraint of 2-stage all-to-all, the input tensor must be divisible by 2 * global_world_size
|
||||
# Otherwise, all-to-all cannot be performed because of shape mismatch.
|
||||
# See more at https://github.com/deepspeedai/DeepSpeed/pull/5056
|
||||
logger.warning(
|
||||
f"qgZ falls back to reduce_scatter because tensor size = {tensor.numel()} is not divisible by (2 * global_world_size) = {2 * global_world_size}. Please consider allocating a new world to enable qgZ"
|
||||
)
|
||||
output_lst[idx] = reduce_scatter_coalesced([tensor])[0]
|
||||
else:
|
||||
intra_quant_group = max(tensor.shape[0], tensor.shape[1], global_world_size)
|
||||
|
||||
inter_quant_group = intra_quant_group // local_world_size
|
||||
intra_quant_int4, intra_q_scales = quantizer_module.swizzle_quant(tensor, intra_quant_group, 4,
|
||||
quantizer_module.Symmetric, 1, num_nodes,
|
||||
local_world_size)
|
||||
local_output = torch.empty_like(intra_quant_int4)
|
||||
scale_output = torch.empty_like(intra_q_scales)
|
||||
all_to_all_single(local_output, intra_quant_int4, group=groups[f'local_{intra_idx}'])
|
||||
all_to_all_single(scale_output, intra_q_scales, group=groups[f'local_{intra_idx}'])
|
||||
global_input_tensor, global_scales = quantizer_module.quantized_reduction(
|
||||
local_output, scale_output, intra_quant_group, inter_quant_group, 4, quantizer_module.Symmetric,
|
||||
local_world_size)
|
||||
global_output = torch.empty_like(global_input_tensor)
|
||||
global_scale_output = torch.empty_like(global_scales)
|
||||
all_to_all_single(global_output, global_input_tensor, group=groups[f'global_{inter_idx}'])
|
||||
all_to_all_single(global_scale_output, global_scales, group=groups[f'global_{inter_idx}'])
|
||||
final_output = quantizer_module.dequantize(global_output, global_scale_output, global_scale_output.numel(),
|
||||
4, quantizer_module.Symmetric)
|
||||
assert final_output.numel(
|
||||
) % num_nodes == 0, f"final_output.numel()={final_output.numel()} is not divisible by num_nodes={num_nodes}"
|
||||
output_lst[idx] = (sum(list(final_output.chunk(num_nodes))) / num_nodes).view(-1)
|
||||
return output_lst
|
||||
|
||||
|
||||
@instrument_w_nvtx
|
||||
@torch.no_grad()
|
||||
def all_to_all_loco_quant_reduce(
|
||||
params: List[Tensor],
|
||||
groups: {},
|
||||
loco_param: Any = None,
|
||||
) -> List[Tensor]:
|
||||
global quantizer_module
|
||||
global loco_idx
|
||||
if quantizer_module is None:
|
||||
quantizer_module = op_builder.QuantizerBuilder().load()
|
||||
local_world_size = get_accelerator().device_count()
|
||||
global_world_size = dist.get_world_size()
|
||||
num_nodes = global_world_size // local_world_size
|
||||
this_rank = dist.get_rank()
|
||||
intra_idx = int(this_rank / local_world_size)
|
||||
inter_idx = this_rank % local_world_size
|
||||
output_lst: List[Tensor] = [None] * len(params)
|
||||
for idx, p in enumerate(params):
|
||||
tensor = p.grad
|
||||
if tensor.dim() == 1:
|
||||
output_lst[idx] = reduce_scatter_coalesced([tensor])[0]
|
||||
elif tensor.numel() % (2 * global_world_size) != 0:
|
||||
# Due to the constraint of 2-stage all-to-all, the input tensor must be divisible by 2 * global_world_size
|
||||
# Otherwise, all-to-all cannot be performed because of shape mismatch.
|
||||
# See more at https://github.com/deepspeedai/DeepSpeed/pull/5056
|
||||
logger.warning(
|
||||
f"qgZ falls back to reduce_scatter because tensor size = {tensor.numel()} is not divisible by (2 * global_world_size) = {2 * global_world_size}. Please consider allocating a new world to enable qgZ"
|
||||
)
|
||||
output_lst[idx] = reduce_scatter_coalesced([tensor])[0]
|
||||
else:
|
||||
err_beta = loco_param['err_beta']
|
||||
reset_T = loco_param['reset_T']
|
||||
if not hasattr(p, 'intra_ef_buf') or loco_idx > reset_T:
|
||||
loco_idx = 0
|
||||
intra_err = torch.zeros_like(p.grad)
|
||||
inter_err = torch.zeros(tensor.numel() // local_world_size, device=tensor.device, dtype=tensor.dtype)
|
||||
else:
|
||||
intra_err = quantizer_module.dequantize(p.intra_ef_buf[0], p.intra_ef_buf[1],
|
||||
p.intra_ef_buf[1].numel(), 8, quantizer_module.Symmetric)
|
||||
inter_err = quantizer_module.dequantize(p.inter_ef_buf[0], p.inter_ef_buf[1],
|
||||
p.inter_ef_buf[1].numel(), 8, quantizer_module.Symmetric)
|
||||
|
||||
intra_quant_group = max(tensor.shape[0], tensor.shape[1], global_world_size)
|
||||
inter_quant_group = intra_quant_group // local_world_size
|
||||
intra_quant_int4, intra_q_scales = quantizer_module.loco_swizzle_quant(tensor, intra_err, err_beta,
|
||||
intra_quant_group, 4,
|
||||
quantizer_module.Symmetric, 1,
|
||||
num_nodes, local_world_size)
|
||||
local_output = torch.empty_like(intra_quant_int4)
|
||||
scale_output = torch.empty_like(intra_q_scales)
|
||||
all_to_all_single(local_output, intra_quant_int4, group=groups[f'local_{intra_idx}'])
|
||||
all_to_all_single(scale_output, intra_q_scales, group=groups[f'local_{intra_idx}'])
|
||||
|
||||
p.intra_ef_buf = quantizer_module.quantize(intra_err, intra_quant_group, 8, quantizer_module.Symmetric)
|
||||
|
||||
global_input_tensor, global_scales = quantizer_module.loco_quantized_reduction(
|
||||
local_output, scale_output, inter_err, err_beta, intra_quant_group, inter_quant_group, 4,
|
||||
quantizer_module.Symmetric, local_world_size)
|
||||
|
||||
global_output = torch.empty_like(global_input_tensor)
|
||||
global_scale_output = torch.empty_like(global_scales)
|
||||
all_to_all_single(global_output, global_input_tensor, group=groups[f'global_{inter_idx}'])
|
||||
all_to_all_single(global_scale_output, global_scales, group=groups[f'global_{inter_idx}'])
|
||||
|
||||
p.inter_ef_buf = quantizer_module.quantize(inter_err, inter_quant_group, 8, quantizer_module.Symmetric)
|
||||
|
||||
final_output = quantizer_module.dequantize(global_output, global_scale_output, global_scale_output.numel(),
|
||||
4, quantizer_module.Symmetric)
|
||||
assert final_output.numel(
|
||||
) % num_nodes == 0, f"final_output.numel()={final_output.numel()} is not divisible by num_nodes={num_nodes}"
|
||||
output_lst[idx] = (sum(list(final_output.chunk(num_nodes))) / num_nodes).view(-1)
|
||||
loco_idx += 1
|
||||
|
||||
return output_lst
|
||||
|
||||
|
||||
@instrument_w_nvtx
|
||||
@torch.no_grad()
|
||||
def reduce_scatter_coalesced(
|
||||
tensors: List[Tensor],
|
||||
group: ProcessGroup = None,
|
||||
) -> List[Tensor]:
|
||||
"""simultaneously reduce-scatter a list of tensors - this can be done more
|
||||
efficiently than individual reduce scatter calls
|
||||
TODO. see if PyTorch team wants a c++ version of this for ProcessGroupNCCL
|
||||
"""
|
||||
this_rank = dist.get_rank(group)
|
||||
world_sz = dist.get_world_size(group)
|
||||
|
||||
partition_lst_for_each_tensor = [None] * len(tensors)
|
||||
for tensor_idx, tensor in enumerate(tensors):
|
||||
flattened_tensor = tensor.view(-1)
|
||||
chunk_sz = math.ceil(tensor.numel() / world_sz)
|
||||
partition_lst_for_each_tensor[tensor_idx] = [
|
||||
flattened_tensor[rank * chunk_sz:rank * chunk_sz + chunk_sz] for rank in range(0, world_sz)
|
||||
]
|
||||
|
||||
padded_partition_sz_for_each_tensor = tuple(math.ceil(t.numel() / world_sz) for t in tensors)
|
||||
|
||||
if len(tensors) == 1 and tensors[0].numel() % world_sz == 0:
|
||||
# if there's only one tensor being reduced and we don't need to pad
|
||||
# we have an opportunity to avoid a memory allocation
|
||||
tensor_partition_flat_buffer = tensors[0].view(-1)
|
||||
else:
|
||||
# interleave tensor partitions such that the correct reduced partitions of each tensor
|
||||
# end up at each rank
|
||||
tensor_partitions_lst_with_padding = []
|
||||
for rank in range(world_sz):
|
||||
for tensor_idx in range(len(tensors)):
|
||||
# add tensor content
|
||||
tensor_chunk = partition_lst_for_each_tensor[tensor_idx][rank]
|
||||
tensor_partitions_lst_with_padding.append(tensor_chunk)
|
||||
|
||||
# add padding if necessary
|
||||
padding_sz = padded_partition_sz_for_each_tensor[tensor_idx] - tensor_chunk.numel()
|
||||
if padding_sz > 0:
|
||||
tensor_partitions_lst_with_padding.append(
|
||||
torch.empty(padding_sz, dtype=tensor_chunk.dtype, device=tensor_chunk.device))
|
||||
|
||||
tensor_partition_flat_buffer = instrument_w_nvtx(torch.cat)(tensor_partitions_lst_with_padding)
|
||||
|
||||
tensor_partition_flat_buffer.div_(world_sz) # pre-divide
|
||||
tensor_partition_buffer_for_each_rank: List[Tensor] = torch.chunk(tensor_partition_flat_buffer, world_sz)
|
||||
|
||||
# batched reduce-scatter call
|
||||
_torch_reduce_scatter_fn(tensor_partition_flat_buffer,
|
||||
tensor_partition_buffer_for_each_rank[this_rank],
|
||||
group=group)
|
||||
|
||||
# reverse procedure of the interleaving done previously, done on the
|
||||
# result of the batched reduce-scatter
|
||||
output_lst: List[Tensor] = [None] * len(tensors)
|
||||
offset = 0
|
||||
for tensor_idx in range(len(tensors)):
|
||||
output_lst[tensor_idx] = tensor_partition_buffer_for_each_rank[this_rank].narrow(
|
||||
0, offset, partition_lst_for_each_tensor[tensor_idx][this_rank].numel())
|
||||
|
||||
offset += padded_partition_sz_for_each_tensor[tensor_idx]
|
||||
return output_lst
|
||||
@@ -0,0 +1,141 @@
|
||||
# Copyright (c) Microsoft Corporation.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# DeepSpeed Team
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import deepspeed.comm as dist
|
||||
from deepspeed.accelerator import get_accelerator
|
||||
from deepspeed.ops.op_builder import PackbitsBuilder
|
||||
from deepspeed.runtime.comm.utils import check_and_handle_empty_buffer
|
||||
|
||||
|
||||
class CompressedBackend(object):
|
||||
|
||||
def __init__(self, mpu=None):
|
||||
if mpu is None:
|
||||
self.world_group = dist.new_group(ranks=range(dist.get_world_size()))
|
||||
else:
|
||||
self.mpu = mpu
|
||||
self.world_group = self.mpu.get_data_parallel_group()
|
||||
self.size = dist.get_world_size(group=self.world_group)
|
||||
self.rank = dist.get_rank(group=self.world_group)
|
||||
self.packer = PackbitsBuilder().load()
|
||||
|
||||
def my_igather(self, rank, size, group, sendbuf, recvbuf, root):
|
||||
req = []
|
||||
if rank == root:
|
||||
for idx in range(size):
|
||||
if idx != rank:
|
||||
req.append(dist.irecv(recvbuf[idx], src=idx, group=group))
|
||||
else:
|
||||
recvbuf[rank] = sendbuf
|
||||
else:
|
||||
req.append(dist.isend(sendbuf, group=group, dst=root))
|
||||
return req
|
||||
|
||||
def my_gather(self, rank, size, group, sendbuf, recvbuf, root):
|
||||
if rank == root:
|
||||
for idx in range(size):
|
||||
if idx != rank:
|
||||
dist.recv(recvbuf[idx], src=idx, group=group)
|
||||
else:
|
||||
recvbuf[rank] = sendbuf
|
||||
else:
|
||||
dist.send(sendbuf, group=group, dst=root)
|
||||
|
||||
def pack(self, buffer, size):
|
||||
# pack float tensor into uint8 tensor
|
||||
packed = self.packer.packbits(buffer.float(), buffer.numel(), self.rank)
|
||||
return packed.reshape(size, -1)
|
||||
|
||||
def unpack(self, buffer, size, dtype):
|
||||
# unpack uint8 to float tensor
|
||||
unpacked = self.packer.unpackbits(buffer, buffer.numel(), self.rank)
|
||||
return unpacked.reshape(size, -1).to(dtype)
|
||||
|
||||
def compressed_allreduce(self, buffer_m: torch.tensor, worker_error, server_error, local_rank):
|
||||
original_shape = buffer_m.size()
|
||||
if len(original_shape) > 1:
|
||||
buffer_m = torch.flatten(buffer_m)
|
||||
|
||||
# align size of original_buffer and error
|
||||
original_size = buffer_m.numel()
|
||||
worker_error_size = worker_error.numel()
|
||||
result = check_and_handle_empty_buffer(buffer_m, original_shape, original_size, worker_error, server_error)
|
||||
if result is not None:
|
||||
return result
|
||||
if original_size != worker_error_size:
|
||||
empty_tensor = torch.zeros(worker_error_size - original_size, device=buffer_m.device)
|
||||
buffer_m = torch.cat([buffer_m, empty_tensor])
|
||||
|
||||
buffer_m.add_(worker_error)
|
||||
worker_scale = torch.linalg.norm(buffer_m) / np.sqrt(torch.numel(buffer_m))
|
||||
|
||||
worker_error.set_(buffer_m - worker_scale * buffer_m.sign().add_(1).bool().float().add_(-0.5).mul_(2.0))
|
||||
|
||||
sign_list_packed_tmp = self.pack(buffer_m, self.size).type(torch.int8)
|
||||
|
||||
recvbuf_sign = torch.zeros([self.size, len(sign_list_packed_tmp[self.rank])],
|
||||
dtype=sign_list_packed_tmp[0].dtype,
|
||||
device=sign_list_packed_tmp.device)
|
||||
|
||||
sign_list_packed = [sign_list_packed_tmp[idx] for idx in range(self.size)]
|
||||
|
||||
recvbuf_scale = [
|
||||
torch.zeros(1, dtype=worker_scale.dtype, device=get_accelerator().current_device_name())
|
||||
for _ in range(self.size)
|
||||
]
|
||||
|
||||
# communication phase 1
|
||||
# all to all for sign
|
||||
dist.all_to_all_single(recvbuf_sign, torch.stack(sign_list_packed), group=self.world_group)
|
||||
# all gather for scale
|
||||
dist.all_gather(recvbuf_scale, worker_scale, group=self.world_group)
|
||||
|
||||
flattened_recvbuf_sign = recvbuf_sign.type(torch.uint8).flatten()
|
||||
compensated_server_m = self.unpack(flattened_recvbuf_sign, self.size, torch.float32) \
|
||||
.mul_(torch.stack(recvbuf_scale).mul_(1 / self.size)).sum(0)
|
||||
|
||||
compensated_server_m.add_(server_error)
|
||||
|
||||
server_scale = torch.linalg.norm(compensated_server_m) / np.sqrt(compensated_server_m.numel())
|
||||
|
||||
server_error.set_(compensated_server_m -
|
||||
server_scale * compensated_server_m.sign().add_(1).bool().float().add_(-0.5).mul_(2.0))
|
||||
|
||||
server_sign_packed = self.pack(compensated_server_m, 1).type(torch.int8)
|
||||
|
||||
# recvbuf_sign_server
|
||||
recvbuf_sign_server_tmp = torch.zeros([self.size, len(server_sign_packed[0])],
|
||||
dtype=recvbuf_sign.dtype,
|
||||
device=server_sign_packed.device)
|
||||
|
||||
recvbuf_sign_server = [recvbuf_sign_server_tmp[idx] for idx in range(self.size)]
|
||||
|
||||
# recvbuf_scale_server
|
||||
recvbuf_scale_server_tmp = torch.zeros([self.size, 1],
|
||||
dtype=worker_scale.dtype,
|
||||
device=server_sign_packed.device)
|
||||
|
||||
recvbuf_scale_server = [recvbuf_scale_server_tmp[idx] for idx in range(self.size)]
|
||||
|
||||
# communication Phase 2
|
||||
dist.all_gather(recvbuf_sign_server, server_sign_packed[0], group=self.world_group)
|
||||
dist.all_gather(recvbuf_scale_server, server_scale, group=self.world_group)
|
||||
|
||||
recvbuf_sign_server = torch.stack(recvbuf_sign_server)
|
||||
|
||||
flattened_recvbuf_sign_server = recvbuf_sign_server.type(torch.uint8).flatten()
|
||||
|
||||
buffer_m.data.copy_(
|
||||
self.unpack(flattened_recvbuf_sign_server, self.size,
|
||||
torch.float32).mul_(recvbuf_scale_server_tmp).flatten().data)
|
||||
|
||||
if original_size != worker_error_size:
|
||||
buffer_m = buffer_m[0:original_size]
|
||||
if len(original_shape) > 1:
|
||||
buffer_m = buffer_m.reshape(original_shape)
|
||||
|
||||
return buffer_m
|
||||
@@ -0,0 +1,129 @@
|
||||
# Copyright (c) Microsoft Corporation.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# DeepSpeed Team
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch_npu
|
||||
|
||||
import deepspeed.comm as dist
|
||||
from deepspeed.runtime.comm.utils import check_and_handle_empty_buffer
|
||||
|
||||
|
||||
class HcclBackend(object):
|
||||
|
||||
def __init__(self, mpu=None):
|
||||
if mpu is None:
|
||||
self.world_group = dist.new_group(ranks=range(dist.get_world_size()))
|
||||
else:
|
||||
self.mpu = mpu
|
||||
self.world_group = self.mpu.get_data_parallel_group()
|
||||
self.size = dist.get_world_size(group=self.world_group)
|
||||
self.rank = dist.get_rank(group=self.world_group)
|
||||
|
||||
def my_igather(self, rank, size, group, sendbuf, recvbuf, root):
|
||||
req = []
|
||||
if rank == root:
|
||||
for idx in range(size):
|
||||
if idx != rank:
|
||||
req.append(dist.irecv(recvbuf[idx], src=idx, group=group))
|
||||
else:
|
||||
recvbuf[rank] = sendbuf
|
||||
else:
|
||||
req.append(dist.isend(sendbuf, group=group, dst=root))
|
||||
return req
|
||||
|
||||
def my_gather(self, rank, size, group, sendbuf, recvbuf, root):
|
||||
if rank == root:
|
||||
for idx in range(size):
|
||||
if idx != rank:
|
||||
dist.recv(recvbuf[idx], src=idx, group=group)
|
||||
else:
|
||||
recvbuf[rank] = sendbuf
|
||||
else:
|
||||
dist.send(sendbuf, group=group, dst=root)
|
||||
|
||||
def compressed_allreduce(self, buffer_m: torch.tensor, worker_error, server_error, local_rank):
|
||||
original_shape = buffer_m.size()
|
||||
if len(original_shape) > 1:
|
||||
buffer_m = torch.flatten(buffer_m)
|
||||
|
||||
# align size of original_buffer and error
|
||||
original_size = buffer_m.numel()
|
||||
worker_error_size = worker_error.numel()
|
||||
result = check_and_handle_empty_buffer(buffer_m, original_shape, original_size, worker_error, server_error)
|
||||
if result is not None:
|
||||
return result
|
||||
if original_size != worker_error_size:
|
||||
empty_tensor = torch.zeros(worker_error_size - original_size, device=buffer_m.device)
|
||||
buffer_m = torch.cat([buffer_m, empty_tensor])
|
||||
|
||||
buffer_m.add_(worker_error)
|
||||
worker_scale = torch.linalg.norm(buffer_m) / np.sqrt(torch.numel(buffer_m))
|
||||
|
||||
worker_error.set_(buffer_m - worker_scale * buffer_m.sign().add_(1).bool().float().add_(-0.5).mul_(2.0))
|
||||
|
||||
sign_list_packed_tmp = torch_npu.npu_sign_bits_pack(buffer_m, self.size).type(torch.int8)
|
||||
|
||||
recvbuf_sign = torch.zeros([self.size, len(sign_list_packed_tmp[self.rank])],
|
||||
dtype=sign_list_packed_tmp[0].dtype,
|
||||
device=sign_list_packed_tmp.device)
|
||||
|
||||
sign_list_packed = [sign_list_packed_tmp[idx] for idx in range(self.size)]
|
||||
|
||||
recvbuf_scale = [
|
||||
torch.zeros(1, dtype=worker_scale.dtype, device=torch.device(local_rank)) for _ in range(self.size)
|
||||
]
|
||||
|
||||
# communication phase 1
|
||||
# all to all for sign
|
||||
dist.all_to_all_single(recvbuf_sign, torch.stack(sign_list_packed), group=self.world_group)
|
||||
# all gather for scale
|
||||
dist.all_gather(recvbuf_scale, worker_scale, group=self.world_group)
|
||||
|
||||
flattened_recvbuf_sign = recvbuf_sign.type(torch.uint8).flatten()
|
||||
compensated_server_m = torch_npu.npu_sign_bits_unpack(flattened_recvbuf_sign, self.size, torch.float32) \
|
||||
.mul_(torch.stack(recvbuf_scale).mul_(1 / self.size)).sum(0)
|
||||
|
||||
compensated_server_m.add_(server_error)
|
||||
|
||||
server_scale = torch.linalg.norm(compensated_server_m) / np.sqrt(compensated_server_m.numel())
|
||||
|
||||
server_error.set_(compensated_server_m -
|
||||
server_scale * compensated_server_m.sign().add_(1).bool().float().add_(-0.5).mul_(2.0))
|
||||
|
||||
server_sign_packed = torch_npu.npu_sign_bits_pack(compensated_server_m, 1).type(torch.int8)
|
||||
|
||||
# recvbuf_sign_server
|
||||
recvbuf_sign_server_tmp = torch.zeros([self.size, len(server_sign_packed[0])],
|
||||
dtype=recvbuf_sign.dtype,
|
||||
device=server_sign_packed.device)
|
||||
|
||||
recvbuf_sign_server = [recvbuf_sign_server_tmp[idx] for idx in range(self.size)]
|
||||
|
||||
# recvbuf_scale_server
|
||||
recvbuf_scale_server_tmp = torch.zeros([self.size, 1],
|
||||
dtype=worker_scale.dtype,
|
||||
device=server_sign_packed.device)
|
||||
|
||||
recvbuf_scale_server = [recvbuf_scale_server_tmp[idx] for idx in range(self.size)]
|
||||
|
||||
# communication Phase 2
|
||||
dist.all_gather(recvbuf_sign_server, server_sign_packed[0], group=self.world_group)
|
||||
dist.all_gather(recvbuf_scale_server, server_scale, group=self.world_group)
|
||||
|
||||
recvbuf_sign_server = torch.stack(recvbuf_sign_server)
|
||||
|
||||
flattened_recvbuf_sign_server = recvbuf_sign_server.type(torch.uint8).flatten()
|
||||
|
||||
buffer_m.data.copy_(
|
||||
torch_npu.npu_sign_bits_unpack(flattened_recvbuf_sign_server, self.size,
|
||||
torch.float32).mul_(recvbuf_scale_server_tmp).flatten().data)
|
||||
|
||||
if original_size != worker_error_size:
|
||||
buffer_m = buffer_m[0:original_size]
|
||||
if len(original_shape) > 1:
|
||||
buffer_m = buffer_m.reshape(original_shape)
|
||||
|
||||
return buffer_m
|
||||
@@ -0,0 +1,219 @@
|
||||
# Copyright (c) Microsoft Corporation.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# DeepSpeed Team
|
||||
|
||||
import torch
|
||||
import cupy
|
||||
import time
|
||||
import numpy as np
|
||||
from mpi4py import MPI
|
||||
|
||||
from deepspeed.runtime.comm.utils import check_and_handle_empty_buffer
|
||||
from deepspeed.runtime.compression.cupy import CupyBackend
|
||||
|
||||
|
||||
class MpiBackend(object):
|
||||
|
||||
def __init__(self, cuda_aware):
|
||||
self.comm = MPI.COMM_WORLD
|
||||
self.rank = self.comm.Get_rank()
|
||||
self.size = self.comm.Get_size()
|
||||
self.cuda_aware = cuda_aware
|
||||
self.compression_backend = CupyBackend()
|
||||
|
||||
def my_igather(self, rank, size, comm, sendbuf, recbuf, root):
|
||||
req = []
|
||||
if rank == root:
|
||||
for idx in range(size):
|
||||
if idx != rank:
|
||||
req.append(comm.Irecv(recbuf[idx], source=idx))
|
||||
else:
|
||||
recbuf[rank] = sendbuf
|
||||
else:
|
||||
req.append(comm.Isend(sendbuf, dest=root))
|
||||
return req
|
||||
|
||||
def gather_cuda(self, rank, world_size, comm, cupy_sign_list_packed, cupy_recvbuf_sign, cupy_worker_scale,
|
||||
cupy_recvbuf_scale):
|
||||
# We do in-place operations on cupy buffers so we do not return any buffers
|
||||
requests = []
|
||||
for idx in range(world_size):
|
||||
req_sign = self.my_igather(rank, world_size, comm, cupy_sign_list_packed[idx], cupy_recvbuf_sign, root=idx)
|
||||
requests += req_sign
|
||||
|
||||
for idx in range(world_size):
|
||||
req_scale = self.my_igather(rank, world_size, comm, cupy_worker_scale, cupy_recvbuf_scale, root=idx)
|
||||
requests += req_scale
|
||||
|
||||
MPI.Request.Waitall(requests)
|
||||
|
||||
def gather_host(self, rank, world_size, comm, cupy_sign_list_packed, cupy_recvbuf_sign, cupy_worker_scale,
|
||||
cupy_recvbuf_scale):
|
||||
|
||||
# In-place operations are not possible for newly created cupy arrays
|
||||
# so we need to return the new buffers
|
||||
numpy_recvbuf_sign = np.zeros([world_size, cupy_sign_list_packed[rank].size],
|
||||
dtype=cupy_sign_list_packed[0].dtype)
|
||||
numpy_recvbuf_scale = np.zeros([world_size, 1], dtype=cupy_worker_scale.dtype)
|
||||
|
||||
# 1. convert from cupy to numpy
|
||||
numpy_sign_list_packed = cupy_sign_list_packed
|
||||
|
||||
for idx in range(world_size):
|
||||
numpy_sign_list_packed[idx] = cupy.asnumpy(cupy_sign_list_packed[idx])
|
||||
|
||||
numpy_worker_scale = cupy.asnumpy(cupy_worker_scale)
|
||||
numpy_recvbuf_scale = cupy.asnumpy(cupy_recvbuf_scale)
|
||||
|
||||
cupy.cuda.get_current_stream().synchronize()
|
||||
|
||||
# 2. use numpy buffers for communication
|
||||
requests = []
|
||||
|
||||
for idx in range(world_size):
|
||||
req_sign = self.my_igather(rank,
|
||||
world_size,
|
||||
comm,
|
||||
numpy_sign_list_packed[idx],
|
||||
numpy_recvbuf_sign,
|
||||
root=idx)
|
||||
requests += req_sign
|
||||
|
||||
for idx in range(world_size):
|
||||
req_scale = self.my_igather(rank, world_size, comm, numpy_worker_scale, numpy_recvbuf_scale, root=idx)
|
||||
requests += req_scale
|
||||
|
||||
MPI.Request.Waitall(requests)
|
||||
|
||||
# 3. Convert back from numpy to cupy
|
||||
cupy_recvbuf_sign = cupy.asarray(numpy_recvbuf_sign)
|
||||
for idx in range(world_size):
|
||||
cupy_sign_list_packed[idx] = cupy.asarray(numpy_sign_list_packed[idx])
|
||||
|
||||
cupy_worker_scale = cupy.asarray(numpy_worker_scale)
|
||||
cupy_recvbuf_scale = cupy.asarray(numpy_recvbuf_scale)
|
||||
cupy.cuda.get_current_stream().synchronize()
|
||||
|
||||
return cupy_sign_list_packed, cupy_recvbuf_sign, cupy_worker_scale, cupy_recvbuf_scale
|
||||
|
||||
def allgather_cuda(self, comm, cupy_server_sign_packed, cupy_recvbuf_sign_server, cupy_server_scale,
|
||||
cupy_recvbuf_scale_server):
|
||||
comm.Allgather(cupy_server_sign_packed, cupy_recvbuf_sign_server)
|
||||
comm.Allgather(cupy_server_scale, cupy_recvbuf_scale_server)
|
||||
|
||||
def allgather_host(self, comm, cupy_server_sign_packed, cupy_recvbuf_sign_server, cupy_server_scale,
|
||||
cupy_recvbuf_scale_server):
|
||||
|
||||
# 1. Convert cupy to numpy
|
||||
numpy_recvbuf_sign_server = np.zeros([comm.Get_size(), cupy_server_sign_packed.size],
|
||||
dtype=cupy_server_sign_packed.dtype)
|
||||
numpy_recvbuf_scale_server = np.zeros([comm.Get_size(), 1], dtype=cupy_server_scale.dtype)
|
||||
|
||||
numpy_server_sign_packed = cupy.asnumpy(cupy_server_sign_packed)
|
||||
numpy_recvbuf_sign_server = cupy.asnumpy(cupy_recvbuf_sign_server)
|
||||
numpy_server_scale = cupy.asnumpy(cupy_server_scale)
|
||||
numpy_recvbuf_scale_server = cupy.asnumpy(cupy_recvbuf_scale_server)
|
||||
cupy.cuda.get_current_stream().synchronize()
|
||||
|
||||
# 2. Communicate numpy buffers
|
||||
comm.Allgather(numpy_server_sign_packed, numpy_recvbuf_sign_server)
|
||||
comm.Allgather(numpy_server_scale, numpy_recvbuf_scale_server)
|
||||
comm.Barrier()
|
||||
|
||||
# 3. Convert numpy back to cupy
|
||||
cupy_server_sign_packed = cupy.asarray(numpy_server_sign_packed)
|
||||
cupy_recvbuf_sign_server = cupy.asarray(numpy_recvbuf_sign_server)
|
||||
cupy_server_scale = cupy.asarray(numpy_server_scale)
|
||||
cupy_recvbuf_scale_server = cupy.asarray(numpy_recvbuf_scale_server)
|
||||
cupy.cuda.get_current_stream().synchronize()
|
||||
|
||||
return cupy_server_sign_packed, cupy_recvbuf_sign_server, cupy_server_scale, cupy_recvbuf_scale_server
|
||||
|
||||
def compressed_allreduce(self, buffer_m: torch.tensor, worker_error, server_error, local_rank):
|
||||
|
||||
all_start_time = time.time()
|
||||
original_shape = buffer_m.size()
|
||||
if len(original_shape) > 1:
|
||||
buffer_m = torch.flatten(buffer_m)
|
||||
original_size = buffer_m.numel()
|
||||
worker_error_size = worker_error.numel()
|
||||
result = check_and_handle_empty_buffer(buffer_m, original_shape, original_size, worker_error, server_error)
|
||||
if result is not None:
|
||||
return result
|
||||
cupy.cuda.Device(local_rank).use()
|
||||
|
||||
if original_size != worker_error_size:
|
||||
empty_tensor = torch.zeros(worker_error_size - original_size, device=buffer_m.device)
|
||||
buffer_m = torch.cat([buffer_m, empty_tensor])
|
||||
|
||||
buffer_m.add_(worker_error)
|
||||
worker_scale = torch.linalg.norm(buffer_m) / np.sqrt(torch.numel(buffer_m))
|
||||
worker_error.set_(buffer_m - worker_scale * buffer_m.sign().add_(1).bool().float().add_(-0.5).mul_(2.0))
|
||||
|
||||
cupy_sign_list_packed = self.compression_backend.compress_by_chunk(
|
||||
self.compression_backend.torch2cupy(buffer_m.sign_().add_(1).bool()), self.size)
|
||||
cupy_worker_scale = self.compression_backend.torch2cupy(worker_scale)
|
||||
|
||||
cupy_recvbuf_sign = cupy.zeros([self.size, cupy_sign_list_packed[self.rank].size],
|
||||
dtype=cupy_sign_list_packed[0].dtype)
|
||||
cupy_recvbuf_scale = cupy.zeros([self.size, 1], dtype=cupy_worker_scale.dtype)
|
||||
|
||||
# Communication Phase 1
|
||||
gather_start = time.time()
|
||||
if self.cuda_aware:
|
||||
self.gather_cuda(self.rank, self.size, self.comm, cupy_sign_list_packed, cupy_recvbuf_sign,
|
||||
cupy_worker_scale, cupy_recvbuf_scale)
|
||||
else:
|
||||
_, cupy_recvbuf_sign, _, cupy_recvbuf_scale = self.gather_host(self.rank, self.size, self.comm,
|
||||
cupy_sign_list_packed, cupy_recvbuf_sign,
|
||||
cupy_worker_scale, cupy_recvbuf_scale)
|
||||
gather_end = time.time()
|
||||
|
||||
# cupy_sign_list_packed, cupy_worker_scale, worker_scale = None, None, None
|
||||
cupy_sign_list_packed = None
|
||||
|
||||
compensated_server_m = self.compression_backend.cupy2torch(
|
||||
(cupy.unpackbits(cupy_recvbuf_sign.flatten())).reshape(self.size, -1)).float().add_(-0.5).mul_(2.0).mul_(
|
||||
self.compression_backend.cupy2torch(cupy_recvbuf_scale).mul_(1 / self.size)).sum(0)
|
||||
compensated_server_m.add_(server_error)
|
||||
server_scale = torch.linalg.norm(compensated_server_m) / np.sqrt(compensated_server_m.numel())
|
||||
server_error.set_(compensated_server_m -
|
||||
server_scale * compensated_server_m.sign().add_(1).bool().float().add_(-0.5).mul_(2.0))
|
||||
|
||||
cupy_server_scale = self.compression_backend.torch2cupy(server_scale)
|
||||
|
||||
cupy_server_sign_packed = self.compression_backend.compress_by_chunk(
|
||||
self.compression_backend.torch2cupy(compensated_server_m.sign_().add_(1).bool()), 1)
|
||||
compensated_server_m = None
|
||||
|
||||
cupy_recvbuf_sign_server = cupy.zeros([self.size, cupy_server_sign_packed[0].size],
|
||||
dtype=cupy_recvbuf_sign.dtype)
|
||||
cupy_recvbuf_scale_server = cupy.zeros([self.size, 1], dtype=cupy_recvbuf_scale.dtype)
|
||||
# cupy_recvbuf_sign, cupy_recvbuf_scale = None, None
|
||||
cupy_recvbuf_sign = None
|
||||
|
||||
# Communication Phase 2
|
||||
if self.cuda_aware:
|
||||
self.allgather_cuda(self.comm, cupy_server_sign_packed[0], cupy_recvbuf_sign_server, cupy_server_scale,
|
||||
cupy_recvbuf_scale_server)
|
||||
else:
|
||||
_, cupy_recvbuf_sign_server, _, cupy_recvbuf_scale_server = self.allgather_host(
|
||||
self.comm, cupy_server_sign_packed[0], cupy_recvbuf_sign_server, cupy_server_scale,
|
||||
cupy_recvbuf_scale_server)
|
||||
|
||||
# cupy_server_sign_packed, cupy_server_scale, server_scale = None, None, None
|
||||
cupy_server_sign_packed = None
|
||||
|
||||
buffer_m.data.copy_(
|
||||
self.compression_backend.cupy2torch((cupy.unpackbits(cupy_recvbuf_sign_server.flatten())).reshape(
|
||||
self.size, -1)).float().add_(-0.5).mul_(2.0).mul_(
|
||||
self.compression_backend.cupy2torch(cupy_recvbuf_scale_server)).flatten().data)
|
||||
if original_size != worker_error_size:
|
||||
buffer_m = buffer_m[0:original_size]
|
||||
if len(original_shape) > 1:
|
||||
buffer_m = buffer_m.reshape(original_shape)
|
||||
|
||||
# cupy_recvbuf_sign_server, cupy_recvbuf_scale_server = None, None
|
||||
|
||||
return buffer_m
|
||||
@@ -0,0 +1,170 @@
|
||||
# Copyright (c) Microsoft Corporation.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# DeepSpeed Team
|
||||
|
||||
import torch
|
||||
import cupy
|
||||
import numpy as np
|
||||
|
||||
import deepspeed.comm as dist
|
||||
from deepspeed.accelerator import get_accelerator
|
||||
from deepspeed.runtime.comm.utils import check_and_handle_empty_buffer
|
||||
from deepspeed.runtime.compression.cupy import CupyBackend
|
||||
from deepspeed.utils.torch import required_torch_version
|
||||
|
||||
|
||||
class NcclBackend(object):
|
||||
|
||||
def __init__(self, mpu=None):
|
||||
if mpu is None:
|
||||
self.world_group = dist.new_group(ranks=range(dist.get_world_size()))
|
||||
else:
|
||||
self.mpu = mpu
|
||||
self.world_group = self.mpu.get_data_parallel_group()
|
||||
self.rank = dist.get_rank(group=self.world_group)
|
||||
self.size = dist.get_world_size(group=self.world_group)
|
||||
self.compression_backend = CupyBackend()
|
||||
self.bool_not_supported = required_torch_version(min_version=1.10)
|
||||
|
||||
def my_igather(self, rank, size, group, sendbuf, recvbuf, root):
|
||||
req = []
|
||||
if rank == root:
|
||||
for idx in range(size):
|
||||
if idx != rank:
|
||||
req.append(dist.irecv(recvbuf[idx], src=idx, group=group))
|
||||
else:
|
||||
recvbuf[rank] = sendbuf
|
||||
else:
|
||||
req.append(dist.isend(sendbuf, group=group, dst=root))
|
||||
return req
|
||||
|
||||
def my_gather(self, rank, size, group, sendbuf, recvbuf, root):
|
||||
if rank == root:
|
||||
for idx in range(size):
|
||||
if idx != rank:
|
||||
dist.recv(recvbuf[idx], src=idx, group=group)
|
||||
else:
|
||||
recvbuf[rank] = sendbuf
|
||||
else:
|
||||
dist.send(sendbuf, group=group, dst=root)
|
||||
|
||||
def compressed_allreduce(self, buffer_m: torch.tensor, worker_error, server_error, local_rank):
|
||||
|
||||
# all_start_time = time.time()
|
||||
original_shape = buffer_m.size()
|
||||
if len(original_shape) > 1:
|
||||
buffer_m = torch.flatten(buffer_m)
|
||||
original_size = buffer_m.numel()
|
||||
worker_error_size = worker_error.numel()
|
||||
result = check_and_handle_empty_buffer(buffer_m, original_shape, original_size, worker_error, server_error)
|
||||
if result is not None:
|
||||
return result
|
||||
cupy.cuda.Device(local_rank).use()
|
||||
|
||||
if original_size != worker_error_size:
|
||||
empty_tensor = torch.zeros(worker_error_size - original_size, device=buffer_m.device)
|
||||
buffer_m = torch.cat([buffer_m, empty_tensor])
|
||||
|
||||
buffer_m.add_(worker_error)
|
||||
worker_scale = torch.linalg.norm(buffer_m) / np.sqrt(buffer_m.numel())
|
||||
worker_error.set_(buffer_m - worker_scale * buffer_m.sign().add_(1).bool().float().add_(-0.5).mul_(2.0))
|
||||
|
||||
if self.bool_not_supported:
|
||||
cupy_sign_list_packed = self.compression_backend.compress_by_chunk(
|
||||
self.compression_backend.torch2cupy(buffer_m.sign_().add_(1).bool().to(dtype=torch.uint8)), self.size)
|
||||
else:
|
||||
cupy_sign_list_packed = self.compression_backend.compress_by_chunk(
|
||||
self.compression_backend.torch2cupy(buffer_m.sign_().add_(1).bool()), self.size)
|
||||
cupy_worker_scale = self.compression_backend.torch2cupy(worker_scale)
|
||||
|
||||
cupy_recvbuf_sign = cupy.zeros([self.size, cupy_sign_list_packed[self.rank].size],
|
||||
dtype=cupy_sign_list_packed[0].dtype)
|
||||
# cupy_recvbuf_scale = cupy.zeros([self.size, 1], dtype=cupy_worker_scale.dtype)
|
||||
|
||||
sign_list_packed = [
|
||||
self.compression_backend.cupy2torch(cupy_sign_list_packed[idx]) for idx in range(self.size)
|
||||
]
|
||||
|
||||
# worker_scale = self.compression_backend.cupy2torch(cupy_worker_scale)
|
||||
recvbuf_sign = self.compression_backend.cupy2torch(cupy_recvbuf_sign)
|
||||
#recvbuf_scale = self.compression_backend.cupy2torch(cupy_recvbuf_scale)
|
||||
recvbuf_scale = [
|
||||
torch.zeros(1, dtype=worker_scale.dtype, device=torch.device(get_accelerator().device_name(local_rank)))
|
||||
for i in range(self.size)
|
||||
]
|
||||
|
||||
# communication phase 1
|
||||
# gather_start = time.time()
|
||||
# Alltoall for sign
|
||||
dist.all_to_all_single(recvbuf_sign, torch.stack(sign_list_packed), group=self.world_group)
|
||||
# Allgather for scale
|
||||
dist.all_gather(recvbuf_scale, worker_scale, group=self.world_group)
|
||||
|
||||
# gather_end = time.time()
|
||||
|
||||
# cupy_sign_list_packed, sign_list_packed, cupy_worker_scale, worker_scale = None, None, None, None
|
||||
cupy_sign_list_packed = None
|
||||
|
||||
cupy_recvbuf_sign = self.compression_backend.torch2cupy(recvbuf_sign)
|
||||
#cupy_recvbuf_scale = self.compression_backend.torch2cupy(torch.stack(recvbuf_scale))
|
||||
|
||||
compensated_server_m = self.compression_backend.cupy2torch(
|
||||
(cupy.unpackbits(cupy_recvbuf_sign.flatten())).reshape(self.size, -1)).float().add_(-0.5).mul_(2.0).mul_(
|
||||
torch.stack(recvbuf_scale).mul_(1 / self.size)).sum(0)
|
||||
compensated_server_m.add_(server_error)
|
||||
server_scale = torch.linalg.norm(compensated_server_m) / np.sqrt(compensated_server_m.numel())
|
||||
server_error.set_(compensated_server_m -
|
||||
server_scale * compensated_server_m.sign().add_(1).bool().float().add_(-0.5).mul_(2.0))
|
||||
|
||||
# cupy_server_scale = self.compression_backend.torch2cupy(server_scale)
|
||||
|
||||
if self.bool_not_supported:
|
||||
cupy_server_sign_packed = self.compression_backend.compress_by_chunk(
|
||||
self.compression_backend.torch2cupy(compensated_server_m.sign_().add_(1).bool().to(dtype=torch.uint8)),
|
||||
1)
|
||||
else:
|
||||
cupy_server_sign_packed = self.compression_backend.compress_by_chunk(
|
||||
self.compression_backend.torch2cupy(compensated_server_m.sign_().add_(1).bool()), 1)
|
||||
compensated_server_m = None
|
||||
|
||||
cupy_recvbuf_sign_server = cupy.zeros([self.size, cupy_server_sign_packed[0].size],
|
||||
dtype=cupy_recvbuf_sign.dtype)
|
||||
# cupy_recvbuf_sign, recvbuf_sign = None, None
|
||||
cupy_recvbuf_sign = None
|
||||
|
||||
server_sign_packed = [self.compression_backend.cupy2torch(cupy_server_sign_packed[0])]
|
||||
recvbuf_sign_server = [
|
||||
self.compression_backend.cupy2torch(cupy_recvbuf_sign_server[idx]) for idx in range(self.size)
|
||||
]
|
||||
|
||||
# server_scale = self.compression_backend.cupy2torch(cupy_server_scale)
|
||||
cupy_recvbuf_scale_server = cupy.zeros([self.size, 1], dtype=cupy_worker_scale.dtype)
|
||||
# cupy_recvbuf_scale, recvbuf_scale = None, None
|
||||
|
||||
recvbuf_scale_server = [
|
||||
self.compression_backend.cupy2torch(cupy_recvbuf_scale_server[idx]) for idx in range(self.size)
|
||||
]
|
||||
|
||||
# Communication Phase 2
|
||||
dist.all_gather(recvbuf_sign_server, server_sign_packed[0], group=self.world_group)
|
||||
dist.all_gather(recvbuf_scale_server, server_scale, group=self.world_group)
|
||||
|
||||
cupy_server_sign_packed = None
|
||||
|
||||
# need to convert from a tensor list to a single tensor
|
||||
# dist.all_gather only provides a tensor list as the recv/output buffer
|
||||
recvbuf_sign_server = torch.stack(recvbuf_sign_server)
|
||||
|
||||
cupy_recvbuf_sign_server = self.compression_backend.torch2cupy(recvbuf_sign_server)
|
||||
|
||||
buffer_m.data.copy_(
|
||||
self.compression_backend.cupy2torch((cupy.unpackbits(cupy_recvbuf_sign_server.flatten())).reshape(
|
||||
self.size, -1)).float().add_(-0.5).mul_(2.0).mul_(
|
||||
self.compression_backend.cupy2torch(cupy_recvbuf_scale_server)).flatten().data)
|
||||
if original_size != worker_error_size:
|
||||
buffer_m = buffer_m[0:original_size]
|
||||
if len(original_shape) > 1:
|
||||
buffer_m = buffer_m.reshape(original_shape)
|
||||
|
||||
return buffer_m
|
||||
@@ -0,0 +1,26 @@
|
||||
# Copyright (c) Microsoft Corporation.
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# DeepSpeed Team
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def check_and_handle_empty_buffer(
|
||||
buffer_m: torch.Tensor,
|
||||
original_shape: torch.Size,
|
||||
original_size: int,
|
||||
worker_error: torch.Tensor,
|
||||
server_error: torch.Tensor,
|
||||
) -> Optional[torch.Tensor]:
|
||||
if original_size == 0:
|
||||
if worker_error.numel():
|
||||
worker_error.zero_()
|
||||
if server_error.numel():
|
||||
server_error.zero_()
|
||||
if len(original_shape) > 1:
|
||||
return buffer_m.reshape(original_shape)
|
||||
return buffer_m
|
||||
return None
|
||||
Reference in New Issue
Block a user