Support thinking budget for Inkling (#33146)

Co-authored-by: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com>
This commit is contained in:
Lifan Shen
2026-08-11 02:01:08 +08:00
committed by GitHub
co-authored by Xinyuan Tong
parent 955aab8db1
commit a2161ce682
2 changed files with 51 additions and 12 deletions
@@ -58,6 +58,16 @@ class DisallowedTokensLogitsProcessor(CustomLogitProcessor):
return logits
def _open_thinking_start(ids: list[int], start_id: int, end_id: int) -> int:
"""Return the index of the start token of the currently open thinking block, or -1."""
for idx in reversed(range(len(ids))):
if ids[idx] == start_id:
return idx
if ids[idx] == end_id:
return -1
return -1
class ThinkingBudgetLogitProcessor(CustomLogitProcessor):
"""A logit processor that controls the length of thinking."""
@@ -85,15 +95,12 @@ class ThinkingBudgetLogitProcessor(CustomLogitProcessor):
cur_ids: list[int] = [*req.origin_input_ids, *req.output_ids]
# Check if out of thinking stage
if (
self.THINKING_START_TOKEN_ID not in cur_ids
or self.THINKING_END_TOKEN_ID in cur_ids
):
start_index = _open_thinking_start(
cur_ids, self.THINKING_START_TOKEN_ID, self.THINKING_END_TOKEN_ID
)
if start_index < 0:
continue
# Find the index of the thinking start token
start_index = cur_ids.index(self.THINKING_START_TOKEN_ID)
# Count the number of tokens after the thinking start token
num_tokens_after_start = len(cur_ids) - start_index - 1
@@ -137,6 +144,14 @@ class DeepSeekR1ThinkingBudgetLogitProcessor(ThinkingBudgetLogitProcessor):
NEW_LINE_TOKEN_ID: int = 201
class InklingThinkingBudgetLogitProcessor(ThinkingBudgetLogitProcessor):
"""A logit processor that controls the length of thinking for Inkling models."""
THINKING_START_TOKEN_ID: int = 200008
THINKING_END_TOKEN_ID: int = 200010
NEW_LINE_TOKEN_ID: int = 198
# Adapted from DeepSeek's implementation: https://github.com/deepseek-ai/DeepSeek-OCR/blob/main/DeepSeek-OCR-master/DeepSeek-OCR-vllm/process/ngram_norepeat.py
class DeepseekOCRNoRepeatNGramLogitProcessor(CustomLogitProcessor):
"""Block n-gram repetitions within a sliding window for DeepSeek-OCR outputs."""
@@ -12,11 +12,17 @@ from unittest.mock import MagicMock
import torch
from sglang.srt.parser.inkling_tokenizer import (
CONTENT_THINKING,
END_MESSAGE,
INKLING_SPECIAL_TOKEN_IDS,
)
from sglang.srt.sampling.custom_logit_processor import (
CustomLogitProcessor,
DeepseekOCRNoRepeatNGramLogitProcessor,
DeepSeekR1ThinkingBudgetLogitProcessor,
DisallowedTokensLogitsProcessor,
InklingThinkingBudgetLogitProcessor,
Qwen3ThinkingBudgetLogitProcessor,
_cache_from_str,
)
@@ -224,21 +230,39 @@ class TestThinkingBudgetLogitProcessor(CustomTestCase):
self.assertEqual(result[1, self.NL].item(), 0.0)
self.assertTrue(torch.isinf(result[1, 0]) and result[1, 0] < 0)
def test_multiple_thinking_start_counts_from_first(self):
"""Test that budget counts from the first THINKING_START occurrence."""
def test_multiple_thinking_start_counts_from_most_recent(self):
"""Test that budget counts from the most recent THINKING_START occurrence."""
req = _make_req(
origin_input_ids=[self.START, 100, 101],
output_ids=[self.START, 200, 201], # second START in output
)
# cur_ids = [START, 100, 101, START, 200, 201]
# First START at index 0, tokens_after_start = 5
# Budget=105 < 10 → no modification
params = [{"thinking_budget": 10, "__req__": req}]
# Most recent START at index 3, tokens_after_start = 2
# Budget=42 < 4 → no modification (index 0 would have given 5)
params = [{"thinking_budget": 4, "__req__": req}]
logits = self._logits()
original = logits.clone()
result = self.processor(logits, params)
self.assertTrue(torch.equal(result, original))
def test_inkling_end_token_in_prompt_does_not_disable_budget(self):
proc = InklingThinkingBudgetLogitProcessor()
start, end = proc.THINKING_START_TOKEN_ID, proc.THINKING_END_TOKEN_ID
self.assertEqual(start, INKLING_SPECIAL_TOKEN_IDS[CONTENT_THINKING])
self.assertEqual(end, INKLING_SPECIAL_TOKEN_IDS[END_MESSAGE])
req = _make_req(
origin_input_ids=[100, end, 101, end],
output_ids=[start] + [100] * 5,
)
params = [{"thinking_budget": 5, "__req__": req}]
for forced_token in (proc.NEW_LINE_TOKEN_ID, end):
result = proc(torch.zeros(1, end + 1), params)
self.assertEqual(
torch.isfinite(result[0]).nonzero().flatten().tolist(), [forced_token]
)
req.output_ids.append(forced_token)
def test_deepseek_r1_variant_forces_end(self):
"""Test DeepSeekR1 variant with its own token IDs."""
proc = DeepSeekR1ThinkingBudgetLogitProcessor()