init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View 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",
)

View 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

View 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()

View 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