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