[Speculative] Support penalty for spec v2 overlap scheduling (#22049)
This commit is contained in:
@@ -23,6 +23,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
ForwardMode,
|
ForwardMode,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
|
from sglang.srt.sampling.penaltylib.repetition_penalty import apply_scaling_penalties
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.speculative.eagle_utils import verify_tree_greedy_func
|
from sglang.srt.speculative.eagle_utils import verify_tree_greedy_func
|
||||||
from sglang.srt.speculative.spec_utils import (
|
from sglang.srt.speculative.spec_utils import (
|
||||||
@@ -89,6 +90,25 @@ class EagleDraftInputV2Mixin:
|
|||||||
# Now seq_lens is correct
|
# Now seq_lens is correct
|
||||||
batch.maybe_wait_verify_done()
|
batch.maybe_wait_verify_done()
|
||||||
|
|
||||||
|
# Accumulate penalty
|
||||||
|
# This is a relaxed version of penalties for speculative decoding.
|
||||||
|
if batch.sampling_info.penalizer_orchestrator.is_required:
|
||||||
|
output_ids = torch.tensor(
|
||||||
|
[
|
||||||
|
(
|
||||||
|
req.output_ids[-1]
|
||||||
|
if len(req.output_ids)
|
||||||
|
else req.origin_input_ids[-1]
|
||||||
|
)
|
||||||
|
for req in batch.reqs
|
||||||
|
],
|
||||||
|
dtype=torch.int64,
|
||||||
|
device=batch.device,
|
||||||
|
)
|
||||||
|
batch.sampling_info.penalizer_orchestrator.cumulate_output_tokens(
|
||||||
|
output_ids
|
||||||
|
)
|
||||||
|
|
||||||
page_size = batch.token_to_kv_pool_allocator.page_size
|
page_size = batch.token_to_kv_pool_allocator.page_size
|
||||||
cur_kv_lens_cpu = []
|
cur_kv_lens_cpu = []
|
||||||
nxt_kv_lens_cpu = []
|
nxt_kv_lens_cpu = []
|
||||||
@@ -292,6 +312,28 @@ class EagleVerifyInputV2Mixin:
|
|||||||
next_token_logits = logits_output.next_token_logits
|
next_token_logits = logits_output.next_token_logits
|
||||||
device = batch.input_ids.device
|
device = batch.input_ids.device
|
||||||
|
|
||||||
|
# Apply penalty
|
||||||
|
# This is a relaxed version of penalties for speculative decoding.
|
||||||
|
if sampling_info.acc_additive_penalties is not None:
|
||||||
|
next_token_logits.add_(
|
||||||
|
torch.repeat_interleave(
|
||||||
|
sampling_info.acc_additive_penalties, self.draft_token_num, dim=0
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if sampling_info.acc_scaling_penalties is not None:
|
||||||
|
apply_scaling_penalties(
|
||||||
|
next_token_logits,
|
||||||
|
torch.repeat_interleave(
|
||||||
|
sampling_info.acc_scaling_penalties, self.draft_token_num, dim=0
|
||||||
|
),
|
||||||
|
)
|
||||||
|
if sampling_info.logit_bias is not None:
|
||||||
|
next_token_logits.add_(
|
||||||
|
torch.repeat_interleave(
|
||||||
|
sampling_info.logit_bias, self.draft_token_num, dim=0
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
# Apply grammar mask if provided
|
# Apply grammar mask if provided
|
||||||
if vocab_mask is not None:
|
if vocab_mask is not None:
|
||||||
assert self.grammar is not None
|
assert self.grammar is not None
|
||||||
|
|||||||
@@ -243,6 +243,41 @@ class TestEagle3ServerBase(CustomTestCase, MatchedStopMixin):
|
|||||||
res = f.result()
|
res = f.result()
|
||||||
self.assertIn("text", res, f"Server error: {res}")
|
self.assertIn("text", res, f"Server error: {res}")
|
||||||
|
|
||||||
|
def test_penalty(self):
|
||||||
|
"""Verify spec v2 handles penalty parameters without crashing."""
|
||||||
|
import concurrent.futures
|
||||||
|
|
||||||
|
args = [
|
||||||
|
{"max_new_tokens": 32},
|
||||||
|
{"max_new_tokens": 16, "frequency_penalty": 2},
|
||||||
|
{"max_new_tokens": 48, "presence_penalty": 1},
|
||||||
|
{"max_new_tokens": 8, "frequency_penalty": 0.4, "presence_penalty": 0.8},
|
||||||
|
{"max_new_tokens": 64, "frequency_penalty": -0.5, "presence_penalty": 0.3},
|
||||||
|
{"max_new_tokens": 24, "min_new_tokens": 8, "frequency_penalty": 0.4},
|
||||||
|
{"max_new_tokens": 32, "repetition_penalty": 1.5},
|
||||||
|
]
|
||||||
|
|
||||||
|
def run_decode(sampling_params):
|
||||||
|
response = requests.post(
|
||||||
|
self.base_url + "/generate",
|
||||||
|
json={
|
||||||
|
"text": "The capital of France is",
|
||||||
|
"sampling_params": sampling_params,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
res = response.json()
|
||||||
|
self.assertIn("text", res, f"Server error: {res}")
|
||||||
|
self.assertIsInstance(
|
||||||
|
res["text"],
|
||||||
|
str,
|
||||||
|
f"Expected 'text' to be str, got {type(res['text']).__name__}: {res}",
|
||||||
|
)
|
||||||
|
|
||||||
|
with concurrent.futures.ThreadPoolExecutor(max_workers=8) as pool:
|
||||||
|
list(pool.map(run_decode, args * 3))
|
||||||
|
assert self.process.poll() is None
|
||||||
|
|
||||||
|
|
||||||
class TestEagle3ServerPage(TestEagle3ServerBase):
|
class TestEagle3ServerPage(TestEagle3ServerBase):
|
||||||
other_launch_args = ["--page-size", "64"]
|
other_launch_args = ["--page-size", "64"]
|
||||||
|
|||||||
Reference in New Issue
Block a user