Revert several PRs (#14958)

Co-authored-by: fzyzcjy <ch271828n@outlook.com>
This commit is contained in:
Yineng Zhang
2025-12-12 11:25:12 -08:00
committed by GitHub
co-authored by fzyzcjy
parent ec242f516e
commit 4b7b5af36a
11 changed files with 38 additions and 500 deletions
-2
View File
@@ -73,8 +73,6 @@ class LoadConfig:
remote_instance_weight_loader_seed_instance_ip: Optional[str] = None
remote_instance_weight_loader_seed_instance_service_port: Optional[int] = None
remote_instance_weight_loader_send_weights_group_ports: Optional[List[int]] = None
remote_instance_weight_loader_backend: Optional[str] = None
remote_instance_weight_loader_transfer_engine: Optional[any] = None
# ModelOpt-specific loading options
modelopt_checkpoint_restore_path: Optional[str] = None
+4 -25
View File
@@ -127,18 +127,13 @@ class Engine(EngineBase):
atexit.register(self.shutdown)
# Launch subprocesses
(
tokenizer_manager,
template_manager,
scheduler_info,
port_args,
remote_instance_transfer_engine_info,
) = _launch_subprocesses(server_args=server_args)
tokenizer_manager, template_manager, scheduler_info, port_args = (
_launch_subprocesses(server_args=server_args)
)
self.tokenizer_manager = tokenizer_manager
self.template_manager = template_manager
self.scheduler_info = scheduler_info
self.port_args = port_args
self.remote_instance_transfer_engine_info = remote_instance_transfer_engine_info
# Initialize ZMQ sockets
context = zmq.Context(2)
@@ -915,7 +910,6 @@ def _launch_subprocesses(
# Wait for the model to finish loading
scheduler_infos = []
remote_instance_transfer_engine_info = {}
for i in range(len(scheduler_pipe_readers)):
try:
data = scheduler_pipe_readers[i].recv()
@@ -932,24 +926,9 @@ def _launch_subprocesses(
"Initialization failed. Please see the error messages above."
)
scheduler_infos.append(data)
if (
"tp_rank" in data
and "remote_instance_transfer_engine_session_id" in data
and "remote_instance_transfer_engine_weights_info_dict" in data
):
remote_instance_transfer_engine_info[data["tp_rank"]] = (
data["remote_instance_transfer_engine_session_id"],
data["remote_instance_transfer_engine_weights_info_dict"],
)
# Assume all schedulers have the same scheduler_info
scheduler_info = scheduler_infos[0]
tokenizer_manager.max_req_input_len = scheduler_info["max_req_input_len"]
return (
tokenizer_manager,
template_manager,
scheduler_info,
port_args,
remote_instance_transfer_engine_info,
)
return tokenizer_manager, template_manager, scheduler_info, port_args
+3 -35
View File
@@ -144,15 +144,6 @@ class _GlobalState:
tokenizer_manager: Union[TokenizerManager, MultiTokenizerRouter, TokenizerWorker]
template_manager: TemplateManager
scheduler_info: Dict
# Dict{
# rank: Tuple(
# session_id,
# Dict{
# name: Tuple (d_ptr, numel, element_size)
# }
# )
# }
remote_instance_transfer_engine_info: Optional[Dict] = None
_global_state: Optional[_GlobalState] = None
@@ -822,24 +813,6 @@ async def send_weights_to_remote_instance(
return ORJSONResponse(content, status_code=HTTPStatus.BAD_REQUEST)
@app.get("/get_remote_instance_transfer_engine_info")
async def get_remote_instance_transfer_engine_info(rank: int = None):
if rank is None or rank < 0:
return Response(status_code=HTTPStatus.BAD_REQUEST)
try:
result = {
"rank": rank,
"remote_instance_transfer_engine_info": _global_state.remote_instance_transfer_engine_info[
rank
],
}
return result
except Exception as e:
logger.error(f"Exception: {e}")
return Response(status_code=HTTPStatus.BAD_REQUEST)
@app.post("/init_weights_update_group")
async def init_weights_update_group(
obj: InitWeightsUpdateGroupReqInput, request: Request
@@ -1413,20 +1386,15 @@ def launch_server(
1. The HTTP server, Engine, and TokenizerManager all run in the main process.
2. Inter-process communication is done through IPC (each process uses a different port) via the ZMQ library.
"""
(
tokenizer_manager,
template_manager,
scheduler_info,
port_args,
remote_instance_transfer_engine_info,
) = _launch_subprocesses(server_args=server_args)
tokenizer_manager, template_manager, scheduler_info, port_args = (
_launch_subprocesses(server_args=server_args)
)
set_global_state(
_GlobalState(
tokenizer_manager=tokenizer_manager,
template_manager=template_manager,
scheduler_info=scheduler_info,
remote_instance_transfer_engine_info=remote_instance_transfer_engine_info,
)
)
+7 -26
View File
@@ -2609,9 +2609,6 @@ class Scheduler(
self.send_to_detokenizer.send_output(recv_req, recv_req)
return None
def get_remote_instance_transfer_engine_info(self):
return self.tp_worker.get_remote_instance_transfer_engine_info()
class IdleSleeper:
"""
@@ -2725,29 +2722,13 @@ def run_scheduler_process(
pp_rank,
dp_rank,
)
if server_args.remote_instance_weight_loader_support_transfer_engine:
(
remote_instance_transfer_engine_session_id,
remote_instance_transfer_engine_weights_info_dict,
) = scheduler.get_remote_instance_transfer_engine_info()
pipe_writer.send(
{
"status": "ready",
"max_total_num_tokens": scheduler.max_total_num_tokens,
"max_req_input_len": scheduler.max_req_input_len,
"tp_rank": tp_rank,
"remote_instance_transfer_engine_session_id": remote_instance_transfer_engine_session_id,
"remote_instance_transfer_engine_weights_info_dict": remote_instance_transfer_engine_weights_info_dict,
}
)
else:
pipe_writer.send(
{
"status": "ready",
"max_total_num_tokens": scheduler.max_total_num_tokens,
"max_req_input_len": scheduler.max_req_input_len,
}
)
pipe_writer.send(
{
"status": "ready",
"max_total_num_tokens": scheduler.max_total_num_tokens,
"max_req_input_len": scheduler.max_req_input_len,
}
)
disaggregation_mode: DisaggregationMode = scheduler.disaggregation_mode
if disaggregation_mode == DisaggregationMode.NULL:
-6
View File
@@ -366,12 +366,6 @@ class TpModelWorker(BaseTpWorker):
can_run_cuda_graph=can_run_cuda_graph,
)
def get_remote_instance_transfer_engine_info(self):
return (
self.model_runner.remote_instance_transfer_engine_session_id,
self.model_runner.remote_instance_transfer_engine_weight_info,
)
def forward_batch_generation(
self,
model_worker_batch: ModelWorkerBatch,
@@ -136,10 +136,9 @@ from sglang.srt.model_executor.input_buffers import GraphInputBuffers
from sglang.srt.model_executor.piecewise_cuda_graph_runner import (
PiecewiseCudaGraphRunner,
)
from sglang.srt.model_loader import get_model
from sglang.srt.model_loader.loader import DefaultModelLoader, get_model_loader
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
RemoteInstanceWeightLoaderBackend,
register_memory_region_v2,
trigger_init_weights_send_group_for_remote_instance_request,
)
from sglang.srt.model_loader.utils import set_default_torch_dtype
@@ -159,7 +158,6 @@ from sglang.srt.utils import (
get_available_gpu_memory,
get_bool_env_var,
get_cpu_ids_by_node,
get_local_ip_auto,
init_custom_process_group,
is_cuda,
is_float4_e2m1fn_x2,
@@ -319,10 +317,6 @@ class ModelRunner:
self.forward_pass_id = 0
self.init_new_workspace = False
self.remote_instance_transfer_engine = None
self.remote_instance_transfer_engine_session_id = ""
self.remote_instance_transfer_engine_weight_info = None
# Apply the rank zero filter to logger
if server_args.show_time_cost:
enable_show_time_cost()
@@ -397,9 +391,6 @@ class ModelRunner:
enable=self.server_args.enable_memory_saver
)
if self.server_args.remote_instance_weight_loader_support_transfer_engine:
self.remote_instance_init_transfer_engine()
if not self.is_draft_worker:
set_global_expert_location_metadata(
compute_initial_expert_location_metadata(
@@ -440,16 +431,6 @@ class ModelRunner:
self.sampler = Sampler()
self.load_model()
if (
self.server_args.remote_instance_weight_loader_support_transfer_engine
and self.remote_instance_transfer_engine_weight_info is None
):
self.remote_instance_transfer_engine_weight_info = (
register_memory_region_v2(
self.model, self.remote_instance_transfer_engine
)
)
# Check if the model is using hybrid SWA
if (
not self.server_args.disable_hybrid_swa_memory
@@ -564,23 +545,6 @@ class ModelRunner:
# Initialize piecewise CUDA graph
self.init_piecewise_cuda_graphs()
def remote_instance_init_transfer_engine(self):
try:
from mooncake.engine import TransferEngine
except ImportError as e:
logger.warning(
"Please install mooncake for using remote instance transfer engine: pip install mooncake"
)
return
self.remote_instance_transfer_engine = TransferEngine()
local_ip = get_local_ip_auto()
self.remote_instance_transfer_engine.initialize(
local_ip, "P2PHANDSHAKE", "rdma", envs.MOONCAKE_DEVICE.value
)
self.remote_instance_transfer_engine_session_id = (
f"{local_ip}:{self.remote_instance_transfer_engine.get_rpc_port()}"
)
def model_specific_adjustment(self):
server_args = self.server_args
@@ -798,8 +762,6 @@ class ModelRunner:
remote_instance_weight_loader_seed_instance_ip=self.server_args.remote_instance_weight_loader_seed_instance_ip,
remote_instance_weight_loader_seed_instance_service_port=self.server_args.remote_instance_weight_loader_seed_instance_service_port,
remote_instance_weight_loader_send_weights_group_ports=self.server_args.remote_instance_weight_loader_send_weights_group_ports,
remote_instance_weight_loader_backend=self.server_args.remote_instance_weight_loader_backend,
remote_instance_weight_loader_transfer_engine=self.remote_instance_transfer_engine,
modelopt_config=modelopt_config,
rl_quant_profile=self.server_args.rl_quant_profile,
)
@@ -808,11 +770,7 @@ class ModelRunner:
self.model_config, self.load_config, self.tp_size
)
if (
self.server_args.load_format == LoadFormat.REMOTE_INSTANCE
and self.server_args.remote_instance_weight_loader_backend
== RemoteInstanceWeightLoaderBackend.NCCL
):
if self.server_args.load_format == LoadFormat.REMOTE_INSTANCE:
if self.tp_rank == 0:
instance_ip = socket.gethostbyname(socket.gethostname())
t = threading.Thread(
@@ -837,18 +795,11 @@ class ModelRunner:
GPU_MEMORY_TYPE_WEIGHTS,
enable_cpu_backup=enable_cpu_backup,
):
self.loader = get_model_loader(
self.model = get_model(
model_config=self.model_config,
load_config=self.load_config,
model_config=self.model_config,
)
self.model = self.loader.load_model(
model_config=self.model_config,
device_config=DeviceConfig(self.device, self.gpu_id),
)
if hasattr(self.loader, "remote_instance_transfer_engine_weight_info"):
self.remote_instance_transfer_engine_weight_info = (
self.loader.remote_instance_transfer_engine_weight_info
)
monkey_patch_vllm_parallel_state(reverse=True)
get_offloader().post_init()
+4 -106
View File
@@ -34,11 +34,6 @@ import huggingface_hub
import numpy as np
import torch
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
RemoteInstanceWeightLoaderBackend,
get_remote_instance_transfer_engine_info_per_rank,
register_memory_region_v2,
)
from sglang.srt.server_args import get_global_server_args
# Try to import accelerate (optional dependency)
@@ -1992,7 +1987,6 @@ class RemoteInstanceModelLoader(BaseModelLoader):
f"Model loader extra config is not supported for "
f"load format {load_config.load_format}"
)
self.remote_instance_transfer_engine_weight_info = None
def download_model(self, model_config: ModelConfig) -> None:
raise NotImplementedError
@@ -2011,19 +2005,16 @@ class RemoteInstanceModelLoader(BaseModelLoader):
f"load format {load_config.load_format}"
)
model_weights = f"instance://{load_config.remote_instance_weight_loader_seed_instance_ip}:{load_config.remote_instance_weight_loader_send_weights_group_ports[load_config.tp_rank]}"
with set_default_torch_dtype(model_config.dtype):
with torch.device(device_config.device):
model = _initialize_model(model_config, self.load_config)
if (
load_config.remote_instance_weight_loader_backend
== RemoteInstanceWeightLoaderBackend.NCCL
):
model_weights = f"instance://{load_config.remote_instance_weight_loader_seed_instance_ip}:{load_config.remote_instance_weight_loader_send_weights_group_ports[load_config.tp_rank]}"
with create_remote_connector(model_weights, device_config.device) as client:
connector_type = get_connector_type(client)
if connector_type == ConnectorType.INSTANCE:
self.load_model_from_remote_instance_by_nccl(
self.load_model_from_remote_instance(
model, client, model_config, device_config
)
else:
@@ -2031,45 +2022,9 @@ class RemoteInstanceModelLoader(BaseModelLoader):
f"Unsupported connector type {connector_type} for "
f"remote tensor model loading."
)
elif (
load_config.remote_instance_weight_loader_backend
== RemoteInstanceWeightLoaderBackend.TRANSFER_ENGINE
):
if load_config.remote_instance_weight_loader_transfer_engine is None:
raise RuntimeError(
"Transfer engine is not initialized for remote instance "
"model loader with `transfer_engine` backend. "
)
logger.info(
"TransferEngine registering memory regions (this may take a few seconds)..."
)
# register memory region
self.remote_instance_transfer_engine_weight_info = (
register_memory_region_v2(
model, load_config.remote_instance_weight_loader_transfer_engine
)
)
logger.info(
"TransferEngine memory regions have been successfully registered."
)
# transfer weights
success = self.load_model_from_remote_instance_by_transfer_engine(
model,
load_config.remote_instance_weight_loader_transfer_engine,
f"http://{load_config.remote_instance_weight_loader_seed_instance_ip}:{load_config.remote_instance_weight_loader_seed_instance_service_port}",
load_config.tp_rank,
)
if not success:
raise RuntimeError(
"Failed to load weights from remote instance via transfer engine."
)
else:
raise ValueError("Invalid remote instance weight loader backend.")
return model.eval()
def load_model_from_remote_instance_by_nccl(
def load_model_from_remote_instance(
self, model, client, model_config: ModelConfig, device_config: DeviceConfig
) -> nn.Module:
load_config = self.load_config
@@ -2120,63 +2075,6 @@ class RemoteInstanceModelLoader(BaseModelLoader):
)
torch.cuda.empty_cache()
def load_model_from_remote_instance_by_transfer_engine(
self, model, transfer_engine, seed_url, tp_rank
) -> bool:
# get remote weights metadata from source instance
seed_transfer_engine_session_id, seed_transfer_engine_weight_info = (
get_remote_instance_transfer_engine_info_per_rank(seed_url, tp_rank)
)
if (
seed_transfer_engine_session_id is None
or seed_transfer_engine_weight_info is None
):
logger.error("Cannot get transfer engine session or weight info.")
return False
# prepare local/remote RDMA keys
seed_ptr_list = []
client_ptr_list = []
client_len_list = []
for name, tensor in model.named_parameters():
weight_info = seed_transfer_engine_weight_info.get(name, None)
if weight_info is None:
logger.error(f"Cannot find weight info for {name}.")
return False
seed_ptr, seed_numel, seed_element_size = weight_info
if (
seed_numel != tensor.numel()
or seed_element_size != tensor.element_size()
):
logger.error(
f"Weight info does not match for {name}, "
f"expected ({seed_numel}, {seed_element_size}), "
f"got ({tensor.numel()}, {tensor.element_size()})"
)
return False
client_ptr = tensor.data_ptr()
client_len = tensor.numel() * tensor.element_size()
seed_ptr_list.append(seed_ptr)
client_ptr_list.append(client_ptr)
client_len_list.append(client_len)
# load weights from source instance through TransferEngine
ret = transfer_engine.batch_transfer_sync_read(
seed_transfer_engine_session_id,
client_ptr_list,
seed_ptr_list,
client_len_list,
)
if ret < 0:
logger.error(f"batch transfer failed, error: {ret}")
return False
if hasattr(model, "post_load_weights"):
model.post_load_weights()
return True
class RemoteModelLoader(BaseModelLoader):
"""Model loader that can load Tensors from remote database."""
@@ -1,21 +1,13 @@
# SPDX-License-Identifier: Apache-2.0
import enum
import logging
import time
from typing import List
import requests
import torch
logger = logging.getLogger(__name__)
class RemoteInstanceWeightLoaderBackend(str, enum.Enum):
NCCL = "nccl"
TRANSFER_ENGINE = "transfer_engine"
def trigger_init_weights_send_group_for_remote_instance_request(
remote_instance_weight_loader_seed_instance_ip: str,
remote_instance_weight_loader_seed_instance_service_port: int,
@@ -75,110 +67,3 @@ def trigger_transferring_weights_request(
except Exception as e:
logger.error(f"Failed to trigger send weights to remote instance request: {e}")
raise
def get_remote_instance_transfer_engine_info_per_rank(seed_url: str, rank: int):
try:
response = requests.get(
f"{seed_url}/get_remote_instance_transfer_engine_info",
params={
"rank": rank,
},
)
if response.status_code == 200:
data = response.json()
if "remote_instance_transfer_engine_info" in data:
return data["remote_instance_transfer_engine_info"]
else:
logger.error(
"Failed to get `remote_instance_transfer_engine_info` in response."
)
return None, None
else:
logger.error(f"request.get failed: {response.status_code}")
return None, None
except Exception as e:
logger.error(f"Exception: {e}")
return None, None
# DEPRECATED. Use register_memory_region_v2 instead.
def register_memory_region(model, transfer_engine):
start_tic = time.time()
weight_mr_dict = {}
for name, weight in model.named_parameters():
ret = transfer_engine.register_memory(
weight.data_ptr(), weight.numel() * weight.element_size()
)
if ret != 0:
raise RuntimeError(
f"register memory failed for weight {name}, error: {ret}"
)
weight_mr_dict[name] = (
weight.data_ptr(),
weight.numel(),
weight.element_size(),
)
end_tic = time.time()
logger.debug(f"Register memory region time: {(end_tic - start_tic):.4f}s")
return weight_mr_dict
def register_memory_region_v2(model, transfer_engine):
start_tic = time.time()
weight_mr_dict = {}
weight_addr_set = set()
for name, weight in model.named_parameters():
weight_mr_dict[name] = (
weight.data_ptr(),
weight.numel(),
weight.element_size(),
)
weight_addr_set.add(weight.data_ptr())
memory_snapshot = torch.cuda.memory.memory_snapshot()
weight_blocks_for_reg_mr = []
# Blocks in each segment have continuous physical addresses,
# so they can be merged for memory registration.
for segment in memory_snapshot:
current_weight_block = None
blocks = segment.get("blocks", [])
for block in blocks:
address = block.get("address", -1)
size = block.get("size", -1)
state = block.get("state", "")
if address < 0 or size < 0 or state == "":
continue
# Only register active allocated memory blocks that hold weights.
if state == "active_allocated":
if address in weight_addr_set:
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)
# Register merged memory blocks that hold weights.
for weight_block in weight_blocks_for_reg_mr:
address, size = weight_block
ret = transfer_engine.register_memory(address, size)
if ret != 0:
raise RuntimeError(
f"register memory failed for weight block at address {address} with size {size}, error: {ret}"
)
end_tic = time.time()
logger.debug(f"Register memory region v2 time: {(end_tic - start_tic):.4f}s")
return weight_mr_dict
+10 -53
View File
@@ -18,7 +18,6 @@ from __future__ import annotations
import argparse
import dataclasses
import importlib
import importlib.util
import json
import logging
import os
@@ -596,8 +595,6 @@ class ServerArgs:
remote_instance_weight_loader_seed_instance_ip: Optional[str] = None
remote_instance_weight_loader_seed_instance_service_port: Optional[int] = None
remote_instance_weight_loader_send_weights_group_ports: Optional[List[int]] = None
remote_instance_weight_loader_backend: Literal["transfer_engine", "nccl"] = "nccl"
remote_instance_weight_loader_support_transfer_engine: bool = False
# For PD-Multiplexing
enable_pdmux: bool = False
@@ -710,9 +707,6 @@ class ServerArgs:
# Handle elastic expert parallelism.
self._handle_elastic_ep()
# Handle remote instance weight loader.
self._handle_remote_instance_weight_loader_support_transfer_engine()
def _handle_deprecated_args(self):
# handle deprecated tool call parsers
deprecated_tool_call_parsers = {"qwen25": "qwen", "glm45": "glm"}
@@ -1971,26 +1965,8 @@ class ServerArgs:
if (
self.remote_instance_weight_loader_seed_instance_ip is None
or self.remote_instance_weight_loader_seed_instance_service_port is None
or self.remote_instance_weight_loader_send_weights_group_ports is None
):
logger.warning(
"Fallback load_format to 'auto' due to incomplete remote instance weight loader settings."
)
self.load_format = "auto"
elif (
self.remote_instance_weight_loader_send_weights_group_ports is None
and self.remote_instance_weight_loader_backend == "nccl"
):
logger.warning(
"Fallback load_format to 'auto' due to incomplete remote instance weight loader NCCL group ports settings."
)
self.load_format = "auto"
elif (
self.enable_memory_saver
and self.remote_instance_weight_loader_backend == "transfer_engine"
):
logger.warning(
"Fallback load_format to 'auto' due to incompatible remote instance weight loader transfer engine backend with memory saver."
)
self.load_format = "auto"
def _handle_disaggregation(self):
@@ -2226,25 +2202,18 @@ class ServerArgs:
self.skip_server_warmup = True
def _handle_remote_instance_weight_loader_support_transfer_engine(self):
try:
importlib.import_module("mooncake")
if importlib.util.find_spec("mooncake.engine") is None:
logger.warning(
f"Failed to import mooncake.engine. Does not support using TransferEngine as remote instance weight loader backend."
)
self.remote_instance_weight_loader_support_transfer_engine = False
elif self.enable_memory_saver:
logger.warning(
"Memory saver is enabled, which is not compatible with TransferEngine. Does not support using TransferEngine as remote instance weight loader backend."
)
self.remote_instance_weight_loader_support_transfer_engine = False
else:
self.remote_instance_weight_loader_support_transfer_engine = True
except ImportError:
if importlib.util.find_spec("mooncake.engine") is None:
logger.warning(
f"Failed to import mooncake. Does not support using TransferEngine as remote instance weight loader backend."
f"Failed to import mooncake.engine. Does not support using TransferEngine as remote instance weight loader backend."
)
self.remote_instance_weight_loader_support_transfer_engine = False
elif self.enable_memory_saver:
logger.warning(
"Memory saver is enabled, which is not compatible with TransferEngine. Does not support using TransferEngine as remote instance weight loader backend."
)
self.remote_instance_weight_loader_support_transfer_engine = False
else:
self.remote_instance_weight_loader_support_transfer_engine = True
@staticmethod
def add_cli_args(parser: argparse.ArgumentParser):
@@ -4087,18 +4056,6 @@ class ServerArgs:
default=ServerArgs.remote_instance_weight_loader_send_weights_group_ports,
help="The communication group ports for loading weights from remote instance.",
)
parser.add_argument(
"--remote-instance-weight-loader-backend",
type=str,
choices=["transfer_engine", "nccl"],
default=ServerArgs.remote_instance_weight_loader_backend,
help="The backend for loading weights from remote instance. Can be 'transfer_engine' or 'nccl'. Default is 'nccl'.",
)
parser.add_argument(
"--remote-instance-weight-loader-support-transfer-engine",
action="store_true",
help="Enable transfer engine support for remote instance weight loader.",
)
# For PD-Multiplexing
parser.add_argument(