Fix Req array token-id concatenation (#26182)
Co-authored-by: Lianmin Zheng <lianminzheng@gmail.com>
This commit is contained in:
co-authored by
Lianmin Zheng
parent
066b4a2180
commit
52f221cce0
@@ -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
|
||||
|
||||
@@ -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 []
|
||||
|
||||
Reference in New Issue
Block a user