Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
64 changes: 43 additions & 21 deletions dingo/common/backend/logprobs.py
Original file line number Diff line number Diff line change
Expand Up @@ -227,22 +227,28 @@ def extract_prompt_logprobs_from_sglang_meta(


_SGLANG_TOP_LOGPROBS_UNSUPPORTED_MSG = (
"Dynamo's SGLang backend does not currently support logprobs >= 1 due to "
"an O(N) per-position detokenization in the upstream sglang tokenizer "
"manager. Use logprobs=0 for chosen-token logprobs, or set "
"DYN_SGL_ALLOW_TOP_LOGPROBS=1 to override at your own risk. "
"Track the upstream fix at https://github.com/sgl-project/sglang/pull/24447."
"SGLang top-k logprobs are disabled by DYN_SGL_ALLOW_TOP_LOGPROBS=0. "
"Set DYN_SGL_ALLOW_TOP_LOGPROBS=1 to enable them. SGLang versions without "
"batched top-token detokenization may incur extra latency for long outputs. "
"See the upstream optimization proposal at "
"https://github.com/sgl-project/sglang/pull/24447."
)

DYN_SGL_ALLOW_TOP_LOGPROBS_ENV = "DYN_SGL_ALLOW_TOP_LOGPROBS"


def sglang_top_logprobs_allowed() -> bool:
"""Read the ``DYN_SGL_ALLOW_TOP_LOGPROBS`` env-var gate."""
return os.environ.get(DYN_SGL_ALLOW_TOP_LOGPROBS_ENV, "").lower() not in (
"",
"""Return whether SGLang top-k logprobs are enabled.

They are enabled by default so valid OpenAI ``logprobs`` requests work.
Set ``DYN_SGL_ALLOW_TOP_LOGPROBS=0`` to restore the opt-out guard on
deployments where SGLang's per-position detokenization cost is a concern.
"""
return os.environ.get(DYN_SGL_ALLOW_TOP_LOGPROBS_ENV, "1").lower() not in (
"0",
"false",
"no",
"off",
)


Expand All @@ -255,10 +261,8 @@ def build_sglang_logprob_kwargs(
``return_logprob`` / ``top_logprobs_num`` / ``logprob_start_len`` kwargs.

Raises ``ValueError`` for ``logprobs >= 1`` when
``allow_top_logprobs`` is ``False``. SGLang's tokenizer manager
detokenizes top-k tokens serially (O(N) per generated token), so
enabling it without a batched detokenize path degrades latency
badly.
``allow_top_logprobs`` is ``False``. This is an opt-out guard for
deployments concerned about per-position top-token detokenization cost.
"""
if not output_options:
return {}
Expand Down Expand Up @@ -298,28 +302,44 @@ def extract_from_sglang_meta(
num_output_logprobs_so_far: int,
*,
return_tokens_as_token_ids: bool = False,
incremental: bool = False,
) -> tuple[Optional[list[float]], Optional[list[list[dict[str, Any]]]], int]:
"""Extract logprobs from SGLang's ``meta_info`` dict.

SGLang's ``output_token_logprobs`` / ``output_top_logprobs`` are
cumulative across stream chunks even though ``output_ids`` is
disjoint — the caller passes the running count to slice the new
entries, and the returned third element is the updated count.
When ``incremental_streaming_output`` is False (the SGLang default),
``output_token_logprobs`` / ``output_top_logprobs`` are cumulative
across stream chunks even though ``output_ids`` is disjoint — the
caller passes the running count to slice the new entries, and the
returned third element is the updated count.

When ``incremental_streaming_output`` is True (Dynamo forces this on),
SGLang sends only the new token's logprobs in each chunk — the arrays
are already disjoint, so no slicing is needed. The returned third
element is ``num_output_logprobs_so_far + len(new_entries)``.
"""
output_token_logprobs = meta_info.get("output_token_logprobs")
if not output_token_logprobs:
return None, None, num_output_logprobs_so_far

new_logprobs = output_token_logprobs[num_output_logprobs_so_far:]
if incremental:
new_logprobs = output_token_logprobs
new_total = num_output_logprobs_so_far + len(output_token_logprobs)
else:
new_logprobs = output_token_logprobs[num_output_logprobs_so_far:]
new_total = len(output_token_logprobs)

if not new_logprobs:
return None, None, num_output_logprobs_so_far
return None, None, new_total

log_probs = [float(entry[0]) for entry in new_logprobs]

top_logprobs: Optional[list[list[dict[str, Any]]]] = None
output_top = meta_info.get("output_top_logprobs")
if output_top:
new_top = output_top[num_output_logprobs_so_far:]
if incremental:
new_top = output_top
else:
new_top = output_top[num_output_logprobs_so_far:]
if new_top:
top_logprobs = []
for position_entries in new_top:
Expand All @@ -330,7 +350,9 @@ def extract_from_sglang_meta(
for rank_idx, entry in enumerate(position_entries):
tok_id = entry[1]
token_str = (
f"token_id:{tok_id}" if return_tokens_as_token_ids else entry[2]
f"token_id:{tok_id}"
if return_tokens_as_token_ids
else entry[2]
)
position_list.append(
{
Expand All @@ -342,4 +364,4 @@ def extract_from_sglang_meta(
)
top_logprobs.append(position_list)

return log_probs, top_logprobs, len(output_token_logprobs)
return log_probs, top_logprobs, new_total
51 changes: 47 additions & 4 deletions dingo/common/backend/tests/test_logprobs.py
Original file line number Diff line number Diff line change
Expand Up @@ -301,13 +301,13 @@ def test_sglang_kwargs_empty_when_no_options():


def test_sglang_kwargs_logprobs_zero_allowed_without_gate():
# The default gate forbids logprobs >= 1; logprobs=0 always works.
# Chosen-token-only logprobs are available even when top-k is disabled.
kwargs = build_sglang_logprob_kwargs({"logprobs": 0}, allow_top_logprobs=False)
assert kwargs == {"return_logprob": True, "top_logprobs_num": 0}


def test_sglang_kwargs_top_logprobs_rejected_without_gate():
with pytest.raises(ValueError, match="does not currently support logprobs >= 1"):
with pytest.raises(ValueError, match="disabled by DYN_SGL_ALLOW_TOP_LOGPROBS=0"):
build_sglang_logprob_kwargs({"logprobs": 2}, allow_top_logprobs=False)


Expand Down Expand Up @@ -336,12 +336,16 @@ def test_sglang_kwargs_both_set_picks_max():


def test_sglang_gate_reads_env(monkeypatch):
monkeypatch.delenv(DYN_SGL_ALLOW_TOP_LOGPROBS_ENV, raising=False)
assert sglang_top_logprobs_allowed() is True
monkeypatch.setenv(DYN_SGL_ALLOW_TOP_LOGPROBS_ENV, "1")
assert sglang_top_logprobs_allowed() is True
monkeypatch.setenv(DYN_SGL_ALLOW_TOP_LOGPROBS_ENV, "0")
assert sglang_top_logprobs_allowed() is False
monkeypatch.delenv(DYN_SGL_ALLOW_TOP_LOGPROBS_ENV, raising=False)
monkeypatch.setenv(DYN_SGL_ALLOW_TOP_LOGPROBS_ENV, "false")
assert sglang_top_logprobs_allowed() is False
monkeypatch.delenv(DYN_SGL_ALLOW_TOP_LOGPROBS_ENV, raising=False)
assert sglang_top_logprobs_allowed() is True


# ---------------------------------------------------------------------------
Expand Down Expand Up @@ -410,6 +414,43 @@ def test_sglang_extract_returns_offset_unchanged_when_no_new_entries():
assert new_total == 1


def test_sglang_extract_incremental_no_slice():
"""incremental=True: arrays are already disjoint, no slicing."""
meta = {
"output_token_logprobs": [(-0.5, 10, "x")],
"output_top_logprobs": [[(-0.5, 10, "x"), (-0.8, 11, "y")]],
}
log_probs, top_logprobs, new_total = extract_from_sglang_meta(
meta, 5, incremental=True
)
assert log_probs == [-0.5]
assert top_logprobs == [
[
{"rank": 1, "token_id": 10, "token": "x", "logprob": -0.5},
{"rank": 2, "token_id": 11, "token": "y", "logprob": -0.8},
]
]
assert new_total == 6


def test_sglang_extract_incremental_empty_array():
"""incremental=True: empty output_token_logprobs returns None."""
meta = {"output_token_logprobs": []}
log_probs, _, new_total = extract_from_sglang_meta(meta, 3, incremental=True)
assert log_probs is None
assert new_total == 3


def test_sglang_extract_incremental_no_top_logprobs():
"""incremental=True: log_probs without top_logprobs."""
meta = {"output_token_logprobs": [(-0.1, 1, "a")]}
log_probs, top_logprobs, new_total = extract_from_sglang_meta(
meta, 0, incremental=True
)
assert log_probs == [-0.1]
assert top_logprobs is None
assert new_total == 1

# ---------------------------------------------------------------------------
# Legacy ↔ unified behavioural parity corner cases.
#
Expand Down Expand Up @@ -656,7 +697,9 @@ def test_parity_sglang_kwargs_rejects_top_logprobs_consistently():
{"prompt_logprobs": 2},
{"logprobs": 0, "prompt_logprobs": 3},
):
with pytest.raises(ValueError, match="does not currently support"):
with pytest.raises(
ValueError, match="SGLang top-k logprobs are disabled"
):
build_sglang_logprob_kwargs(opts, allow_top_logprobs=False)


Expand Down
Loading
Loading