init v0.23.0

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

View File

@@ -0,0 +1,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

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

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

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

View 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

View 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

View 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