[Whisper] Automatic language detection via structured generation (#22997)
Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
This commit is contained in:
co-authored by
Xinyuan Tong
parent
c2ec64f243
commit
a3fc982ba7
+1
-1
@@ -1,3 +1,3 @@
|
|||||||
[codespell]
|
[codespell]
|
||||||
ignore-words-list = ans, als, hel, boostrap, childs, te, vas, hsa, ment, cann, thi, makro, wil, rouge, PRIS, ather, MIS, medias, allready, inout, nd
|
ignore-words-list = ans, als, hel, boostrap, childs, te, vas, hsa, ment, cann, thi, makro, wil, rouge, PRIS, ather, MIS, medias, allready, inout, nd, fo, visibles
|
||||||
skip = *.json, *.jsonl, *.patch, *.txt, *.lock
|
skip = *.json, *.jsonl, *.patch, *.txt, *.lock
|
||||||
|
|||||||
@@ -80,7 +80,10 @@ class OpenAIServingTranscription(OpenAIServingBase):
|
|||||||
raw_request: Request = None,
|
raw_request: Request = None,
|
||||||
) -> tuple[GenerateReqInput, TranscriptionRequest]:
|
) -> tuple[GenerateReqInput, TranscriptionRequest]:
|
||||||
"""Convert transcription request to internal format."""
|
"""Convert transcription request to internal format."""
|
||||||
sampling_params = self._adapter.build_sampling_params(request)
|
if getattr(request, "_fused_autodetect", False):
|
||||||
|
sampling_params = self._adapter.build_fused_autodetect_params(request)
|
||||||
|
else:
|
||||||
|
sampling_params = self._adapter.build_sampling_params(request)
|
||||||
adapted_request = GenerateReqInput(
|
adapted_request = GenerateReqInput(
|
||||||
text="", # Empty text — the multimodal processor sets proper decoder/prompt tokens
|
text="", # Empty text — the multimodal processor sets proper decoder/prompt tokens
|
||||||
audio_data=request.audio_data,
|
audio_data=request.audio_data,
|
||||||
@@ -125,6 +128,22 @@ class OpenAIServingTranscription(OpenAIServingBase):
|
|||||||
# Calculate audio duration for usage reporting
|
# Calculate audio duration for usage reporting
|
||||||
audio_duration_s = self._get_audio_duration(audio_data)
|
audio_duration_s = self._get_audio_duration(audio_data)
|
||||||
|
|
||||||
|
# When language is not specified and the adapter supports detection,
|
||||||
|
# use a single fused request: SGLang's structured generation (regex)
|
||||||
|
# constrains the first 3 decode tokens to the forced prefix while
|
||||||
|
# allowing free transcription afterwards — one encoder pass, no
|
||||||
|
# extra round-trip. The adapter picks the regex variant based on
|
||||||
|
# whether timestamps were requested, so fused covers all four
|
||||||
|
# combinations of (stream, timestamp_granularities):
|
||||||
|
# * non-streaming: parse_fused_output strips the prefix and
|
||||||
|
# scrubs trailing/embedded special tokens.
|
||||||
|
# * streaming: the handler buffers until the sentinel,
|
||||||
|
# re-anchors, and scrubs each delta via
|
||||||
|
# adapter.strip_special_tokens.
|
||||||
|
# verbose_json segment timing still comes from _parse_segments
|
||||||
|
# over output_ids, which is unaffected by the string-level scrub.
|
||||||
|
use_fused = language is None and self._adapter.supports_language_detection
|
||||||
|
|
||||||
# Build request
|
# Build request
|
||||||
request = TranscriptionRequest(
|
request = TranscriptionRequest(
|
||||||
audio_data=audio_data,
|
audio_data=audio_data,
|
||||||
@@ -136,6 +155,13 @@ class OpenAIServingTranscription(OpenAIServingBase):
|
|||||||
stream=stream,
|
stream=stream,
|
||||||
audio_duration_s=audio_duration_s,
|
audio_duration_s=audio_duration_s,
|
||||||
)
|
)
|
||||||
|
if use_fused:
|
||||||
|
request._fused_autodetect = True
|
||||||
|
# Stash the variant alongside the flag so the adapter dispatch in
|
||||||
|
# parse_fused_output and the build_fused_autodetect_params regex
|
||||||
|
# selection see the same boolean — and we don't recompute it on
|
||||||
|
# every cumulative-text snapshot in streaming.
|
||||||
|
request._fused_ts_variant = bool(timestamp_granularities)
|
||||||
|
|
||||||
# Use the base class handle_request pattern
|
# Use the base class handle_request pattern
|
||||||
return await self.handle_request(request, raw_request)
|
return await self.handle_request(request, raw_request)
|
||||||
@@ -161,6 +187,27 @@ class OpenAIServingTranscription(OpenAIServingBase):
|
|||||||
return self.create_error_response(str(e))
|
return self.create_error_response(str(e))
|
||||||
|
|
||||||
text = self._adapter.postprocess_text(ret.get("text", ""))
|
text = self._adapter.postprocess_text(ret.get("text", ""))
|
||||||
|
|
||||||
|
# For fused auto-detect, parse_fused_output returns the scrubbed
|
||||||
|
# user-visible text. On parse failure (FSM abort, truncation) it
|
||||||
|
# returns (None, None) and we fall back to a best-effort scrub —
|
||||||
|
# the language stays unset rather than reporting a bogus detection.
|
||||||
|
if getattr(request, "_fused_autodetect", False):
|
||||||
|
lang, visible = self._adapter.parse_fused_output(
|
||||||
|
text, ts_variant=getattr(request, "_fused_ts_variant", False)
|
||||||
|
)
|
||||||
|
if visible is None:
|
||||||
|
logger.warning(
|
||||||
|
"Fused auto-detect parse failed on non-streaming response; "
|
||||||
|
"falling back to raw-text scrub."
|
||||||
|
)
|
||||||
|
text = self._adapter.strip_special_tokens(text)
|
||||||
|
else:
|
||||||
|
text = visible
|
||||||
|
if lang is not None:
|
||||||
|
request.language = lang
|
||||||
|
logger.info("Auto-detected language: '%s'", lang)
|
||||||
|
|
||||||
usage = TranscriptionUsage(seconds=int(math.ceil(request.audio_duration_s)))
|
usage = TranscriptionUsage(seconds=int(math.ceil(request.audio_duration_s)))
|
||||||
|
|
||||||
# Build response based on format
|
# Build response based on format
|
||||||
@@ -204,11 +251,33 @@ class OpenAIServingTranscription(OpenAIServingBase):
|
|||||||
request: TranscriptionRequest,
|
request: TranscriptionRequest,
|
||||||
raw_request: Request,
|
raw_request: Request,
|
||||||
) -> AsyncGenerator[str, None]:
|
) -> AsyncGenerator[str, None]:
|
||||||
"""Generate streaming transcription response."""
|
"""Generate streaming transcription response.
|
||||||
|
|
||||||
|
In fused auto-detect mode, each cumulative-text snapshot is passed
|
||||||
|
through ``parse_fused_output`` — which returns ``(None, None)``
|
||||||
|
while the forced prefix is still arriving and ``(lang, visible)``
|
||||||
|
once it's in. ``visible`` is already stripped of the prefix and
|
||||||
|
scrubbed of embedded special tokens, and it grows monotonically
|
||||||
|
across snapshots, so deltas are a plain suffix slice.
|
||||||
|
"""
|
||||||
created_time = int(time.time())
|
created_time = int(time.time())
|
||||||
request_id = f"{self._request_id_prefix()}{uuid.uuid4().hex}"
|
request_id = f"{self._request_id_prefix()}{uuid.uuid4().hex}"
|
||||||
model = request.model
|
model = request.model
|
||||||
stream_buffer = ""
|
visible_buffer = ""
|
||||||
|
|
||||||
|
fused_mode = getattr(request, "_fused_autodetect", False)
|
||||||
|
ts_variant = getattr(request, "_fused_ts_variant", False)
|
||||||
|
# When ``incremental_streaming_output`` is enabled, each chunk's
|
||||||
|
# ``content["text"]`` is the new delta from the detokenizer, not
|
||||||
|
# the cumulative text. Always reconstruct cumulative text locally
|
||||||
|
# so the rest of the loop (prefix parse + visible-buffer slice)
|
||||||
|
# works uniformly under either mode.
|
||||||
|
incremental = getattr(
|
||||||
|
self.tokenizer_manager.server_args,
|
||||||
|
"incremental_streaming_output",
|
||||||
|
False,
|
||||||
|
)
|
||||||
|
cumulative_text = ""
|
||||||
|
|
||||||
try:
|
try:
|
||||||
async for content in self.tokenizer_manager.generate_request(
|
async for content in self.tokenizer_manager.generate_request(
|
||||||
@@ -217,10 +286,44 @@ class OpenAIServingTranscription(OpenAIServingBase):
|
|||||||
finish_reason = content["meta_info"]["finish_reason"]
|
finish_reason = content["meta_info"]["finish_reason"]
|
||||||
finish_reason_type = finish_reason["type"] if finish_reason else None
|
finish_reason_type = finish_reason["type"] if finish_reason else None
|
||||||
|
|
||||||
# Calculate delta (new text since last chunk)
|
chunk_text = content.get("text", "")
|
||||||
current_text = content.get("text", "")
|
if incremental:
|
||||||
delta = current_text[len(stream_buffer) :]
|
cumulative_text += chunk_text
|
||||||
stream_buffer = current_text
|
else:
|
||||||
|
cumulative_text = chunk_text
|
||||||
|
|
||||||
|
if fused_mode:
|
||||||
|
lang, visible = self._adapter.parse_fused_output(
|
||||||
|
cumulative_text, ts_variant=ts_variant
|
||||||
|
)
|
||||||
|
if visible is None:
|
||||||
|
# Prefix not yet locatable. Keep buffering until the
|
||||||
|
# stream ends.
|
||||||
|
if not finish_reason_type:
|
||||||
|
continue
|
||||||
|
# Stream ended before the forced prefix was parseable —
|
||||||
|
# emit an SSE error frame so the client can distinguish
|
||||||
|
# this from "silent audio, zero transcription" and raise
|
||||||
|
# a real error instead of quietly succeeding.
|
||||||
|
logger.warning(
|
||||||
|
"Fused auto-detect stream finished before prefix "
|
||||||
|
"was parseable; returning detection-failed error."
|
||||||
|
)
|
||||||
|
error = self.create_streaming_error_response(
|
||||||
|
"language auto-detect failed: forced-prefix sentinel "
|
||||||
|
"was not produced before stream end"
|
||||||
|
)
|
||||||
|
yield f"data: {error}\n\n"
|
||||||
|
yield "data: [DONE]\n\n"
|
||||||
|
return
|
||||||
|
if lang is not None and request.language is None:
|
||||||
|
request.language = lang
|
||||||
|
logger.info("Auto-detected language: '%s'", lang)
|
||||||
|
else:
|
||||||
|
visible = cumulative_text
|
||||||
|
|
||||||
|
delta = visible[len(visible_buffer) :]
|
||||||
|
visible_buffer = visible
|
||||||
|
|
||||||
# Send content delta if there's new text
|
# Send content delta if there's new text
|
||||||
if delta:
|
if delta:
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import List
|
from typing import List, Optional
|
||||||
|
|
||||||
from sglang.srt.entrypoints.openai.protocol import (
|
from sglang.srt.entrypoints.openai.protocol import (
|
||||||
TranscriptionRequest,
|
TranscriptionRequest,
|
||||||
@@ -22,6 +22,60 @@ class TranscriptionAdapter(ABC):
|
|||||||
def build_sampling_params(self, request: TranscriptionRequest) -> dict:
|
def build_sampling_params(self, request: TranscriptionRequest) -> dict:
|
||||||
"""Return the ``sampling_params`` dict for ``GenerateReqInput``."""
|
"""Return the ``sampling_params`` dict for ``GenerateReqInput``."""
|
||||||
|
|
||||||
|
@property
|
||||||
|
def supports_language_detection(self) -> bool:
|
||||||
|
"""Whether this model supports automatic language detection.
|
||||||
|
|
||||||
|
When True, the adapter must implement the fused autodetect methods
|
||||||
|
and the standalone detection methods below.
|
||||||
|
"""
|
||||||
|
return False
|
||||||
|
|
||||||
|
# -- Fused detect+transcribe (used by the server) ----------------------
|
||||||
|
|
||||||
|
def build_fused_autodetect_params(self, request) -> dict:
|
||||||
|
"""Return ``sampling_params`` dict for a fused detect+transcribe request.
|
||||||
|
|
||||||
|
Uses structured generation (``regex``) to constrain the output prefix
|
||||||
|
to a valid language + task token sequence while allowing free
|
||||||
|
transcription afterwards — all in a single request.
|
||||||
|
"""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def parse_fused_output(
|
||||||
|
text: str, *, ts_variant: bool = False
|
||||||
|
) -> tuple[Optional[str], Optional[str]]:
|
||||||
|
"""Parse the fused output into ``(language_code, user_visible_text)``.
|
||||||
|
|
||||||
|
Called by both streaming and non-streaming handlers with the same
|
||||||
|
contract. ``ts_variant`` indicates which forced-prefix shape was
|
||||||
|
requested (the caller knows from ``request.timestamp_granularities``);
|
||||||
|
adapters use it to disambiguate variants whose detokenized prefix
|
||||||
|
differs in shape from their token-id prefix.
|
||||||
|
|
||||||
|
* ``(None, None)`` — the forced prefix is not yet locatable.
|
||||||
|
Streaming callers keep buffering; non-streaming / end-of-stream
|
||||||
|
callers treat this as a parse failure and fall back to
|
||||||
|
``strip_special_tokens`` on the raw text.
|
||||||
|
* ``(lang, visible)`` — prefix parsed. ``visible`` is fully
|
||||||
|
user-visible (prefix removed, embedded special tokens scrubbed).
|
||||||
|
It must grow monotonically across cumulative streaming snapshots
|
||||||
|
so callers can compute deltas against it directly.
|
||||||
|
"""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def strip_special_tokens(text: str) -> str:
|
||||||
|
"""Best-effort scrub of model-specific special-token strings.
|
||||||
|
|
||||||
|
Used as a fallback when ``parse_fused_output`` reports a parse
|
||||||
|
failure (e.g. FSM abort). Default is an identity pass-through;
|
||||||
|
adapters that request generation with ``skip_special_tokens=False``
|
||||||
|
should override to strip their special-token syntax.
|
||||||
|
"""
|
||||||
|
return text
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def supports_chunked_streaming(self) -> bool:
|
def supports_chunked_streaming(self) -> bool:
|
||||||
"""Whether this model uses chunk-based streaming instead of token-level streaming."""
|
"""Whether this model uses chunk-based streaming instead of token-level streaming."""
|
||||||
|
|||||||
@@ -1,6 +1,10 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from typing import List
|
import logging
|
||||||
|
import re
|
||||||
|
from typing import List, Optional
|
||||||
|
|
||||||
|
from transformers.models.whisper.tokenization_whisper import LANGUAGES
|
||||||
|
|
||||||
from sglang.srt.entrypoints.openai.protocol import (
|
from sglang.srt.entrypoints.openai.protocol import (
|
||||||
TranscriptionRequest,
|
TranscriptionRequest,
|
||||||
@@ -13,6 +17,109 @@ from sglang.srt.entrypoints.openai.transcription_adapters.base import (
|
|||||||
register_transcription_adapter,
|
register_transcription_adapter,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Sampling-params key the adapter plants and the multimodal processor pops
|
||||||
|
# to flip the decoder prompt from the explicit 4-token forced sequence to
|
||||||
|
# the bare ``<|startoftranscript|>`` (so the FSM regex drives token 1-3
|
||||||
|
# instead). Centralized so adapter / processor / warmup all reference the
|
||||||
|
# same string.
|
||||||
|
FUSED_AUTODETECT_FLAG = "_detect_language"
|
||||||
|
|
||||||
|
# The complete set of Whisper language tokens as they appear in the tokenizer
|
||||||
|
# vocab (<|xx|> / <|xxx|>). Sourced from the upstream ``LANGUAGES`` dict in
|
||||||
|
# ``transformers.models.whisper.tokenization_whisper`` so newly-added tokens
|
||||||
|
# (e.g. ``yue`` in Whisper v3) automatically propagate.
|
||||||
|
#
|
||||||
|
# Intentionally wider than ``processors.whisper.ISO639_1_SUPPORTED_LANGS``
|
||||||
|
# (the narrower input-validation set used by ``normalize_language_to_code``)
|
||||||
|
# — for the FSM regex we want every language the model was trained on so we
|
||||||
|
# don't silently force a wrong nearest-match code on audio in languages the
|
||||||
|
# model *can* detect but the input dict doesn't list (yue/Cantonese,
|
||||||
|
# jw/Javanese, haw/Hawaiian, ba/Bashkir, su/Sundanese, ...). Codes whose
|
||||||
|
# ``<|xxx|>`` token isn't in an older checkpoint's vocab are harmless —
|
||||||
|
# xgrammar simply leaves that regex branch with no admissible tokens.
|
||||||
|
WHISPER_LANG_TOKEN_CODES: frozenset[str] = frozenset(LANGUAGES.keys())
|
||||||
|
|
||||||
|
# Two forced-prefix variants, picked at request build time based on whether
|
||||||
|
# the client asked for timestamp_granularities:
|
||||||
|
# * notimestamps variant: <|lang|><|transcribe|><|notimestamps|> text...
|
||||||
|
# — drops segment/word timing, used when the client doesn't request it.
|
||||||
|
# * timestamps variant: <|lang|><|transcribe|><|0.00|> text <|X.XX|> ...
|
||||||
|
# — <|0.00|> anchors the first segment at t=0, and the model naturally
|
||||||
|
# emits further timestamp tokens between segments. _parse_segments
|
||||||
|
# reconstructs segments from output_ids afterwards.
|
||||||
|
# sorted() gives a deterministic regex string so the warmup-compiled FSM is
|
||||||
|
# reused across server restarts.
|
||||||
|
_LANG_ALT = "|".join(re.escape(c) for c in sorted(WHISPER_LANG_TOKEN_CODES))
|
||||||
|
_LANG_PREFIX = r"<\|(" + _LANG_ALT + r")\|>"
|
||||||
|
WHISPER_AUTODETECT_REGEX = (
|
||||||
|
_LANG_PREFIX + r"<\|transcribe\|>" + r"<\|notimestamps\|>" + r"[\s\S]*"
|
||||||
|
)
|
||||||
|
WHISPER_AUTODETECT_TS_REGEX = (
|
||||||
|
_LANG_PREFIX + r"<\|transcribe\|>" + r"<\|0\.00\|>" + r"[\s\S]*"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Forced-prefix patterns, one per FSM variant. Each is anchored at start
|
||||||
|
# and rejects anything missing ``<|transcribe|>`` so a bypassed FSM or a
|
||||||
|
# mid-stream snapshot can't slip through as a valid detection. The two
|
||||||
|
# patterns differ in what the third forced token is decoded *as*:
|
||||||
|
#
|
||||||
|
# * ``_FUSED_PREFIX_RE_NOTS`` — non-timestamps variant. The third token
|
||||||
|
# is ``<|notimestamps|>`` (id 50364), which detokenizes to its literal
|
||||||
|
# string. Mid-stream snapshots stuck at ``<|en|><|transcribe|>`` (the
|
||||||
|
# third token hasn't fired yet) correctly miss this regex, so the
|
||||||
|
# streaming handler can detect FSM-abort and surface an error.
|
||||||
|
#
|
||||||
|
# * ``_FUSED_PREFIX_RE_TS`` — timestamps variant. The third token is
|
||||||
|
# ``<|0.00|>`` (id 50365), which Whisper's tokenizer decodes to the
|
||||||
|
# *empty string* even with ``skip_special_tokens=False`` (only
|
||||||
|
# ``<|notimestamps|>`` survives detokenization; every ``<|X.XX|>``
|
||||||
|
# maps to ``""``). So the regex must accept just
|
||||||
|
# ``<|en|><|transcribe|>`` and rely on the FSM having already
|
||||||
|
# constrained ``output_ids[2] == 50365``. ``_parse_segments`` reads
|
||||||
|
# the timestamps from ``output_ids`` directly, so segment timing is
|
||||||
|
# unaffected.
|
||||||
|
_FUSED_PREFIX_RE_NOTS = re.compile(
|
||||||
|
r"^" + _LANG_PREFIX + r"<\|transcribe\|><\|notimestamps\|>"
|
||||||
|
)
|
||||||
|
_FUSED_PREFIX_RE_TS = re.compile(r"^" + _LANG_PREFIX + r"<\|transcribe\|>")
|
||||||
|
|
||||||
|
# Fixed Whisper control tokens (see transformers.models.whisper vocab).
|
||||||
|
# <|startoftranscript|> / <|startofprev|> / <|startoflm|> only appear at
|
||||||
|
# the decoder prompt and never in generated output, but they are cheap to
|
||||||
|
# include and harmless if they ever leak.
|
||||||
|
_WHISPER_CONTROL_TOKENS = frozenset(
|
||||||
|
{
|
||||||
|
"endoftext",
|
||||||
|
"startoftranscript",
|
||||||
|
"startofprev",
|
||||||
|
"startoflm",
|
||||||
|
"translate",
|
||||||
|
"transcribe",
|
||||||
|
"notimestamps",
|
||||||
|
"nospeech",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
# Scrubs only actual Whisper special-token literals: language codes
|
||||||
|
# (WHISPER_LANG_TOKEN_CODES), control tokens (_WHISPER_CONTROL_TOKENS),
|
||||||
|
# and timestamp tokens (<|X.XX|>, where X.XX matches the
|
||||||
|
# ``{ts_base + k * 0.02}`` schema the model emits). A broad
|
||||||
|
# ``<\|[^|]+\|>`` would eat legitimate spoken content on audio that
|
||||||
|
# pronounces angle-bracket / pipe sequences (e.g. someone reading
|
||||||
|
# ``<|endoftext|>`` out loud). Used to scrub trailing ``<|endoftext|>``
|
||||||
|
# and embedded ``<|X.XX|>`` timestamp tokens from the user-visible text
|
||||||
|
# in fused-autodetect responses, where ``skip_special_tokens=False`` is
|
||||||
|
# needed to preserve the language prefix for parsing but would otherwise
|
||||||
|
# leak other special tokens downstream.
|
||||||
|
_SPECIAL_TOKEN_RE = re.compile(
|
||||||
|
r"<\|(?:"
|
||||||
|
+ "|".join(sorted(WHISPER_LANG_TOKEN_CODES | _WHISPER_CONTROL_TOKENS))
|
||||||
|
+ r"|\d+\.\d{2}"
|
||||||
|
+ r")\|>"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@register_transcription_adapter("Whisper")
|
@register_transcription_adapter("Whisper")
|
||||||
class WhisperAdapter(TranscriptionAdapter):
|
class WhisperAdapter(TranscriptionAdapter):
|
||||||
@@ -29,6 +136,108 @@ class WhisperAdapter(TranscriptionAdapter):
|
|||||||
params["timestamp_granularities"] = request.timestamp_granularities
|
params["timestamp_granularities"] = request.timestamp_granularities
|
||||||
return params
|
return params
|
||||||
|
|
||||||
|
# -- language detection ------------------------------------------------
|
||||||
|
|
||||||
|
@property
|
||||||
|
def supports_language_detection(self) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
def build_fused_autodetect_params(self, request: TranscriptionRequest) -> dict:
|
||||||
|
"""Build sampling params for a single fused detect+transcribe request.
|
||||||
|
|
||||||
|
Uses SGLang's native structured generation (``regex``) to constrain
|
||||||
|
the first 3 decode tokens. Picks the regex variant based on whether
|
||||||
|
the client requested ``timestamp_granularities``:
|
||||||
|
|
||||||
|
* no timestamps: ``<|lang|><|transcribe|><|notimestamps|>text``
|
||||||
|
* with timestamps: ``<|lang|><|transcribe|><|0.00|>text<|X.XX|>...``
|
||||||
|
— ``<|0.00|>`` anchors segment 0 at t=0; the model naturally
|
||||||
|
emits further timestamp tokens between segments and
|
||||||
|
``_parse_segments`` reconstructs them from ``output_ids``.
|
||||||
|
|
||||||
|
Either way, detection and transcription run in a single encoder
|
||||||
|
pass with no extra HTTP round-trip.
|
||||||
|
"""
|
||||||
|
ts_variant = bool(request.timestamp_granularities)
|
||||||
|
params: dict = {
|
||||||
|
"temperature": request.temperature,
|
||||||
|
# Fused auto-detect decoder prompt is just <|startoftranscript|>
|
||||||
|
# (1 token, see processors/whisper.py). Whisper's
|
||||||
|
# max_target_positions is 448, so max_new_tokens caps at 447:
|
||||||
|
# 1 prompt + 3 forced prefix + up to 444 free transcription = 448.
|
||||||
|
"max_new_tokens": 447,
|
||||||
|
"regex": (
|
||||||
|
WHISPER_AUTODETECT_TS_REGEX if ts_variant else WHISPER_AUTODETECT_REGEX
|
||||||
|
),
|
||||||
|
"skip_special_tokens": False,
|
||||||
|
# parse_fused_output matches a zero-space forced prefix
|
||||||
|
# (``<|en|><|transcribe|><|notimestamps|>`` glued together).
|
||||||
|
# Fast Whisper tokenizers decode adjacent added tokens with no
|
||||||
|
# space, but slow ones insert a space between them. Force
|
||||||
|
# spaces_between_special_tokens=False so the parse regex is
|
||||||
|
# correct regardless of tokenizer variant.
|
||||||
|
"spaces_between_special_tokens": False,
|
||||||
|
FUSED_AUTODETECT_FLAG: True,
|
||||||
|
}
|
||||||
|
if ts_variant:
|
||||||
|
params["timestamp_granularities"] = request.timestamp_granularities
|
||||||
|
return params
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def parse_fused_output(
|
||||||
|
text: str, *, ts_variant: bool = False
|
||||||
|
) -> tuple[Optional[str], Optional[str]]:
|
||||||
|
"""Parse fused output into ``(language_code, user_visible_text)``.
|
||||||
|
|
||||||
|
Matches the forced prefix the FSM emits. ``ts_variant`` selects
|
||||||
|
which shape to expect — the caller knows from
|
||||||
|
``request.timestamp_granularities`` which regex was sent to the
|
||||||
|
FSM and so which decoded shape to look for:
|
||||||
|
|
||||||
|
* ``ts_variant=False`` — ``<|en|><|transcribe|><|notimestamps|> Hello...``
|
||||||
|
* ``ts_variant=True`` — ``<|en|><|transcribe|> Hello...`` (``<|0.00|>``
|
||||||
|
is in ``output_ids`` but Whisper detokenizes it to the empty string).
|
||||||
|
|
||||||
|
Return cases:
|
||||||
|
|
||||||
|
* ``(None, None)`` — the prefix isn't fully in yet (mid-stream
|
||||||
|
snapshot before ``<|transcribe|>`` lands, or before
|
||||||
|
``<|notimestamps|>`` lands in the no-ts variant) or the prefix
|
||||||
|
is malformed. Streaming callers keep buffering; non-streaming /
|
||||||
|
end-of-stream callers treat this as a parse failure and fall
|
||||||
|
back to a best-effort scrub of the raw text.
|
||||||
|
* ``(lang, visible)`` — prefix fully parsed. ``visible`` is the
|
||||||
|
transcription with the forced prefix removed, any embedded
|
||||||
|
special tokens (``<|X.XX|>``, ``<|endoftext|>``) scrubbed, and
|
||||||
|
surrounding whitespace trimmed. It grows monotonically across
|
||||||
|
streaming chunks because Whisper's special tokens detokenize
|
||||||
|
atomically, so callers can compute deltas against it directly.
|
||||||
|
"""
|
||||||
|
pattern = _FUSED_PREFIX_RE_TS if ts_variant else _FUSED_PREFIX_RE_NOTS
|
||||||
|
m = pattern.match(text)
|
||||||
|
if not m:
|
||||||
|
return None, None
|
||||||
|
transcription = text[m.end() :]
|
||||||
|
# Scrub any remaining special tokens. skip_special_tokens=False is
|
||||||
|
# set on fused requests so the language prefix survives for
|
||||||
|
# parsing, but that also preserves trailing <|endoftext|> and, in
|
||||||
|
# the timestamps variant, embedded <|X.XX|> segment tokens. Those
|
||||||
|
# are unwanted in the user-visible text (verbose_json gets its
|
||||||
|
# segments from _parse_segments over output_ids instead).
|
||||||
|
transcription = _SPECIAL_TOKEN_RE.sub("", transcription)
|
||||||
|
return m.group(1), transcription.strip()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def strip_special_tokens(text: str) -> str:
|
||||||
|
"""Remove any ``<|...|>`` special-token strings from *text*.
|
||||||
|
|
||||||
|
Used as the best-effort scrub on FSM abort / parse failure when
|
||||||
|
the full ``parse_fused_output`` path can't locate the prefix.
|
||||||
|
"""
|
||||||
|
return _SPECIAL_TOKEN_RE.sub("", text)
|
||||||
|
|
||||||
|
# -- end language detection --------------------------------------------
|
||||||
|
|
||||||
def build_verbose_response(
|
def build_verbose_response(
|
||||||
self,
|
self,
|
||||||
request: TranscriptionRequest,
|
request: TranscriptionRequest,
|
||||||
@@ -40,7 +249,11 @@ class WhisperAdapter(TranscriptionAdapter):
|
|||||||
output_ids = ret.get("output_ids", [])
|
output_ids = ret.get("output_ids", [])
|
||||||
parsed_text, segments = self._parse_segments(output_ids, tokenizer)
|
parsed_text, segments = self._parse_segments(output_ids, tokenizer)
|
||||||
return TranscriptionVerboseResponse(
|
return TranscriptionVerboseResponse(
|
||||||
language=request.language or "en",
|
# Pass None through when fused auto-detect failed to parse a
|
||||||
|
# language — the client should see detection-failed, not a silent
|
||||||
|
# English default. For explicit-language requests request.language
|
||||||
|
# is already set by the caller.
|
||||||
|
language=request.language,
|
||||||
duration=round(request.audio_duration_s, 2),
|
duration=round(request.audio_duration_s, 2),
|
||||||
text=parsed_text or text,
|
text=parsed_text or text,
|
||||||
segments=segments,
|
segments=segments,
|
||||||
|
|||||||
@@ -38,6 +38,73 @@ async def execute_warmups(
|
|||||||
await _warmup_registry[warmup_name](disaggregation_mode, tokenizer_manager)
|
await _warmup_registry[warmup_name](disaggregation_mode, tokenizer_manager)
|
||||||
|
|
||||||
|
|
||||||
|
@warmup("whisper_autodetect")
|
||||||
|
async def whisper_autodetect(
|
||||||
|
disaggregation_mode: str, tokenizer_manager: TokenizerManager
|
||||||
|
):
|
||||||
|
"""Pre-compile the xgrammar FSM for both Whisper auto-detect regexes.
|
||||||
|
|
||||||
|
The first request that uses each structured-generation regex incurs a
|
||||||
|
~15-20s compilation cost. xgrammar caches compiled grammars by the
|
||||||
|
exact regex string, so we warm both the notimestamps and timestamps
|
||||||
|
variants here — otherwise the first ``language=None +
|
||||||
|
timestamp_granularities`` request would still pay the full spike.
|
||||||
|
"""
|
||||||
|
# A short silent audio encoded as base64 WAV (0.1s, 16kHz, mono) —
|
||||||
|
# soundfile produces the WAV header + PCM data from a list of floats.
|
||||||
|
import base64
|
||||||
|
import io
|
||||||
|
|
||||||
|
import soundfile as sf
|
||||||
|
|
||||||
|
from sglang.srt.entrypoints.openai.transcription_adapters.whisper import (
|
||||||
|
FUSED_AUTODETECT_FLAG,
|
||||||
|
WHISPER_AUTODETECT_REGEX,
|
||||||
|
WHISPER_AUTODETECT_TS_REGEX,
|
||||||
|
)
|
||||||
|
|
||||||
|
sr, dur = 16000, 0.1
|
||||||
|
n = int(sr * dur)
|
||||||
|
buf = io.BytesIO()
|
||||||
|
sf.write(buf, [0.0] * n, sr, format="WAV")
|
||||||
|
audio_b64 = base64.b64encode(buf.getvalue()).decode()
|
||||||
|
audio_data_uri = f"data:audio/wav;base64,{audio_b64}"
|
||||||
|
|
||||||
|
for variant_name, regex in (
|
||||||
|
("notimestamps", WHISPER_AUTODETECT_REGEX),
|
||||||
|
("timestamps", WHISPER_AUTODETECT_TS_REGEX),
|
||||||
|
):
|
||||||
|
logger.info(
|
||||||
|
"Compiling Whisper auto-detect regex FSM (%s, one-time, ~15-20s)...",
|
||||||
|
variant_name,
|
||||||
|
)
|
||||||
|
req = GenerateReqInput(
|
||||||
|
text="",
|
||||||
|
audio_data=audio_data_uri,
|
||||||
|
sampling_params={
|
||||||
|
"max_new_tokens": 4,
|
||||||
|
"temperature": 0,
|
||||||
|
"regex": regex,
|
||||||
|
"skip_special_tokens": False,
|
||||||
|
"spaces_between_special_tokens": False,
|
||||||
|
FUSED_AUTODETECT_FLAG: True,
|
||||||
|
},
|
||||||
|
modalities=["audio"],
|
||||||
|
)
|
||||||
|
# PD prefill servers assert req.bootstrap_room is not None in the
|
||||||
|
# default follow_bootstrap_room scheduler; the fake values match
|
||||||
|
# what the voice_chat warmup uses for the same reason.
|
||||||
|
if disaggregation_mode != "null":
|
||||||
|
req.bootstrap_room = 0
|
||||||
|
req.bootstrap_host = FAKE_BOOTSTRAP_HOST
|
||||||
|
# Drain the generator so the FSM is fully installed and any
|
||||||
|
# downstream exception surfaces instead of being swallowed after
|
||||||
|
# the first yield.
|
||||||
|
async for _ in tokenizer_manager.generate_request(req, None):
|
||||||
|
pass
|
||||||
|
logger.info("Whisper auto-detect regex FSMs compiled.")
|
||||||
|
|
||||||
|
|
||||||
@warmup("voice_chat")
|
@warmup("voice_chat")
|
||||||
async def voice_chat(disaggregation_mode: str, tokenizer_manager: TokenizerManager):
|
async def voice_chat(disaggregation_mode: str, tokenizer_manager: TokenizerManager):
|
||||||
# this warms up the fused_moe triton kernels and caches them
|
# this warms up the fused_moe triton kernels and caches them
|
||||||
|
|||||||
@@ -1,6 +1,9 @@
|
|||||||
import logging
|
import logging
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any, Dict, Optional
|
||||||
|
|
||||||
|
from sglang.srt.entrypoints.openai.transcription_adapters.whisper import (
|
||||||
|
FUSED_AUTODETECT_FLAG,
|
||||||
|
)
|
||||||
from sglang.srt.managers.schedule_batch import (
|
from sglang.srt.managers.schedule_batch import (
|
||||||
Modality,
|
Modality,
|
||||||
MultimodalDataItem,
|
MultimodalDataItem,
|
||||||
@@ -86,10 +89,14 @@ def normalize_language_to_code(language: Optional[str]) -> Optional[str]:
|
|||||||
|
|
||||||
Args:
|
Args:
|
||||||
language: Language as full name (e.g., 'English', 'Spanish') or
|
language: Language as full name (e.g., 'English', 'Spanish') or
|
||||||
ISO 639-1 code (e.g., 'en', 'es')
|
ISO 639-1 code (e.g., 'en', 'es'). Three-letter Whisper
|
||||||
|
codes the model supports but that aren't in
|
||||||
|
ISO639_1_SUPPORTED_LANGS (e.g., 'yue', 'haw', 'jw') are
|
||||||
|
also accepted so that a code returned by fused autodetect
|
||||||
|
round-trips cleanly when reused as ``language=`` later.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
ISO 639-1 code or None if input is None
|
Whisper language code or None if input is None
|
||||||
"""
|
"""
|
||||||
if language is None:
|
if language is None:
|
||||||
return None
|
return None
|
||||||
@@ -104,6 +111,18 @@ def normalize_language_to_code(language: Optional[str]) -> Optional[str]:
|
|||||||
if language_lower in LANG_NAME_TO_CODE:
|
if language_lower in LANG_NAME_TO_CODE:
|
||||||
return LANG_NAME_TO_CODE[language_lower]
|
return LANG_NAME_TO_CODE[language_lower]
|
||||||
|
|
||||||
|
# Fused autodetect's FSM regex covers the full Whisper language-token
|
||||||
|
# vocab (see WHISPER_LANG_TOKEN_CODES), which is wider than the
|
||||||
|
# English-name-keyed ISO639_1_SUPPORTED_LANGS dict. Accept any code in
|
||||||
|
# that wider set too so that detection -> reuse-as-input round-trips.
|
||||||
|
# Lazy import to avoid top-level cycle with the openai entrypoint.
|
||||||
|
from sglang.srt.entrypoints.openai.transcription_adapters.whisper import (
|
||||||
|
WHISPER_LANG_TOKEN_CODES,
|
||||||
|
)
|
||||||
|
|
||||||
|
if language_lower in WHISPER_LANG_TOKEN_CODES:
|
||||||
|
return language_lower
|
||||||
|
|
||||||
# Not recognized
|
# Not recognized
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Language '{language}' not recognized. "
|
f"Language '{language}' not recognized. "
|
||||||
@@ -128,7 +147,21 @@ class WhisperProcessor(BaseMultimodalProcessor):
|
|||||||
if language is None:
|
if language is None:
|
||||||
language = "en" # Default to English
|
language = "en" # Default to English
|
||||||
language_token = f"<|{language}|>"
|
language_token = f"<|{language}|>"
|
||||||
return self._tokenizer.convert_tokens_to_ids(language_token)
|
token_id = self._tokenizer.convert_tokens_to_ids(language_token)
|
||||||
|
# normalize_language_to_code accepts the full Whisper language-token
|
||||||
|
# vocab (including yue/haw/jw) so fused autodetect output round-trips.
|
||||||
|
# Older checkpoints (v1/v2) don't have every newer token in their
|
||||||
|
# vocab, in which case convert_tokens_to_ids returns the unk id.
|
||||||
|
# Raise a clean error here instead of silently feeding unk into the
|
||||||
|
# decoder and producing garbage.
|
||||||
|
unk_id = getattr(self._tokenizer, "unk_token_id", None)
|
||||||
|
if token_id is None or (unk_id is not None and token_id == unk_id):
|
||||||
|
raise ValueError(
|
||||||
|
f"Language '{language}' is not in this Whisper model's vocabulary. "
|
||||||
|
f"The '{language_token}' token may have been added in a later "
|
||||||
|
f"Whisper version than the loaded checkpoint."
|
||||||
|
)
|
||||||
|
return token_id
|
||||||
|
|
||||||
async def process_mm_data_async(
|
async def process_mm_data_async(
|
||||||
self,
|
self,
|
||||||
@@ -146,41 +179,20 @@ class WhisperProcessor(BaseMultimodalProcessor):
|
|||||||
f"Whisper expects exactly 1 audio input, got {len(audio_data)}"
|
f"Whisper expects exactly 1 audio input, got {len(audio_data)}"
|
||||||
)
|
)
|
||||||
|
|
||||||
audios = [load_audio(audio) for audio in audio_data]
|
# Check if this is a fused auto-detect request (decoder prompt = [SOT] only,
|
||||||
|
# structured generation handles the rest via regex constraint).
|
||||||
# For Whisper, ALWAYS use the proper transcription token sequence
|
detect_language = self._pop_sampling_param(request_obj, FUSED_AUTODETECT_FLAG)
|
||||||
# and IGNORE any text prompt - Whisper is a pure speech-to-text model
|
# timestamp_granularities is a transcription-level field; it must be
|
||||||
# The decoder_start_token_id and forced_decoder_ids from generation config
|
# popped in both branches or it leaks into SamplingParams(**kwargs)
|
||||||
# set up: <|startoftranscript|> <|lang|> <|task|> [<|notimestamps|> or <|0.00|>]
|
# downstream and TypeErrors. In the fused branch the FSM regex was
|
||||||
|
# already picked in build_fused_autodetect_params based on this value,
|
||||||
language = normalize_language_to_code(
|
# so we only need to keep it here to pick the timestamp_token_id for
|
||||||
self._pop_sampling_param(request_obj, "language")
|
# the explicit-language branch.
|
||||||
)
|
|
||||||
language_token_id = self._get_language_token_id(language)
|
|
||||||
timestamp_granularities = self._pop_sampling_param(
|
timestamp_granularities = self._pop_sampling_param(
|
||||||
request_obj, "timestamp_granularities"
|
request_obj, "timestamp_granularities"
|
||||||
)
|
)
|
||||||
|
|
||||||
# Build decoder input tokens
|
audios = [load_audio(audio) for audio in audio_data]
|
||||||
decoder_start_token_id = getattr(
|
|
||||||
self.hf_config, "decoder_start_token_id", 50258
|
|
||||||
)
|
|
||||||
transcribe_token_id = self._tokenizer.convert_tokens_to_ids("<|transcribe|>")
|
|
||||||
|
|
||||||
# Use <|0.00|> to enable timestamp generation, or <|notimestamps|> to disable
|
|
||||||
if timestamp_granularities:
|
|
||||||
timestamp_token_id = self._tokenizer.convert_tokens_to_ids("<|0.00|>")
|
|
||||||
else:
|
|
||||||
timestamp_token_id = self._tokenizer.convert_tokens_to_ids(
|
|
||||||
"<|notimestamps|>"
|
|
||||||
)
|
|
||||||
|
|
||||||
input_ids = [
|
|
||||||
decoder_start_token_id,
|
|
||||||
language_token_id,
|
|
||||||
transcribe_token_id,
|
|
||||||
timestamp_token_id,
|
|
||||||
]
|
|
||||||
|
|
||||||
# Whisper expects input features padded to max_length (3000 frames = 30 seconds)
|
# Whisper expects input features padded to max_length (3000 frames = 30 seconds)
|
||||||
# This is the standard context length for Whisper
|
# This is the standard context length for Whisper
|
||||||
@@ -191,6 +203,47 @@ class WhisperProcessor(BaseMultimodalProcessor):
|
|||||||
return_tensors="pt",
|
return_tensors="pt",
|
||||||
)["input_features"][0]
|
)["input_features"][0]
|
||||||
|
|
||||||
|
# Whisper is a pure speech-to-text model; text prompts are ignored.
|
||||||
|
# The full decoder sequence is:
|
||||||
|
# <|startoftranscript|> <|lang|> <|transcribe|> [<|notimestamps|> | <|0.00|>]
|
||||||
|
#
|
||||||
|
# When language is known, we build this prefix explicitly below.
|
||||||
|
# When auto-detecting (_detect_language=True), we feed only <|startoftranscript|>
|
||||||
|
# and let SGLang's structured generation (regex) constrain the model to produce
|
||||||
|
# <|lang|><|transcribe|><|notimestamps|> as the first 3 decode tokens — this is
|
||||||
|
# equivalent to HuggingFace's forced_decoder_ids but uses SGLang's native API.
|
||||||
|
|
||||||
|
decoder_start_token_id = getattr(
|
||||||
|
self.hf_config, "decoder_start_token_id", 50258
|
||||||
|
)
|
||||||
|
|
||||||
|
if detect_language:
|
||||||
|
input_ids = [decoder_start_token_id]
|
||||||
|
else:
|
||||||
|
language = normalize_language_to_code(
|
||||||
|
self._pop_sampling_param(request_obj, "language")
|
||||||
|
)
|
||||||
|
language_token_id = self._get_language_token_id(language)
|
||||||
|
|
||||||
|
transcribe_token_id = self._tokenizer.convert_tokens_to_ids(
|
||||||
|
"<|transcribe|>"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Use <|0.00|> to enable timestamp generation, or <|notimestamps|> to disable
|
||||||
|
if timestamp_granularities:
|
||||||
|
timestamp_token_id = self._tokenizer.convert_tokens_to_ids("<|0.00|>")
|
||||||
|
else:
|
||||||
|
timestamp_token_id = self._tokenizer.convert_tokens_to_ids(
|
||||||
|
"<|notimestamps|>"
|
||||||
|
)
|
||||||
|
|
||||||
|
input_ids = [
|
||||||
|
decoder_start_token_id,
|
||||||
|
language_token_id,
|
||||||
|
transcribe_token_id,
|
||||||
|
timestamp_token_id,
|
||||||
|
]
|
||||||
|
|
||||||
return MultimodalProcessorOutput(
|
return MultimodalProcessorOutput(
|
||||||
input_ids=input_ids,
|
input_ids=input_ids,
|
||||||
mm_items=[
|
mm_items=[
|
||||||
|
|||||||
@@ -6,7 +6,9 @@ Usage:
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import io
|
import io
|
||||||
|
import json
|
||||||
import unittest
|
import unittest
|
||||||
|
from typing import List, Optional
|
||||||
|
|
||||||
import requests
|
import requests
|
||||||
|
|
||||||
@@ -19,7 +21,7 @@ from sglang.test.test_utils import (
|
|||||||
popen_launch_server,
|
popen_launch_server,
|
||||||
)
|
)
|
||||||
|
|
||||||
register_cuda_ci(est_time=51, suite="stage-b-test-1-gpu-small")
|
register_cuda_ci(est_time=60, suite="stage-b-test-1-gpu-small")
|
||||||
|
|
||||||
WHISPER_MODEL = "openai/whisper-large-v3"
|
WHISPER_MODEL = "openai/whisper-large-v3"
|
||||||
AUDIO_URL = "https://raw.githubusercontent.com/sgl-project/sgl-test-files/refs/heads/main/audios/Trump_WEF_2018_10s.mp3"
|
AUDIO_URL = "https://raw.githubusercontent.com/sgl-project/sgl-test-files/refs/heads/main/audios/Trump_WEF_2018_10s.mp3"
|
||||||
@@ -54,20 +56,65 @@ class TestServingTranscription(CustomTestCase):
|
|||||||
if hasattr(cls, "process") and cls.process:
|
if hasattr(cls, "process") and cls.process:
|
||||||
kill_process_tree(cls.process.pid)
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
def _transcribe(self, language="en"):
|
def _transcribe(
|
||||||
"""Send a transcription request and return the JSON response."""
|
self,
|
||||||
|
language: Optional[str] = "en",
|
||||||
|
response_format: Optional[str] = None,
|
||||||
|
timestamp_granularities: Optional[List[str]] = None,
|
||||||
|
):
|
||||||
|
"""Send a non-streaming transcription request and return the JSON response.
|
||||||
|
|
||||||
|
Passing ``language=None`` omits the field entirely, which exercises
|
||||||
|
the fused auto-detect path.
|
||||||
|
"""
|
||||||
audio_bytes = download_audio_bytes()
|
audio_bytes = download_audio_bytes()
|
||||||
|
data = {"model": "whisper"}
|
||||||
|
if language is not None:
|
||||||
|
data["language"] = language
|
||||||
|
if response_format is not None:
|
||||||
|
data["response_format"] = response_format
|
||||||
|
if timestamp_granularities is not None:
|
||||||
|
# Form-encoded list fields repeat the key
|
||||||
|
data["timestamp_granularities[]"] = timestamp_granularities
|
||||||
response = requests.post(
|
response = requests.post(
|
||||||
self.base_url + "/v1/audio/transcriptions",
|
self.base_url + "/v1/audio/transcriptions",
|
||||||
files={"file": ("audio.mp3", io.BytesIO(audio_bytes), "audio/mpeg")},
|
files={"file": ("audio.mp3", io.BytesIO(audio_bytes), "audio/mpeg")},
|
||||||
data={
|
data=data,
|
||||||
"model": "whisper",
|
|
||||||
"language": language,
|
|
||||||
},
|
|
||||||
)
|
)
|
||||||
self.assertEqual(response.status_code, 200, response.text)
|
self.assertEqual(response.status_code, 200, response.text)
|
||||||
return response.json()
|
return response.json()
|
||||||
|
|
||||||
|
def _transcribe_stream(self, language: Optional[str] = None) -> List[str]:
|
||||||
|
"""Send a streaming transcription request and return the delta strings."""
|
||||||
|
audio_bytes = download_audio_bytes()
|
||||||
|
data = {"model": "whisper", "stream": "true"}
|
||||||
|
if language is not None:
|
||||||
|
data["language"] = language
|
||||||
|
with requests.post(
|
||||||
|
self.base_url + "/v1/audio/transcriptions",
|
||||||
|
files={"file": ("audio.mp3", io.BytesIO(audio_bytes), "audio/mpeg")},
|
||||||
|
data=data,
|
||||||
|
stream=True,
|
||||||
|
timeout=120,
|
||||||
|
) as response:
|
||||||
|
self.assertEqual(response.status_code, 200, response.text)
|
||||||
|
deltas: List[str] = []
|
||||||
|
for raw in response.iter_lines():
|
||||||
|
if not raw:
|
||||||
|
continue
|
||||||
|
line = raw.decode("utf-8")
|
||||||
|
if not line.startswith("data: "):
|
||||||
|
continue
|
||||||
|
payload = line[len("data: ") :].strip()
|
||||||
|
if payload == "[DONE]":
|
||||||
|
break
|
||||||
|
obj = json.loads(payload)
|
||||||
|
for choice in obj.get("choices", []):
|
||||||
|
content = (choice.get("delta") or {}).get("content")
|
||||||
|
if content:
|
||||||
|
deltas.append(content)
|
||||||
|
return deltas
|
||||||
|
|
||||||
def test_basic_transcription(self):
|
def test_basic_transcription(self):
|
||||||
"""Test that transcription returns a valid non-empty response."""
|
"""Test that transcription returns a valid non-empty response."""
|
||||||
result = self._transcribe()
|
result = self._transcribe()
|
||||||
@@ -103,6 +150,81 @@ class TestServingTranscription(CustomTestCase):
|
|||||||
f"Transcription {i + 1} differs from first transcription",
|
f"Transcription {i + 1} differs from first transcription",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# -- fused auto-detect (language=None) ---------------------------------
|
||||||
|
# The clip is English, so the fused path must both produce a valid
|
||||||
|
# transcription AND expose "en" as the detected language. None of the
|
||||||
|
# deltas / text fields should leak Whisper special tokens.
|
||||||
|
|
||||||
|
def test_auto_detect_language_verbose_json(self):
|
||||||
|
"""language omitted + verbose_json returns detected language + clean text."""
|
||||||
|
result = self._transcribe(language=None, response_format="verbose_json")
|
||||||
|
self.assertEqual(result.get("language"), "en")
|
||||||
|
text = result.get("text", "")
|
||||||
|
self.assertTrue(len(text) > 0, "Transcription should not be empty")
|
||||||
|
self.assertNotIn("<|", text, f"Special token leaked into text: {text!r}")
|
||||||
|
# Sanity-check content against the same keywords the English test uses.
|
||||||
|
keywords = ["privilege", "leader", "science", "art"]
|
||||||
|
matches = [kw for kw in keywords if kw in text.lower()]
|
||||||
|
self.assertGreaterEqual(
|
||||||
|
len(matches),
|
||||||
|
2,
|
||||||
|
f"Expected at least 2 of {keywords} in auto-detected transcription, "
|
||||||
|
f"found {matches}. Full text: {text!r}",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_auto_detect_matches_explicit_english(self):
|
||||||
|
"""Auto-detected (language=None) text should match explicit language=en."""
|
||||||
|
auto = self._transcribe(language=None).get("text", "")
|
||||||
|
explicit = self._transcribe(language="en").get("text", "")
|
||||||
|
self.assertEqual(
|
||||||
|
auto.strip(),
|
||||||
|
explicit.strip(),
|
||||||
|
"Auto-detect should produce the same transcription as language=en "
|
||||||
|
"on an English clip.",
|
||||||
|
)
|
||||||
|
self.assertNotIn("<|", auto)
|
||||||
|
|
||||||
|
def test_auto_detect_with_segment_timestamps(self):
|
||||||
|
"""language=None + timestamp_granularities uses the timestamps fused regex."""
|
||||||
|
result = self._transcribe(
|
||||||
|
language=None,
|
||||||
|
response_format="verbose_json",
|
||||||
|
timestamp_granularities=["segment"],
|
||||||
|
)
|
||||||
|
self.assertEqual(result.get("language"), "en")
|
||||||
|
segments = result.get("segments") or []
|
||||||
|
self.assertGreater(len(segments), 0, "Expected at least one segment")
|
||||||
|
for seg in segments:
|
||||||
|
self.assertIn("start", seg)
|
||||||
|
self.assertIn("end", seg)
|
||||||
|
self.assertIn("text", seg)
|
||||||
|
self.assertGreaterEqual(seg["end"], seg["start"])
|
||||||
|
self.assertNotIn(
|
||||||
|
"<|", seg["text"], f"Special token leaked into segment: {seg!r}"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_auto_detect_streaming(self):
|
||||||
|
"""language=None + stream=True: deltas scrubbed, concat matches non-streaming.
|
||||||
|
|
||||||
|
Verified against a real server: sglang's streaming path for Whisper
|
||||||
|
produces clean deltas (complete words, no BPE fragmentation), so the
|
||||||
|
fused path only needs to hide the forced prefix — which this PR
|
||||||
|
does. Asserts both the prefix-leak guard and text equivalence.
|
||||||
|
"""
|
||||||
|
deltas = self._transcribe_stream(language=None)
|
||||||
|
self.assertTrue(len(deltas) > 0, "Expected at least one streamed delta")
|
||||||
|
for d in deltas:
|
||||||
|
self.assertNotIn(
|
||||||
|
"<|", d, f"Special token leaked into streaming delta: {d!r}"
|
||||||
|
)
|
||||||
|
streamed = "".join(deltas).strip()
|
||||||
|
reference = self._transcribe(language=None).get("text", "").strip()
|
||||||
|
self.assertEqual(
|
||||||
|
streamed,
|
||||||
|
reference,
|
||||||
|
"Streamed auto-detect text should match the non-streaming result.",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -0,0 +1,303 @@
|
|||||||
|
"""Unit tests for OpenAIServingTranscription's streaming fused-autodetect path.
|
||||||
|
|
||||||
|
Exercises the streaming handler: buffer deltas until the forced-prefix
|
||||||
|
sentinel lands, emit the scrubbed user-visible text, and never leak
|
||||||
|
Whisper special tokens. Covers both streaming modes — cumulative
|
||||||
|
(``incremental_streaming_output=False``, the default) and incremental
|
||||||
|
(``incremental_streaming_output=True``).
|
||||||
|
|
||||||
|
The tests mock ``TokenizerManager.generate_request`` to yield synthetic
|
||||||
|
``text`` chunks for each of the happy, abort, and boundary cases.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from sglang.test.test_utils import maybe_stub_sgl_kernel
|
||||||
|
|
||||||
|
maybe_stub_sgl_kernel() # must precede any import that pulls in sgl_kernel
|
||||||
|
|
||||||
|
import json
|
||||||
|
import unittest
|
||||||
|
from typing import List
|
||||||
|
from unittest.mock import Mock
|
||||||
|
|
||||||
|
from sglang.srt.entrypoints.openai.protocol import TranscriptionRequest
|
||||||
|
from sglang.srt.entrypoints.openai.serving_transcription import (
|
||||||
|
OpenAIServingTranscription,
|
||||||
|
)
|
||||||
|
from sglang.srt.managers.io_struct import GenerateReqInput
|
||||||
|
from sglang.srt.utils import get_or_create_event_loop
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=4, suite="stage-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
def _chunk(text: str, finish: str = None) -> dict:
|
||||||
|
"""Shape of what TokenizerManager.generate_request yields per step."""
|
||||||
|
return {
|
||||||
|
"text": text,
|
||||||
|
"meta_info": {
|
||||||
|
"finish_reason": {"type": finish} if finish else None,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class _MockTokenizerManager:
|
||||||
|
"""Minimal mock satisfying OpenAIServingTranscription.__init__ and stream loop."""
|
||||||
|
|
||||||
|
def __init__(self, stream_chunks: List[dict]):
|
||||||
|
self.model_config = Mock()
|
||||||
|
self.model_config.hf_config = Mock()
|
||||||
|
self.model_config.hf_config.architectures = ["WhisperForConditionalGeneration"]
|
||||||
|
# Not a real ServerArgs, so base class sets allowed_custom_labels=None.
|
||||||
|
# Default tests assume cumulative-text streaming (the sglang upstream
|
||||||
|
# default); tests for incremental_streaming_output=True override this.
|
||||||
|
self.server_args = Mock(incremental_streaming_output=False)
|
||||||
|
self.tokenizer = Mock()
|
||||||
|
self._stream_chunks = stream_chunks
|
||||||
|
|
||||||
|
def generate_request(self, adapted_request, raw_request):
|
||||||
|
chunks = self._stream_chunks
|
||||||
|
|
||||||
|
async def gen():
|
||||||
|
for c in chunks:
|
||||||
|
yield c
|
||||||
|
|
||||||
|
return gen()
|
||||||
|
|
||||||
|
def create_abort_task(self, adapted_request):
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _deltas_from_sse(sse_lines: List[str]) -> List[str]:
|
||||||
|
"""Extract ``choices[0].delta.content`` strings from a list of SSE frames."""
|
||||||
|
out = []
|
||||||
|
for line in sse_lines:
|
||||||
|
if not line.startswith("data: "):
|
||||||
|
continue
|
||||||
|
payload = line[len("data: ") :].strip()
|
||||||
|
if payload == "[DONE]":
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
obj = json.loads(payload)
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
continue
|
||||||
|
for choice in obj.get("choices", []):
|
||||||
|
content = (choice.get("delta") or {}).get("content")
|
||||||
|
if content:
|
||||||
|
out.append(content)
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
class TestStreamingFusedAutodetect(CustomTestCase):
|
||||||
|
"""_generate_transcription_stream with _fused_autodetect=True."""
|
||||||
|
|
||||||
|
def _run_stream(
|
||||||
|
self, chunks: List[dict], fused: bool = True, ts_variant: bool = False
|
||||||
|
):
|
||||||
|
tm = _MockTokenizerManager(chunks)
|
||||||
|
serving = OpenAIServingTranscription(tm)
|
||||||
|
|
||||||
|
kwargs = {"model": "whisper", "stream": True}
|
||||||
|
if ts_variant:
|
||||||
|
kwargs["timestamp_granularities"] = ["segment"]
|
||||||
|
request = TranscriptionRequest(**kwargs)
|
||||||
|
if fused:
|
||||||
|
request._fused_autodetect = True
|
||||||
|
request._fused_ts_variant = ts_variant
|
||||||
|
adapted = GenerateReqInput(text="", modalities=["audio"])
|
||||||
|
raw_request = Mock()
|
||||||
|
|
||||||
|
async def drive():
|
||||||
|
frames = []
|
||||||
|
async for frame in serving._generate_transcription_stream(
|
||||||
|
adapted, request, raw_request
|
||||||
|
):
|
||||||
|
frames.append(frame)
|
||||||
|
return frames
|
||||||
|
|
||||||
|
loop = get_or_create_event_loop()
|
||||||
|
frames = loop.run_until_complete(drive())
|
||||||
|
return request, frames
|
||||||
|
|
||||||
|
def test_prefix_stripped_and_language_extracted(self):
|
||||||
|
chunks = [
|
||||||
|
_chunk("<|en|>"),
|
||||||
|
_chunk("<|en|><|transcribe|>"),
|
||||||
|
_chunk("<|en|><|transcribe|><|notimestamps|>"),
|
||||||
|
_chunk("<|en|><|transcribe|><|notimestamps|> Hello"),
|
||||||
|
_chunk("<|en|><|transcribe|><|notimestamps|> Hello world", finish="stop"),
|
||||||
|
]
|
||||||
|
request, frames = self._run_stream(chunks)
|
||||||
|
deltas = _deltas_from_sse(frames)
|
||||||
|
self.assertEqual(deltas, ["Hello", " world"])
|
||||||
|
self.assertEqual(request.language, "en")
|
||||||
|
# No delta ever starts with the forced prefix or leading whitespace.
|
||||||
|
self.assertFalse(any("<|" in d for d in deltas))
|
||||||
|
self.assertFalse(deltas[0].startswith(" "))
|
||||||
|
|
||||||
|
def test_non_english_language_extracted(self):
|
||||||
|
chunks = [
|
||||||
|
_chunk("<|zh|><|transcribe|><|notimestamps|>你好"),
|
||||||
|
_chunk("<|zh|><|transcribe|><|notimestamps|>你好世界", finish="stop"),
|
||||||
|
]
|
||||||
|
request, frames = self._run_stream(chunks)
|
||||||
|
self.assertEqual(request.language, "zh")
|
||||||
|
self.assertEqual(_deltas_from_sse(frames), ["你好", "世界"])
|
||||||
|
|
||||||
|
def test_fsm_abort_before_sentinel_emits_error_frame(self):
|
||||||
|
# Sentinel never arrives; stream terminates on finish_reason. The
|
||||||
|
# handler must surface this as a real SSE error frame so the client
|
||||||
|
# can distinguish "detection failed" from "silent audio with zero
|
||||||
|
# transcription". language stays unset.
|
||||||
|
chunks = [
|
||||||
|
_chunk("<|en|>"),
|
||||||
|
_chunk("<|en|><|transcribe|>", finish="length"),
|
||||||
|
]
|
||||||
|
request, frames = self._run_stream(chunks)
|
||||||
|
self.assertEqual(_deltas_from_sse(frames), [])
|
||||||
|
error_frames = [f for f in frames if f.startswith("data: ") and '"error"' in f]
|
||||||
|
self.assertTrue(
|
||||||
|
error_frames, f"expected an SSE error frame, got frames={frames!r}"
|
||||||
|
)
|
||||||
|
self.assertIn("language auto-detect failed", error_frames[0])
|
||||||
|
self.assertIsNone(request.language)
|
||||||
|
|
||||||
|
def test_non_fused_stream_passes_through(self):
|
||||||
|
# When _fused_autodetect is False, no buffering or anchoring happens.
|
||||||
|
chunks = [
|
||||||
|
_chunk("Hello"),
|
||||||
|
_chunk("Hello world", finish="stop"),
|
||||||
|
]
|
||||||
|
request, frames = self._run_stream(chunks, fused=False)
|
||||||
|
self.assertEqual(_deltas_from_sse(frames), ["Hello", " world"])
|
||||||
|
|
||||||
|
def test_streaming_ts_variant_sentinel_at_chunk_boundary(self):
|
||||||
|
# The <|0.00|> sentinel can land in its own chunk ahead of any
|
||||||
|
# transcription text, and the trailing-space arrives later. The
|
||||||
|
# handler must buffer silently until a non-whitespace char shows
|
||||||
|
# up (so the first delta doesn't leak a leading space) and then
|
||||||
|
# scrub subsequent embedded timestamp tokens.
|
||||||
|
chunks = [
|
||||||
|
_chunk("<|en|>"),
|
||||||
|
_chunk("<|en|><|transcribe|>"),
|
||||||
|
_chunk("<|en|><|transcribe|><|0.00|>"), # sentinel alone
|
||||||
|
_chunk("<|en|><|transcribe|><|0.00|> "), # + whitespace only
|
||||||
|
_chunk("<|en|><|transcribe|><|0.00|> Hello"), # first word
|
||||||
|
_chunk("<|en|><|transcribe|><|0.00|> Hello<|5.00|> World"),
|
||||||
|
_chunk(
|
||||||
|
"<|en|><|transcribe|><|0.00|> Hello<|5.00|> World<|endoftext|>",
|
||||||
|
finish="stop",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
request, frames = self._run_stream(chunks, ts_variant=True)
|
||||||
|
deltas = _deltas_from_sse(frames)
|
||||||
|
self.assertEqual(request.language, "en")
|
||||||
|
self.assertFalse(any("<|" in d for d in deltas))
|
||||||
|
# No delta starts with a leading space (the one Whisper emits
|
||||||
|
# between <|0.00|> and "Hello" was consumed by the defer-on-
|
||||||
|
# whitespace path).
|
||||||
|
self.assertFalse(deltas[0].startswith(" "))
|
||||||
|
self.assertEqual("".join(deltas), "Hello World")
|
||||||
|
|
||||||
|
def test_streaming_timestamps_variant_scrubs_embedded_segment_tokens(self):
|
||||||
|
# Streaming + timestamp_granularities + language=None uses the fused
|
||||||
|
# timestamps variant (<|0.00|> sentinel). Segment-boundary tokens
|
||||||
|
# <|5.00|>, <|10.00|> land mid-stream; each delta must have them
|
||||||
|
# scrubbed before reaching the client. Auto-detection still works
|
||||||
|
# — the SSE stream carries clean text, and callers who want
|
||||||
|
# segment timing can use response_format=verbose_json which builds
|
||||||
|
# segments from output_ids on a separate path.
|
||||||
|
chunks = [
|
||||||
|
_chunk("<|en|><|transcribe|><|0.00|> Hello"),
|
||||||
|
_chunk("<|en|><|transcribe|><|0.00|> Hello<|5.00|> World"),
|
||||||
|
_chunk(
|
||||||
|
"<|en|><|transcribe|><|0.00|> Hello<|5.00|> World<|10.00|><|endoftext|>",
|
||||||
|
finish="stop",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
request, frames = self._run_stream(chunks, ts_variant=True)
|
||||||
|
deltas = _deltas_from_sse(frames)
|
||||||
|
self.assertEqual(request.language, "en")
|
||||||
|
self.assertFalse(any("<|" in d for d in deltas))
|
||||||
|
self.assertEqual("".join(deltas), "Hello World")
|
||||||
|
|
||||||
|
def test_trailing_endoftext_scrubbed_from_last_delta(self):
|
||||||
|
# skip_special_tokens=False means the detokenizer may emit
|
||||||
|
# <|endoftext|> at the tail. The fused streaming path must scrub it
|
||||||
|
# per-delta so clients never see special tokens in SSE chunks.
|
||||||
|
chunks = [
|
||||||
|
_chunk("<|en|><|transcribe|><|notimestamps|> Hello"),
|
||||||
|
_chunk(
|
||||||
|
"<|en|><|transcribe|><|notimestamps|> Hello world<|endoftext|>",
|
||||||
|
finish="stop",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
_, frames = self._run_stream(chunks)
|
||||||
|
deltas = _deltas_from_sse(frames)
|
||||||
|
self.assertEqual(deltas, ["Hello", " world"])
|
||||||
|
self.assertFalse(any("<|" in d for d in deltas))
|
||||||
|
|
||||||
|
|
||||||
|
class TestStreamingIncrementalOutputMode(CustomTestCase):
|
||||||
|
"""Server runs with ``incremental_streaming_output=True``.
|
||||||
|
|
||||||
|
In that mode each chunk's ``content["text"]`` is the new delta from the
|
||||||
|
detokenizer, not the cumulative text. The handler must accumulate
|
||||||
|
locally into ``cumulative_text`` — otherwise the subsequent
|
||||||
|
``visible[len(visible_buffer):]`` slice would strip characters the
|
||||||
|
server already sent as a delta.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def _run_incremental_stream(self, chunk_deltas, fused=False):
|
||||||
|
"""Server in incremental mode: yield per-chunk delta, not cumulative."""
|
||||||
|
chunks = [
|
||||||
|
_chunk(d, finish=("stop" if i == len(chunk_deltas) - 1 else None))
|
||||||
|
for i, d in enumerate(chunk_deltas)
|
||||||
|
]
|
||||||
|
tm = _MockTokenizerManager(chunks)
|
||||||
|
tm.server_args = Mock(incremental_streaming_output=True)
|
||||||
|
serving = OpenAIServingTranscription(tm)
|
||||||
|
|
||||||
|
request = TranscriptionRequest(model="whisper", stream=True)
|
||||||
|
if fused:
|
||||||
|
request._fused_autodetect = True
|
||||||
|
adapted = GenerateReqInput(text="", modalities=["audio"])
|
||||||
|
|
||||||
|
async def drive():
|
||||||
|
frames = []
|
||||||
|
async for f in serving._generate_transcription_stream(
|
||||||
|
adapted, request, Mock()
|
||||||
|
):
|
||||||
|
frames.append(f)
|
||||||
|
return frames
|
||||||
|
|
||||||
|
return request, get_or_create_event_loop().run_until_complete(drive())
|
||||||
|
|
||||||
|
def test_incremental_non_fused_emits_each_delta_verbatim(self):
|
||||||
|
# sglang.private default: each content["text"] IS the new delta, so
|
||||||
|
# the handler should NOT slice it. Client should see exactly what
|
||||||
|
# the detokenizer emitted.
|
||||||
|
deltas_in = [" The", " President", ":", " Thank", " you"]
|
||||||
|
_, frames = self._run_incremental_stream(deltas_in, fused=False)
|
||||||
|
self.assertEqual(_deltas_from_sse(frames), deltas_in)
|
||||||
|
|
||||||
|
def test_incremental_fused_autodetect_still_strips_prefix(self):
|
||||||
|
# Incremental + fused: the handler must accumulate to find the
|
||||||
|
# sentinel, then emit only the post-prefix portion per chunk.
|
||||||
|
deltas_in = [
|
||||||
|
"<|en|>",
|
||||||
|
"<|transcribe|>",
|
||||||
|
"<|notimestamps|>",
|
||||||
|
" Hello",
|
||||||
|
" world",
|
||||||
|
]
|
||||||
|
request, frames = self._run_incremental_stream(deltas_in, fused=True)
|
||||||
|
emitted = _deltas_from_sse(frames)
|
||||||
|
# Prefix never leaks, and concat matches the expected transcription.
|
||||||
|
self.assertFalse(any("<|" in d for d in emitted))
|
||||||
|
self.assertEqual("".join(emitted), "Hello world")
|
||||||
|
self.assertEqual(request.language, "en")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,318 @@
|
|||||||
|
"""Unit tests for the Whisper transcription adapter.
|
||||||
|
|
||||||
|
Focused on ``WhisperAdapter.parse_fused_output`` — a pure static method
|
||||||
|
that parses the fused auto-detect output into ``(language, user_visible_text)``.
|
||||||
|
``visible=None`` means "forced prefix not yet locatable; streaming callers
|
||||||
|
should keep buffering, non-streaming callers should fall back to a
|
||||||
|
best-effort scrub".
|
||||||
|
"""
|
||||||
|
|
||||||
|
import re
|
||||||
|
import unittest
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from sglang.srt.entrypoints.openai.protocol import TranscriptionRequest
|
||||||
|
from sglang.srt.entrypoints.openai.transcription_adapters.whisper import (
|
||||||
|
WHISPER_AUTODETECT_REGEX,
|
||||||
|
WHISPER_AUTODETECT_TS_REGEX,
|
||||||
|
WHISPER_LANG_TOKEN_CODES,
|
||||||
|
WhisperAdapter,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=2, suite="stage-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class TestWhisperParseFusedOutput(CustomTestCase):
|
||||||
|
"""parse_fused_output: (language, visible) where visible=None means defer."""
|
||||||
|
|
||||||
|
def test_happy_english(self):
|
||||||
|
lang, visible = WhisperAdapter.parse_fused_output(
|
||||||
|
"<|en|><|transcribe|><|notimestamps|> Hello world"
|
||||||
|
)
|
||||||
|
self.assertEqual((lang, visible), ("en", "Hello world"))
|
||||||
|
|
||||||
|
def test_happy_non_english(self):
|
||||||
|
lang, visible = WhisperAdapter.parse_fused_output(
|
||||||
|
"<|zh|><|transcribe|><|notimestamps|>你好世界"
|
||||||
|
)
|
||||||
|
self.assertEqual((lang, visible), ("zh", "你好世界"))
|
||||||
|
|
||||||
|
def test_missing_language_prefix_defers(self):
|
||||||
|
# Partial prefix or raw untagged text — streaming callers should
|
||||||
|
# keep buffering; non-streaming callers fall back to best-effort.
|
||||||
|
self.assertEqual(
|
||||||
|
WhisperAdapter.parse_fused_output("raw untagged output"), (None, None)
|
||||||
|
)
|
||||||
|
self.assertEqual(WhisperAdapter.parse_fused_output(""), (None, None))
|
||||||
|
|
||||||
|
def test_missing_sentinel_defers(self):
|
||||||
|
# Reviewer's repro: <|zh|> Hi — language tag in but no sentinel.
|
||||||
|
self.assertEqual(WhisperAdapter.parse_fused_output("<|zh|> Hi"), (None, None))
|
||||||
|
|
||||||
|
def test_truncated_after_transcribe_defers(self):
|
||||||
|
self.assertEqual(
|
||||||
|
WhisperAdapter.parse_fused_output("<|en|><|transcribe|>"), (None, None)
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_unsupported_language_code_defers(self):
|
||||||
|
# FSM regex only allows ISO639_1_SUPPORTED_LANGS. A bypassed-FSM
|
||||||
|
# <|xx|> must not leak through as a valid detection.
|
||||||
|
self.assertEqual(
|
||||||
|
WhisperAdapter.parse_fused_output("<|xx|><|transcribe|><|notimestamps|>hi"),
|
||||||
|
(None, None),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_malformed_prefix_without_transcribe_defers(self):
|
||||||
|
# The parse must match the exact 3-token forced prefix, not
|
||||||
|
# "lang tag + sentinel somewhere". A bypassed-FSM string that
|
||||||
|
# skips <|transcribe|> must not parse as a valid detection.
|
||||||
|
self.assertEqual(
|
||||||
|
WhisperAdapter.parse_fused_output("<|en|>junk<|notimestamps|>text"),
|
||||||
|
(None, None),
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
WhisperAdapter.parse_fused_output("<|en|><|0.00|> text"),
|
||||||
|
(None, None),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_sentinel_in_but_whitespace_only_returns_empty_visible(self):
|
||||||
|
# Prefix arrived at a chunk boundary before the first word. The
|
||||||
|
# .strip() collapses to "" so streaming callers see no delta yet;
|
||||||
|
# the language is still reported as soon as the sentinel lands.
|
||||||
|
self.assertEqual(
|
||||||
|
WhisperAdapter.parse_fused_output("<|en|><|transcribe|><|notimestamps|>"),
|
||||||
|
("en", ""),
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
WhisperAdapter.parse_fused_output("<|en|><|transcribe|><|notimestamps|> "),
|
||||||
|
("en", ""),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_trailing_endoftext_scrubbed(self):
|
||||||
|
lang, visible = WhisperAdapter.parse_fused_output(
|
||||||
|
"<|en|><|transcribe|><|notimestamps|> Hello world<|endoftext|>"
|
||||||
|
)
|
||||||
|
self.assertEqual((lang, visible), ("en", "Hello world"))
|
||||||
|
|
||||||
|
def test_embedded_timestamp_tokens_scrubbed(self):
|
||||||
|
# Defensive: in the ts variant Whisper's tokenizer normally
|
||||||
|
# decodes <|X.XX|> tokens to "" so they never reach this path,
|
||||||
|
# but if a future tokenizer leaks them through they must be
|
||||||
|
# scrubbed from the user-visible text. verbose_json segment
|
||||||
|
# timing comes from _parse_segments over output_ids on a
|
||||||
|
# separate path.
|
||||||
|
lang, visible = WhisperAdapter.parse_fused_output(
|
||||||
|
"<|en|><|transcribe|><|0.00|> Hello<|5.00|> world<|10.00|><|endoftext|>",
|
||||||
|
ts_variant=True,
|
||||||
|
)
|
||||||
|
self.assertEqual((lang, visible), ("en", "Hello world"))
|
||||||
|
|
||||||
|
def test_ts_variant_realistic_decoded_text(self):
|
||||||
|
# Real Whisper tokenizer decodes every <|X.XX|> timestamp token
|
||||||
|
# (id 50365+) to "" even with skip_special_tokens=False, so for
|
||||||
|
# the ts variant the cumulative text is just <|en|><|transcribe|>
|
||||||
|
# followed directly by the BPE-decoded transcription. Asserts
|
||||||
|
# that the parser handles this shape — without ts_variant=True
|
||||||
|
# it would (correctly) defer because <|notimestamps|> is missing.
|
||||||
|
lang, visible = WhisperAdapter.parse_fused_output(
|
||||||
|
"<|en|><|transcribe|> Hello world<|endoftext|>", ts_variant=True
|
||||||
|
)
|
||||||
|
self.assertEqual((lang, visible), ("en", "Hello world"))
|
||||||
|
# Same input under non-ts contract correctly defers.
|
||||||
|
self.assertEqual(
|
||||||
|
WhisperAdapter.parse_fused_output(
|
||||||
|
"<|en|><|transcribe|> Hello world<|endoftext|>"
|
||||||
|
),
|
||||||
|
(None, None),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_visible_grows_monotonically_across_snapshots(self):
|
||||||
|
# Streaming property: cumulative text produces cumulative visible.
|
||||||
|
snapshots = [
|
||||||
|
"<|en|><|transcribe|>",
|
||||||
|
"<|en|><|transcribe|><|notimestamps|>",
|
||||||
|
"<|en|><|transcribe|><|notimestamps|> Hello",
|
||||||
|
"<|en|><|transcribe|><|notimestamps|> Hello world",
|
||||||
|
"<|en|><|transcribe|><|notimestamps|> Hello world<|endoftext|>",
|
||||||
|
]
|
||||||
|
visibles = [WhisperAdapter.parse_fused_output(s)[1] for s in snapshots]
|
||||||
|
# (None, "", "Hello", "Hello world", "Hello world")
|
||||||
|
self.assertEqual(visibles, [None, "", "Hello", "Hello world", "Hello world"])
|
||||||
|
# Every non-None entry is a prefix of the next non-None entry.
|
||||||
|
real = [v for v in visibles if v is not None]
|
||||||
|
for a, b in zip(real, real[1:]):
|
||||||
|
self.assertTrue(b.startswith(a), f"monotonicity broken: {a!r} -> {b!r}")
|
||||||
|
|
||||||
|
|
||||||
|
class TestWhisperLangTokenCoverage(CustomTestCase):
|
||||||
|
"""The FSM regex must cover every Whisper language token, not just the
|
||||||
|
narrower ISO639_1_SUPPORTED_LANGS set used for input validation."""
|
||||||
|
|
||||||
|
def test_three_letter_codes_parse(self):
|
||||||
|
# yue (Cantonese, v3), haw (Hawaiian), jw (Javanese, two-letter but
|
||||||
|
# missing from ISO639_1_SUPPORTED_LANGS) — reviewer's flagged examples.
|
||||||
|
for code in ("yue", "haw", "jw"):
|
||||||
|
with self.subTest(lang=code):
|
||||||
|
lang, visible = WhisperAdapter.parse_fused_output(
|
||||||
|
f"<|{code}|><|transcribe|><|notimestamps|> Hi"
|
||||||
|
)
|
||||||
|
self.assertEqual(lang, code)
|
||||||
|
self.assertEqual(visible, "Hi")
|
||||||
|
|
||||||
|
def test_known_whisper_langs_in_allowlist(self):
|
||||||
|
# Spot-check: codes the reviewer named + common 3-letter tokens.
|
||||||
|
for code in ("yue", "haw", "jw", "su", "ba", "tt", "ln", "lo"):
|
||||||
|
self.assertIn(code, WHISPER_LANG_TOKEN_CODES)
|
||||||
|
|
||||||
|
def test_fsm_regex_includes_three_letter_alternatives(self):
|
||||||
|
# Defensive: the regex alternation must spell out the 3-letter codes
|
||||||
|
# so xgrammar's FSM admits the <|yue|> / <|haw|> single-token path.
|
||||||
|
for code in ("yue", "haw"):
|
||||||
|
self.assertIn(re.escape(code), WHISPER_AUTODETECT_REGEX)
|
||||||
|
self.assertIn(re.escape(code), WHISPER_AUTODETECT_TS_REGEX)
|
||||||
|
|
||||||
|
def test_autodetect_codes_round_trip_through_input_validator(self):
|
||||||
|
# A code returned by fused autodetect must be accepted as
|
||||||
|
# ``language=`` on a follow-up request. Before the fix,
|
||||||
|
# ``normalize_language_to_code("yue")`` raised ValueError even
|
||||||
|
# though verbose_json could report ``"yue"`` from the same server.
|
||||||
|
from sglang.srt.multimodal.processors.whisper import (
|
||||||
|
normalize_language_to_code,
|
||||||
|
)
|
||||||
|
|
||||||
|
for code in ("yue", "haw", "jw", "ba", "su", "tt"):
|
||||||
|
with self.subTest(lang=code):
|
||||||
|
self.assertEqual(normalize_language_to_code(code), code)
|
||||||
|
|
||||||
|
def test_unknown_language_token_id_raises_clean_error(self):
|
||||||
|
# Some Whisper codes (yue, v3-only) aren't in older checkpoints'
|
||||||
|
# vocabs. The explicit-language path must raise a clean ValueError
|
||||||
|
# in that case instead of silently feeding the unk token into the
|
||||||
|
# decoder and producing garbage. Mocks cover both "returns None"
|
||||||
|
# and "returns unk_token_id" tokenizer behaviors.
|
||||||
|
from unittest.mock import Mock
|
||||||
|
|
||||||
|
from sglang.srt.multimodal.processors.whisper import WhisperProcessor
|
||||||
|
|
||||||
|
proc = WhisperProcessor.__new__(WhisperProcessor)
|
||||||
|
# Tokenizer where <|yue|> is not in the vocab → returns unk_id.
|
||||||
|
tok = Mock()
|
||||||
|
tok.convert_tokens_to_ids = Mock(return_value=100) # arbitrary unk
|
||||||
|
tok.unk_token_id = 100
|
||||||
|
proc._tokenizer = tok
|
||||||
|
with self.assertRaises(ValueError) as ctx:
|
||||||
|
proc._get_language_token_id("yue")
|
||||||
|
self.assertIn("yue", str(ctx.exception))
|
||||||
|
|
||||||
|
# Known code (English) on the same tokenizer still works.
|
||||||
|
tok.convert_tokens_to_ids = Mock(return_value=50259) # <|en|>
|
||||||
|
self.assertEqual(proc._get_language_token_id("en"), 50259)
|
||||||
|
|
||||||
|
# Some tokenizers return None for unknown tokens instead of unk_id.
|
||||||
|
tok2 = Mock()
|
||||||
|
tok2.convert_tokens_to_ids = Mock(return_value=None)
|
||||||
|
tok2.unk_token_id = 100
|
||||||
|
proc._tokenizer = tok2
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
proc._get_language_token_id("yue")
|
||||||
|
|
||||||
|
|
||||||
|
class TestWhisperStripSpecialTokens(CustomTestCase):
|
||||||
|
"""Fallback scrub used when parse_fused_output defers."""
|
||||||
|
|
||||||
|
def test_strips_all_whisper_specials(self):
|
||||||
|
self.assertEqual(
|
||||||
|
WhisperAdapter.strip_special_tokens(
|
||||||
|
"<|en|><|transcribe|><|0.00|>hi<|5.00|>world<|endoftext|>"
|
||||||
|
),
|
||||||
|
"hiworld",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_identity_on_plain_text(self):
|
||||||
|
self.assertEqual(
|
||||||
|
WhisperAdapter.strip_special_tokens("plain text"), "plain text"
|
||||||
|
)
|
||||||
|
self.assertEqual(WhisperAdapter.strip_special_tokens(""), "")
|
||||||
|
|
||||||
|
def test_preserves_spoken_angle_bracket_sequences(self):
|
||||||
|
# The scrub must only remove actual Whisper special-token literals
|
||||||
|
# (lang / control / <|X.XX|> timestamps), not arbitrary ``<|...|>``
|
||||||
|
# patterns that can appear in transcribed speech (someone reading a
|
||||||
|
# token name aloud, an AI-safety demo, code dictation, etc.).
|
||||||
|
self.assertEqual(
|
||||||
|
WhisperAdapter.strip_special_tokens("the token <|foo|> is unused"),
|
||||||
|
"the token <|foo|> is unused",
|
||||||
|
)
|
||||||
|
# Real specials still scrubbed even when interleaved with bogus ones.
|
||||||
|
self.assertEqual(
|
||||||
|
WhisperAdapter.strip_special_tokens(
|
||||||
|
"<|en|>hello <|foo|> world<|endoftext|>"
|
||||||
|
),
|
||||||
|
"hello <|foo|> world",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_parse_preserves_spoken_angle_bracket_sequences(self):
|
||||||
|
# Same for the per-chunk scrub inside parse_fused_output.
|
||||||
|
lang, visible = WhisperAdapter.parse_fused_output(
|
||||||
|
"<|en|><|transcribe|><|notimestamps|> look at <|foo|><|endoftext|>"
|
||||||
|
)
|
||||||
|
self.assertEqual((lang, visible), ("en", "look at <|foo|>"))
|
||||||
|
|
||||||
|
|
||||||
|
class TestWhisperBuildFusedAutodetectParams(CustomTestCase):
|
||||||
|
"""build_fused_autodetect_params picks the right regex + propagates ts param."""
|
||||||
|
|
||||||
|
def _request(self, **kwargs: Any) -> TranscriptionRequest:
|
||||||
|
base: dict[str, Any] = dict(model="whisper", temperature=0.0)
|
||||||
|
base.update(kwargs)
|
||||||
|
return TranscriptionRequest(**base)
|
||||||
|
|
||||||
|
def test_no_timestamps_uses_notimestamps_regex(self):
|
||||||
|
params = WhisperAdapter().build_fused_autodetect_params(self._request())
|
||||||
|
self.assertEqual(params["regex"], WHISPER_AUTODETECT_REGEX)
|
||||||
|
self.assertNotIn("timestamp_granularities", params)
|
||||||
|
|
||||||
|
def test_timestamps_uses_ts_regex_and_propagates_granularities(self):
|
||||||
|
req = self._request(timestamp_granularities=["segment"])
|
||||||
|
params = WhisperAdapter().build_fused_autodetect_params(req)
|
||||||
|
self.assertEqual(params["regex"], WHISPER_AUTODETECT_TS_REGEX)
|
||||||
|
self.assertEqual(params["timestamp_granularities"], ["segment"])
|
||||||
|
|
||||||
|
def test_empty_timestamps_list_uses_notimestamps_regex(self):
|
||||||
|
# Empty list is falsy — treat as "no timestamps requested".
|
||||||
|
req = self._request(timestamp_granularities=[])
|
||||||
|
params = WhisperAdapter().build_fused_autodetect_params(req)
|
||||||
|
self.assertEqual(params["regex"], WHISPER_AUTODETECT_REGEX)
|
||||||
|
self.assertNotIn("timestamp_granularities", params)
|
||||||
|
|
||||||
|
def test_spaces_between_special_tokens_is_false(self):
|
||||||
|
# parse_fused_output assumes a zero-space forced prefix. Slow
|
||||||
|
# Whisper tokenizers otherwise insert a space between adjacent
|
||||||
|
# special tokens, which would silently break the parse path.
|
||||||
|
for req in (
|
||||||
|
self._request(),
|
||||||
|
self._request(timestamp_granularities=["segment"]),
|
||||||
|
):
|
||||||
|
params = WhisperAdapter().build_fused_autodetect_params(req)
|
||||||
|
self.assertIs(params["spaces_between_special_tokens"], False)
|
||||||
|
|
||||||
|
def test_fused_params_survive_sampling_params_construction(self):
|
||||||
|
# Regression: the multimodal processor's fused branch used to skip
|
||||||
|
# popping `timestamp_granularities`, leaking the key into
|
||||||
|
# SamplingParams(**kwargs) → TypeError on any language=None +
|
||||||
|
# timestamp_granularities request. Mirrors what the processor does
|
||||||
|
# before constructing SamplingParams.
|
||||||
|
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||||
|
|
||||||
|
req = self._request(timestamp_granularities=["segment"])
|
||||||
|
params = WhisperAdapter().build_fused_autodetect_params(req)
|
||||||
|
# Fields the processor pops before SamplingParams(**kwargs).
|
||||||
|
params.pop("_detect_language", None)
|
||||||
|
params.pop("timestamp_granularities", None)
|
||||||
|
SamplingParams(**params)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user