Extract init_torch_distributed and refactor into functions (#31152)
This commit is contained in:
@@ -0,0 +1,291 @@
|
|||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import time
|
||||||
|
from typing import List, Optional
|
||||||
|
|
||||||
|
import msgspec
|
||||||
|
import torch
|
||||||
|
import torch.distributed as dist
|
||||||
|
|
||||||
|
from sglang.srt.configs.model_config import ModelConfig
|
||||||
|
from sglang.srt.distributed import (
|
||||||
|
get_default_distributed_backend,
|
||||||
|
get_pp_group,
|
||||||
|
get_tp_group,
|
||||||
|
get_world_group,
|
||||||
|
init_distributed_environment,
|
||||||
|
initialize_model_parallel,
|
||||||
|
set_custom_all_reduce,
|
||||||
|
set_mscclpp_all_reduce,
|
||||||
|
set_torch_symm_mem_all_reduce,
|
||||||
|
)
|
||||||
|
from sglang.srt.environ import envs
|
||||||
|
from sglang.srt.layers.dp_attention import initialize_dp_attention
|
||||||
|
from sglang.srt.platforms import current_platform
|
||||||
|
from sglang.srt.runtime_context import get_parallel
|
||||||
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
from sglang.srt.utils import (
|
||||||
|
cpu_has_amx_support,
|
||||||
|
get_available_gpu_memory,
|
||||||
|
is_host_cpu_arm64,
|
||||||
|
is_npu,
|
||||||
|
monkey_patch_p2p_access_check,
|
||||||
|
)
|
||||||
|
from sglang.srt.utils.network import NetworkAddress
|
||||||
|
from sglang.srt.utils.patch_torch import register_sgl_tp_rank
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_is_cpu_amx_available = cpu_has_amx_support()
|
||||||
|
_is_cpu_arm64 = is_host_cpu_arm64()
|
||||||
|
|
||||||
|
|
||||||
|
class TorchDistributedResult(msgspec.Struct, frozen=True, kw_only=True):
|
||||||
|
tp_group: object
|
||||||
|
pp_group: object
|
||||||
|
attention_tp_group: object
|
||||||
|
pre_model_load_memory: float
|
||||||
|
|
||||||
|
|
||||||
|
def init_torch_distributed(
|
||||||
|
*,
|
||||||
|
server_args: ServerArgs,
|
||||||
|
model_config: ModelConfig,
|
||||||
|
device: str,
|
||||||
|
gpu_id: int,
|
||||||
|
tp_rank: int,
|
||||||
|
tp_size: int,
|
||||||
|
pp_rank: int,
|
||||||
|
pp_size: int,
|
||||||
|
dp_size: int,
|
||||||
|
attn_cp_size: int,
|
||||||
|
moe_ep_size: int,
|
||||||
|
moe_dp_size: int,
|
||||||
|
dcp_size: int,
|
||||||
|
dist_port: int,
|
||||||
|
is_draft_worker: bool,
|
||||||
|
local_omp_cpuid: Optional[List[int]],
|
||||||
|
):
|
||||||
|
tic = time.perf_counter()
|
||||||
|
logger.info("Init torch distributed begin.")
|
||||||
|
|
||||||
|
try:
|
||||||
|
torch.get_device_module(device).set_device(gpu_id)
|
||||||
|
except Exception:
|
||||||
|
logger.warning(
|
||||||
|
f"Context: {device=} {gpu_id=} {os.environ.get('CUDA_VISIBLE_DEVICES')=} {tp_rank=} {tp_size=}"
|
||||||
|
)
|
||||||
|
raise
|
||||||
|
|
||||||
|
backend = _resolve_backend(device=device, server_args=server_args, gpu_id=gpu_id)
|
||||||
|
|
||||||
|
before_avail_memory = get_available_gpu_memory(device, gpu_id)
|
||||||
|
if not server_args.enable_p2p_check:
|
||||||
|
monkey_patch_p2p_access_check()
|
||||||
|
|
||||||
|
dist_init_method = _resolve_dist_init_method(
|
||||||
|
server_args=server_args, dist_port=dist_port
|
||||||
|
)
|
||||||
|
_set_all_reduce_flags(server_args=server_args)
|
||||||
|
|
||||||
|
if not is_draft_worker:
|
||||||
|
if device == "cpu":
|
||||||
|
_init_cpu_threads_env(
|
||||||
|
tp_size=tp_size, tp_rank=tp_rank, local_omp_cpuid=local_omp_cpuid
|
||||||
|
)
|
||||||
|
|
||||||
|
# Only initialize the distributed environment on the target model worker.
|
||||||
|
_init_parallel_groups(
|
||||||
|
backend=backend,
|
||||||
|
dist_init_method=dist_init_method,
|
||||||
|
server_args=server_args,
|
||||||
|
model_config=model_config,
|
||||||
|
gpu_id=gpu_id,
|
||||||
|
tp_rank=tp_rank,
|
||||||
|
tp_size=tp_size,
|
||||||
|
pp_rank=pp_rank,
|
||||||
|
pp_size=pp_size,
|
||||||
|
dp_size=dp_size,
|
||||||
|
attn_cp_size=attn_cp_size,
|
||||||
|
moe_ep_size=moe_ep_size,
|
||||||
|
moe_dp_size=moe_dp_size,
|
||||||
|
dcp_size=dcp_size,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Pre-warm NCCL/RCCL/HCCL to eliminate cold-start latency in first request
|
||||||
|
# Controlled by --pre-warm-nccl flag (default: enabled on AMD GPUs)
|
||||||
|
if server_args.pre_warm_nccl and (
|
||||||
|
tp_size > 1 or pp_size > 1 or moe_ep_size > 1
|
||||||
|
):
|
||||||
|
_prewarm_nccl(tp_size=tp_size, pp_size=pp_size, moe_ep_size=moe_ep_size)
|
||||||
|
|
||||||
|
pre_model_load_memory = get_available_gpu_memory(
|
||||||
|
device,
|
||||||
|
gpu_id,
|
||||||
|
distributed=get_world_group().world_size > 1,
|
||||||
|
cpu_group=get_world_group().cpu_group,
|
||||||
|
)
|
||||||
|
tp_group = get_tp_group()
|
||||||
|
pp_group = get_pp_group()
|
||||||
|
attention_tp_group = get_parallel().attn_tp_group
|
||||||
|
|
||||||
|
# Check memory for tensor parallelism
|
||||||
|
local_gpu_memory = get_available_gpu_memory(device, gpu_id)
|
||||||
|
if tp_size > 1 and not is_draft_worker:
|
||||||
|
_check_tp_memory_balance(
|
||||||
|
pre_model_load_memory=pre_model_load_memory,
|
||||||
|
local_gpu_memory=local_gpu_memory,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
f"Init torch distributed ends. elapsed={time.perf_counter() - tic:.2f} s, "
|
||||||
|
f"mem usage={(before_avail_memory - local_gpu_memory):.2f} GB"
|
||||||
|
)
|
||||||
|
return TorchDistributedResult(
|
||||||
|
tp_group=tp_group,
|
||||||
|
pp_group=pp_group,
|
||||||
|
attention_tp_group=attention_tp_group,
|
||||||
|
pre_model_load_memory=pre_model_load_memory,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_backend(*, device: str, server_args: ServerArgs, gpu_id: int) -> str:
|
||||||
|
backend = get_default_distributed_backend(device)
|
||||||
|
if device == "cuda" and server_args.elastic_ep_backend == "mooncake":
|
||||||
|
backend = "mooncake"
|
||||||
|
if server_args.mooncake_ib_device:
|
||||||
|
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
|
||||||
|
get_ib_devices_for_gpu,
|
||||||
|
)
|
||||||
|
|
||||||
|
ib_device_for_gpu = get_ib_devices_for_gpu(
|
||||||
|
server_args.mooncake_ib_device, gpu_id
|
||||||
|
)
|
||||||
|
mooncake_ib_device = (
|
||||||
|
ib_device_for_gpu.split(",") if ib_device_for_gpu else []
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
from mooncake import ep as mooncake_ep
|
||||||
|
|
||||||
|
mooncake_ep.set_device_filter(mooncake_ib_device)
|
||||||
|
except:
|
||||||
|
pass # A warning will be raised in `init_distributed_environment`
|
||||||
|
return backend
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_dist_init_method(*, server_args: ServerArgs, dist_port: int) -> str:
|
||||||
|
# Allow external orchestrators (e.g. trainpi) to override the distributed
|
||||||
|
# init method. When set to "env://", torch uses MASTER_ADDR/MASTER_PORT
|
||||||
|
# env-vars and an externally-created TCPStore, completely avoiding port
|
||||||
|
# conflicts with intra-host collocation.
|
||||||
|
dist_init_method_override = envs.SGLANG_DISTRIBUTED_INIT_METHOD_OVERRIDE.get()
|
||||||
|
if dist_init_method_override:
|
||||||
|
dist_init_method = dist_init_method_override
|
||||||
|
elif server_args.dist_init_addr:
|
||||||
|
na = NetworkAddress.parse(server_args.dist_init_addr)
|
||||||
|
dist_init_method = na.to_tcp()
|
||||||
|
else:
|
||||||
|
dist_init_method = NetworkAddress(
|
||||||
|
server_args.host or "127.0.0.1", dist_port
|
||||||
|
).to_tcp()
|
||||||
|
return dist_init_method
|
||||||
|
|
||||||
|
|
||||||
|
def _set_all_reduce_flags(*, server_args: ServerArgs) -> None:
|
||||||
|
set_custom_all_reduce(not server_args.disable_custom_all_reduce)
|
||||||
|
set_mscclpp_all_reduce(server_args.enable_mscclpp)
|
||||||
|
set_torch_symm_mem_all_reduce(server_args.enable_torch_symm_mem)
|
||||||
|
|
||||||
|
|
||||||
|
def _init_cpu_threads_env(
|
||||||
|
*, tp_size: int, tp_rank: int, local_omp_cpuid: Optional[List[int]]
|
||||||
|
) -> None:
|
||||||
|
if _is_cpu_amx_available or _is_cpu_arm64:
|
||||||
|
# Bind OpenMP threads to CPU cores
|
||||||
|
torch.ops.sgl_kernel.init_cpu_threads_env(local_omp_cpuid)
|
||||||
|
|
||||||
|
# Set local size to hint SGLang to use shared memory based AllReduce
|
||||||
|
os.environ["LOCAL_SIZE"] = str(tp_size)
|
||||||
|
torch.ops.sgl_kernel.initialize(tp_size, tp_rank)
|
||||||
|
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"init_cpu_threads_env and shared memory based AllReduce is disabled, only intel amx backend and arm64 are supported"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _init_parallel_groups(
|
||||||
|
*,
|
||||||
|
backend: str,
|
||||||
|
dist_init_method: str,
|
||||||
|
server_args: ServerArgs,
|
||||||
|
model_config: ModelConfig,
|
||||||
|
gpu_id: int,
|
||||||
|
tp_rank: int,
|
||||||
|
tp_size: int,
|
||||||
|
pp_rank: int,
|
||||||
|
pp_size: int,
|
||||||
|
dp_size: int,
|
||||||
|
attn_cp_size: int,
|
||||||
|
moe_ep_size: int,
|
||||||
|
moe_dp_size: int,
|
||||||
|
dcp_size: int,
|
||||||
|
) -> None:
|
||||||
|
init_distributed_environment(
|
||||||
|
backend=backend,
|
||||||
|
world_size=tp_size * pp_size,
|
||||||
|
rank=tp_size * pp_rank + tp_rank,
|
||||||
|
local_rank=gpu_id,
|
||||||
|
distributed_init_method=dist_init_method,
|
||||||
|
timeout=server_args.dist_timeout,
|
||||||
|
moe_a2a_backend=server_args.moe_a2a_backend,
|
||||||
|
recovered_rank=server_args.elastic_ep_rejoin,
|
||||||
|
)
|
||||||
|
initialize_model_parallel(
|
||||||
|
tensor_model_parallel_size=tp_size,
|
||||||
|
attention_data_parallel_size=dp_size,
|
||||||
|
pipeline_model_parallel_size=pp_size,
|
||||||
|
expert_model_parallel_size=moe_ep_size,
|
||||||
|
attention_context_model_parallel_size=attn_cp_size,
|
||||||
|
moe_data_model_parallel_size=moe_dp_size,
|
||||||
|
decode_context_parallel_size=dcp_size,
|
||||||
|
duplicate_tp_group=server_args.enable_pdmux,
|
||||||
|
enable_symm_mem=server_args.enable_symm_mem,
|
||||||
|
recovered_rank=server_args.elastic_ep_rejoin,
|
||||||
|
)
|
||||||
|
initialize_dp_attention(
|
||||||
|
server_args=server_args,
|
||||||
|
model_config=model_config,
|
||||||
|
)
|
||||||
|
if is_npu():
|
||||||
|
register_sgl_tp_rank(gpu_id)
|
||||||
|
|
||||||
|
|
||||||
|
def _prewarm_nccl(*, tp_size: int, pp_size: int, moe_ep_size: int) -> None:
|
||||||
|
warmup_start = time.perf_counter()
|
||||||
|
tp_group_handle = get_tp_group().device_group
|
||||||
|
|
||||||
|
# Single warmup all_reduce to initialize NCCL/RCCL/HCCL communicator
|
||||||
|
warmup_tensor = torch.zeros(1, device=torch.cuda.current_device())
|
||||||
|
dist.all_reduce(warmup_tensor, group=tp_group_handle)
|
||||||
|
current_platform.synchronize()
|
||||||
|
|
||||||
|
warmup_elapsed = time.perf_counter() - warmup_start
|
||||||
|
logger.info(
|
||||||
|
f"NCCL/RCCL/HCCL warmup completed in {warmup_elapsed:.3f}s "
|
||||||
|
f"(tp_size={tp_size}, pp_size={pp_size}, ep_size={moe_ep_size})"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _check_tp_memory_balance(
|
||||||
|
*, pre_model_load_memory: float, local_gpu_memory: float
|
||||||
|
) -> None:
|
||||||
|
if pre_model_load_memory < local_gpu_memory * 0.9:
|
||||||
|
msg = "The memory capacity is unbalanced. Some GPUs may be occupied by other processes. "
|
||||||
|
msg += (
|
||||||
|
f"{pre_model_load_memory=}, {local_gpu_memory=}, {local_gpu_memory * 0.9=}"
|
||||||
|
)
|
||||||
|
if envs.SGLANG_ENABLE_TP_MEMORY_INBALANCE_CHECK.get():
|
||||||
|
raise RuntimeError(msg)
|
||||||
|
else:
|
||||||
|
logger.warning(msg)
|
||||||
@@ -47,15 +47,9 @@ from sglang.srt.debug_utils.tensor_dump_forward_hook import (
|
|||||||
register_forward_hook_for_model,
|
register_forward_hook_for_model,
|
||||||
)
|
)
|
||||||
from sglang.srt.distributed import (
|
from sglang.srt.distributed import (
|
||||||
get_default_distributed_backend,
|
bootstrap,
|
||||||
get_pp_group,
|
|
||||||
get_tp_group,
|
get_tp_group,
|
||||||
get_world_group,
|
get_world_group,
|
||||||
init_distributed_environment,
|
|
||||||
initialize_model_parallel,
|
|
||||||
set_custom_all_reduce,
|
|
||||||
set_mscclpp_all_reduce,
|
|
||||||
set_torch_symm_mem_all_reduce,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||||
prealloc_symmetric_memory_pool,
|
prealloc_symmetric_memory_pool,
|
||||||
@@ -105,9 +99,6 @@ from sglang.srt.layers.attention.tbo_backend import TboAttnBackend
|
|||||||
from sglang.srt.layers.cp.utils import (
|
from sglang.srt.layers.cp.utils import (
|
||||||
get_cp_strategy,
|
get_cp_strategy,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.dp_attention import (
|
|
||||||
initialize_dp_attention,
|
|
||||||
)
|
|
||||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||||
from sglang.srt.layers.moe.hash_topk import HashTopK
|
from sglang.srt.layers.moe.hash_topk import HashTopK
|
||||||
from sglang.srt.layers.moe.topk import TopK
|
from sglang.srt.layers.moe.topk import TopK
|
||||||
@@ -164,7 +155,7 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.model_loader.utils import resolve_language_model
|
from sglang.srt.model_loader.utils import resolve_language_model
|
||||||
from sglang.srt.platforms import current_platform
|
from sglang.srt.platforms import current_platform
|
||||||
from sglang.srt.runtime_context import get_flags, get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_flags, get_server_args
|
||||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||||
from sglang.srt.server_args import ( # noqa: F401 (re-export)
|
from sglang.srt.server_args import ( # noqa: F401 (re-export)
|
||||||
CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS,
|
CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS,
|
||||||
@@ -197,7 +188,6 @@ from sglang.srt.utils import (
|
|||||||
is_host_cpu_arm64,
|
is_host_cpu_arm64,
|
||||||
is_npu,
|
is_npu,
|
||||||
log_info_on_rank0,
|
log_info_on_rank0,
|
||||||
monkey_patch_p2p_access_check,
|
|
||||||
numa_utils,
|
numa_utils,
|
||||||
require_gathered_buffer,
|
require_gathered_buffer,
|
||||||
reserve_rope_cache_for_long_sequences,
|
reserve_rope_cache_for_long_sequences,
|
||||||
@@ -212,7 +202,6 @@ from sglang.srt.utils.offloader import (
|
|||||||
get_offloader,
|
get_offloader,
|
||||||
set_offloader,
|
set_offloader,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.patch_torch import register_sgl_tp_rank
|
|
||||||
from sglang.srt.utils.profile_utils import build_step_span_name
|
from sglang.srt.utils.profile_utils import build_step_span_name
|
||||||
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
||||||
from sglang.srt.utils.weight_checker import WeightChecker
|
from sglang.srt.utils.weight_checker import WeightChecker
|
||||||
@@ -475,7 +464,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
|
|
||||||
# Get available memory before model loading.
|
# Get available memory before model loading.
|
||||||
# Stored for later use by alloc_memory_pool().
|
# Stored for later use by alloc_memory_pool().
|
||||||
self.pre_model_load_memory = self.init_torch_distributed()
|
self.init_torch_distributed()
|
||||||
|
|
||||||
# Initialize MooncakeTransferEngine
|
# Initialize MooncakeTransferEngine
|
||||||
self.init_shared_mooncake_transfer_engine()
|
self.init_shared_mooncake_transfer_engine()
|
||||||
@@ -1109,150 +1098,28 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def init_torch_distributed(self):
|
def init_torch_distributed(self):
|
||||||
tic = time.perf_counter()
|
result = bootstrap.init_torch_distributed(
|
||||||
logger.info("Init torch distributed begin.")
|
server_args=self.server_args,
|
||||||
|
model_config=self.model_config,
|
||||||
try:
|
device=self.device,
|
||||||
torch.get_device_module(self.device).set_device(self.gpu_id)
|
gpu_id=self.gpu_id,
|
||||||
except Exception:
|
tp_rank=self.tp_rank,
|
||||||
logger.warning(
|
tp_size=self.tp_size,
|
||||||
f"Context: {self.device=} {self.gpu_id=} {os.environ.get('CUDA_VISIBLE_DEVICES')=} {self.tp_rank=} {self.tp_size=}"
|
pp_rank=self.pp_rank,
|
||||||
)
|
pp_size=self.pp_size,
|
||||||
raise
|
dp_size=self.attn_dp_size,
|
||||||
|
attn_cp_size=self.attn_cp_size,
|
||||||
backend = get_default_distributed_backend(self.device)
|
moe_ep_size=self.moe_ep_size,
|
||||||
if self.device == "cuda" and self.server_args.elastic_ep_backend == "mooncake":
|
moe_dp_size=self.moe_dp_size,
|
||||||
backend = "mooncake"
|
dcp_size=self.dcp_size,
|
||||||
if self.server_args.mooncake_ib_device:
|
dist_port=self.dist_port,
|
||||||
from sglang.srt.distributed.device_communicators.mooncake_transfer_engine import (
|
is_draft_worker=self.is_draft_worker,
|
||||||
get_ib_devices_for_gpu,
|
local_omp_cpuid=self.local_omp_cpuid if self.device == "cpu" else None,
|
||||||
)
|
|
||||||
|
|
||||||
ib_device_for_gpu = get_ib_devices_for_gpu(
|
|
||||||
self.server_args.mooncake_ib_device, self.gpu_id
|
|
||||||
)
|
|
||||||
mooncake_ib_device = (
|
|
||||||
ib_device_for_gpu.split(",") if ib_device_for_gpu else []
|
|
||||||
)
|
|
||||||
try:
|
|
||||||
from mooncake import ep as mooncake_ep
|
|
||||||
|
|
||||||
mooncake_ep.set_device_filter(mooncake_ib_device)
|
|
||||||
except:
|
|
||||||
pass # A warning will be raised in `init_distributed_environment`
|
|
||||||
|
|
||||||
before_avail_memory = get_available_gpu_memory(self.device, self.gpu_id)
|
|
||||||
if not self.server_args.enable_p2p_check:
|
|
||||||
monkey_patch_p2p_access_check()
|
|
||||||
|
|
||||||
# Allow external orchestrators (e.g. trainpi) to override the distributed
|
|
||||||
# init method. When set to "env://", torch uses MASTER_ADDR/MASTER_PORT
|
|
||||||
# env-vars and an externally-created TCPStore, completely avoiding port
|
|
||||||
# conflicts with intra-host collocation.
|
|
||||||
dist_init_method_override = envs.SGLANG_DISTRIBUTED_INIT_METHOD_OVERRIDE.get()
|
|
||||||
if dist_init_method_override:
|
|
||||||
dist_init_method = dist_init_method_override
|
|
||||||
elif self.server_args.dist_init_addr:
|
|
||||||
na = NetworkAddress.parse(self.server_args.dist_init_addr)
|
|
||||||
dist_init_method = na.to_tcp()
|
|
||||||
else:
|
|
||||||
dist_init_method = NetworkAddress(
|
|
||||||
self.server_args.host or "127.0.0.1", self.dist_port
|
|
||||||
).to_tcp()
|
|
||||||
set_custom_all_reduce(not self.server_args.disable_custom_all_reduce)
|
|
||||||
set_mscclpp_all_reduce(self.server_args.enable_mscclpp)
|
|
||||||
set_torch_symm_mem_all_reduce(self.server_args.enable_torch_symm_mem)
|
|
||||||
|
|
||||||
if not self.is_draft_worker:
|
|
||||||
if self.device == "cpu":
|
|
||||||
if _is_cpu_amx_available or _is_cpu_arm64:
|
|
||||||
# Bind OpenMP threads to CPU cores
|
|
||||||
torch.ops.sgl_kernel.init_cpu_threads_env(self.local_omp_cpuid)
|
|
||||||
|
|
||||||
# Set local size to hint SGLang to use shared memory based AllReduce
|
|
||||||
os.environ["LOCAL_SIZE"] = str(self.tp_size)
|
|
||||||
torch.ops.sgl_kernel.initialize(self.tp_size, self.tp_rank)
|
|
||||||
|
|
||||||
else:
|
|
||||||
logger.warning(
|
|
||||||
"init_cpu_threads_env and shared memory based AllReduce is disabled, only intel amx backend and arm64 are supported"
|
|
||||||
)
|
|
||||||
|
|
||||||
# Only initialize the distributed environment on the target model worker.
|
|
||||||
init_distributed_environment(
|
|
||||||
backend=backend,
|
|
||||||
world_size=self.tp_size * self.pp_size,
|
|
||||||
rank=self.tp_size * self.pp_rank + self.tp_rank,
|
|
||||||
local_rank=self.gpu_id,
|
|
||||||
distributed_init_method=dist_init_method,
|
|
||||||
timeout=self.server_args.dist_timeout,
|
|
||||||
moe_a2a_backend=self.server_args.moe_a2a_backend,
|
|
||||||
recovered_rank=self.server_args.elastic_ep_rejoin,
|
|
||||||
)
|
|
||||||
initialize_model_parallel(
|
|
||||||
tensor_model_parallel_size=self.tp_size,
|
|
||||||
attention_data_parallel_size=self.attn_dp_size,
|
|
||||||
pipeline_model_parallel_size=self.pp_size,
|
|
||||||
expert_model_parallel_size=self.moe_ep_size,
|
|
||||||
attention_context_model_parallel_size=self.attn_cp_size,
|
|
||||||
moe_data_model_parallel_size=self.moe_dp_size,
|
|
||||||
decode_context_parallel_size=self.dcp_size,
|
|
||||||
duplicate_tp_group=self.server_args.enable_pdmux,
|
|
||||||
enable_symm_mem=self.server_args.enable_symm_mem,
|
|
||||||
recovered_rank=self.server_args.elastic_ep_rejoin,
|
|
||||||
)
|
|
||||||
initialize_dp_attention(
|
|
||||||
server_args=self.server_args,
|
|
||||||
model_config=self.model_config,
|
|
||||||
)
|
|
||||||
if is_npu():
|
|
||||||
register_sgl_tp_rank(self.gpu_id)
|
|
||||||
|
|
||||||
# Pre-warm NCCL/RCCL/HCCL to eliminate cold-start latency in first request
|
|
||||||
# Controlled by --pre-warm-nccl flag (default: enabled on AMD GPUs)
|
|
||||||
if self.server_args.pre_warm_nccl and (
|
|
||||||
self.tp_size > 1 or self.pp_size > 1 or self.moe_ep_size > 1
|
|
||||||
):
|
|
||||||
warmup_start = time.perf_counter()
|
|
||||||
tp_group_handle = get_tp_group().device_group
|
|
||||||
|
|
||||||
# Single warmup all_reduce to initialize NCCL/RCCL/HCCL communicator
|
|
||||||
warmup_tensor = torch.zeros(1, device=torch.cuda.current_device())
|
|
||||||
dist.all_reduce(warmup_tensor, group=tp_group_handle)
|
|
||||||
current_platform.synchronize()
|
|
||||||
|
|
||||||
warmup_elapsed = time.perf_counter() - warmup_start
|
|
||||||
logger.info(
|
|
||||||
f"NCCL/RCCL/HCCL warmup completed in {warmup_elapsed:.3f}s "
|
|
||||||
f"(tp_size={self.tp_size}, pp_size={self.pp_size}, ep_size={self.moe_ep_size})"
|
|
||||||
)
|
|
||||||
|
|
||||||
pre_model_load_memory = get_available_gpu_memory(
|
|
||||||
self.device,
|
|
||||||
self.gpu_id,
|
|
||||||
distributed=get_world_group().world_size > 1,
|
|
||||||
cpu_group=get_world_group().cpu_group,
|
|
||||||
)
|
)
|
||||||
self.tp_group = get_tp_group()
|
self.tp_group = result.tp_group
|
||||||
self.pp_group = get_pp_group()
|
self.pp_group = result.pp_group
|
||||||
self.attention_tp_group = get_parallel().attn_tp_group
|
self.attention_tp_group = result.attention_tp_group
|
||||||
|
self.pre_model_load_memory = result.pre_model_load_memory
|
||||||
# Check memory for tensor parallelism
|
|
||||||
local_gpu_memory = get_available_gpu_memory(self.device, self.gpu_id)
|
|
||||||
if self.tp_size > 1 and not self.is_draft_worker:
|
|
||||||
if pre_model_load_memory < local_gpu_memory * 0.9:
|
|
||||||
msg = "The memory capacity is unbalanced. Some GPUs may be occupied by other processes. "
|
|
||||||
msg += f"{pre_model_load_memory=}, {local_gpu_memory=}, {local_gpu_memory * 0.9=}"
|
|
||||||
if envs.SGLANG_ENABLE_TP_MEMORY_INBALANCE_CHECK.get():
|
|
||||||
raise RuntimeError(msg)
|
|
||||||
else:
|
|
||||||
logger.warning(msg)
|
|
||||||
|
|
||||||
logger.info(
|
|
||||||
f"Init torch distributed ends. elapsed={time.perf_counter() - tic:.2f} s, "
|
|
||||||
f"mem usage={(before_avail_memory - local_gpu_memory):.2f} GB"
|
|
||||||
)
|
|
||||||
return pre_model_load_memory
|
|
||||||
|
|
||||||
def init_shared_mooncake_transfer_engine(self):
|
def init_shared_mooncake_transfer_engine(self):
|
||||||
"""
|
"""
|
||||||
|
|||||||
Reference in New Issue
Block a user