fix bench_speculative bug (#13197)
This commit is contained in:
@@ -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:
|
||||||
|
|||||||
Reference in New Issue
Block a user