Remove sampling info events and overlap thread file (#11300)

This commit is contained in:
Liangsheng Yin
2025-10-07 21:34:25 +08:00
committed by GitHub
parent 79d3495177
commit 501dfa6b42
9 changed files with 13 additions and 393 deletions

View File

@@ -44,12 +44,9 @@ class SamplingBatchInfo:
vocab_mask: Optional[torch.Tensor] = None
apply_mask_func: Optional[Callable[[torch.Tensor, torch.Tensor], None]] = None
# An event used for overlap schedule
sampling_info_done: Optional[threading.Event] = None
# Penalizer
penalizer_orchestrator: Optional[penaltylib.BatchedPenalizerOrchestrator] = None
linear_penalty: torch.Tensor = None
acc_linear_penalties: torch.Tensor = None # Used in the overlap mode
# Whether any request has custom logit processor
has_custom_logit_processor: bool = False
@@ -217,19 +214,19 @@ class SamplingBatchInfo:
def update_penalties(self):
if self.penalizer_orchestrator.is_required:
self.linear_penalty = torch.zeros(
self.acc_linear_penalties = torch.zeros(
(len(self.temperatures), self.vocab_size),
dtype=torch.float32,
device=self.temperatures.device,
)
self.penalizer_orchestrator.apply(self.linear_penalty)
self.penalizer_orchestrator.apply(self.acc_linear_penalties)
else:
self.linear_penalty = None
self.acc_linear_penalties = None
def apply_logits_bias(self, logits: torch.Tensor):
if self.linear_penalty is not None:
if self.acc_linear_penalties is not None:
# Used in the overlap mode
logits.add_(self.linear_penalty)
logits.add_(self.acc_linear_penalties)
if self.penalizer_orchestrator and self.penalizer_orchestrator.is_required:
# Used in the non-overlap mode
@@ -373,11 +370,7 @@ class SamplingBatchInfo:
def copy_for_forward(self):
# Accumulate the penalty into a pre-allocated buffer to get rid of the dependency of `penalizer_orchestrator` later
self.update_penalties()
return dataclasses.replace(
self,
sampling_info_done=threading.Event(),
penalizer_orchestrator=None,
)
return dataclasses.replace(self, penalizer_orchestrator=None)
def merge_bias_tensor(