qwen 3.8 rebase (#35758)
Co-authored-by: cherichy <cherichy@outlook.com> Co-authored-by: guangyunh-nv <guangyunh@nvidia.com> Co-authored-by: jiahanc <jiahanc@nvidia.com> Co-authored-by: jinyangyuan-nvidia <joyuan@nvidia.com> Co-authored-by: Cheng Hang <chang@nvidia.com> Co-authored-by: Yicheng Qiang <yqiang@nvidia.com> Co-authored-by: Sam Li <lsam@nvidia.com> Co-authored-by: Tom-Zheng <tizheng@nvidia.com> Co-authored-by: Yangmin Li <yangminl@nvidia.com> Co-authored-by: xiaoweiw-nv <xiaoweiw@nvidia.com> Co-authored-by: Zheng Li <lizheng.cs@zju.edu.cn> Co-authored-by: yizhang2077 <1109276519@qq.com> Co-authored-by: Ke Bao <ispobaoke@gmail.com> Co-authored-by: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com> Co-authored-by: Yuhao Yang <47235274+yhyang201@users.noreply.github.com> Co-authored-by: Zijie Xia <zijie.xia@radixark.ai> Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
cherichy
guangyunh-nv
jiahanc
jinyangyuan-nvidia
Cheng Hang
Yicheng Qiang
Sam Li
Tom-Zheng
Yangmin Li
xiaoweiw-nv
Zheng Li
yizhang2077
Ke Bao
Xinyuan Tong
Yuhao Yang
Zijie Xia
github-actions[bot]
parent
ca8cc101b8
commit
5f216fc33f
@@ -8,6 +8,7 @@ from torch import nn
|
||||
|
||||
from sglang.kernels.ops.sampling.murmur_hash import murmur_hash32
|
||||
from sglang.srt.distributed import get_tp_group
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
is_dp_attention_enabled,
|
||||
)
|
||||
@@ -68,6 +69,18 @@ _CUSTOM_SAMPLER_FACTORIES: Dict[str, Callable[[], "Sampler"]] = {}
|
||||
_BUILT_IN_SAMPLING_BACKENDS = {"flashinfer", "pytorch", "ascend"}
|
||||
|
||||
|
||||
def _trace_e2e_sampler(stage: str, **fields) -> None:
|
||||
if not envs.SGLANG_TRACE_SAMPLER_E2E.get():
|
||||
return
|
||||
try:
|
||||
parallel = get_parallel()
|
||||
rank = f"dp={parallel.attn_dp_rank} tp={parallel.tp_rank}"
|
||||
except Exception:
|
||||
rank = "rank=unknown"
|
||||
details = " ".join(f"{key}={value}" for key, value in fields.items())
|
||||
print(f"SGLANG_TRACE_SAMPLER_E2E {rank} stage={stage} {details}", flush=True)
|
||||
|
||||
|
||||
class Sampler(nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
@@ -119,15 +132,23 @@ class Sampler(nn.Module):
|
||||
to get the unique seed for each position.
|
||||
"""
|
||||
logits = logits_output.next_token_logits
|
||||
_trace_e2e_sampler(
|
||||
"forward_enter",
|
||||
logits_shape=tuple(logits.shape),
|
||||
all_greedy=sampling_info.is_all_greedy,
|
||||
)
|
||||
|
||||
if _is_hip and logits.shape[0] == 0:
|
||||
return torch.empty((0,), dtype=torch.int64, device=logits.device)
|
||||
|
||||
# Preprocess logits (custom processors and NaN handling)
|
||||
_trace_e2e_sampler("preprocess_enter")
|
||||
logits = self._preprocess_logits(logits, sampling_info)
|
||||
_trace_e2e_sampler("preprocess_returned")
|
||||
return_sampling_mask = any(sampling_info.return_sampling_masks or [])
|
||||
|
||||
if sampling_info.is_all_greedy:
|
||||
_trace_e2e_sampler("greedy_enter")
|
||||
if _use_aiter and not _disable_aiter_greedy_sample:
|
||||
batch_next_token_ids = torch.empty(
|
||||
logits.shape[0], device=logits.device, dtype=torch.int32
|
||||
@@ -135,6 +156,9 @@ class Sampler(nn.Module):
|
||||
_aiter_greedy_sample(batch_next_token_ids, logits)
|
||||
else:
|
||||
batch_next_token_ids = torch.argmax(logits, -1)
|
||||
_trace_e2e_sampler(
|
||||
"greedy_returned", output_shape=tuple(batch_next_token_ids.shape)
|
||||
)
|
||||
if return_sampling_mask:
|
||||
self._attach_greedy_sampling_mask_to_output(
|
||||
logits_output, sampling_info, batch_next_token_ids
|
||||
@@ -243,8 +267,11 @@ class Sampler(nn.Module):
|
||||
)
|
||||
logprob_result.write_output_to(logits_output)
|
||||
|
||||
_trace_e2e_sampler("token_sync_enter")
|
||||
self._sync_token_ids_across_tp(batch_next_token_ids, sampling_info)
|
||||
_trace_e2e_sampler("token_sync_returned")
|
||||
|
||||
_trace_e2e_sampler("forward_returned")
|
||||
return batch_next_token_ids
|
||||
|
||||
def _sample_from_probs(
|
||||
|
||||
Reference in New Issue
Block a user