[PP] Fix dynamic chunking strategy for PP (#15372)
Signed-off-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
@@ -10,6 +10,7 @@ from typing import TYPE_CHECKING, Dict, List, Optional, Tuple
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
import torch.distributed
|
import torch.distributed
|
||||||
|
from tqdm import tqdm
|
||||||
|
|
||||||
from sglang.srt.disaggregation.base.conn import KVPoll
|
from sglang.srt.disaggregation.base.conn import KVPoll
|
||||||
from sglang.srt.disaggregation.utils import DisaggregationMode, poll_and_all_reduce
|
from sglang.srt.disaggregation.utils import DisaggregationMode, poll_and_all_reduce
|
||||||
@@ -532,13 +533,11 @@ class SchedulerPPMixin:
|
|||||||
latencies: List[float] = []
|
latencies: List[float] = []
|
||||||
|
|
||||||
if self.pp_group.is_first_rank:
|
if self.pp_group.is_first_rank:
|
||||||
logger.info("Profiling prefill latency for dynamic chunk sizing...")
|
|
||||||
|
|
||||||
# Create requests with different lengths: base_chunk_size // (2**i) for i in range(10)
|
|
||||||
input_ids_list = []
|
input_ids_list = []
|
||||||
for i in range(32):
|
for i in range(128):
|
||||||
chunk_size = self.chunked_prefill_size - i * (
|
chunk_size = int(
|
||||||
self.chunked_prefill_size // 32
|
self.chunked_prefill_size * 1.25
|
||||||
|
- i * (self.chunked_prefill_size * 1.25 // 128)
|
||||||
)
|
)
|
||||||
if chunk_size <= 0:
|
if chunk_size <= 0:
|
||||||
break
|
break
|
||||||
@@ -551,9 +550,13 @@ class SchedulerPPMixin:
|
|||||||
temperature=0,
|
temperature=0,
|
||||||
max_new_tokens=1,
|
max_new_tokens=1,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Create and profile requests
|
# Create and profile requests
|
||||||
for i, input_ids in enumerate(input_ids_list):
|
for i, input_ids in enumerate(
|
||||||
|
tqdm(
|
||||||
|
input_ids_list,
|
||||||
|
desc="Profiling prefill latency for dynamic chunking",
|
||||||
|
)
|
||||||
|
):
|
||||||
req = Req(
|
req = Req(
|
||||||
rid=str(i),
|
rid=str(i),
|
||||||
origin_input_text="",
|
origin_input_text="",
|
||||||
@@ -1338,8 +1341,8 @@ class ChunkSizePredictor:
|
|||||||
)
|
)
|
||||||
calculated_chunk_size = int(smoothed_chunk_size)
|
calculated_chunk_size = int(smoothed_chunk_size)
|
||||||
|
|
||||||
# Align to page_size (round down to nearest multiple)
|
# Align to page_size (minimum alignment size is 64)
|
||||||
alignment_size = max(page_size, 1)
|
alignment_size = max(page_size, 64)
|
||||||
dynamic_chunk_size = (calculated_chunk_size // alignment_size) * alignment_size
|
dynamic_chunk_size = (calculated_chunk_size // alignment_size) * alignment_size
|
||||||
|
|
||||||
# Ensure aligned size is at least alignment_size
|
# Ensure aligned size is at least alignment_size
|
||||||
|
|||||||
Reference in New Issue
Block a user