Make the functions in logits_processor.py and sampler.py more modular (#17885)
This commit is contained in:
@@ -126,52 +126,9 @@ class Sampler(nn.Module):
|
||||
probs = logits
|
||||
del logits
|
||||
|
||||
if can_sample_directly_from_probs:
|
||||
# when we don't need top-k, top-p, or min-p sampling, we can directly sample from the probs
|
||||
batch_next_token_ids = sampling_from_probs_torch(
|
||||
probs,
|
||||
sampling_seed=sampling_info.sampling_seed,
|
||||
positions=positions,
|
||||
)
|
||||
else:
|
||||
if get_global_server_args().sampling_backend == "flashinfer":
|
||||
if sampling_info.need_min_p_sampling:
|
||||
probs = top_k_renorm_prob(probs, sampling_info.top_ks)
|
||||
probs = top_p_renorm_prob(probs, sampling_info.top_ps)
|
||||
batch_next_token_ids = min_p_sampling_from_probs(
|
||||
probs, sampling_info.min_ps
|
||||
)
|
||||
else:
|
||||
batch_next_token_ids = top_k_top_p_sampling_from_probs(
|
||||
probs.contiguous(),
|
||||
sampling_info.top_ks,
|
||||
sampling_info.top_ps,
|
||||
filter_apply_order="joint",
|
||||
check_nan=self.use_nan_detection,
|
||||
)
|
||||
elif get_global_server_args().sampling_backend == "pytorch":
|
||||
# A slower fallback implementation with torch native operations.
|
||||
batch_next_token_ids = top_k_top_p_min_p_sampling_from_probs_torch(
|
||||
probs,
|
||||
sampling_info.top_ks,
|
||||
sampling_info.top_ps,
|
||||
sampling_info.min_ps,
|
||||
sampling_info.need_min_p_sampling,
|
||||
sampling_info.sampling_seed,
|
||||
positions,
|
||||
)
|
||||
elif get_global_server_args().sampling_backend == "ascend":
|
||||
batch_next_token_ids = top_k_top_p_min_p_sampling_from_probs_ascend(
|
||||
probs,
|
||||
sampling_info.top_ks,
|
||||
sampling_info.top_ps,
|
||||
sampling_info.min_ps,
|
||||
sampling_info.need_min_p_sampling,
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid sampling backend: {get_global_server_args().sampling_backend}"
|
||||
)
|
||||
batch_next_token_ids = self._sample_from_probs(
|
||||
probs, sampling_info, positions, can_sample_directly_from_probs
|
||||
)
|
||||
|
||||
if return_logprob:
|
||||
if get_global_server_args().rl_on_policy_target is not None:
|
||||
@@ -188,23 +145,77 @@ class Sampler(nn.Module):
|
||||
|
||||
# Attach logprobs to logits_output (in-place modification)
|
||||
if return_logprob:
|
||||
if any(x > 0 for x in top_logprobs_nums):
|
||||
(
|
||||
logits_output.next_token_top_logprobs_val,
|
||||
logits_output.next_token_top_logprobs_idx,
|
||||
) = get_top_logprobs(logprobs, top_logprobs_nums)
|
||||
|
||||
if any(x is not None for x in token_ids_logprobs):
|
||||
(
|
||||
logits_output.next_token_token_ids_logprobs_val,
|
||||
logits_output.next_token_token_ids_logprobs_idx,
|
||||
) = get_token_ids_logprobs(logprobs, token_ids_logprobs)
|
||||
|
||||
logits_output.next_token_logprobs = logprobs[
|
||||
torch.arange(len(batch_next_token_ids), device=sampling_info.device),
|
||||
self._attach_logprobs_to_output(
|
||||
logits_output,
|
||||
logprobs,
|
||||
top_logprobs_nums,
|
||||
token_ids_logprobs,
|
||||
sampling_info,
|
||||
batch_next_token_ids,
|
||||
]
|
||||
)
|
||||
|
||||
self._sync_token_ids_across_tp(batch_next_token_ids, sampling_info)
|
||||
|
||||
return batch_next_token_ids
|
||||
|
||||
def _sample_from_probs(
|
||||
self,
|
||||
probs: torch.Tensor,
|
||||
sampling_info: SamplingBatchInfo,
|
||||
positions: torch.Tensor,
|
||||
can_sample_directly_from_probs: bool,
|
||||
) -> torch.Tensor:
|
||||
if can_sample_directly_from_probs:
|
||||
# when we don't need top-k, top-p, or min-p sampling, we can directly sample from the probs
|
||||
batch_next_token_ids = sampling_from_probs_torch(
|
||||
probs,
|
||||
sampling_seed=sampling_info.sampling_seed,
|
||||
positions=positions,
|
||||
)
|
||||
else:
|
||||
if get_global_server_args().sampling_backend == "flashinfer":
|
||||
if sampling_info.need_min_p_sampling:
|
||||
probs = top_k_renorm_prob(probs, sampling_info.top_ks)
|
||||
probs = top_p_renorm_prob(probs, sampling_info.top_ps)
|
||||
batch_next_token_ids = min_p_sampling_from_probs(
|
||||
probs, sampling_info.min_ps
|
||||
)
|
||||
else:
|
||||
batch_next_token_ids = top_k_top_p_sampling_from_probs(
|
||||
probs.contiguous(),
|
||||
sampling_info.top_ks,
|
||||
sampling_info.top_ps,
|
||||
filter_apply_order="joint",
|
||||
check_nan=self.use_nan_detection,
|
||||
)
|
||||
elif get_global_server_args().sampling_backend == "pytorch":
|
||||
# A slower fallback implementation with torch native operations.
|
||||
batch_next_token_ids = top_k_top_p_min_p_sampling_from_probs_torch(
|
||||
probs,
|
||||
sampling_info.top_ks,
|
||||
sampling_info.top_ps,
|
||||
sampling_info.min_ps,
|
||||
sampling_info.need_min_p_sampling,
|
||||
sampling_info.sampling_seed,
|
||||
positions,
|
||||
)
|
||||
elif get_global_server_args().sampling_backend == "ascend":
|
||||
batch_next_token_ids = top_k_top_p_min_p_sampling_from_probs_ascend(
|
||||
probs,
|
||||
sampling_info.top_ks,
|
||||
sampling_info.top_ps,
|
||||
sampling_info.min_ps,
|
||||
sampling_info.need_min_p_sampling,
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid sampling backend: {get_global_server_args().sampling_backend}"
|
||||
)
|
||||
return batch_next_token_ids
|
||||
|
||||
def _sync_token_ids_across_tp(
|
||||
self, batch_next_token_ids: torch.Tensor, sampling_info: SamplingBatchInfo
|
||||
):
|
||||
if SYNC_TOKEN_IDS_ACROSS_TP or sampling_info.grammars:
|
||||
# For performance reasons, SGLang does not sync the final token IDs across TP ranks by default.
|
||||
# This saves one all-reduce, but the correctness of this approach depends on the determinism of several operators:
|
||||
@@ -219,7 +230,32 @@ class Sampler(nn.Module):
|
||||
group=self.tp_sync_group,
|
||||
)
|
||||
|
||||
return batch_next_token_ids
|
||||
def _attach_logprobs_to_output(
|
||||
self,
|
||||
logits_output: LogitsProcessorOutput,
|
||||
logprobs: torch.Tensor,
|
||||
top_logprobs_nums: List[int],
|
||||
token_ids_logprobs: List[List[int]],
|
||||
sampling_info: SamplingBatchInfo,
|
||||
batch_next_token_ids: torch.Tensor,
|
||||
):
|
||||
# Attach logprobs to logits_output (in-place modification)
|
||||
if any(x > 0 for x in top_logprobs_nums):
|
||||
(
|
||||
logits_output.next_token_top_logprobs_val,
|
||||
logits_output.next_token_top_logprobs_idx,
|
||||
) = get_top_logprobs(logprobs, top_logprobs_nums)
|
||||
|
||||
if any(x is not None for x in token_ids_logprobs):
|
||||
(
|
||||
logits_output.next_token_token_ids_logprobs_val,
|
||||
logits_output.next_token_token_ids_logprobs_idx,
|
||||
) = get_token_ids_logprobs(logprobs, token_ids_logprobs)
|
||||
|
||||
logits_output.next_token_logprobs = logprobs[
|
||||
torch.arange(len(batch_next_token_ids), device=sampling_info.device),
|
||||
batch_next_token_ids,
|
||||
]
|
||||
|
||||
def compute_logprobs_only(
|
||||
self,
|
||||
|
||||
Reference in New Issue
Block a user