Support thinking budget for Inkling (#33146)
Co-authored-by: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com>
This commit is contained in:
co-authored by
Xinyuan Tong
parent
955aab8db1
commit
a2161ce682
@@ -58,6 +58,16 @@ class DisallowedTokensLogitsProcessor(CustomLogitProcessor):
|
|||||||
return logits
|
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):
|
class ThinkingBudgetLogitProcessor(CustomLogitProcessor):
|
||||||
"""A logit processor that controls the length of thinking."""
|
"""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]
|
cur_ids: list[int] = [*req.origin_input_ids, *req.output_ids]
|
||||||
|
|
||||||
# Check if out of thinking stage
|
# Check if out of thinking stage
|
||||||
if (
|
start_index = _open_thinking_start(
|
||||||
self.THINKING_START_TOKEN_ID not in cur_ids
|
cur_ids, self.THINKING_START_TOKEN_ID, self.THINKING_END_TOKEN_ID
|
||||||
or self.THINKING_END_TOKEN_ID in cur_ids
|
)
|
||||||
):
|
if start_index < 0:
|
||||||
continue
|
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
|
# Count the number of tokens after the thinking start token
|
||||||
num_tokens_after_start = len(cur_ids) - start_index - 1
|
num_tokens_after_start = len(cur_ids) - start_index - 1
|
||||||
|
|
||||||
@@ -137,6 +144,14 @@ class DeepSeekR1ThinkingBudgetLogitProcessor(ThinkingBudgetLogitProcessor):
|
|||||||
NEW_LINE_TOKEN_ID: int = 201
|
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
|
# 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):
|
class DeepseekOCRNoRepeatNGramLogitProcessor(CustomLogitProcessor):
|
||||||
"""Block n-gram repetitions within a sliding window for DeepSeek-OCR outputs."""
|
"""Block n-gram repetitions within a sliding window for DeepSeek-OCR outputs."""
|
||||||
|
|||||||
@@ -12,11 +12,17 @@ from unittest.mock import MagicMock
|
|||||||
|
|
||||||
import torch
|
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 (
|
from sglang.srt.sampling.custom_logit_processor import (
|
||||||
CustomLogitProcessor,
|
CustomLogitProcessor,
|
||||||
DeepseekOCRNoRepeatNGramLogitProcessor,
|
DeepseekOCRNoRepeatNGramLogitProcessor,
|
||||||
DeepSeekR1ThinkingBudgetLogitProcessor,
|
DeepSeekR1ThinkingBudgetLogitProcessor,
|
||||||
DisallowedTokensLogitsProcessor,
|
DisallowedTokensLogitsProcessor,
|
||||||
|
InklingThinkingBudgetLogitProcessor,
|
||||||
Qwen3ThinkingBudgetLogitProcessor,
|
Qwen3ThinkingBudgetLogitProcessor,
|
||||||
_cache_from_str,
|
_cache_from_str,
|
||||||
)
|
)
|
||||||
@@ -224,21 +230,39 @@ class TestThinkingBudgetLogitProcessor(CustomTestCase):
|
|||||||
self.assertEqual(result[1, self.NL].item(), 0.0)
|
self.assertEqual(result[1, self.NL].item(), 0.0)
|
||||||
self.assertTrue(torch.isinf(result[1, 0]) and result[1, 0] < 0)
|
self.assertTrue(torch.isinf(result[1, 0]) and result[1, 0] < 0)
|
||||||
|
|
||||||
def test_multiple_thinking_start_counts_from_first(self):
|
def test_multiple_thinking_start_counts_from_most_recent(self):
|
||||||
"""Test that budget counts from the first THINKING_START occurrence."""
|
"""Test that budget counts from the most recent THINKING_START occurrence."""
|
||||||
req = _make_req(
|
req = _make_req(
|
||||||
origin_input_ids=[self.START, 100, 101],
|
origin_input_ids=[self.START, 100, 101],
|
||||||
output_ids=[self.START, 200, 201], # second START in output
|
output_ids=[self.START, 200, 201], # second START in output
|
||||||
)
|
)
|
||||||
# cur_ids = [START, 100, 101, START, 200, 201]
|
# cur_ids = [START, 100, 101, START, 200, 201]
|
||||||
# First START at index 0, tokens_after_start = 5
|
# Most recent START at index 3, tokens_after_start = 2
|
||||||
# Budget=10 → 5 < 10 → no modification
|
# Budget=4 → 2 < 4 → no modification (index 0 would have given 5)
|
||||||
params = [{"thinking_budget": 10, "__req__": req}]
|
params = [{"thinking_budget": 4, "__req__": req}]
|
||||||
logits = self._logits()
|
logits = self._logits()
|
||||||
original = logits.clone()
|
original = logits.clone()
|
||||||
result = self.processor(logits, params)
|
result = self.processor(logits, params)
|
||||||
self.assertTrue(torch.equal(result, original))
|
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):
|
def test_deepseek_r1_variant_forces_end(self):
|
||||||
"""Test DeepSeekR1 variant with its own token IDs."""
|
"""Test DeepSeekR1 variant with its own token IDs."""
|
||||||
proc = DeepSeekR1ThinkingBudgetLogitProcessor()
|
proc = DeepSeekR1ThinkingBudgetLogitProcessor()
|
||||||
|
|||||||
Reference in New Issue
Block a user