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

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

View 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

View 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

View 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

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

View 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

View 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