[1/n] Fix hanging during DeepGemm Warmup (#14493)
This commit is contained in:
@@ -14,7 +14,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.deep_gemm_wrapper.configurer import ENABLE_JIT_DEEPGEMM
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils import ceil_div, get_bool_env_var
|
||||
from sglang.srt.utils import ceil_div, get_available_gpu_memory, get_bool_env_var
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -132,8 +132,35 @@ def _compile_deep_gemm_one_type_all(
|
||||
m_alignment = deep_gemm.get_mk_alignment_for_contiguous_layout()
|
||||
m_list = sorted(list(set(m for m in m_list if m % m_alignment == 0)))
|
||||
|
||||
# Here the precompilation is only run on the first rank, so gpu_id should be 0
|
||||
memory_budget = get_available_gpu_memory(device="cuda", gpu_id=0)
|
||||
|
||||
# If the memory budget is less memory requirement, we need to reduce max_m to avoid out of memory, which might further cause hanging during warmup
|
||||
max_m = max(m_list)
|
||||
required_memory = _BaseWarmupExecutor.get_memory_requirement(
|
||||
kernel_type, max_m=max_m, n=n, k=k, num_groups=num_groups
|
||||
)
|
||||
logger.info(
|
||||
f"Required memory for warmup: {required_memory}GB, Available memory: {memory_budget}GB"
|
||||
)
|
||||
if memory_budget < required_memory:
|
||||
# TODO: Maybe compute the max_m based on the memory budget
|
||||
while (
|
||||
_BaseWarmupExecutor.get_memory_requirement(
|
||||
kernel_type, max_m=max_m, n=n, k=k, num_groups=num_groups
|
||||
)
|
||||
> memory_budget
|
||||
and max_m > 4096
|
||||
):
|
||||
max_m = max_m // 2
|
||||
logger.warning(
|
||||
f"Available memory {memory_budget}GB is less than required memory {required_memory}GB for warmup, reducing max_m to {max_m} to avoid out of memory"
|
||||
)
|
||||
m_list = [m for m in m_list if m <= max_m]
|
||||
|
||||
# Need some methods to estimate needed memory for warmup
|
||||
executor = _BaseWarmupExecutor.create(
|
||||
kernel_type, max_m=max(m_list), n=n, k=k, num_groups=num_groups
|
||||
kernel_type, max_m=max_m, n=n, k=k, num_groups=num_groups
|
||||
)
|
||||
|
||||
old_compile_mode = deep_gemm.get_compile_mode()
|
||||
@@ -161,6 +188,26 @@ class _BaseWarmupExecutor:
|
||||
DeepGemmKernelType.GROUPED_GEMM_NT_F8F8BF16_MASKED: _GroupedMaskedWarmupExecutor,
|
||||
}[kernel_type](**kwargs)
|
||||
|
||||
@staticmethod
|
||||
def get_memory_requirement(
|
||||
kernel_type: DeepGemmKernelType, max_m: int, n: int, k: int, num_groups: int
|
||||
) -> int:
|
||||
# Return the required memory space in GB for warmup executor
|
||||
_GB = 1 << 30
|
||||
if kernel_type == DeepGemmKernelType.GEMM_NT_F8F8BF16:
|
||||
return (max_m * k + n * k + max_m * n * 2) / _GB
|
||||
elif kernel_type == DeepGemmKernelType.GROUPED_GEMM_NT_F8F8BF16_CONTIG:
|
||||
return (max_m * k + num_groups * n * k + max_m * 4 + max_m * n * 2) / _GB
|
||||
elif kernel_type == DeepGemmKernelType.GROUPED_GEMM_NT_F8F8BF16_MASKED:
|
||||
return (
|
||||
num_groups * max_m * k
|
||||
+ num_groups * n * k
|
||||
+ num_groups * 4
|
||||
+ num_groups * max_m * n * 2
|
||||
) / _GB
|
||||
else:
|
||||
raise ValueError(f"Invalid kernel type: {kernel_type}")
|
||||
|
||||
def execute(self, m):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
Reference in New Issue
Block a user