support non-disturbing remote-instance-weight-loader (#13125)
Signed-off-by: Anqi Shen <amy.saq@antgroup.com>
This commit is contained in:
@@ -0,0 +1,36 @@
|
|||||||
|
# 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. |
|
||||||
|
|
||||||
|
### NCCL as backend
|
||||||
|
|
||||||
|
```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
|
||||||
|
|
||||||
|
```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
|
||||||
|
```
|
||||||
@@ -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
|
||||||
|
|||||||
@@ -127,13 +127,18 @@ class Engine(EngineBase):
|
|||||||
atexit.register(self.shutdown)
|
atexit.register(self.shutdown)
|
||||||
|
|
||||||
# Launch subprocesses
|
# Launch subprocesses
|
||||||
tokenizer_manager, template_manager, scheduler_info, port_args = (
|
(
|
||||||
_launch_subprocesses(server_args=server_args)
|
tokenizer_manager,
|
||||||
)
|
template_manager,
|
||||||
|
scheduler_info,
|
||||||
|
port_args,
|
||||||
|
remote_instance_transfer_engine_info,
|
||||||
|
) = _launch_subprocesses(server_args=server_args)
|
||||||
self.tokenizer_manager = tokenizer_manager
|
self.tokenizer_manager = tokenizer_manager
|
||||||
self.template_manager = template_manager
|
self.template_manager = template_manager
|
||||||
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 = remote_instance_transfer_engine_info
|
||||||
|
|
||||||
# Initialize ZMQ sockets
|
# Initialize ZMQ sockets
|
||||||
context = zmq.Context(2)
|
context = zmq.Context(2)
|
||||||
@@ -910,6 +915,7 @@ def _launch_subprocesses(
|
|||||||
|
|
||||||
# Wait for the model to finish loading
|
# Wait for the model to finish loading
|
||||||
scheduler_infos = []
|
scheduler_infos = []
|
||||||
|
remote_instance_transfer_engine_info = {}
|
||||||
for i in range(len(scheduler_pipe_readers)):
|
for i in range(len(scheduler_pipe_readers)):
|
||||||
try:
|
try:
|
||||||
data = scheduler_pipe_readers[i].recv()
|
data = scheduler_pipe_readers[i].recv()
|
||||||
@@ -926,9 +932,24 @@ def _launch_subprocesses(
|
|||||||
"Initialization failed. Please see the error messages above."
|
"Initialization failed. Please see the error messages above."
|
||||||
)
|
)
|
||||||
scheduler_infos.append(data)
|
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
|
# Assume all schedulers have the same scheduler_info
|
||||||
scheduler_info = scheduler_infos[0]
|
scheduler_info = scheduler_infos[0]
|
||||||
tokenizer_manager.max_req_input_len = scheduler_info["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_info,
|
||||||
|
port_args,
|
||||||
|
remote_instance_transfer_engine_info,
|
||||||
|
)
|
||||||
|
|||||||
@@ -144,6 +144,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
|
||||||
@@ -813,6 +822,24 @@ 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)
|
||||||
|
|
||||||
|
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
|
||||||
@@ -1386,15 +1413,20 @@ 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 = (
|
(
|
||||||
_launch_subprocesses(server_args=server_args)
|
tokenizer_manager,
|
||||||
)
|
template_manager,
|
||||||
|
scheduler_info,
|
||||||
|
port_args,
|
||||||
|
remote_instance_transfer_engine_info,
|
||||||
|
) = _launch_subprocesses(server_args=server_args)
|
||||||
|
|
||||||
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,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -2573,6 +2573,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:
|
||||||
"""
|
"""
|
||||||
@@ -2686,6 +2689,22 @@ def run_scheduler_process(
|
|||||||
pp_rank,
|
pp_rank,
|
||||||
dp_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(
|
pipe_writer.send(
|
||||||
{
|
{
|
||||||
"status": "ready",
|
"status": "ready",
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -65,6 +65,7 @@ from sglang.srt.distributed import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.distributed.parallel_state import monkey_patch_vllm_parallel_state
|
from sglang.srt.distributed.parallel_state import monkey_patch_vllm_parallel_state
|
||||||
from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager
|
from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.eplb.eplb_manager import EPLBManager
|
from sglang.srt.eplb.eplb_manager import EPLBManager
|
||||||
from sglang.srt.eplb.expert_distribution import (
|
from sglang.srt.eplb.expert_distribution import (
|
||||||
ExpertDistributionRecorder,
|
ExpertDistributionRecorder,
|
||||||
@@ -135,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_v2,
|
||||||
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
|
||||||
@@ -157,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 +322,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 +400,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_support_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 +443,16 @@ class ModelRunner:
|
|||||||
self.sampler = Sampler()
|
self.sampler = Sampler()
|
||||||
self.load_model()
|
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
|
# 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 +567,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 +801,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 +811,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 +840,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_v2,
|
||||||
|
)
|
||||||
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)
|
||||||
@@ -1974,6 +1979,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
|
||||||
@@ -1992,16 +1998,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:
|
||||||
@@ -2009,9 +2018,45 @@ 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_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()
|
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
|
||||||
@@ -2062,6 +2107,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,13 +1,21 @@
|
|||||||
# SPDX-License-Identifier: Apache-2.0
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
|
||||||
|
import enum
|
||||||
import logging
|
import logging
|
||||||
|
import time
|
||||||
from typing import List
|
from typing import List
|
||||||
|
|
||||||
import requests
|
import requests
|
||||||
|
import torch
|
||||||
|
|
||||||
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 +75,110 @@ 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
|
||||||
|
|
||||||
|
|
||||||
|
# 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
|
||||||
|
|||||||
@@ -17,6 +17,8 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import argparse
|
import argparse
|
||||||
import dataclasses
|
import dataclasses
|
||||||
|
import importlib
|
||||||
|
import importlib.util
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
@@ -593,6 +595,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_support_transfer_engine: bool = False
|
||||||
|
|
||||||
# For PD-Multiplexing
|
# For PD-Multiplexing
|
||||||
enable_pdmux: bool = False
|
enable_pdmux: bool = False
|
||||||
@@ -704,6 +708,9 @@ class ServerArgs:
|
|||||||
# Handle elastic expert parallelism.
|
# Handle elastic expert parallelism.
|
||||||
self._handle_elastic_ep()
|
self._handle_elastic_ep()
|
||||||
|
|
||||||
|
# Handle remote instance weight loader.
|
||||||
|
self._handle_remote_instance_weight_loader_support_transfer_engine()
|
||||||
|
|
||||||
def _handle_deprecated_args(self):
|
def _handle_deprecated_args(self):
|
||||||
# handle deprecated tool call parsers
|
# handle deprecated tool call parsers
|
||||||
deprecated_tool_call_parsers = {"qwen25": "qwen", "glm45": "glm"}
|
deprecated_tool_call_parsers = {"qwen25": "qwen", "glm45": "glm"}
|
||||||
@@ -1882,8 +1889,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 (
|
||||||
|
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"
|
self.load_format = "auto"
|
||||||
|
|
||||||
def _handle_disaggregation(self):
|
def _handle_disaggregation(self):
|
||||||
@@ -2118,6 +2143,20 @@ 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):
|
||||||
|
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
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def add_cli_args(parser: argparse.ArgumentParser):
|
def add_cli_args(parser: argparse.ArgumentParser):
|
||||||
|
|
||||||
@@ -3947,6 +3986,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-support-transfer-engine",
|
||||||
|
action="store_true",
|
||||||
|
help="Enable transfer engine support for remote instance weight loader.",
|
||||||
|
)
|
||||||
|
|
||||||
# For PD-Multiplexing
|
# For PD-Multiplexing
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
|
|||||||
@@ -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,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -159,6 +161,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 +189,7 @@ 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,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
host, _, port = DEFAULT_URL_FOR_TEST.rpartition(":")
|
host, _, port = DEFAULT_URL_FOR_TEST.rpartition(":")
|
||||||
@@ -213,6 +217,8 @@ 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,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
@@ -250,9 +256,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 +283,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 +348,36 @@ 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, ["Sever"], "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, ["Sever"], "transfer_engine"),
|
||||||
|
(
|
||||||
|
2,
|
||||||
|
2,
|
||||||
|
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
|
||||||
|
["Engine", "Server"],
|
||||||
|
"transfer_engine",
|
||||||
|
),
|
||||||
]
|
]
|
||||||
|
|
||||||
truncate_size = 10
|
truncate_size = 10
|
||||||
@@ -365,7 +395,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 +412,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