fix bench_speculative bug (#13197)

This commit is contained in:
Lzhang-hub
2025-11-20 17:09:04 +08:00
committed by GitHub
parent c8ede0e93c
commit 2847e5c4b4
+17 -7
View File
@@ -13,6 +13,7 @@ import json
import os import os
import time import time
from types import SimpleNamespace from types import SimpleNamespace
from typing import List
import numpy as np import numpy as np
import requests import requests
@@ -54,17 +55,18 @@ class FakeTokenizer:
return [] return []
def send_one_batch(base_url, num_prompts, batch_size, tokenizer, is_multimodal): def send_one_batch(base_url, num_prompts, batch_size, processor, is_multimodal):
# format: (prompt, input_len, output len). We set input_len as a dummy value 0. # format: (prompt, input_len, output len). We set input_len as a dummy value 0.
if is_multimodal: if is_multimodal:
backend = "sglang-oai-chat" backend = "sglang-oai-chat"
api_url = f"{base_url}/v1/chat/completions" api_url = f"{base_url}/v1/chat/completions"
input_requests = sample_mmmu_requests( input_requests = sample_mmmu_requests(
num_prompts, num_prompts,
tokenizer, processor,
backend=backend, backend=backend,
fixed_output_len=512, fixed_output_len=512,
) )
tokenizer = processor.tokenizer
else: else:
padded_prompts = (prompts * ((num_prompts + len(prompts) - 1) // len(prompts)))[ padded_prompts = (prompts * ((num_prompts + len(prompts) - 1) // len(prompts)))[
:num_prompts :num_prompts
@@ -74,6 +76,7 @@ def send_one_batch(base_url, num_prompts, batch_size, tokenizer, is_multimodal):
] ]
backend = "sglang" backend = "sglang"
api_url = f"{base_url}/generate" api_url = f"{base_url}/generate"
tokenizer = processor
# We need to set some dummy values in order to call `benchmark` below. # We need to set some dummy values in order to call `benchmark` below.
args = SimpleNamespace( args = SimpleNamespace(
@@ -227,14 +230,21 @@ def main(args, server_args):
}, },
) )
tokenizer = AutoTokenizer.from_pretrained( if args.is_multimodal:
args.model_path, trust_remote_code=server_args.trust_remote_code from transformers import AutoProcessor
)
processor = AutoProcessor.from_pretrained(
args.model_path, trust_remote_code=server_args.trust_remote_code
)
else:
processor = AutoTokenizer.from_pretrained(
args.model_path, trust_remote_code=server_args.trust_remote_code
)
try: try:
# Warmup # Warmup
send_one_batch( send_one_batch(
base_url, batch_size, batch_size, tokenizer, args.is_multimodal base_url, batch_size, batch_size, processor, args.is_multimodal
) )
# Benchmark # Benchmark
@@ -242,7 +252,7 @@ def main(args, server_args):
base_url, base_url,
max(args.num_prompts, batch_size), max(args.num_prompts, batch_size),
batch_size, batch_size,
tokenizer, processor,
args.is_multimodal, args.is_multimodal,
) )
finally: finally: