[Feature Restoration] repetition_penalty is essential for GLM-V models (#21258)

Co-authored-by: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com>
Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
Co-authored-by: hnyls2002 <lsyincs@gmail.com>
Co-authored-by: Liangsheng Yin <hnyls2002@gmail.com>
This commit is contained in:
Yuxuan Zhang
2026-03-31 23:29:49 -07:00
committed by GitHub
co-authored by Xinyuan Tong Xinyuan Tong hnyls2002 Liangsheng Yin
parent cffc95edf4
commit 72d3d8f4cf
7 changed files with 189 additions and 41 deletions
@@ -2,10 +2,12 @@ from sglang.srt.sampling.penaltylib.frequency_penalty import BatchedFrequencyPen
from sglang.srt.sampling.penaltylib.min_new_tokens import BatchedMinNewTokensPenalizer from sglang.srt.sampling.penaltylib.min_new_tokens import BatchedMinNewTokensPenalizer
from sglang.srt.sampling.penaltylib.orchestrator import BatchedPenalizerOrchestrator from sglang.srt.sampling.penaltylib.orchestrator import BatchedPenalizerOrchestrator
from sglang.srt.sampling.penaltylib.presence_penalty import BatchedPresencePenalizer from sglang.srt.sampling.penaltylib.presence_penalty import BatchedPresencePenalizer
from sglang.srt.sampling.penaltylib.repetition_penalty import BatchedRepetitionPenalizer
__all__ = [ __all__ = [
"BatchedFrequencyPenalizer", "BatchedFrequencyPenalizer",
"BatchedMinNewTokensPenalizer", "BatchedMinNewTokensPenalizer",
"BatchedPresencePenalizer", "BatchedPresencePenalizer",
"BatchedPenalizerOrchestrator", "BatchedPenalizerOrchestrator",
"BatchedRepetitionPenalizer",
] ]
@@ -52,19 +52,56 @@ class BatchedPenalizerOrchestrator:
for penalizer in self.penalizers.values(): for penalizer in self.penalizers.values():
penalizer.cumulate_output_tokens(output_ids=output_ids) penalizer.cumulate_output_tokens(output_ids=output_ids)
def apply(self, logits: torch.Tensor) -> torch.Tensor: def apply(self, logits: torch.Tensor, repeat: Optional[int] = None):
""" """
Apply the penalizers to the logits. Apply all penalizers to the logits in-place.
Note that it may apply the penalizers in-place.
Args: Args:
logits (torch.Tensor): The logits to apply the penalizers to. logits: The logits tensor to apply penalties to.
repeat: If set (speculative decoding), per-request penalties are
Returns: expanded via repeat_interleave to match the draft token layout.
torch.Tensor: The logits after applying the penalizers. Additive penalties are captured into a zeros tensor, expanded,
then added; scaling penalties are accumulated, expanded, then
applied directly.
""" """
if repeat is None:
for penalizer in self.penalizers.values():
penalizer.apply(logits)
else:
# Additive: capture into zeros, expand, add
bs = logits.shape[0] // repeat
additive = torch.zeros(
(bs, logits.shape[1]), dtype=torch.float32, device=logits.device
)
self.accumulate_additive_penalties(additive)
logits.add_(torch.repeat_interleave(additive, repeat, dim=0))
# Scaling: accumulate, expand, apply
accumulated = self.accumulate_scaling_penalties()
if accumulated is not None:
from sglang.srt.sampling.penaltylib.repetition_penalty import (
apply_scaling_penalties,
)
expanded = torch.repeat_interleave(accumulated, repeat, dim=0)
apply_scaling_penalties(logits, expanded)
def accumulate_additive_penalties(self, logits: torch.Tensor):
"""Apply only additive (non-multiplicative) penalizers."""
for penalizer in self.penalizers.values(): for penalizer in self.penalizers.values():
penalizer.apply(logits) if not penalizer.is_multiplicative:
penalizer.apply(logits)
def accumulate_scaling_penalties(self) -> Optional[torch.Tensor]:
"""Accumulate all multiplicative penalty tensors into one, or None if none active."""
result = None
for penalizer in self.penalizers.values():
if not penalizer._is_prepared or not penalizer.is_multiplicative:
continue
if result is None:
result = penalizer.get_scaling_penalties().clone()
else:
result *= penalizer.get_scaling_penalties()
return result
def filter(self, keep_indices: torch.Tensor): def filter(self, keep_indices: torch.Tensor):
""" """
@@ -132,6 +169,8 @@ class _BatchedPenalizer(abc.ABC):
An abstract class for a batched penalizer. An abstract class for a batched penalizer.
""" """
is_multiplicative: bool = False
def __init__(self, orchestrator: BatchedPenalizerOrchestrator): def __init__(self, orchestrator: BatchedPenalizerOrchestrator):
self._orchestrator_ref: weakref.ReferenceType[BatchedPenalizerOrchestrator] = ( self._orchestrator_ref: weakref.ReferenceType[BatchedPenalizerOrchestrator] = (
weakref.ref(orchestrator) weakref.ref(orchestrator)
@@ -227,6 +266,13 @@ class _BatchedPenalizer(abc.ABC):
""" """
pass pass
def get_scaling_penalties(self) -> torch.Tensor:
"""
Return the accumulated scaling penalty tensor for multiplicative penalizers.
Only meaningful when is_multiplicative is True. Subclasses should override.
"""
raise NotImplementedError
@abc.abstractmethod @abc.abstractmethod
def _filter(self, keep_indices: torch.Tensor): def _filter(self, keep_indices: torch.Tensor):
""" """
@@ -0,0 +1,78 @@
import torch
from sglang.srt.sampling.penaltylib.orchestrator import _BatchedPenalizer
from sglang.srt.utils import get_compiler_backend
@torch.compile(dynamic=True, backend=get_compiler_backend())
def apply_scaling_penalties(logits, scaling_penalties):
logits[:] = torch.where(
logits < 0,
logits * scaling_penalties,
logits / scaling_penalties,
)
class BatchedRepetitionPenalizer(_BatchedPenalizer):
"""
Repetition penalizer penalizes tokens based on their presence in the generated output.
"""
is_multiplicative: bool = True
def _is_required(self) -> bool:
return any(
req.sampling_params.repetition_penalty != 1.0
for req in self.orchestrator.reqs()
)
def _prepare(self):
self.cumulated_repetition_penalties = torch.ones(
(len(self.orchestrator.reqs()), self.orchestrator.vocab_size),
dtype=torch.float32,
device=self.orchestrator.device,
)
self.repetition_penalties = (
torch.tensor(
data=[
req.sampling_params.repetition_penalty
for req in self.orchestrator.reqs()
],
dtype=torch.float32,
device=self.orchestrator.device,
)
).unsqueeze_(1)
def _cumulate_output_tokens(self, output_ids: torch.Tensor):
self.cumulated_repetition_penalties.scatter_(
dim=1,
index=output_ids.unsqueeze(1),
src=self.repetition_penalties,
)
def _apply(self, logits: torch.Tensor) -> torch.Tensor:
apply_scaling_penalties(logits, self.cumulated_repetition_penalties)
return logits
def get_scaling_penalties(self) -> torch.Tensor:
return self.cumulated_repetition_penalties
def _filter(self, keep_indices: torch.Tensor):
self.repetition_penalties = self.repetition_penalties[keep_indices]
self.cumulated_repetition_penalties = self.cumulated_repetition_penalties[
keep_indices
]
def _merge(self, their: "BatchedRepetitionPenalizer"):
self.repetition_penalties = torch.cat(
[self.repetition_penalties, their.repetition_penalties], dim=0
)
self.cumulated_repetition_penalties = torch.cat(
[self.cumulated_repetition_penalties, their.cumulated_repetition_penalties],
dim=0,
)
def _teardown(self) -> None:
for name in ("repetition_penalties", "cumulated_repetition_penalties"):
if hasattr(self, name):
delattr(self, name)
@@ -8,6 +8,7 @@ import torch
import sglang.srt.sampling.penaltylib as penaltylib import sglang.srt.sampling.penaltylib as penaltylib
from sglang.srt.sampling.custom_logit_processor import CustomLogitProcessor from sglang.srt.sampling.custom_logit_processor import CustomLogitProcessor
from sglang.srt.sampling.penaltylib.repetition_penalty import apply_scaling_penalties
from sglang.srt.sampling.sampling_params import TOP_K_ALL from sglang.srt.sampling.sampling_params import TOP_K_ALL
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
@@ -46,7 +47,10 @@ class SamplingBatchInfo:
# Penalizer # Penalizer
penalizer_orchestrator: Optional[penaltylib.BatchedPenalizerOrchestrator] = None penalizer_orchestrator: Optional[penaltylib.BatchedPenalizerOrchestrator] = None
acc_linear_penalties: torch.Tensor = None # Used in the overlap mode acc_additive_penalties: Optional[torch.Tensor] = None # Used in the overlap mode
acc_scaling_penalties: Optional[torch.Tensor] = (
None # Used in the overlap mode for repetition penalty
)
# Whether any request has custom logit processor # Whether any request has custom logit processor
has_custom_logit_processor: bool = False has_custom_logit_processor: bool = False
@@ -159,6 +163,7 @@ class SamplingBatchInfo:
penaltylib.BatchedFrequencyPenalizer, penaltylib.BatchedFrequencyPenalizer,
penaltylib.BatchedMinNewTokensPenalizer, penaltylib.BatchedMinNewTokensPenalizer,
penaltylib.BatchedPresencePenalizer, penaltylib.BatchedPresencePenalizer,
penaltylib.BatchedRepetitionPenalizer,
}, },
) )
@@ -229,19 +234,29 @@ class SamplingBatchInfo:
def update_penalties(self): def update_penalties(self):
if self.penalizer_orchestrator.is_required: if self.penalizer_orchestrator.is_required:
self.acc_linear_penalties = torch.zeros( self.acc_additive_penalties = torch.zeros(
(len(self.temperatures), self.vocab_size), (len(self.temperatures), self.vocab_size),
dtype=torch.float32, dtype=torch.float32,
device=self.temperatures.device, device=self.temperatures.device,
) )
self.penalizer_orchestrator.apply(self.acc_linear_penalties) self.penalizer_orchestrator.accumulate_additive_penalties(
self.acc_additive_penalties
)
self.acc_scaling_penalties = (
self.penalizer_orchestrator.accumulate_scaling_penalties()
)
else: else:
self.acc_linear_penalties = None self.acc_additive_penalties = None
self.acc_scaling_penalties = None
def apply_logits_bias(self, logits: torch.Tensor): def apply_logits_bias(self, logits: torch.Tensor):
if self.acc_linear_penalties is not None: if self.acc_additive_penalties is not None:
# Used in the overlap mode # Used in the overlap mode
logits.add_(self.acc_linear_penalties) logits.add_(self.acc_additive_penalties)
if self.acc_scaling_penalties is not None:
# Used in the overlap mode
apply_scaling_penalties(logits, self.acc_scaling_penalties)
if self.penalizer_orchestrator and self.penalizer_orchestrator.is_required: if self.penalizer_orchestrator and self.penalizer_orchestrator.is_required:
# Used in the non-overlap mode # Used in the non-overlap mode
+8 -8
View File
@@ -284,15 +284,15 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
or sampling_info.logit_bias is not None or sampling_info.logit_bias is not None
): ):
# This is a relaxed version of penalties for speculative decoding. # This is a relaxed version of penalties for speculative decoding.
linear_penalty = torch.zeros( sampling_info.penalizer_orchestrator.apply(
(bs, logits_output.next_token_logits.shape[1]), logits_output.next_token_logits, repeat=self.draft_token_num
dtype=torch.float32,
device=batch.device,
)
sampling_info.apply_logits_bias(linear_penalty)
logits_output.next_token_logits.add_(
torch.repeat_interleave(linear_penalty, self.draft_token_num, dim=0)
) )
if sampling_info.logit_bias is not None:
logits_output.next_token_logits.add_(
torch.repeat_interleave(
sampling_info.logit_bias, self.draft_token_num, dim=0
)
)
# Apply grammar mask # Apply grammar mask
if vocab_mask is not None: if vocab_mask is not None:
+12 -9
View File
@@ -398,17 +398,20 @@ class NgramVerifyInput(SpecInput):
) )
# Apply penalty # Apply penalty
if sampling_info.penalizer_orchestrator.is_required: if (
sampling_info.penalizer_orchestrator.is_required
or sampling_info.logit_bias is not None
):
# This is a relaxed version of penalties for speculative decoding. # This is a relaxed version of penalties for speculative decoding.
linear_penalty = torch.zeros( sampling_info.penalizer_orchestrator.apply(
(bs, logits_output.next_token_logits.shape[1]), logits_output.next_token_logits, repeat=self.draft_token_num
dtype=torch.float32,
device=self.device,
)
sampling_info.apply_logits_bias(linear_penalty)
logits_output.next_token_logits.add_(
torch.repeat_interleave(linear_penalty, self.draft_token_num, dim=0)
) )
if sampling_info.logit_bias is not None:
logits_output.next_token_logits.add_(
torch.repeat_interleave(
sampling_info.logit_bias, self.draft_token_num, dim=0
)
)
# Apply grammar mask # Apply grammar mask
if vocab_mask is not None: if vocab_mask is not None:
@@ -142,10 +142,10 @@ class TestMergeCustomLogitProcessor(CustomTestCase):
# apply_logits_bias # apply_logits_bias
class TestApplyLogitsBias(CustomTestCase): class TestApplyLogitsBias(CustomTestCase):
def test_applies_linear_penalties(self): def test_applies_additive_penalties(self):
"""Test that pre-accumulated linear penalties are added to logits.""" """Test that pre-accumulated additive penalties are added to logits."""
info = _make_info(batch_size=1) info = _make_info(batch_size=1)
info.acc_linear_penalties = torch.tensor([[-1.0] * VOCAB_SIZE]) info.acc_additive_penalties = torch.tensor([[-1.0] * VOCAB_SIZE])
logits = torch.zeros(1, VOCAB_SIZE) logits = torch.zeros(1, VOCAB_SIZE)
info.apply_logits_bias(logits) info.apply_logits_bias(logits)
self.assertAlmostEqual(logits[0, 0].item(), -1.0, places=5) self.assertAlmostEqual(logits[0, 0].item(), -1.0, places=5)
@@ -181,7 +181,7 @@ class TestApplyLogitsBias(CustomTestCase):
def test_no_bias_no_change(self): def test_no_bias_no_change(self):
"""Test that logits stay unchanged when no bias sources are set.""" """Test that logits stay unchanged when no bias sources are set."""
info = _make_info(batch_size=1) info = _make_info(batch_size=1)
info.acc_linear_penalties = None info.acc_additive_penalties = None
info.logit_bias = None info.logit_bias = None
info.vocab_mask = None info.vocab_mask = None
logits = torch.zeros(1, VOCAB_SIZE) logits = torch.zeros(1, VOCAB_SIZE)
@@ -194,20 +194,24 @@ class TestApplyLogitsBias(CustomTestCase):
class TestUpdatePenalties(CustomTestCase): class TestUpdatePenalties(CustomTestCase):
def test_required_creates_penalties_tensor(self): def test_required_creates_penalties_tensor(self):
"""Test that update_penalties allocates a zero tensor and calls orchestrator.apply.""" """Test that update_penalties allocates a zero tensor and calls orchestrator methods."""
orch = MagicMock(is_required=True) orch = MagicMock(is_required=True)
orch.accumulate_scaling_penalties.return_value = None
info = _make_info(batch_size=2, penalizer_orchestrator=orch) info = _make_info(batch_size=2, penalizer_orchestrator=orch)
info.update_penalties() info.update_penalties()
self.assertIsNotNone(info.acc_linear_penalties) self.assertIsNotNone(info.acc_additive_penalties)
self.assertEqual(info.acc_linear_penalties.shape, (2, VOCAB_SIZE)) self.assertEqual(info.acc_additive_penalties.shape, (2, VOCAB_SIZE))
orch.apply.assert_called_once() orch.accumulate_additive_penalties.assert_called_once_with(
info.acc_additive_penalties
)
orch.accumulate_scaling_penalties.assert_called_once()
def test_not_required_sets_none(self): def test_not_required_sets_none(self):
"""Test that update_penalties sets acc_linear_penalties to None when not required.""" """Test that update_penalties sets acc_additive_penalties to None when not required."""
orch = MagicMock(is_required=False) orch = MagicMock(is_required=False)
info = _make_info(batch_size=2, penalizer_orchestrator=orch) info = _make_info(batch_size=2, penalizer_orchestrator=orch)
info.update_penalties() info.update_penalties()
self.assertIsNone(info.acc_linear_penalties) self.assertIsNone(info.acc_additive_penalties)
# update_regex_vocab_mask # update_regex_vocab_mask