Add per-rank staggered weight loading for improved TP I/O concurrency (#26937)
Signed-off-by: zhaoshang <zhaoshangsjtu@linux.alibaba.com> Co-authored-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
co-authored by
Shangming Cai
parent
1ebc7438ac
commit
3e681d7fff
+3
-1
@@ -163,7 +163,9 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
f"Cannot find any model weights with `{model_name_or_path}`"
|
f"Cannot find any model weights with `{model_name_or_path}`"
|
||||||
)
|
)
|
||||||
|
|
||||||
if envs.SGLANG_SORT_WEIGHT_FILES.get():
|
# Sort weight files when SGLANG_SORT_WEIGHT_FILES >= 0 (default).
|
||||||
|
# Staggering is not applicable to text-encoder loading (no TP split).
|
||||||
|
if envs.SGLANG_SORT_WEIGHT_FILES.get() >= 0:
|
||||||
hf_weights_files.sort()
|
hf_weights_files.sort()
|
||||||
|
|
||||||
return hf_folder, hf_weights_files, use_safetensors
|
return hf_folder, hf_weights_files, use_safetensors
|
||||||
|
|||||||
@@ -196,7 +196,14 @@ class Envs:
|
|||||||
|
|
||||||
# Model & File Download
|
# Model & File Download
|
||||||
SGLANG_USE_MODELSCOPE = EnvBool(False)
|
SGLANG_USE_MODELSCOPE = EnvBool(False)
|
||||||
SGLANG_SORT_WEIGHT_FILES = EnvBool(False)
|
# Controls weight-file ordering for load-time I/O optimization.
|
||||||
|
# -1 : no sorting, no staggering; preserves original file order.
|
||||||
|
# 0 : sort files only; maximizes ordering but may reduce cross-rank I/O concurrency.
|
||||||
|
# k>0: sort files and stagger per-rank order with factor k.
|
||||||
|
# Files are processed in groups of (tp_size * k), and rank r starts each
|
||||||
|
# group at offset (r * k), improving multi-rank I/O concurrency while
|
||||||
|
# keeping access relatively ordered.
|
||||||
|
SGLANG_SORT_WEIGHT_FILES = EnvInt(0)
|
||||||
SGLANG_DISABLED_MODEL_ARCHS = EnvTuple(tuple())
|
SGLANG_DISABLED_MODEL_ARCHS = EnvTuple(tuple())
|
||||||
SGLANG_PREFETCH_BLOCK_SIZE_MB = EnvInt(16)
|
SGLANG_PREFETCH_BLOCK_SIZE_MB = EnvInt(16)
|
||||||
SGLANG_GEMMA_OUT_OF_PLACE_POSITION_MUTATION = EnvBool(False)
|
SGLANG_GEMMA_OUT_OF_PLACE_POSITION_MUTATION = EnvBool(False)
|
||||||
|
|||||||
@@ -494,8 +494,22 @@ class DefaultModelLoader(BaseModelLoader):
|
|||||||
f"Cannot find any model weights with `{model_name_or_path}`"
|
f"Cannot find any model weights with `{model_name_or_path}`"
|
||||||
)
|
)
|
||||||
|
|
||||||
if envs.SGLANG_SORT_WEIGHT_FILES.get():
|
# Sort and optionally stagger weight files (see SGLANG_SORT_WEIGHT_FILES).
|
||||||
|
# k=-1: no sort; k=0: sort only; k>0: sort + stagger by (tp_rank*k).
|
||||||
|
k = envs.SGLANG_SORT_WEIGHT_FILES.get()
|
||||||
|
if k >= 0:
|
||||||
hf_weights_files.sort()
|
hf_weights_files.sort()
|
||||||
|
if k > 0:
|
||||||
|
tp_size = get_tensor_model_parallel_world_size()
|
||||||
|
if tp_size > 1:
|
||||||
|
tp_rank = get_tensor_model_parallel_rank()
|
||||||
|
group_size = tp_size * k
|
||||||
|
staggered: List[str] = []
|
||||||
|
for i in range(0, len(hf_weights_files), group_size):
|
||||||
|
group = hf_weights_files[i : i + group_size]
|
||||||
|
n = len(group)
|
||||||
|
staggered.extend(group[(j + tp_rank * k) % n] for j in range(n))
|
||||||
|
hf_weights_files = staggered
|
||||||
|
|
||||||
return hf_folder, hf_weights_files, use_safetensors
|
return hf_folder, hf_weights_files, use_safetensors
|
||||||
|
|
||||||
|
|||||||
@@ -913,13 +913,13 @@ def safetensors_weights_iterator(
|
|||||||
not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0
|
not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0
|
||||||
)
|
)
|
||||||
|
|
||||||
sorted_files = sorted(hf_weights_files)
|
|
||||||
|
|
||||||
if prefetch and not disable_mmap:
|
if prefetch and not disable_mmap:
|
||||||
_prefetch_all_checkpoints(sorted_files, num_threads=prefetch_num_threads)
|
_prefetch_all_checkpoints(
|
||||||
|
sorted(hf_weights_files), num_threads=prefetch_num_threads
|
||||||
|
)
|
||||||
|
|
||||||
for st_file in tqdm(
|
for st_file in tqdm(
|
||||||
sorted_files,
|
hf_weights_files,
|
||||||
desc="Loading safetensors checkpoint shards",
|
desc="Loading safetensors checkpoint shards",
|
||||||
disable=not enable_tqdm,
|
disable=not enable_tqdm,
|
||||||
bar_format=BAR_FORMAT,
|
bar_format=BAR_FORMAT,
|
||||||
@@ -1051,9 +1051,10 @@ def buffered_multi_thread_safetensors_weights_iterator(
|
|||||||
max_workers loading concurrently + 1 prefetched and ready to yield.
|
max_workers loading concurrently + 1 prefetched and ready to yield.
|
||||||
Peak CPU RAM ≈ (max_workers + 2) × shard_file_size.
|
Peak CPU RAM ≈ (max_workers + 2) × shard_file_size.
|
||||||
"""
|
"""
|
||||||
sorted_files = sorted(hf_weights_files)
|
|
||||||
if prefetch and not disable_mmap:
|
if prefetch and not disable_mmap:
|
||||||
_prefetch_all_checkpoints(sorted_files, num_threads=prefetch_num_threads)
|
_prefetch_all_checkpoints(
|
||||||
|
sorted(hf_weights_files), num_threads=prefetch_num_threads
|
||||||
|
)
|
||||||
enable_tqdm = (
|
enable_tqdm = (
|
||||||
not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0
|
not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0
|
||||||
)
|
)
|
||||||
@@ -1071,7 +1072,7 @@ def buffered_multi_thread_safetensors_weights_iterator(
|
|||||||
buffer_size = max_workers + 1
|
buffer_size = max_workers + 1
|
||||||
|
|
||||||
with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
|
with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
|
||||||
file_iter = iter(sorted_files)
|
file_iter = iter(hf_weights_files)
|
||||||
pending: collections.deque = collections.deque()
|
pending: collections.deque = collections.deque()
|
||||||
|
|
||||||
# Seed the buffer.
|
# Seed the buffer.
|
||||||
|
|||||||
Reference in New Issue
Block a user