32
vllm_ascend/distributed/weight_transfer/__init__.py
Normal file
32
vllm_ascend/distributed/weight_transfer/__init__.py
Normal file
@@ -0,0 +1,32 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
# This file is a part of the vllm-ascend project.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
#
|
||||
|
||||
from vllm.distributed.weight_transfer.factory import WeightTransferEngineFactory
|
||||
|
||||
|
||||
def register_engine():
|
||||
"""Register Ascend weight transfer engines as vLLM plugins."""
|
||||
WeightTransferEngineFactory.register_engine(
|
||||
"hccl",
|
||||
"vllm_ascend.distributed.weight_transfer.hccl_engine",
|
||||
"HCCLWeightTransferEngine",
|
||||
)
|
||||
WeightTransferEngineFactory.register_engine(
|
||||
"npu_ipc",
|
||||
"vllm_ascend.distributed.weight_transfer.npu_ipc_engine",
|
||||
"NPUIPCWeightTransferEngine",
|
||||
)
|
||||
336
vllm_ascend/distributed/weight_transfer/hccl_engine.py
Normal file
336
vllm_ascend/distributed/weight_transfer/hccl_engine.py
Normal file
@@ -0,0 +1,336 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""HCCL-based weight transfer engine."""
|
||||
|
||||
from collections.abc import Callable, Iterator
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import torch
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm_ascend.distributed.device_communicators.pyhccl import PyHcclCommunicator
|
||||
|
||||
from vllm.config.parallel import ParallelConfig
|
||||
from vllm.config.weight_transfer import WeightTransferConfig
|
||||
from vllm.distributed.weight_transfer.base import (
|
||||
WeightTransferEngine,
|
||||
WeightTransferInitInfo,
|
||||
WeightTransferUpdateInfo,
|
||||
)
|
||||
|
||||
from vllm_ascend.distributed.weight_transfer.packed_tensor import (
|
||||
DEFAULT_PACKED_BUFFER_SIZE_BYTES,
|
||||
DEFAULT_PACKED_NUM_BUFFERS,
|
||||
packed_broadcast_consumer,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class HCCLWeightTransferInitInfo(WeightTransferInitInfo):
|
||||
"""Initialization info for HCCL weight transfer backend."""
|
||||
|
||||
master_address: str
|
||||
"""IP address of the trainer (rank 0) for HCCL process group setup."""
|
||||
master_port: int
|
||||
"""Port on the trainer for HCCL process group setup."""
|
||||
rank_offset: int
|
||||
"""Offset added to each vLLM worker's rank within the HCCL group.
|
||||
Typically 1 (trainer is rank 0, workers start at rank 1)."""
|
||||
world_size: int
|
||||
"""Total number of participants in the HCCL group (trainer + all workers)."""
|
||||
|
||||
|
||||
@dataclass
|
||||
class HCCLTrainerSendWeightsArgs:
|
||||
"""Arguments for HCCL trainer_send_weights method."""
|
||||
|
||||
group: Any
|
||||
"""Process group (PyHcclCommunicator) for HCCL communication."""
|
||||
src: int = 0
|
||||
"""Source rank (default 0, trainer is typically rank 0)."""
|
||||
post_iter_func: Callable[[tuple[str, torch.Tensor]], torch.Tensor] | None = None
|
||||
"""Optional function to apply to each (name, tensor) pair before broadcasting.
|
||||
If None, extracts just the tensor."""
|
||||
packed: bool = False
|
||||
"""Whether to use packed tensor broadcasting for efficiency.
|
||||
When True, multiple tensors are batched together before broadcasting
|
||||
to reduce HCCL communication overhead."""
|
||||
stream: torch.npu.Stream | None = None
|
||||
"""ACL stream to use for broadcasting if packed is False.
|
||||
If packed is True, new streams will be created for each buffer."""
|
||||
packed_buffer_size_bytes: int = DEFAULT_PACKED_BUFFER_SIZE_BYTES
|
||||
"""Size in bytes for each packed tensor buffer.
|
||||
Must match the value used in HCCLWeightTransferUpdateInfo."""
|
||||
packed_num_buffers: int = DEFAULT_PACKED_NUM_BUFFERS
|
||||
"""Number of buffers for double/triple buffering during packed transfer.
|
||||
Must match the value used in HCCLWeightTransferUpdateInfo."""
|
||||
|
||||
|
||||
@dataclass
|
||||
class HCCLWeightTransferUpdateInfo(WeightTransferUpdateInfo):
|
||||
"""Update info for HCCL weight transfer backend."""
|
||||
|
||||
names: list[str]
|
||||
"""Names of the parameters to transfer (e.g. ``model.layers.0.weight``)."""
|
||||
dtype_names: list[str]
|
||||
"""Torch dtype names (e.g. ``bfloat16``, ``float32``) for each parameter."""
|
||||
shapes: list[list[int]]
|
||||
"""Shapes of each parameter as integer lists."""
|
||||
packed: bool = False
|
||||
"""Whether to use packed tensor broadcasting for efficiency.
|
||||
When True, multiple tensors are batched together before broadcasting
|
||||
to reduce HCCL communication overhead."""
|
||||
packed_buffer_size_bytes: int = DEFAULT_PACKED_BUFFER_SIZE_BYTES
|
||||
"""Size in bytes for each packed tensor buffer.
|
||||
Both producer and consumer must use the same value."""
|
||||
packed_num_buffers: int = DEFAULT_PACKED_NUM_BUFFERS
|
||||
"""Number of buffers for double/triple buffering during packed transfer.
|
||||
Both producer and consumer must use the same value."""
|
||||
|
||||
def __post_init__(self):
|
||||
"""Validate that all lists have the same length."""
|
||||
num_params = len(self.names)
|
||||
if len(self.dtype_names) != num_params:
|
||||
raise ValueError(
|
||||
f"`dtype_names` should be of the same size as `names`: "
|
||||
f"got {len(self.dtype_names)} and {len(self.names)}"
|
||||
)
|
||||
if len(self.shapes) != num_params:
|
||||
raise ValueError(
|
||||
f"`shapes` should be of the same size as `names`: got {len(self.shapes)} and {len(self.names)}"
|
||||
)
|
||||
|
||||
|
||||
class HCCLWeightTransferEngine(WeightTransferEngine[HCCLWeightTransferInitInfo, HCCLWeightTransferUpdateInfo]):
|
||||
"""
|
||||
Weight transfer engine using HCCL for communication between trainer and workers.
|
||||
|
||||
This implementation uses HCCL broadcast operations to transfer weights from
|
||||
the trainer (rank 0) to all inference workers in a process group.
|
||||
"""
|
||||
|
||||
# Define backend-specific dataclass types
|
||||
init_info_cls = HCCLWeightTransferInitInfo
|
||||
update_info_cls = HCCLWeightTransferUpdateInfo
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: WeightTransferConfig,
|
||||
parallel_config: ParallelConfig,
|
||||
model: torch.nn.Module | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Initialize the HCCL weight transfer engine.
|
||||
|
||||
Args:
|
||||
config: The configuration for the weight transfer engine
|
||||
parallel_config: The configuration for the parallel setup
|
||||
model: The local model instance which will receive the weights.
|
||||
"""
|
||||
super().__init__(config, parallel_config, model)
|
||||
self.model_update_group: PyHcclCommunicator | None = None
|
||||
|
||||
def init_transfer_engine(self, init_info: HCCLWeightTransferInitInfo) -> None:
|
||||
"""
|
||||
Initialize HCCL process group with the trainer.
|
||||
|
||||
Args:
|
||||
init_info: HCCL initialization info containing master address, port,
|
||||
rank offset, and world size
|
||||
"""
|
||||
|
||||
# Calculate the global rank in the trainer-worker process group
|
||||
# Must account for data parallel to get unique ranks across all workers
|
||||
dp_rank = self.parallel_config.data_parallel_index
|
||||
world_size_per_dp = self.parallel_config.world_size # TP * PP
|
||||
rank_within_dp = self.parallel_config.rank
|
||||
|
||||
# Unique rank across all DP groups
|
||||
worker_rank = dp_rank * world_size_per_dp + rank_within_dp
|
||||
rank = worker_rank + init_info.rank_offset
|
||||
# Create stateless process group
|
||||
device = torch.accelerator.current_device_index()
|
||||
self.model_update_group = HCCLWeightTransferEngine._stateless_init_process_group(
|
||||
init_info.master_address,
|
||||
init_info.master_port,
|
||||
rank,
|
||||
init_info.world_size,
|
||||
device=device,
|
||||
)
|
||||
|
||||
def receive_weights(
|
||||
self,
|
||||
update_info: HCCLWeightTransferUpdateInfo,
|
||||
load_weights: Callable[[list[tuple[str, torch.Tensor]]], None],
|
||||
) -> None:
|
||||
"""
|
||||
Receive weights from trainer via HCCL broadcast and load them incrementally.
|
||||
|
||||
If update_info.packed is True, uses packed tensor broadcasting for
|
||||
efficient transfer of multiple weights in batches. Otherwise, uses simple
|
||||
one-by-one broadcasting.
|
||||
|
||||
Args:
|
||||
update_info: HCCL update info containing parameter names, dtypes, shapes,
|
||||
and packed flag
|
||||
load_weights: Callable that loads weights into the model. Called
|
||||
incrementally for each batch of weights to avoid OOM.
|
||||
"""
|
||||
if self.model_update_group is None:
|
||||
raise RuntimeError("HCCL weight transfer not initialized. Call init_transfer_engine() first.")
|
||||
|
||||
if update_info.packed:
|
||||
# Build iterator of (name, (shape, dtype)) from update_info
|
||||
def state_dict_info_iterator():
|
||||
for name, dtype_name, shape in zip(update_info.names, update_info.dtype_names, update_info.shapes):
|
||||
dtype = getattr(torch, dtype_name)
|
||||
yield (name, (shape, dtype))
|
||||
|
||||
packed_broadcast_consumer(
|
||||
iterator=state_dict_info_iterator(),
|
||||
group=self.model_update_group,
|
||||
src=0,
|
||||
post_unpack_func=load_weights,
|
||||
buffer_size_bytes=update_info.packed_buffer_size_bytes,
|
||||
num_buffers=update_info.packed_num_buffers,
|
||||
)
|
||||
else:
|
||||
# Use simple one-by-one broadcasting
|
||||
for name, dtype_name, shape in zip(update_info.names, update_info.dtype_names, update_info.shapes):
|
||||
dtype = getattr(torch, dtype_name)
|
||||
weight = torch.empty(shape, dtype=dtype, device="npu")
|
||||
self.model_update_group.broadcast(weight, src=0, stream=torch.npu.current_stream())
|
||||
load_weights([(name, weight)])
|
||||
del weight
|
||||
|
||||
def shutdown(self) -> None:
|
||||
if self.model_update_group is not None:
|
||||
# Clean up the communicator by removing the reference
|
||||
self.model_update_group = None
|
||||
|
||||
@staticmethod
|
||||
def trainer_send_weights(
|
||||
iterator: Iterator[tuple[str, torch.Tensor]],
|
||||
trainer_args: dict[str, Any] | HCCLTrainerSendWeightsArgs,
|
||||
) -> None:
|
||||
"""Broadcast weights from trainer to vLLM workers.
|
||||
|
||||
Args:
|
||||
iterator: Iterator of model parameters. Returns (name, tensor) tuples
|
||||
trainer_args: Dictionary or HCCLTrainerSendWeightsArgs instance containing
|
||||
HCCL-specific arguments. If a dict, should contain keys from
|
||||
HCCLTrainerSendWeightsArgs.
|
||||
|
||||
Example:
|
||||
>>> from vllm.distributed.weight_transfer.hccl_engine import (
|
||||
... HCCLWeightTransferEngine,
|
||||
... HCCLTrainerSendWeightsArgs,
|
||||
... )
|
||||
>>> param_iter = ((n, p) for n, p in model.named_parameters())
|
||||
>>> args = HCCLTrainerSendWeightsArgs(group=group, packed=True)
|
||||
>>> HCCLWeightTransferEngine.trainer_send_weights(param_iter, args)
|
||||
"""
|
||||
# Parse trainer args - accept either dict or dataclass instance
|
||||
if isinstance(trainer_args, dict):
|
||||
args = HCCLTrainerSendWeightsArgs(**trainer_args)
|
||||
else:
|
||||
args = trainer_args
|
||||
|
||||
if args.post_iter_func is None:
|
||||
# Default: extract just the tensor from (name, tensor) tuple
|
||||
post_iter_func = lambda x: x[1]
|
||||
else:
|
||||
post_iter_func = args.post_iter_func
|
||||
|
||||
if args.packed:
|
||||
# Use packed tensor broadcasting for efficiency
|
||||
from vllm_ascend.distributed.weight_transfer.packed_tensor import (
|
||||
packed_broadcast_producer,
|
||||
)
|
||||
|
||||
packed_broadcast_producer(
|
||||
iterator=iterator,
|
||||
group=args.group,
|
||||
src=args.src,
|
||||
post_iter_func=post_iter_func,
|
||||
buffer_size_bytes=args.packed_buffer_size_bytes,
|
||||
num_buffers=args.packed_num_buffers,
|
||||
)
|
||||
else:
|
||||
# Use simple one-by-one broadcasting
|
||||
for item in iterator:
|
||||
tensor = post_iter_func(item)
|
||||
args.group.broadcast(
|
||||
tensor,
|
||||
src=args.src,
|
||||
stream=args.stream or torch.npu.current_stream(),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def trainer_init(
|
||||
init_info: HCCLWeightTransferInitInfo | dict,
|
||||
) -> "PyHcclCommunicator":
|
||||
"""
|
||||
Initialize HCCL process group for trainer-side weight transfer.
|
||||
|
||||
The trainer is always rank 0 in the process group. Uses the current
|
||||
Ascend device (torch.accelerator.current_device_index()).
|
||||
|
||||
Args:
|
||||
init_info: Either an HCCLWeightTransferInitInfo object or a dict with keys:
|
||||
- master_address: str
|
||||
- master_port: int
|
||||
- world_size: int
|
||||
|
||||
Returns:
|
||||
PyHcclCommunicator for weight transfer.
|
||||
|
||||
Example:
|
||||
>>> from vllm.distributed.weight_transfer.hccl_engine import (
|
||||
... HCCLWeightTransferEngine,
|
||||
... )
|
||||
>>> group = HCCLWeightTransferEngine.trainer_init(
|
||||
... dict(
|
||||
... master_address=master_address,
|
||||
... master_port=master_port,
|
||||
... world_size=world_size,
|
||||
... ),
|
||||
... )
|
||||
"""
|
||||
if isinstance(init_info, dict):
|
||||
master_address = init_info["master_address"]
|
||||
master_port = init_info["master_port"]
|
||||
world_size = init_info["world_size"]
|
||||
else:
|
||||
# HCCLWeightTransferInitInfo object
|
||||
master_address = init_info.master_address
|
||||
master_port = init_info.master_port
|
||||
world_size = init_info.world_size
|
||||
|
||||
# Trainer is always rank 0
|
||||
device = torch.accelerator.current_device_index()
|
||||
return HCCLWeightTransferEngine._stateless_init_process_group(
|
||||
master_address,
|
||||
master_port,
|
||||
0,
|
||||
world_size,
|
||||
device,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _stateless_init_process_group(master_address, master_port, rank, world_size, device):
|
||||
"""
|
||||
vLLM provides `StatelessProcessGroup` to create a process group
|
||||
without considering the global process group in torch.distributed.
|
||||
It is recommended to create `StatelessProcessGroup`, and then initialize
|
||||
the data-plane communication (HCCL) between external (train processes)
|
||||
and vLLM workers.
|
||||
"""
|
||||
from vllm.distributed.utils import StatelessProcessGroup
|
||||
|
||||
from vllm_ascend.distributed.device_communicators.pyhccl import PyHcclCommunicator
|
||||
|
||||
pg = StatelessProcessGroup.create(host=master_address, port=master_port, rank=rank, world_size=world_size)
|
||||
pyhccl = PyHcclCommunicator(pg, device=device)
|
||||
return pyhccl
|
||||
405
vllm_ascend/distributed/weight_transfer/npu_ipc_engine.py
Normal file
405
vllm_ascend/distributed/weight_transfer/npu_ipc_engine.py
Normal file
@@ -0,0 +1,405 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""NPU IPC-based weight transfer engine using Ascend IPC for communication."""
|
||||
|
||||
import os
|
||||
import pickle
|
||||
import socket
|
||||
from collections.abc import Callable, Iterator
|
||||
from dataclasses import asdict, dataclass
|
||||
from functools import lru_cache
|
||||
from typing import Any
|
||||
|
||||
import pybase64 as base64
|
||||
import requests
|
||||
import torch
|
||||
from torch.multiprocessing.reductions import reduce_tensor
|
||||
from vllm import envs
|
||||
from vllm.config.parallel import ParallelConfig
|
||||
from vllm.config.weight_transfer import WeightTransferConfig
|
||||
from vllm.distributed.weight_transfer.base import (
|
||||
WeightTransferEngine,
|
||||
WeightTransferInitInfo,
|
||||
)
|
||||
from vllm.distributed.weight_transfer.ipc_engine import (
|
||||
IPCTrainerSendWeightsArgs,
|
||||
IPCWeightTransferUpdateInfo,
|
||||
)
|
||||
|
||||
from vllm_ascend.distributed.weight_transfer.packed_tensor import (
|
||||
packed_npu_ipc_consumer,
|
||||
packed_npu_ipc_producer,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class NPUIPCTrainerSendWeightsArgs(IPCTrainerSendWeightsArgs):
|
||||
"""NPU IPC variant — inherits all fields and validation from the CUDA IPC
|
||||
base class. Only the ``send_mode`` callable type is widened to accept the
|
||||
NPU update-info type."""
|
||||
|
||||
send_mode: str | Callable[["NPUIPCWeightTransferUpdateInfo"], None]
|
||||
|
||||
|
||||
@dataclass
|
||||
class NPUIPCWeightTransferInitInfo(WeightTransferInitInfo):
|
||||
"""Initialization info for NPU IPC weight transfer backend.
|
||||
|
||||
No initialization needed for NPU IPC.
|
||||
"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class NPUIPCWeightTransferUpdateInfo(IPCWeightTransferUpdateInfo):
|
||||
"""NPU IPC variant — inherits all fields and validation from the CUDA IPC
|
||||
base class. No overrides needed; the field types and ``__post_init__`` are
|
||||
identical."""
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def get_ip() -> str:
|
||||
try:
|
||||
# try to get ip from network interface
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_DGRAM) as s:
|
||||
s.connect(("8.8.8.8", 80))
|
||||
return s.getsockname()[0]
|
||||
except Exception: # noqa: BLE001
|
||||
# fallback to get ip from hostname
|
||||
return socket.gethostbyname(socket.gethostname())
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def npu_generate_uuid() -> str:
|
||||
"""Generate a unique identifier for the current process's physical NPU chip.
|
||||
|
||||
Returns ``{host_ip}-{physical_chip_id}`` where ``host_ip`` is the local
|
||||
machine's IP address and ``physical_chip_id`` is derived from the current
|
||||
logical device index mapped through ``ASCEND_RT_VISIBLE_DEVICES``.
|
||||
|
||||
On Ascend NPU, ``torch.accelerator.current_device_index()`` returns the
|
||||
*logical* device index. When ``ASCEND_RT_VISIBLE_DEVICES`` is set, it
|
||||
maps logical indices to physical chip IDs (e.g., ``ASCEND_RT_VISIBLE_DEVICES=2,3``
|
||||
means logical device 0 → physical chip 2, logical device 1 → physical chip 3).
|
||||
If the env var is not set, the logical index is used directly as the
|
||||
physical chip ID (identity mapping).
|
||||
|
||||
The result is cached because it is constant for the lifetime of the
|
||||
process. Both the trainer and inference worker processes co-located
|
||||
on the same physical NPU chip will produce the same UUID, which is
|
||||
required for NPU IPC handle matching.
|
||||
"""
|
||||
logical_device = torch.accelerator.current_device_index()
|
||||
visible_devices = os.environ.get("ASCEND_RT_VISIBLE_DEVICES", None)
|
||||
if visible_devices:
|
||||
physical_device = int(visible_devices.split(",")[logical_device].strip())
|
||||
else:
|
||||
physical_device = logical_device
|
||||
return f"{get_ip()}-{physical_device}"
|
||||
|
||||
|
||||
class NPUIPCWeightTransferEngine(WeightTransferEngine[NPUIPCWeightTransferInitInfo, NPUIPCWeightTransferUpdateInfo]):
|
||||
"""
|
||||
Weight transfer engine using NPU IPC for communication between
|
||||
trainer and workers.
|
||||
|
||||
This implementation uses Ascend NPU IPC to transfer weights from the
|
||||
trainer (rank 0) to all inference workers. IPC handles are used to
|
||||
share memory between processes on the same node.
|
||||
|
||||
Requires ``torch_npu`` to be imported (which patches
|
||||
``torch.multiprocessing.reductions.reduce_tensor`` to support
|
||||
NPU tensors via ``_share_npu_()`` / ``rebuild_npu_tensor``).
|
||||
"""
|
||||
|
||||
init_info_cls = NPUIPCWeightTransferInitInfo
|
||||
update_info_cls = NPUIPCWeightTransferUpdateInfo
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: WeightTransferConfig,
|
||||
parallel_config: ParallelConfig,
|
||||
model: torch.nn.Module | None = None,
|
||||
) -> None:
|
||||
super().__init__(config, parallel_config, model)
|
||||
|
||||
def parse_update_info(self, update_dict: dict[str, Any]) -> NPUIPCWeightTransferUpdateInfo:
|
||||
"""Parse update dict, deserializing pickled IPC handles if present.
|
||||
|
||||
HTTP transport sends IPC handles as a base64-encoded pickle under the
|
||||
key ``ipc_handles_pickled``. This method deserializes them back into
|
||||
``ipc_handles`` before constructing the typed dataclass, keeping
|
||||
serialization concerns out of the dataclass itself.
|
||||
|
||||
Requires ``VLLM_ALLOW_INSECURE_SERIALIZATION=1`` because the
|
||||
payload is deserialized via ``pickle.loads``.
|
||||
"""
|
||||
if "ipc_handles_pickled" in update_dict:
|
||||
if "ipc_handles" in update_dict:
|
||||
raise ValueError("Cannot specify both `ipc_handles` and `ipc_handles_pickled`")
|
||||
|
||||
if not envs.VLLM_ALLOW_INSECURE_SERIALIZATION:
|
||||
raise ValueError(
|
||||
"Refusing to deserialize `ipc_handles_pickled` without VLLM_ALLOW_INSECURE_SERIALIZATION=1"
|
||||
)
|
||||
|
||||
pickled = update_dict.pop("ipc_handles_pickled")
|
||||
update_dict["ipc_handles"] = pickle.loads(base64.b64decode(pickled))
|
||||
|
||||
return super().parse_update_info(update_dict)
|
||||
|
||||
def init_transfer_engine(self, init_info: NPUIPCWeightTransferInitInfo) -> None:
|
||||
"""No initialization needed for NPU IPC backend."""
|
||||
pass
|
||||
|
||||
def receive_weights(
|
||||
self,
|
||||
update_info: NPUIPCWeightTransferUpdateInfo,
|
||||
load_weights: Callable[[list[tuple[str, torch.Tensor]]], None],
|
||||
) -> None:
|
||||
"""Receive weights from the trainer via NPU IPC handles.
|
||||
|
||||
Args:
|
||||
update_info: NPU IPC update info containing parameter names,
|
||||
dtypes, shapes, and IPC handles.
|
||||
load_weights: Callable that loads weights into the model.
|
||||
"""
|
||||
device_index = torch.accelerator.current_device_index()
|
||||
physical_npu_id = npu_generate_uuid()
|
||||
|
||||
if update_info.packed:
|
||||
assert update_info.tensor_sizes is not None
|
||||
assert isinstance(update_info.ipc_handles, dict)
|
||||
weights = packed_npu_ipc_consumer(
|
||||
ipc_handle=update_info.ipc_handles,
|
||||
physical_npu_id=physical_npu_id,
|
||||
names=update_info.names,
|
||||
shapes=update_info.shapes,
|
||||
dtype_names=update_info.dtype_names,
|
||||
tensor_sizes=update_info.tensor_sizes,
|
||||
device_index=device_index,
|
||||
)
|
||||
load_weights(weights)
|
||||
else:
|
||||
# Lazy import: ``rebuild_npu_tensor`` lives in ``torch_npu`` and
|
||||
# must not be imported at module load time on non-NPU hosts.
|
||||
from torch_npu.multiprocessing.reductions import rebuild_npu_tensor
|
||||
|
||||
assert isinstance(update_info.ipc_handles, list)
|
||||
weights = []
|
||||
for name, ipc_handle in zip(
|
||||
update_info.names,
|
||||
update_info.ipc_handles,
|
||||
):
|
||||
if physical_npu_id not in ipc_handle:
|
||||
raise ValueError(
|
||||
f"IPC handle not found for NPU UUID {physical_npu_id}. "
|
||||
f"Available UUIDs: {list(ipc_handle.keys())}. "
|
||||
f"This may indicate that the trainer and worker are "
|
||||
f"not co-located on the same physical NPU (node)."
|
||||
)
|
||||
|
||||
args = ipc_handle[physical_npu_id]
|
||||
list_args = list(args)
|
||||
# Index 6 is the device_index parameter in torch's
|
||||
# IPC handle tuple (rebuild_npu_tensor). Update it
|
||||
# to the current device since the logical index can
|
||||
# differ between sender and receiver.
|
||||
list_args[6] = device_index
|
||||
weight = rebuild_npu_tensor(*list_args)
|
||||
weights.append((name, weight))
|
||||
|
||||
load_weights(weights)
|
||||
|
||||
def shutdown(self) -> None:
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
def trainer_send_weights(
|
||||
iterator: Iterator[tuple[str, torch.Tensor]],
|
||||
trainer_args: dict[str, Any] | NPUIPCTrainerSendWeightsArgs,
|
||||
) -> None:
|
||||
"""Send weights from trainer to inference workers via NPU IPC.
|
||||
|
||||
Supports two transport modes ('ray' and 'http') and two transfer
|
||||
strategies:
|
||||
- Non-packed (default): all weights in a single API call.
|
||||
- Packed (packed=True): chunked transfer with bounded NPU memory.
|
||||
|
||||
For multi-NPU training, all ranks must call this method in
|
||||
parallel. IPC handles are all-gathered across ranks and merged
|
||||
so that each vLLM worker can find its own NPU UUID. Only rank 0
|
||||
sends the payload to vLLM.
|
||||
|
||||
.. note::
|
||||
This method calls ``update_weights`` internally. The caller must
|
||||
handle ``pause`` / ``start_weight_update`` / ``finish_weight_update``
|
||||
/ ``resume`` before and after this method.
|
||||
|
||||
Args:
|
||||
iterator: Iterator of (name, tensor) pairs.
|
||||
trainer_args: NPUIPCTrainerSendWeightsArgs or equivalent dict.
|
||||
"""
|
||||
args = NPUIPCTrainerSendWeightsArgs(**trainer_args) if isinstance(trainer_args, dict) else trainer_args
|
||||
npu_uuid = npu_generate_uuid()
|
||||
if args.packed:
|
||||
NPUIPCWeightTransferEngine._send_packed(iterator, args, npu_uuid)
|
||||
else:
|
||||
NPUIPCWeightTransferEngine._send_unpacked(iterator, args, npu_uuid)
|
||||
|
||||
@staticmethod
|
||||
def _is_rank_zero() -> bool:
|
||||
"""Return True if this is rank 0 or no distributed group exists."""
|
||||
if not torch.distributed.is_initialized():
|
||||
return True
|
||||
return torch.distributed.get_rank() == 0
|
||||
|
||||
@staticmethod
|
||||
def _all_gather_and_merge_handles(
|
||||
handles: list[dict[str, tuple]],
|
||||
) -> list[dict[str, tuple]]:
|
||||
"""All-gather and merge IPC handle dicts across ranks.
|
||||
|
||||
Each rank contributes a list of ``{npu_uuid: ipc_args}`` dicts.
|
||||
Rank 0 collects and merges per-index; other ranks receive a list
|
||||
of empty dicts. No-op when no distributed group exists.
|
||||
"""
|
||||
if not torch.distributed.is_initialized() or torch.distributed.get_world_size() == 1:
|
||||
return handles
|
||||
|
||||
world_size = torch.distributed.get_world_size()
|
||||
gathered: list[list[dict[str, tuple]] | None] = [None] * world_size
|
||||
torch.distributed.all_gather_object(gathered, handles)
|
||||
torch.distributed.barrier()
|
||||
torch.npu.synchronize()
|
||||
|
||||
if torch.distributed.get_rank() == 0:
|
||||
merged: list[dict[str, tuple]] = []
|
||||
for param_idx in range(len(handles)):
|
||||
m: dict[str, tuple] = {}
|
||||
for rank_handles in gathered:
|
||||
if rank_handles is not None:
|
||||
m.update(rank_handles[param_idx])
|
||||
merged.append(m)
|
||||
return merged
|
||||
return [{} for _ in handles]
|
||||
|
||||
@staticmethod
|
||||
def _post_send_sync() -> None:
|
||||
"""Barrier + synchronize after a send; no-op if single-NPU."""
|
||||
if torch.distributed.is_initialized() and torch.distributed.get_world_size() > 1:
|
||||
torch.distributed.barrier()
|
||||
torch.npu.synchronize()
|
||||
|
||||
@staticmethod
|
||||
def _send_unpacked(
|
||||
iterator: Iterator[tuple[str, torch.Tensor]],
|
||||
args: NPUIPCTrainerSendWeightsArgs,
|
||||
npu_uuid: str,
|
||||
) -> None:
|
||||
"""Send all weights in a single API call (non-packed mode)."""
|
||||
names: list[str] = []
|
||||
dtype_names: list[str] = []
|
||||
shapes: list[list[int]] = []
|
||||
ipc_handles: list[dict[str, tuple]] = []
|
||||
# Hold strong refs to every contiguous copy until the send + post-send
|
||||
# sync completes. ``reduce_tensor``'s returned args do NOT keep
|
||||
# storage alive.
|
||||
weight_refs: list[torch.Tensor] = []
|
||||
|
||||
for name, tensor in iterator:
|
||||
names.append(name)
|
||||
dtype_names.append(str(tensor.dtype).split(".")[-1])
|
||||
shapes.append(list(tensor.shape))
|
||||
|
||||
weight = tensor.detach().contiguous()
|
||||
weight_refs.append(weight)
|
||||
# Store only the rebuild args (drop the func); the consumer rebuilds
|
||||
# with the well-known ``rebuild_npu_tensor``, mirroring upstream's
|
||||
# CUDA IPC engine.
|
||||
_, ipc_args = reduce_tensor(weight)
|
||||
ipc_handles.append({npu_uuid: ipc_args})
|
||||
|
||||
ipc_handles = NPUIPCWeightTransferEngine._all_gather_and_merge_handles(ipc_handles)
|
||||
|
||||
if NPUIPCWeightTransferEngine._is_rank_zero():
|
||||
NPUIPCWeightTransferEngine._do_send(
|
||||
args=args,
|
||||
names=names,
|
||||
dtype_names=dtype_names,
|
||||
shapes=shapes,
|
||||
ipc_handles=ipc_handles,
|
||||
)
|
||||
|
||||
NPUIPCWeightTransferEngine._post_send_sync()
|
||||
|
||||
@staticmethod
|
||||
def _send_packed(
|
||||
iterator: Iterator[tuple[str, torch.Tensor]],
|
||||
args: NPUIPCTrainerSendWeightsArgs,
|
||||
npu_uuid: str,
|
||||
) -> None:
|
||||
"""Send weights in bounded-memory chunks (packed mode)."""
|
||||
post_iter_func: Callable = lambda item: item[1]
|
||||
|
||||
for chunk in packed_npu_ipc_producer(
|
||||
iterator=iterator,
|
||||
npu_uuid=npu_uuid,
|
||||
post_iter_func=post_iter_func,
|
||||
buffer_size_bytes=args.packed_buffer_size_bytes,
|
||||
):
|
||||
ipc_handle = NPUIPCWeightTransferEngine._all_gather_and_merge_handles([chunk["ipc_handle"]])[0]
|
||||
|
||||
if NPUIPCWeightTransferEngine._is_rank_zero():
|
||||
NPUIPCWeightTransferEngine._do_send(
|
||||
args=args,
|
||||
names=chunk["names"],
|
||||
dtype_names=chunk["dtype_names"],
|
||||
shapes=chunk["shapes"],
|
||||
ipc_handles=ipc_handle,
|
||||
tensor_sizes=chunk["tensor_sizes"],
|
||||
packed=True,
|
||||
)
|
||||
|
||||
NPUIPCWeightTransferEngine._post_send_sync()
|
||||
|
||||
@staticmethod
|
||||
def _do_send(
|
||||
args: NPUIPCTrainerSendWeightsArgs,
|
||||
names: list[str],
|
||||
dtype_names: list[str],
|
||||
shapes: list[list[int]],
|
||||
ipc_handles: list[dict[str, tuple]] | dict[str, tuple],
|
||||
tensor_sizes: list[int] | None = None,
|
||||
packed: bool = False,
|
||||
) -> None:
|
||||
"""Send a single update payload via the configured transport."""
|
||||
update_fields: dict[str, Any] = {
|
||||
"names": names,
|
||||
"dtype_names": dtype_names,
|
||||
"shapes": shapes,
|
||||
"packed": packed,
|
||||
}
|
||||
if tensor_sizes is not None:
|
||||
update_fields["tensor_sizes"] = tensor_sizes
|
||||
|
||||
update_fields["ipc_handles"] = ipc_handles
|
||||
update_info = NPUIPCWeightTransferUpdateInfo(**update_fields)
|
||||
|
||||
if callable(args.send_mode):
|
||||
args.send_mode(update_info)
|
||||
elif args.send_mode == "ray":
|
||||
import ray
|
||||
|
||||
handles = args.llm_handle if isinstance(args.llm_handle, list) else [args.llm_handle]
|
||||
ray.get([h.update_weights.remote(dict(update_info=asdict(update_info))) for h in handles])
|
||||
elif args.send_mode == "http":
|
||||
pickled_handles = base64.b64encode(pickle.dumps(ipc_handles)).decode("utf-8")
|
||||
http_fields = {k: v for k, v in update_fields.items() if k != "ipc_handles"}
|
||||
http_fields["ipc_handles_pickled"] = pickled_handles
|
||||
|
||||
url = f"{args.url}/update_weights"
|
||||
payload = {"update_info": http_fields}
|
||||
response = requests.post(url, json=payload, timeout=300)
|
||||
response.raise_for_status()
|
||||
325
vllm_ascend/distributed/weight_transfer/packed_tensor.py
Normal file
325
vllm_ascend/distributed/weight_transfer/packed_tensor.py
Normal file
@@ -0,0 +1,325 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Packed tensor utilities for efficient weight transfer."""
|
||||
|
||||
import math
|
||||
from collections.abc import Callable, Iterator
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch.multiprocessing.reductions import reduce_tensor
|
||||
|
||||
# Default values for packed tensor configuration.
|
||||
# These are imported by HCCLWeightTransferUpdateInfo and trainer_send_weights.
|
||||
DEFAULT_PACKED_BUFFER_SIZE_BYTES = 1024 * 1024 * 1024 # 1GB
|
||||
DEFAULT_PACKED_NUM_BUFFERS = 2
|
||||
|
||||
|
||||
def packed_broadcast_producer(
|
||||
iterator: Iterator[tuple[str, torch.Tensor]],
|
||||
group: Any,
|
||||
src: int,
|
||||
post_iter_func: Callable[[tuple[str, torch.Tensor]], torch.Tensor],
|
||||
buffer_size_bytes: int = DEFAULT_PACKED_BUFFER_SIZE_BYTES,
|
||||
num_buffers: int = DEFAULT_PACKED_NUM_BUFFERS,
|
||||
) -> None:
|
||||
"""Broadcast tensors in a packed manner from trainer to workers.
|
||||
|
||||
Args:
|
||||
iterator: Iterator of model parameters. Returns a tuple of (name, tensor)
|
||||
group: Process group (PyHcclCommunicator)
|
||||
src: Source rank (0 in current implementation)
|
||||
post_iter_func: Function to apply to each (name, tensor) pair before
|
||||
packing, should return a tensor
|
||||
buffer_size_bytes: Size in bytes for each packed tensor buffer.
|
||||
Both producer and consumer must use the same value.
|
||||
num_buffers: Number of buffers for double/triple buffering.
|
||||
Both producer and consumer must use the same value.
|
||||
"""
|
||||
target_packed_tensor_size = buffer_size_bytes
|
||||
|
||||
streams = [torch.npu.Stream() for _ in range(num_buffers)]
|
||||
buffer_idx = 0
|
||||
|
||||
packing_tensor_list: list[list[torch.Tensor]] = [[] for _ in range(num_buffers)]
|
||||
packing_tensor_sizes: list[int] = [0 for _ in range(num_buffers)]
|
||||
packed_tensors: list[torch.Tensor] = [torch.empty(0, dtype=torch.uint8, device="npu") for _ in range(num_buffers)]
|
||||
|
||||
done = False
|
||||
while not done:
|
||||
# Synchronize the current stream (waits for previous
|
||||
# iteration's work on this buffer to finish)
|
||||
streams[buffer_idx].synchronize()
|
||||
# Start tasks for the new buffer in a new stream
|
||||
with torch.npu.stream(streams[buffer_idx]):
|
||||
# Initialize the packing tensor list and sizes
|
||||
packing_tensor_list[buffer_idx] = []
|
||||
packing_tensor_sizes[buffer_idx] = 0
|
||||
# Pack the tensors
|
||||
while True:
|
||||
try:
|
||||
item = next(iterator)
|
||||
except StopIteration:
|
||||
done = True
|
||||
break
|
||||
# Apply post processing and convert to linearized uint8 tensor
|
||||
tensor = post_iter_func(item).contiguous().view(torch.uint8).view(-1)
|
||||
packing_tensor_list[buffer_idx].append(tensor)
|
||||
packing_tensor_sizes[buffer_idx] += tensor.numel()
|
||||
if packing_tensor_sizes[buffer_idx] > target_packed_tensor_size:
|
||||
break
|
||||
if len(packing_tensor_list[buffer_idx]) > 0:
|
||||
# Pack the tensors
|
||||
packed_tensors[buffer_idx] = torch.cat(packing_tensor_list[buffer_idx], dim=0)
|
||||
|
||||
if len(packing_tensor_list[buffer_idx]) == 0:
|
||||
# No more tensors — nothing left to broadcast
|
||||
break
|
||||
|
||||
# torch.cat runs on the custom stream. Synchronize before
|
||||
# broadcasting on the default stream so the packed data is ready.
|
||||
streams[buffer_idx].synchronize()
|
||||
group.broadcast(packed_tensors[buffer_idx], src=src)
|
||||
|
||||
# Move to the next buffer
|
||||
buffer_idx = (buffer_idx + 1) % num_buffers
|
||||
|
||||
# Ensure the last broadcast on the default stream has completed
|
||||
# before returning, so NPU tensor cleanup at exit doesn't hang.
|
||||
torch.npu.current_stream().synchronize()
|
||||
|
||||
|
||||
def packed_broadcast_consumer(
|
||||
iterator: Iterator[tuple[str, tuple[list[int], torch.dtype]]],
|
||||
group: Any,
|
||||
src: int,
|
||||
post_unpack_func: Callable[[list[tuple[str, torch.Tensor]]], None],
|
||||
buffer_size_bytes: int = DEFAULT_PACKED_BUFFER_SIZE_BYTES,
|
||||
num_buffers: int = DEFAULT_PACKED_NUM_BUFFERS,
|
||||
) -> None:
|
||||
"""Consume packed tensors and unpack them into a list of tensors.
|
||||
|
||||
Args:
|
||||
iterator: Iterator of parameter metadata. Returns (name, (shape, dtype))
|
||||
group: Process group (PyHcclCommunicator)
|
||||
src: Source rank (0 in current implementation)
|
||||
post_unpack_func: Function to apply to each list of (name, tensor) after
|
||||
unpacking
|
||||
buffer_size_bytes: Size in bytes for each packed tensor buffer.
|
||||
Both producer and consumer must use the same value.
|
||||
num_buffers: Number of buffers for double/triple buffering.
|
||||
Both producer and consumer must use the same value.
|
||||
"""
|
||||
|
||||
def unpack_tensor(
|
||||
packed_tensor: torch.Tensor,
|
||||
names: list[str],
|
||||
shapes: list[list[int]],
|
||||
dtypes: list[torch.dtype],
|
||||
tensor_sizes: list[int],
|
||||
) -> list[tuple[str, torch.Tensor]]:
|
||||
"""Unpack a packed uint8 tensor into a list of typed tensors."""
|
||||
unpacked_tensors = packed_tensor.split(tensor_sizes)
|
||||
unpacked_list = [
|
||||
(name, tensor.contiguous().view(dtype).view(*shape))
|
||||
for name, shape, dtype, tensor in zip(names, shapes, dtypes, unpacked_tensors)
|
||||
]
|
||||
return unpacked_list
|
||||
|
||||
target_packed_tensor_size = buffer_size_bytes
|
||||
|
||||
streams = [torch.npu.Stream() for _ in range(num_buffers)]
|
||||
default_stream = torch.npu.current_stream()
|
||||
buffer_idx = 0
|
||||
|
||||
packing_tensor_meta_data: list[list[tuple[str, list[int], torch.dtype, int]]] = [[] for _ in range(num_buffers)]
|
||||
packing_tensor_sizes: list[int] = [0 for _ in range(num_buffers)]
|
||||
packed_tensors: list[torch.Tensor] = [torch.empty(0, dtype=torch.uint8, device="npu") for _ in range(num_buffers)]
|
||||
|
||||
done = False
|
||||
while not done:
|
||||
# Synchronize the current stream (waits for previous
|
||||
# iteration's load_weights on this buffer to finish)
|
||||
streams[buffer_idx].synchronize()
|
||||
with torch.npu.stream(streams[buffer_idx]):
|
||||
# Collect parameter metadata for this buffer
|
||||
packing_tensor_meta_data[buffer_idx] = []
|
||||
packing_tensor_sizes[buffer_idx] = 0
|
||||
while True:
|
||||
try:
|
||||
name, (shape, dtype) = next(iterator)
|
||||
except StopIteration:
|
||||
done = True
|
||||
break
|
||||
tensor_size = math.prod(shape) * dtype.itemsize
|
||||
packing_tensor_meta_data[buffer_idx].append((name, shape, dtype, tensor_size))
|
||||
packing_tensor_sizes[buffer_idx] += tensor_size
|
||||
if packing_tensor_sizes[buffer_idx] > target_packed_tensor_size:
|
||||
break
|
||||
if len(packing_tensor_meta_data[buffer_idx]) > 0:
|
||||
packed_tensors[buffer_idx] = torch.empty(
|
||||
packing_tensor_sizes[buffer_idx],
|
||||
dtype=torch.uint8,
|
||||
device="npu",
|
||||
)
|
||||
|
||||
if len(packing_tensor_meta_data[buffer_idx]) == 0:
|
||||
break
|
||||
|
||||
# Broadcast on the default stream.
|
||||
group.broadcast(packed_tensors[buffer_idx], src=src)
|
||||
|
||||
# Synchronize the default stream so broadcast completes before
|
||||
# load_weights (running on the custom stream) reads the data.
|
||||
default_stream.synchronize()
|
||||
|
||||
# Unpack and load weights on the custom stream
|
||||
with torch.npu.stream(streams[buffer_idx]):
|
||||
names, shapes, dtypes, tensor_sizes = zip(*packing_tensor_meta_data[buffer_idx])
|
||||
post_unpack_func(
|
||||
unpack_tensor(
|
||||
packed_tensors[buffer_idx],
|
||||
list(names),
|
||||
list(shapes),
|
||||
list(dtypes),
|
||||
list(tensor_sizes),
|
||||
)
|
||||
)
|
||||
|
||||
# Move to the next buffer
|
||||
buffer_idx = (buffer_idx + 1) % num_buffers
|
||||
|
||||
# Wait for all in-flight load_weights (on custom streams) to finish.
|
||||
# Otherwise NPU tensor cleanup at exit may hang.
|
||||
for s in streams:
|
||||
s.synchronize()
|
||||
|
||||
|
||||
# ── NPU IPC packed transfer ────────────────────────────────────────────
|
||||
|
||||
|
||||
def packed_npu_ipc_producer(
|
||||
iterator: Iterator[tuple[str, torch.Tensor]],
|
||||
npu_uuid: str,
|
||||
post_iter_func: Callable[[tuple[str, torch.Tensor]], torch.Tensor],
|
||||
buffer_size_bytes: int = DEFAULT_PACKED_BUFFER_SIZE_BYTES,
|
||||
) -> Iterator[dict[str, Any]]:
|
||||
"""Pack tensors into a reusable NPU IPC buffer and yield chunks.
|
||||
|
||||
Allocates a single NPU buffer of ``buffer_size_bytes`` and registers
|
||||
it for IPC once via ``reduce_tensor``. Each chunk's packed data is
|
||||
copied into this buffer before yielding, so only one IPC-shared
|
||||
allocation exists for the lifetime of the transfer.
|
||||
|
||||
Args:
|
||||
iterator: Iterator of (name, tensor) pairs.
|
||||
npu_uuid: Physical NPU UUID string for this rank.
|
||||
post_iter_func: Applied to each (name, tensor) before packing.
|
||||
buffer_size_bytes: Exact capacity of the reusable IPC buffer.
|
||||
"""
|
||||
ipc_buffer = torch.empty(buffer_size_bytes, dtype=torch.uint8, device="npu")
|
||||
# Store only the rebuild args (drop the func); the consumer rebuilds with
|
||||
# the well-known ``rebuild_npu_tensor``, mirroring upstream's CUDA IPC engine.
|
||||
_, ipc_args = reduce_tensor(ipc_buffer)
|
||||
|
||||
names: list[str] = []
|
||||
shapes: list[list[int]] = []
|
||||
dtypes: list[torch.dtype] = []
|
||||
tensor_sizes: list[int] = []
|
||||
total_bytes = 0
|
||||
|
||||
for name, orig_tensor in iterator:
|
||||
flat = post_iter_func((name, orig_tensor)).contiguous().view(torch.uint8).view(-1)
|
||||
|
||||
if flat.numel() > buffer_size_bytes:
|
||||
raise ValueError(
|
||||
f"Tensor '{name}' has size {flat.numel()} bytes, "
|
||||
f"which exceeds buffer_size_bytes={buffer_size_bytes}. "
|
||||
f"Increase buffer_size_bytes to at least {flat.numel()}."
|
||||
)
|
||||
|
||||
if total_bytes and total_bytes + flat.numel() > buffer_size_bytes:
|
||||
torch.npu.current_stream().synchronize()
|
||||
yield {
|
||||
"names": names,
|
||||
"shapes": shapes,
|
||||
"dtype_names": [str(d).split(".")[-1] for d in dtypes],
|
||||
"tensor_sizes": tensor_sizes,
|
||||
"ipc_handle": {npu_uuid: ipc_args},
|
||||
}
|
||||
names, shapes, dtypes, tensor_sizes = [], [], [], []
|
||||
total_bytes = 0
|
||||
|
||||
ipc_buffer[total_bytes : total_bytes + flat.numel()].copy_(flat)
|
||||
names.append(name)
|
||||
shapes.append(list(orig_tensor.shape))
|
||||
dtypes.append(orig_tensor.dtype)
|
||||
tensor_sizes.append(flat.numel())
|
||||
total_bytes += flat.numel()
|
||||
|
||||
if total_bytes:
|
||||
torch.npu.current_stream().synchronize()
|
||||
yield {
|
||||
"names": names,
|
||||
"shapes": shapes,
|
||||
"dtype_names": [str(d).split(".")[-1] for d in dtypes],
|
||||
"tensor_sizes": tensor_sizes,
|
||||
"ipc_handle": {npu_uuid: ipc_args},
|
||||
}
|
||||
|
||||
|
||||
def packed_npu_ipc_consumer(
|
||||
ipc_handle: dict[str, tuple],
|
||||
physical_npu_id: str,
|
||||
names: list[str],
|
||||
shapes: list[list[int]],
|
||||
dtype_names: list[str],
|
||||
tensor_sizes: list[int],
|
||||
device_index: int,
|
||||
) -> list[tuple[str, torch.Tensor]]:
|
||||
"""Unpack a single packed IPC chunk into named tensors.
|
||||
|
||||
Reconstructs the packed buffer via the IPC handle, unpacks into
|
||||
individual tensors, and clones each into independent storage before
|
||||
returning. The clone is required because the producer reuses one
|
||||
IPC buffer across chunks.
|
||||
|
||||
Args:
|
||||
ipc_handle: Mapping of NPU UUID to a ``rebuild_npu_tensor`` args tuple
|
||||
from ``reduce_tensor``.
|
||||
physical_npu_id: Physical NPU UUID string for the current process.
|
||||
names: Parameter names in the packed buffer.
|
||||
shapes: Parameter shapes.
|
||||
dtype_names: Parameter dtype name strings (e.g. "float16").
|
||||
tensor_sizes: Size in bytes of each parameter in the packed buffer.
|
||||
device_index: Local NPU device index.
|
||||
"""
|
||||
# Lazy import: ``rebuild_npu_tensor`` lives in ``torch_npu`` and must not be
|
||||
# imported at module load time on non-NPU hosts.
|
||||
from torch_npu.multiprocessing.reductions import rebuild_npu_tensor
|
||||
|
||||
if physical_npu_id not in ipc_handle:
|
||||
raise ValueError(
|
||||
f"IPC handle not found for NPU UUID {physical_npu_id}. Available UUIDs: {list(ipc_handle.keys())}"
|
||||
)
|
||||
|
||||
args = ipc_handle[physical_npu_id]
|
||||
list_args = list(args)
|
||||
# Index 6 of the args from reduce_tensor is the device_index.
|
||||
# Overwrite it with the receiver's device index.
|
||||
list_args[6] = device_index
|
||||
packed = rebuild_npu_tensor(*list_args)
|
||||
|
||||
content_size = sum(tensor_sizes)
|
||||
packed = packed[:content_size]
|
||||
|
||||
dtypes = [getattr(torch, dn) for dn in dtype_names]
|
||||
weights: list[tuple[str, torch.Tensor]] = []
|
||||
offset = 0
|
||||
for name, shape, dtype, size in zip(names, shapes, dtypes, tensor_sizes):
|
||||
raw = packed[offset : offset + size]
|
||||
tensor = raw.contiguous().view(dtype).view(*shape).clone()
|
||||
weights.append((name, tensor))
|
||||
offset += size
|
||||
|
||||
return weights
|
||||
Reference in New Issue
Block a user