0
vllm_ascend/model_loader/__init__.py
Normal file
0
vllm_ascend/model_loader/__init__.py
Normal file
20
vllm_ascend/model_loader/netloader/__init__.py
Normal file
20
vllm_ascend/model_loader/netloader/__init__.py
Normal file
@@ -0,0 +1,20 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
|
||||
def register_netloader():
|
||||
"""Register the NetLoader plugin."""
|
||||
from .netloader import ModelNetLoaderElastic # noqa
|
||||
161
vllm_ascend/model_loader/netloader/executor/elastic_load.py
Normal file
161
vllm_ascend/model_loader/netloader/executor/elastic_load.py
Normal file
@@ -0,0 +1,161 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
import torch
|
||||
import torch_npu
|
||||
from vllm.logger import logger
|
||||
|
||||
from .netloader_pg import destroy_stateless_process_group, stateless_init_process_group
|
||||
|
||||
|
||||
class P2PLoad:
|
||||
"""
|
||||
Class for receiving model parameters in a distributed manner using HCCL backend.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
world_name: str,
|
||||
source_ip: str,
|
||||
source_port: int,
|
||||
group_name: str = "netloader",
|
||||
):
|
||||
"""
|
||||
Initializes the P2PLoad instance.
|
||||
|
||||
Parameters:
|
||||
- world_name: The name of the distributed group.
|
||||
- source_ip: The IP address of the source node.
|
||||
- source_port: The port number for the source node.
|
||||
- group_name: Name of the HCCL process group.
|
||||
"""
|
||||
self.world_name = world_name
|
||||
self.source_ip = source_ip
|
||||
self.source_port = source_port
|
||||
self.group_name = group_name
|
||||
|
||||
def load(self, model):
|
||||
"""
|
||||
Loads the model parameters using HCCL backend.
|
||||
|
||||
Parameters:
|
||||
- model: The model whose parameters are to be loaded.
|
||||
|
||||
Returns:
|
||||
- The model if loading is successful, otherwise None.
|
||||
"""
|
||||
model_device = next(model.parameters()).device
|
||||
logger.info(
|
||||
"Start init_process_group, name: %s, addr: %s:%s", self.world_name, self.source_ip, self.source_port
|
||||
)
|
||||
receiver_pg = None
|
||||
loaded_model = None
|
||||
try:
|
||||
receiver_pg = stateless_init_process_group(
|
||||
host=self.world_name.split(":")[0],
|
||||
port=self.source_port,
|
||||
rank=0,
|
||||
world_size=2,
|
||||
group_name=self.group_name,
|
||||
)
|
||||
logger.info(
|
||||
"Finish init_process_group, name: %s, addr: %s:%s", self.world_name, self.source_ip, self.source_port
|
||||
)
|
||||
|
||||
logger.info("Start recv, name: %s, addr: %s:%s", self.world_name, self.source_ip, self.source_port)
|
||||
logger.info("Model device: %s", model_device)
|
||||
|
||||
trans_stream = torch_npu.npu.Stream()
|
||||
with torch_npu.npu.stream(trans_stream):
|
||||
for name, param in model.named_parameters():
|
||||
if len(param.shape) == 0:
|
||||
continue
|
||||
receiver_pg.recv([param], 1, 0).wait()
|
||||
torch.distributed.barrier(group=receiver_pg, device_ids=[model_device.index])
|
||||
|
||||
torch_npu.npu.synchronize(trans_stream)
|
||||
|
||||
logger.info("Finish recv, name: %s, addr: %s:%s", self.world_name, self.source_ip, self.source_port)
|
||||
loaded_model = model
|
||||
except Exception as e:
|
||||
logger.error("Failed to recv model: %s", e)
|
||||
finally:
|
||||
if receiver_pg:
|
||||
destroy_stateless_process_group(receiver_pg)
|
||||
return loaded_model
|
||||
|
||||
|
||||
class P2PSend:
|
||||
"""
|
||||
Class for sending model parameters in a distributed manner using HCCL backend.
|
||||
"""
|
||||
|
||||
def __init__(self, listen_ip: str, listen_port: int, comm_name: str, group_name: str = "netloader"):
|
||||
"""
|
||||
Initializes the P2PSend instance.
|
||||
|
||||
Parameters:
|
||||
- listen_ip: The IP address to listen on.
|
||||
- listen_port: The port number to listen on.
|
||||
- comm_name: The name of the communication group.
|
||||
- group_name: Name of the HCCL process group.
|
||||
"""
|
||||
self.listen_ip = listen_ip
|
||||
self.listen_port = listen_port
|
||||
self.comm_name = comm_name
|
||||
self.group_name = group_name
|
||||
|
||||
def send(self, model, int8_params: dict):
|
||||
"""
|
||||
Sends the model parameters using HCCL backend.
|
||||
|
||||
Parameters:
|
||||
- model: The model whose parameters are to be sent.
|
||||
- int8_params: Dictionary of parameters that are in int8 format.
|
||||
"""
|
||||
model_device = next(model.parameters()).device
|
||||
torch.npu.set_device(model_device)
|
||||
logger.info("Start init_process_group, name: %s, addr: %s:%s", self.comm_name, self.listen_ip, self.listen_port)
|
||||
sender_pg = None
|
||||
try:
|
||||
sender_pg = stateless_init_process_group(
|
||||
host=self.comm_name.split(":")[0],
|
||||
port=self.listen_port,
|
||||
rank=1,
|
||||
world_size=2,
|
||||
group_name=self.group_name,
|
||||
)
|
||||
logger.info(
|
||||
"Finish init_process_group, name: %s, addr: %s:%s", self.comm_name, self.listen_ip, self.listen_port
|
||||
)
|
||||
logger.info("Start send, name: %s, addr: %s:%s", self.comm_name, self.listen_ip, self.listen_port)
|
||||
logger.info("Model device: %s", model_device)
|
||||
|
||||
trans_stream = torch_npu.npu.Stream()
|
||||
with torch_npu.npu.stream(trans_stream):
|
||||
for name, param in model.named_parameters():
|
||||
if "aclnn_input_scale" in name:
|
||||
continue
|
||||
if name in int8_params:
|
||||
sender_pg.send([int8_params[name].to(model_device)], 0, 0).wait()
|
||||
else:
|
||||
sender_pg.send([param.contiguous()], 0, 0).wait()
|
||||
torch.distributed.barrier(group=sender_pg, device_ids=[model_device.index])
|
||||
torch_npu.npu.synchronize(trans_stream)
|
||||
logger.info("Finish send, name: %s, addr: %s:%s", self.comm_name, self.listen_ip, self.listen_port)
|
||||
finally:
|
||||
if sender_pg:
|
||||
destroy_stateless_process_group(sender_pg)
|
||||
180
vllm_ascend/model_loader/netloader/executor/netloader_pg.py
Normal file
180
vllm_ascend/model_loader/netloader/executor/netloader_pg.py
Normal file
@@ -0,0 +1,180 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
import gc
|
||||
import ipaddress
|
||||
from datetime import timedelta
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch_npu
|
||||
from torch._C._distributed_c10d import _DEFAULT_PG_TIMEOUT, _register_process_group, _unregister_process_group
|
||||
from torch.distributed import ProcessGroup, is_hccl_available
|
||||
from torch.distributed.distributed_c10d import Backend, BackendConfig, PrefixStore, _world
|
||||
from torch.distributed.rendezvous import rendezvous
|
||||
from torch_npu._C._distributed_c10d import ProcessGroupHCCL
|
||||
from vllm.logger import logger
|
||||
|
||||
|
||||
def stateless_init_process_group(
|
||||
host: str,
|
||||
port: int,
|
||||
world_size: int,
|
||||
rank: int,
|
||||
timeout: timedelta = _DEFAULT_PG_TIMEOUT,
|
||||
group_name: str = "",
|
||||
pg_options: Any | None = None,
|
||||
) -> ProcessGroup:
|
||||
"""
|
||||
Initializes a stateless process group.
|
||||
|
||||
Args:
|
||||
host: Hostname.
|
||||
port: Port number.
|
||||
world_size: Size of the process group.
|
||||
rank: Rank of the current process.
|
||||
timeout: Timeout duration, defaults to _DEFAULT_PG_TIMEOUT.
|
||||
group_name: Name of the process group, defaults to an empty string.
|
||||
pg_options: Options for the process group, defaults to None.
|
||||
|
||||
Returns:
|
||||
ProcessGroup: The initialized process group.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If world_size is not positive, or if rank is not within
|
||||
[0, world_size - 1], or if HCCL is unavailable.
|
||||
TypeError: If timeout is not a timedelta type.
|
||||
ValueError: If group_name already exists.
|
||||
"""
|
||||
|
||||
# Check if world_size is positive
|
||||
if not world_size > 0:
|
||||
raise RuntimeError("world_size must be positive")
|
||||
# Check if rank is within [0, world_size - 1]
|
||||
if not (rank >= 0 and rank <= world_size - 1):
|
||||
raise RuntimeError("rank should be a number between 0 and ``world_size``-1")
|
||||
# Check if HCCL is available
|
||||
if not is_hccl_available():
|
||||
raise RuntimeError("HCCL is not available")
|
||||
# Check if timeout is a timedelta type
|
||||
if not isinstance(timeout, timedelta):
|
||||
raise TypeError(f"Expected timeout argument to be of type datetime.timedelta, got {timeout}")
|
||||
# Check if group_name already exists
|
||||
if group_name in _world.pg_names.values():
|
||||
raise ValueError(
|
||||
f"The specified group name {group_name} has already been created, please use a different group name"
|
||||
)
|
||||
|
||||
# Function to check if an IPv6 address is valid
|
||||
def is_valid_ipv6_address(address: str) -> bool:
|
||||
try:
|
||||
ipaddress.IPv6Address(address)
|
||||
return True
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
# Function to get TCP URI
|
||||
def get_tcp_uri(ip: str, port: int) -> str:
|
||||
if is_valid_ipv6_address(ip):
|
||||
return f"tcp://[{ip}]:{port}"
|
||||
else:
|
||||
return f"tcp://{ip}:{port}"
|
||||
|
||||
# Get initialization method
|
||||
init_method = get_tcp_uri(host, port)
|
||||
# Create Backend object
|
||||
backend = Backend("hccl")
|
||||
# Use rendezvous function to get store, rank, and world_size
|
||||
store, rank, world_size = next(rendezvous(init_method, rank, world_size, timeout=timeout))
|
||||
|
||||
# Set timeout for store
|
||||
store.set_timeout(timeout)
|
||||
# Create PrefixStore object
|
||||
prefix_store = PrefixStore(f"{init_method}/{group_name}/", store)
|
||||
# Set group_rank and group_size
|
||||
group_rank = rank
|
||||
group_size = world_size
|
||||
# Create ProcessGroup object
|
||||
pg: ProcessGroup = ProcessGroup(
|
||||
prefix_store,
|
||||
group_rank,
|
||||
group_size,
|
||||
)
|
||||
# Create BackendConfig object
|
||||
backend_config = BackendConfig(backend)
|
||||
# Set default backend for ProcessGroup
|
||||
pg._set_default_backend(Backend.backend_type_map[backend])
|
||||
|
||||
# Check if pg_options is None or not of type ProcessGroupHCCL.Options
|
||||
if pg_options is None or not isinstance(pg_options, torch_npu._C._distributed_c10d.ProcessGroupHCCL.Options):
|
||||
pg_options = torch_npu._C._distributed_c10d.ProcessGroupHCCL.Options()
|
||||
# Set attributes for pg_options
|
||||
pg_options.is_high_priority_stream = False
|
||||
pg_options._timeout = timeout
|
||||
pg_options.global_ranks_in_group = []
|
||||
pg_options.group_id = f"{init_method}/{group_name}/"
|
||||
# Create ProcessGroupHCCL object
|
||||
backend_class = ProcessGroupHCCL(prefix_store, group_rank, group_size, pg_options)
|
||||
# Set sequence number for backend_class
|
||||
backend_class._set_sequence_number_for_group()
|
||||
# Set backend_type
|
||||
backend_type = ProcessGroup.BackendType.CUSTOM
|
||||
# Register backend
|
||||
pg._register_backend(torch.device("npu"), backend_type, backend_class)
|
||||
|
||||
# Set group_desc and pg_tag
|
||||
group_desc = "undefined"
|
||||
assert group_name is not None
|
||||
assert group_desc is not None
|
||||
pg._set_group_name(group_name)
|
||||
pg._set_group_desc(group_desc)
|
||||
|
||||
# Update attributes in _world
|
||||
_world.pg_group_ranks[pg] = {i: i for i in range(world_size)}
|
||||
_world.pg_map[pg] = (backend, prefix_store)
|
||||
_world.pg_names[pg] = group_name
|
||||
_register_process_group(group_name, pg)
|
||||
_world.pg_backend_config[pg] = str(backend_config)
|
||||
return pg
|
||||
|
||||
|
||||
def destroy_stateless_process_group(pg: ProcessGroup, manual_gc: bool = False):
|
||||
"""
|
||||
Destroy a stateless process group.
|
||||
|
||||
Args:
|
||||
pg: Process group to be destroyed.
|
||||
manual_gc: Whether to manually perform garbage collection, defaults to False.
|
||||
"""
|
||||
# Shutdown the process group
|
||||
pg.shutdown()
|
||||
# Remove related attributes from _world
|
||||
_world.pg_map.pop(pg, None)
|
||||
_world.pg_names.pop(pg, None)
|
||||
_world.pg_group_ranks.pop(pg, None)
|
||||
_world.pg_backend_config.pop(pg, None)
|
||||
# Check if pg is in keys of _world.pg_coalesce_state
|
||||
if pg in _world.pg_coalesce_state:
|
||||
logger.warning(
|
||||
"Some coalesced collectives haven't been launched when ProcessGroup is destroyed. They will be cleaned."
|
||||
)
|
||||
del _world.pg_coalesce_state[pg]
|
||||
# Unregister the process group
|
||||
_unregister_process_group(pg.group_name)
|
||||
|
||||
# If manual_gc is True, perform garbage collection
|
||||
if manual_gc:
|
||||
gc.collect()
|
||||
422
vllm_ascend/model_loader/netloader/interaction/elastic.py
Normal file
422
vllm_ascend/model_loader/netloader/interaction/elastic.py
Normal file
@@ -0,0 +1,422 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
import json
|
||||
import socket
|
||||
import threading
|
||||
from contextlib import suppress
|
||||
|
||||
import regex as re
|
||||
import torch
|
||||
from vllm.logger import logger
|
||||
|
||||
from ..executor.elastic_load import P2PSend
|
||||
from ..utils import find_free_port
|
||||
|
||||
|
||||
class ElasticClient:
|
||||
"""
|
||||
Class for handling the client-side logic of Netloader of models.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, sources: list[str], device_id: int, model_path: str, tp: int, pp: int, group_name: str = "netloader"
|
||||
):
|
||||
"""
|
||||
Initializes the ElasticClient instance.
|
||||
|
||||
Parameters:
|
||||
- sources: List of source addresses in the format IP:port.
|
||||
- device_id: The ID of the current device.
|
||||
- model_path: The path to the model.
|
||||
- tp: Tensor parallel size.
|
||||
- pp: Pipeline parallel size.
|
||||
- group_name: Name of the HCCL process group.
|
||||
"""
|
||||
self.sources = sources
|
||||
self.device_id = device_id
|
||||
self.model_path = model_path
|
||||
self.tp = tp
|
||||
self.pp = pp
|
||||
self.group_name = group_name
|
||||
|
||||
self.s: socket.socket | None = None
|
||||
self.ack: tuple[str, int] | None = None
|
||||
self.server_addr: str | None = None
|
||||
self.server_port: int | None = None
|
||||
|
||||
for source in self.sources:
|
||||
try:
|
||||
ip, port_str = source.split(":")
|
||||
port = int(port_str)
|
||||
except Exception as e:
|
||||
logger.info("IP format error: %s, detail: %s", source, e)
|
||||
continue
|
||||
|
||||
self.server_addr = ip
|
||||
self.server_port = port
|
||||
|
||||
try:
|
||||
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
||||
logger.info("Start connection to server: %s:%s", self.server_addr, self.server_port)
|
||||
sock.connect((self.server_addr, self.server_port))
|
||||
logger.info("Finish connection to server: %s:%s", self.server_addr, self.server_port)
|
||||
sock.settimeout(60)
|
||||
|
||||
self.s = sock
|
||||
self.ack = self.register(device_id, model_path, tp, pp)
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error("Connect to %s fails, detail: %s", source, e)
|
||||
if sock is not None:
|
||||
with suppress(Exception):
|
||||
sock.close()
|
||||
self.s = None
|
||||
self.ack = None
|
||||
self.server_addr = None
|
||||
self.server_port = None
|
||||
|
||||
if self.s is None:
|
||||
sources_str = ", ".join(self.sources[:2])
|
||||
if len(self.sources) > 2:
|
||||
sources_str += f", ... (total {len(self.sources)})"
|
||||
logger.error(
|
||||
"All sources exhausted, no connection established for device_id=%s, model_path=%s, sources=[%s]",
|
||||
device_id,
|
||||
model_path,
|
||||
sources_str,
|
||||
)
|
||||
|
||||
def close(self) -> None:
|
||||
"""
|
||||
Closes the socket connection.
|
||||
"""
|
||||
if self.s is not None:
|
||||
try:
|
||||
self.s.close()
|
||||
except Exception as e:
|
||||
logger.error("Error closing socket: %s", e)
|
||||
finally:
|
||||
self.s = None
|
||||
|
||||
def __enter__(self) -> "ElasticClient":
|
||||
"""
|
||||
Context manager enter method.
|
||||
"""
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb) -> None:
|
||||
"""
|
||||
Context manager exit method.
|
||||
"""
|
||||
self.close()
|
||||
|
||||
def __del__(self):
|
||||
"""
|
||||
Destructor method to ensure socket is closed.
|
||||
"""
|
||||
with suppress(Exception):
|
||||
self.close()
|
||||
|
||||
def send_str(self, data_str: str) -> None:
|
||||
"""
|
||||
Sends a string over the socket connection.
|
||||
|
||||
Parameters:
|
||||
- data_str: The string to be sent.
|
||||
"""
|
||||
if self.s is None:
|
||||
raise RuntimeError("Socket was not created correctly.")
|
||||
self.s.send(data_str.encode("utf-8"))
|
||||
|
||||
def recv_str(self, buffer_size: int = 1024) -> str:
|
||||
"""
|
||||
Receives a string over the socket connection.
|
||||
|
||||
Parameters:
|
||||
- buffer_size: The size of the buffer for receiving data.
|
||||
|
||||
Returns:
|
||||
- The received string.
|
||||
"""
|
||||
if self.s is None:
|
||||
raise RuntimeError("Socket was not created correctly.")
|
||||
data_str = self.s.recv(buffer_size).decode("utf-8")
|
||||
return data_str
|
||||
|
||||
def register(self, device_id: int, model_path: str, tp: int, pp: int) -> tuple[str, int]:
|
||||
"""
|
||||
Registers the client with the server.
|
||||
|
||||
Parameters:
|
||||
- device_id: The ID of the current device.
|
||||
- model_path: The path to the model.
|
||||
- tp: Tensor parallel size.
|
||||
- pp: Pipeline parallel size.
|
||||
|
||||
Returns:
|
||||
- A tuple containing the communication name and port.
|
||||
"""
|
||||
free_port = find_free_port()
|
||||
data = {
|
||||
"label": "JOIN",
|
||||
"content": {
|
||||
"device_id": device_id,
|
||||
"model_path": model_path,
|
||||
"tp": tp,
|
||||
"pp": pp,
|
||||
"port": free_port,
|
||||
"group_name": self.group_name,
|
||||
},
|
||||
}
|
||||
|
||||
try:
|
||||
self.send_str(json.dumps(data))
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Send data {data} to server fails, detail: {e}")
|
||||
|
||||
try:
|
||||
ack_str = self.recv_str()
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Receive data from server fails, detail: {e}")
|
||||
|
||||
try:
|
||||
ack = json.loads(ack_str)
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Receive data {ack_str} cannot be converted to JSON format, detail: {e}")
|
||||
|
||||
logger.info("Receive ack: %s", ack)
|
||||
|
||||
if (
|
||||
"label" in ack
|
||||
and ack["label"] == "JOIN_ACK"
|
||||
and "content" in ack
|
||||
and ack["content"] is not None
|
||||
and "name" in ack["content"]
|
||||
):
|
||||
return (ack["content"]["name"], free_port)
|
||||
elif "label" in ack and ack["label"] == "JOIN_NACK" and "content" in ack:
|
||||
raise RuntimeError(f"Receive nack from server, reason: {ack['content']}")
|
||||
else:
|
||||
raise RuntimeError(f"Receive ack {ack} from server does not contain required fields")
|
||||
|
||||
|
||||
class ElasticServer:
|
||||
"""
|
||||
Class for handling the server-side logic of Netloader of models.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
addr: str,
|
||||
port: int,
|
||||
model,
|
||||
device_id: int,
|
||||
model_path: str,
|
||||
tp: int,
|
||||
pp: int,
|
||||
int8_cache: str,
|
||||
int8_cache_name: list[str] | None,
|
||||
group_name: str = "netloader",
|
||||
):
|
||||
"""
|
||||
Initializes the ElasticServer instance.
|
||||
|
||||
Parameters:
|
||||
- addr: The IP address to listen on.
|
||||
- port: The port number to listen on.
|
||||
- model: The model to be served.
|
||||
- device_id: The ID of the current device (i.e. global rank).
|
||||
- model_path: The path to the model.
|
||||
- tp: Tensor parallel size.
|
||||
- pp: Pipeline parallel size.
|
||||
- int8_cache: The type of caching for int8 parameters (HBM, DRAM, or no).
|
||||
- int8_cache_name: List of parameter names to be cached.
|
||||
- group_name: Name of the HCCL process group.
|
||||
"""
|
||||
self.addr = addr
|
||||
self.port = port
|
||||
self.s = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
||||
self.s.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
self.s.bind((self.addr, self.port))
|
||||
self.s.listen(256)
|
||||
|
||||
self.model = model
|
||||
self.device_id = device_id
|
||||
self.model_path = model_path
|
||||
self.tp = tp
|
||||
self.pp = pp
|
||||
self.group_name = group_name
|
||||
|
||||
self.original_int8 = {}
|
||||
int8_pattern = "|".join(map(re.escape, int8_cache_name)) if int8_cache_name is not None else "(?:)"
|
||||
for name, param in self.model.named_parameters():
|
||||
if param.dtype == torch.int8:
|
||||
if int8_cache == "hbm":
|
||||
if int8_cache_name is None or (
|
||||
int8_cache_name is not None and re.search(int8_pattern, name) is not None
|
||||
):
|
||||
try:
|
||||
self.original_int8[name] = param.data.clone().detach()
|
||||
except RuntimeError as e:
|
||||
logger.error("Failed to cache int8 tensor %s to HBM, change to DRAM, due to %s", name, e)
|
||||
self.original_int8[name] = param.data.cpu()
|
||||
|
||||
elif int8_cache == "dram":
|
||||
if int8_cache_name is None or (
|
||||
int8_cache_name is not None and re.search(int8_pattern, name) is not None
|
||||
):
|
||||
self.original_int8[name] = param.data.cpu()
|
||||
elif int8_cache == "no":
|
||||
pass
|
||||
else:
|
||||
logger.warning(
|
||||
"int8_cache should be selected in [HBM, DRAM], but got %s, change to no cache", int8_cache
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"Server %s:%s starts, device id: %s, model path: %s, tp: %s, pp: %s, int8 params %s are saved to %s",
|
||||
self.addr,
|
||||
self.port,
|
||||
self.device_id,
|
||||
self.model_path,
|
||||
self.tp,
|
||||
self.pp,
|
||||
list(self.original_int8),
|
||||
int8_cache,
|
||||
)
|
||||
|
||||
def __del__(self):
|
||||
"""
|
||||
Destructor method to ensure socket is closed.
|
||||
"""
|
||||
if self.s is not None:
|
||||
with suppress(Exception):
|
||||
self.s.close()
|
||||
|
||||
def start(self):
|
||||
"""
|
||||
Starts the server to handle incoming connections.
|
||||
"""
|
||||
handler_thread = threading.Thread(target=self.elastic_client_handler)
|
||||
handler_thread.daemon = True
|
||||
handler_thread.start()
|
||||
|
||||
def elastic_client_handler(self):
|
||||
"""
|
||||
Handles incoming client connections.
|
||||
"""
|
||||
while True:
|
||||
conn, addr = self.s.accept()
|
||||
logger.info("Accept new connection from %s:%s...", *addr)
|
||||
self.register_handler(conn, addr)
|
||||
|
||||
def register_handler(self, conn, addr, buffer_size=1024):
|
||||
"""
|
||||
Handles the registration of a client.
|
||||
|
||||
Parameters:
|
||||
- conn: The connection socket.
|
||||
- addr: The address of the client.
|
||||
- buffer_size: The size of the buffer for receiving data.
|
||||
"""
|
||||
data_str = conn.recv(buffer_size).decode("utf-8")
|
||||
if not data_str:
|
||||
return
|
||||
try:
|
||||
data = json.loads(data_str)
|
||||
except Exception:
|
||||
logger.error("Failed to load %s as JSON string from %s", data_str, addr)
|
||||
conn.close()
|
||||
return
|
||||
|
||||
def is_valid_data(data):
|
||||
"""
|
||||
Validates the received data.
|
||||
|
||||
Parameters:
|
||||
- data: The data to be validated.
|
||||
|
||||
Returns:
|
||||
- True if the data is valid, otherwise False.
|
||||
"""
|
||||
if not isinstance(data, dict):
|
||||
return False
|
||||
if data.get("label") != "JOIN":
|
||||
return False
|
||||
content = data.get("content")
|
||||
if not isinstance(content, dict):
|
||||
return False
|
||||
required_keys = ["device_id", "model_path", "tp", "pp", "port"]
|
||||
if not all(k in content for k in required_keys):
|
||||
return False
|
||||
port = content["port"]
|
||||
return isinstance(port, int) or (isinstance(port, str) and port.isdigit())
|
||||
|
||||
comm_name = None
|
||||
if is_valid_data(data):
|
||||
device_id = int(data["content"]["device_id"])
|
||||
model_path = data["content"]["model_path"]
|
||||
tp = int(data["content"]["tp"])
|
||||
pp = int(data["content"]["pp"])
|
||||
|
||||
if (
|
||||
int(self.device_id) == device_id
|
||||
and self.model_path == model_path
|
||||
and int(self.tp) == tp
|
||||
and int(self.pp) == pp
|
||||
):
|
||||
comm_name = str(addr[0]) + ":" + str(addr[1])
|
||||
ack = {"label": "JOIN_ACK", "content": {"name": comm_name}}
|
||||
else:
|
||||
server_desc = (int(self.device_id), self.model_path, int(self.tp), int(self.pp))
|
||||
client_desc = (device_id, model_path, tp, pp)
|
||||
msg = f"Received data {client_desc} does not consist with this server {server_desc}"
|
||||
logger.warning(msg)
|
||||
ack = {
|
||||
"label": "JOIN_NACK",
|
||||
"content": msg,
|
||||
}
|
||||
else:
|
||||
logger.warning("Received data does not contain required fields: %s", data)
|
||||
ack = {"label": "JOIN_NACK", "content": f"Received data does not contain required fields: {data}"}
|
||||
|
||||
try:
|
||||
ack_str = json.dumps(ack).encode("utf-8")
|
||||
except Exception as e:
|
||||
logger.error("Failed to convert %s to JSON format, details: %s", ack, e)
|
||||
conn.close()
|
||||
return
|
||||
|
||||
try:
|
||||
conn.send(ack_str)
|
||||
except Exception as e:
|
||||
logger.error("Failed to send %s to %s, details: %s", ack, addr, e)
|
||||
conn.close()
|
||||
return
|
||||
|
||||
if ack["content"] and isinstance(ack["content"], dict) and "name" in ack["content"]:
|
||||
try:
|
||||
p2psend = P2PSend(
|
||||
self.addr,
|
||||
data["content"]["port"],
|
||||
ack["content"]["name"],
|
||||
data["content"].get("group_name", "netloader"),
|
||||
)
|
||||
p2psend.send(self.model, self.original_int8)
|
||||
except Exception as e:
|
||||
logger.error("P2PSend Failed to send model to %s, details: %s", self.addr, e)
|
||||
conn.close()
|
||||
77
vllm_ascend/model_loader/netloader/load.py
Normal file
77
vllm_ascend/model_loader/netloader/load.py
Normal file
@@ -0,0 +1,77 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
import time
|
||||
|
||||
from vllm.logger import logger
|
||||
|
||||
from .executor.elastic_load import P2PLoad
|
||||
from .interaction.elastic import ElasticClient
|
||||
|
||||
|
||||
def elastic_load(
|
||||
model,
|
||||
device_id: int,
|
||||
model_path: str,
|
||||
sources: list,
|
||||
tp: int,
|
||||
pp: int,
|
||||
group_name: str = "netloader",
|
||||
):
|
||||
"""
|
||||
Loads a model using elastic loading across multiple devices.
|
||||
|
||||
Parameters:
|
||||
- model: The model instance to be loaded.
|
||||
- device_id: The ID of the current device (i.e. global rank).
|
||||
- model_path: The path to the model file.
|
||||
- sources: A list of source configurations, each containing device_id and sources.
|
||||
- tp: Tensor parallel size, indicating the number of devices for tensor parallelism.
|
||||
- pp: Pipeline parallel size, indicating the number of devices for pipeline parallelism.
|
||||
- group_name: Name of the HCCL process group.
|
||||
|
||||
Returns:
|
||||
- The loaded model if successful, otherwise None.
|
||||
"""
|
||||
|
||||
# Filter sources for the current device
|
||||
sources_this_device = []
|
||||
for s in sources:
|
||||
if isinstance(s, dict) and "device_id" in s and s["device_id"] == device_id and isinstance(s["sources"], list):
|
||||
sources_this_device += s["sources"]
|
||||
if len(sources_this_device) == 0:
|
||||
return None
|
||||
|
||||
try:
|
||||
# Initialize the interaction layer with the ElasticClient
|
||||
with ElasticClient(sources_this_device, device_id, model_path, tp, pp, group_name) as client_interaction_layer:
|
||||
if client_interaction_layer.s is None or client_interaction_layer.server_addr is None:
|
||||
raise RuntimeError("Failed to initialize ElasticClient: socket or server_addr is None")
|
||||
ack = client_interaction_layer.ack
|
||||
if ack is None:
|
||||
raise RuntimeError("ElasticClient.register did not return ack")
|
||||
|
||||
t0 = time.perf_counter()
|
||||
elastic_loader = P2PLoad(ack[0], client_interaction_layer.server_addr, ack[1], group_name)
|
||||
model_loaded = elastic_loader.load(model=model)
|
||||
if model_loaded is None:
|
||||
logger.error("Failed to load model")
|
||||
return None
|
||||
logger.info("Finish elastic load (duration: %ss)", time.perf_counter() - t0)
|
||||
return model_loaded
|
||||
except Exception as e:
|
||||
logger.info("elastic_load error: %s", e)
|
||||
return None
|
||||
443
vllm_ascend/model_loader/netloader/netloader.py
Normal file
443
vllm_ascend/model_loader/netloader/netloader.py
Normal file
@@ -0,0 +1,443 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
import gc
|
||||
import json
|
||||
import time
|
||||
from copy import deepcopy
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from vllm.config import LoadConfig, ModelConfig, VllmConfig
|
||||
from vllm.logger import logger
|
||||
from vllm.model_executor.model_loader import register_model_loader
|
||||
from vllm.model_executor.model_loader.base_loader import BaseModelLoader
|
||||
from vllm.model_executor.model_loader.default_loader import DefaultModelLoader
|
||||
from vllm.model_executor.model_loader.utils import initialize_model, process_weights_after_loading
|
||||
from vllm.utils.torch_utils import set_default_torch_dtype
|
||||
|
||||
from .interaction.elastic import ElasticServer
|
||||
from .load import elastic_load
|
||||
from .utils import find_free_port, is_valid_path_prefix
|
||||
|
||||
DRAFT_PORT_OFFSET = 10000
|
||||
|
||||
try:
|
||||
# Older vLLM versions may not expose the current-config accessor.
|
||||
from vllm.config import get_current_vllm_config
|
||||
except ImportError:
|
||||
get_current_vllm_config = None
|
||||
|
||||
|
||||
@register_model_loader("netloader")
|
||||
class ModelNetLoaderElastic(BaseModelLoader):
|
||||
"""
|
||||
A model loader that uses elastic loading for loading weights.
|
||||
"""
|
||||
|
||||
source: list[dict] | None
|
||||
model_path: str | None
|
||||
listen_port: int | None
|
||||
int8_cache: str
|
||||
int8_cache_name: list[str] | None
|
||||
output_prefix: str | None
|
||||
|
||||
def __init__(self, load_config: LoadConfig):
|
||||
"""
|
||||
Initializes the ModelNetLoaderElastic with configuration.
|
||||
|
||||
Parameters:
|
||||
- load_config: Configuration for loading the model.
|
||||
"""
|
||||
super().__init__(load_config)
|
||||
|
||||
config = None
|
||||
|
||||
# Try to read config file at first
|
||||
extra = load_config.model_loader_extra_config
|
||||
|
||||
if extra is not None and not isinstance(extra, dict):
|
||||
err_msg = "NetLoader requires --model-loader-extra-config to be a JSON object."
|
||||
logger.error(err_msg)
|
||||
raise RuntimeError(err_msg)
|
||||
|
||||
if extra and "CONFIG_FILE" in extra:
|
||||
try:
|
||||
logger.info("Reading configs in file %s ...", load_config.model_loader_extra_config["CONFIG_FILE"])
|
||||
with open(extra["CONFIG_FILE"]) as f:
|
||||
config = json.load(f)
|
||||
except FileNotFoundError:
|
||||
logger.error("CONFIG_FILE not found")
|
||||
except json.JSONDecodeError:
|
||||
logger.error("CONFIG_FILE is not a valid JSON file")
|
||||
except Exception as e:
|
||||
logger.error("Unexpected error while reading CONFIG_FILE: %s", e)
|
||||
|
||||
if config is None and extra:
|
||||
logger.info("Reading configs in model_loader_extra_config ...")
|
||||
config = extra
|
||||
config = config or {}
|
||||
|
||||
for key, attr, checker, caster, default in [
|
||||
("SOURCE", "source", lambda v: isinstance(v, list), lambda v: v, None),
|
||||
("MODEL", "model_path", lambda v: isinstance(v, str), lambda v: v, None),
|
||||
(
|
||||
"LISTEN_PORT",
|
||||
"listen_port",
|
||||
lambda v: isinstance(v, int) or (isinstance(v, str) and v.isdigit()),
|
||||
lambda v: int(v),
|
||||
None,
|
||||
),
|
||||
(
|
||||
"INT8_CACHE",
|
||||
"int8_cache",
|
||||
lambda v: isinstance(v, str) and v.lower() in ["hbm", "dram", "no"],
|
||||
lambda v: v.lower(),
|
||||
"no",
|
||||
),
|
||||
("INT8_CACHE_NAME", "int8_cache_name", lambda v: isinstance(v, list), lambda v: v, None),
|
||||
(
|
||||
"OUTPUT_PREFIX",
|
||||
"output_prefix",
|
||||
lambda v: isinstance(v, str) and is_valid_path_prefix(v),
|
||||
lambda v: v,
|
||||
None,
|
||||
),
|
||||
]:
|
||||
v = config.get(key, default)
|
||||
if not checker(v):
|
||||
v = default
|
||||
else:
|
||||
v = caster(v)
|
||||
setattr(self, attr, v)
|
||||
|
||||
logger.info(
|
||||
"Initializing elastic Netloader with config: "
|
||||
"MODEL=%s, LISTEN_PORT=%s,"
|
||||
"SOURCE=%s, INT8_CACHE=%s, INT8_CACHE_NAME=%s,"
|
||||
"OUTPUT_PREFIX=%s)",
|
||||
self.model_path,
|
||||
self.listen_port,
|
||||
self.source,
|
||||
self.int8_cache,
|
||||
self.int8_cache_name,
|
||||
self.output_prefix,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _is_draft_model(model_config: ModelConfig) -> bool:
|
||||
"""Check whether the model_config corresponds to a draft model for speculative decoding."""
|
||||
return getattr(model_config, "runner_type", None) == "draft"
|
||||
|
||||
@staticmethod
|
||||
def _sync_target_netloader_before_draft(vllm_config: VllmConfig) -> None:
|
||||
if getattr(vllm_config, "speculative_config", None) is None:
|
||||
return
|
||||
if not torch.distributed.is_available() or not torch.distributed.is_initialized():
|
||||
return
|
||||
|
||||
logger.info("Waiting for all target netloader ranks before loading draft model")
|
||||
barrier_start = time.perf_counter()
|
||||
torch.distributed.barrier()
|
||||
logger.info(
|
||||
"Target netloader barrier before draft model time: %s",
|
||||
time.perf_counter() - barrier_start,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_static_forward_context(vllm_config: VllmConfig):
|
||||
compilation_config = getattr(vllm_config, "compilation_config", None)
|
||||
static_forward_context = getattr(compilation_config, "static_forward_context", None)
|
||||
if static_forward_context is None or not hasattr(static_forward_context, "clear"):
|
||||
return None
|
||||
return static_forward_context
|
||||
|
||||
@staticmethod
|
||||
def _clear_static_forward_context(vllm_config: VllmConfig) -> None:
|
||||
"""Clear static layer registrations before rebuilding the model on fallback."""
|
||||
candidates = [("vllm_config", vllm_config)]
|
||||
if get_current_vllm_config is not None:
|
||||
try:
|
||||
candidates.append(("current_vllm_config", get_current_vllm_config()))
|
||||
except Exception as e:
|
||||
logger.debug("Failed to get current vLLM config while clearing static context: %s", e)
|
||||
|
||||
cleared_contexts = []
|
||||
seen_context_ids = set()
|
||||
for source, config in candidates:
|
||||
static_forward_context = ModelNetLoaderElastic._get_static_forward_context(config)
|
||||
if static_forward_context is None:
|
||||
continue
|
||||
|
||||
context_id = id(static_forward_context)
|
||||
if context_id in seen_context_ids:
|
||||
continue
|
||||
seen_context_ids.add(context_id)
|
||||
|
||||
try:
|
||||
context_size = str(len(static_forward_context))
|
||||
except TypeError:
|
||||
context_size = "unknown"
|
||||
static_forward_context.clear()
|
||||
cleared_contexts.append(f"{source}:{context_size}")
|
||||
|
||||
if cleared_contexts:
|
||||
logger.info("Cleared static_forward_context before fallback: %s", cleared_contexts)
|
||||
|
||||
def load_model(self, vllm_config: VllmConfig, model_config: ModelConfig, prefix: str = "") -> nn.Module:
|
||||
"""
|
||||
Loads the model using the specified configuration.
|
||||
|
||||
Parameters:
|
||||
- vllm_config: Configuration for the VLLM.
|
||||
- model_config: Configuration for the model.
|
||||
- prefix: Module prefix for pipeline parallelism (e.g., "model.layers.0.").
|
||||
|
||||
Returns:
|
||||
- The loaded model.
|
||||
"""
|
||||
|
||||
device_config = vllm_config.device_config
|
||||
parallel_config = vllm_config.parallel_config
|
||||
|
||||
need_process_weights_after_loading = False
|
||||
|
||||
if self.model_path is None:
|
||||
self.model_path = model_config.model
|
||||
logger.info("model_path is set to %s", self.model_path)
|
||||
|
||||
device_id = torch.distributed.get_rank()
|
||||
is_draft = self._is_draft_model(model_config)
|
||||
|
||||
if is_draft:
|
||||
logger.info("Loading draft model via netloader, model_path: %s", model_config.model)
|
||||
else:
|
||||
logger.info("Loading target model via netloader, model_path: %s", model_config.model)
|
||||
|
||||
if (
|
||||
self.source is None
|
||||
or not isinstance(self.source, list)
|
||||
or device_id
|
||||
not in [
|
||||
one_device["device_id"]
|
||||
for one_device in self.source
|
||||
if isinstance(one_device, dict) and "device_id" in one_device
|
||||
]
|
||||
):
|
||||
logger.warning("Did not get valid source info, use DefaultModelLoader")
|
||||
model, need_process_weights_after_loading = self.revert_to_default(
|
||||
model_config, vllm_config, device_config, prefix
|
||||
)
|
||||
|
||||
else:
|
||||
target_device = torch.device(device_config.device)
|
||||
|
||||
_quant_config = getattr(vllm_config, "quant_config", None)
|
||||
_quant_config = deepcopy(_quant_config) if _quant_config is not None else None
|
||||
model_config_backup = deepcopy(model_config)
|
||||
|
||||
with set_default_torch_dtype(model_config.dtype):
|
||||
with target_device:
|
||||
model = initialize_model(vllm_config=vllm_config, model_config=model_config, prefix=prefix)
|
||||
|
||||
start_elastic_load = time.perf_counter()
|
||||
|
||||
sources = self.source
|
||||
if is_draft:
|
||||
sources = [
|
||||
{
|
||||
"device_id": s["device_id"],
|
||||
"sources": [
|
||||
f"{parts[0]}:{int(parts[1]) + DRAFT_PORT_OFFSET}"
|
||||
for addr in s.get("sources", [])
|
||||
if isinstance(addr, str)
|
||||
and len(parts := addr.rsplit(":", 1)) == 2
|
||||
and parts[1].isdigit()
|
||||
],
|
||||
}
|
||||
for s in self.source
|
||||
if isinstance(s, dict) and "device_id" in s
|
||||
]
|
||||
|
||||
model = elastic_load(
|
||||
model=model,
|
||||
device_id=device_id,
|
||||
model_path=model_config.model,
|
||||
sources=sources,
|
||||
tp=parallel_config.tensor_parallel_size,
|
||||
pp=parallel_config.pipeline_parallel_size,
|
||||
group_name="netloader_draft" if is_draft else "netloader",
|
||||
)
|
||||
end_elastic_load = time.perf_counter()
|
||||
logger.info("Elastic load time: %s, rank: %s", end_elastic_load - start_elastic_load, device_id)
|
||||
need_process_weights_after_loading = True
|
||||
|
||||
if model is None:
|
||||
logger.warning("Netloader elastic loading fails, use load format DefaultModelLoader")
|
||||
|
||||
if hasattr(vllm_config, "quant_config"):
|
||||
vllm_config.quant_config = _quant_config
|
||||
model_config = model_config_backup
|
||||
|
||||
del model
|
||||
gc.collect()
|
||||
if device_config.device_type == "npu":
|
||||
logger.info("Empty NPU cache")
|
||||
torch.npu.empty_cache()
|
||||
elif device_config.device_type == "cuda":
|
||||
logger.info("Empty CUDA cache")
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Clear registrations from the failed initialize_model
|
||||
self._clear_static_forward_context(vllm_config)
|
||||
|
||||
model, need_process_weights_after_loading = self.revert_to_default(
|
||||
model_config, vllm_config, device_config, prefix
|
||||
)
|
||||
|
||||
start_elastic_server = time.perf_counter()
|
||||
# start elastic server
|
||||
if model is not None and (
|
||||
(self.listen_port and self.listen_port in range(1024, 65535)) or (self.listen_port is None)
|
||||
):
|
||||
from vllm.utils.network_utils import get_ip
|
||||
|
||||
driver_ip = get_ip()
|
||||
|
||||
if driver_ip == "0.0.0.0":
|
||||
logger.error("Driver IP is not set, skip to start Netloader server")
|
||||
else:
|
||||
if self.listen_port is None:
|
||||
listen_port = find_free_port()
|
||||
else:
|
||||
listen_port = self.listen_port + device_id
|
||||
if is_draft:
|
||||
listen_port += DRAFT_PORT_OFFSET
|
||||
self.listen_port = listen_port
|
||||
|
||||
group_name = "netloader_draft" if is_draft else "netloader"
|
||||
|
||||
logger.info(
|
||||
"Start elastic Netloader server, rank: %s, listen port: %s:%s, group: %s",
|
||||
device_id,
|
||||
driver_ip,
|
||||
listen_port,
|
||||
group_name,
|
||||
)
|
||||
|
||||
if self.output_prefix is not None and not is_draft:
|
||||
try:
|
||||
with open(self.output_prefix + str(device_id) + ".txt", "w") as file:
|
||||
file.write(f"{driver_ip}:{listen_port}")
|
||||
logger.info(
|
||||
"Successfully wrote server address to file: %s", self.output_prefix + str(device_id)
|
||||
)
|
||||
except FileNotFoundError:
|
||||
logger.error("File path %s does not exist.", self.output_prefix + str(device_id))
|
||||
except PermissionError:
|
||||
logger.error("No permission to write to file %s.", self.output_prefix + str(device_id))
|
||||
except OSError as e:
|
||||
logger.error(
|
||||
"I/O error occurred while writing to file %s: %s", self.output_prefix + str(device_id), e
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error("Unknown error: %s", e)
|
||||
|
||||
try:
|
||||
server_int8_cache = "hbm" if is_draft and self.int8_cache != "no" else self.int8_cache
|
||||
elastic_server = ElasticServer(
|
||||
driver_ip,
|
||||
listen_port,
|
||||
model,
|
||||
device_id,
|
||||
model_config.model,
|
||||
parallel_config.tensor_parallel_size,
|
||||
parallel_config.pipeline_parallel_size,
|
||||
server_int8_cache,
|
||||
self.int8_cache_name,
|
||||
group_name=group_name,
|
||||
)
|
||||
elastic_server.start()
|
||||
if is_draft:
|
||||
self._draft_elastic_server = elastic_server
|
||||
else:
|
||||
self._target_elastic_server = elastic_server
|
||||
except Exception as e:
|
||||
logger.error("Failed to start Netloader server for rank: %s, details: %s", device_id, e)
|
||||
else:
|
||||
logger.info("Skip to start Netloader server")
|
||||
|
||||
end_elastic_server = time.perf_counter()
|
||||
logger.info("Elastic server start time: %s, rank: %s", end_elastic_server - start_elastic_server, device_id)
|
||||
|
||||
if need_process_weights_after_loading:
|
||||
process_weights_after_loading(model, model_config, torch.device(device_config.device))
|
||||
|
||||
if not is_draft:
|
||||
self._sync_target_netloader_before_draft(vllm_config)
|
||||
|
||||
if model is None:
|
||||
logger.error("NetLoader elastic loads model fails")
|
||||
raise RuntimeError("NetLoader elastic loads model fails")
|
||||
|
||||
return model.eval()
|
||||
|
||||
def revert_to_default(self, model_config, vllm_config, device_config, prefix: str = "") -> tuple[nn.Module, bool]:
|
||||
"""
|
||||
Reverts to the default model loading logic when elastic loading fails or is not applicable.
|
||||
|
||||
This method resets the loader's extra config and load format to defaults,
|
||||
then delegates model loading to a DefaultModelLoader.
|
||||
If quantization is enabled, it will load the model and then run the
|
||||
processing of weights (i.e. applying quantization adjustments) before returning.
|
||||
|
||||
Parameters:
|
||||
- model_config: Configuration describing model architecture, quantization, etc.
|
||||
- vllm_config: Configuration for vLLM (device, parallelism, dtype, etc).
|
||||
- device_config: Configuration for the target device (device type, device id, etc).
|
||||
- prefix: Module prefix for pipeline parallelism.
|
||||
|
||||
Returns:
|
||||
- A tuple (model, need_process_weights_after_loading):
|
||||
* model: The loaded `nn.Module` under default loading logic.
|
||||
* need_process_weights_after_loading: A boolean flag indicating whether
|
||||
weights post-processing (e.g. quantization adjustments) still needs to be applied.
|
||||
"""
|
||||
load_config = deepcopy(self.load_config)
|
||||
load_config.model_loader_extra_config = {}
|
||||
load_config.load_format = "auto"
|
||||
default_model_loader = DefaultModelLoader(load_config)
|
||||
|
||||
if model_config.quantization is None:
|
||||
model = default_model_loader.load_model(vllm_config=vllm_config, model_config=model_config, prefix=prefix)
|
||||
need_process_weights_after_loading = False
|
||||
else:
|
||||
logger.warning("Quantization is set, netloader use DefaultModelLoader with process_weights_after_loading ")
|
||||
need_process_weights_after_loading = True
|
||||
target_device = torch.device(device_config.device)
|
||||
with set_default_torch_dtype(model_config.dtype):
|
||||
with target_device:
|
||||
model = initialize_model(vllm_config=vllm_config, model_config=model_config, prefix=prefix)
|
||||
default_model_loader.load_weights(model, model_config)
|
||||
model = model.eval()
|
||||
|
||||
return model, need_process_weights_after_loading
|
||||
|
||||
def download_model(self, model_config: ModelConfig) -> None:
|
||||
pass
|
||||
|
||||
def load_weights(self, model: nn.Module, model_config: ModelConfig) -> None:
|
||||
pass
|
||||
63
vllm_ascend/model_loader/netloader/utils.py
Normal file
63
vllm_ascend/model_loader/netloader/utils.py
Normal file
@@ -0,0 +1,63 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
import os
|
||||
import socket
|
||||
|
||||
import regex as re
|
||||
from vllm.logger import logger
|
||||
|
||||
|
||||
def find_free_port():
|
||||
"""
|
||||
Finds a free port on the local machine.
|
||||
|
||||
Returns:
|
||||
- A free port number.
|
||||
"""
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||
s.bind(("", 0))
|
||||
return s.getsockname()[1]
|
||||
|
||||
|
||||
def is_valid_path_prefix(path_prefix):
|
||||
"""
|
||||
Checks if the provided path prefix is valid.
|
||||
|
||||
Parameters:
|
||||
- path_prefix: The path prefix to validate.
|
||||
|
||||
Returns:
|
||||
- True if the path prefix is valid, otherwise False.
|
||||
"""
|
||||
if not path_prefix:
|
||||
return False
|
||||
|
||||
if re.search(r'[<>:"|?*]', path_prefix):
|
||||
logger.warning("The path prefix %s contains illegal characters.", path_prefix)
|
||||
return False
|
||||
|
||||
if path_prefix.startswith("/") or path_prefix.startswith("\\"):
|
||||
if not os.path.exists(os.path.dirname(path_prefix)):
|
||||
logger.warning("The directory for the path prefix %s does not exist.", os.path.dirname(path_prefix))
|
||||
return False
|
||||
else:
|
||||
if not os.path.exists(os.path.dirname(os.path.abspath(path_prefix))):
|
||||
logger.warning(
|
||||
"The directory for the path prefix %s does not exist.", os.path.dirname(os.path.abspath(path_prefix))
|
||||
)
|
||||
return False
|
||||
return True
|
||||
20
vllm_ascend/model_loader/rfork/__init__.py
Normal file
20
vllm_ascend/model_loader/rfork/__init__.py
Normal file
@@ -0,0 +1,20 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
|
||||
def register_rforkloader() -> None:
|
||||
"""Register the RFork model loader plugin."""
|
||||
from .rfork_loader import RForkModelLoader # noqa: F401
|
||||
333
vllm_ascend/model_loader/rfork/rfork_loader.py
Normal file
333
vllm_ascend/model_loader/rfork/rfork_loader.py
Normal file
@@ -0,0 +1,333 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
import gc
|
||||
import os
|
||||
import time
|
||||
from copy import copy
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.nn import Module
|
||||
from vllm.config import ModelConfig, VllmConfig
|
||||
from vllm.config.load import LoadConfig
|
||||
from vllm.distributed import get_tensor_model_parallel_rank
|
||||
from vllm.distributed.parallel_state import get_ep_group, get_pp_group
|
||||
from vllm.logger import logger
|
||||
from vllm.model_executor.model_loader import register_model_loader
|
||||
from vllm.model_executor.model_loader.base_loader import BaseModelLoader
|
||||
from vllm.model_executor.model_loader.utils import (
|
||||
initialize_model,
|
||||
process_weights_after_loading,
|
||||
)
|
||||
from vllm.utils.torch_utils import set_default_torch_dtype
|
||||
|
||||
from vllm_ascend.model_loader.rfork.rfork_worker import RForkWorker
|
||||
|
||||
|
||||
def _is_mtp_hf_config(hf_config: object | None) -> bool:
|
||||
if hf_config is None:
|
||||
return False
|
||||
|
||||
model_type = getattr(hf_config, "model_type", None)
|
||||
if isinstance(model_type, str) and model_type.lower().endswith("_mtp"):
|
||||
return True
|
||||
|
||||
architectures = getattr(hf_config, "architectures", None)
|
||||
if isinstance(architectures, str):
|
||||
architectures = [architectures]
|
||||
if not isinstance(architectures, (list, tuple)):
|
||||
return False
|
||||
|
||||
return any(isinstance(architecture, str) and architecture.endswith("MTPModel") for architecture in architectures)
|
||||
|
||||
|
||||
def _is_draft_model_config(model_config: object | None) -> bool:
|
||||
if model_config is None:
|
||||
return False
|
||||
if getattr(model_config, "runner_type", None) == "draft":
|
||||
return True
|
||||
|
||||
return any(
|
||||
_is_mtp_hf_config(getattr(model_config, hf_config_attr, None))
|
||||
for hf_config_attr in ("hf_config", "hf_text_config")
|
||||
)
|
||||
|
||||
|
||||
def _is_draft_model(vllm_config: VllmConfig, model_config: ModelConfig | None = None) -> bool:
|
||||
return (
|
||||
_is_draft_model_config(model_config)
|
||||
or _is_draft_model_config(getattr(vllm_config, "model_config", None))
|
||||
or _is_draft_model_config(getattr(vllm_config, "scheduler_config", None))
|
||||
)
|
||||
|
||||
|
||||
def _get_rfork_worker_attr(vllm_config: VllmConfig, model_config: ModelConfig) -> str:
|
||||
return "rfork_draft_worker" if _is_draft_model(vllm_config, model_config) else "rfork_worker"
|
||||
|
||||
|
||||
def _get_ep_rank(vllm_config: VllmConfig) -> int | None:
|
||||
parallel_config = vllm_config.parallel_config
|
||||
if not parallel_config.enable_expert_parallel or getattr(parallel_config, "is_moe_model", None) is False:
|
||||
return None
|
||||
|
||||
try:
|
||||
return get_ep_group().rank_in_group
|
||||
except AssertionError as e:
|
||||
raise RuntimeError("Expert parallelism is enabled, but the EP group is not initialized.") from e
|
||||
|
||||
|
||||
def _get_pp_rank(vllm_config: VllmConfig) -> int | None:
|
||||
if getattr(vllm_config.parallel_config, "pipeline_parallel_size", 1) <= 1:
|
||||
return None
|
||||
|
||||
try:
|
||||
return get_pp_group().rank_in_group
|
||||
except AssertionError as e:
|
||||
raise RuntimeError("Pipeline parallelism is enabled, but the PP group is not initialized.") from e
|
||||
|
||||
|
||||
def _make_fallback_load_config(load_config: LoadConfig) -> LoadConfig:
|
||||
fallback_load_config = copy(load_config)
|
||||
fallback_load_config.load_format = "auto"
|
||||
fallback_load_config.model_loader_extra_config = {}
|
||||
return fallback_load_config
|
||||
|
||||
|
||||
def _is_layer_sharding_enabled(vllm_config: VllmConfig) -> bool:
|
||||
additional_config = getattr(vllm_config, "additional_config", None) or {}
|
||||
return bool(additional_config.get("layer_sharding"))
|
||||
|
||||
|
||||
def _is_dynamic_eplb_enabled(vllm_config: VllmConfig) -> bool:
|
||||
parallel_config = getattr(vllm_config, "parallel_config", None)
|
||||
if bool(getattr(parallel_config, "enable_eplb", False)):
|
||||
return True
|
||||
|
||||
additional_config = getattr(vllm_config, "additional_config", None) or {}
|
||||
eplb_config = additional_config.get("eplb_config", {})
|
||||
if not isinstance(eplb_config, dict):
|
||||
return False
|
||||
return bool(eplb_config.get("dynamic_eplb") or eplb_config.get("expert_map_record_path"))
|
||||
|
||||
|
||||
@register_model_loader("rfork")
|
||||
class RForkModelLoader(BaseModelLoader):
|
||||
def __init__(self, load_config: LoadConfig):
|
||||
super().__init__(load_config)
|
||||
config = load_config.model_loader_extra_config
|
||||
if config is None:
|
||||
config = {}
|
||||
elif not isinstance(config, dict):
|
||||
err_msg = "RFork requires --model-loader-extra-config to be a JSON object."
|
||||
logger.error(err_msg)
|
||||
raise RuntimeError(err_msg)
|
||||
|
||||
def _get_extra_config(key: str, default: str = "") -> str:
|
||||
value = config.get(key)
|
||||
if value is None or not isinstance(value, str):
|
||||
value = os.environ.get(key.upper())
|
||||
return value if isinstance(value, str) and value else default
|
||||
|
||||
def _get_extra_config_float(key: str, default: float) -> float:
|
||||
value = config.get(key)
|
||||
if value is None or isinstance(value, bool) or not isinstance(value, (int, float, str)):
|
||||
value = os.environ.get(key.upper())
|
||||
parsed_value = default
|
||||
if isinstance(value, (int, float)):
|
||||
parsed_value = float(value)
|
||||
elif isinstance(value, str) and value:
|
||||
try:
|
||||
parsed_value = float(value)
|
||||
except ValueError:
|
||||
return default
|
||||
|
||||
if parsed_value <= 0:
|
||||
return default
|
||||
|
||||
return parsed_value
|
||||
|
||||
self.model_url = _get_extra_config("model_url", "")
|
||||
self.model_deploy_strategy_name = _get_extra_config("model_deploy_strategy_name", "")
|
||||
self.scheduler_url = _get_extra_config("rfork_scheduler_url", "")
|
||||
self.seed_timeout_sec = _get_extra_config_float("rfork_seed_timeout_sec", 5.0)
|
||||
self.seed_key_separator = _get_extra_config("rfork_seed_key_separator", "$")
|
||||
|
||||
logger.info(
|
||||
"Initializing rfork with config: "
|
||||
"MODEL_URL=%s, MODEL_DEPLOY_STRATEGY_NAME=%s, "
|
||||
"SCHEDULER_URL=%s, SEED_TIMEOUT_SEC=%s, "
|
||||
"SEED_KEY_SEPARATOR=%s",
|
||||
self.model_url,
|
||||
self.model_deploy_strategy_name,
|
||||
self.scheduler_url,
|
||||
self.seed_timeout_sec,
|
||||
self.seed_key_separator,
|
||||
)
|
||||
|
||||
def download_model(self, model_config: ModelConfig) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
def load_weights(self, model: nn.Module, model_config: ModelConfig) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
def _ensure_rfork_worker(self, vllm_config: VllmConfig, model_config: ModelConfig) -> RForkWorker:
|
||||
worker_attr = _get_rfork_worker_attr(vllm_config, model_config)
|
||||
rfork_worker = getattr(self.load_config, worker_attr, None)
|
||||
if rfork_worker is None:
|
||||
kv_transfer_config = vllm_config.kv_transfer_config
|
||||
disaggregation_mode = "kv_both" if kv_transfer_config is None else str(kv_transfer_config.kv_role)
|
||||
is_draft_model = _is_draft_model(vllm_config, model_config)
|
||||
device_id = torch.distributed.get_rank()
|
||||
pp_rank = _get_pp_rank(vllm_config)
|
||||
ep_rank = _get_ep_rank(vllm_config)
|
||||
rfork_worker = RForkWorker(
|
||||
disaggregation_mode=disaggregation_mode,
|
||||
node_rank=vllm_config.parallel_config.node_rank,
|
||||
tp_rank=get_tensor_model_parallel_rank(),
|
||||
device_id=device_id,
|
||||
scheduler_url=self.scheduler_url,
|
||||
model_url=self.model_url,
|
||||
model_deploy_strategy_name=self.model_deploy_strategy_name,
|
||||
seed_timeout_sec=self.seed_timeout_sec,
|
||||
seed_key_separator=self.seed_key_separator,
|
||||
is_draft_model=is_draft_model,
|
||||
pp_rank=pp_rank,
|
||||
ep_rank=ep_rank,
|
||||
)
|
||||
setattr(self.load_config, worker_attr, rfork_worker)
|
||||
logger.info(
|
||||
"RFork worker initialized, load_format=rfork, is_draft_model=%s, worker_attr=%s",
|
||||
is_draft_model,
|
||||
worker_attr,
|
||||
)
|
||||
return rfork_worker
|
||||
|
||||
def _requires_processed_layout_transfer(self, model_config: ModelConfig) -> bool:
|
||||
return getattr(model_config, "quantization", None) is not None
|
||||
|
||||
def load_model(
|
||||
self,
|
||||
vllm_config: VllmConfig,
|
||||
model_config: ModelConfig,
|
||||
prefix: str = "",
|
||||
) -> Module | None:
|
||||
device_config = vllm_config.device_config
|
||||
load_config = self.load_config
|
||||
load_device = device_config.device if load_config.device is None else load_config.device
|
||||
target_device = torch.device(load_device)
|
||||
|
||||
with set_default_torch_dtype(model_config.dtype):
|
||||
need_del = False
|
||||
bypass_reason = None
|
||||
if _is_layer_sharding_enabled(vllm_config):
|
||||
bypass_reason = "additional_config.layer_sharding"
|
||||
elif _is_dynamic_eplb_enabled(vllm_config):
|
||||
bypass_reason = "dynamic EPLB"
|
||||
|
||||
if bypass_reason is not None:
|
||||
logger.warning(
|
||||
"RFork transfer is disabled when %s is enabled; using the default model loader.",
|
||||
bypass_reason,
|
||||
)
|
||||
fallback_load_config = _make_fallback_load_config(self.load_config)
|
||||
|
||||
from vllm.model_executor.model_loader import get_model
|
||||
|
||||
try:
|
||||
return get_model(
|
||||
vllm_config=vllm_config,
|
||||
model_config=model_config,
|
||||
load_config=fallback_load_config,
|
||||
prefix=prefix,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("RFork disabled for %s, but default loader failed.", bypass_reason)
|
||||
raise
|
||||
|
||||
rfork_worker = self._ensure_rfork_worker(vllm_config, model_config)
|
||||
processed_layout_transfer = self._requires_processed_layout_transfer(model_config)
|
||||
try:
|
||||
if not rfork_worker.is_seed_available():
|
||||
raise RuntimeError("seed is not available.")
|
||||
|
||||
with target_device:
|
||||
model = initialize_model(
|
||||
vllm_config=vllm_config,
|
||||
model_config=model_config,
|
||||
prefix=prefix,
|
||||
)
|
||||
need_del = True
|
||||
|
||||
if processed_layout_transfer:
|
||||
logger.info("RFork uses post-load tensor layout transfer for quantized model.")
|
||||
process_weights_after_loading(model, model_config, target_device)
|
||||
|
||||
weight_load_start_time = time.perf_counter()
|
||||
if not rfork_worker.pre_transfer(model):
|
||||
raise RuntimeError("pre_transfer failed.")
|
||||
if not rfork_worker.transfer(model):
|
||||
raise RuntimeError("transfer failed.")
|
||||
if not rfork_worker.post_transfer():
|
||||
raise RuntimeError("post_transfer failed.")
|
||||
logger.info(
|
||||
"Loading model weights took %.2f seconds",
|
||||
time.perf_counter() - weight_load_start_time,
|
||||
)
|
||||
|
||||
rfork_worker.start_seed_service(model)
|
||||
if not processed_layout_transfer:
|
||||
process_weights_after_loading(model, model_config, target_device)
|
||||
|
||||
return model.eval()
|
||||
except Exception as e:
|
||||
logger.warning("RFork transfer failed: %s, clean up and fall back to default loader", e)
|
||||
|
||||
rfork_worker.post_transfer()
|
||||
rfork_worker.reset_transfer_state()
|
||||
|
||||
if need_del:
|
||||
del model
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
for _ in range(3):
|
||||
gc.collect()
|
||||
torch.npu.empty_cache()
|
||||
|
||||
fallback_load_config = _make_fallback_load_config(self.load_config)
|
||||
|
||||
from vllm.model_executor.model_loader import get_model
|
||||
|
||||
try:
|
||||
model = get_model(
|
||||
vllm_config=vllm_config,
|
||||
model_config=model_config,
|
||||
load_config=fallback_load_config,
|
||||
prefix=prefix,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("RFork fallback default loader failed.")
|
||||
raise
|
||||
|
||||
try:
|
||||
rfork_worker.reset_transfer_state()
|
||||
rfork_worker.start_seed_service(model)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"Fallback model loaded, but start_seed_service failed: %s",
|
||||
e,
|
||||
)
|
||||
return model
|
||||
144
vllm_ascend/model_loader/rfork/rfork_worker.py
Normal file
144
vllm_ascend/model_loader/rfork/rfork_worker.py
Normal file
@@ -0,0 +1,144 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
import threading
|
||||
|
||||
from vllm.logger import logger
|
||||
|
||||
from vllm_ascend.model_loader.rfork.seed_protocol import RForkSeedProtocol
|
||||
from vllm_ascend.model_loader.rfork.seed_server import start_rfork_server
|
||||
from vllm_ascend.model_loader.rfork.transfer_backend import (
|
||||
RForkTransferBackend,
|
||||
)
|
||||
|
||||
|
||||
class RForkWorker:
|
||||
def __init__(
|
||||
self,
|
||||
disaggregation_mode: str,
|
||||
node_rank: int,
|
||||
tp_rank: int,
|
||||
device_id: int,
|
||||
scheduler_url: str,
|
||||
model_url: str,
|
||||
model_deploy_strategy_name: str,
|
||||
seed_timeout_sec: float = 30.0,
|
||||
seed_key_separator: str = "$",
|
||||
is_draft_model: bool = False,
|
||||
pp_rank: int | None = None,
|
||||
ep_rank: int | None = None,
|
||||
):
|
||||
self.device_id = device_id
|
||||
self.rfork_seed = None
|
||||
self.transfer_backend = RForkTransferBackend()
|
||||
self.ready_to_start_seed_service = False
|
||||
self.seed_service_started = False
|
||||
self.seed_timeout_sec = seed_timeout_sec
|
||||
self.seed_protocol = RForkSeedProtocol(
|
||||
disaggregation_mode=disaggregation_mode,
|
||||
node_rank=node_rank,
|
||||
tp_rank=tp_rank,
|
||||
scheduler_url=scheduler_url,
|
||||
model_url=model_url,
|
||||
model_deploy_strategy_name=model_deploy_strategy_name,
|
||||
seed_key_separator=seed_key_separator,
|
||||
is_draft_worker=is_draft_model,
|
||||
pp_rank=pp_rank,
|
||||
ep_rank=ep_rank,
|
||||
)
|
||||
|
||||
def is_seed_available(self) -> bool:
|
||||
self.rfork_seed = self.seed_protocol.get_seed()
|
||||
return self.rfork_seed is not None
|
||||
|
||||
def pre_transfer(self, model) -> bool:
|
||||
try:
|
||||
assert self.transfer_backend.is_initialized(), "transfer_backend is not initialized, cannot pre_transfer."
|
||||
result = self.transfer_backend.register_memory_region(model)
|
||||
self.ready_to_start_seed_service = result
|
||||
return result
|
||||
except AssertionError as e:
|
||||
logger.exception("Pre-transfer failed for device_id=%s: %s", self.device_id, e)
|
||||
return False
|
||||
|
||||
def reset_transfer_state(self) -> None:
|
||||
try:
|
||||
self.transfer_backend.unregister_memory_region()
|
||||
except Exception as e:
|
||||
logger.warning("Failed to unregister rfork memory region: %s", e)
|
||||
self.ready_to_start_seed_service = False
|
||||
|
||||
def transfer(self, model) -> bool:
|
||||
try:
|
||||
assert self.transfer_backend.is_initialized(), "transfer_backend is not initialized, cannot transfer."
|
||||
assert self.rfork_seed is not None, "rfork seed is None, cannot transfer."
|
||||
return self.transfer_backend.recv_from_source(
|
||||
model=model,
|
||||
seed_instance_ip=self.rfork_seed["seed_ip"],
|
||||
seed_instance_service_port=self.rfork_seed["seed_port"],
|
||||
local_seed_key=self.seed_protocol.get_local_seed_key(),
|
||||
)
|
||||
except AssertionError as e:
|
||||
logger.exception(
|
||||
"Transfer failed for device_id=%s: %s",
|
||||
self.device_id,
|
||||
e,
|
||||
)
|
||||
return False
|
||||
|
||||
def post_transfer(self):
|
||||
if self.rfork_seed is None:
|
||||
logger.info("rfork seed is None, no need to release.")
|
||||
return True
|
||||
self.seed_protocol.release_seed(self.rfork_seed)
|
||||
self.rfork_seed = None
|
||||
return True
|
||||
|
||||
def start_seed_service(self, model):
|
||||
if self.seed_service_started:
|
||||
logger.info("Seed service already started, skipping.")
|
||||
return
|
||||
|
||||
if not self.ready_to_start_seed_service:
|
||||
if not self.pre_transfer(model):
|
||||
logger.warning(
|
||||
"start_seed_service aborted for device_id=%s: pre_transfer failed",
|
||||
self.device_id,
|
||||
)
|
||||
return
|
||||
|
||||
port = start_rfork_server(
|
||||
self.seed_protocol.get_local_seed_key(),
|
||||
(
|
||||
self.transfer_backend.rfork_transfer_engine_session_id,
|
||||
self.transfer_backend.rfork_transfer_engine_weights_info_dict,
|
||||
self.transfer_backend.rfork_transfer_engine_weights_shape_dict,
|
||||
),
|
||||
health_timeout_sec=self.seed_timeout_sec,
|
||||
)
|
||||
if port <= 0:
|
||||
logger.warning("start_seed_service failed for device_id=%s", self.device_id)
|
||||
return
|
||||
|
||||
self.rfork_heartbeat_thread = threading.Thread(
|
||||
target=self.seed_protocol.report_seed,
|
||||
args=(port,),
|
||||
daemon=True,
|
||||
name="RForkHeartbeat",
|
||||
)
|
||||
self.rfork_heartbeat_thread.start()
|
||||
logger.info("Seed service started for device_id=%s, port=%s", self.device_id, port)
|
||||
self.seed_service_started = True
|
||||
242
vllm_ascend/model_loader/rfork/seed_protocol.py
Normal file
242
vllm_ascend/model_loader/rfork/seed_protocol.py
Normal file
@@ -0,0 +1,242 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
import time
|
||||
from urllib.error import HTTPError
|
||||
|
||||
import requests
|
||||
from vllm.logger import logger
|
||||
from vllm.utils.network_utils import get_ip
|
||||
|
||||
REQUEST_TIMEOUT_SEC = 10.0
|
||||
HEARTBEAT_LOG_EVERY_N = 4
|
||||
|
||||
|
||||
def get_local_seed_key(
|
||||
disaggregation_mode: str,
|
||||
node_rank: int,
|
||||
tp_rank: int,
|
||||
model_url: str,
|
||||
model_deploy_strategy_name: str,
|
||||
seed_key_separator: str = "$",
|
||||
is_draft_worker: bool = False,
|
||||
pp_rank: int | None = None,
|
||||
ep_rank: int | None = None,
|
||||
) -> str:
|
||||
if not model_url or not model_deploy_strategy_name:
|
||||
err_msg = (
|
||||
f"RFork seed key is not set: model_url={model_url!r}, "
|
||||
f"model_deploy_strategy_name={model_deploy_strategy_name!r}. "
|
||||
"Ensure model_loader_extra_config contains "
|
||||
"`model_url` and `model_deploy_strategy_name`, or set "
|
||||
"MODEL_URL and MODEL_DEPLOY_STRATEGY_NAME."
|
||||
)
|
||||
logger.error(err_msg)
|
||||
raise RuntimeError(err_msg)
|
||||
|
||||
seed_key = f"{model_url}{seed_key_separator}{model_deploy_strategy_name}"
|
||||
key_parts = [disaggregation_mode, str(node_rank)]
|
||||
if pp_rank is not None:
|
||||
key_parts.append(f"pp{pp_rank}")
|
||||
key_parts.append(str(tp_rank))
|
||||
if ep_rank is not None:
|
||||
key_parts.append(f"ep{ep_rank}")
|
||||
if is_draft_worker:
|
||||
key_parts.append("draft")
|
||||
return f"{seed_key}{seed_key_separator}{seed_key_separator.join(key_parts)}"
|
||||
|
||||
|
||||
class RForkSeedProtocol:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
disaggregation_mode: str,
|
||||
node_rank: int,
|
||||
tp_rank: int,
|
||||
scheduler_url: str,
|
||||
model_url: str,
|
||||
model_deploy_strategy_name: str,
|
||||
seed_key_separator: str = "$",
|
||||
is_draft_worker: bool = False,
|
||||
pp_rank: int | None = None,
|
||||
ep_rank: int | None = None,
|
||||
):
|
||||
self.disaggregation_mode = disaggregation_mode
|
||||
self.node_rank = node_rank
|
||||
self.tp_rank = tp_rank
|
||||
self.pp_rank = pp_rank
|
||||
self.ep_rank = ep_rank
|
||||
self.scheduler_url = scheduler_url
|
||||
self.model_url = model_url
|
||||
self.model_deploy_strategy_name = model_deploy_strategy_name
|
||||
self.seed_key_separator = seed_key_separator
|
||||
self.is_draft_worker = is_draft_worker
|
||||
|
||||
self._local_seed_key = get_local_seed_key(
|
||||
disaggregation_mode=self.disaggregation_mode,
|
||||
node_rank=self.node_rank,
|
||||
tp_rank=self.tp_rank,
|
||||
model_url=self.model_url,
|
||||
model_deploy_strategy_name=self.model_deploy_strategy_name,
|
||||
seed_key_separator=self.seed_key_separator,
|
||||
is_draft_worker=self.is_draft_worker,
|
||||
pp_rank=self.pp_rank,
|
||||
ep_rank=self.ep_rank,
|
||||
)
|
||||
|
||||
def get_local_seed_key(self) -> str:
|
||||
return self._local_seed_key
|
||||
|
||||
@staticmethod
|
||||
def _request_timeout_sec() -> float:
|
||||
return REQUEST_TIMEOUT_SEC
|
||||
|
||||
def _ensure_scheduler_url_set(self) -> None:
|
||||
if not self.scheduler_url:
|
||||
raise RuntimeError(
|
||||
"rfork_scheduler_url is not set. Set it through model_loader_extra_config or RFORK_SCHEDULER_URL."
|
||||
)
|
||||
|
||||
def get_seed(self):
|
||||
try:
|
||||
self._ensure_scheduler_url_set()
|
||||
seed_key = self.get_local_seed_key()
|
||||
response = requests.get(
|
||||
f"{self.scheduler_url}/get_seed",
|
||||
headers={
|
||||
"SEED_KEY": seed_key,
|
||||
},
|
||||
timeout=self._request_timeout_sec(),
|
||||
)
|
||||
if response.status_code != 200:
|
||||
raise RuntimeError(
|
||||
f"Failed to get seed from the planner, {response.status_code}, seed_key={seed_key!r}"
|
||||
)
|
||||
|
||||
seed_ip = response.headers.get("SEED_IP")
|
||||
seed_port = response.headers.get("SEED_PORT")
|
||||
user_id = response.headers.get("USER_ID")
|
||||
seed_rank = response.headers.get("SEED_RANK")
|
||||
logger.debug(
|
||||
"seed_ip: %s, seed_port: %s, user_id: %s, seed_rank: %s",
|
||||
seed_ip,
|
||||
seed_port,
|
||||
user_id,
|
||||
seed_rank,
|
||||
)
|
||||
return {
|
||||
"seed_ip": seed_ip,
|
||||
"seed_port": seed_port,
|
||||
"user_id": user_id,
|
||||
"seed_rank": seed_rank,
|
||||
}
|
||||
|
||||
except RuntimeError as e:
|
||||
logger.warning("get_seed from scheduler RuntimeError: %s", e)
|
||||
return None
|
||||
except HTTPError as e:
|
||||
logger.exception("get_seed from scheduler HTTPError: %s", e)
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.exception("get_seed from scheduler Exception: %s", e)
|
||||
return None
|
||||
|
||||
def release_seed(self, seed) -> bool:
|
||||
try:
|
||||
self._ensure_scheduler_url_set()
|
||||
user_id = seed["user_id"]
|
||||
seed_ip = seed["seed_ip"]
|
||||
seed_port = str(seed["seed_port"])
|
||||
seed_rank = str(seed["seed_rank"])
|
||||
|
||||
response = requests.post(
|
||||
f"{self.scheduler_url}/put_seed",
|
||||
headers={
|
||||
"SEED_IP": seed_ip,
|
||||
"SEED_PORT": seed_port,
|
||||
"USER_ID": user_id,
|
||||
"SEED_RANK": seed_rank,
|
||||
},
|
||||
timeout=self._request_timeout_sec(),
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
raise RuntimeError(f"Failed to release seed to the planner, {response.status_code}")
|
||||
return True
|
||||
except RuntimeError as e:
|
||||
logger.exception("release_seed to planner RuntimeError: %s", e)
|
||||
return False
|
||||
except HTTPError as e:
|
||||
logger.exception("release_seed to planner HTTPError: %s", e)
|
||||
return False
|
||||
except Exception as e:
|
||||
logger.exception("release_seed to planner Exception: %s", e)
|
||||
return False
|
||||
|
||||
def report_seed(self, port: int, sleep_interval: int = 30):
|
||||
heartbeat_idx = 0
|
||||
log_every_n = HEARTBEAT_LOG_EVERY_N
|
||||
try:
|
||||
self._ensure_scheduler_url_set()
|
||||
seed_ip = get_ip()
|
||||
seed_key = self.get_local_seed_key()
|
||||
logger.debug("[rfork_heartbeat] reporting seed key: %s", seed_key)
|
||||
except Exception as e:
|
||||
logger.exception("report_seed setup Exception: %s", e)
|
||||
return
|
||||
|
||||
while True:
|
||||
heartbeat_idx += 1
|
||||
result = False
|
||||
try:
|
||||
response = requests.post(
|
||||
f"{self.scheduler_url}/add_seed",
|
||||
headers={
|
||||
"SEED_KEY": seed_key,
|
||||
"SEED_IP": seed_ip,
|
||||
"SEED_PORT": str(port),
|
||||
"SEED_RANK": str(self.tp_rank),
|
||||
"SEED_REFCNT": str(0),
|
||||
},
|
||||
timeout=self._request_timeout_sec(),
|
||||
)
|
||||
if response.status_code == 200:
|
||||
result = True
|
||||
except HTTPError as e:
|
||||
logger.warning("report_seed to planner HTTPError: %s", e)
|
||||
except Exception as e:
|
||||
logger.warning("report_seed to planner Exception: %s", e)
|
||||
|
||||
# Keep heartbeat frequency unchanged, but reduce log noise.
|
||||
# Always print failures immediately; keep success in debug logs.
|
||||
if result:
|
||||
if heartbeat_idx % log_every_n == 0:
|
||||
logger.debug(
|
||||
"[rfork_heartbeat] report seed to planner result: %s (%d/%d), seed_key=%s",
|
||||
result,
|
||||
heartbeat_idx % log_every_n if heartbeat_idx % log_every_n != 0 else log_every_n,
|
||||
log_every_n,
|
||||
seed_key,
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"[rfork_heartbeat] report seed to planner result: %s (%d/%d), seed_key=%s",
|
||||
result,
|
||||
heartbeat_idx % log_every_n if heartbeat_idx % log_every_n != 0 else log_every_n,
|
||||
log_every_n,
|
||||
seed_key,
|
||||
)
|
||||
time.sleep(sleep_interval)
|
||||
141
vllm_ascend/model_loader/rfork/seed_server.py
Normal file
141
vllm_ascend/model_loader/rfork/seed_server.py
Normal file
@@ -0,0 +1,141 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
import queue
|
||||
import socket
|
||||
import threading
|
||||
import time
|
||||
from http import HTTPStatus
|
||||
|
||||
import requests
|
||||
import uvicorn
|
||||
from fastapi import FastAPI
|
||||
from fastapi.responses import Response
|
||||
from vllm.logger import logger
|
||||
|
||||
|
||||
def start_fastapi_server(
|
||||
port_queue: queue.Queue[int],
|
||||
local_seed_key,
|
||||
info,
|
||||
):
|
||||
logger.debug("[RFork Seed] Preparing socket with dynamic port...")
|
||||
|
||||
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
||||
sock.bind(("0.0.0.0", 0))
|
||||
_, port = sock.getsockname()
|
||||
logger.debug("[RFork Seed] Assigned dynamic port: %s", port)
|
||||
|
||||
app = FastAPI()
|
||||
rfork_transfer_engine_info = info
|
||||
rfork_transfer_engine_shape_info = None
|
||||
if isinstance(info, (list, tuple)) and len(info) == 3:
|
||||
rfork_transfer_engine_info = (info[0], info[1])
|
||||
rfork_transfer_engine_shape_info = info[2]
|
||||
|
||||
@app.get("/get_rfork_transfer_engine_info")
|
||||
def get_rfork_transfer_engine_info(seed_key: str):
|
||||
if seed_key == local_seed_key:
|
||||
return {"rfork_transfer_engine_info": rfork_transfer_engine_info}
|
||||
return {"rfork_transfer_engine_info": None}
|
||||
|
||||
@app.get("/get_rfork_transfer_engine_shape_info")
|
||||
def get_rfork_transfer_engine_shape_info(seed_key: str):
|
||||
if seed_key == local_seed_key:
|
||||
return {"rfork_transfer_engine_shape_info": rfork_transfer_engine_shape_info}
|
||||
return {"rfork_transfer_engine_shape_info": None}
|
||||
|
||||
@app.get("/rfork_fetch_seed")
|
||||
def rfork_fetch_seed():
|
||||
return {"status": "ok"}
|
||||
|
||||
@app.get("/health_check_with_key")
|
||||
def health_check_with_key(seed_key: str):
|
||||
if seed_key == local_seed_key:
|
||||
return Response(status_code=HTTPStatus.OK)
|
||||
return Response(status_code=HTTPStatus.BAD_REQUEST)
|
||||
|
||||
config = uvicorn.Config(app, host=None, port=None, log_level="warning")
|
||||
server = uvicorn.Server(config)
|
||||
|
||||
try:
|
||||
port_queue.put(port)
|
||||
except Exception as e:
|
||||
logger.error("[RFork Seed] Failed to send port via queue: %s", e)
|
||||
sock.close()
|
||||
return
|
||||
|
||||
logger.debug("[RFork Seed] FastAPI server starting on port %s...", port)
|
||||
server.run(sockets=[sock])
|
||||
sock.close()
|
||||
|
||||
|
||||
def start_rfork_server(local_seed_key, rfork_transfer_engine_info, health_timeout_sec: float = 30.0) -> int:
|
||||
port_queue: queue.Queue[int] = queue.Queue()
|
||||
process = threading.Thread(
|
||||
target=start_fastapi_server,
|
||||
args=(port_queue, local_seed_key, rfork_transfer_engine_info),
|
||||
daemon=True,
|
||||
)
|
||||
process.start()
|
||||
|
||||
try:
|
||||
port = port_queue.get(timeout=15)
|
||||
if port == -1:
|
||||
raise RuntimeError("Child process failed to start server")
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"[RFork Seed] start server error for seed_key=%s: %s",
|
||||
local_seed_key,
|
||||
e,
|
||||
)
|
||||
return -1
|
||||
|
||||
deadline = time.time() + health_timeout_sec
|
||||
healthy = False
|
||||
retry_count = 0
|
||||
last_error = None
|
||||
while time.time() < deadline:
|
||||
time.sleep(0.01)
|
||||
url = f"http://127.0.0.1:{port}/health_check_with_key"
|
||||
try:
|
||||
response = requests.get(
|
||||
url,
|
||||
params={"seed_key": local_seed_key},
|
||||
timeout=10,
|
||||
)
|
||||
if response.status_code == 200:
|
||||
healthy = True
|
||||
break
|
||||
last_error = f"unexpected status code {response.status_code} from health check"
|
||||
except Exception as e:
|
||||
last_error = str(e)
|
||||
retry_count += 1
|
||||
if healthy:
|
||||
if retry_count > 1:
|
||||
logger.info(
|
||||
"[RFork Seed] health check passed after %d retries for port %s",
|
||||
retry_count - 1,
|
||||
port,
|
||||
)
|
||||
return port
|
||||
logger.error(
|
||||
"[RFork Seed] health check timed out after %.1fs for port %s, last error: %s",
|
||||
health_timeout_sec,
|
||||
port,
|
||||
last_error,
|
||||
)
|
||||
return -1
|
||||
585
vllm_ascend/model_loader/rfork/transfer_backend.py
Normal file
585
vllm_ascend/model_loader/rfork/transfer_backend.py
Normal file
@@ -0,0 +1,585 @@
|
||||
#
|
||||
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
#
|
||||
# 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.
|
||||
#
|
||||
|
||||
import time
|
||||
from bisect import bisect_left
|
||||
from typing import Any
|
||||
|
||||
import requests
|
||||
import torch
|
||||
from torch import nn
|
||||
from vllm.logger import logger
|
||||
from vllm.utils.network_utils import get_ip, get_open_port, join_host_port
|
||||
|
||||
MAX_TRANSFER_CHUNK_BYTES = 1024**3
|
||||
MAX_TRANSFER_CHUNK_WEIGHTS = 512
|
||||
|
||||
|
||||
def _normalize_weight_shape(shape: Any) -> tuple[int, ...] | None:
|
||||
if shape is None:
|
||||
return None
|
||||
if not isinstance(shape, (list, tuple)):
|
||||
return None
|
||||
if not all(isinstance(dim, int) and dim >= 0 for dim in shape):
|
||||
return None
|
||||
return tuple(shape)
|
||||
|
||||
|
||||
def _parse_weight_info(weight_info: Any):
|
||||
if not isinstance(weight_info, (list, tuple)) or len(weight_info) not in (3, 4):
|
||||
return None
|
||||
|
||||
seed_ptr, seed_len, seed_size = weight_info[:3]
|
||||
if not all(isinstance(value, int) for value in (seed_ptr, seed_len, seed_size)):
|
||||
return None
|
||||
|
||||
seed_shape = None
|
||||
if len(weight_info) == 4:
|
||||
seed_shape = _normalize_weight_shape(weight_info[3])
|
||||
if seed_shape is None:
|
||||
return None
|
||||
|
||||
return seed_ptr, seed_len, seed_size, seed_shape
|
||||
|
||||
|
||||
def _reshape_tensor_to_seed_shape(
|
||||
name: str,
|
||||
tensor: torch.Tensor,
|
||||
seed_shape: tuple[int, ...] | None,
|
||||
reshape_events: list[tuple[str, tuple[int, ...], tuple[int, ...]]] | None = None,
|
||||
) -> bool:
|
||||
if seed_shape is None or tuple(tensor.shape) == seed_shape:
|
||||
return True
|
||||
|
||||
if tensor.numel() != _numel_from_shape(seed_shape):
|
||||
logger.error(
|
||||
"Weight shape mismatch for %s, local shape %s cannot view as seed shape %s",
|
||||
name,
|
||||
tuple(tensor.shape),
|
||||
seed_shape,
|
||||
)
|
||||
return False
|
||||
|
||||
local_shape = tuple(tensor.shape)
|
||||
try:
|
||||
tensor.data = tensor.data.view(seed_shape)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Failed to reshape RFork tensor %s from %s to seed shape %s: %s",
|
||||
name,
|
||||
local_shape,
|
||||
seed_shape,
|
||||
e,
|
||||
)
|
||||
return False
|
||||
|
||||
if reshape_events is not None:
|
||||
reshape_events.append((name, local_shape, seed_shape))
|
||||
return True
|
||||
|
||||
|
||||
def _update_registered_weight_shape(
|
||||
weight_shape_dict: dict[str, tuple[int, ...]] | None,
|
||||
name: str,
|
||||
tensor: torch.Tensor,
|
||||
) -> None:
|
||||
if isinstance(weight_shape_dict, dict):
|
||||
weight_shape_dict[name] = tuple(tensor.shape)
|
||||
|
||||
|
||||
def _numel_from_shape(shape: tuple[int, ...]) -> int:
|
||||
numel = 1
|
||||
for dim in shape:
|
||||
numel *= dim
|
||||
return numel
|
||||
|
||||
|
||||
def _is_transferable_tensor(tensor: torch.Tensor) -> bool:
|
||||
return not tensor.is_meta and tensor.numel() > 0 and _is_tensor_on_transfer_device(tensor)
|
||||
|
||||
|
||||
def _is_tensor_on_transfer_device(tensor: torch.Tensor) -> bool:
|
||||
return tensor.device.type == "npu"
|
||||
|
||||
|
||||
def _iter_tensors_in_value(prefix: str, value: Any, visited_object_ids: set[int], scan_objects: bool = False):
|
||||
if isinstance(value, torch.Tensor):
|
||||
yield prefix, value
|
||||
return
|
||||
|
||||
if isinstance(value, (nn.Module, str, bytes)) or callable(value):
|
||||
return
|
||||
|
||||
if isinstance(value, (list, tuple)):
|
||||
for index, item in enumerate(value):
|
||||
yield from _iter_tensors_in_value(f"{prefix}.{index}", item, visited_object_ids, scan_objects)
|
||||
return
|
||||
|
||||
if isinstance(value, dict):
|
||||
for key, item in value.items():
|
||||
yield from _iter_tensors_in_value(f"{prefix}.{key}", item, visited_object_ids, scan_objects)
|
||||
return
|
||||
|
||||
if not scan_objects or not hasattr(value, "__dict__"):
|
||||
return
|
||||
|
||||
value_id = id(value)
|
||||
if value_id in visited_object_ids:
|
||||
return
|
||||
visited_object_ids.add(value_id)
|
||||
for attr_name, attr_value in vars(value).items():
|
||||
if attr_name.startswith("_"):
|
||||
continue
|
||||
yield from _iter_tensors_in_value(f"{prefix}.{attr_name}", attr_value, visited_object_ids, scan_objects)
|
||||
|
||||
|
||||
def _try_collect_transferable_tensor(
|
||||
name: str,
|
||||
tensor: torch.Tensor,
|
||||
seen_data_ptrs: set[int],
|
||||
collected_tensors: list[tuple[str, torch.Tensor]],
|
||||
) -> tuple[bool, bool]:
|
||||
if not _is_transferable_tensor(tensor):
|
||||
return False, False
|
||||
|
||||
data_ptr = tensor.data_ptr()
|
||||
if data_ptr in seen_data_ptrs:
|
||||
return False, True
|
||||
|
||||
seen_data_ptrs.add(data_ptr)
|
||||
collected_tensors.append((name, tensor))
|
||||
return True, False
|
||||
|
||||
|
||||
def _collect_transferable_tensors(model: nn.Module) -> list[tuple[str, torch.Tensor]]:
|
||||
seen_data_ptrs: set[int] = set()
|
||||
collected_tensors: list[tuple[str, torch.Tensor]] = []
|
||||
|
||||
for name, tensor in model.named_parameters():
|
||||
_try_collect_transferable_tensor(
|
||||
name,
|
||||
tensor,
|
||||
seen_data_ptrs,
|
||||
collected_tensors,
|
||||
)
|
||||
|
||||
for name, tensor in model.named_buffers():
|
||||
_try_collect_transferable_tensor(
|
||||
name,
|
||||
tensor,
|
||||
seen_data_ptrs,
|
||||
collected_tensors,
|
||||
)
|
||||
|
||||
# Some Ascend post-load paths replace checkpoint parameters with runtime
|
||||
# tensors stored as plain module attributes, e.g. MLA/SFA W_UV and W_UK_T.
|
||||
for module_prefix, module in model.named_modules():
|
||||
for attr_name, attr_value in vars(module).items():
|
||||
if attr_name.startswith("_") or isinstance(attr_value, nn.Module):
|
||||
continue
|
||||
|
||||
scan_objects = attr_name == "impl"
|
||||
for tensor_name, tensor in _iter_tensors_in_value(attr_name, attr_value, set(), scan_objects):
|
||||
full_name = f"{module_prefix}.{tensor_name}" if module_prefix else tensor_name
|
||||
_try_collect_transferable_tensor(
|
||||
full_name,
|
||||
tensor,
|
||||
seen_data_ptrs,
|
||||
collected_tensors,
|
||||
)
|
||||
return collected_tensors
|
||||
|
||||
|
||||
def _iter_transferable_tensors(model: nn.Module):
|
||||
yield from _collect_transferable_tensors(model)
|
||||
|
||||
|
||||
def _block_contains_weight_ptr(address: int, size: int, sorted_weight_ptrs: list[int]) -> bool:
|
||||
index = bisect_left(sorted_weight_ptrs, address)
|
||||
return index < len(sorted_weight_ptrs) and sorted_weight_ptrs[index] < address + size
|
||||
|
||||
|
||||
def _iter_transfer_chunks(
|
||||
weight_names: list[str],
|
||||
seed_ptr_list: list[int],
|
||||
client_ptr_list: list[int],
|
||||
client_len_list: list[int],
|
||||
):
|
||||
chunk_start = 0
|
||||
chunk_bytes = 0
|
||||
chunk_weights = 0
|
||||
|
||||
for index, length in enumerate(client_len_list):
|
||||
should_flush = chunk_weights > 0 and (
|
||||
chunk_bytes + length > MAX_TRANSFER_CHUNK_BYTES or chunk_weights >= MAX_TRANSFER_CHUNK_WEIGHTS
|
||||
)
|
||||
if should_flush:
|
||||
yield (
|
||||
weight_names[chunk_start:index],
|
||||
seed_ptr_list[chunk_start:index],
|
||||
client_ptr_list[chunk_start:index],
|
||||
client_len_list[chunk_start:index],
|
||||
)
|
||||
chunk_start = index
|
||||
chunk_bytes = 0
|
||||
chunk_weights = 0
|
||||
|
||||
chunk_bytes += length
|
||||
chunk_weights += 1
|
||||
|
||||
if chunk_weights > 0:
|
||||
yield (
|
||||
weight_names[chunk_start:],
|
||||
seed_ptr_list[chunk_start:],
|
||||
client_ptr_list[chunk_start:],
|
||||
client_len_list[chunk_start:],
|
||||
)
|
||||
|
||||
|
||||
class RForkTransferBackend:
|
||||
def __init__(self):
|
||||
self.rfork_transfer_engine: Any | None = None
|
||||
self.rfork_transfer_engine_session_id = None
|
||||
self.rfork_transfer_engine_weights_info_dict = None
|
||||
self.rfork_transfer_engine_weights_shape_dict = None
|
||||
self.registered_weight_blocks = []
|
||||
self._registered_transferable_tensors: list[tuple[str, torch.Tensor]] | None = None
|
||||
self._is_initialized = False
|
||||
self.init_transfer_engine()
|
||||
|
||||
def init_transfer_engine(self):
|
||||
try:
|
||||
from yr.datasystem import TransferEngine # type: ignore[import-not-found]
|
||||
except ImportError as e:
|
||||
err_msg = (
|
||||
"Failed to import TransferEngine from yr.datasystem. "
|
||||
"Please install @yuanrong-datasystem/transfer_engine."
|
||||
)
|
||||
logger.error(err_msg)
|
||||
raise ImportError(err_msg) from e
|
||||
|
||||
transfer_engine = TransferEngine()
|
||||
local_hostname = join_host_port(get_ip(), get_open_port())
|
||||
ret = transfer_engine.initialize(local_hostname, "ascend", f"npu:{torch.npu.current_device()}")
|
||||
if ret.is_error():
|
||||
err_msg = (
|
||||
f"TransferEngine initialization failed: "
|
||||
f"initialize({local_hostname}, 'ascend', "
|
||||
f"'npu:{int(torch.npu.current_device())}') -> {ret.to_string()}"
|
||||
)
|
||||
logger.error(err_msg)
|
||||
raise RuntimeError(err_msg)
|
||||
|
||||
self.rfork_transfer_engine = transfer_engine
|
||||
self.rfork_transfer_engine_session_id = local_hostname
|
||||
self._is_initialized = True
|
||||
|
||||
def is_initialized(self) -> bool:
|
||||
return self._is_initialized
|
||||
|
||||
def _get_transfer_engine(self) -> Any:
|
||||
if self.rfork_transfer_engine is None:
|
||||
raise RuntimeError("TransferEngine is not initialized.")
|
||||
return self.rfork_transfer_engine
|
||||
|
||||
def register_memory_region(self, model):
|
||||
transfer_engine = self._get_transfer_engine()
|
||||
start_reg_mr_time = time.perf_counter()
|
||||
self._registered_transferable_tensors = None
|
||||
|
||||
weight_mr_dict = {}
|
||||
weight_shape_dict = {}
|
||||
weight_addr_set = set()
|
||||
transferable_tensors = list(_iter_transferable_tensors(model))
|
||||
for name, weight in transferable_tensors:
|
||||
weight_mr_dict[name] = (
|
||||
weight.data_ptr(),
|
||||
weight.numel(),
|
||||
weight.element_size(),
|
||||
)
|
||||
weight_shape_dict[name] = tuple(weight.shape)
|
||||
weight_addr_set.add(weight.data_ptr())
|
||||
|
||||
sorted_weight_ptrs = sorted(weight_addr_set)
|
||||
|
||||
memory_snapshot = torch.npu.memory.memory_snapshot()
|
||||
|
||||
weight_blocks_for_reg_mr = []
|
||||
for segment in memory_snapshot:
|
||||
current_weight_block = None
|
||||
for block in segment.get("blocks", []):
|
||||
address = block.get("address", -1)
|
||||
size = block.get("size", -1)
|
||||
state = block.get("state", "")
|
||||
if address < 0 or size < 0 or state == "":
|
||||
continue
|
||||
if state == "active_allocated" and _block_contains_weight_ptr(address, size, sorted_weight_ptrs):
|
||||
if current_weight_block is None:
|
||||
current_weight_block = (address, size)
|
||||
elif current_weight_block[0] + current_weight_block[1] == address:
|
||||
current_weight_block = (
|
||||
current_weight_block[0],
|
||||
current_weight_block[1] + size,
|
||||
)
|
||||
else:
|
||||
weight_blocks_for_reg_mr.append(current_weight_block)
|
||||
current_weight_block = (address, size)
|
||||
if current_weight_block is not None:
|
||||
weight_blocks_for_reg_mr.append(current_weight_block)
|
||||
|
||||
addresses, sizes = zip(*weight_blocks_for_reg_mr) if weight_blocks_for_reg_mr else ((), ())
|
||||
ret = transfer_engine.batch_register_memory(addresses, sizes)
|
||||
if ret.is_error():
|
||||
self._registered_transferable_tensors = None
|
||||
logger.error(
|
||||
"batch_register_memory failed for %d blocks, ret: %s",
|
||||
len(weight_blocks_for_reg_mr),
|
||||
ret.to_string(),
|
||||
)
|
||||
return False
|
||||
|
||||
self.rfork_transfer_engine_weights_info_dict = weight_mr_dict
|
||||
self.rfork_transfer_engine_weights_shape_dict = weight_shape_dict
|
||||
self.registered_weight_blocks = weight_blocks_for_reg_mr
|
||||
self._registered_transferable_tensors = transferable_tensors
|
||||
logger.info(
|
||||
"register_memory_region time: %.4fs, weights: %d",
|
||||
time.perf_counter() - start_reg_mr_time,
|
||||
len(weight_mr_dict),
|
||||
)
|
||||
return True
|
||||
|
||||
def unregister_memory_region(self) -> bool:
|
||||
transfer_engine = self._get_transfer_engine()
|
||||
start_unreg_mr_time = time.perf_counter()
|
||||
if not self.registered_weight_blocks:
|
||||
self.rfork_transfer_engine_weights_info_dict = None
|
||||
self.rfork_transfer_engine_weights_shape_dict = None
|
||||
self._registered_transferable_tensors = None
|
||||
logger.debug("unregister_memory_region skipped because no blocks are registered.")
|
||||
return True
|
||||
|
||||
ret = transfer_engine.batch_unregister_memory([address for address, _ in self.registered_weight_blocks])
|
||||
if ret.is_error():
|
||||
logger.error(
|
||||
"batch_unregister_memory failed for %d blocks, ret: %s",
|
||||
len(self.registered_weight_blocks),
|
||||
ret.to_string(),
|
||||
)
|
||||
return False
|
||||
self.rfork_transfer_engine_weights_info_dict = None
|
||||
self.rfork_transfer_engine_weights_shape_dict = None
|
||||
self.registered_weight_blocks = []
|
||||
self._registered_transferable_tensors = None
|
||||
logger.info(
|
||||
"unregister_memory_region time: %.4fs",
|
||||
time.perf_counter() - start_unreg_mr_time,
|
||||
)
|
||||
return True
|
||||
|
||||
def recv_from_source(
|
||||
self,
|
||||
model,
|
||||
seed_instance_ip,
|
||||
seed_instance_service_port,
|
||||
local_seed_key,
|
||||
):
|
||||
transfer_engine = self._get_transfer_engine()
|
||||
seed_url = f"http://{seed_instance_ip}:{seed_instance_service_port}"
|
||||
seed_session_id, seed_weight_info, seed_weight_shapes = get_remote_instance_transfer_engine_info(
|
||||
seed_url,
|
||||
local_seed_key,
|
||||
)
|
||||
if seed_session_id is None or seed_weight_info is None:
|
||||
self._registered_transferable_tensors = None
|
||||
logger.error("Cannot get transfer engine session or weight info.")
|
||||
return False
|
||||
|
||||
transferable_tensors = getattr(self, "_registered_transferable_tensors", None)
|
||||
if transferable_tensors is None:
|
||||
transferable_tensors = list(_iter_transferable_tensors(model))
|
||||
|
||||
seed_ptr_list = []
|
||||
client_ptr_list = []
|
||||
client_len_list = []
|
||||
weight_names = []
|
||||
reshape_events: list[tuple[str, tuple[int, ...], tuple[int, ...]]] = []
|
||||
try:
|
||||
for name, tensor in transferable_tensors:
|
||||
weight_info = seed_weight_info.get(name, None)
|
||||
if weight_info is None:
|
||||
logger.error("Cannot find weight info for %s.", name)
|
||||
return False
|
||||
|
||||
parsed_weight_info = _parse_weight_info(weight_info)
|
||||
if parsed_weight_info is None:
|
||||
logger.error("Invalid weight info for %s: %s", name, weight_info)
|
||||
return False
|
||||
|
||||
seed_ptr, seed_len, seed_size, seed_shape = parsed_weight_info
|
||||
if seed_shape is None and isinstance(seed_weight_shapes, dict):
|
||||
seed_shape = _normalize_weight_shape(seed_weight_shapes.get(name))
|
||||
if seed_len != tensor.numel() or seed_size != tensor.element_size():
|
||||
logger.error(
|
||||
"Weight info mismatch for %s, expected (%s, %s), got (%s, %s)",
|
||||
name,
|
||||
seed_len,
|
||||
seed_size,
|
||||
tensor.numel(),
|
||||
tensor.element_size(),
|
||||
)
|
||||
return False
|
||||
|
||||
if not _reshape_tensor_to_seed_shape(name, tensor, seed_shape, reshape_events):
|
||||
return False
|
||||
_update_registered_weight_shape(
|
||||
self.rfork_transfer_engine_weights_shape_dict,
|
||||
name,
|
||||
tensor,
|
||||
)
|
||||
|
||||
seed_ptr_list.append(seed_ptr)
|
||||
client_ptr_list.append(tensor.data_ptr())
|
||||
client_len_list.append(tensor.numel() * tensor.element_size())
|
||||
weight_names.append(name)
|
||||
finally:
|
||||
self._registered_transferable_tensors = None
|
||||
transferable_tensors = None
|
||||
|
||||
if reshape_events:
|
||||
sample_events = ", ".join(
|
||||
f"{name}: {local_shape}->{seed_shape}" for name, local_shape, seed_shape in reshape_events[:3]
|
||||
)
|
||||
if len(reshape_events) > 3:
|
||||
sample_events += ", ..."
|
||||
logger.debug(
|
||||
"RFork reshaped %d tensors to match seed shapes: %s",
|
||||
len(reshape_events),
|
||||
sample_events,
|
||||
)
|
||||
|
||||
transfer_chunks = list(
|
||||
_iter_transfer_chunks(
|
||||
weight_names,
|
||||
seed_ptr_list,
|
||||
client_ptr_list,
|
||||
client_len_list,
|
||||
)
|
||||
)
|
||||
total_transfer_bytes = sum(client_len_list)
|
||||
|
||||
transfer_start_time = time.perf_counter()
|
||||
logger.info(
|
||||
"transfer weights starts, weights: %d, chunks: %d, total bytes: %.2f GiB",
|
||||
len(client_len_list),
|
||||
len(transfer_chunks),
|
||||
total_transfer_bytes / (1024**3),
|
||||
)
|
||||
for index, (chunk_names, chunk_seed_ptrs, chunk_client_ptrs, chunk_lengths) in enumerate(transfer_chunks, 1):
|
||||
chunk_start_time = time.perf_counter()
|
||||
logger.debug(
|
||||
"transfer weights chunk %d/%d starts, weights: %d, bytes: %.2f GiB, first: %s, last: %s",
|
||||
index,
|
||||
len(transfer_chunks),
|
||||
len(chunk_lengths),
|
||||
sum(chunk_lengths) / (1024**3),
|
||||
chunk_names[0],
|
||||
chunk_names[-1],
|
||||
)
|
||||
ret = transfer_engine.batch_transfer_sync_read(
|
||||
seed_session_id,
|
||||
chunk_client_ptrs,
|
||||
chunk_seed_ptrs,
|
||||
chunk_lengths,
|
||||
)
|
||||
if ret.is_error():
|
||||
logger.error(
|
||||
"Failed to transfer weights chunk %d/%d, first: %s, last: %s, ret=%s",
|
||||
index,
|
||||
len(transfer_chunks),
|
||||
chunk_names[0],
|
||||
chunk_names[-1],
|
||||
ret.to_string(),
|
||||
)
|
||||
return False
|
||||
logger.debug(
|
||||
"transfer weights chunk %d/%d done, time: %.4fs",
|
||||
index,
|
||||
len(transfer_chunks),
|
||||
time.perf_counter() - chunk_start_time,
|
||||
)
|
||||
transfer_time = time.perf_counter() - transfer_start_time
|
||||
logger.info("transfer weights time: %.4fs", transfer_time)
|
||||
return True
|
||||
|
||||
|
||||
def get_remote_instance_transfer_engine_info(seed_url: str, local_seed_key: str):
|
||||
try:
|
||||
response = requests.get(
|
||||
f"{seed_url}/get_rfork_transfer_engine_info",
|
||||
params={"seed_key": local_seed_key},
|
||||
)
|
||||
if response.status_code != 200:
|
||||
logger.error(
|
||||
"GET %s/get_rfork_transfer_engine_info failed: %s",
|
||||
seed_url,
|
||||
response.status_code,
|
||||
)
|
||||
return None, None, None
|
||||
|
||||
data = response.json()
|
||||
info = data.get("rfork_transfer_engine_info", None)
|
||||
if info is not None and isinstance(info, list) and len(info) == 2:
|
||||
shape_info = get_remote_instance_weight_shape_info(seed_url, local_seed_key)
|
||||
return info[0], info[1], shape_info
|
||||
|
||||
logger.error(
|
||||
"Failed to get rfork_transfer_engine_info in response from %s.",
|
||||
seed_url,
|
||||
)
|
||||
return None, None, None
|
||||
except Exception as e:
|
||||
logger.error("Exception getting transfer engine info from %s: %s", seed_url, e)
|
||||
return None, None, None
|
||||
|
||||
|
||||
def get_remote_instance_weight_shape_info(seed_url: str, local_seed_key: str):
|
||||
try:
|
||||
response = requests.get(
|
||||
f"{seed_url}/get_rfork_transfer_engine_shape_info",
|
||||
params={"seed_key": local_seed_key},
|
||||
)
|
||||
if response.status_code != 200:
|
||||
logger.debug(
|
||||
"GET %s/get_rfork_transfer_engine_shape_info failed: %s",
|
||||
seed_url,
|
||||
response.status_code,
|
||||
)
|
||||
return None
|
||||
|
||||
data = response.json()
|
||||
info = data.get("rfork_transfer_engine_shape_info", None)
|
||||
if info is None or isinstance(info, dict):
|
||||
return info
|
||||
|
||||
logger.error(
|
||||
"Failed to get rfork_transfer_engine_shape_info in response from %s.",
|
||||
seed_url,
|
||||
)
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.debug("Exception getting transfer engine shape info from %s: %s", seed_url, e)
|
||||
return None
|
||||
Reference in New Issue
Block a user