support non disturbing remote instance weight loader v2 (#14997)
Signed-off-by: Anqi Shen <amy.saq@antgroup.com>
This commit is contained in:
@@ -0,0 +1,49 @@
|
|||||||
|
# R-Fork
|
||||||
|
|
||||||
|
R-Fork (Tensor Remote Fork) is a novel weight loading methodology that leverages efficient inter-node GPU-to-GPU data transfer path to load tensors from a running SGLang instance to a new instance with zero-copy. It can significantly optimize the SGLang instance boot-up time by reducing model weights loading from several minutes to mere seconds.
|
||||||
|
|
||||||
|
To learn more details about R-Fork, please check **<a href=https://lmsys.org/blog/2025-12-10-rfork/> R-Fork blog </a>**
|
||||||
|
|
||||||
|
## Usage
|
||||||
|
|
||||||
|
| Argument | Usage |
|
||||||
|
|--------------|--------------------------------------------|
|
||||||
|
| load-format | set to `remote_instance` to enable R-Fork. |
|
||||||
|
| remote-instance-weight-loader-backend | `nccl` or `transfer_engine`, default value is `nccl` |
|
||||||
|
| remote-instance-weight-loader-seed-instance-ip | IP address of the seed instance who will provide the model weight |
|
||||||
|
| remote-instance-weight-loader-seed-instance-service-port | the port that the seed instance's HTTP server is listening on |
|
||||||
|
| remote-instance-weight-loader-send-weights-group-ports | the list of available ports on the seed instance that will be used to build NCCL communication groups between seed and client instance. This argument is only needed by `nccl` backend. |
|
||||||
|
| remote-instance-weight-loader-start-seed-via-transfer-engine | set to start seed service that supports TransferEngine as backend. It is needed for seed instances when using `transfer_engine` as backend. |
|
||||||
|
|
||||||
|
### NCCL as backend
|
||||||
|
|
||||||
|
seed instance:
|
||||||
|
```shell
|
||||||
|
python -m sglang.launch_server [args]
|
||||||
|
```
|
||||||
|
|
||||||
|
client instance:
|
||||||
|
```shell
|
||||||
|
python -m sglang.launch_server [args] \
|
||||||
|
--load-format remote_instance \
|
||||||
|
--remote-instance-weight-loader-seed-instance-ip [seed_instance_ip] \
|
||||||
|
--remote-instance-weight-loader-seed-instance-service-port [seed_instance_service_port] \
|
||||||
|
--remote-instance-weight-loader-send-weights-group-ports [send_weights_nccl_group_ports_list] \
|
||||||
|
--remote-instance-weight-loader-backend nccl
|
||||||
|
```
|
||||||
|
|
||||||
|
### TransferEngine as backend
|
||||||
|
|
||||||
|
seed instance:
|
||||||
|
```shell
|
||||||
|
python -m sglang.launch_server [args] \
|
||||||
|
--remote-instance-weight-loader-start-seed-via-transfer-engine
|
||||||
|
```
|
||||||
|
|
||||||
|
```shell
|
||||||
|
python -m sglang.launch_server [args] \
|
||||||
|
--load-format remote_instance \
|
||||||
|
--remote-instance-weight-loader-seed-instance-ip [seed_instance_ip] \
|
||||||
|
--remote-instance-weight-loader-seed-instance-service-port [seed_instance_service_port] \
|
||||||
|
--remote-instance-weight-loader-backend transfer_engine
|
||||||
|
```
|
||||||
@@ -2,7 +2,7 @@
|
|||||||
import enum
|
import enum
|
||||||
import logging
|
import logging
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import List, Optional, Union
|
from typing import Any, List, Optional, Union
|
||||||
|
|
||||||
import orjson
|
import orjson
|
||||||
|
|
||||||
@@ -73,6 +73,8 @@ class LoadConfig:
|
|||||||
remote_instance_weight_loader_seed_instance_ip: Optional[str] = None
|
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_seed_instance_service_port: Optional[int] = None
|
||||||
remote_instance_weight_loader_send_weights_group_ports: Optional[List[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-specific loading options
|
||||||
modelopt_checkpoint_restore_path: Optional[str] = None
|
modelopt_checkpoint_restore_path: Optional[str] = None
|
||||||
|
|||||||
@@ -63,6 +63,9 @@ from sglang.srt.managers.multi_tokenizer_mixin import MultiTokenizerRouter
|
|||||||
from sglang.srt.managers.scheduler import run_scheduler_process
|
from sglang.srt.managers.scheduler import run_scheduler_process
|
||||||
from sglang.srt.managers.template_manager import TemplateManager
|
from sglang.srt.managers.template_manager import TemplateManager
|
||||||
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||||
|
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
||||||
|
parse_remote_instance_transfer_engine_info_from_scheduler_infos,
|
||||||
|
)
|
||||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||||
from sglang.srt.tracing.trace import process_tracing_init, trace_set_thread_info
|
from sglang.srt.tracing.trace import process_tracing_init, trace_set_thread_info
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
@@ -173,10 +176,9 @@ def _launch_subprocesses(
|
|||||||
scheduler_infos.append(data)
|
scheduler_infos.append(data)
|
||||||
|
|
||||||
# Get back some info from scheduler to tokenizer_manager
|
# Get back some info from scheduler to tokenizer_manager
|
||||||
scheduler_info = scheduler_infos[0]
|
tokenizer_manager.max_req_input_len = scheduler_infos[0]["max_req_input_len"]
|
||||||
tokenizer_manager.max_req_input_len = scheduler_info["max_req_input_len"]
|
|
||||||
|
|
||||||
return tokenizer_manager, template_manager, scheduler_info, port_args
|
return tokenizer_manager, template_manager, scheduler_infos, port_args
|
||||||
|
|
||||||
|
|
||||||
class Engine(EngineBase):
|
class Engine(EngineBase):
|
||||||
@@ -221,13 +223,21 @@ class Engine(EngineBase):
|
|||||||
atexit.register(self.shutdown)
|
atexit.register(self.shutdown)
|
||||||
|
|
||||||
# Launch subprocesses
|
# Launch subprocesses
|
||||||
tokenizer_manager, template_manager, scheduler_info, port_args = (
|
tokenizer_manager, template_manager, scheduler_infos, port_args = (
|
||||||
self.launch_subprocesses_func(server_args=server_args)
|
self.launch_subprocesses_func(server_args=server_args)
|
||||||
)
|
)
|
||||||
self.tokenizer_manager = tokenizer_manager
|
self.tokenizer_manager = tokenizer_manager
|
||||||
self.template_manager = template_manager
|
self.template_manager = template_manager
|
||||||
|
|
||||||
|
scheduler_info = scheduler_infos[0]
|
||||||
self.scheduler_info = scheduler_info
|
self.scheduler_info = scheduler_info
|
||||||
|
|
||||||
self.port_args = port_args
|
self.port_args = port_args
|
||||||
|
self.remote_instance_transfer_engine_info = (
|
||||||
|
parse_remote_instance_transfer_engine_info_from_scheduler_infos(
|
||||||
|
scheduler_infos
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
# Initialize ZMQ sockets
|
# Initialize ZMQ sockets
|
||||||
context = zmq.Context(2)
|
context = zmq.Context(2)
|
||||||
|
|||||||
@@ -123,6 +123,9 @@ from sglang.srt.managers.multi_tokenizer_mixin import (
|
|||||||
from sglang.srt.managers.template_manager import TemplateManager
|
from sglang.srt.managers.template_manager import TemplateManager
|
||||||
from sglang.srt.managers.tokenizer_manager import ServerStatus, TokenizerManager
|
from sglang.srt.managers.tokenizer_manager import ServerStatus, TokenizerManager
|
||||||
from sglang.srt.metrics.func_timer import enable_func_timer
|
from sglang.srt.metrics.func_timer import enable_func_timer
|
||||||
|
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
||||||
|
parse_remote_instance_transfer_engine_info_from_scheduler_infos,
|
||||||
|
)
|
||||||
from sglang.srt.parser.reasoning_parser import ReasoningParser
|
from sglang.srt.parser.reasoning_parser import ReasoningParser
|
||||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||||
from sglang.srt.tracing.trace import process_tracing_init, trace_set_thread_info
|
from sglang.srt.tracing.trace import process_tracing_init, trace_set_thread_info
|
||||||
@@ -152,6 +155,15 @@ class _GlobalState:
|
|||||||
tokenizer_manager: Union[TokenizerManager, MultiTokenizerRouter, TokenizerWorker]
|
tokenizer_manager: Union[TokenizerManager, MultiTokenizerRouter, TokenizerWorker]
|
||||||
template_manager: TemplateManager
|
template_manager: TemplateManager
|
||||||
scheduler_info: Dict
|
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
|
_global_state: Optional[_GlobalState] = None
|
||||||
@@ -825,6 +837,30 @@ async def send_weights_to_remote_instance(
|
|||||||
return ORJSONResponse(content, status_code=HTTPStatus.BAD_REQUEST)
|
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)
|
||||||
|
|
||||||
|
if (
|
||||||
|
_global_state.remote_instance_transfer_engine_info is None
|
||||||
|
or len(_global_state.remote_instance_transfer_engine_info) == 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")
|
@app.post("/init_weights_update_group")
|
||||||
async def init_weights_update_group(
|
async def init_weights_update_group(
|
||||||
obj: InitWeightsUpdateGroupReqInput, request: Request
|
obj: InitWeightsUpdateGroupReqInput, request: Request
|
||||||
@@ -1615,15 +1651,24 @@ def launch_server(
|
|||||||
1. The HTTP server, Engine, and TokenizerManager all run in the main process.
|
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.
|
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 = (
|
tokenizer_manager, template_manager, scheduler_infos, port_args = (
|
||||||
launch_subprocesses_func(server_args=server_args)
|
launch_subprocesses_func(server_args=server_args)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
scheduler_info = scheduler_infos[0]
|
||||||
|
remote_instance_transfer_engine_info = None
|
||||||
|
if server_args.remote_instance_weight_loader_use_transfer_engine():
|
||||||
|
remote_instance_transfer_engine_info = (
|
||||||
|
parse_remote_instance_transfer_engine_info_from_scheduler_infos(
|
||||||
|
scheduler_infos
|
||||||
|
)
|
||||||
|
)
|
||||||
set_global_state(
|
set_global_state(
|
||||||
_GlobalState(
|
_GlobalState(
|
||||||
tokenizer_manager=tokenizer_manager,
|
tokenizer_manager=tokenizer_manager,
|
||||||
template_manager=template_manager,
|
template_manager=template_manager,
|
||||||
scheduler_info=scheduler_info,
|
scheduler_info=scheduler_info,
|
||||||
|
remote_instance_transfer_engine_info=remote_instance_transfer_engine_info,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -2656,6 +2656,9 @@ class Scheduler(
|
|||||||
self.send_to_detokenizer.send_output(recv_req, recv_req)
|
self.send_to_detokenizer.send_output(recv_req, recv_req)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
def get_remote_instance_transfer_engine_info(self):
|
||||||
|
return self.tp_worker.get_remote_instance_transfer_engine_info()
|
||||||
|
|
||||||
|
|
||||||
class IdleSleeper:
|
class IdleSleeper:
|
||||||
"""
|
"""
|
||||||
@@ -2769,14 +2772,25 @@ def run_scheduler_process(
|
|||||||
pp_rank,
|
pp_rank,
|
||||||
dp_rank,
|
dp_rank,
|
||||||
)
|
)
|
||||||
pipe_writer.send(
|
result_dict = {
|
||||||
{
|
|
||||||
"status": "ready",
|
"status": "ready",
|
||||||
"max_total_num_tokens": scheduler.max_total_num_tokens,
|
"max_total_num_tokens": scheduler.max_total_num_tokens,
|
||||||
"max_req_input_len": scheduler.max_req_input_len,
|
"max_req_input_len": scheduler.max_req_input_len,
|
||||||
}
|
}
|
||||||
|
if server_args.remote_instance_weight_loader_use_transfer_engine():
|
||||||
|
(
|
||||||
|
remote_instance_transfer_engine_session_id,
|
||||||
|
remote_instance_transfer_engine_weights_info_dict,
|
||||||
|
) = scheduler.get_remote_instance_transfer_engine_info()
|
||||||
|
result_dict.update(
|
||||||
|
{
|
||||||
|
"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,
|
||||||
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
pipe_writer.send(result_dict)
|
||||||
disaggregation_mode: DisaggregationMode = scheduler.disaggregation_mode
|
disaggregation_mode: DisaggregationMode = scheduler.disaggregation_mode
|
||||||
if disaggregation_mode == DisaggregationMode.NULL:
|
if disaggregation_mode == DisaggregationMode.NULL:
|
||||||
if scheduler.enable_pdmux:
|
if scheduler.enable_pdmux:
|
||||||
|
|||||||
@@ -366,6 +366,12 @@ class TpModelWorker(BaseTpWorker):
|
|||||||
can_run_cuda_graph=can_run_cuda_graph,
|
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(
|
def forward_batch_generation(
|
||||||
self,
|
self,
|
||||||
model_worker_batch: ModelWorkerBatch,
|
model_worker_batch: ModelWorkerBatch,
|
||||||
|
|||||||
@@ -136,9 +136,10 @@ from sglang.srt.model_executor.input_buffers import GraphInputBuffers
|
|||||||
from sglang.srt.model_executor.piecewise_cuda_graph_runner import (
|
from sglang.srt.model_executor.piecewise_cuda_graph_runner import (
|
||||||
PiecewiseCudaGraphRunner,
|
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.loader import DefaultModelLoader, get_model_loader
|
||||||
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
||||||
|
RemoteInstanceWeightLoaderBackend,
|
||||||
|
register_memory_region,
|
||||||
trigger_init_weights_send_group_for_remote_instance_request,
|
trigger_init_weights_send_group_for_remote_instance_request,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_loader.utils import set_default_torch_dtype
|
from sglang.srt.model_loader.utils import set_default_torch_dtype
|
||||||
@@ -158,6 +159,7 @@ from sglang.srt.utils import (
|
|||||||
get_available_gpu_memory,
|
get_available_gpu_memory,
|
||||||
get_bool_env_var,
|
get_bool_env_var,
|
||||||
get_cpu_ids_by_node,
|
get_cpu_ids_by_node,
|
||||||
|
get_local_ip_auto,
|
||||||
init_custom_process_group,
|
init_custom_process_group,
|
||||||
is_cuda,
|
is_cuda,
|
||||||
is_float4_e2m1fn_x2,
|
is_float4_e2m1fn_x2,
|
||||||
@@ -319,6 +321,10 @@ class ModelRunner:
|
|||||||
self.forward_pass_id = 0
|
self.forward_pass_id = 0
|
||||||
self.init_new_workspace = False
|
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
|
# Apply the rank zero filter to logger
|
||||||
if server_args.show_time_cost:
|
if server_args.show_time_cost:
|
||||||
enable_show_time_cost()
|
enable_show_time_cost()
|
||||||
@@ -393,6 +399,9 @@ class ModelRunner:
|
|||||||
enable=self.server_args.enable_memory_saver
|
enable=self.server_args.enable_memory_saver
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if self.server_args.remote_instance_weight_loader_use_transfer_engine():
|
||||||
|
self.remote_instance_init_transfer_engine()
|
||||||
|
|
||||||
if not self.is_draft_worker:
|
if not self.is_draft_worker:
|
||||||
set_global_expert_location_metadata(
|
set_global_expert_location_metadata(
|
||||||
compute_initial_expert_location_metadata(
|
compute_initial_expert_location_metadata(
|
||||||
@@ -433,6 +442,15 @@ class ModelRunner:
|
|||||||
self.sampler = Sampler()
|
self.sampler = Sampler()
|
||||||
self.load_model()
|
self.load_model()
|
||||||
|
|
||||||
|
if (
|
||||||
|
self.server_args.remote_instance_weight_loader_use_transfer_engine()
|
||||||
|
and self.remote_instance_transfer_engine is not None
|
||||||
|
and self.remote_instance_transfer_engine_weight_info is None
|
||||||
|
):
|
||||||
|
self.remote_instance_transfer_engine_weight_info = register_memory_region(
|
||||||
|
self.model, self.remote_instance_transfer_engine
|
||||||
|
)
|
||||||
|
|
||||||
# Check if the model is using hybrid SWA
|
# Check if the model is using hybrid SWA
|
||||||
if (
|
if (
|
||||||
not self.server_args.disable_hybrid_swa_memory
|
not self.server_args.disable_hybrid_swa_memory
|
||||||
@@ -547,6 +565,23 @@ class ModelRunner:
|
|||||||
# Initialize piecewise CUDA graph
|
# Initialize piecewise CUDA graph
|
||||||
self.init_piecewise_cuda_graphs()
|
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):
|
def model_specific_adjustment(self):
|
||||||
server_args = self.server_args
|
server_args = self.server_args
|
||||||
|
|
||||||
@@ -764,6 +799,8 @@ 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_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_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_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,
|
modelopt_config=modelopt_config,
|
||||||
rl_quant_profile=self.server_args.rl_quant_profile,
|
rl_quant_profile=self.server_args.rl_quant_profile,
|
||||||
)
|
)
|
||||||
@@ -772,7 +809,11 @@ class ModelRunner:
|
|||||||
self.model_config, self.load_config, self.tp_size
|
self.model_config, self.load_config, self.tp_size
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.server_args.load_format == LoadFormat.REMOTE_INSTANCE:
|
if (
|
||||||
|
self.server_args.load_format == LoadFormat.REMOTE_INSTANCE
|
||||||
|
and self.server_args.remote_instance_weight_loader_backend
|
||||||
|
== RemoteInstanceWeightLoaderBackend.NCCL
|
||||||
|
):
|
||||||
if self.tp_rank == 0:
|
if self.tp_rank == 0:
|
||||||
instance_ip = socket.gethostbyname(socket.gethostname())
|
instance_ip = socket.gethostbyname(socket.gethostname())
|
||||||
t = threading.Thread(
|
t = threading.Thread(
|
||||||
@@ -797,11 +838,18 @@ class ModelRunner:
|
|||||||
GPU_MEMORY_TYPE_WEIGHTS,
|
GPU_MEMORY_TYPE_WEIGHTS,
|
||||||
enable_cpu_backup=enable_cpu_backup,
|
enable_cpu_backup=enable_cpu_backup,
|
||||||
):
|
):
|
||||||
self.model = get_model(
|
self.loader = get_model_loader(
|
||||||
model_config=self.model_config,
|
|
||||||
load_config=self.load_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),
|
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)
|
monkey_patch_vllm_parallel_state(reverse=True)
|
||||||
|
|
||||||
get_offloader().post_init()
|
get_offloader().post_init()
|
||||||
|
|||||||
@@ -34,6 +34,11 @@ import huggingface_hub
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
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,
|
||||||
|
)
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
|
|
||||||
# Try to import accelerate (optional dependency)
|
# Try to import accelerate (optional dependency)
|
||||||
@@ -1987,6 +1992,7 @@ class RemoteInstanceModelLoader(BaseModelLoader):
|
|||||||
f"Model loader extra config is not supported for "
|
f"Model loader extra config is not supported for "
|
||||||
f"load format {load_config.load_format}"
|
f"load format {load_config.load_format}"
|
||||||
)
|
)
|
||||||
|
self.remote_instance_transfer_engine_weight_info = None
|
||||||
|
|
||||||
def download_model(self, model_config: ModelConfig) -> None:
|
def download_model(self, model_config: ModelConfig) -> None:
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
@@ -2005,16 +2011,19 @@ class RemoteInstanceModelLoader(BaseModelLoader):
|
|||||||
f"load format {load_config.load_format}"
|
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 set_default_torch_dtype(model_config.dtype):
|
||||||
with torch.device(device_config.device):
|
with torch.device(device_config.device):
|
||||||
model = _initialize_model(model_config, self.load_config)
|
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:
|
with create_remote_connector(model_weights, device_config.device) as client:
|
||||||
connector_type = get_connector_type(client)
|
connector_type = get_connector_type(client)
|
||||||
if connector_type == ConnectorType.INSTANCE:
|
if connector_type == ConnectorType.INSTANCE:
|
||||||
self.load_model_from_remote_instance(
|
self.load_model_from_remote_instance_by_nccl(
|
||||||
model, client, model_config, device_config
|
model, client, model_config, device_config
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -2022,9 +2031,43 @@ class RemoteInstanceModelLoader(BaseModelLoader):
|
|||||||
f"Unsupported connector type {connector_type} for "
|
f"Unsupported connector type {connector_type} for "
|
||||||
f"remote tensor model loading."
|
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(
|
||||||
|
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()
|
return model.eval()
|
||||||
|
|
||||||
def load_model_from_remote_instance(
|
def load_model_from_remote_instance_by_nccl(
|
||||||
self, model, client, model_config: ModelConfig, device_config: DeviceConfig
|
self, model, client, model_config: ModelConfig, device_config: DeviceConfig
|
||||||
) -> nn.Module:
|
) -> nn.Module:
|
||||||
load_config = self.load_config
|
load_config = self.load_config
|
||||||
@@ -2075,6 +2118,63 @@ class RemoteInstanceModelLoader(BaseModelLoader):
|
|||||||
)
|
)
|
||||||
torch.cuda.empty_cache()
|
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):
|
class RemoteModelLoader(BaseModelLoader):
|
||||||
"""Model loader that can load Tensors from remote database."""
|
"""Model loader that can load Tensors from remote database."""
|
||||||
|
|||||||
@@ -1,6 +1,10 @@
|
|||||||
# SPDX-License-Identifier: Apache-2.0
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
import enum
|
||||||
|
import importlib
|
||||||
|
import importlib.util
|
||||||
import logging
|
import logging
|
||||||
|
import time
|
||||||
from typing import List
|
from typing import List
|
||||||
|
|
||||||
import requests
|
import requests
|
||||||
@@ -8,6 +12,11 @@ import requests
|
|||||||
logger = logging.getLogger(__name__)
|
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(
|
def trigger_init_weights_send_group_for_remote_instance_request(
|
||||||
remote_instance_weight_loader_seed_instance_ip: str,
|
remote_instance_weight_loader_seed_instance_ip: str,
|
||||||
remote_instance_weight_loader_seed_instance_service_port: int,
|
remote_instance_weight_loader_seed_instance_service_port: int,
|
||||||
@@ -67,3 +76,133 @@ def trigger_transferring_weights_request(
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Failed to trigger send weights to remote instance request: {e}")
|
logger.error(f"Failed to trigger send weights to remote instance request: {e}")
|
||||||
raise
|
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
|
||||||
|
|
||||||
|
|
||||||
|
def parse_remote_instance_transfer_engine_info_from_scheduler_infos(scheduler_infos):
|
||||||
|
remote_instance_transfer_engine_info = {}
|
||||||
|
for data in scheduler_infos:
|
||||||
|
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"],
|
||||||
|
)
|
||||||
|
return remote_instance_transfer_engine_info
|
||||||
|
|
||||||
|
|
||||||
|
def register_memory_region(model, transfer_engine):
|
||||||
|
if importlib.util.find_spec("torch") is None:
|
||||||
|
return register_memory_region_v1(model, transfer_engine)
|
||||||
|
else:
|
||||||
|
return register_memory_region_v2(model, transfer_engine)
|
||||||
|
|
||||||
|
|
||||||
|
def register_memory_region_v1(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())
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
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
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ from __future__ import annotations
|
|||||||
import argparse
|
import argparse
|
||||||
import dataclasses
|
import dataclasses
|
||||||
import importlib
|
import importlib
|
||||||
|
import importlib.util
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
@@ -614,6 +615,8 @@ class ServerArgs:
|
|||||||
remote_instance_weight_loader_seed_instance_ip: Optional[str] = None
|
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_seed_instance_service_port: Optional[int] = None
|
||||||
remote_instance_weight_loader_send_weights_group_ports: Optional[List[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_start_seed_via_transfer_engine: bool = False
|
||||||
|
|
||||||
# For PD-Multiplexing
|
# For PD-Multiplexing
|
||||||
enable_pdmux: bool = False
|
enable_pdmux: bool = False
|
||||||
@@ -690,6 +693,9 @@ class ServerArgs:
|
|||||||
# Handle speculative decoding logic.
|
# Handle speculative decoding logic.
|
||||||
self._handle_speculative_decoding()
|
self._handle_speculative_decoding()
|
||||||
|
|
||||||
|
# Handle remote instance weight loader.
|
||||||
|
self._handle_remote_instance_weight_loader_start_seed_via_transfer_engine()
|
||||||
|
|
||||||
# Handle model loading format.
|
# Handle model loading format.
|
||||||
self._handle_load_format()
|
self._handle_load_format()
|
||||||
|
|
||||||
@@ -2107,8 +2113,26 @@ class ServerArgs:
|
|||||||
if (
|
if (
|
||||||
self.remote_instance_weight_loader_seed_instance_ip is None
|
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_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 (
|
||||||
|
not self.validate_transfer_engine()
|
||||||
|
and self.remote_instance_weight_loader_backend == "transfer_engine"
|
||||||
|
):
|
||||||
|
logger.warning(
|
||||||
|
"Fallback load_format to 'auto' due to 'transfer_engine' backend is not supported."
|
||||||
|
)
|
||||||
self.load_format = "auto"
|
self.load_format = "auto"
|
||||||
|
|
||||||
def _handle_encoder_disaggregation(self):
|
def _handle_encoder_disaggregation(self):
|
||||||
@@ -2366,19 +2390,12 @@ class ServerArgs:
|
|||||||
self.disable_cuda_graph = True
|
self.disable_cuda_graph = True
|
||||||
self.skip_server_warmup = True
|
self.skip_server_warmup = True
|
||||||
|
|
||||||
def _handle_remote_instance_weight_loader_support_transfer_engine(self):
|
def _handle_remote_instance_weight_loader_start_seed_via_transfer_engine(self):
|
||||||
if importlib.util.find_spec("mooncake.engine") is None:
|
# Check whether TransferEngine can be used when users want to start seed service that supports TransferEngine backend.
|
||||||
logger.warning(
|
if self.remote_instance_weight_loader_start_seed_via_transfer_engine:
|
||||||
f"Failed to import mooncake.engine. Does not support using TransferEngine as remote instance weight loader backend."
|
self.remote_instance_weight_loader_start_seed_via_transfer_engine = (
|
||||||
|
self.validate_transfer_engine()
|
||||||
)
|
)
|
||||||
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
|
@staticmethod
|
||||||
def add_cli_args(parser: argparse.ArgumentParser):
|
def add_cli_args(parser: argparse.ArgumentParser):
|
||||||
@@ -4277,6 +4294,18 @@ class ServerArgs:
|
|||||||
default=ServerArgs.remote_instance_weight_loader_send_weights_group_ports,
|
default=ServerArgs.remote_instance_weight_loader_send_weights_group_ports,
|
||||||
help="The communication group ports for loading weights from remote instance.",
|
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-start-seed-via-transfer-engine",
|
||||||
|
action="store_true",
|
||||||
|
help="Start seed server via transfer engine backend for remote instance weight loader.",
|
||||||
|
)
|
||||||
|
|
||||||
# For PD-Multiplexing
|
# For PD-Multiplexing
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
@@ -4782,6 +4811,33 @@ class ServerArgs:
|
|||||||
original_server_arg_mem_fraction * final_overall_factor
|
original_server_arg_mem_fraction * final_overall_factor
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def validate_transfer_engine(self):
|
||||||
|
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."
|
||||||
|
)
|
||||||
|
return 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."
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
else:
|
||||||
|
return True
|
||||||
|
|
||||||
|
def remote_instance_weight_loader_use_transfer_engine(self):
|
||||||
|
# Use TransferEngine as seed backend.
|
||||||
|
if self.remote_instance_weight_loader_start_seed_via_transfer_engine:
|
||||||
|
return True
|
||||||
|
# Use TransferEngine as client backend.
|
||||||
|
elif (
|
||||||
|
self.load_format == "remote_instance"
|
||||||
|
and self.remote_instance_weight_loader_backend == "transfer_engine"
|
||||||
|
):
|
||||||
|
return True
|
||||||
|
else:
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
# NOTE: This is a global variable to hold the server args for scheduler.
|
# NOTE: This is a global variable to hold the server args for scheduler.
|
||||||
_global_server_args: Optional[ServerArgs] = None
|
_global_server_args: Optional[ServerArgs] = None
|
||||||
|
|||||||
@@ -62,6 +62,7 @@ def init_process(
|
|||||||
seed_instance_group_base_port,
|
seed_instance_group_base_port,
|
||||||
event_seed_ready,
|
event_seed_ready,
|
||||||
event_dst_ready_list,
|
event_dst_ready_list,
|
||||||
|
remote_instance_loader_backend,
|
||||||
):
|
):
|
||||||
torch.cuda.set_device(rank)
|
torch.cuda.set_device(rank)
|
||||||
|
|
||||||
@@ -90,6 +91,7 @@ def init_process(
|
|||||||
tp_size,
|
tp_size,
|
||||||
event_seed_ready,
|
event_seed_ready,
|
||||||
event_dst_ready_list,
|
event_dst_ready_list,
|
||||||
|
remote_instance_loader_backend,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -122,6 +124,7 @@ def init_process_seed(
|
|||||||
str(rank),
|
str(rank),
|
||||||
"--tp-size",
|
"--tp-size",
|
||||||
str(tp_size),
|
str(tp_size),
|
||||||
|
"--remote-instance-weight-loader-start-seed-via-transfer-engine",
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
@@ -159,6 +162,7 @@ def init_process_dst(
|
|||||||
tp_size,
|
tp_size,
|
||||||
event_seed_ready,
|
event_seed_ready,
|
||||||
event_dst_ready_list,
|
event_dst_ready_list,
|
||||||
|
remote_instance_loader_backend,
|
||||||
):
|
):
|
||||||
torch.cuda.set_device(rank * tp_size)
|
torch.cuda.set_device(rank * tp_size)
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
@@ -186,6 +190,10 @@ def init_process_dst(
|
|||||||
remote_instance_weight_loader_seed_instance_service_port=seed_instance_service_port,
|
remote_instance_weight_loader_seed_instance_service_port=seed_instance_service_port,
|
||||||
remote_instance_weight_loader_send_weights_group_ports=ports,
|
remote_instance_weight_loader_send_weights_group_ports=ports,
|
||||||
load_format="remote_instance",
|
load_format="remote_instance",
|
||||||
|
remote_instance_weight_loader_backend=remote_instance_loader_backend,
|
||||||
|
remote_instance_weight_loader_start_seed_via_transfer_engine=(
|
||||||
|
remote_instance_loader_backend == "transfer_engine"
|
||||||
|
),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
host, _, port = DEFAULT_URL_FOR_TEST.rpartition(":")
|
host, _, port = DEFAULT_URL_FOR_TEST.rpartition(":")
|
||||||
@@ -213,6 +221,9 @@ def init_process_dst(
|
|||||||
f"[{','.join(str(port) for port in ports)}]",
|
f"[{','.join(str(port) for port in ports)}]",
|
||||||
"--load-format",
|
"--load-format",
|
||||||
"remote_instance",
|
"remote_instance",
|
||||||
|
"--remote-instance-weight-loader-backend",
|
||||||
|
remote_instance_loader_backend,
|
||||||
|
"--remote-instance-weight-loader-start-seed-via-transfer-engine",
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
@@ -250,9 +261,10 @@ def test_load_weights_from_remote_instance(
|
|||||||
seed_instance_ip,
|
seed_instance_ip,
|
||||||
seed_instance_service_port,
|
seed_instance_service_port,
|
||||||
seed_instance_group_base_port,
|
seed_instance_group_base_port,
|
||||||
|
remote_instance_loader_backend,
|
||||||
):
|
):
|
||||||
print(
|
print(
|
||||||
f"Testing model: {model_name} tp_size: {tp_size}, dp_size: {dp_size} backend: {backends}"
|
f"Testing model: {model_name} tp_size: {tp_size}, dp_size: {dp_size} backend: {backends} remote_instance_loader_backend: {remote_instance_loader_backend}"
|
||||||
)
|
)
|
||||||
param_queue = mp.Queue()
|
param_queue = mp.Queue()
|
||||||
results = {}
|
results = {}
|
||||||
@@ -276,6 +288,7 @@ def test_load_weights_from_remote_instance(
|
|||||||
seed_instance_group_base_port,
|
seed_instance_group_base_port,
|
||||||
event_seed_ready,
|
event_seed_ready,
|
||||||
event_dst_ready_list,
|
event_dst_ready_list,
|
||||||
|
remote_instance_loader_backend,
|
||||||
),
|
),
|
||||||
nprocs=1 + dp_size,
|
nprocs=1 + dp_size,
|
||||||
join=False,
|
join=False,
|
||||||
@@ -340,14 +353,42 @@ class TestLoadWeightsFromRemoteInstance(CustomTestCase):
|
|||||||
# test_suits : tp, dp, model_name, backend, dst_instance_id
|
# test_suits : tp, dp, model_name, backend, dst_instance_id
|
||||||
if is_in_ci():
|
if is_in_ci():
|
||||||
mode = random.choice(["Engine", "Server"])
|
mode = random.choice(["Engine", "Server"])
|
||||||
|
remote_instance_loader_backend = random.choice(["nccl", "transfer_engine"])
|
||||||
test_suits = [
|
test_suits = [
|
||||||
(1, 1, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, [mode]),
|
(
|
||||||
|
1,
|
||||||
|
1,
|
||||||
|
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
|
||||||
|
[mode],
|
||||||
|
remote_instance_loader_backend,
|
||||||
|
),
|
||||||
]
|
]
|
||||||
else:
|
else:
|
||||||
test_suits = [
|
test_suits = [
|
||||||
(1, 1, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, ["Engine"]),
|
(1, 1, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, ["Engine"], "nccl"),
|
||||||
(1, 1, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, ["Sever"]),
|
(1, 1, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, ["Server"], "nccl"),
|
||||||
(2, 2, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, ["Engine", "Server"]),
|
(2, 2, DEFAULT_SMALL_MODEL_NAME_FOR_TEST, ["Engine", "Server"], "nccl"),
|
||||||
|
(
|
||||||
|
1,
|
||||||
|
1,
|
||||||
|
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
|
||||||
|
["Engine"],
|
||||||
|
"transfer_engine",
|
||||||
|
),
|
||||||
|
(
|
||||||
|
1,
|
||||||
|
1,
|
||||||
|
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
|
||||||
|
["Server"],
|
||||||
|
"transfer_engine",
|
||||||
|
),
|
||||||
|
(
|
||||||
|
2,
|
||||||
|
2,
|
||||||
|
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
|
||||||
|
["Engine", "Server"],
|
||||||
|
"transfer_engine",
|
||||||
|
),
|
||||||
]
|
]
|
||||||
|
|
||||||
truncate_size = 10
|
truncate_size = 10
|
||||||
@@ -365,7 +406,13 @@ class TestLoadWeightsFromRemoteInstance(CustomTestCase):
|
|||||||
"model.norm.weight",
|
"model.norm.weight",
|
||||||
]
|
]
|
||||||
|
|
||||||
for tp_size, dp_size, model_name, backends in test_suits:
|
for (
|
||||||
|
tp_size,
|
||||||
|
dp_size,
|
||||||
|
model_name,
|
||||||
|
backends,
|
||||||
|
remote_instance_loader_backend,
|
||||||
|
) in test_suits:
|
||||||
test_load_weights_from_remote_instance(
|
test_load_weights_from_remote_instance(
|
||||||
tp_size,
|
tp_size,
|
||||||
dp_size,
|
dp_size,
|
||||||
@@ -376,6 +423,7 @@ class TestLoadWeightsFromRemoteInstance(CustomTestCase):
|
|||||||
"127.0.0.1",
|
"127.0.0.1",
|
||||||
DEFAULT_PORT_FOR_SRT_TEST_RUNNER + 1000,
|
DEFAULT_PORT_FOR_SRT_TEST_RUNNER + 1000,
|
||||||
60000,
|
60000,
|
||||||
|
remote_instance_loader_backend,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user