under test, not sure no errors
This commit is contained in:
@@ -104,15 +104,11 @@ def _dp_all_gather(
|
|||||||
"""All-gather ``tensor`` along ``dim`` across the DP process group."""
|
"""All-gather ``tensor`` along ``dim`` across the DP process group."""
|
||||||
if world_size <= 1:
|
if world_size <= 1:
|
||||||
return tensor
|
return tensor
|
||||||
try:
|
from vllm.distributed import get_dp_group
|
||||||
from vllm.distributed import get_dp_group
|
group = get_dp_group()
|
||||||
group = get_dp_group()
|
gathered = [torch.empty_like(tensor) for _ in range(world_size)]
|
||||||
gathered = [torch.empty_like(tensor) for _ in range(world_size)]
|
torch.distributed.all_gather(gathered, tensor, group=group)
|
||||||
torch.distributed.all_gather(gathered, tensor, group=group)
|
return torch.cat(gathered, dim=dim)
|
||||||
return torch.cat(gathered, dim=dim)
|
|
||||||
except (ImportError, RuntimeError):
|
|
||||||
# Fallback: repeat for testing without actual distributed backend
|
|
||||||
return tensor.repeat(world_size, *([1] * (tensor.dim() - 1)))
|
|
||||||
|
|
||||||
|
|
||||||
def _dp_all_gather_variable(
|
def _dp_all_gather_variable(
|
||||||
@@ -123,20 +119,17 @@ def _dp_all_gather_variable(
|
|||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""Variable-length all-gather: each rank contributes a different number
|
"""Variable-length all-gather: each rank contributes a different number
|
||||||
of tokens. Returns a compact concatenation without padding."""
|
of tokens. Returns a compact concatenation without padding."""
|
||||||
try:
|
from vllm.distributed import get_dp_group
|
||||||
from vllm.distributed import get_dp_group
|
group = get_dp_group()
|
||||||
group = get_dp_group()
|
world_size = len(token_counts)
|
||||||
world_size = len(token_counts)
|
hidden_dim = tensor.shape[1] if tensor.dim() > 1 else 1
|
||||||
hidden_dim = tensor.shape[1] if tensor.dim() > 1 else 1
|
recv_tensors = []
|
||||||
recv_tensors = []
|
for i, count in enumerate(token_counts):
|
||||||
for i, count in enumerate(token_counts):
|
if i == dp_rank:
|
||||||
if i == dp_rank:
|
recv_tensors.append(tensor[:count])
|
||||||
recv_tensors.append(tensor[:count])
|
else:
|
||||||
else:
|
recv_tensors.append(
|
||||||
recv_tensors.append(
|
torch.empty(count, hidden_dim, dtype=tensor.dtype, device=tensor.device)
|
||||||
torch.empty(count, hidden_dim, dtype=tensor.dtype, device=tensor.device)
|
)
|
||||||
)
|
torch.distributed.all_gather(recv_tensors, tensor[:token_counts[dp_rank]], group=group)
|
||||||
torch.distributed.all_gather(recv_tensors, tensor[:token_counts[dp_rank]], group=group)
|
return torch.cat(recv_tensors, dim=0)
|
||||||
return torch.cat(recv_tensors, dim=0)
|
|
||||||
except (ImportError, RuntimeError):
|
|
||||||
return tensor
|
|
||||||
|
|||||||
@@ -120,13 +120,8 @@ class ModelExecutor:
|
|||||||
).lower()
|
).lower()
|
||||||
graph_disabled = graph_backend in ("", "off", "none", "0")
|
graph_disabled = graph_backend in ("", "off", "none", "0")
|
||||||
if graph_disabled and config.get("enable_graph", False):
|
if graph_disabled and config.get("enable_graph", False):
|
||||||
# Default to ACL graph on NPU platforms
|
import torch_npu # noqa: F401
|
||||||
try:
|
return "aclgraph"
|
||||||
import torch_npu # noqa: F401
|
|
||||||
|
|
||||||
return "aclgraph"
|
|
||||||
except ImportError:
|
|
||||||
pass
|
|
||||||
return graph_backend
|
return graph_backend
|
||||||
|
|
||||||
@torch.inference_mode()
|
@torch.inference_mode()
|
||||||
|
|||||||
@@ -141,15 +141,12 @@ def _dp_all_gather(
|
|||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
if world_size <= 1:
|
if world_size <= 1:
|
||||||
return tensor
|
return tensor
|
||||||
try:
|
from vllm.distributed import get_dp_group
|
||||||
from vllm.distributed import get_dp_group
|
|
||||||
|
|
||||||
group = get_dp_group()
|
group = get_dp_group()
|
||||||
gathered = [torch.empty_like(tensor) for _ in range(world_size)]
|
gathered = [torch.empty_like(tensor) for _ in range(world_size)]
|
||||||
torch.distributed.all_gather(gathered, tensor, group=group)
|
torch.distributed.all_gather(gathered, tensor, group=group)
|
||||||
return torch.cat(gathered, dim=dim)
|
return torch.cat(gathered, dim=dim)
|
||||||
except (ImportError, RuntimeError):
|
|
||||||
return tensor.repeat(world_size, *([1] * (tensor.dim() - 1)))
|
|
||||||
|
|
||||||
|
|
||||||
def _dp_all_gather_variable(
|
def _dp_all_gather_variable(
|
||||||
@@ -158,25 +155,22 @@ def _dp_all_gather_variable(
|
|||||||
dp_rank: int,
|
dp_rank: int,
|
||||||
group_name: str = "dp",
|
group_name: str = "dp",
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
try:
|
from vllm.distributed import get_dp_group
|
||||||
from vllm.distributed import get_dp_group
|
|
||||||
|
|
||||||
group = get_dp_group()
|
group = get_dp_group()
|
||||||
world_size = len(token_counts)
|
world_size = len(token_counts)
|
||||||
hidden_dim = tensor.shape[1] if tensor.dim() > 1 else 1
|
hidden_dim = tensor.shape[1] if tensor.dim() > 1 else 1
|
||||||
recv_tensors = []
|
recv_tensors = []
|
||||||
for i, count in enumerate(token_counts):
|
for i, count in enumerate(token_counts):
|
||||||
if i == dp_rank:
|
if i == dp_rank:
|
||||||
recv_tensors.append(tensor[:count])
|
recv_tensors.append(tensor[:count])
|
||||||
else:
|
else:
|
||||||
recv_tensors.append(
|
recv_tensors.append(
|
||||||
torch.empty(
|
torch.empty(
|
||||||
count, hidden_dim, dtype=tensor.dtype, device=tensor.device
|
count, hidden_dim, dtype=tensor.dtype, device=tensor.device
|
||||||
)
|
|
||||||
)
|
)
|
||||||
torch.distributed.all_gather(
|
)
|
||||||
recv_tensors, tensor[: token_counts[dp_rank]], group=group
|
torch.distributed.all_gather(
|
||||||
)
|
recv_tensors, tensor[: token_counts[dp_rank]], group=group
|
||||||
return torch.cat(recv_tensors, dim=0)
|
)
|
||||||
except (ImportError, RuntimeError):
|
return torch.cat(recv_tensors, dim=0)
|
||||||
return tensor
|
|
||||||
|
|||||||
Reference in New Issue
Block a user