Fix Req array token-id concatenation (#26182)

Co-authored-by: Lianmin Zheng <lianminzheng@gmail.com>
This commit is contained in:
Mohammad Miadh Angkad
2026-06-06 19:59:51 -07:00
committed by GitHub
co-authored by Lianmin Zheng
parent 066b4a2180
commit 52f221cce0
8 changed files with 71 additions and 50 deletions
+3 -2
View File
@@ -56,6 +56,7 @@ import logging
import multiprocessing
import os
import time
from array import array
from types import SimpleNamespace
from typing import Optional, Tuple
@@ -367,7 +368,7 @@ def prepare_inputs_for_correctness_test(bench_args, tokenizer, custom_prompts):
req = Req(
rid=i,
origin_input_text=prompts[i],
origin_input_ids=tmp_input_ids,
origin_input_ids=array("q", tmp_input_ids),
sampling_params=sampling_params,
)
req.fill_ids = req.origin_input_ids
@@ -412,7 +413,7 @@ def prepare_synthetic_inputs_for_latency_test(
req = Req(
rid=i,
origin_input_text="",
origin_input_ids=list(input_ids[i]),
origin_input_ids=array("q", input_ids[i]),
sampling_params=sampling_params,
)
req.fill_ids = req.origin_input_ids
+2 -2
View File
@@ -688,9 +688,9 @@ class Req(ReqDllmMixin):
):
# Input and output info
self.rid = rid
self.origin_input_ids = array("q", origin_input_ids)
self.origin_input_ids = origin_input_ids
self.origin_input_ids_unpadded = (
array("q", origin_input_ids_unpadded)
origin_input_ids_unpadded
if origin_input_ids_unpadded
else self.origin_input_ids
) # Before image padding
@@ -1,5 +1,6 @@
import json
from abc import ABC, abstractmethod
from array import array
from functools import lru_cache
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set
@@ -165,7 +166,7 @@ class DeepseekOCRNoRepeatNGramLogitProcessor(CustomLogitProcessor):
if ngram_size <= 0 or window_size <= 0:
continue
sequence: List[int] = req.origin_input_ids + req.output_ids
sequence = req.origin_input_ids + req.output_ids
if len(sequence) < ngram_size:
continue
@@ -175,14 +176,14 @@ class DeepseekOCRNoRepeatNGramLogitProcessor(CustomLogitProcessor):
continue
if ngram_size > 1:
current_prefix = tuple(sequence[-(ngram_size - 1) :])
current_prefix = sequence[-(ngram_size - 1) :]
else:
current_prefix = tuple()
current_prefix = array("q")
banned_tokens: Set[int] = set()
for idx in range(search_start, search_end):
ngram = sequence[idx : idx + ngram_size]
if ngram_size == 1 or tuple(ngram[:-1]) == current_prefix:
if ngram_size == 1 or ngram[:-1] == current_prefix:
banned_tokens.add(ngram[-1])
whitelist_ids = params.get("whitelist_token_ids") or []