diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index b4af57e72..74d5a10bc 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -680,6 +680,10 @@ class SchedulerDisaggregationPrefillMixin: if copy_done is not None: copy_done.synchronize() + auxiliary_output_starts = ( + self.batch_result_processor.snapshot_auxiliary_output_starts(batch, result) + ) + auxiliary_output = result.auxiliary_host_output if result.routed_experts_output is not None: result.routed_experts_output.finalize() result.routed_experts_output = None @@ -819,6 +823,13 @@ class SchedulerDisaggregationPrefillMixin: self.send_kv_chunk(req, last_chunk=False, end_idx=req.tmp_end_idx) req.time_stats.set_last_chunked_prefill_finish_time() + if auxiliary_output is not None: + self.batch_result_processor.consume_auxiliary_output( + batch, + auxiliary_output, + auxiliary_output_starts, + ) + can_run_cuda_graph = result.can_run_cuda_graph self.metrics_reporter.report_prefill_stats( batch=batch, diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index 2487072b6..7a678e944 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -53,6 +53,7 @@ from sglang.srt.model_executor.forward_batch_info import ( ForwardMode, ) from sglang.srt.runtime_context import get_exec, get_parallel +from sglang.srt.sampling.sampling_observer import DeviceAuxiliaryOutput from sglang.srt.utils.common import ( is_cpu, is_npu, @@ -145,6 +146,9 @@ class LogitsProcessorOutput: # They should be moved to GenerationBatchResult to keep this class clean. mm_input_embeds: Optional[torch.Tensor] = None + # Scheduler-local output copied alongside the ordinary generation result. + auxiliary_device_output: Optional[DeviceAuxiliaryOutput] = None + @dataclasses.dataclass class LogitsMetadata: diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 7c4e7166c..bbda59882 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -2180,7 +2180,7 @@ class Scheduler( ) def init_output_streamer(self) -> None: - self.output_streamer = SchedulerOutputStreamer( + self.output_streamer = self.get_output_streamer_class()( send_to_detokenizer=self.ipc_channels.send_to_detokenizer, tree_cache=self.tree_cache, ps=self.ps, @@ -2192,6 +2192,9 @@ class Scheduler( rust_server=self.rust_server, ) + def get_output_streamer_class(self) -> type[SchedulerOutputStreamer]: + return SchedulerOutputStreamer + def init_batch_result_processor(self) -> None: self.batch_result_processor = SchedulerBatchResultProcessor( is_generation=self.is_generation, @@ -3837,6 +3840,7 @@ class Scheduler( batch_result = self.tp_worker.forward_batch_split_prefill(batch) self._relay_forward_payload(batch.req_pool_indices, batch_result) batch.input_ids = None + self._copy_auxiliary_output_to_cpu(batch, batch_result) elif not batch.spec_algorithm.is_none(): # Non-overlap: drive the V2 worker synchronously (no # future_map relay / on_publish). @@ -3874,6 +3878,7 @@ class Scheduler( self._relay_forward_payload(batch.req_pool_indices, batch_result) batch.input_ids = None self.update_cache_from_scheduler(batch, batch_result) + self._copy_auxiliary_output_to_cpu(batch, batch_result) # These 2 values are needed for processing the output, but the values can be # modified by overlap schedule. So we have to copy them here so that @@ -3965,6 +3970,32 @@ class Scheduler( return self.future_map.stash(future_indices, payload) + def _copy_auxiliary_output_to_cpu( + self, + batch: ScheduleBatch, + result: GenerationBatchResult, + ) -> None: + logits_output = result.logits_output + if ( + logits_output is None + or logits_output.auxiliary_device_output is None + or result.auxiliary_host_output is not None + ): + return + # PP transports the device output to the first rank before copying it. + if self.ps.pp_size > 1: + return + if result.copy_done is not None: + raise RuntimeError( + "generation result has an uncopied auxiliary output after its " + "device-to-host copy was scheduled" + ) + result.copy_done = self.device_module.Event() + result.copy_to_cpu( + return_logprob=batch.return_logprob, + return_hidden_states=batch.return_hidden_states, + ) + def launch_batch_sample_if_needed( self, batch_result: GenerationBatchResult, cur_batch: ScheduleBatch ) -> Union[GenerationBatchResult]: diff --git a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py index 93e37a010..42245d088 100644 --- a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py @@ -40,6 +40,7 @@ from sglang.srt.runtime_context import ( mamba_track_grid, max_speculative_num_draft_tokens, ) +from sglang.srt.sampling.sampling_observer import CommittedTokens from sglang.srt.speculative.base_spec_worker import BaseSpecWorker from sglang.srt.state_capturer.indexer_topk import get_global_indexer_capturer from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer @@ -68,6 +69,7 @@ if TYPE_CHECKING: from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.observability.metrics_collector import SchedulerMetricsCollector + from sglang.srt.sampling.sampling_observer import HostAuxiliaryOutput from sglang.srt.server_args import ServerArgs logger = logging.getLogger(__name__) @@ -190,6 +192,51 @@ class SchedulerBatchResultProcessor: elem = elem.copy() req.customized_info[k].append(elem) + @staticmethod + def _visible_output_len(req: Req) -> int: + return req.finished_len if req.finished_len is not None else len(req.output_ids) + + @classmethod + def snapshot_auxiliary_output_starts( + cls, + batch: ScheduleBatch, + result: GenerationBatchResult, + ) -> Optional[List[int]]: + if result.auxiliary_host_output is None: + return None + return [cls._visible_output_len(req) for req in batch.reqs] + + @classmethod + def _build_auxiliary_commits( + cls, + batch: ScheduleBatch, + output_starts: List[int], + ) -> List[Optional[CommittedTokens]]: + commits: List[Optional[CommittedTokens]] = [] + for req, output_start in zip(batch.reqs, output_starts, strict=True): + output_end = cls._visible_output_len(req) + if output_end < output_start: + raise RuntimeError("committed output length moved backwards") + if output_end == output_start: + commits.append(None) + continue + commits.append( + CommittedTokens( + output_index=output_start, + token_ids=tuple(req.output_ids[output_start:output_end]), + ) + ) + return commits + + @classmethod + def consume_auxiliary_output( + cls, + batch: ScheduleBatch, + output: HostAuxiliaryOutput, + output_starts: List[int], + ) -> None: + output.consume(batch, cls._build_auxiliary_commits(batch, output_starts)) + def process_batch_result_prefill( self, batch: ScheduleBatch, @@ -201,6 +248,10 @@ class SchedulerBatchResultProcessor: if self.is_generation: if result.copy_done is not None: result.copy_done.synchronize() + auxiliary_output_starts = self.snapshot_auxiliary_output_starts( + batch, result + ) + auxiliary_output = result.auxiliary_host_output if result.routed_experts_output is not None: result.routed_experts_output.finalize() result.routed_experts_output = None @@ -325,6 +376,13 @@ class SchedulerBatchResultProcessor: req.time_stats.set_last_chunked_prefill_finish_time() + if auxiliary_output is not None: + self.consume_auxiliary_output( + batch, + auxiliary_output, + auxiliary_output_starts, + ) + else: # embedding or reward model if result.copy_done is not None: result.copy_done.synchronize() @@ -809,6 +867,8 @@ class SchedulerBatchResultProcessor: ): if result.copy_done is not None: result.copy_done.synchronize() + auxiliary_output_starts = self.snapshot_auxiliary_output_starts(batch, result) + auxiliary_output = result.auxiliary_host_output if result.routed_experts_output is not None: result.routed_experts_output.finalize() result.routed_experts_output = None @@ -903,6 +963,13 @@ class SchedulerBatchResultProcessor: self._accept_grammar_tokens(req, next_token_id) req.grammar.finished = req.finished() + if auxiliary_output is not None: + self.consume_auxiliary_output( + batch, + auxiliary_output, + auxiliary_output_starts, + ) + self.output_streamer.stream_output(batch.reqs, batch.return_logprob) self.token_to_kv_pool_allocator.free_group_end() diff --git a/python/sglang/srt/managers/scheduler_components/output_streamer.py b/python/sglang/srt/managers/scheduler_components/output_streamer.py index b5b303720..8fdd669e6 100644 --- a/python/sglang/srt/managers/scheduler_components/output_streamer.py +++ b/python/sglang/srt/managers/scheduler_components/output_streamer.py @@ -6,6 +6,7 @@ from typing import ( TYPE_CHECKING, Any, Callable, + ClassVar, List, Optional, ) @@ -43,6 +44,8 @@ DEFAULT_FORCE_STREAM_INTERVAL = envs.SGLANG_FORCE_STREAM_INTERVAL.get() @dataclass(kw_only=True, slots=True) class SchedulerOutputStreamer: + has_additional_customized_info: ClassVar[bool] = False + send_to_detokenizer: zmq.Socket tree_cache: BasePrefixCache ps: ParallelState @@ -57,6 +60,13 @@ class SchedulerOutputStreamer: rust_server: Optional[RustServer] = None _test_stream_output_count: int = 0 + def __post_init__(self) -> None: + if self.has_additional_customized_info and self.rust_server is not None: + raise ValueError( + "additional customized generation output is not supported by " + "Rust egress" + ) + def _get_storage_backend_type(self) -> str: """Get storage backend type from tree_cache.""" storage_backend_type = "none" @@ -170,6 +180,23 @@ class SchedulerOutputStreamer: acc.accept(req=req) self._maybe_log_time_stats(req=req) + if ( + self.has_additional_customized_info + and self.should_build_additional_customized_info() + ): + additional_customized_info = self.build_additional_customized_info( + acc.output_reqs + ) + for key, values in additional_customized_info.items(): + if key in acc.customized_info: + raise ValueError(f"duplicate customized_info key: {key}") + if len(values) != len(acc.output_reqs): + raise ValueError( + f"customized_info key {key!r} returned {len(values)} values " + f"for {len(acc.output_reqs)} requests" + ) + acc.customized_info[key] = values + # Send to detokenizer payload = acc.to_payload( dp_rank=self.ps.dp_rank, @@ -181,6 +208,21 @@ class SchedulerOutputStreamer: else: self.send_to_detokenizer.send_output(payload) + def build_additional_customized_info(self, reqs: List[Req]) -> dict[str, list]: + """Return fields aligned with the emitted requests in ``reqs``. + + Subclasses must set ``has_additional_customized_info`` to opt in. Each + returned value must have one entry per request. A matching + ``HostAuxiliaryOutput.consume`` call has already observed the tokens + committed in this scheduler step, so implementations can read buffered + per-request state here. + """ + return {} + + def should_build_additional_customized_info(self) -> bool: + """Return whether this invocation needs subclass-provided output fields.""" + return True + def _maybe_log_time_stats(self, *, req: Req) -> None: if ( req.finished() @@ -275,6 +317,7 @@ class _GenerationStreamAccumulator: default_force_stream_interval: int get_cached_tokens_details: Callable[[Req], Optional[CachedTokensDetails]] rids: list = field(default_factory=list) + output_reqs: list[Req] = field(default_factory=list) http_worker_ipcs: list = field(default_factory=list) finished_reasons: list = field(default_factory=list) decoded_texts: list = field(default_factory=list) @@ -393,6 +436,7 @@ class _GenerationStreamAccumulator: send_token_offset = req.send_token_offset send_output_token_logprobs_offset = req.send_output_token_logprobs_offset self.rids.append(req.rid) + self.output_reqs.append(req) self.finished_reasons.append( req.finished_reason.to_json() if req.finished_reason else None ) diff --git a/python/sglang/srt/managers/scheduler_pp_mixin.py b/python/sglang/srt/managers/scheduler_pp_mixin.py index 7fa52ae77..4ac4eef1c 100644 --- a/python/sglang/srt/managers/scheduler_pp_mixin.py +++ b/python/sglang/srt/managers/scheduler_pp_mixin.py @@ -23,6 +23,7 @@ from sglang.srt.layers.dp_attention import ( is_dp_attention_enabled, set_is_extend_in_batch, ) +from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.managers.overlap_utils import RelayPayload from sglang.srt.managers.schedule_batch import FINISH_ABORT, Req, ScheduleBatch from sglang.srt.managers.utils import ( @@ -37,6 +38,10 @@ from sglang.srt.model_executor.forward_batch_info import ( ) from sglang.srt.observability.req_time_stats import set_time_batch from sglang.srt.runtime_context import get_disagg, get_parallel +from sglang.srt.sampling.sampling_observer_pp import ( + add_auxiliary_output_to_pp_tensors, + pop_auxiliary_output_from_pp_tensors, +) from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.utils import DynamicGradMode, broadcast_pyobj, point_to_point_pyobj from sglang.srt.utils.common import get_device_module, is_xpu @@ -1026,6 +1031,12 @@ class SchedulerPPMixin: **tensor_dict, **logprob_dict, } + auxiliary_output = ( + result.logits_output.auxiliary_device_output + if result.logits_output is not None + else None + ) + add_auxiliary_output_to_pp_tensors(tensor_dict, auxiliary_output) return tensor_dict def _pp_send_dict_to_next_stage( @@ -1150,6 +1161,16 @@ class SchedulerPPMixin: extend_input_len_per_req, extend_logprob_start_len_per_req, ) = get_logprob_from_pp_outputs(pp_outputs) + if self.pp_group.is_first_rank: + observer = self.tp_worker.model_runner.sampling_observer + auxiliary_output = pop_auxiliary_output_from_pp_tensors( + pp_outputs.tensors, + observer, + ) + if auxiliary_output is not None: + if logits_output is None: + logits_output = LogitsProcessorOutput(next_token_logits=None) + logits_output.auxiliary_device_output = auxiliary_output next_token_ids = pp_outputs["next_token_ids"].to(torch.int64) # PP rank 0 also relays into output_tokens_buf so the next iter's # resolve_forward_inputs finds these tokens for the decode portion @@ -1166,6 +1187,7 @@ class SchedulerPPMixin: extend_logprob_start_len_per_req=extend_logprob_start_len_per_req, can_run_cuda_graph=mb_metadata.can_run_cuda_graph, ) + output_result.copy_auxiliary_output_to_cpu() return output_result def _pp_process_batch_result( diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 1341cb31d..3bef7c669 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -2252,11 +2252,13 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): i ] if customized_info is not None: - for k, v in customized_info.items(): - if k not in state.customized_info_accumulated: - state.customized_info_accumulated[k] = [] - state.customized_info_accumulated[k].extend(v[i]) - meta_info[k] = state.customized_info_accumulated[k] + self.update_request_meta_info( + meta_info, + state, + customized_info, + i, + recv_obj.finished_reasons[i], + ) # Add multimodal prompt token counts only for requests that # actually consumed them, so plain-text meta_info stays unchanged. @@ -2456,6 +2458,34 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): for s in pending_notify.values(): s.event.set() + @staticmethod + def _accumulate_request_meta_info( + meta_info: dict, + state: ReqState, + key: str, + values: list, + ) -> None: + accumulated = state.customized_info_accumulated.setdefault(key, []) + accumulated.extend(values) + meta_info[key] = accumulated + + def update_request_meta_info( + self, + meta_info: dict, + state: ReqState, + customized_info: dict, + index: int, + finish_reason: Optional[dict], + ) -> None: + """Accumulate metadata; subclasses may use finish_reason for terminal data.""" + for key, values in customized_info.items(): + self._accumulate_request_meta_info( + meta_info, + state, + key, + values[index], + ) + def add_logprob_to_meta_info( self, meta_info: dict, diff --git a/python/sglang/srt/managers/utils.py b/python/sglang/srt/managers/utils.py index 42db8b4e0..aece0c5e2 100644 --- a/python/sglang/srt/managers/utils.py +++ b/python/sglang/srt/managers/utils.py @@ -21,6 +21,7 @@ from sglang.srt.state_capturer.base import TopkCaptureOutput if TYPE_CHECKING: from sglang.srt.managers.scheduler import GenerationBatchResult + from sglang.srt.sampling.sampling_observer import HostAuxiliaryOutput from sglang.srt.speculative.eagle_info import EagleDraftInput @@ -108,6 +109,8 @@ class GenerationBatchResult: fpm_start_event: Optional[torch.cuda.Event] = None fpm_end_event: Optional[torch.cuda.Event] = None + auxiliary_host_output: Optional[HostAuxiliaryOutput] = None + @property def has_sampled_token_ids(self) -> bool: """True when this iter sampled token ids; False when none were produced @@ -170,8 +173,18 @@ class GenerationBatchResult: if holder is not None: holder.map_device_tensors(_async_d2h) + self.copy_auxiliary_output_to_cpu() + self.copy_done.record() + def copy_auxiliary_output_to_cpu(self) -> None: + if self.logits_output is None or self.auxiliary_host_output is not None: + return + device_output = self.logits_output.auxiliary_device_output + if device_output is not None: + self.auxiliary_host_output = device_output.copy_to_host(_async_d2h) + self.logits_output.auxiliary_device_output = None + @classmethod def from_pp_proxy( cls, logits_output, next_pp_outputs: PPProxyTensors, can_run_cuda_graph diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index e9ea1f666..59583db17 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -182,6 +182,7 @@ from sglang.srt.runtime_context import ( set_global_dwdp_manager, ) from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo +from sglang.srt.sampling.sampling_observer import SamplingObserver from sglang.srt.server_args import ( # noqa: F401 (re-export) CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS, ServerArgs, @@ -284,6 +285,23 @@ def resolve_draft_attention_backend( class ModelRunner: """ModelRunner runs the forward passes of the models.""" + @property + def sampling_observer(self) -> Optional[SamplingObserver]: + return self._sampling_observer + + @sampling_observer.setter + def sampling_observer(self, observer: Optional[SamplingObserver]) -> None: + if observer is not None and not self.supports_sampling_observer(): + raise ValueError( + "sampling observers are not supported by the configured " + "sampling path" + ) + self._sampling_observer = observer + + def supports_sampling_observer(self) -> bool: + """Whether this runner's sampling path publishes observer output.""" + return self.server_args.dllm_algorithm is None and self.spec_algorithm.is_none() + def __init__( self, model_config: ModelConfig, @@ -355,6 +373,7 @@ class ModelRunner: self.init_new_workspace = False self.draft_model_idx = draft_model_idx self.enable_hisparse = server_args.enable_hisparse + self._sampling_observer: Optional[SamplingObserver] = None self.init_startup_observability() @@ -1743,14 +1762,24 @@ class ModelRunner: return ModelRunnerOutput(logits_output=ret, can_run_graph=can_run_graph) def _preprocess_logits( - self, logits_output: LogitsProcessorOutput, sampling_info: SamplingBatchInfo + self, + logits_output: LogitsProcessorOutput, + sampling_info: SamplingBatchInfo, + observer: Optional[SamplingObserver] = None, ): # NOTE: In overlap mode, the function update_regex_vocab_mask (in sample) # was executed after we processed last batch's results. # Calculate logits bias and apply it to next_token_logits. sampling_info.update_regex_vocab_mask() - sampling_info.apply_logits_bias(logits_output.next_token_logits) + observer_state = None + if observer is not None: + observer_state = sampling_info.apply_logits_bias_with_observer( + logits_output.next_token_logits, + observer=observer, + ) + else: + sampling_info.apply_logits_bias(logits_output.next_token_logits) # Release the vocab_mask GPU tensor immediately after it has been applied # to the logits. In overlap scheduling, the sampling_info (and its @@ -1758,6 +1787,7 @@ class ModelRunner: # batch_record_buf until the next iteration, causing a steady VRAM leak # when structured output (grammar) is used. sampling_info.grammar_mask = None + return observer_state def sample( self, @@ -1773,7 +1803,22 @@ class ModelRunner: Returns: A list of next_token_ids """ - self._preprocess_logits(logits_output, forward_batch.sampling_info) + # LogitsProcessorOutput is normally invocation-scoped, but CUDA graph + # runners may reuse backing objects. Never leak an auxiliary result from + # a previous replay into a request with no observer state. + logits_output.auxiliary_device_output = None + observer = self.sampling_observer + # Preserve two-argument overrides when observation is inactive. + if observer is not None and observer.is_active(forward_batch.sampling_info): + observer_state = self._preprocess_logits( + logits_output, + forward_batch.sampling_info, + observer=observer, + ) + else: + observer_state = self._preprocess_logits( + logits_output, forward_batch.sampling_info + ) # Sample the next tokens next_token_ids = self.sampler( @@ -1789,6 +1834,11 @@ class ModelRunner: else forward_batch.seq_lens - 1 ), ) + if observer_state is not None: + logits_output.auxiliary_device_output = observer.after_sample( + observer_state, + next_token_ids, + ) self.ngram_embedding_manager.update_after_decode( next_token_ids=next_token_ids, forward_batch=forward_batch, @@ -1811,10 +1861,10 @@ class ModelRunner: logits_output: The logits output from the model forward forward_batch: The forward batch that generates logits_output """ + logits_output.auxiliary_device_output = None if not forward_batch.token_ids_logprobs: return - # Preprocess logits (same as in sample method) self._preprocess_logits(logits_output, forward_batch.sampling_info) # Delegate to sampler for logprob-only computation diff --git a/python/sglang/srt/sampling/sampling_batch_info.py b/python/sglang/srt/sampling/sampling_batch_info.py index c07b0a05d..9bf2c9683 100644 --- a/python/sglang/srt/sampling/sampling_batch_info.py +++ b/python/sglang/srt/sampling/sampling_batch_info.py @@ -20,6 +20,7 @@ from sglang.srt.utils.common import is_pin_memory_available if TYPE_CHECKING: from sglang.srt.managers.schedule_batch import ScheduleBatch + from sglang.srt.sampling.sampling_observer import SamplingObserver logger = logging.getLogger(__name__) @@ -280,7 +281,7 @@ class SamplingBatchInfo: self.acc_additive_penalties = None self.acc_scaling_penalties = None - def apply_logits_bias(self, logits: torch.Tensor): + def _apply_pre_grammar_logits_transforms(self, logits: torch.Tensor) -> None: if self.acc_additive_penalties is not None: # Used in the overlap mode logits.add_(self.acc_additive_penalties) @@ -293,11 +294,32 @@ class SamplingBatchInfo: # Used in the non-overlap mode self.penalizer_orchestrator.apply(logits) + def _apply_post_grammar_logits_transforms(self, logits: torch.Tensor) -> None: + if self.logit_bias is not None: + logits.add_(self.logit_bias) + + def apply_logits_bias(self, logits: torch.Tensor): + self._apply_pre_grammar_logits_transforms(logits) + if self.grammar_mask is not None: self.grammar_mask.apply(logits) - if self.logit_bias is not None: - logits.add_(self.logit_bias) + self._apply_post_grammar_logits_transforms(logits) + + def apply_logits_bias_with_observer( + self, + logits: torch.Tensor, + observer: SamplingObserver, + ) -> Any: + self._apply_pre_grammar_logits_transforms(logits) + observer_state = observer.before_grammar(logits, self) + + if self.grammar_mask is not None: + self.grammar_mask.apply(logits) + + self._apply_post_grammar_logits_transforms(logits) + + return observer_state def filter_batch(self, keep_indices: List[int], keep_indices_device: torch.Tensor): self.penalizer_orchestrator.filter(keep_indices_device) diff --git a/python/sglang/srt/sampling/sampling_observer.py b/python/sglang/srt/sampling/sampling_observer.py new file mode 100644 index 000000000..135d196ee --- /dev/null +++ b/python/sglang/srt/sampling/sampling_observer.py @@ -0,0 +1,79 @@ +"""Extension contracts for sampling-time auxiliary response metadata. + +Out-of-tree integrations can install these hooks from an SGLang plugin by +extending ``ModelRunner``, ``Scheduler``, and ``TokenizerManager`` through the +plugin hook registry. The model runner installs a ``SamplingObserver``; the +scheduler selects a ``SchedulerOutputStreamer`` subclass; and the tokenizer +manager consumes the resulting customized response fields. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any, Callable, Optional, Protocol, Sequence + +import torch + +if TYPE_CHECKING: + from sglang.srt.managers.schedule_batch import ScheduleBatch + from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo + + +@dataclass(frozen=True) +class CommittedTokens: + output_index: int + token_ids: tuple[int, ...] + + +class DeviceAuxiliaryOutput(Protocol): + """Device output copied later by the scheduler. + + Tensors must not alias CUDA-graph static buffers that a later replay can + overwrite before the scheduler-side copy completes. + """ + + def copy_to_host( + self, copy_tensor: Callable[[torch.Tensor], torch.Tensor] + ) -> HostAuxiliaryOutput: ... + + +class HostAuxiliaryOutput(Protocol): + """Scheduler-side result produced by ``DeviceAuxiliaryOutput``. + + ``consume`` runs after sampled tokens have been committed to each request + and immediately before response streaming. ``commits`` is aligned with + ``batch.reqs`` and identifies only the newly visible tokens. Implementations + can buffer per-request values for a ``SchedulerOutputStreamer`` subclass to + expose through customized response metadata. + """ + + def consume( + self, + batch: ScheduleBatch, + commits: Sequence[Optional[CommittedTokens]], + ) -> None: ... + + +class SamplingObserver(Protocol): + """Invocation-scoped hooks around the production grammar mask and sampler. + + Returning ``None`` from ``before_grammar`` skips ``after_sample``. Install + an observer through ``ModelRunner.sampling_observer`` from a model-runner + subclass or plugin hook. Specialized sampling paths must override + ``ModelRunner.supports_sampling_observer`` and publish equivalent auxiliary + output before installing one. + """ + + def is_active(self, sampling_info: SamplingBatchInfo) -> bool: ... + + def before_grammar( + self, + logits: torch.Tensor, + sampling_info: SamplingBatchInfo, + ) -> Any: ... + + def after_sample( + self, state: Any, token_ids: torch.Tensor + ) -> Optional[DeviceAuxiliaryOutput]: + """Return graph-safe device output for later scheduler-side copying.""" + ... diff --git a/python/sglang/srt/sampling/sampling_observer_pp.py b/python/sglang/srt/sampling/sampling_observer_pp.py new file mode 100644 index 000000000..4de76b969 --- /dev/null +++ b/python/sglang/srt/sampling/sampling_observer_pp.py @@ -0,0 +1,86 @@ +from __future__ import annotations + +from typing import ( + Any, + Mapping, + MutableMapping, + Optional, + Protocol, + runtime_checkable, +) + +import torch + +from sglang.srt.sampling.sampling_observer import ( + DeviceAuxiliaryOutput, + SamplingObserver, +) + + +@runtime_checkable +class PipelineParallelAuxiliaryOutput(Protocol): + def to_pp_tensors(self) -> Mapping[str, torch.Tensor]: ... + + +@runtime_checkable +class PipelineParallelSamplingObserver(Protocol): + def from_pp_tensors( + self, tensors: Mapping[str, torch.Tensor] + ) -> DeviceAuxiliaryOutput: ... + + +_OUTPUT_PREFIX = "__sampling_observer_output__." + + +def add_auxiliary_output_to_pp_tensors( + tensors: MutableMapping[str, Any], + output: Optional[DeviceAuxiliaryOutput], +) -> None: + if output is None: + return + if not isinstance(output, PipelineParallelAuxiliaryOutput): + raise RuntimeError( + "auxiliary output does not support pipeline-parallel transport" + ) + + output_tensors = output.to_pp_tensors() + if not output_tensors: + raise RuntimeError("auxiliary PP output must contain at least one tensor") + + for name, tensor in output_tensors.items(): + if not isinstance(name, str) or not name: + raise RuntimeError("auxiliary PP tensor names must be non-empty strings") + if not torch.is_tensor(tensor): + raise RuntimeError(f"auxiliary PP output {name!r} is not a tensor") + key = f"{_OUTPUT_PREFIX}{name}" + if key in tensors: + raise RuntimeError(f"duplicate auxiliary PP tensor {name!r}") + tensors[key] = tensor + + +def pop_auxiliary_output_from_pp_tensors( + tensors: MutableMapping[str, Any], + observer: Optional[SamplingObserver], +) -> Optional[DeviceAuxiliaryOutput]: + output_tensors = { + key.removeprefix(_OUTPUT_PREFIX): value + for key, value in tensors.items() + if key.startswith(_OUTPUT_PREFIX) + } + if not output_tensors: + return None + if observer is None: + raise RuntimeError("received auxiliary PP output without a sampling observer") + if not isinstance(observer, PipelineParallelSamplingObserver): + raise RuntimeError( + "sampling observer does not support pipeline-parallel transport" + ) + if any(not torch.is_tensor(tensor) for tensor in output_tensors.values()): + raise RuntimeError("received a non-tensor auxiliary PP output") + + output = observer.from_pp_tensors(output_tensors) + if output is None: + raise RuntimeError("sampling observer did not reconstruct its PP output") + for name in output_tensors: + del tensors[f"{_OUTPUT_PREFIX}{name}"] + return output diff --git a/test/registered/unit/managers/test_batch_result_processor_hidden_states.py b/test/registered/unit/managers/test_batch_result_processor_hidden_states.py index c96081bbd..82d1beee5 100644 --- a/test/registered/unit/managers/test_batch_result_processor_hidden_states.py +++ b/test/registered/unit/managers/test_batch_result_processor_hidden_states.py @@ -132,6 +132,7 @@ class TestPrefillHiddenStateOffsets(CustomTestCase): ) result = SimpleNamespace( copy_done=None, + auxiliary_host_output=None, routed_experts_output=None, indexer_topk_output=None, logits_output=SimpleNamespace( @@ -180,6 +181,7 @@ class TestDecodeHiddenStateRetention(CustomTestCase): def result(hidden_states): return SimpleNamespace( copy_done=None, + auxiliary_host_output=None, routed_experts_output=None, indexer_topk_output=None, logits_output=SimpleNamespace(hidden_states=hidden_states), diff --git a/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py b/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py index 8173cd23d..2cc5f4d46 100644 --- a/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py +++ b/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py @@ -80,6 +80,7 @@ def _make_processor() -> SchedulerBatchResultProcessor: def _make_result(): return SimpleNamespace( copy_done=None, + auxiliary_host_output=None, routed_experts_output=None, indexer_topk_output=None, logits_output=SimpleNamespace(hidden_states=None, customized_info=None), diff --git a/test/registered/unit/managers/test_generation_auxiliary_output.py b/test/registered/unit/managers/test_generation_auxiliary_output.py new file mode 100644 index 000000000..ae371fd12 --- /dev/null +++ b/test/registered/unit/managers/test_generation_auxiliary_output.py @@ -0,0 +1,731 @@ +from dataclasses import dataclass +from types import SimpleNamespace +from unittest.mock import Mock, patch + +import pytest +import torch + +from sglang.srt.disaggregation.prefill import SchedulerDisaggregationPrefillMixin +from sglang.srt.layers.logits_processor import LogitsProcessorOutput +from sglang.srt.managers.scheduler import Scheduler +from sglang.srt.managers.scheduler_components.batch_result_processor import ( + SchedulerBatchResultProcessor, +) +from sglang.srt.managers.scheduler_pp_mixin import PPBatchMetadata +from sglang.srt.managers.utils import GenerationBatchResult +from sglang.srt.model_executor.forward_batch_info import PPProxyTensors +from sglang.srt.model_executor.model_runner import ModelRunner +from sglang.srt.speculative.spec_info import SpeculativeAlgorithm +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=1, suite="base-a-test-cpu") + + +@dataclass +class HostOutput: + values: torch.Tensor + + +class DeviceOutput: + def __init__(self, values: torch.Tensor): + self.values = values + self.copy_count = 0 + + def copy_to_host(self, copy_tensor): + self.copy_count += 1 + return HostOutput(copy_tensor(self.values)) + + def to_pp_tensors(self): + return {"values": self.values} + + +class HostOnlyDeviceOutput: + def __init__(self, values: torch.Tensor): + self.values = values + + def copy_to_host(self, copy_tensor): + return HostOutput(copy_tensor(self.values)) + + +class Observer: + def __init__(self): + self.received_tensors = None + + def from_pp_tensors(self, tensors): + self.received_tensors = tensors + return DeviceOutput(tensors["values"]) + + +class CopyDone: + def __init__(self): + self.record_count = 0 + + def record(self): + self.record_count += 1 + + +def _model_runner_for_sampling_path( + *, + spec_algorithm=SpeculativeAlgorithm.NONE, + dllm_algorithm=None, +): + runner = object.__new__(ModelRunner) + runner.server_args = SimpleNamespace(dllm_algorithm=dllm_algorithm) + runner.spec_algorithm = spec_algorithm + runner._sampling_observer = None + return runner + + +def test_auxiliary_output_releases_device_holder_after_copy(): + device_output = DeviceOutput(torch.tensor([1.0, 2.0])) + logits_output = LogitsProcessorOutput( + next_token_logits=None, + auxiliary_device_output=device_output, + ) + result = GenerationBatchResult( + logits_output=logits_output, + next_token_ids=torch.tensor([7]), + copy_done=CopyDone(), + ) + + result.copy_to_cpu(return_logprob=False) + + assert logits_output.auxiliary_device_output is None + assert result.auxiliary_host_output is not device_output + assert device_output.copy_count == 1 + assert result.copy_done.record_count == 1 + + +def test_non_pp_auxiliary_output_only_requires_host_copy_support(): + device_output = HostOnlyDeviceOutput(torch.tensor([1.0, 2.0])) + result = GenerationBatchResult( + logits_output=LogitsProcessorOutput( + next_token_logits=None, + auxiliary_device_output=device_output, + ), + next_token_ids=torch.tensor([7]), + copy_done=CopyDone(), + ) + + result.copy_to_cpu(return_logprob=False) + + assert torch.equal(result.auxiliary_host_output.values, device_output.values) + + +def test_auxiliary_host_outputs_are_owned_by_each_generation_result(): + logits_output = LogitsProcessorOutput(next_token_logits=None) + first_device = DeviceOutput(torch.tensor([1.0])) + logits_output.auxiliary_device_output = first_device + first = GenerationBatchResult( + logits_output=logits_output, + next_token_ids=torch.tensor([1]), + copy_done=CopyDone(), + ) + first.copy_to_cpu(return_logprob=False) + + second_device = DeviceOutput(torch.tensor([2.0])) + logits_output.auxiliary_device_output = second_device + second = GenerationBatchResult( + logits_output=logits_output, + next_token_ids=torch.tensor([2]), + copy_done=CopyDone(), + ) + second.copy_to_cpu(return_logprob=False) + + assert first.auxiliary_host_output.values.tolist() == [1.0] + assert second.auxiliary_host_output.values.tolist() == [2.0] + assert first_device.copy_count == second_device.copy_count == 1 + + +def test_sampling_clears_stale_device_output_when_observer_produces_no_state(): + runner = _model_runner_for_sampling_path() + runner.sampling_observer = SimpleNamespace( + is_active=lambda sampling_info: True, + after_sample=lambda state, token_ids: None, + ) + runner._preprocess_logits = Mock(return_value=None) + runner.sampler = lambda *args, **kwargs: torch.tensor([3]) + runner.ngram_embedding_manager = SimpleNamespace( + update_after_decode=lambda **kwargs: None + ) + logits_output = LogitsProcessorOutput( + next_token_logits=torch.zeros(1, 4), + auxiliary_device_output=DeviceOutput(torch.tensor([99.0])), + ) + forward_batch = SimpleNamespace( + sampling_info=object(), + return_logprob=False, + top_logprobs_nums=None, + token_ids_logprobs=None, + positions=torch.tensor([0]), + seq_lens=torch.tensor([1]), + forward_mode=SimpleNamespace(is_decode=lambda: True), + ) + + ModelRunner.sample(runner, logits_output, forward_batch) + + assert logits_output.auxiliary_device_output is None + runner._preprocess_logits.assert_called_once_with( + logits_output, + forward_batch.sampling_info, + observer=runner.sampling_observer, + ) + + +def test_sampling_publishes_observer_output_for_the_sampled_tokens(): + state = object() + device_output = DeviceOutput(torch.tensor([4.0])) + observer = SimpleNamespace( + is_active=Mock(return_value=True), + after_sample=Mock(return_value=device_output), + ) + runner = _model_runner_for_sampling_path() + runner.sampling_observer = observer + runner._preprocess_logits = Mock(return_value=state) + sampled_tokens = torch.tensor([3]) + runner.sampler = Mock(return_value=sampled_tokens) + runner.ngram_embedding_manager = SimpleNamespace( + update_after_decode=lambda **kwargs: None + ) + logits_output = LogitsProcessorOutput(next_token_logits=torch.zeros(1, 4)) + forward_batch = SimpleNamespace( + sampling_info=object(), + return_logprob=False, + top_logprobs_nums=None, + token_ids_logprobs=None, + positions=torch.tensor([0]), + seq_lens=torch.tensor([1]), + forward_mode=SimpleNamespace(is_decode=lambda: True), + ) + + ModelRunner.sample(runner, logits_output, forward_batch) + + assert logits_output.auxiliary_device_output is device_output + runner._preprocess_logits.assert_called_once_with( + logits_output, + forward_batch.sampling_info, + observer=observer, + ) + observer.is_active.assert_called_once_with(forward_batch.sampling_info) + observer.after_sample.assert_called_once_with(state, sampled_tokens) + + +@pytest.mark.parametrize( + ("spec_algorithm", "dllm_algorithm"), + [ + (SpeculativeAlgorithm.EAGLE, None), + (SpeculativeAlgorithm.NONE, "dream"), + ], +) +def test_sampling_observer_rejects_sampling_paths_that_bypass_hooks( + spec_algorithm, + dllm_algorithm, +): + runner = _model_runner_for_sampling_path( + spec_algorithm=spec_algorithm, + dllm_algorithm=dllm_algorithm, + ) + + with pytest.raises(ValueError, match="configured sampling path"): + runner.sampling_observer = Observer() + + +def test_custom_sampling_path_can_enable_sampling_observer(): + class SupportedModelRunner(ModelRunner): + def supports_sampling_observer(self): + return True + + runner = object.__new__(SupportedModelRunner) + observer = Observer() + + runner.sampling_observer = observer + + assert runner.sampling_observer is observer + + +@pytest.mark.parametrize("has_inactive_observer", [False, True]) +def test_sampling_without_active_observer_preserves_preprocess_override( + has_inactive_observer, +): + observer = ( + SimpleNamespace(is_active=Mock(return_value=False)) + if has_inactive_observer + else None + ) + runner = _model_runner_for_sampling_path() + runner.sampling_observer = observer + runner._preprocess_logits = Mock( + side_effect=lambda logits_output, sampling_info: None + ) + runner.sampler = Mock(return_value=torch.tensor([3])) + runner.ngram_embedding_manager = SimpleNamespace( + update_after_decode=lambda **kwargs: None + ) + logits_output = LogitsProcessorOutput(next_token_logits=torch.zeros(1, 4)) + forward_batch = SimpleNamespace( + sampling_info=object(), + return_logprob=False, + top_logprobs_nums=None, + token_ids_logprobs=None, + positions=torch.tensor([0]), + seq_lens=torch.tensor([1]), + forward_mode=SimpleNamespace(is_decode=lambda: True), + ) + + ModelRunner.sample(runner, logits_output, forward_batch) + + runner._preprocess_logits.assert_called_once_with( + logits_output, forward_batch.sampling_info + ) + if observer is not None: + observer.is_active.assert_called_once_with(forward_batch.sampling_info) + + +def test_preprocess_logits_without_observer_uses_standard_path(): + runner = object.__new__(ModelRunner) + logits_output = LogitsProcessorOutput(next_token_logits=torch.zeros(1, 4)) + grammar_mask = object() + sampling_info = SimpleNamespace( + grammar_mask=grammar_mask, + update_regex_vocab_mask=Mock(), + apply_logits_bias=Mock(), + apply_logits_bias_with_observer=Mock(), + ) + + state = ModelRunner._preprocess_logits( + runner, + logits_output, + sampling_info, + ) + + assert state is None + sampling_info.update_regex_vocab_mask.assert_called_once_with() + sampling_info.apply_logits_bias.assert_called_once_with( + logits_output.next_token_logits + ) + sampling_info.apply_logits_bias_with_observer.assert_not_called() + assert sampling_info.grammar_mask is None + + +def test_active_observer_uses_observer_logits_preprocessing(): + runner = object.__new__(ModelRunner) + observer = SimpleNamespace() + observer_state = object() + logits_output = LogitsProcessorOutput(next_token_logits=torch.zeros(1, 4)) + sampling_info = SimpleNamespace( + grammar_mask=object(), + update_regex_vocab_mask=Mock(), + apply_logits_bias=Mock(), + apply_logits_bias_with_observer=Mock(return_value=observer_state), + ) + + state = ModelRunner._preprocess_logits( + runner, + logits_output, + sampling_info, + observer=observer, + ) + + assert state is observer_state + sampling_info.update_regex_vocab_mask.assert_called_once_with() + sampling_info.apply_logits_bias.assert_not_called() + sampling_info.apply_logits_bias_with_observer.assert_called_once_with( + logits_output.next_token_logits, + observer=observer, + ) + assert sampling_info.grammar_mask is None + + +def test_scheduler_copies_auxiliary_output_for_non_overlap_results(): + event = object() + scheduler = object.__new__(Scheduler) + scheduler.ps = SimpleNamespace(pp_size=1) + scheduler.device_module = SimpleNamespace(Event=Mock(return_value=event)) + result = SimpleNamespace( + logits_output=SimpleNamespace(auxiliary_device_output=object()), + auxiliary_host_output=None, + copy_done=None, + copy_to_cpu=Mock(), + ) + batch = SimpleNamespace(return_logprob=False, return_hidden_states=False) + + Scheduler._copy_auxiliary_output_to_cpu(scheduler, batch, result) + + assert result.copy_done is event + result.copy_to_cpu.assert_called_once_with( + return_logprob=False, + return_hidden_states=False, + ) + + +def test_scheduler_preserves_pipeline_parallel_output_for_transport(): + scheduler = object.__new__(Scheduler) + scheduler.ps = SimpleNamespace(pp_size=2) + scheduler.device_module = SimpleNamespace(Event=Mock()) + result = SimpleNamespace( + logits_output=SimpleNamespace(auxiliary_device_output=object()), + auxiliary_host_output=None, + copy_done=None, + copy_to_cpu=Mock(), + ) + batch = SimpleNamespace(return_logprob=False, return_hidden_states=False) + + Scheduler._copy_auxiliary_output_to_cpu(scheduler, batch, result) + + assert result.copy_done is None + result.copy_to_cpu.assert_not_called() + scheduler.device_module.Event.assert_not_called() + + +def test_pdmux_split_prefill_schedules_auxiliary_output_copy(): + device_output = DeviceOutput(torch.tensor([1.0])) + result = GenerationBatchResult( + logits_output=LogitsProcessorOutput( + next_token_logits=None, + auxiliary_device_output=device_output, + ), + next_token_ids=torch.tensor([7]), + ) + copy_done = CopyDone() + scheduler = object.__new__(Scheduler) + scheduler.forward_ct = 0 + scheduler._sched_idled = False + scheduler.scripted_scheduler_hook = None + scheduler.profiler_manager = SimpleNamespace(_profile_batch_predicate=Mock()) + scheduler.forward_sleep_time = None + scheduler.disaggregation_mode = None + scheduler.is_generation = True + scheduler.enable_overlap = False + scheduler.enable_pdmux = True + scheduler.ps = SimpleNamespace(pp_size=1) + scheduler.tp_worker = SimpleNamespace( + forward_batch_split_prefill=Mock(return_value=result) + ) + scheduler.future_map = object() + scheduler._relay_forward_payload = Mock() + scheduler.device_module = SimpleNamespace(Event=Mock(return_value=copy_done)) + scheduler.enable_dp_attention = False + batch = SimpleNamespace( + forward_mode=SimpleNamespace( + is_prebuilt=lambda: False, + is_split_prefill=lambda: True, + ), + reqs=[], + req_pool_indices=torch.tensor([3]), + input_ids=torch.tensor([5]), + return_logprob=False, + return_hidden_states=False, + ) + + with patch( + "sglang.srt.managers.scheduler.resolve_forward_inputs" + ) as resolve_forward_inputs: + output_result = Scheduler.run_batch(scheduler, batch) + + resolve_forward_inputs.assert_called_once_with(batch, scheduler.future_map) + assert output_result is result + assert result.auxiliary_host_output.values.tolist() == [1.0] + assert copy_done.record_count == 1 + + +def test_disaggregated_prefill_consumes_auxiliary_output_after_commit(): + host_output = HostOutput(torch.tensor([1.0])) + copy_done = SimpleNamespace(synchronize=Mock()) + result = GenerationBatchResult( + logits_output=None, + next_token_ids=torch.tensor([7]), + next_draft_input=None, + copy_done=copy_done, + auxiliary_host_output=host_output, + ) + req = SimpleNamespace( + output_ids=[], + finished_len=None, + inflight_middle_chunks=0, + pending_bootstrap=False, + return_logprob=False, + return_sampling_mask=False, + grammar=None, + time_stats=SimpleNamespace( + set_prefill_finished_time=Mock(), + set_prefill_transfer_queue_entry_time=Mock(), + ), + ) + batch = SimpleNamespace( + reqs=[req], + spec_info=None, + prefill_stats=None, + dp_cooperation_info=None, + ) + snapshot_auxiliary_output_starts = Mock( + side_effect=SchedulerBatchResultProcessor.snapshot_auxiliary_output_starts + ) + processor = SimpleNamespace( + move_logprobs_to_cpu=Mock(), + consume_auxiliary_output=Mock(), + snapshot_auxiliary_output_starts=snapshot_auxiliary_output_starts, + ) + scheduler = SimpleNamespace( + batch_result_processor=processor, + spec_algorithm=SimpleNamespace(is_eagle=lambda: False), + tree_cache=object(), + disagg_prefill_inflight_queue=[], + send_kv_chunk=Mock(), + metrics_reporter=SimpleNamespace(report_prefill_stats=Mock()), + ) + + with patch("sglang.srt.disaggregation.prefill.maybe_cache_unfinished_req"): + SchedulerDisaggregationPrefillMixin.process_batch_result_disagg_prefill( + scheduler, + batch, + result, + ) + + assert req.output_ids == [7] + snapshot_auxiliary_output_starts.assert_called_once_with(batch, result) + processor.consume_auxiliary_output.assert_called_once_with( + batch, + host_output, + [0], + ) + + +def test_logprob_only_reuses_preprocessing_without_observer_lifecycle(): + runner = object.__new__(ModelRunner) + runner._preprocess_logits = Mock() + runner.sampler = SimpleNamespace(compute_logprobs_only=Mock()) + logits_output = LogitsProcessorOutput(next_token_logits=torch.zeros(1, 4)) + sampling_info = object() + forward_batch = SimpleNamespace( + sampling_info=sampling_info, + top_logprobs_nums=None, + token_ids_logprobs=[1], + ) + + ModelRunner.compute_logprobs_only(runner, logits_output, forward_batch) + + runner._preprocess_logits.assert_called_once_with(logits_output, sampling_info) + runner.sampler.compute_logprobs_only.assert_called_once() + + +def test_logprob_only_clears_stale_output_before_early_return(): + runner = object.__new__(ModelRunner) + runner.sampler = SimpleNamespace(compute_logprobs_only=Mock()) + logits_output = LogitsProcessorOutput( + next_token_logits=None, + auxiliary_device_output=DeviceOutput(torch.tensor([99.0])), + ) + forward_batch = SimpleNamespace(token_ids_logprobs=None) + + ModelRunner.compute_logprobs_only(runner, logits_output, forward_batch) + + assert logits_output.auxiliary_device_output is None + runner.sampler.compute_logprobs_only.assert_not_called() + + +def test_pipeline_parallel_auxiliary_output_round_trip(): + device_output = DeviceOutput(torch.tensor([1.0, 2.0])) + result = GenerationBatchResult( + logits_output=LogitsProcessorOutput( + next_token_logits=None, + auxiliary_device_output=device_output, + ), + next_token_ids=torch.tensor([7]), + ) + batch = SimpleNamespace( + return_logprob=False, + req_pool_indices=torch.tensor([3]), + input_ids=torch.tensor([5]), + ) + + tensors = Scheduler._pp_prepare_tensor_dict( + object.__new__(Scheduler), result, batch + ) + observer = Observer() + receiver = object.__new__(Scheduler) + receiver.pp_group = SimpleNamespace(is_first_rank=True) + receiver.tp_worker = SimpleNamespace( + model_runner=SimpleNamespace(sampling_observer=observer) + ) + receiver.future_map = SimpleNamespace(stash=Mock()) + + output_result = Scheduler._pp_prep_batch_result( + receiver, + batch, + PPBatchMetadata(can_run_cuda_graph=True), + PPProxyTensors(tensors), + ) + + assert set(observer.received_tensors) == {"values"} + assert torch.equal(observer.received_tensors["values"], device_output.values) + assert output_result.logits_output.auxiliary_device_output is not device_output + assert torch.equal(output_result.auxiliary_host_output.values, device_output.values) + assert all("sampling_observer_output" not in key for key in tensors) + receiver.future_map.stash.assert_called_once() + + +def test_pipeline_parallel_auxiliary_output_stays_packed_before_first_rank(): + device_output = DeviceOutput(torch.tensor([1.0])) + result = GenerationBatchResult( + logits_output=LogitsProcessorOutput( + next_token_logits=None, + auxiliary_device_output=device_output, + ), + next_token_ids=torch.tensor([7]), + ) + batch = SimpleNamespace( + return_logprob=False, + req_pool_indices=torch.tensor([3]), + input_ids=torch.tensor([5]), + ) + tensors = Scheduler._pp_prepare_tensor_dict( + object.__new__(Scheduler), result, batch + ) + receiver = object.__new__(Scheduler) + receiver.pp_group = SimpleNamespace(is_first_rank=False) + receiver.future_map = SimpleNamespace(stash=Mock()) + + output_result = Scheduler._pp_prep_batch_result( + receiver, + batch, + PPBatchMetadata(can_run_cuda_graph=False), + PPProxyTensors(tensors), + ) + + assert output_result.logits_output is None + assert any("sampling_observer_output" in key for key in tensors) + + +def test_pipeline_parallel_auxiliary_output_requires_receiver_observer(): + device_output = DeviceOutput(torch.tensor([1.0])) + result = GenerationBatchResult( + logits_output=LogitsProcessorOutput( + next_token_logits=None, + auxiliary_device_output=device_output, + ), + next_token_ids=torch.tensor([7]), + ) + batch = SimpleNamespace(return_logprob=False) + tensors = Scheduler._pp_prepare_tensor_dict( + object.__new__(Scheduler), result, batch + ) + receiver = object.__new__(Scheduler) + receiver.pp_group = SimpleNamespace(is_first_rank=True) + receiver.tp_worker = SimpleNamespace( + model_runner=SimpleNamespace(sampling_observer=None) + ) + + with pytest.raises(RuntimeError, match="without a sampling observer"): + Scheduler._pp_prep_batch_result( + receiver, + batch, + PPBatchMetadata(can_run_cuda_graph=False), + PPProxyTensors(tensors), + ) + + +def test_pipeline_parallel_auxiliary_output_requires_transport_support(): + result = GenerationBatchResult( + logits_output=LogitsProcessorOutput( + next_token_logits=None, + auxiliary_device_output=HostOnlyDeviceOutput(torch.tensor([1.0])), + ), + next_token_ids=torch.tensor([7]), + ) + batch = SimpleNamespace(return_logprob=False) + + with pytest.raises(RuntimeError, match="does not support pipeline-parallel"): + Scheduler._pp_prepare_tensor_dict(object.__new__(Scheduler), result, batch) + + +def test_pipeline_parallel_auxiliary_output_requires_transport_observer(): + result = GenerationBatchResult( + logits_output=LogitsProcessorOutput( + next_token_logits=None, + auxiliary_device_output=DeviceOutput(torch.tensor([1.0])), + ), + next_token_ids=torch.tensor([7]), + ) + batch = SimpleNamespace(return_logprob=False) + tensors = Scheduler._pp_prepare_tensor_dict( + object.__new__(Scheduler), result, batch + ) + receiver = object.__new__(Scheduler) + receiver.pp_group = SimpleNamespace(is_first_rank=True) + receiver.tp_worker = SimpleNamespace( + model_runner=SimpleNamespace(sampling_observer=SimpleNamespace()) + ) + + with pytest.raises(RuntimeError, match="does not support pipeline-parallel"): + Scheduler._pp_prep_batch_result( + receiver, + batch, + PPBatchMetadata(can_run_cuda_graph=False), + PPProxyTensors(tensors), + ) + + +def test_auxiliary_output_snapshot_uses_visible_request_lengths(): + batch = SimpleNamespace( + reqs=[ + SimpleNamespace(output_ids=[10, 11, 12], finished_len=2), + SimpleNamespace(output_ids=[20], finished_len=None), + ] + ) + result = SimpleNamespace(auxiliary_host_output=None) + + assert ( + SchedulerBatchResultProcessor.snapshot_auxiliary_output_starts(batch, result) + is None + ) + + result.auxiliary_host_output = object() + + assert SchedulerBatchResultProcessor.snapshot_auxiliary_output_starts( + batch, result + ) == [2, 1] + + +def test_auxiliary_commit_uses_the_scheduler_visible_prefix(): + req = SimpleNamespace(output_ids=[10, 11, 12], finished_len=2) + + commits = SchedulerBatchResultProcessor._build_auxiliary_commits( + SimpleNamespace(reqs=[req]), + output_starts=[1], + ) + + assert commits[0].output_index == 1 + assert commits[0].token_ids == (11,) + + +def test_auxiliary_commit_discards_samples_outside_the_visible_output(): + req = SimpleNamespace(output_ids=[10], finished_len=None) + + commits = SchedulerBatchResultProcessor._build_auxiliary_commits( + SimpleNamespace(reqs=[req]), + output_starts=[1], + ) + + assert commits == [None] + + +def test_auxiliary_output_consumes_only_newly_visible_tokens(): + req = SimpleNamespace(output_ids=[10, 11, 12], finished_len=2) + output = Mock() + batch = SimpleNamespace(reqs=[req]) + + SchedulerBatchResultProcessor.consume_auxiliary_output( + batch, + output, + output_starts=[1], + ) + + commits = output.consume.call_args.args[1] + assert commits[0].output_index == 1 + assert commits[0].token_ids == (11,) + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/unit/managers/test_output_streamer_customized_info.py b/test/registered/unit/managers/test_output_streamer_customized_info.py index b0be57dcb..24d4d2ef7 100644 --- a/test/registered/unit/managers/test_output_streamer_customized_info.py +++ b/test/registered/unit/managers/test_output_streamer_customized_info.py @@ -1,9 +1,11 @@ import unittest from types import SimpleNamespace +from unittest.mock import Mock, patch from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.managers.io_struct import unwrap_from_pickle from sglang.srt.managers.scheduler_components.output_streamer import ( + SchedulerOutputStreamer, _GenerationStreamAccumulator, ) from sglang.srt.speculative.spec_info import SpeculativeAlgorithm @@ -13,10 +15,20 @@ register_cpu_ci(est_time=1, suite="base-a-test-cpu") class _FakeReq: - def __init__(self, rid, output_ids, customized_info=None): + def __init__( + self, + rid, + output_ids, + customized_info=None, + *, + finished=False, + ): self.rid = rid self.http_worker_ipc = None - self.finished_reason = None + self._finished = finished + self.finished_reason = ( + SimpleNamespace(to_json=lambda: {"type": "stop"}) if finished else None + ) self.finished_output = False self.finished_len = None self.stream = False @@ -37,6 +49,10 @@ class _FakeReq: self.cached_tokens = 0 self.retraction_count = 0 self.time_stats = None + self.return_hidden_states = False + self.return_routed_experts = False + self.return_indexer_topk = False + self.return_sampling_mask = False self.mm_image_tokens = 0 self.mm_audio_tokens = 0 self.mm_video_tokens = 0 @@ -44,7 +60,7 @@ class _FakeReq: self.customized_info = customized_info def finished(self): - return False + return self._finished def init_incremental_detokenize(self): return self.output_ids_through_stop, 0 @@ -54,8 +70,23 @@ class _FakeReq: class TestOutputStreamerCustomizedInfo(unittest.TestCase): - def test_customized_info_is_padded_for_mixed_batches(self): - accumulator = _GenerationStreamAccumulator( + def setUp(self): + serving_patch = patch( + "sglang.srt.managers.scheduler_components.output_streamer.get_serving", + return_value=SimpleNamespace(stream_interval=1), + ) + observability_patch = patch( + "sglang.srt.managers.scheduler_components.output_streamer.get_observability", + return_value=SimpleNamespace(enable_request_time_stats_logging=False), + ) + serving_patch.start() + observability_patch.start() + self.addCleanup(serving_patch.stop) + self.addCleanup(observability_patch.stop) + + @staticmethod + def _accumulator(): + return _GenerationStreamAccumulator( return_logprob=False, return_hidden_states=False, return_routed_experts=False, @@ -67,6 +98,9 @@ class TestOutputStreamerCustomizedInfo(unittest.TestCase): get_cached_tokens_details=lambda req: None, ) + def test_customized_info_is_padded_for_mixed_batches(self): + accumulator = self._accumulator() + accumulator.accept(req=_FakeReq("r0", [10, 11])) accumulator.accept( req=_FakeReq( @@ -90,6 +124,215 @@ class TestOutputStreamerCustomizedInfo(unittest.TestCase): [[None, None], [None, None, None], [300]], ) + def test_additional_customized_info_uses_the_existing_payload(self): + class Streamer(SchedulerOutputStreamer): + has_additional_customized_info = True + + def get_cached_tokens_details(self, req): + return None + + def build_additional_customized_info(self, reqs): + return {"request_info": [[req.rid] for req in reqs]} + + outputs = [] + streamer = Streamer( + send_to_detokenizer=SimpleNamespace(send_output=outputs.append), + tree_cache=None, + ps=SimpleNamespace(dp_rank=0, attn_tp_rank=0), + server_args=SimpleNamespace( + stream_interval=1, + enable_request_time_stats_logging=False, + ), + is_generation=True, + spec_algorithm=SpeculativeAlgorithm.NONE, + disaggregation_mode=DisaggregationMode.NULL, + enable_hicache_storage=lambda: False, + ) + + streamer._stream_output_generation([_FakeReq("r0", [], finished=True)], False) + + self.assertEqual(len(outputs), 1) + self.assertEqual( + unwrap_from_pickle(outputs[0].customized_info), + {"request_info": [["r0"]]}, + ) + + def test_additional_customized_info_only_indexes_emitted_requests(self): + class Streamer(SchedulerOutputStreamer): + has_additional_customized_info = True + + def get_cached_tokens_details(self, req): + return None + + def build_additional_customized_info(self, reqs): + return {"request_info": [[req.rid] for req in reqs]} + + outputs = [] + streamer = Streamer( + send_to_detokenizer=SimpleNamespace(send_output=outputs.append), + tree_cache=None, + ps=SimpleNamespace(dp_rank=0, attn_tp_rank=0), + server_args=SimpleNamespace( + stream_interval=1, + enable_request_time_stats_logging=False, + ), + is_generation=True, + spec_algorithm=SpeculativeAlgorithm.NONE, + disaggregation_mode=DisaggregationMode.NULL, + enable_hicache_storage=lambda: False, + ) + quiet = _FakeReq("quiet", [10, 11]) + quiet.stream = True + quiet.sampling_params.stream_interval = 2 + terminal = _FakeReq("terminal", [20], finished=True) + + streamer._stream_output_generation([quiet, terminal], False) + + self.assertEqual(outputs[0].rids, ["terminal"]) + self.assertEqual( + unwrap_from_pickle(outputs[0].customized_info), + {"request_info": [["terminal"]]}, + ) + + def test_additional_customized_info_handles_suppressed_request_last(self): + class Streamer(SchedulerOutputStreamer): + has_additional_customized_info = True + + def get_cached_tokens_details(self, req): + return None + + def build_additional_customized_info(self, reqs): + return {"request_info": [[req.rid] for req in reqs]} + + outputs = [] + streamer = Streamer( + send_to_detokenizer=SimpleNamespace(send_output=outputs.append), + tree_cache=None, + ps=SimpleNamespace(dp_rank=0, attn_tp_rank=0), + server_args=SimpleNamespace( + stream_interval=1, + enable_request_time_stats_logging=False, + ), + is_generation=True, + spec_algorithm=SpeculativeAlgorithm.NONE, + disaggregation_mode=DisaggregationMode.NULL, + enable_hicache_storage=lambda: False, + ) + terminal = _FakeReq("terminal", [20], finished=True) + quiet = _FakeReq("quiet", [10, 11]) + quiet.stream = True + quiet.sampling_params.stream_interval = 2 + + streamer._stream_output_generation([terminal, quiet], False) + + self.assertEqual(outputs[0].rids, ["terminal"]) + self.assertEqual( + unwrap_from_pickle(outputs[0].customized_info), + {"request_info": [["terminal"]]}, + ) + + def test_additional_customized_info_preserves_duplicate_rid_requests(self): + accepted_reqs = [] + + class Streamer(SchedulerOutputStreamer): + has_additional_customized_info = True + + def get_cached_tokens_details(self, req): + return None + + def build_additional_customized_info(self, reqs): + accepted_reqs.extend(reqs) + return {"request_info": [[req.rid] for req in reqs]} + + outputs = [] + streamer = Streamer( + send_to_detokenizer=SimpleNamespace(send_output=outputs.append), + tree_cache=None, + ps=SimpleNamespace(dp_rank=0, attn_tp_rank=0), + server_args=SimpleNamespace( + stream_interval=1, + enable_request_time_stats_logging=False, + ), + is_generation=True, + spec_algorithm=SpeculativeAlgorithm.NONE, + disaggregation_mode=DisaggregationMode.NULL, + enable_hicache_storage=lambda: False, + ) + first = _FakeReq("duplicate", [10], finished=True) + second = _FakeReq("duplicate", [20], finished=True) + + streamer._stream_output_generation([first, second], False) + + self.assertEqual(accepted_reqs, [first, second]) + + def test_additional_customized_info_hook_is_opt_in(self): + class Streamer(SchedulerOutputStreamer): + build_additional_customized_info = Mock() + + def get_cached_tokens_details(self, req): + return None + + outputs = [] + streamer = Streamer( + send_to_detokenizer=SimpleNamespace(send_output=outputs.append), + tree_cache=None, + ps=SimpleNamespace(dp_rank=0, attn_tp_rank=0), + server_args=SimpleNamespace(), + is_generation=True, + spec_algorithm=SpeculativeAlgorithm.NONE, + disaggregation_mode=DisaggregationMode.NULL, + enable_hicache_storage=lambda: False, + ) + + streamer._stream_output_generation([_FakeReq("r0", [], finished=True)], False) + + Streamer.build_additional_customized_info.assert_not_called() + self.assertIsNone(outputs[0].customized_info) + + def test_additional_customized_info_hook_can_skip_inactive_batches(self): + class Streamer(SchedulerOutputStreamer): + has_additional_customized_info = True + build_additional_customized_info = Mock() + should_build_additional_customized_info = Mock(return_value=False) + + def get_cached_tokens_details(self, req): + return None + + outputs = [] + streamer = Streamer( + send_to_detokenizer=SimpleNamespace(send_output=outputs.append), + tree_cache=None, + ps=SimpleNamespace(dp_rank=0, attn_tp_rank=0), + server_args=SimpleNamespace(), + is_generation=True, + spec_algorithm=SpeculativeAlgorithm.NONE, + disaggregation_mode=DisaggregationMode.NULL, + enable_hicache_storage=lambda: False, + ) + + streamer._stream_output_generation([_FakeReq("r0", [], finished=True)], False) + + Streamer.should_build_additional_customized_info.assert_called_once_with() + Streamer.build_additional_customized_info.assert_not_called() + self.assertIsNone(outputs[0].customized_info) + + def test_additional_customized_info_rejects_rust_egress(self): + class Streamer(SchedulerOutputStreamer): + has_additional_customized_info = True + + with self.assertRaisesRegex(ValueError, "Rust egress"): + Streamer( + send_to_detokenizer=SimpleNamespace(), + tree_cache=None, + ps=SimpleNamespace(), + server_args=SimpleNamespace(), + is_generation=True, + spec_algorithm=SpeculativeAlgorithm.NONE, + disaggregation_mode=DisaggregationMode.NULL, + enable_hicache_storage=lambda: False, + rust_server=object(), + ) + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/sampling/test_sampling_batch_info.py b/test/registered/unit/sampling/test_sampling_batch_info.py index 7453fe2d8..83abf7f90 100644 --- a/test/registered/unit/sampling/test_sampling_batch_info.py +++ b/test/registered/unit/sampling/test_sampling_batch_info.py @@ -198,6 +198,77 @@ class TestApplyLogitsBias(CustomTestCase): info.apply_logits_bias(logits) self.assertTrue(torch.equal(logits, original)) + def test_apply_logits_bias_without_penalizer_orchestrator(self): + info = _make_info(batch_size=1, penalizer_orchestrator=None) + logits = torch.zeros(1, VOCAB_SIZE) + + info.apply_logits_bias(logits) + + self.assertTrue(torch.equal(logits, torch.zeros_like(logits))) + + def test_observer_sees_production_constraint_boundary(self): + events = [] + + class Observer: + def before_grammar(self, logits, sampling_info): + events.append(("before", logits.clone())) + return object() + + grammar = MagicMock() + grammar.apply_vocab_mask.side_effect = lambda logits, vocab_mask: logits.fill_( + -4.0 + ) + info = _make_info(batch_size=1) + info.acc_additive_penalties = torch.ones(1, VOCAB_SIZE) + info.grammar_mask = GrammarMask(grammar, torch.ones(1, VOCAB_SIZE)) + info.logit_bias = torch.full((1, VOCAB_SIZE), 2.0) + logits = torch.zeros(1, VOCAB_SIZE) + + state = info.apply_logits_bias_with_observer(logits, observer=Observer()) + + self.assertIsNotNone(state) + self.assertTrue(torch.equal(events[0][1], torch.ones_like(logits))) + self.assertEqual(len(events), 1) + self.assertTrue(torch.equal(logits, torch.full_like(logits, -2.0))) + grammar.apply_vocab_mask.assert_called_once() + + def test_observer_path_preserves_production_logit_transforms(self): + class Observer: + def before_grammar(self, logits, sampling_info): + return object() + + def make_info(): + grammar = MagicMock() + grammar.apply_vocab_mask.side_effect = ( + lambda logits, vocab_mask: logits.add_(vocab_mask) + ) + info = _make_info(batch_size=1) + info.acc_additive_penalties = torch.linspace( + -0.5, 0.5, VOCAB_SIZE + ).unsqueeze(0) + info.acc_scaling_penalties = torch.linspace(1.0, 1.5, VOCAB_SIZE).unsqueeze( + 0 + ) + info.grammar_mask = GrammarMask( + grammar, + torch.linspace(-2.0, 0.0, VOCAB_SIZE).unsqueeze(0), + ) + info.logit_bias = torch.linspace(0.0, 1.0, VOCAB_SIZE).unsqueeze(0) + return info + + ordinary = make_info() + observed = make_info() + ordinary_logits = torch.linspace(-3.0, 3.0, VOCAB_SIZE).unsqueeze(0) + observed_logits = ordinary_logits.clone() + + ordinary.apply_logits_bias(ordinary_logits) + observed.apply_logits_bias_with_observer( + observed_logits, + observer=Observer(), + ) + + self.assertTrue(torch.equal(observed_logits, ordinary_logits)) + # update_penalties class TestUpdatePenalties(CustomTestCase):