Fix Whisper transcription for audio over 30 seconds (#33604)

This commit is contained in:
Shenxiu Liu
2026-08-15 23:51:05 +08:00
committed by GitHub
parent 4beb157e87
commit fb97be4359
9 changed files with 1279 additions and 46 deletions
@@ -0,0 +1,102 @@
"""Energy-aware audio chunking for ASR models with a bounded input window.
Whisper-style encoders ingest a fixed-length window (30 s of audio = 3000
mel frames); the feature extractor silently truncates anything longer, so
long audio must be split into independent chunks before prompting. Cutting
blindly at the window boundary can land mid-word and corrupt the
transcription on both sides of the seam, so the splitter searches the tail
of each window for the quietest stretch (lowest RMS energy) and cuts there.
This matches the chunking behavior of vLLM's
``OpenAISpeechToText._split_audio`` / ``_find_split_point``: chunks are
contiguous and non-overlapping — the "search window" is only the region in
which the cut is allowed to land, and stitching the transcripts back is a
plain in-order concatenation.
"""
from __future__ import annotations
import io
import math
from typing import List, Tuple
import numpy as np
import soundfile as sf
from sglang.srt.utils import load_audio
# Region at the tail of each max-length window in which to search for a
# low-energy split point, in seconds.
SPLIT_SEARCH_WINDOW_S = 1.0
# RMS energy is evaluated over strides of this many samples (100 ms at
# 16 kHz); the quietest stride in the search region wins.
MIN_ENERGY_WINDOW_SIZE = 1600
def find_split_point(wav: np.ndarray, start_idx: int, end_idx: int) -> int:
"""Return the start index of the quietest energy window in
``wav[start_idx:end_idx]``.
Energy is RMS over consecutive ``MIN_ENERGY_WINDOW_SIZE``-sample
strides; the returned index is absolute (into ``wav``). The loop bound
intentionally leaves the final stride of the search region unevaluated:
it mirrors vLLM's ``_find_split_point`` verbatim so split points stay
bit-identical with the old vLLM transcription endpoint.
"""
segment = wav[start_idx:end_idx]
min_energy = math.inf
quietest_idx = start_idx
for i in range(0, len(segment) - MIN_ENERGY_WINDOW_SIZE, MIN_ENERGY_WINDOW_SIZE):
window = segment[i : i + MIN_ENERGY_WINDOW_SIZE]
energy = (window**2).mean() ** 0.5
if energy < min_energy:
quietest_idx = i + start_idx
min_energy = energy
return quietest_idx
def split_audio_energy_aware(
audio_data: bytes,
max_clip_s: float,
sample_rate: int = 16000,
) -> Tuple[List[bytes], List[float]]:
"""Split audio into WAV chunks no longer than ``max_clip_s`` seconds.
Decodes (and resamples to ``sample_rate``) the input, walks it in
``max_clip_s`` strides, and cuts each chunk at the lowest-energy point
within the final ``SPLIT_SEARCH_WINDOW_S`` of the stride so cuts land
on pauses instead of mid-word. Chunks are contiguous and
non-overlapping; concatenated they reproduce the full waveform.
Returns ``(chunk_wav_bytes, chunk_start_offsets_s)`` where
``chunk_start_offsets_s[i]`` is the start time of chunk ``i`` in the
original audio.
"""
if not audio_data:
raise ValueError("audio_data is empty")
audio = load_audio(audio_data, sr=sample_rate, mono=True)
chunk_size = int(sample_rate * max_clip_s)
search_size = int(sample_rate * SPLIT_SEARCH_WINDOW_S)
total = audio.shape[-1]
raw_chunks: List[np.ndarray] = []
offsets_s: List[float] = []
i = 0
while i < total:
offsets_s.append(i / sample_rate)
if i + chunk_size >= total:
raw_chunks.append(audio[i:])
break
search_start = i + chunk_size - search_size
search_end = min(i + chunk_size, total)
split_point = find_split_point(audio, search_start, search_end)
raw_chunks.append(audio[i:split_point])
i = split_point
chunks: List[bytes] = []
for chunk in raw_chunks:
buf = io.BytesIO()
sf.write(buf, chunk, sample_rate, format="WAV")
chunks.append(buf.getvalue())
return chunks, offsets_s
@@ -32,6 +32,7 @@ from typing import TYPE_CHECKING, AsyncGenerator, List, Optional, Union
from fastapi import Request, WebSocket
from fastapi.responses import ORJSONResponse, Response, StreamingResponse
from sglang.srt.entrypoints.openai.audio_chunking import split_audio_energy_aware
from sglang.srt.entrypoints.openai.protocol import (
DeltaMessage,
ErrorResponse,
@@ -113,9 +114,17 @@ class OpenAIServingTranscription(OpenAIServingBase):
info = sf.info(io.BytesIO(audio_data))
return info.duration
except Exception as e:
logger.warning(f"Could not calculate audio duration: {e}")
return 0.0
except Exception:
# soundfile can't parse some containers (e.g. mp3 on older
# libsndfile builds); fall back to a full decode.
try:
from sglang.srt.utils import load_audio
audio = load_audio(audio_data, sr=16000, mono=True)
return audio.shape[-1] / 16000.0
except Exception as e:
logger.warning(f"Could not calculate audio duration: {e}")
return 0.0
async def create_transcription(
self,
@@ -135,8 +144,63 @@ class OpenAIServingTranscription(OpenAIServingBase):
ORJSONResponse,
]:
"""Main entry point for transcription requests."""
# Calculate audio duration for usage reporting
audio_duration_s = self._get_audio_duration(audio_data)
# Calculate audio duration for usage reporting. Run in a thread:
# the fallback path decodes the full file and would block the
# event loop on long inputs.
audio_duration_s = await asyncio.to_thread(self._get_audio_duration, audio_data)
# Audio longer than the model's encoder window (30 s for Whisper)
# would be silently truncated by the feature extractor. Split it
# into chunks cut at low-energy points (pauses) so the cut never
# lands mid-word; each chunk is transcribed as an independent
# request and the texts are stitched back in order.
audio_chunks: Optional[List[bytes]] = None
chunk_offsets_s: Optional[List[float]] = None
max_clip_s = self._adapter.max_audio_clip_s
if max_clip_s is not None and audio_duration_s > max_clip_s:
split_error = (
f"Failed to split audio longer than {max_clip_s:g} seconds into "
"supported chunks."
)
try:
# In a thread: decodes + re-encodes the whole file, which
# would otherwise block the event loop on long inputs.
audio_chunks, chunk_offsets_s = await asyncio.to_thread(
split_audio_energy_aware, audio_data, max_clip_s
)
except Exception as e:
logger.warning(
"Failed to split %.1fs audio into chunks of <=%ss: %s",
audio_duration_s,
max_clip_s,
e,
)
return self.create_error_response(split_error)
else:
if (
not audio_chunks
or len(audio_chunks) <= 1
or chunk_offsets_s is None
or len(chunk_offsets_s) != len(audio_chunks)
):
logger.error(
"Audio splitter returned an invalid result for %.1fs audio: "
"%d chunks and %s offsets",
audio_duration_s,
len(audio_chunks or []),
(
"no"
if chunk_offsets_s is None
else str(len(chunk_offsets_s))
),
)
return self.create_error_response(split_error)
logger.info(
"Split %.1fs audio into %d chunks of <=%ss for transcription",
audio_duration_s,
len(audio_chunks),
max_clip_s,
)
# When language is not specified and the adapter supports detection,
# use a single fused request: SGLang's structured generation (regex)
@@ -172,6 +236,9 @@ class OpenAIServingTranscription(OpenAIServingBase):
# 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)
if audio_chunks is not None and len(audio_chunks) > 1:
request._audio_chunks = audio_chunks
request._chunk_offsets_s = chunk_offsets_s
# Use the base class handle_request pattern
return await self.handle_request(request, raw_request)
@@ -189,6 +256,11 @@ class OpenAIServingTranscription(OpenAIServingBase):
Response,
]:
"""Handle non-streaming transcription request."""
if getattr(request, "_audio_chunks", None):
return await self._handle_chunked_non_streaming_request(
adapted_request, request, raw_request
)
try:
ret = await self.tokenizer_manager.generate_request(
adapted_request, raw_request
@@ -196,27 +268,7 @@ class OpenAIServingTranscription(OpenAIServingBase):
except ValueError as e:
return self.create_error_response(str(e))
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 strip_special_tokens —
# 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)
text = self._finalize_text(request, ret.get("text", ""))
usage = TranscriptionUsage(seconds=int(math.ceil(request.audio_duration_s)))
@@ -233,6 +285,161 @@ class OpenAIServingTranscription(OpenAIServingBase):
# Default JSON format
return TranscriptionResponse(text=text, usage=usage)
def _finalize_text(
self, request: TranscriptionRequest, raw_text: str, strip: bool = True
) -> str:
"""Postprocess one generation's raw text into user-visible text.
For fused auto-detect requests, parse_fused_output returns the
scrubbed user-visible text and the detected language is recorded on
the request (first non-empty parsed chunk wins for chunked audio). 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.
``strip=False`` keeps the model-emitted boundary whitespace in the
fused path; chunk stitching needs it as the natural separator.
"""
text = self._adapter.postprocess_text(raw_text)
if not getattr(request, "_fused_autodetect", False):
return text
lang, visible = self._adapter.parse_fused_output(
text,
ts_variant=getattr(request, "_fused_ts_variant", False),
strip=strip,
)
if visible is None:
logger.warning(
"Fused auto-detect parse failed on non-streaming response; "
"falling back to raw-text scrub."
)
return self._adapter.strip_special_tokens(text)
if lang is not None and visible.strip() and request.language is None:
request.language = lang
logger.info("Auto-detected language: '%s'", lang)
return visible
def _build_chunk_request(
self,
adapted_request: GenerateReqInput,
chunk_audio: bytes,
stream: bool,
) -> GenerateReqInput:
"""Clone the adapted request for one audio chunk.
``sampling_params`` must be a fresh dict per chunk: the multimodal
processor pops transcription-level keys (language,
timestamp_granularities, the fused-autodetect flag) out of it while
building each chunk's decoder prompt.
"""
sampling_params = adapted_request.sampling_params
assert isinstance(sampling_params, dict)
chunk_request = GenerateReqInput(
text="",
audio_data=chunk_audio,
sampling_params=dict(sampling_params),
stream=stream,
modalities=["audio"],
routing_key=adapted_request.routing_key,
)
chunk_request.received_time = adapted_request.received_time
return chunk_request
def _abort_chunk_requests(self, chunk_requests: List[GenerateReqInput]) -> None:
"""Abort chunk generations engine-side.
The scheduler keeps decoding until a request is aborted by rid (a
no-op for chunks that already finished). A request whose rid has not
been assigned yet cannot be aborted by the immediate pass, so a second
abort fires after the dispatch window — the same approach as
``TokenizerManager.create_abort_task``.
"""
def abort_assigned_rids():
# Read rids at call time: a chunk task that had not started
# executing during the first pass gets its rid assigned later,
# so the delayed pass must not reuse an earlier snapshot.
for chunk_request in chunk_requests:
if isinstance(chunk_request.rid, str):
self.tokenizer_manager.abort_request(chunk_request.rid)
abort_assigned_rids()
async def abort_after_dispatch_window():
await asyncio.sleep(2)
abort_assigned_rids()
asyncio.create_task(abort_after_dispatch_window())
async def _handle_chunked_non_streaming_request(
self,
adapted_request: GenerateReqInput,
request: TranscriptionRequest,
raw_request: Request,
) -> Union[
TranscriptionResponse,
TranscriptionVerboseResponse,
ErrorResponse,
ORJSONResponse,
Response,
]:
"""Transcribe pre-split long audio (duration > max_audio_clip_s).
Each chunk is an independent generation. Chunks run sequentially so
one user-controlled upload cannot fan out into an unbounded number of
tokenizer/GPU requests. Results are stitched in chunk (= audio) order.
"""
chunk_requests = [
self._build_chunk_request(adapted_request, chunk_audio, stream=False)
for chunk_audio in request._audio_chunks
]
rets = []
try:
for chunk_request in chunk_requests:
ret = await self.tokenizer_manager.generate_request(
chunk_request, raw_request
).__anext__()
rets.append(ret)
except BaseException as e:
# Abort the in-flight request on failures or parent cancellation.
# Later chunks have not been dispatched.
self._abort_chunk_requests(chunk_requests)
if isinstance(e, ValueError):
return self.create_error_response(str(e))
raise
fused = getattr(request, "_fused_autodetect", False)
# Plain in-order concatenation: each chunk's model-emitted boundary
# whitespace (a leading space for spaced scripts, nothing for
# spaceless scripts like zh/ja/th) is the correct seam separator,
# so the fused path keeps it (strip=False) instead of inventing
# one. Only the full fused text is trimmed at the ends, matching
# the single-request fused response.
texts = [
self._finalize_text(request, ret.get("text", ""), strip=False)
for ret in rets
]
text = "".join(texts)
if fused:
text = text.strip()
usage = TranscriptionUsage(seconds=int(math.ceil(request.audio_duration_s)))
if request.response_format == "text":
return Response(content=text, media_type="text/plain")
if request.response_format == "verbose_json":
return self._adapter.build_verbose_response_chunked(
request,
text,
rets,
request._chunk_offsets_s,
self.tokenizer_manager.tokenizer,
usage,
)
return TranscriptionResponse(text=text, usage=usage)
async def _handle_streaming_request(
self,
adapted_request: GenerateReqInput,
@@ -249,6 +456,15 @@ class OpenAIServingTranscription(OpenAIServingBase):
),
media_type="text/event-stream",
)
if getattr(request, "_audio_chunks", None):
# Long audio pre-split into chunks, transcribed sequentially.
# No background abort_task: the in-flight chunk is aborted in
# the generator's finally on teardown, and disconnect is checked
# between chunks.
return StreamingResponse(
self._generate_long_audio_stream(adapted_request, request, raw_request),
media_type="text/event-stream",
)
return StreamingResponse(
self._generate_transcription_stream(adapted_request, request, raw_request),
media_type="text/event-stream",
@@ -326,7 +542,11 @@ class OpenAIServingTranscription(OpenAIServingBase):
yield f"data: {error}\n\n"
yield "data: [DONE]\n\n"
return
if lang is not None and request.language is None:
if (
lang is not None
and visible.strip()
and request.language is None
):
request.language = lang
logger.info("Auto-detected language: '%s'", lang)
else:
@@ -369,6 +589,136 @@ class OpenAIServingTranscription(OpenAIServingBase):
yield "data: [DONE]\n\n"
async def _generate_long_audio_stream(
self,
adapted_request: GenerateReqInput,
request: TranscriptionRequest,
raw_request: Request,
) -> AsyncGenerator[str, None]:
"""Stream transcription of long audio pre-split into chunks.
Chunks are transcribed sequentially (one streaming request at a
time, like ``_generate_chunked_asr_stream``), so the client sees
the transcript in audio order with a single finish frame after the
last chunk. The first abnormal chunk finish_reason (length/abort)
wins so a truncated non-final chunk isn't masked by later chunks
stopping cleanly. Fused auto-detect applies per chunk; the reported
language is the first non-empty chunk's detection. Disconnect is
checked between chunks, and the in-flight chunk is aborted on teardown.
``request._chunk_offsets_s`` is unused here — only the non-streaming
verbose_json path needs segment timing.
"""
created_time = int(time.time())
request_id = f"{self._request_id_prefix()}{uuid.uuid4().hex}"
model = request.model
fused_mode = getattr(request, "_fused_autodetect", False)
ts_variant = getattr(request, "_fused_ts_variant", False)
incremental = getattr(
self.tokenizer_manager.server_args,
"incremental_streaming_output",
False,
)
def _frame(delta: Optional[str], finish_reason: Optional[str] = None) -> str:
chunk = TranscriptionStreamResponse(
id=request_id,
created=created_time,
model=model,
choices=[
TranscriptionStreamChoice(
delta=DeltaMessage(content=delta) if delta else DeltaMessage(),
finish_reason=finish_reason,
)
],
)
return f"data: {chunk.model_dump_json()}\n\n"
finish_reason_type = None
in_flight: Optional[GenerateReqInput] = None
emitted_text = False
try:
for chunk_audio in request._audio_chunks:
if await raw_request.is_disconnected():
break
in_flight = self._build_chunk_request(
adapted_request, chunk_audio, stream=True
)
cumulative_text = ""
visible_buffer = ""
chunk_finish_reason = None
strip_chunk = not emitted_text
async for content in self.tokenizer_manager.generate_request(
in_flight, raw_request
):
finish_reason = content["meta_info"]["finish_reason"]
chunk_finish_reason = (
finish_reason["type"] if finish_reason else None
)
chunk_text = content.get("text", "")
if incremental:
cumulative_text += chunk_text
else:
cumulative_text = chunk_text
if fused_mode:
# Strip the first chunk that emits visible text (no
# leading space at stream start, like the single-request
# path). Later visible chunks keep their model-emitted
# boundary whitespace, which is the correct seam
# separator for spaced and spaceless scripts alike.
lang, visible = self._adapter.parse_fused_output(
cumulative_text,
ts_variant=ts_variant,
strip=strip_chunk,
)
if visible is None:
if not chunk_finish_reason:
continue
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 visible.strip()
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
if delta:
yield _frame(delta)
emitted_text = True
in_flight = None
if finish_reason_type in (None, "stop") and chunk_finish_reason:
finish_reason_type = chunk_finish_reason
yield _frame(None, finish_reason=finish_reason_type or "stop")
except ValueError as e:
error = self.create_streaming_error_response(str(e))
yield f"data: {error}\n\n"
finally:
# Abort the chunk still decoding if the stream is torn down
# (client disconnect / generator close) so it doesn't keep
# running in the scheduler. Only one chunk is ever in flight.
if in_flight is not None and isinstance(in_flight.rid, str):
self.tokenizer_manager.abort_request(in_flight.rid)
yield "data: [DONE]\n\n"
async def _generate_chunked_asr_stream(
self,
adapted_request: GenerateReqInput,
@@ -44,7 +44,7 @@ class TranscriptionAdapter(ABC):
@staticmethod
def parse_fused_output(
text: str, *, ts_variant: bool = False
text: str, *, ts_variant: bool = False, strip: bool = True
) -> tuple[Optional[str], Optional[str]]:
"""Parse the fused output into ``(language_code, user_visible_text)``.
@@ -76,6 +76,17 @@ class TranscriptionAdapter(ABC):
"""
return text
@property
def max_audio_clip_s(self) -> Optional[float]:
"""Maximum audio duration (seconds) the model can ingest per request.
Audio longer than this is split at low-energy points into
contiguous chunks (see ``audio_chunking.split_audio_energy_aware``),
transcribed as independent requests, and stitched back in order.
``None`` disables chunking.
"""
return None
@property
def supports_chunked_streaming(self) -> bool:
"""Whether this model uses chunk-based streaming instead of token-level streaming."""
@@ -125,6 +136,31 @@ class TranscriptionAdapter(ABC):
) -> TranscriptionVerboseResponse:
"""Build a ``verbose_json`` response with segments / timestamps."""
def build_verbose_response_chunked(
self,
request: TranscriptionRequest,
text: str,
rets: List[dict],
chunk_offsets_s: List[float],
tokenizer,
usage: TranscriptionUsage,
) -> TranscriptionVerboseResponse:
"""Build a ``verbose_json`` response from multiple chunk results.
Called instead of ``build_verbose_response`` when the audio was
longer than ``max_audio_clip_s`` and split into chunks;
``chunk_offsets_s[i]`` is chunk *i*'s start time in the original
audio. The default implementation returns the stitched text with
no segment timing.
"""
return TranscriptionVerboseResponse(
language=request.language,
duration=round(request.audio_duration_s, 2),
text=text,
segments=[],
usage=usage,
)
_ADAPTER_REGISTRY: dict[str, type[TranscriptionAdapter]] = {}
_DEFAULT_ADAPTER_KEY = "Whisper"
@@ -126,6 +126,13 @@ class WhisperAdapter(TranscriptionAdapter):
TIMESTAMP_BASE_TOKEN_ID = 50365 # <|0.00|>
TIMESTAMP_BASE_OFFSET = 0.02 # each token step = 0.02 s
@property
def max_audio_clip_s(self) -> Optional[float]:
# Whisper's encoder ingests a fixed 30 s window (3000 mel frames);
# the feature extractor truncates anything longer, so longer audio
# must be split into chunks before prompting.
return 30.0
def build_sampling_params(self, request: TranscriptionRequest) -> dict:
params: dict = {
"temperature": request.temperature,
@@ -185,7 +192,7 @@ class WhisperAdapter(TranscriptionAdapter):
@staticmethod
def parse_fused_output(
text: str, *, ts_variant: bool = False
text: str, *, ts_variant: bool = False, strip: bool = True
) -> tuple[Optional[str], Optional[str]]:
"""Parse fused output into ``(language_code, user_visible_text)``.
@@ -209,9 +216,16 @@ class WhisperAdapter(TranscriptionAdapter):
* ``(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.
(with ``strip=True``, the default) surrounding whitespace
trimmed. It grows monotonically across streaming chunks because
Whisper's special tokens detokenize atomically, so callers can
compute deltas against it directly.
``strip=False`` preserves the model-emitted boundary whitespace.
Long-audio chunk stitching relies on it: a chunk's own leading
space (or its absence, for spaceless scripts like zh/ja/th) is the
correct separator at the chunk seam, so an artificial one must not
be invented after stripping.
"""
pattern = _FUSED_PREFIX_RE_TS if ts_variant else _FUSED_PREFIX_RE_NOTS
m = pattern.match(text)
@@ -225,7 +239,7 @@ class WhisperAdapter(TranscriptionAdapter):
# 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()
return m.group(1), transcription.strip() if strip else transcription
@staticmethod
def strip_special_tokens(text: str) -> str:
@@ -260,9 +274,41 @@ class WhisperAdapter(TranscriptionAdapter):
usage=usage,
)
def build_verbose_response_chunked(
self,
request: TranscriptionRequest,
text: str,
rets: List[dict],
chunk_offsets_s: List[float],
tokenizer,
usage: TranscriptionUsage,
) -> TranscriptionVerboseResponse:
segments: list[TranscriptionSegment] = []
for ret, offset_s in zip(rets, chunk_offsets_s):
_, part_segments = self._parse_segments(
ret.get("output_ids", []),
tokenizer,
time_offset_s=offset_s,
seg_id_start=len(segments),
)
segments.extend(part_segments)
return TranscriptionVerboseResponse(
language=request.language,
duration=round(request.audio_duration_s, 2),
# The serving layer already stitched model text using each
# chunk's emitted boundary whitespace. Reconstructing it from
# output_ids with an ASCII join corrupts spaceless scripts.
text=text,
segments=segments,
usage=usage,
)
@staticmethod
def _parse_segments(
output_ids: List[int], tokenizer
output_ids: List[int],
tokenizer,
time_offset_s: float = 0.0,
seg_id_start: int = 0,
) -> tuple[str, List[TranscriptionSegment]]:
"""Parse Whisper timestamp tokens from *output_ids* into segments.
@@ -273,6 +319,11 @@ class WhisperAdapter(TranscriptionAdapter):
Each timestamp token marks the end of the current segment; its value
also becomes the start of the next segment.
``time_offset_s`` / ``seg_id_start`` shift segment times and ids for
audio that was split into chunks — timestamp tokens are relative to
the chunk's own 30 s window, so each chunk's segments are offset by
the chunk's start time within the original audio.
"""
eos_token_id = getattr(tokenizer, "eos_token_id", 50257)
ts_base = WhisperAdapter.TIMESTAMP_BASE_TOKEN_ID
@@ -281,12 +332,13 @@ class WhisperAdapter(TranscriptionAdapter):
segments: list[TranscriptionSegment] = []
full_text_parts: list[str] = []
current_text_tokens: list[int] = []
current_start = 0.0 # First segment starts at 0.0 (from prompt <|0.00|>)
seg_id = 0
# First segment starts at the chunk start (prompt anchors <|0.00|>)
current_start = time_offset_s
seg_id = seg_id_start
for token_id in output_ids:
if token_id >= ts_base:
timestamp = (token_id - ts_base) * ts_step
timestamp = (token_id - ts_base) * ts_step + time_offset_s
if current_text_tokens:
seg_text = tokenizer.decode(
+2 -2
View File
@@ -998,7 +998,7 @@ def get_compiler_backend(mode=None) -> str:
if hasattr(torch, "npu") and torch.npu.is_available():
try:
import torchair
import torchair.ge_concrete_graph.ge_converter.experimental.patch_for_hcom_allreduce
import torchair.ge_concrete_graph.ge_converter.experimental.patch_for_hcom_allreduce # noqa: F401
from torchair.configs.compiler_config import CompilerConfig
except ImportError:
raise ImportError(
@@ -1687,7 +1687,7 @@ CLIENT_MEDIA_EXCEPTIONS = (
def load_audio(
audio_file: str, sr: Optional[int] = None, mono: bool = True
audio_file: Union[str, bytes], sr: Optional[int] = None, mono: bool = True
) -> np.ndarray:
if sr is None:
sr = 16000
@@ -10,9 +10,11 @@ import json
import unittest
from typing import List, Optional
import numpy as np
import requests
import soundfile as sf
from sglang.srt.utils import kill_process_tree
from sglang.srt.utils import kill_process_tree, load_audio
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
@@ -21,7 +23,7 @@ from sglang.test.test_utils import (
popen_launch_server,
)
register_cuda_ci(est_time=60, stage="base-b", runner_config="1-gpu-small")
register_cuda_ci(est_time=90, stage="base-b", runner_config="1-gpu-small")
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"
@@ -34,6 +36,21 @@ def download_audio_bytes(url=AUDIO_URL):
return response.content
def long_audio_wav_bytes(prefix_silence_s: float = 30.0) -> bytes:
"""A 40 s WAV: silence, then the 10 s speech clip.
The speech sits entirely past Whisper's 30 s encoder window, so without
long-audio chunking the feature extractor silently truncates it away
and the transcript contains none of the spoken content.
"""
sr = 16000
speech = load_audio(download_audio_bytes(), sr=sr, mono=True).astype(np.float32)
wav = np.concatenate([np.zeros(int(prefix_silence_s * sr), np.float32), speech])
buf = io.BytesIO()
sf.write(buf, wav, sr, format="WAV")
return buf.getvalue()
class TestServingTranscription(CustomTestCase):
"""Test Whisper transcription via /v1/audio/transcriptions endpoint."""
@@ -61,13 +78,15 @@ class TestServingTranscription(CustomTestCase):
language: Optional[str] = "en",
response_format: Optional[str] = None,
timestamp_granularities: Optional[List[str]] = None,
audio_bytes: Optional[bytes] = 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()
if audio_bytes is None:
audio_bytes = download_audio_bytes()
data = {"model": "whisper"}
if language is not None:
data["language"] = language
@@ -84,9 +103,14 @@ class TestServingTranscription(CustomTestCase):
self.assertEqual(response.status_code, 200, response.text)
return response.json()
def _transcribe_stream(self, language: Optional[str] = None) -> List[str]:
def _transcribe_stream(
self,
language: Optional[str] = None,
audio_bytes: Optional[bytes] = None,
) -> List[str]:
"""Send a streaming transcription request and return the delta strings."""
audio_bytes = download_audio_bytes()
if audio_bytes is None:
audio_bytes = download_audio_bytes()
data = {"model": "whisper", "stream": "true"}
if language is not None:
data["language"] = language
@@ -225,6 +249,56 @@ class TestServingTranscription(CustomTestCase):
"Streamed auto-detect text should match the non-streaming result.",
)
# -- long audio (> 30 s encoder window) --------------------------------
# 30 s of silence followed by the 10 s speech clip: every spoken word is
# past Whisper's encoder window, so these tests fail outright unless the
# server splits long audio into chunks and stitches the transcripts.
KEYWORDS = ["privilege", "leader", "science", "art"]
def _assert_keywords(self, text: str):
matches = [kw for kw in self.KEYWORDS if kw in text.lower()]
self.assertGreaterEqual(
len(matches),
2,
f"Expected at least 2 of {self.KEYWORDS}, found {matches}. "
f"Full text: {text!r}",
)
def test_long_audio_transcribes_past_30s(self):
"""Speech past the 30 s window must appear in the transcript."""
result = self._transcribe(audio_bytes=long_audio_wav_bytes())
self._assert_keywords(result["text"])
# Usage reports the full audio duration, not one chunk's.
self.assertGreaterEqual(result["usage"]["seconds"], 40)
def test_long_audio_verbose_json_segment_offsets(self):
"""Segment timestamps are offset by each chunk's start time."""
result = self._transcribe(
response_format="verbose_json",
timestamp_granularities=["segment"],
audio_bytes=long_audio_wav_bytes(),
)
self._assert_keywords(result.get("text", ""))
segments = result.get("segments") or []
self.assertGreater(len(segments), 0, "Expected at least one segment")
# The speech starts at t=30 s; its segments must be reported in
# original-audio time, which is unreachable within a single 30 s
# window.
self.assertGreater(
max(seg["end"] for seg in segments),
30.0,
f"Expected segment timing past the 30 s window, got {segments!r}",
)
def test_long_audio_streaming(self):
"""Streaming long audio emits the post-30 s content as deltas."""
deltas = self._transcribe_stream(
language="en", audio_bytes=long_audio_wav_bytes()
)
self.assertTrue(len(deltas) > 0, "Expected at least one streamed delta")
self._assert_keywords("".join(deltas))
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,122 @@
"""Unit tests for energy-aware audio chunking (long-audio transcription).
Whisper's encoder window is 30 s; longer audio must be split before
prompting. The splitter must cut at low-energy points (pauses) inside the
tail search window of each stride — never blindly at the stride boundary,
which could land mid-word — and the chunks must be contiguous,
non-overlapping, and reproduce the full waveform when concatenated.
"""
from sglang.test.test_utils import maybe_stub_sgl_kernel
maybe_stub_sgl_kernel() # must precede any import that pulls in sgl_kernel
import io
import unittest
import numpy as np
import soundfile as sf
from sglang.srt.entrypoints.openai.audio_chunking import (
find_split_point,
split_audio_energy_aware,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=4, suite="base-a-test-cpu")
SR = 16000
def _tone_with_silences(duration_s: float, silences: list, sr: int = SR) -> np.ndarray:
"""A 440 Hz tone with zeroed-out gaps at the given (start_s, end_s) spans."""
t = np.arange(int(duration_s * sr)) / sr
wav = (0.5 * np.sin(2 * np.pi * 440.0 * t)).astype(np.float32)
for start_s, end_s in silences:
wav[int(start_s * sr) : int(end_s * sr)] = 0.0
return wav
def _wav_bytes(wav: np.ndarray, sr: int = SR) -> bytes:
buf = io.BytesIO()
sf.write(buf, wav, sr, format="WAV")
return buf.getvalue()
def _decode(chunk: bytes) -> np.ndarray:
data, sr = sf.read(io.BytesIO(chunk), dtype="float32")
assert sr == SR
return data
class TestFindSplitPoint(CustomTestCase):
def test_picks_quietest_window(self):
# Loud tone with one silent 200 ms gap; the split point must land
# inside the gap.
wav = _tone_with_silences(4.0, [(2.0, 2.2)])
idx = find_split_point(wav, int(1.0 * SR), int(3.0 * SR))
self.assertGreaterEqual(idx, int(2.0 * SR))
self.assertLess(idx, int(2.2 * SR))
def test_uniform_energy_returns_search_start(self):
wav = _tone_with_silences(4.0, [])
idx = find_split_point(wav, int(1.0 * SR), int(3.0 * SR))
# All windows are equally loud; the first one wins.
self.assertEqual(idx, int(1.0 * SR))
class TestSplitAudioEnergyAware(CustomTestCase):
def test_short_audio_single_chunk(self):
wav = _tone_with_silences(10.0, [])
chunks, offsets = split_audio_energy_aware(_wav_bytes(wav), max_clip_s=30.0)
self.assertEqual(len(chunks), 1)
self.assertEqual(offsets, [0.0])
self.assertEqual(len(_decode(chunks[0])), len(wav))
def test_long_audio_splits_at_silence(self):
# 70 s tone with silence gaps planted inside each stride's search
# window (the last 1 s before the 30 s boundary): the first cut must
# land in [29.2, 29.5], which puts the second search window at
# ~[58.2, 59.2] where the second gap [58.5, 58.8] lives.
gaps = [(29.2, 29.5), (58.5, 58.8)]
wav = _tone_with_silences(70.0, gaps)
chunks, offsets = split_audio_energy_aware(_wav_bytes(wav), max_clip_s=30.0)
self.assertEqual(len(chunks), 3)
self.assertEqual(offsets[0], 0.0)
# Each cut lands inside its silence gap, not at the blind 30 s mark.
for offset, (gap_start, gap_end) in zip(offsets[1:], gaps):
self.assertGreaterEqual(offset, gap_start)
self.assertLess(offset, gap_end)
decoded = [_decode(c) for c in chunks]
# No chunk exceeds the model's window.
for d in decoded:
self.assertLessEqual(len(d), 30 * SR)
# Chunks are contiguous and non-overlapping: offsets line up with
# cumulative chunk lengths and the total sample count is preserved.
cumulative = 0
for d, offset in zip(decoded, offsets):
self.assertEqual(cumulative, int(round(offset * SR)))
cumulative += len(d)
self.assertEqual(cumulative, len(wav))
# Concatenation reproduces the waveform (modulo PCM16 quantization).
stitched = np.concatenate(decoded)
np.testing.assert_allclose(stitched, wav, atol=2.0 / 32768)
def test_no_silence_still_makes_progress(self):
# A constant-energy tone has no preferred split point; the splitter
# must still terminate with bounded chunks covering all samples.
wav = _tone_with_silences(65.0, [])
chunks, offsets = split_audio_energy_aware(_wav_bytes(wav), max_clip_s=30.0)
self.assertGreaterEqual(len(chunks), 3)
decoded = [_decode(c) for c in chunks]
for d in decoded:
self.assertLessEqual(len(d), 30 * SR)
self.assertEqual(sum(len(d) for d in decoded), len(wav))
self.assertEqual(len(offsets), len(chunks))
if __name__ == "__main__":
unittest.main()
@@ -14,12 +14,20 @@ from sglang.test.test_utils import maybe_stub_sgl_kernel
maybe_stub_sgl_kernel() # must precede any import that pulls in sgl_kernel
import asyncio
import io
import json
import unittest
from typing import List
from unittest.mock import Mock
from unittest.mock import AsyncMock, Mock, patch
from sglang.srt.entrypoints.openai.protocol import TranscriptionRequest
import numpy as np
import soundfile as sf
from sglang.srt.entrypoints.openai.protocol import (
TranscriptionRequest,
TranscriptionResponse,
)
from sglang.srt.entrypoints.openai.serving_transcription import (
OpenAIServingTranscription,
)
@@ -241,6 +249,426 @@ class TestStreamingFusedAutodetect(CustomTestCase):
self.assertFalse(any("<|" in d for d in deltas))
class _MockChunkTokenizerManager:
"""Mock TM scripting one result list per dispatched request, in order."""
def __init__(self, results_per_request: List):
self.model_config = Mock()
self.model_config.hf_config = Mock()
self.model_config.hf_config.architectures = ["WhisperForConditionalGeneration"]
self.server_args = Mock(
incremental_streaming_output=False,
asr_max_concurrent_sessions=32,
)
self.request_logger = Mock(log_requests=False)
self.tokenizer = Mock()
self.requests: List[GenerateReqInput] = []
self.aborted: List[str] = []
self._results = results_per_request
self.active_dispatches = 0
self.max_active_dispatches = 0
def generate_request(self, adapted_request, raw_request):
idx = len(self.requests)
# Mimic the real generate_request assigning a rid (via
# normalize_batch_and_arguments) so abort-by-rid is exercisable.
adapted_request.rid = f"rid{idx}"
self.requests.append(adapted_request)
results = self._results[idx]
async def gen():
self.active_dispatches += 1
self.max_active_dispatches = max(
self.max_active_dispatches, self.active_dispatches
)
# Give concurrently scheduled generators a chance to overlap at
# dispatch, then record that this request reached the engine.
await asyncio.sleep(0)
self.active_dispatches -= 1
for r in results:
# A bare exception entry simulates a chunk request failing.
if isinstance(r, BaseException):
raise r
yield r
return gen()
def abort_request(self, rid: str = "", abort_all: bool = False):
self.aborted.append(rid)
def create_abort_task(self, adapted_request):
return None
def _long_wav_bytes(duration_s: float = 65.0) -> bytes:
"""A 16 kHz tone with silence gaps inside each 30 s stride's split-search
window, so the energy-aware splitter cuts at ~29.2 s and ~58.5 s."""
sr = 16000
t = np.arange(int(duration_s * sr)) / sr
wav = (0.5 * np.sin(2 * np.pi * 440.0 * t)).astype(np.float32)
for gap_start, gap_end in ((29.2, 29.5), (58.5, 58.8)):
wav[int(gap_start * sr) : int(gap_end * sr)] = 0.0
buf = io.BytesIO()
sf.write(buf, wav, sr, format="WAV")
return buf.getvalue()
class TestLongAudioChunkedNonStreaming(CustomTestCase):
"""Audio longer than Whisper's 30 s window must be split into chunk
requests and the transcripts stitched in order — without chunking the
feature extractor silently truncates everything past 30 s."""
def _create_transcription(self, tm, audio_bytes, language="en", **kwargs):
serving = OpenAIServingTranscription(tm)
loop = get_or_create_event_loop()
return loop.run_until_complete(
serving.create_transcription(
audio_data=audio_bytes,
model="whisper",
language=language,
response_format=kwargs.pop("response_format", "json"),
temperature=0.0,
stream=False,
raw_request=Mock(),
**kwargs,
)
)
def test_long_audio_is_chunked_and_stitched(self):
texts = [" part one.", " part two.", " part three."]
tm = _MockChunkTokenizerManager(
[
[{"text": t, "meta_info": {"finish_reason": {"type": "stop"}}}]
for t in texts
]
)
result = self._create_transcription(tm, _long_wav_bytes(65.0))
self.assertIsInstance(result, TranscriptionResponse, result)
# In-order plain concatenation (vLLM-parity stitching).
self.assertEqual(result.text, " part one. part two. part three.")
self.assertEqual(result.usage.seconds, 65)
self.assertEqual(len(tm.requests), 3)
self.assertEqual(tm.max_active_dispatches, 1)
# Each chunk request is independent: own audio payload, own
# sampling_params dict (the multimodal processor pops keys out of
# it per request), stream=False, audio modality.
params_ids = {id(req.sampling_params) for req in tm.requests}
self.assertEqual(len(params_ids), len(tm.requests))
total_samples = 0
for req in tm.requests:
self.assertFalse(req.stream)
self.assertEqual(req.modalities, ["audio"])
data, sr = sf.read(io.BytesIO(req.audio_data), dtype="float32")
self.assertEqual(sr, 16000)
self.assertLessEqual(len(data), 30 * 16000)
total_samples += len(data)
self.assertEqual(total_samples, 65 * 16000)
def test_chunk_failure_returns_error_and_stops_dispatch(self):
# When one chunk request fails, create_transcription must return an
# error response (not a partial transcript), abort the in-flight
# request, and leave later chunks undispatched.
tm = _MockChunkTokenizerManager(
[
[
{
"text": " part one.",
"meta_info": {"finish_reason": {"type": "stop"}},
}
],
[ValueError("chunk boom")],
[
{
"text": " part three.",
"meta_info": {"finish_reason": {"type": "stop"}},
}
],
]
)
result = self._create_transcription(tm, _long_wav_bytes(65.0))
# Error response, not a TranscriptionResponse.
self.assertNotIsInstance(result, TranscriptionResponse)
# Sequential dispatch means the third chunk never reaches the engine.
self.assertEqual(len(tm.requests), 2)
self.assertIn(tm.requests[-1].rid, tm.aborted)
self.assertEqual(tm.max_active_dispatches, 1)
def test_split_failure_returns_error_without_dispatch(self):
tm = _MockChunkTokenizerManager([])
with patch(
"sglang.srt.entrypoints.openai.serving_transcription."
"split_audio_energy_aware",
side_effect=RuntimeError("decode failed"),
):
result = self._create_transcription(tm, _long_wav_bytes(65.0))
self.assertEqual(result.status_code, 400)
self.assertIn("Failed to split audio", json.loads(result.body)["message"])
self.assertEqual(tm.requests, [])
def test_short_audio_stays_unchunked(self):
tm = _MockChunkTokenizerManager(
[[{"text": " short.", "meta_info": {"finish_reason": {"type": "stop"}}}]]
)
result = self._create_transcription(tm, _long_wav_bytes(10.0))
self.assertIsInstance(result, TranscriptionResponse, result)
self.assertEqual(result.text, " short.")
self.assertEqual(len(tm.requests), 1)
def test_chunked_fused_autodetect_first_chunk_language_wins(self):
# language=None → fused auto-detect per chunk. Each chunk carries
# its own forced prefix; the stitched text must strip all of them,
# and the reported language comes from the first chunk.
tm = _MockChunkTokenizerManager(
[
[
{
"text": "<|en|><|transcribe|><|notimestamps|> part one.",
"meta_info": {"finish_reason": {"type": "stop"}},
}
],
[
{
"text": "<|fr|><|transcribe|><|notimestamps|> part two.",
"meta_info": {"finish_reason": {"type": "stop"}},
}
],
[
{
"text": "<|en|><|transcribe|><|notimestamps|> part three.",
"meta_info": {"finish_reason": {"type": "stop"}},
}
],
]
)
result = self._create_transcription(tm, _long_wav_bytes(65.0), language=None)
self.assertIsInstance(result, TranscriptionResponse, result)
self.assertEqual(result.text, "part one. part two. part three.")
self.assertNotIn("<|", result.text)
self.assertEqual(len(tm.requests), 3)
# Every chunk request kept the fused regex constraint.
for req in tm.requests:
self.assertIn("regex", req.sampling_params)
def test_chunked_fused_spaceless_script_not_space_joined(self):
# zh/ja/th chunk texts carry no boundary whitespace; stitching must
# not inject an ASCII space the model never emitted.
tm = _MockChunkTokenizerManager(
[
[
{
"text": "<|zh|><|transcribe|><|notimestamps|>你好",
"meta_info": {"finish_reason": {"type": "stop"}},
}
],
[
{
"text": "<|zh|><|transcribe|><|notimestamps|>世界",
"meta_info": {"finish_reason": {"type": "stop"}},
}
],
]
)
result = self._create_transcription(tm, _long_wav_bytes(40.0), language=None)
self.assertIsInstance(result, TranscriptionResponse, result)
self.assertEqual(result.text, "你好世界")
self.assertEqual(len(tm.requests), 2)
def test_chunked_fused_language_uses_first_nonempty_chunk(self):
tm = _MockChunkTokenizerManager(
[
[
{
"text": "<|fr|><|transcribe|><|notimestamps|>",
"output_ids": [],
"meta_info": {"finish_reason": {"type": "stop"}},
}
],
[
{
"text": "<|en|><|transcribe|><|notimestamps|> Hello",
"output_ids": [],
"meta_info": {"finish_reason": {"type": "stop"}},
}
],
]
)
result = self._create_transcription(
tm,
_long_wav_bytes(40.0),
language=None,
response_format="verbose_json",
)
self.assertEqual(result.text, "Hello")
self.assertEqual(result.language, "en")
class TestLongAudioChunkedStreaming(CustomTestCase):
"""_generate_long_audio_stream: chunks transcribed sequentially, deltas
emitted in audio order, exactly one finish frame."""
def _run_stream(self, results_per_request, fused=False, n_chunks=2):
tm = _MockChunkTokenizerManager(results_per_request)
serving = OpenAIServingTranscription(tm)
request = TranscriptionRequest(model="whisper", stream=True)
if fused:
request._fused_autodetect = True
request._fused_ts_variant = False
request._audio_chunks = [b"chunk%d" % i for i in range(n_chunks)]
# Streaming has no segment timing, so the offsets are intentionally
# not set here — only the non-streaming verbose_json path reads them.
adapted = GenerateReqInput(
text="", modalities=["audio"], sampling_params={"temperature": 0.0}
)
raw_request = Mock()
raw_request.is_disconnected = AsyncMock(return_value=False)
async def drive():
frames = []
async for frame in serving._generate_long_audio_stream(
adapted, request, raw_request
):
frames.append(frame)
return frames
loop = get_or_create_event_loop()
return tm, request, loop.run_until_complete(drive())
@staticmethod
def _finish_reasons(frames: List[str]) -> List[str]:
out = []
for line in frames:
if not line.startswith("data: ") or line.strip() == "data: [DONE]":
continue
obj = json.loads(line[len("data: ") :])
for choice in obj.get("choices", []):
if choice.get("finish_reason"):
out.append(choice["finish_reason"])
return out
def test_chunks_streamed_in_order_with_single_finish(self):
tm, _, frames = self._run_stream(
[
[_chunk(" Hello"), _chunk(" Hello world", finish="stop")],
[_chunk(" Again"), _chunk(" Again done", finish="stop")],
]
)
self.assertEqual(
_deltas_from_sse(frames), [" Hello", " world", " Again", " done"]
)
self.assertEqual(self._finish_reasons(frames), ["stop"])
self.assertEqual(frames[-1], "data: [DONE]\n\n")
# Both chunk requests were dispatched and streamed.
self.assertEqual(len(tm.requests), 2)
self.assertTrue(all(req.stream for req in tm.requests))
def test_disconnect_between_chunks_stops_and_aborts(self):
# Client disconnects after the first chunk: the second chunk is
# never dispatched, and the in-flight request from chunk 0 is
# already done so nothing is left decoding.
tm = _MockChunkTokenizerManager(
[
[_chunk(" Hello", finish="stop")],
[_chunk(" world", finish="stop")],
]
)
serving = OpenAIServingTranscription(tm)
request = TranscriptionRequest(model="whisper", stream=True)
request._audio_chunks = [b"chunk0", b"chunk1"]
adapted = GenerateReqInput(
text="", modalities=["audio"], sampling_params={"temperature": 0.0}
)
raw_request = Mock()
# Connected for the first chunk, disconnected before the second.
raw_request.is_disconnected = AsyncMock(side_effect=[False, True])
async def drive():
return [
f
async for f in serving._generate_long_audio_stream(
adapted, request, raw_request
)
]
frames = get_or_create_event_loop().run_until_complete(drive())
self.assertEqual(_deltas_from_sse(frames), [" Hello"])
# Only the first chunk was dispatched.
self.assertEqual(len(tm.requests), 1)
def test_abnormal_chunk_finish_reason_not_masked(self):
# A non-final chunk truncated at the token cap (finish="length")
# must surface in the single final frame even though later chunks
# stop cleanly — otherwise silently missing transcript content
# reads as success.
_, _, frames = self._run_stream(
[
[_chunk(" Hello", finish="length")],
[_chunk(" world", finish="stop")],
]
)
self.assertEqual(_deltas_from_sse(frames), [" Hello", " world"])
self.assertEqual(self._finish_reasons(frames), ["length"])
def test_fused_chunks_strip_prefixes_and_preserve_boundary_space(self):
tm, request, frames = self._run_stream(
[
[
_chunk("<|fr|><|transcribe|>"),
_chunk(
"<|fr|><|transcribe|><|notimestamps|> Bonjour", finish="stop"
),
],
[
_chunk(
"<|fr|><|transcribe|><|notimestamps|> le monde",
finish="stop",
)
],
],
fused=True,
)
deltas = _deltas_from_sse(frames)
self.assertFalse(any("<|" in d for d in deltas))
# No leading space at stream start; the later chunk's own leading
# space is the seam separator.
self.assertFalse(deltas[0].startswith(" "))
self.assertEqual("".join(deltas), "Bonjour le monde")
self.assertEqual(request.language, "fr")
self.assertEqual(self._finish_reasons(frames), ["stop"])
def test_fused_spaceless_script_chunks_not_space_joined(self):
# zh/ja/th transcripts carry no boundary whitespace; the seam must
# not inject an ASCII space the model never emitted.
_, request, frames = self._run_stream(
[
[_chunk("<|zh|><|transcribe|><|notimestamps|>你好", finish="stop")],
[_chunk("<|zh|><|transcribe|><|notimestamps|>世界", finish="stop")],
],
fused=True,
)
self.assertEqual("".join(_deltas_from_sse(frames)), "你好世界")
self.assertEqual(request.language, "zh")
def test_fused_leading_silence_uses_first_nonempty_chunk_language(self):
_, request, frames = self._run_stream(
[
[_chunk("<|fr|><|transcribe|><|notimestamps|>", finish="stop")],
[
_chunk(
"<|en|><|transcribe|><|notimestamps|> Hello",
finish="stop",
)
],
],
fused=True,
)
self.assertEqual("".join(_deltas_from_sse(frames)), "Hello")
self.assertEqual(request.language, "en")
class TestStreamingIncrementalOutputMode(CustomTestCase):
"""Server runs with ``incremental_streaming_output=True``.
@@ -314,5 +314,74 @@ class TestWhisperBuildFusedAutodetectParams(CustomTestCase):
SamplingParams(**params)
class _FakeTokenizer:
"""Decodes token ids to a deterministic placeholder string."""
eos_token_id = 50257
def decode(self, ids, skip_special_tokens=True):
return " " + " ".join(f"tok{i}" for i in ids)
class TestWhisperChunkedVerboseResponse(CustomTestCase):
"""build_verbose_response_chunked for audio split into >30 s chunks.
Whisper timestamp tokens are relative to each chunk's own 30 s window,
so every chunk's segments must be shifted by the chunk's start offset
in the original audio and segment ids must keep counting across chunks.
"""
TS = WhisperAdapter.TIMESTAMP_BASE_TOKEN_ID # <|0.00|>
def test_segments_offset_by_chunk_start(self):
request = TranscriptionRequest(
model="whisper", language="en", audio_duration_s=57.06
)
from sglang.srt.entrypoints.openai.protocol import TranscriptionUsage
usage = TranscriptionUsage(seconds=58)
rets = [
# chunk 0: one closed segment [0.00 → 5.00]
{"output_ids": [100, 101, self.TS + 250]},
# chunk 1: closed segment [0.00 → 2.00] + trailing unclosed text
{"output_ids": [200, self.TS + 100, 300]},
]
resp = WhisperAdapter().build_verbose_response_chunked(
request, "你好世界", rets, [0.0, 29.3], _FakeTokenizer(), usage
)
self.assertEqual(resp.language, "en")
self.assertEqual(resp.duration, 57.06)
self.assertEqual(resp.text, "你好世界")
self.assertEqual([s.id for s in resp.segments], [0, 1, 2])
self.assertEqual(
[(s.start, s.end) for s in resp.segments],
[(0.0, 5.0), (29.3, 31.3), (31.3, 31.3)],
)
def test_single_chunk_at_zero_offset_matches_unchunked(self):
request = TranscriptionRequest(
model="whisper", language="en", audio_duration_s=10.0
)
from sglang.srt.entrypoints.openai.protocol import TranscriptionUsage
usage = TranscriptionUsage(seconds=10)
output_ids = [100, self.TS + 50]
text, segments = WhisperAdapter._parse_segments(output_ids, _FakeTokenizer())
chunked = WhisperAdapter().build_verbose_response_chunked(
request,
text,
[{"output_ids": output_ids}],
[0.0],
_FakeTokenizer(),
usage,
)
self.assertEqual(chunked.text, text)
self.assertEqual(
[(s.id, s.start, s.end) for s in chunked.segments],
[(s.id, s.start, s.end) for s in segments],
)
if __name__ == "__main__":
unittest.main()