Optimize delayed sample and mrope position computation (#32637)
This commit is contained in:
@@ -1148,6 +1148,10 @@ class Envs:
|
||||
SGLANG_OPT_USE_FUSED_STORE_CACHE = EnvBool(True)
|
||||
SGLANG_OPT_USE_JIT_NORM = EnvBool(True)
|
||||
SGLANG_OPT_USE_MULTI_STREAM_OVERLAP = EnvBool(True)
|
||||
# Force delay_sample_func for all overlap decode (not just grammar mode),
|
||||
# allowing CPU result processing to overlap with subsequent forward computation
|
||||
# and reducing the impact of sampling overhead on the critical path.
|
||||
SGLANG_ENABLE_DELAY_SAMPLE = EnvBool(False)
|
||||
|
||||
# CUDA graph
|
||||
SGLANG_PREP_IN_CUDA_GRAPH = EnvBool(True)
|
||||
|
||||
@@ -3887,6 +3887,12 @@ class Scheduler(
|
||||
assert _batch_result is batch_result
|
||||
# Delay-sample is non-spec only; relays the sampled bonus tokens.
|
||||
self._relay_forward_payload(batch_result.future_indices, batch_result)
|
||||
|
||||
# Run device-to-host copy on a separate stream to avoid blocking the
|
||||
# forward stream. The copy waits for the sampled result and can overlap
|
||||
# with subsequent forward computation.
|
||||
self.copy_stream.wait_stream(self.forward_stream)
|
||||
with self.copy_stream_ctx:
|
||||
batch_result.copy_to_cpu(
|
||||
return_logprob=cur_batch.return_logprob,
|
||||
return_hidden_states=cur_batch.return_hidden_states,
|
||||
|
||||
@@ -23,6 +23,7 @@ import torch
|
||||
|
||||
from sglang.srt.distributed import get_pp_group, get_world_group
|
||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.managers.io_struct import (
|
||||
DestroyWeightsUpdateGroupReqInput,
|
||||
GetWeightsByNameReqInput,
|
||||
@@ -610,10 +611,18 @@ class TpModelWorker(BaseTpWorker):
|
||||
# Skip sampling; spec_v2 worker fires its own publish post-verify.
|
||||
return batch_result
|
||||
|
||||
# Delay sampling only for normal generation requests.
|
||||
# Keep the existing grammar behavior unchanged.
|
||||
if (
|
||||
self.enable_overlap
|
||||
and not self.enable_spec
|
||||
and forward_batch.sampling_info.grammars is not None
|
||||
and (
|
||||
forward_batch.sampling_info.grammars is not None
|
||||
or (
|
||||
envs.SGLANG_ENABLE_DELAY_SAMPLE.get()
|
||||
and not forward_batch.is_prefill_only
|
||||
)
|
||||
)
|
||||
):
|
||||
|
||||
def sample_batch_func():
|
||||
|
||||
@@ -1161,57 +1161,129 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
return mrope_positions
|
||||
|
||||
def _compute_mrope_positions(self, model_runner: ModelRunner, batch: ScheduleBatch):
|
||||
# batch_size * [3 * seq_len]
|
||||
batch_size = self.seq_lens_cpu.shape[0]
|
||||
mrope_positions_list = [[]] * batch_size
|
||||
# mrope_positions shape: [3, total_seq_len]
|
||||
# The first dimension corresponds to temporal, height and width positions.
|
||||
forward_mode = self.forward_mode
|
||||
|
||||
if forward_mode.is_decode():
|
||||
self._compute_mrope_positions_decode(model_runner, batch)
|
||||
return
|
||||
|
||||
if not forward_mode.is_extend(include_draft_extend_v2=True):
|
||||
# Should not reach here: spec/draft_extend_v2 modes go through
|
||||
# compute_spec_mrope_positions in init_new.
|
||||
raise RuntimeError(
|
||||
f"_compute_mrope_positions called with unsupported forward_mode: {forward_mode}"
|
||||
)
|
||||
|
||||
self._compute_mrope_positions_extend(model_runner, batch)
|
||||
|
||||
def _compute_mrope_positions_decode(
|
||||
self, model_runner: ModelRunner, batch: ScheduleBatch
|
||||
):
|
||||
seq_lens_cpu = self.seq_lens_cpu
|
||||
batch_size = seq_lens_cpu.shape[0]
|
||||
mm_inputs = batch.multimodal_inputs
|
||||
rl_on_policy_target = get_exec().deterministic.rl_on_policy_target
|
||||
seq_lens_int64 = self.seq_lens.to(torch.int64)
|
||||
|
||||
# Some multimodal models (e.g. image generation models) provide
|
||||
# precomputed MRoPE positions. In this case, positions cannot be
|
||||
# reconstructed only from mrope_position_delta and need the original
|
||||
# per-request computation path.
|
||||
has_precomputed_mrope = any(
|
||||
mm is not None
|
||||
and mm.mrope_positions is not None
|
||||
and mm.mrope_positions.shape[1] >= seq_lens_cpu[i]
|
||||
for i, mm in enumerate(mm_inputs)
|
||||
)
|
||||
has_multimodal_input = any(mm is not None for mm in mm_inputs)
|
||||
|
||||
if rl_on_policy_target is not None or not has_multimodal_input:
|
||||
# Text-only
|
||||
positions_1d = seq_lens_int64 - 1
|
||||
self.mrope_positions = positions_1d.unsqueeze(0).repeat(3, 1)
|
||||
return
|
||||
|
||||
if not has_precomputed_mrope:
|
||||
# Compute all requests together to avoid per-request Python loops.
|
||||
deltas_list = [0] * batch_size
|
||||
for i, mm in enumerate(mm_inputs):
|
||||
deltas_list[i] = mm.mrope_position_delta.item() if mm is not None else 0
|
||||
deltas = torch.tensor(
|
||||
deltas_list, dtype=torch.int64, device=model_runner.device
|
||||
)
|
||||
positions_1d = (deltas - 1) + seq_lens_int64
|
||||
self.mrope_positions = positions_1d.unsqueeze(0).repeat(3, 1)
|
||||
return
|
||||
|
||||
# Fallback path for models with precomputed spatial MRoPE positions.
|
||||
# These positions cannot be derived from mrope_position_delta.
|
||||
mrope_positions_list = [None] * batch_size
|
||||
for batch_idx in range(batch_size):
|
||||
mm_input = batch.multimodal_inputs[batch_idx]
|
||||
if self.forward_mode.is_decode():
|
||||
# 3 * N
|
||||
if mm_input is None or rl_on_policy_target is not None:
|
||||
mrope_positions_list[batch_idx] = torch.full(
|
||||
(3, 1),
|
||||
self.seq_lens_cpu[batch_idx] - 1,
|
||||
dtype=torch.int64,
|
||||
)
|
||||
else:
|
||||
mrope_positions = self._expand_mrope_from_input(
|
||||
mm_input, self.seq_lens_cpu[batch_idx]
|
||||
)
|
||||
mrope_positions_list[batch_idx] = mrope_positions
|
||||
elif self.forward_mode.is_extend(include_draft_extend_v2=True):
|
||||
extend_seq_len, extend_prefix_len = (
|
||||
batch.extend_lens[batch_idx],
|
||||
batch.prefix_lens[batch_idx],
|
||||
mm_input = mm_inputs[batch_idx]
|
||||
if mm_input is None:
|
||||
# text only
|
||||
mrope_positions = torch.full(
|
||||
(3, 1),
|
||||
seq_lens_cpu[batch_idx] - 1,
|
||||
dtype=torch.int64,
|
||||
)
|
||||
if mm_input is None or rl_on_policy_target is not None:
|
||||
# text only
|
||||
mrope_positions = torch.tensor(
|
||||
[
|
||||
[
|
||||
pos
|
||||
for pos in range(
|
||||
extend_prefix_len,
|
||||
extend_prefix_len + extend_seq_len,
|
||||
)
|
||||
]
|
||||
]
|
||||
* 3
|
||||
)
|
||||
else:
|
||||
mrope_positions = mm_input.mrope_positions[
|
||||
:,
|
||||
extend_prefix_len : extend_prefix_len + extend_seq_len,
|
||||
]
|
||||
if mrope_positions.numel() == 0:
|
||||
mrope_positions = self._expand_mrope_from_input(
|
||||
mm_input, self.seq_lens_cpu[batch_idx]
|
||||
)
|
||||
mrope_positions_list[batch_idx] = mrope_positions
|
||||
else:
|
||||
mrope_positions = self._expand_mrope_from_input(
|
||||
mm_input, seq_lens_cpu[batch_idx]
|
||||
)
|
||||
mrope_positions_list[batch_idx] = mrope_positions
|
||||
|
||||
self.mrope_positions = torch.cat(
|
||||
[pos for pos in mrope_positions_list],
|
||||
mrope_positions_list,
|
||||
dim=1,
|
||||
).to(
|
||||
dtype=torch.int64,
|
||||
device=model_runner.device,
|
||||
non_blocking=True,
|
||||
)
|
||||
|
||||
def _compute_mrope_positions_extend(
|
||||
self, model_runner: ModelRunner, batch: ScheduleBatch
|
||||
):
|
||||
seq_lens_cpu = self.seq_lens_cpu
|
||||
batch_size = seq_lens_cpu.shape[0]
|
||||
mm_inputs = batch.multimodal_inputs
|
||||
rl_on_policy_target = get_exec().deterministic.rl_on_policy_target
|
||||
extend_lens = batch.extend_lens
|
||||
prefix_lens = batch.prefix_lens
|
||||
|
||||
# Optimize text-only branch with torch.arange
|
||||
mrope_positions_list = [None] * batch_size
|
||||
for batch_idx in range(batch_size):
|
||||
mm_input = mm_inputs[batch_idx]
|
||||
extend_seq_len = extend_lens[batch_idx]
|
||||
extend_prefix_len = prefix_lens[batch_idx]
|
||||
if mm_input is None or rl_on_policy_target is not None:
|
||||
# text only
|
||||
mrope_positions = (
|
||||
torch.arange(
|
||||
extend_prefix_len,
|
||||
extend_prefix_len + extend_seq_len,
|
||||
dtype=torch.int64,
|
||||
)
|
||||
.unsqueeze(0)
|
||||
.repeat(3, 1)
|
||||
)
|
||||
else:
|
||||
mrope_positions = mm_input.mrope_positions[
|
||||
:,
|
||||
extend_prefix_len : extend_prefix_len + extend_seq_len,
|
||||
]
|
||||
if mrope_positions.numel() == 0:
|
||||
mrope_positions = self._expand_mrope_from_input(
|
||||
mm_input, seq_lens_cpu[batch_idx]
|
||||
)
|
||||
mrope_positions_list[batch_idx] = mrope_positions
|
||||
|
||||
self.mrope_positions = torch.cat(
|
||||
mrope_positions_list,
|
||||
dim=1,
|
||||
).to(dtype=torch.int64, device=model_runner.device, non_blocking=True)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user