[Benchmark] Remove 22 unmaintained benchmarks (#34520)
This commit is contained in:
@@ -1,130 +0,0 @@
|
||||
# Benchmark with lots of common prefixes. Used to benchmark prefix caching performance.
|
||||
#
|
||||
# Launch a server:
|
||||
# python -m sglang.launch_server --model-path meta-llama/Llama-2-7b-chat-hf --port 30000 --log-level-http warning
|
||||
|
||||
import random
|
||||
import string
|
||||
import time
|
||||
|
||||
from tqdm import tqdm
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
import sglang as sgl
|
||||
from sglang import set_default_backend
|
||||
from sglang.lang.backend.runtime_endpoint import RuntimeEndpoint
|
||||
|
||||
|
||||
def generate_random_string(token_length: int) -> str:
|
||||
random_string = "".join(
|
||||
random.choices(string.ascii_letters + string.digits, k=token_length * 100)
|
||||
)
|
||||
tokenized_output = tokenizer.encode(random_string, add_special_tokens=False)[
|
||||
:token_length
|
||||
]
|
||||
|
||||
if len(tokenized_output) < token_length:
|
||||
tokenized_output = tokenized_output + [tokenizer.pad_token_id] * (
|
||||
token_length - len(tokenized_output)
|
||||
)
|
||||
|
||||
decoded_string = tokenizer.decode(tokenized_output, skip_special_tokens=False)
|
||||
return decoded_string
|
||||
|
||||
|
||||
def generate_unique_prefix(base_text, index):
|
||||
return str(index) + base_text[len(str(index)) :]
|
||||
|
||||
|
||||
@sgl.function
|
||||
def text_qa(s, question, gen_len):
|
||||
s += "Q: " + question + "\n"
|
||||
s += "A:" + sgl.gen("answer", stop="\n", temperature=0, max_tokens=gen_len)
|
||||
|
||||
|
||||
def prepare_prompts(num_prefix, num_samples_per_prefix, prefix_length, suffix_length):
|
||||
base_prefix = generate_random_string(prefix_length)
|
||||
|
||||
tot_input_len = 0
|
||||
all_prompts = []
|
||||
for i in tqdm(range(num_prefix), desc="prepare prompts"):
|
||||
unique_prefix = generate_unique_prefix(base_prefix, i)
|
||||
prompt_list = []
|
||||
for j in range(num_samples_per_prefix):
|
||||
suffix = generate_random_string(suffix_length)
|
||||
prompt = unique_prefix + suffix
|
||||
prompt_list.append(prompt)
|
||||
tot_input_len += len(tokenizer.encode(prompt))
|
||||
all_prompts.append(prompt_list)
|
||||
return all_prompts, tot_input_len
|
||||
|
||||
|
||||
def test_batch_by_batch(all_prompts, gen_len):
|
||||
backend.flush_cache()
|
||||
|
||||
tot_time = 0
|
||||
for i in range(len(all_prompts)):
|
||||
tic = time.perf_counter()
|
||||
text_qa.run_batch(
|
||||
list(zip(all_prompts[i], [gen_len] * len(all_prompts[i]))),
|
||||
)
|
||||
tot_time += time.perf_counter() - tic
|
||||
|
||||
return tot_time
|
||||
|
||||
|
||||
def test_batch_by_batch_with_hint(all_prompts, gen_len):
|
||||
backend.flush_cache()
|
||||
|
||||
tot_time = 0
|
||||
for i in range(len(all_prompts)):
|
||||
tic = time.perf_counter()
|
||||
# Send a hint to cache the prefix
|
||||
text_qa.run_batch(list(zip(all_prompts[i][:1], [gen_len])))
|
||||
# Send the batch
|
||||
text_qa.run_batch(list(zip(all_prompts[i], [gen_len] * len(all_prompts[i]))))
|
||||
|
||||
tot_time += time.perf_counter() - tic
|
||||
|
||||
return tot_time
|
||||
|
||||
|
||||
def test_send_all(all_prompts, gen_len):
|
||||
backend.flush_cache()
|
||||
|
||||
all_prompts = [x for prompt_list in all_prompts for x in prompt_list]
|
||||
|
||||
tic = time.perf_counter()
|
||||
text_qa.run_batch(
|
||||
list(zip(all_prompts, [gen_len] * len(all_prompts))),
|
||||
)
|
||||
tot_time = time.perf_counter() - tic
|
||||
|
||||
return tot_time
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
tokenizer = AutoTokenizer.from_pretrained("hf-internal-testing/llama-tokenizer")
|
||||
backend = RuntimeEndpoint("http://127.0.0.1:30000")
|
||||
set_default_backend(backend)
|
||||
|
||||
random.seed(0)
|
||||
num_prefix = 10
|
||||
num_samples_per_prefix = 32
|
||||
prefix_length = 1024
|
||||
suffix_length = 128
|
||||
gen_len = 1
|
||||
all_prompts, tot_input_len = prepare_prompts(
|
||||
num_prefix, num_samples_per_prefix, prefix_length, suffix_length
|
||||
)
|
||||
|
||||
print(f"Total input token length: {tot_input_len}\n")
|
||||
|
||||
cost = test_batch_by_batch(all_prompts, gen_len)
|
||||
print(f"Latency of test_batch_by_batch : {cost:.4f} s\n")
|
||||
|
||||
cost = test_batch_by_batch_with_hint(all_prompts, gen_len)
|
||||
print(f"Latency of test_batch_by_batch_with_hint: {cost:.4f} s\n")
|
||||
|
||||
cost = test_send_all(all_prompts, gen_len)
|
||||
print(f"Latency of test_send_all : {cost:.4f} s\n")
|
||||
@@ -1,193 +0,0 @@
|
||||
import concurrent.futures
|
||||
import os
|
||||
import random
|
||||
import time
|
||||
from concurrent.futures import ProcessPoolExecutor
|
||||
from statistics import mean
|
||||
|
||||
import requests
|
||||
from tqdm import tqdm
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from sglang.lang.backend.runtime_endpoint import RuntimeEndpoint
|
||||
|
||||
###############################################################################
|
||||
# CONFIG
|
||||
###############################################################################
|
||||
ENDPOINT_URL = "http://127.0.0.1:30000"
|
||||
TOKENIZER_DIR = "/models/meta-llama/Llama-3.2-3B"
|
||||
|
||||
# Benchmark configurations
|
||||
NUM_REQUESTS = 10 # Total number of requests (each with BATCH_SIZE prompts)
|
||||
NUM_TOKENS = 32000 # Tokens per prompt
|
||||
BATCH_SIZE = 8 # Number of prompts per request
|
||||
GEN_TOKENS = 0 # Tokens to generate per prompt
|
||||
|
||||
|
||||
###############################################################################
|
||||
# REQUEST GENERATION (in parallel)
|
||||
###############################################################################
|
||||
def generate_random_prompt(index, tokenizer_dir, num_tokens):
|
||||
"""Generate a single random prompt with specified token count."""
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_dir)
|
||||
vocab_size = tokenizer.vocab_size
|
||||
|
||||
def generate_random_text(num_toks):
|
||||
random_token_ids = [random.randint(0, vocab_size - 1) for _ in range(num_toks)]
|
||||
return tokenizer.decode(random_token_ids, clean_up_tokenization_spaces=True)
|
||||
|
||||
random_text = generate_random_text(num_tokens)
|
||||
return f"Prompt {index}: {random_text}"
|
||||
|
||||
|
||||
def prepare_all_prompts(num_requests, batch_size, num_tokens, tokenizer_dir):
|
||||
"""Generate prompts for all requests in parallel."""
|
||||
total_prompts = num_requests * batch_size
|
||||
all_prompts = [None] * total_prompts
|
||||
max_workers = min(os.cpu_count() or 1, total_prompts)
|
||||
|
||||
with ProcessPoolExecutor(max_workers=max_workers) as executor:
|
||||
futures = [
|
||||
executor.submit(generate_random_prompt, i, tokenizer_dir, num_tokens)
|
||||
for i in range(total_prompts)
|
||||
]
|
||||
for future in tqdm(
|
||||
concurrent.futures.as_completed(futures),
|
||||
total=total_prompts,
|
||||
desc="Generating prompts",
|
||||
):
|
||||
index = futures.index(future)
|
||||
all_prompts[index] = future.result()
|
||||
|
||||
batched_prompts = [
|
||||
all_prompts[i * batch_size : (i + 1) * batch_size] for i in range(num_requests)
|
||||
]
|
||||
|
||||
print(
|
||||
f"Generated {total_prompts} prompts with {num_tokens} tokens each, grouped into {num_requests} requests of {batch_size} prompts.\n"
|
||||
)
|
||||
return batched_prompts
|
||||
|
||||
|
||||
###############################################################################
|
||||
# HTTP CALLS
|
||||
###############################################################################
|
||||
def send_batch_request(endpoint, prompts, gen_tokens, request_id):
|
||||
"""Send a batch of prompts to the /generate endpoint synchronously."""
|
||||
sampling_params = {
|
||||
"max_new_tokens": gen_tokens,
|
||||
"temperature": 0.7,
|
||||
"stop": "\n",
|
||||
}
|
||||
data = {"text": prompts, "sampling_params": sampling_params}
|
||||
|
||||
start_time = time.perf_counter()
|
||||
try:
|
||||
response = requests.post(
|
||||
endpoint.base_url + "/generate", json=data, timeout=3600
|
||||
)
|
||||
if response.status_code != 200:
|
||||
error = response.json()
|
||||
raise RuntimeError(f"Request {request_id} failed: {error}")
|
||||
result = response.json()
|
||||
elapsed_time = (time.perf_counter() - start_time) * 1000 # Convert to ms
|
||||
avg_per_prompt = elapsed_time / len(prompts) if prompts else 0
|
||||
return request_id, elapsed_time, avg_per_prompt, True, len(prompts)
|
||||
except Exception as e:
|
||||
print(f"[Request] Error for request {request_id}: {e}")
|
||||
return request_id, 0, 0, False, len(prompts)
|
||||
|
||||
|
||||
def run_benchmark(endpoint, batched_prompts, batch_size, gen_tokens):
|
||||
"""Run the benchmark sequentially."""
|
||||
results = []
|
||||
num_requests = len(batched_prompts)
|
||||
|
||||
# Record start time for total latency
|
||||
benchmark_start_time = time.perf_counter()
|
||||
|
||||
for i, batch_prompts in enumerate(batched_prompts):
|
||||
request_id = i + 1
|
||||
assert (
|
||||
len(batch_prompts) == batch_size
|
||||
), f"Request {request_id} should have {batch_size} prompts, got {len(batch_prompts)}"
|
||||
|
||||
print(
|
||||
f"[Request] Sending request {request_id}/{num_requests} with {len(batch_prompts)} prompts at {int(time.time()*1000)}"
|
||||
)
|
||||
result = send_batch_request(endpoint, batch_prompts, gen_tokens, request_id)
|
||||
results.append(result)
|
||||
|
||||
# Calculate total latency
|
||||
total_latency = (time.perf_counter() - benchmark_start_time) * 1000 # Convert to ms
|
||||
|
||||
return results, total_latency
|
||||
|
||||
|
||||
###############################################################################
|
||||
# RESULTS
|
||||
###############################################################################
|
||||
def process_results(results, total_latency, num_requests):
|
||||
"""Process and display benchmark results."""
|
||||
total_time = 0
|
||||
successful_requests = 0
|
||||
failed_requests = 0
|
||||
request_latencies = []
|
||||
per_prompt_latencies = []
|
||||
total_prompts = 0
|
||||
|
||||
for request_id, elapsed_time, avg_per_prompt, success, batch_size in results:
|
||||
if success:
|
||||
successful_requests += 1
|
||||
total_prompts += batch_size
|
||||
request_latencies.append(elapsed_time)
|
||||
per_prompt_latencies.append(avg_per_prompt)
|
||||
total_time += elapsed_time / 1000 # Convert to seconds
|
||||
else:
|
||||
failed_requests += 1
|
||||
|
||||
avg_request_latency = mean(request_latencies) if request_latencies else 0
|
||||
avg_per_prompt_latency = mean(per_prompt_latencies) if per_prompt_latencies else 0
|
||||
throughput = total_prompts / total_time if total_time > 0 else 0
|
||||
|
||||
print("\nBenchmark Summary:")
|
||||
print(f" Total requests sent: {len(results)}")
|
||||
print(f" Total prompts sent: {total_prompts}")
|
||||
print(f" Successful requests: {successful_requests}")
|
||||
print(f" Failed requests: {failed_requests}")
|
||||
print(f" Total latency (all requests): {total_latency:.2f} ms")
|
||||
print(f" Avg per request latency: {avg_request_latency:.2f} ms")
|
||||
print(f" Avg per prompt latency: {avg_per_prompt_latency:.2f} ms")
|
||||
print(f" Throughput: {throughput:.2f} prompts/second\n")
|
||||
|
||||
|
||||
###############################################################################
|
||||
# MAIN
|
||||
###############################################################################
|
||||
def main():
|
||||
# Initialize endpoint
|
||||
endpoint = RuntimeEndpoint(ENDPOINT_URL)
|
||||
|
||||
# Generate prompts
|
||||
batched_prompts = prepare_all_prompts(
|
||||
NUM_REQUESTS, BATCH_SIZE, NUM_TOKENS, TOKENIZER_DIR
|
||||
)
|
||||
|
||||
# Flush cache before benchmark
|
||||
# endpoint.flush_cache()
|
||||
|
||||
# Run benchmark
|
||||
print(
|
||||
f"Starting benchmark: NUM_TOKENS={NUM_TOKENS}, BATCH_SIZE={BATCH_SIZE}, NUM_REQUESTS={NUM_REQUESTS}\n"
|
||||
)
|
||||
results, total_latency = run_benchmark(
|
||||
endpoint, batched_prompts, BATCH_SIZE, GEN_TOKENS
|
||||
)
|
||||
|
||||
# Process and display results
|
||||
process_results(results, total_latency, NUM_REQUESTS)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
random.seed(0)
|
||||
main()
|
||||
@@ -1,237 +0,0 @@
|
||||
import argparse
|
||||
import random
|
||||
import time
|
||||
from statistics import mean
|
||||
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from sglang.srt.utils.patch_tokenizer import patch_tokenizer
|
||||
|
||||
|
||||
def main():
|
||||
args = parse_args()
|
||||
|
||||
print("Tokenizer Benchmark: Sequential vs Batch Processing")
|
||||
print("-" * 60)
|
||||
print(f"Tokenizer: {args.tokenizer}")
|
||||
print(f"Functions: {', '.join(args.function)}")
|
||||
print(f"Tokens per prompt: {args.num_tokens}")
|
||||
print(f"Number of runs per batch size: {args.num_runs}")
|
||||
print(f"Batch mode: {', '.join(args.batch_mode)}")
|
||||
print("-" * 60)
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(args.tokenizer, trust_remote_code=True)
|
||||
tokenizer = patch_tokenizer(tokenizer)
|
||||
max_batch_size = max(args.batch_sizes)
|
||||
|
||||
token_ids = generate_random_token_ids(
|
||||
num_prompts=max_batch_size, num_tokens=args.num_tokens, tokenizer=tokenizer
|
||||
)
|
||||
|
||||
if "encode" in args.function:
|
||||
prompts = [
|
||||
tokenizer.decode(ids, clean_up_tokenization_spaces=True)
|
||||
for ids in token_ids
|
||||
]
|
||||
run_benchmark(
|
||||
name="encode",
|
||||
data=prompts,
|
||||
sequential_fn=lambda batch: [tokenizer.encode(p) for p in batch],
|
||||
batch_fn=lambda batch: tokenizer(batch),
|
||||
batch_sizes=args.batch_sizes,
|
||||
num_runs=args.num_runs,
|
||||
batch_mode=args.batch_mode,
|
||||
)
|
||||
|
||||
if "decode" in args.function:
|
||||
# mimic DetokenizerManager's usual case
|
||||
decode_kwargs = dict(
|
||||
skip_special_tokens=True,
|
||||
spaces_between_special_tokens=True,
|
||||
)
|
||||
run_benchmark(
|
||||
name="decode",
|
||||
data=token_ids,
|
||||
sequential_fn=lambda batch: [
|
||||
tokenizer.decode(ids, **decode_kwargs) for ids in batch
|
||||
],
|
||||
batch_fn=lambda batch: tokenizer.batch_decode(batch, **decode_kwargs),
|
||||
batch_sizes=args.batch_sizes,
|
||||
num_runs=args.num_runs,
|
||||
batch_mode=args.batch_mode,
|
||||
)
|
||||
|
||||
|
||||
def run_benchmark(
|
||||
*, name, data, sequential_fn, batch_fn, batch_sizes, num_runs, batch_mode
|
||||
):
|
||||
print("\n" + "=" * 60)
|
||||
print(f"{name.upper()} BENCHMARK")
|
||||
print("=" * 60)
|
||||
|
||||
results = [
|
||||
benchmark(
|
||||
data=data,
|
||||
batch_size=bs,
|
||||
sequential_fn=sequential_fn,
|
||||
batch_fn=batch_fn,
|
||||
num_runs=num_runs,
|
||||
batch_mode=batch_mode,
|
||||
)
|
||||
for bs in batch_sizes
|
||||
]
|
||||
print_results(results=results, func_name=name, batch_mode=batch_mode)
|
||||
|
||||
|
||||
def benchmark(*, data, batch_size, sequential_fn, batch_fn, num_runs, batch_mode):
|
||||
batch_data = data[:batch_size]
|
||||
run_single = "single" in batch_mode
|
||||
run_batch = "batch" in batch_mode
|
||||
|
||||
out = {"batch_size": batch_size}
|
||||
|
||||
if run_single:
|
||||
sequential_times = measure_times(
|
||||
fn=lambda: sequential_fn(batch_data), num_runs=num_runs
|
||||
)
|
||||
out |= {
|
||||
"avg_sequential_ms": mean(sequential_times),
|
||||
"sequential_runs": sequential_times,
|
||||
}
|
||||
|
||||
if run_batch:
|
||||
batch_times = measure_times(fn=lambda: batch_fn(batch_data), num_runs=num_runs)
|
||||
out |= {
|
||||
"avg_batch_ms": mean(batch_times),
|
||||
"batch_runs": batch_times,
|
||||
}
|
||||
|
||||
if run_single and run_batch:
|
||||
out["speedup_factor"] = (
|
||||
out["avg_sequential_ms"] / out["avg_batch_ms"]
|
||||
if out["avg_batch_ms"] > 0
|
||||
else 0
|
||||
)
|
||||
|
||||
return out
|
||||
|
||||
|
||||
def print_results(*, results, func_name, batch_mode):
|
||||
run_single = "single" in batch_mode
|
||||
run_batch = "batch" in batch_mode
|
||||
|
||||
for r in results:
|
||||
print(f"\nBatch size: {r['batch_size']}")
|
||||
if run_single:
|
||||
print_runs(
|
||||
label=f"Sequential {func_name}",
|
||||
runs=r["sequential_runs"],
|
||||
avg=r["avg_sequential_ms"],
|
||||
)
|
||||
if run_batch:
|
||||
print_runs(
|
||||
label=f"Batch {func_name}", runs=r["batch_runs"], avg=r["avg_batch_ms"]
|
||||
)
|
||||
if run_single and run_batch:
|
||||
print(f" Speedup factor: {r['speedup_factor']:.2f}x")
|
||||
|
||||
print("\n" + "=" * 60)
|
||||
print(f"SUMMARY: {func_name.upper()}")
|
||||
print("=" * 60)
|
||||
|
||||
headers = ["Batch Size"]
|
||||
if run_single:
|
||||
headers.append("Sequential (ms)")
|
||||
if run_batch:
|
||||
headers.append("Batch (ms)")
|
||||
if run_single and run_batch:
|
||||
headers.append("Speedup")
|
||||
print("".join(f"{h:<18}" for h in headers))
|
||||
print("-" * (18 * len(headers)))
|
||||
|
||||
for r in results:
|
||||
row = [f"{r['batch_size']}"]
|
||||
if run_single:
|
||||
row.append(f"{r['avg_sequential_ms']:.2f} ms")
|
||||
if run_batch:
|
||||
row.append(f"{r['avg_batch_ms']:.2f} ms")
|
||||
if run_single and run_batch:
|
||||
row.append(f"{r['speedup_factor']:.2f}x")
|
||||
print("".join(f"{v:<18}" for v in row))
|
||||
|
||||
|
||||
def print_runs(*, label, runs, avg):
|
||||
print(f" {label}:")
|
||||
for i, t in enumerate(runs):
|
||||
print(f" Run {i+1}: {t:.2f} ms")
|
||||
print(f" Average: {avg:.2f} ms")
|
||||
|
||||
|
||||
def measure_times(*, fn, num_runs):
|
||||
times = []
|
||||
for _ in range(num_runs):
|
||||
start = time.perf_counter()
|
||||
fn()
|
||||
times.append((time.perf_counter() - start) * 1000)
|
||||
return times
|
||||
|
||||
|
||||
def generate_random_token_ids(*, num_prompts, num_tokens, tokenizer):
|
||||
vocab_size = tokenizer.vocab_size
|
||||
print(f"Generating {num_prompts} random sequences with {num_tokens} tokens each...")
|
||||
return [
|
||||
[random.randint(0, vocab_size - 1) for _ in range(num_tokens)]
|
||||
for _ in range(num_prompts)
|
||||
]
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Tokenizer Benchmark: Sequential vs Batch Processing"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--tokenizer",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Tokenizer name or path (e.g. nvidia/Kimi-K2-Thinking-NVFP4)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--function",
|
||||
type=str,
|
||||
nargs="+",
|
||||
choices=["encode", "decode"],
|
||||
default=["encode", "decode"],
|
||||
help="Functions to benchmark (default: encode decode)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-tokens",
|
||||
type=int,
|
||||
default=20000,
|
||||
help="Number of tokens per prompt (default: 20000)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch-sizes",
|
||||
type=int,
|
||||
nargs="+",
|
||||
default=[1, 2, 4, 8],
|
||||
help="Batch sizes to test (default: 1 2 4 8)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--batch-mode",
|
||||
nargs="+",
|
||||
choices=["single", "batch"],
|
||||
default=["single", "batch"],
|
||||
help="Benchmark modes to run (default: single batch)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-runs",
|
||||
type=int,
|
||||
default=5,
|
||||
help="Number of runs per batch size (default: 5)",
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
random.seed(0)
|
||||
main()
|
||||
@@ -1,89 +0,0 @@
|
||||
## How to reproduce the benchmark results for SGLang v0.3.0 compared to vLLM v0.6.0
|
||||
|
||||
In short, with multi step enabled, in online scenarios that we benchmarked, the Median TTFT of vLLM is **3 times** that of SGLang, and the Median ITL is **10 times** that of SGLang. Lower Median TTFT and ITL are better. vLLM's multi-step optimization did not improve throughput while ensuring lower Median TTFT and ITL. Also, under maximum throughput benchmark, if vLLM does not set gpu util to 0.95 separately and uses the default configuration instead, its maximum throughput is **lower** than that of SGLang.
|
||||
|
||||
## Online benchmark results
|
||||
|
||||
### Llama 3.1 8B Instruct 1 x A100 80G
|
||||
|
||||
| RPS | Num prompts | Engine | Median E2E Latency | Median TTFT | Median TPOT | Median ITL |
|
||||
|------|-------------|--------|--------------------|-------------|-------------|------------|
|
||||
| 4 | 1200 | SGLang | 1564.17 | **31.98** | 13.17 | **11.93** |
|
||||
| 4 | 1200 | vLLM | 1691.97 | **100.48** | 14.14 | **129.32** |
|
||||
| 8 | 2400 | SGLang | 2175.02 | **35.68** | 17.85 | **14.41** |
|
||||
| 8 | 2400 | vLLM | 2137.16 | **120.39** | 17.09 | **158.63** |
|
||||
|
||||
### Llama 3.1 70B Insruct 4 x H100 80G
|
||||
|
||||
| RPS | Num Prompts | Engine | Median E2E Latency | Median TTFT | Median TPOT | Median ITL |
|
||||
|------|-------------|--------|--------------------|-------------|-------------|------------|
|
||||
| 4 | 1200 | SGLang | 3005.24 | **53.94** | 25.03 | **21.67** |
|
||||
| 4 | 1200 | vLLM | 2915.60 | **179.15** | 23.58 | **231.23** |
|
||||
| 8 | 2400 | SGLang | 4064.98 | **58.11** | 33.07 | **24.45** |
|
||||
| 8 | 2400 | vLLM | 3752.38 | **207.12** | 29.15 | **275.32** |
|
||||
|
||||
## Offline benchmark results
|
||||
|
||||
### Llama 3.1 8B Instruct 1 x A100 80G
|
||||
|
||||
| RPS | Num Prompts | Engine | Request throughput | Output token throughput |
|
||||
|------|-------------|--------|--------------------|-------------------------|
|
||||
| inf | 5000 | SGLang | 22.03 | **4281.51** |
|
||||
| inf | 5000 | vLLM | 21.27 | **4132.37** |
|
||||
|
||||
### Llama 3.1 70B Insruct 4 x H100 80G
|
||||
|
||||
| RPS | Num Prompts | Engine | Request throughput | Output token throughput |
|
||||
|------|-------------|--------|--------------------|-------------------------|
|
||||
| inf | 5000 | SGLang | 19.84 | **3856.01** |
|
||||
| inf | 5000 | vLLM | 19.04 | **3700.64** |
|
||||
|
||||
## Installation
|
||||
|
||||
```bash
|
||||
# install sglang v0.3.0
|
||||
pip install --upgrade pip
|
||||
pip install "sglang[all]"==0.3.0
|
||||
pip install flashinfer -i https://flashinfer.ai/whl/cu121/torch2.4/
|
||||
|
||||
# install vllm v0.6.0
|
||||
pip install vllm==0.6.0
|
||||
```
|
||||
|
||||
## Notes
|
||||
|
||||
We referred to the reproduction method in https://github.com/vllm-project/vllm/issues/8176, and added the `--num-scheduler-steps 10` parameter when starting the vLLM server. The `gpu_memory_utilization` of vLLM is by default 0.9 at both TP 1 and TP 4, while SGLang's `mem_frac` is 0.88 at TP 1 and 0.85 at TP 4, so we manually set it to 0.88 at TP 4.
|
||||
|
||||
## Online benchmarks
|
||||
|
||||
```bash
|
||||
# Llama 3.1 8B Instruct on 1 x A100
|
||||
python -m sglang.launch_server --model-path meta-llama/Llama-3.1-8B-Instruct --enable-torch-compile --disable-radix-cache
|
||||
python -m vllm.entrypoints.openai.api_server --model meta-llama/Llama-3.1-8B-Instruct --disable-log-requests --num-scheduler-steps 10 --max_model_len 4096
|
||||
|
||||
# Llama 3.1 70B Instruct on 4 x H100
|
||||
python -m sglang.launch_server --model-path meta-llama/Llama-3.1-70B-Instruct --disable-radix-cache --tp 4
|
||||
python -m vllm.entrypoints.openai.api_server --model meta-llama/Llama-3.1-70B-Instruct --disable-log-requests --num-scheduler-steps 10 --tensor 4 --max_model_len 4096
|
||||
|
||||
# bench serving
|
||||
python3 -m sglang.bench_serving --backend sglang --dataset-name sharegpt --num-prompts 1200 --request-rate 4
|
||||
python3 -m sglang.bench_serving --backend sglang --dataset-name sharegpt --num-prompts 2400 --request-rate 8
|
||||
python3 -m sglang.bench_serving --backend vllm --dataset-name sharegpt --num-prompts 1200 --request-rate 4
|
||||
python3 -m sglang.bench_serving --backend vllm --dataset-name sharegpt --num-prompts 2400 --request-rate 8
|
||||
```
|
||||
|
||||
## Offline benchmarks
|
||||
|
||||
```bash
|
||||
# Llama 3.1 8B Instruct on 1 x A100
|
||||
python -m sglang.launch_server --model-path meta-llama/Llama-3.1-8B-Instruct --enable-torch-compile --disable-radix-cache
|
||||
python -m vllm.entrypoints.openai.api_server --model meta-llama/Llama-3.1-8B-Instruct --disable-log-requests --num-scheduler-steps 10 --max_model_len 4096
|
||||
|
||||
# Llama 3.1 70B Instruct on 4 x H100
|
||||
python -m sglang.launch_server --model-path meta-llama/Llama-3.1-70B-Instruct --disable-radix-cache --tp 4 --mem-frac 0.88
|
||||
python -m vllm.entrypoints.openai.api_server --model meta-llama/Llama-3.1-70B-Instruct --disable-log-requests --num-scheduler-steps 10 --tensor 4 --max_model_len 4096
|
||||
|
||||
# bench serving
|
||||
python3 -m sglang.bench_serving --backend sglang --dataset-name sharegpt --num-prompts 5000
|
||||
python3 -m sglang.bench_serving --backend vllm --dataset-name sharegpt --num-prompts 5000
|
||||
```
|
||||
@@ -1,15 +0,0 @@
|
||||
## Download data
|
||||
```
|
||||
git lfs clone https://huggingface.co/datasets/ceval/ceval-exam
|
||||
```
|
||||
|
||||
## Run benchmark
|
||||
|
||||
### Benchmark sglang
|
||||
```
|
||||
python -m sglang.launch_server --model-path ramblingpolymath/Qwen3-32B-W8A8 --port 30000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_sglang.py
|
||||
```
|
||||
@@ -1,138 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
import re
|
||||
import time
|
||||
|
||||
import numpy as np
|
||||
from datasets import load_dataset
|
||||
|
||||
from sglang.lang.api import set_default_backend
|
||||
from sglang.test.test_utils import (
|
||||
add_common_sglang_args_and_parse,
|
||||
select_sglang_backend,
|
||||
)
|
||||
|
||||
choices = ["A", "B", "C", "D"]
|
||||
|
||||
|
||||
def get_one_example(line, include_answer):
|
||||
res = line["question"]
|
||||
res += f"\nA. {line['A']}"
|
||||
res += f"\nB. {line['B']}"
|
||||
res += f"\nC. {line['C']}"
|
||||
res += f"\nD. {line['D']}"
|
||||
|
||||
if include_answer:
|
||||
res += f"\nAnswer: {line['answer']} \n\n"
|
||||
return res
|
||||
|
||||
|
||||
def get_few_shot_examples(lines):
|
||||
res = ""
|
||||
for line in lines:
|
||||
res += get_one_example(line, True) + "\n\n"
|
||||
return res
|
||||
|
||||
|
||||
def get_answer_value(response):
|
||||
pattern = r"(Answer:|answer:|答案是|答案是:|正确答案是:|答案:|Assistant:)\s*([A-D])(?![\w])"
|
||||
match = re.search(pattern, response)
|
||||
|
||||
if match:
|
||||
return match.group(2)
|
||||
|
||||
return random.choice(choices)
|
||||
|
||||
|
||||
def main(args):
|
||||
# Read data && Construct prompts
|
||||
arguments = []
|
||||
labels = []
|
||||
examples = "examples:\n"
|
||||
data_path = args.data_path
|
||||
for subject in os.listdir(data_path):
|
||||
subject_path = os.path.join(data_path, subject)
|
||||
if os.path.isdir(subject_path) and subject != ".git":
|
||||
dataset = load_dataset(data_path, name=subject)
|
||||
dev_lines_temp = dataset["dev"]
|
||||
val_lines_temp = dataset["val"]
|
||||
few_shot_examples = get_few_shot_examples(dev_lines_temp)
|
||||
examples += f"{few_shot_examples}"
|
||||
for val_line in val_lines_temp:
|
||||
arguments.append(
|
||||
{
|
||||
"examples": few_shot_examples,
|
||||
"question": get_one_example(val_line, False),
|
||||
}
|
||||
)
|
||||
labels.append(val_line["answer"])
|
||||
|
||||
#####################################
|
||||
######### SGL Program Begin #########
|
||||
#####################################
|
||||
|
||||
import sglang as sgl
|
||||
|
||||
@sgl.function
|
||||
def few_shot_ceval(s, examples, question):
|
||||
s += examples + question + sgl.gen("Answer")
|
||||
|
||||
#####################################
|
||||
########## SGL Program End ##########
|
||||
#####################################
|
||||
|
||||
num_questions = args.num_questions if args.num_questions else len(arguments)
|
||||
|
||||
# Select backend
|
||||
set_default_backend(select_sglang_backend(args))
|
||||
|
||||
# Run requests
|
||||
tic = time.perf_counter()
|
||||
states = few_shot_ceval.run_batch(
|
||||
arguments[:num_questions],
|
||||
temperature=0,
|
||||
num_threads=args.parallel,
|
||||
progress_bar=True,
|
||||
)
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
preds = [get_answer_value(states[i]["Answer"]) for i in range(num_questions)]
|
||||
|
||||
# Compute accuracy
|
||||
acc = np.mean(np.array(preds) == np.array(labels[:num_questions]))
|
||||
|
||||
# Compute speed
|
||||
num_output_tokens = sum(
|
||||
s.get_meta_info("Answer")["completion_tokens"] for s in states
|
||||
)
|
||||
output_throughput = num_output_tokens / latency
|
||||
|
||||
# Print results
|
||||
print(f"Accuracy: {acc:.3f}")
|
||||
print(f"Latency: {latency:.3f} s")
|
||||
print(f"Output throughput: {output_throughput:.3f} token/s")
|
||||
|
||||
# Write results
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "ceval",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"accuracy": round(acc, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"other": {
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="ceval/ceval-exam")
|
||||
parser.add_argument("--num-questions", type=int, default=None)
|
||||
args = add_common_sglang_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,51 +0,0 @@
|
||||
## Install
|
||||
|
||||
```
|
||||
pip3 install dspy-ai
|
||||
```
|
||||
|
||||
Turn off cache at https://github.com/stanfordnlp/dspy/blob/34d8420383ec752037aa271825c1d3bf391e1277/dsp/modules/cache_utils.py#L10.
|
||||
```
|
||||
cache_turn_on = False
|
||||
```
|
||||
|
||||
or set the environment variable
|
||||
|
||||
```
|
||||
export DSP_CACHEBOOL=false
|
||||
```
|
||||
|
||||
## Benchmark SGLang
|
||||
```
|
||||
python -m sglang.launch_server --model-path meta-llama/Llama-2-7b-chat-hf --port 30000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_dspy_intro.py --backend sglang
|
||||
```
|
||||
|
||||
|
||||
## Benchmark TGI
|
||||
```
|
||||
docker run --name tgi --rm -ti --gpus all --network host \
|
||||
-v /home/ubuntu/model_weights/Llama-2-7b-chat-hf:/Llama-2-7b-chat-hf \
|
||||
ghcr.io/huggingface/text-generation-inference:1.3.0 \
|
||||
--model-id /Llama-2-7b-chat-hf --num-shard 1 --trust-remote-code \
|
||||
--max-input-length 2048 --max-total-tokens 4096 \
|
||||
--port 24000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_dspy_intro.py --backend tgi
|
||||
```
|
||||
|
||||
|
||||
|
||||
## Benchmark vLLM
|
||||
```
|
||||
python3 -m vllm.entrypoints.openai.api_server --model meta-llama/Llama-2-7b-chat-hf --disable-log-requests --port 21000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_dspy_intro.py --backend vllm
|
||||
```
|
||||
@@ -1,192 +0,0 @@
|
||||
"""
|
||||
Adapted from
|
||||
https://github.com/stanfordnlp/dspy/blob/34d8420383ec752037aa271825c1d3bf391e1277/intro.ipynb#L9
|
||||
"""
|
||||
|
||||
import argparse
|
||||
|
||||
import dspy
|
||||
from dspy.datasets import HotPotQA
|
||||
|
||||
|
||||
class BasicQA(dspy.Signature):
|
||||
"""Answer questions with short factoid answers."""
|
||||
|
||||
question = dspy.InputField()
|
||||
answer = dspy.OutputField(desc="often between 1 and 5 words")
|
||||
|
||||
|
||||
class GenerateAnswer(dspy.Signature):
|
||||
"""Answer questions with short factoid answers."""
|
||||
|
||||
context = dspy.InputField(desc="may contain relevant facts")
|
||||
question = dspy.InputField()
|
||||
answer = dspy.OutputField(desc="often between 1 and 5 words")
|
||||
|
||||
|
||||
class RAG(dspy.Module):
|
||||
def __init__(self, num_passages=3):
|
||||
super().__init__()
|
||||
|
||||
self.retrieve = dspy.Retrieve(k=num_passages)
|
||||
self.generate_answer = dspy.ChainOfThought(GenerateAnswer)
|
||||
|
||||
def forward(self, question):
|
||||
context = self.retrieve(question).passages
|
||||
prediction = self.generate_answer(context=context, question=question)
|
||||
return dspy.Prediction(context=context, answer=prediction.answer)
|
||||
|
||||
|
||||
def main(args):
|
||||
# lm = dspy.OpenAI(model='gpt-3.5-turbo')
|
||||
if args.backend == "tgi":
|
||||
lm = dspy.HFClientTGI(
|
||||
model="meta-llama/Llama-2-7b-chat-hf",
|
||||
port=args.port,
|
||||
url="http://localhost",
|
||||
)
|
||||
elif args.backend == "sglang":
|
||||
lm = dspy.HFClientSGLang(
|
||||
model="meta-llama/Llama-2-7b-chat-hf",
|
||||
port=args.port,
|
||||
url="http://localhost",
|
||||
)
|
||||
elif args.backend == "vllm":
|
||||
lm = dspy.HFClientVLLM(
|
||||
model="meta-llama/Llama-2-7b-chat-hf",
|
||||
port=args.port,
|
||||
url="http://localhost",
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Invalid backend: {args.backend}")
|
||||
|
||||
colbertv2_wiki17_abstracts = dspy.ColBERTv2(
|
||||
url="http://20.102.90.50:2017/wiki17_abstracts"
|
||||
)
|
||||
dspy.settings.configure(lm=lm, rm=colbertv2_wiki17_abstracts)
|
||||
|
||||
# Load the dataset.
|
||||
dataset = HotPotQA(
|
||||
train_seed=1, train_size=20, eval_seed=2023, dev_size=args.dev_size, test_size=0
|
||||
)
|
||||
|
||||
# Tell DSPy that the 'question' field is the input. Any other fields are labels and/or metadata.
|
||||
trainset = [x.with_inputs("question") for x in dataset.train]
|
||||
devset = [x.with_inputs("question") for x in dataset.dev]
|
||||
|
||||
print(len(trainset), len(devset))
|
||||
|
||||
train_example = trainset[0]
|
||||
print(f"Question: {train_example.question}")
|
||||
print(f"Answer: {train_example.answer}")
|
||||
|
||||
dev_example = devset[18]
|
||||
print(f"Question: {dev_example.question}")
|
||||
print(f"Answer: {dev_example.answer}")
|
||||
print(f"Relevant Wikipedia Titles: {dev_example.gold_titles}")
|
||||
|
||||
print(
|
||||
f"For this dataset, training examples have input keys {train_example.inputs().keys()} and label keys {train_example.labels().keys()}"
|
||||
)
|
||||
print(
|
||||
f"For this dataset, dev examples have input keys {dev_example.inputs().keys()} and label keys {dev_example.labels().keys()}"
|
||||
)
|
||||
|
||||
# Define the predictor.
|
||||
generate_answer = dspy.Predict(BasicQA)
|
||||
|
||||
# Call the predictor on a particular input.
|
||||
pred = generate_answer(question=dev_example.question)
|
||||
|
||||
# Print the input and the prediction.
|
||||
print(f"Question: {dev_example.question}")
|
||||
print(f"Predicted Answer: {pred.answer}")
|
||||
|
||||
lm.inspect_history(n=1)
|
||||
|
||||
# Define the predictor. Notice we're just changing the class. The signature BasicQA is unchanged.
|
||||
generate_answer_with_chain_of_thought = dspy.ChainOfThought(BasicQA)
|
||||
|
||||
# Call the predictor on the same input.
|
||||
pred = generate_answer_with_chain_of_thought(question=dev_example.question)
|
||||
|
||||
# Print the input, the chain of thought, and the prediction.
|
||||
print(f"Question: {dev_example.question}")
|
||||
print(f"Thought: {pred.rationale.split('.', 1)[1].strip()}")
|
||||
print(f"Predicted Answer: {pred.answer}")
|
||||
|
||||
retrieve = dspy.Retrieve(k=3)
|
||||
topK_passages = retrieve(dev_example.question).passages
|
||||
|
||||
print(
|
||||
f"Top {retrieve.k} passages for question: {dev_example.question} \n",
|
||||
"-" * 30,
|
||||
"\n",
|
||||
)
|
||||
|
||||
for idx, passage in enumerate(topK_passages):
|
||||
print(f"{idx+1}]", passage, "\n")
|
||||
|
||||
retrieve("When was the first FIFA World Cup held?").passages[0]
|
||||
|
||||
from dspy.teleprompt import BootstrapFewShot
|
||||
|
||||
# Validation logic: check that the predicted answer is correct.
|
||||
# Also check that the retrieved context does actually contain that answer.
|
||||
def validate_context_and_answer(example, pred, trace=None):
|
||||
answer_EM = dspy.evaluate.answer_exact_match(example, pred)
|
||||
answer_PM = dspy.evaluate.answer_passage_match(example, pred)
|
||||
return answer_EM and answer_PM
|
||||
|
||||
# Set up a basic teleprompter, which will compile our RAG program.
|
||||
teleprompter = BootstrapFewShot(metric=validate_context_and_answer)
|
||||
|
||||
# Compile!
|
||||
compiled_rag = teleprompter.compile(RAG(), trainset=trainset)
|
||||
|
||||
# Ask any question you like to this simple RAG program.
|
||||
my_question = "What castle did David Gregory inherit?"
|
||||
|
||||
# Get the prediction. This contains `pred.context` and `pred.answer`.
|
||||
pred = compiled_rag(my_question)
|
||||
|
||||
# Print the contexts and the answer.
|
||||
print(f"Question: {my_question}")
|
||||
print(f"Predicted Answer: {pred.answer}")
|
||||
print(f"Retrieved Contexts (truncated): {[c[:200] + '...' for c in pred.context]}")
|
||||
|
||||
from dspy.evaluate.evaluate import Evaluate
|
||||
|
||||
# Set up the `evaluate_on_hotpotqa` function. We'll use this many times below.
|
||||
evaluate_on_hotpotqa = Evaluate(
|
||||
devset=devset,
|
||||
num_threads=args.num_threads,
|
||||
display_progress=True,
|
||||
display_table=5,
|
||||
)
|
||||
|
||||
# Evaluate the `compiled_rag` program with the `answer_exact_match` metric.
|
||||
metric = dspy.evaluate.answer_exact_match
|
||||
evaluate_on_hotpotqa(compiled_rag, metric=metric)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--port", type=int)
|
||||
parser.add_argument("--num-threads", type=int, default=32)
|
||||
parser.add_argument("--dev-size", type=int, default=150)
|
||||
parser.add_argument(
|
||||
"--backend", type=str, choices=["sglang", "tgi", "vllm"], default="sglang"
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.port is None:
|
||||
default_port = {
|
||||
"vllm": 21000,
|
||||
"lightllm": 22000,
|
||||
"tgi": 24000,
|
||||
"sglang": 30000,
|
||||
}
|
||||
args.port = default_port.get(args.backend, None)
|
||||
|
||||
main(args)
|
||||
@@ -1,38 +0,0 @@
|
||||
## Download the dataset
|
||||
|
||||
```
|
||||
wget -O agent_calls.jsonl https://drive.google.com/uc?export=download&id=19qLpD45e9JGTKF2cUjJJegwzSUEZEKht
|
||||
```
|
||||
|
||||
## Run benchmark
|
||||
|
||||
Ensure that this benchmark is run in a serial manner (using --parallel 1) to preserve any potential dependencies between requests.
|
||||
|
||||
### Benchmark sglang
|
||||
```
|
||||
python -m sglang.launch_server --model-path meta-llama/Llama-2-7b-chat-hf --port 30000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_sglang.py --num-events 1000 --parallel 1
|
||||
```
|
||||
|
||||
### Benchmark vllm
|
||||
```
|
||||
python3 -m vllm.entrypoints.api_server --tokenizer-mode auto --model meta-llama/Llama-2-7b-chat-hf --disable-log-requests --port 21000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --num-events 1000 --backend vllm --parallel 1
|
||||
```
|
||||
|
||||
### Benchmark guidance
|
||||
```
|
||||
python3 bench_other.py --num-events 1000 --backend guidance --parallel 1 --n-ctx 4096 --model-path path/to/gguf
|
||||
```
|
||||
|
||||
### Benchmark lmql
|
||||
|
||||
```
|
||||
python3 bench_other.py --num-events 1000 --backend lmql --parallel 1
|
||||
```
|
||||
@@ -1,300 +0,0 @@
|
||||
import sglang as sgl
|
||||
|
||||
# here are the top five agent functions contributing ~70% LLM calls
|
||||
# reference: https://github.com/joonspk-research/generative_agents/
|
||||
|
||||
|
||||
@sgl.function
|
||||
def poignancy_event(s, persona_name, persona_iss, event):
|
||||
s += "Here is a brief description of " + persona_name + ".\n"
|
||||
s += persona_iss + "\n"
|
||||
s += "On the scale of 1 to 10, where 1 is purely mundane (e.g., brushing teeth, making bed) and 10 is extremely poignant (e.g., a break up, college acceptance), rate the likely poignancy of the following event for"
|
||||
s += persona_name + ".\n\n"
|
||||
s += "Event: " + event
|
||||
s += "Rate (return a number between 1 to 10):"
|
||||
s += sgl.gen(name="Rate", max_tokens=2)
|
||||
|
||||
|
||||
def poignancy_event_prompt(persona_name, persona_iss, event):
|
||||
# return prompt and max_tokens
|
||||
s = ""
|
||||
s += "Here is a brief description of " + persona_name + ".\n"
|
||||
s += persona_iss + "\n"
|
||||
s += "On the scale of 1 to 10, where 1 is purely mundane (e.g., brushing teeth, making bed) and 10 is extremely poignant (e.g., a break up, college acceptance), rate the likely poignancy of the following event for"
|
||||
s += persona_name + ".\n\n"
|
||||
s += "Event: " + event
|
||||
s += "Rate (return a number between 1 to 10):"
|
||||
return {"prompt": s, "max_tokens": 2, "stop": None}
|
||||
|
||||
|
||||
@sgl.function
|
||||
def generate_event_triple(s, persona_name, action):
|
||||
s += """Task: Turn the input into (subject, predicate, object).
|
||||
Input: Sam Johnson is eating breakfast.
|
||||
Output: (Dolores Murphy, eat, breakfast)
|
||||
---
|
||||
Input: Joon Park is brewing coffee.
|
||||
Output: (Joon Park, brew, coffee)
|
||||
---
|
||||
Input: Jane Cook is sleeping.
|
||||
Output: (Jane Cook, is, sleep)
|
||||
---
|
||||
Input: Michael Bernstein is writing email on a computer.
|
||||
Output: (Michael Bernstein, write, email)
|
||||
---
|
||||
Input: Percy Liang is teaching students in a classroom.
|
||||
Output: (Percy Liang, teach, students)
|
||||
---
|
||||
Input: Merrie Morris is running on a treadmill.
|
||||
Output: (Merrie Morris, run, treadmill)
|
||||
---"""
|
||||
s += persona_name + "is" + action + ".\n"
|
||||
s += "(" + persona_name + ","
|
||||
s += sgl.gen(name="Triple", max_tokens=20, stop=")")
|
||||
|
||||
|
||||
def generate_event_triple_prompt(persona_name, action):
|
||||
s = ""
|
||||
s += """Task: Turn the input into (subject, predicate, object).
|
||||
Input: Sam Johnson is eating breakfast.
|
||||
Output: (Dolores Murphy, eat, breakfast)
|
||||
---
|
||||
Input: Joon Park is brewing coffee.
|
||||
Output: (Joon Park, brew, coffee)
|
||||
---
|
||||
Input: Jane Cook is sleeping.
|
||||
Output: (Jane Cook, is, sleep)
|
||||
---
|
||||
Input: Michael Bernstein is writing email on a computer.
|
||||
Output: (Michael Bernstein, write, email)
|
||||
---
|
||||
Input: Percy Liang is teaching students in a classroom.
|
||||
Output: (Percy Liang, teach, students)
|
||||
---
|
||||
Input: Merrie Morris is running on a treadmill.
|
||||
Output: (Merrie Morris, run, treadmill)
|
||||
---"""
|
||||
s += persona_name + "is" + action + ".\n"
|
||||
s += "(" + persona_name + ","
|
||||
return {"prompt": s, "max_tokens": 20, "stop": ")"}
|
||||
|
||||
|
||||
@sgl.function
|
||||
def generate_pronunciatio(s, action):
|
||||
s += "Convert an action description to an emoji (important: use two or less emojis).\n"
|
||||
s += "Action description: " + action + ".\n"
|
||||
s += "Emoji:" + sgl.gen(name="Emoji", max_tokens=6)
|
||||
|
||||
|
||||
def generate_pronunciatio_prompt(action):
|
||||
s = ""
|
||||
s += "Convert an action description to an emoji (important: use two or less emojis).\n"
|
||||
s += "Action description: " + action + ".\n"
|
||||
s += "Emoji:"
|
||||
return {"prompt": s, "max_tokens": 6, "stop": None}
|
||||
|
||||
|
||||
@sgl.function
|
||||
def action_location_sector(
|
||||
s,
|
||||
persona_name,
|
||||
living_sector,
|
||||
living_sector_areas,
|
||||
current_sector,
|
||||
current_sector_areas,
|
||||
daily_plan,
|
||||
sector_options,
|
||||
current_action,
|
||||
next_action,
|
||||
):
|
||||
s += """Task -- choose an appropriate area from the area options for a task at hand.
|
||||
Sam Kim lives in {Sam Kim's house} that has Sam Kim's room, bathroom, kitchen.
|
||||
Sam Kim is currently in {Sam Kim's house} that has Sam Kim's room, bathroom, kitchen.
|
||||
Area options: {Sam Kim's house, The Rose and Crown Pub, Hobbs Cafe, Oak Hill College, Johnson Park, Harvey Oak Supply Store, The Willows Market and Pharmacy}.
|
||||
* Stay in the current area if the activity can be done there. Only go out if the activity needs to take place in another place.
|
||||
* Must be one of the "Area options," verbatim.
|
||||
For taking a walk, Sam Kim should go to the following area: {Johnson Park}
|
||||
---
|
||||
Jane Anderson lives in {Oak Hill College Student Dormatory} that has Jane Anderson's room.
|
||||
Jane Anderson is currently in {Oak Hill College} that has a classroom, library
|
||||
Area options: {Oak Hill College Student Dormatory, The Rose and Crown Pub, Hobbs Cafe, Oak Hill College, Johnson Park, Harvey Oak Supply Store, The Willows Market and Pharmacy}.
|
||||
* Stay in the current area if the activity can be done there. Only go out if the activity needs to take place in another place.
|
||||
* Must be one of the "Area options," verbatim.
|
||||
For eating dinner, Jane Anderson should go to the following area: {Hobbs Cafe}
|
||||
---"""
|
||||
s += (
|
||||
persona_name
|
||||
+ " lives in "
|
||||
+ living_sector
|
||||
+ " that has "
|
||||
+ living_sector_areas
|
||||
+ ".\n"
|
||||
)
|
||||
s += (
|
||||
persona_name
|
||||
+ " is currently in "
|
||||
+ current_sector
|
||||
+ " that has "
|
||||
+ current_sector_areas
|
||||
+ ".\n"
|
||||
)
|
||||
s += daily_plan + ".\n"
|
||||
s += "Area options: " + sector_options + ".\n"
|
||||
s += """* Stay in the current area if the activity can be done there. Only go out if the activity needs to take place in another place.
|
||||
* Must be one of the "Area options," verbatim.\n"""
|
||||
s += (
|
||||
persona_name
|
||||
+ " is "
|
||||
+ current_action
|
||||
+ ". For "
|
||||
+ next_action
|
||||
+ ", "
|
||||
+ persona_name
|
||||
+ " should go to the following area: {"
|
||||
)
|
||||
s += sgl.gen(name="Location", max_tokens=10, stop="}")
|
||||
|
||||
|
||||
def action_location_sector_prompt(
|
||||
persona_name,
|
||||
living_sector,
|
||||
living_sector_areas,
|
||||
current_sector,
|
||||
current_sector_areas,
|
||||
daily_plan,
|
||||
sector_options,
|
||||
current_action,
|
||||
next_action,
|
||||
):
|
||||
s = ""
|
||||
s += """Task -- choose an appropriate area from the area options for a task at hand.
|
||||
Sam Kim lives in {Sam Kim's house} that has Sam Kim's room, bathroom, kitchen.
|
||||
Sam Kim is currently in {Sam Kim's house} that has Sam Kim's room, bathroom, kitchen.
|
||||
Area options: {Sam Kim's house, The Rose and Crown Pub, Hobbs Cafe, Oak Hill College, Johnson Park, Harvey Oak Supply Store, The Willows Market and Pharmacy}.
|
||||
* Stay in the current area if the activity can be done there. Only go out if the activity needs to take place in another place.
|
||||
* Must be one of the "Area options," verbatim.
|
||||
For taking a walk, Sam Kim should go to the following area: {Johnson Park}
|
||||
---
|
||||
Jane Anderson lives in {Oak Hill College Student Dormatory} that has Jane Anderson's room.
|
||||
Jane Anderson is currently in {Oak Hill College} that has a classroom, library
|
||||
Area options: {Oak Hill College Student Dormatory, The Rose and Crown Pub, Hobbs Cafe, Oak Hill College, Johnson Park, Harvey Oak Supply Store, The Willows Market and Pharmacy}.
|
||||
* Stay in the current area if the activity can be done there. Only go out if the activity needs to take place in another place.
|
||||
* Must be one of the "Area options," verbatim.
|
||||
For eating dinner, Jane Anderson should go to the following area: {Hobbs Cafe}
|
||||
---"""
|
||||
s += (
|
||||
persona_name
|
||||
+ " lives in "
|
||||
+ living_sector
|
||||
+ " that has "
|
||||
+ living_sector_areas
|
||||
+ ".\n"
|
||||
)
|
||||
s += (
|
||||
persona_name
|
||||
+ " is currently in "
|
||||
+ current_sector
|
||||
+ " that has "
|
||||
+ current_sector_areas
|
||||
+ ".\n"
|
||||
)
|
||||
s += daily_plan + ".\n"
|
||||
s += "Area options: " + sector_options + ".\n"
|
||||
s += """* Stay in the current area if the activity can be done there. Only go out if the activity needs to take place in another place.
|
||||
* Must be one of the "Area options," verbatim.\n"""
|
||||
s += (
|
||||
persona_name
|
||||
+ " is "
|
||||
+ current_action
|
||||
+ ". For "
|
||||
+ next_action
|
||||
+ ", "
|
||||
+ persona_name
|
||||
+ " should go to the following area: {"
|
||||
)
|
||||
return {"prompt": s, "max_tokens": 10, "stop": "}"}
|
||||
|
||||
|
||||
@sgl.function
|
||||
def action_location_object(
|
||||
s, persona_name, target_sector, target_sector_areas, current_action, next_action
|
||||
):
|
||||
s += """
|
||||
Jane Anderson is in kitchen in Jane Anderson's house.
|
||||
Jane Anderson is going to Jane Anderson's house that has the following areas: {kitchen, bedroom, bathroom}
|
||||
Stay in the current area if the activity can be done there. Never go into other people's rooms unless necessary.
|
||||
For cooking, Jane Anderson should go to the following area in Jane Anderson's house:
|
||||
Answer: {kitchen}
|
||||
---
|
||||
Tom Watson is in common room in Tom Watson's apartment.
|
||||
Tom Watson is going to Hobbs Cafe that has the following areas: {cafe}
|
||||
Stay in the current area if the activity can be done there. Never go into other people's rooms unless necessary.
|
||||
For getting coffee, Tom Watson should go to the following area in Hobbs Cafe:
|
||||
Answer: {cafe}
|
||||
---"""
|
||||
s += (
|
||||
persona_name
|
||||
+ " is going to "
|
||||
+ target_sector
|
||||
+ " that has the following areas: {"
|
||||
+ target_sector_areas
|
||||
+ "}\n"
|
||||
)
|
||||
s += """* Stay in the current area if the activity can be done there.
|
||||
* NEVER go into other people's rooms unless necessary."""
|
||||
s += (
|
||||
persona_name
|
||||
+ " is "
|
||||
+ current_action
|
||||
+ ". For "
|
||||
+ next_action
|
||||
+ ", "
|
||||
+ persona_name
|
||||
+ "should go to the following area in "
|
||||
+ target_sector
|
||||
)
|
||||
s += " (MUST pick one of {" + target_sector_areas + "}):\n"
|
||||
s += "Answer: {" + sgl.gen(name="Area", max_tokens=5, stop="}")
|
||||
|
||||
|
||||
def action_location_object_prompt(
|
||||
persona_name, target_sector, target_sector_areas, current_action, next_action
|
||||
):
|
||||
s = ""
|
||||
s += """
|
||||
Jane Anderson is in kitchen in Jane Anderson's house.
|
||||
Jane Anderson is going to Jane Anderson's house that has the following areas: {kitchen, bedroom, bathroom}
|
||||
Stay in the current area if the activity can be done there. Never go into other people's rooms unless necessary.
|
||||
For cooking, Jane Anderson should go to the following area in Jane Anderson's house:
|
||||
Answer: {kitchen}
|
||||
---
|
||||
Tom Watson is in common room in Tom Watson's apartment.
|
||||
Tom Watson is going to Hobbs Cafe that has the following areas: {cafe}
|
||||
Stay in the current area if the activity can be done there. Never go into other people's rooms unless necessary.
|
||||
For getting coffee, Tom Watson should go to the following area in Hobbs Cafe:
|
||||
Answer: {cafe}
|
||||
---"""
|
||||
s += (
|
||||
persona_name
|
||||
+ " is going to "
|
||||
+ target_sector
|
||||
+ " that has the following areas: {"
|
||||
+ target_sector_areas
|
||||
+ "}\n"
|
||||
)
|
||||
s += """* Stay in the current area if the activity can be done there.
|
||||
* NEVER go into other people's rooms unless necessary."""
|
||||
s += (
|
||||
persona_name
|
||||
+ " is "
|
||||
+ current_action
|
||||
+ ". For "
|
||||
+ next_action
|
||||
+ ", "
|
||||
+ persona_name
|
||||
+ "should go to the following area in "
|
||||
+ target_sector
|
||||
)
|
||||
s += " (MUST pick one of {" + target_sector_areas + "}):\n"
|
||||
s += "Answer: {"
|
||||
return {"prompt": s, "max_tokens": 5, "stop": "}"}
|
||||
@@ -1,80 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
|
||||
from agent_functions import (
|
||||
action_location_object_prompt,
|
||||
action_location_sector_prompt,
|
||||
generate_event_triple_prompt,
|
||||
generate_pronunciatio_prompt,
|
||||
poignancy_event_prompt,
|
||||
)
|
||||
from tqdm import tqdm
|
||||
|
||||
from sglang.test.test_utils import add_common_other_args_and_parse, get_call_generate
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
|
||||
def main(args):
|
||||
lines = read_jsonl(args.data_path)[: args.num_events]
|
||||
mapping = {
|
||||
"poignancy_event": poignancy_event_prompt,
|
||||
"generate_event_triple": generate_event_triple_prompt,
|
||||
"generate_pronunciatio": generate_pronunciatio_prompt,
|
||||
"action_location_sector": action_location_sector_prompt,
|
||||
"action_location_object": action_location_object_prompt,
|
||||
}
|
||||
|
||||
arguments = [mapping[k](**v) for l in lines for k, v in l.items()]
|
||||
states = []
|
||||
|
||||
# Select backend
|
||||
call_generate = get_call_generate(args)
|
||||
|
||||
def get_one_answer(arg):
|
||||
answer = call_generate(**arg, temperature=0)
|
||||
states.append(answer)
|
||||
|
||||
async def get_one_answer_async(arg):
|
||||
answer = await call_generate(**arg, temperature=0)
|
||||
states.append(answer)
|
||||
|
||||
tic = time.perf_counter()
|
||||
# we always sequentially execute agent calls to maintain its dependency
|
||||
if args.backend != "lmql":
|
||||
for arg in tqdm(arguments):
|
||||
get_one_answer(arg)
|
||||
else:
|
||||
import asyncio
|
||||
|
||||
loop = asyncio.get_event_loop()
|
||||
for arg in tqdm(arguments):
|
||||
loop.run_until_complete(get_one_answer_async(arg))
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
print(f"Latency: {latency:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "Generative Agents",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
# to pack weighted functions as a single agent
|
||||
"num_requests": len(arguments) / len(mapping),
|
||||
"other": {
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="agent_calls.jsonl")
|
||||
parser.add_argument("--num-events", type=int, default=10)
|
||||
args = add_common_other_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,74 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
|
||||
from agent_functions import (
|
||||
action_location_object,
|
||||
action_location_sector,
|
||||
generate_event_triple,
|
||||
generate_pronunciatio,
|
||||
poignancy_event,
|
||||
)
|
||||
|
||||
import sglang as sgl
|
||||
from sglang.test.test_utils import (
|
||||
add_common_sglang_args_and_parse,
|
||||
select_sglang_backend,
|
||||
)
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
|
||||
def main(args):
|
||||
lines = read_jsonl(args.data_path)[: args.num_events]
|
||||
mapping = {
|
||||
"poignancy_event": poignancy_event,
|
||||
"generate_event_triple": generate_event_triple,
|
||||
"generate_pronunciatio": generate_pronunciatio,
|
||||
"action_location_sector": action_location_sector,
|
||||
"action_location_object": action_location_object,
|
||||
}
|
||||
arguments = [{mapping[k]: v for k, v in l.items()} for l in lines]
|
||||
|
||||
# Select backend
|
||||
backend = select_sglang_backend(args)
|
||||
sgl.set_default_backend(backend)
|
||||
|
||||
states = []
|
||||
# Run requests
|
||||
tic = time.perf_counter()
|
||||
for a in arguments:
|
||||
# only a single key in the dict
|
||||
for func, arg in a.items():
|
||||
result = func.run(**arg)
|
||||
result.sync()
|
||||
states.append(result)
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
# Compute accuracy
|
||||
print(f"Latency: {latency:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "Generative Agents",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
# to pack weighted functions as a single agent
|
||||
"num_requests": len(arguments) / len(mapping),
|
||||
"other": {
|
||||
"num_events": args.num_events,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="agent_calls.jsonl")
|
||||
parser.add_argument("--num-events", type=int, default=10)
|
||||
args = add_common_sglang_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,163 +0,0 @@
|
||||
# How to reproduce the result of GPT-OSS with SGLang
|
||||
|
||||
### Install the latest SGLang
|
||||
|
||||
```bash
|
||||
git clone https://github.com/sgl-project/sglang.git
|
||||
cd sglang
|
||||
git checkout v0.5.1.post3
|
||||
|
||||
pip install --upgrade pip
|
||||
pip install -e "python[all]"
|
||||
```
|
||||
|
||||
### Reproduce the benchmark throughput result (Batch Size 1)
|
||||
|
||||
Launch Command
|
||||
|
||||
```bash
|
||||
# MXFP4 120B on H100
|
||||
python3 -m sglang.launch_server --model openai/gpt-oss-120b --tp 8 --attention-backend triton
|
||||
|
||||
# BF16 120B on H100
|
||||
python3 -m sglang.launch_server --model lmsys/gpt-oss-120b-bf16 --tp 8 --attention-backend triton
|
||||
|
||||
# MXFP4 120B on B200
|
||||
python3 -m sglang.launch_server --model openai/gpt-oss-120b --tp 4
|
||||
|
||||
# BF16 120B on B200
|
||||
python3 -m sglang.launch_server --model lmsys/gpt-oss-120b-bf16 --tp 4
|
||||
```
|
||||
|
||||
Benchmark Command
|
||||
|
||||
```bash
|
||||
|
||||
# MXFP4 120B on H100
|
||||
python3 -m sglang.bench_one_batch_server --model openai/gpt-oss-120b --base-url http://localhost:30000 --batch-size 1 --input-len 1024 --output-len 512 --show-report
|
||||
```
|
||||
|
||||
### Reproduce the benchmark throughput result (Batch Size 32)
|
||||
|
||||
Launch Command
|
||||
|
||||
```bash
|
||||
# MXFP4 120B on H100
|
||||
python3 -m sglang.launch_server --model openai/gpt-oss-120b --tp 8
|
||||
|
||||
# BF16 120B on H100
|
||||
python3 -m sglang.launch_server --model lmsys/gpt-oss-120b-bf16 --tp 8
|
||||
|
||||
# MXFP4 120B on B200
|
||||
python3 -m sglang.launch_server --model openai/gpt-oss-120b --tp 4
|
||||
|
||||
# BF16 120B on B200
|
||||
python3 -m sglang.launch_server --model lmsys/gpt-oss-120b-bf16 --tp 4
|
||||
```
|
||||
|
||||
Benchmark Command
|
||||
|
||||
```bash
|
||||
python3 -m sglang.bench_one_batch_server --model openai/gpt-oss-120b --base-url http://localhost:30000 --batch-size 32 --input-len 1024 8192 --output-len 512 --show-report
|
||||
```
|
||||
|
||||
### Reproduce the evaluation result
|
||||
|
||||
Install gpt-oss
|
||||
|
||||
```bash
|
||||
git clone https://github.com/openai/gpt-oss.git
|
||||
cd gpt-oss
|
||||
pip install -e .
|
||||
```
|
||||
|
||||
Evaluation Command
|
||||
|
||||
```bash
|
||||
DATASET=gpqa
|
||||
BASE_URL=YOUR_BASE_URL
|
||||
OPENAI_API_KEY=dummy python -m gpt_oss.evals \
|
||||
--base-url ${BASE_URL}/v1 \
|
||||
--model dummy \
|
||||
--reasoning-effort low,medium,high \
|
||||
--eval $DATASET \
|
||||
--n-threads 1000
|
||||
```
|
||||
|
||||
### Reproduce the benchmark result of acceptance length
|
||||
> Note: On B200, if top k is 1, set `--attention-backend trtllm_mha`
|
||||
```bash
|
||||
git clone https://github.com/sgl-project/SpecForge.git
|
||||
cd SpecForge/benchmarks
|
||||
config_list=(
|
||||
"1,0,0,0"
|
||||
"1,3,1,4"
|
||||
"1,5,4,8"
|
||||
)
|
||||
python3 bench_model_speedup.py \
|
||||
--model-path openai/gpt-oss-120b \
|
||||
--speculative-draft-model-path lmsys/EAGLE3-gpt-oss-120b-bf16 \
|
||||
--port 20001 \
|
||||
--trust-remote-code \
|
||||
--mem-fraction-static 0.8 \
|
||||
--tp-size 4 \
|
||||
--attention-backend fa3 \
|
||||
--config-list "${config_list[@]}" \
|
||||
--benchmark-list mtbench:80 gsm8k:200 humaneval:200 math500:200 \
|
||||
--output lmsys_gpt-oss-120b_Eagle3_result.jsonl
|
||||
|
||||
python3 bench_model_speedup.py \
|
||||
--model-path openai/gpt-oss-120b \
|
||||
--speculative-draft-model-path nvidia/gpt-oss-120b-Eagle3 \
|
||||
--port 20001 \
|
||||
--trust-remote-code \
|
||||
--mem-fraction-static 0.8 \
|
||||
--tp-size 4 \
|
||||
--attention-backend fa3 \
|
||||
--config-list "${config_list[@]}" \
|
||||
--benchmark-list mtbench:80 gsm8k:200 humaneval:200 math500:200 \
|
||||
--output nv_gpt-oss-120b_Eagle3_result.jsonl
|
||||
```
|
||||
|
||||
### Reproduce the result of speculative decoding speedup
|
||||
|
||||
Launch Command
|
||||
|
||||
```bash
|
||||
# On Hopper:
|
||||
# - Tree decoding (topk > 1) and chain decoding (topk = 1) are supported on both FA3 and Triton backends.
|
||||
python3 -m sglang.launch_server --model openai/gpt-oss-120b --speculative-algorithm EAGLE3 --speculative-draft-model-path lmsys/EAGLE3-gpt-oss-120b-bf16 --speculative-num-steps 3 --speculative-eagle-topk 1 --speculative-num-draft-tokens 4 --tp 4
|
||||
python3 -m sglang.launch_server --model openai/gpt-oss-120b --speculative-algorithm EAGLE3 --speculative-draft-model-path lmsys/EAGLE3-gpt-oss-120b-bf16 --speculative-num-steps 5 --speculative-eagle-topk 4 --speculative-num-draft-tokens 8 --tp 4
|
||||
|
||||
# On Blackwell:
|
||||
# - Chain decoding (topk = 1) is supported on TRTLLM-MHA backend. Tree decoding (topk > 1) is in progress, stay tuned!
|
||||
# - Both tree decoding (topk > 1) and chain decoding (topk = 1) are supported on the Triton backend.
|
||||
python3 -m sglang.launch_server --model openai/gpt-oss-120b --speculative-algo EAGLE3 --speculative-draft-model-path lmsys/EAGLE3-gpt-oss-120b-bf16 --speculative-num-steps 3 --speculative-eagle-topk 1 --speculative-num-draft-tokens 4 --tp 4
|
||||
python3 -m sglang.launch_server --model openai/gpt-oss-120b --speculative-algo EAGLE3 --speculative-draft-model-path lmsys/EAGLE3-gpt-oss-120b-bf16 --speculative-num-steps 5 --speculative-eagle-topk 4 --speculative-num-draft-tokens 8 --attention-backend triton --tp 4
|
||||
```
|
||||
|
||||
Benchmark Command
|
||||
|
||||
```bash
|
||||
config_list=(
|
||||
"1,0,0,0"
|
||||
"1,3,1,4"
|
||||
"1,5,4,8"
|
||||
)
|
||||
python3 bench_model_speedup.py \
|
||||
--model-path openai/gpt-oss-120b \
|
||||
--speculative-draft-model-path lmsys/EAGLE3-gpt-oss-120b-bf16 \
|
||||
--port 20001 \
|
||||
--trust-remote-code \
|
||||
--mem-fraction-static 0.8 \
|
||||
--tp-size 4 \
|
||||
--attention-backend fa3 \
|
||||
--config-list "${config_list[@]}" \
|
||||
--benchmark-list gsm8k:200 humaneval:200 math500:200 \
|
||||
--output lmsys_gpt-oss-120b_Eagle3_result.jsonl
|
||||
```
|
||||
|
||||
We can gain the best speedup with the following settings:
|
||||
|
||||
- **1.39x** speedup with the `--speculative-num-steps 3 --speculative-eagle-topk 1 --speculative-num-draft-tokens 4` setting.
|
||||
- **1.52x** speedup with the `--speculative-num-steps 5 --speculative-eagle-topk 4 --speculative-num-draft-tokens 8` setting.
|
||||
@@ -18,40 +18,3 @@ python -m sglang.launch_server --model-path meta-llama/Llama-2-7b-chat-hf --port
|
||||
```
|
||||
python3 bench_sglang.py --num-questions 200
|
||||
```
|
||||
|
||||
|
||||
### Benchmark vllm
|
||||
```
|
||||
python3 -m vllm.entrypoints.api_server --tokenizer-mode auto --model meta-llama/Llama-2-7b-chat-hf --disable-log-requests --port 21000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --num-questions 200 --backend vllm
|
||||
```
|
||||
|
||||
|
||||
### Benchmark lightllm
|
||||
```
|
||||
# A10G
|
||||
python -m lightllm.server.api_server --tokenizer_mode auto --model_dir ~/model_weights/llama-2-7b-chat-hf --max_total_token_num 16000 --port 22000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --num-questions 200 --backend lightllm
|
||||
```
|
||||
|
||||
|
||||
### Benchmark guidance
|
||||
```
|
||||
python3 bench_other.py --num-questions 200 --backend guidance --parallel 1 --n-ctx 4096 --model-path path/to/gguf
|
||||
```
|
||||
|
||||
|
||||
### Benchmark lmql
|
||||
```
|
||||
CUDA_VISIBLE_DEVICES=0,1 lmql serve-model meta-llama/Llama-2-7b-chat-hf --cuda --port 23000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --num-questions 100 --backend lmql --parallel 2
|
||||
```
|
||||
|
||||
@@ -1,164 +0,0 @@
|
||||
import argparse
|
||||
import ast
|
||||
import asyncio
|
||||
import json
|
||||
import re
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
import numpy as np
|
||||
from datasets import load_dataset
|
||||
from tqdm import tqdm
|
||||
|
||||
from sglang.test.test_utils import add_common_other_args_and_parse, get_call_generate
|
||||
from sglang.utils import download_and_cache_file, dump_state_text, read_jsonl
|
||||
|
||||
INVALID = -9999999
|
||||
|
||||
|
||||
def get_one_example(lines, i, include_answer):
|
||||
ret = "Question: " + lines[i]["question"] + "\nAnswer:"
|
||||
if include_answer:
|
||||
ret += " " + lines[i]["answer"]
|
||||
return ret
|
||||
|
||||
|
||||
def get_few_shot_examples(lines, k):
|
||||
ret = ""
|
||||
for i in range(k):
|
||||
ret += get_one_example(lines, i, True) + "\n\n"
|
||||
return ret
|
||||
|
||||
|
||||
def get_answer_value(answer_str):
|
||||
answer_str = answer_str.replace(",", "")
|
||||
numbers = re.findall(r"\d+", answer_str)
|
||||
if len(numbers) < 1:
|
||||
return INVALID
|
||||
try:
|
||||
return ast.literal_eval(numbers[-1])
|
||||
except SyntaxError:
|
||||
return INVALID
|
||||
|
||||
|
||||
def main(args):
|
||||
# Select backend
|
||||
call_generate = get_call_generate(args)
|
||||
|
||||
# Read data
|
||||
if args.platinum:
|
||||
print("Loading GSM8K Platinum dataset from HuggingFace...")
|
||||
dataset = load_dataset("madrylab/gsm8k-platinum", "main", split="test")
|
||||
lines = [
|
||||
{"question": item["question"], "answer": item["answer"]} for item in dataset
|
||||
]
|
||||
else:
|
||||
url = "https://raw.githubusercontent.com/openai/grade-school-math/master/grade_school_math/data/test.jsonl"
|
||||
filename = download_and_cache_file(url)
|
||||
lines = list(read_jsonl(filename))
|
||||
|
||||
# Construct prompts
|
||||
num_questions = args.num_questions
|
||||
num_shots = args.num_shots
|
||||
few_shot_examples = get_few_shot_examples(lines, num_shots)
|
||||
|
||||
questions = []
|
||||
labels = []
|
||||
for i in range(len(lines[:num_questions])):
|
||||
questions.append(get_one_example(lines, i, False))
|
||||
labels.append(get_answer_value(lines[i]["answer"]))
|
||||
assert all(l != INVALID for l in labels)
|
||||
|
||||
states = [None] * len(labels)
|
||||
|
||||
# Run requests
|
||||
if args.backend != "lmql":
|
||||
# Use thread pool
|
||||
def get_one_answer(i):
|
||||
answer = call_generate(
|
||||
prompt=few_shot_examples + questions[i],
|
||||
temperature=0,
|
||||
max_tokens=256,
|
||||
stop=["Question", "Assistant:", "<|separator|>"],
|
||||
)
|
||||
states[i] = answer
|
||||
|
||||
tic = time.perf_counter()
|
||||
if args.parallel == 1:
|
||||
for i in tqdm(range(len(questions))):
|
||||
get_one_answer(i)
|
||||
else:
|
||||
with ThreadPoolExecutor(args.parallel) as executor:
|
||||
list(
|
||||
tqdm(
|
||||
executor.map(get_one_answer, list(range(len(questions)))),
|
||||
total=len(questions),
|
||||
)
|
||||
)
|
||||
|
||||
else:
|
||||
# Use asyncio
|
||||
async def batched_call(batch_size):
|
||||
for i in range(0, len(questions), batch_size):
|
||||
tasks = []
|
||||
for q in questions[i : i + batch_size]:
|
||||
tasks.append(
|
||||
call_generate(
|
||||
few_shot_examples + q,
|
||||
temperature=0,
|
||||
max_tokens=256,
|
||||
stop="Question",
|
||||
)
|
||||
)
|
||||
rets = await asyncio.gather(*tasks)
|
||||
for j in range(len(rets)):
|
||||
states[i + j] = rets[j]
|
||||
|
||||
tic = time.perf_counter()
|
||||
asyncio.run(batched_call(batch_size=args.parallel))
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
preds = []
|
||||
for i in range(len(states)):
|
||||
preds.append(get_answer_value(states[i]))
|
||||
|
||||
# Compute accuracy
|
||||
acc = np.mean(np.array(preds) == np.array(labels))
|
||||
invalid = np.mean(np.array(preds) == INVALID)
|
||||
|
||||
# Print results
|
||||
print(f"Accuracy: {acc:.3f}")
|
||||
print(f"Invalid: {invalid:.3f}")
|
||||
print(f"Latency: {latency:.3f} s")
|
||||
|
||||
# Dump results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "gsm8k-platinum" if args.platinum else "gsm8k",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"accuracy": round(acc, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"other": {
|
||||
"num_questions": args.num_questions,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--num-shots", type=int, default=5)
|
||||
parser.add_argument("--data-path", type=str, default="test.jsonl")
|
||||
parser.add_argument("--num-questions", type=int, default=200)
|
||||
parser.add_argument(
|
||||
"--platinum",
|
||||
action="store_true",
|
||||
help="Use GSM8K Platinum dataset (drop-in replacement with corrected labels)",
|
||||
)
|
||||
args = add_common_other_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -8,40 +8,3 @@ python -m sglang.launch_server --model-path meta-llama/Llama-2-7b-chat-hf --port
|
||||
```
|
||||
python3 bench_sglang.py --num-questions 200
|
||||
```
|
||||
|
||||
|
||||
### Benchmark vllm
|
||||
```
|
||||
python3 -m vllm.entrypoints.api_server --tokenizer-mode auto --model meta-llama/Llama-2-7b-chat-hf --disable-log-requests --port 21000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --num-questions 200 --backend vllm
|
||||
```
|
||||
|
||||
|
||||
### Benchmark lightllm
|
||||
```
|
||||
# A10G
|
||||
python -m lightllm.server.api_server --tokenizer_mode auto --model_dir ~/model_weights/llama-2-7b-chat-hf --max_total_token_num 16000 --port 22000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --num-questions 200 --backend lightllm
|
||||
```
|
||||
|
||||
|
||||
### Benchmark guidance
|
||||
```
|
||||
CUDA_VISIBLE_DEVICES=0,1 python3 bench_other.py --num-questions 200 --backend guidance --parallel 1 --n-ctx 4096 --model-path path/to/gguf
|
||||
```
|
||||
|
||||
|
||||
### Benchmark lmql
|
||||
```
|
||||
lmql serve-model meta-llama/Llama-2-7b-chat-hf --cuda --port 23000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --num-questions 200 --backend lmql --port 23000 --parallel 1
|
||||
```
|
||||
|
||||
@@ -1,118 +0,0 @@
|
||||
import argparse
|
||||
import asyncio
|
||||
import json
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
|
||||
from sglang.test.test_utils import add_common_other_args_and_parse, get_call_select
|
||||
from sglang.utils import download_and_cache_file, read_jsonl
|
||||
|
||||
|
||||
def get_one_example(lines, i, include_answer):
|
||||
ret = lines[i]["activity_label"] + ": " + lines[i]["ctx"] + " "
|
||||
if include_answer:
|
||||
ret += lines[i]["endings"][lines[i]["label"]]
|
||||
return ret
|
||||
|
||||
|
||||
def get_few_shot_examples(lines, k):
|
||||
ret = ""
|
||||
for i in range(k):
|
||||
ret += get_one_example(lines, i, True) + "\n\n"
|
||||
return ret
|
||||
|
||||
|
||||
def main(args):
|
||||
# Select backend
|
||||
call_select = get_call_select(args)
|
||||
|
||||
# Read data
|
||||
url = "https://raw.githubusercontent.com/rowanz/hellaswag/master/data/hellaswag_val.jsonl"
|
||||
filename = download_and_cache_file(url)
|
||||
lines = list(read_jsonl(filename))
|
||||
|
||||
# Construct prompts
|
||||
num_questions = args.num_questions
|
||||
num_shots = args.num_shots
|
||||
few_shot_examples = get_few_shot_examples(lines, num_shots)
|
||||
|
||||
questions = []
|
||||
choices = []
|
||||
labels = []
|
||||
for i in range(len(lines[:num_questions])):
|
||||
questions.append(get_one_example(lines, i, False))
|
||||
choices.append(lines[i]["endings"])
|
||||
labels.append(lines[i]["label"])
|
||||
|
||||
preds = [None] * len(labels)
|
||||
|
||||
# Run requests
|
||||
if args.backend != "lmql":
|
||||
# Use thread pool
|
||||
def get_one_answer(i):
|
||||
preds[i] = call_select(
|
||||
context=few_shot_examples + questions[i], choices=choices[i]
|
||||
)
|
||||
|
||||
tic = time.perf_counter()
|
||||
if args.parallel == 1:
|
||||
for i in tqdm(range(len(questions))):
|
||||
get_one_answer(i)
|
||||
else:
|
||||
with ThreadPoolExecutor(args.parallel) as executor:
|
||||
list(
|
||||
tqdm(
|
||||
executor.map(get_one_answer, list(range(len(questions)))),
|
||||
total=len(questions),
|
||||
)
|
||||
)
|
||||
else:
|
||||
# Use asyncio
|
||||
async def batched_call(batch_size):
|
||||
for i in range(0, len(questions), batch_size):
|
||||
tasks = []
|
||||
for q, c in zip(
|
||||
questions[i : i + batch_size], choices[i : i + batch_size]
|
||||
):
|
||||
tasks.append(call_select(context=few_shot_examples + q, choices=c))
|
||||
rets = await asyncio.gather(*tasks)
|
||||
for j in range(len(rets)):
|
||||
preds[i + j] = rets[j]
|
||||
|
||||
tic = time.perf_counter()
|
||||
asyncio.run(batched_call(batch_size=args.parallel))
|
||||
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
# Compute accuracy
|
||||
acc = np.mean(np.array(preds) == np.array(labels))
|
||||
print(f"Latency: {latency:.3f}")
|
||||
print(f"Accuracy: {acc:.3f}")
|
||||
|
||||
# Write results
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "hellaswag",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"accuracy": round(acc, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"other": {
|
||||
"num_questions": args.num_questions,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--num-shots", type=int, default=20)
|
||||
parser.add_argument("--data-path", type=str, default="hellaswag_val.jsonl")
|
||||
parser.add_argument("--num-questions", type=int, default=200)
|
||||
args = add_common_other_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,60 +0,0 @@
|
||||
## Run benchmark
|
||||
|
||||
### Build dataset
|
||||
```
|
||||
pip install wikipedia
|
||||
python3 build_dataset.py
|
||||
```
|
||||
|
||||
### Dependencies
|
||||
|
||||
```
|
||||
llama_cpp_python 0.2.19
|
||||
guidance 0.1.10
|
||||
vllm 0.2.5
|
||||
outlines 0.0.22
|
||||
```
|
||||
|
||||
### Benchmark sglang
|
||||
|
||||
Run Llama-7B
|
||||
|
||||
```
|
||||
python3 -m sglang.launch_server --model-path meta-llama/Llama-2-7b-chat-hf --port 30000
|
||||
```
|
||||
|
||||
Run Mixtral-8x7B
|
||||
|
||||
```
|
||||
python3 -m sglang.launch_server --model-path mistralai/Mixtral-8x7B-Instruct-v0.1 --port 30000 --tp-size 8
|
||||
```
|
||||
|
||||
Benchmark
|
||||
|
||||
```
|
||||
python3 bench_sglang.py --num-questions 10
|
||||
```
|
||||
|
||||
|
||||
### Benchmark Outlines + vLLM
|
||||
|
||||
Run Llama-7B
|
||||
|
||||
```
|
||||
python3 -m outlines.serve.serve --tokenizer-mode auto --model meta-llama/Llama-2-7b-chat-hf --disable-log-requests --port 21000
|
||||
```
|
||||
|
||||
Benchmark
|
||||
|
||||
```
|
||||
python3 bench_other.py --backend outlines --num-questions 10
|
||||
```
|
||||
|
||||
|
||||
### Benchmark guidance
|
||||
|
||||
Run Llama-7B and benchmark
|
||||
|
||||
```
|
||||
python3 bench_other.py --backend guidance --num-questions 10 --parallel 1 --n-ctx 4096 --model-path path/to/gguf
|
||||
```
|
||||
@@ -1,98 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from functools import partial
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
from sglang.lang.ir import REGEX_FLOAT, REGEX_INT, REGEX_STR
|
||||
from sglang.test.test_utils import add_common_other_args_and_parse, get_call_generate
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
REGEX_LIST = r"\[(" + REGEX_STR + ", )*" + REGEX_STR + r"\]"
|
||||
|
||||
|
||||
# fmt: off
|
||||
def json_decode(document, generate):
|
||||
s = "Please extract the information of a city from the following wikipedia page.\n"
|
||||
s += "Page begin.\n" + document + "Page end.\n"
|
||||
s += "Here is the name, country, and symbol of the city in JSON format.\n"
|
||||
s += "{\n"
|
||||
s += ' "name": '
|
||||
s += generate(s, max_tokens=8, regex=REGEX_STR + ",") + "\n"
|
||||
s += ' "country": '
|
||||
s += generate(s, max_tokens=8, regex=REGEX_STR + ",") + "\n"
|
||||
s += ' "latitude": '
|
||||
s += generate(s, max_tokens=8, regex=REGEX_FLOAT + ",") + "\n"
|
||||
s += ' "population": '
|
||||
s += generate(s, max_tokens=8, regex=REGEX_INT + ",") + "\n"
|
||||
s += ' "top 3 landmarks": '
|
||||
s += generate(s, max_tokens=24, regex=REGEX_LIST) + "\n"
|
||||
s += "}\n"
|
||||
|
||||
return s
|
||||
# fmt: on
|
||||
|
||||
|
||||
def main(args):
|
||||
lines = read_jsonl(args.data_path)
|
||||
arguments = []
|
||||
for i in range(len(lines[: args.num_questions])):
|
||||
arguments.append(
|
||||
{
|
||||
"document": lines[i]["document"],
|
||||
}
|
||||
)
|
||||
states = [None] * len(arguments)
|
||||
|
||||
# Select backend
|
||||
call_generate = partial(get_call_generate(args), temperature=0)
|
||||
|
||||
# Run requests
|
||||
def get_one_answer(i):
|
||||
states[i] = json_decode(generate=call_generate, **arguments[i])
|
||||
|
||||
tic = time.perf_counter()
|
||||
if args.parallel == 1:
|
||||
for i in tqdm(range(len(arguments))):
|
||||
get_one_answer(i)
|
||||
else:
|
||||
with ThreadPoolExecutor(args.parallel) as executor:
|
||||
rets = list(
|
||||
tqdm(
|
||||
executor.map(get_one_answer, list(range(len(arguments)))),
|
||||
total=len(arguments),
|
||||
)
|
||||
)
|
||||
for _ in rets:
|
||||
pass
|
||||
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
# Compute accuracy
|
||||
print(f"Latency: {latency:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "json_decode_regex",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"other": {
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="questions.jsonl")
|
||||
parser.add_argument("--num-questions", type=int, default=20)
|
||||
args = add_common_other_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,101 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
|
||||
import sglang as sgl
|
||||
from sglang.lang.ir import REGEX_FLOAT, REGEX_INT, REGEX_STR
|
||||
from sglang.test.test_utils import (
|
||||
add_common_sglang_args_and_parse,
|
||||
select_sglang_backend,
|
||||
)
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
REGEX_LIST = r"\[(" + REGEX_STR + ", )*" + REGEX_STR + r"\]"
|
||||
|
||||
# fmt: off
|
||||
@sgl.function
|
||||
def json_warm_up(s):
|
||||
s += "The information about Hogwarts is in the following JSON format.\n"
|
||||
with s.var_scope("json_output"):
|
||||
s += "{\n"
|
||||
s += ' "name": ' + sgl.gen("name", max_tokens=8, regex=REGEX_STR + ",") + "\n"
|
||||
s += ' "country": ' + sgl.gen("country", max_tokens=8, regex=REGEX_STR + ",") + "\n"
|
||||
s += ' "latitude": ' + sgl.gen("latitude", max_tokens=8, regex=REGEX_FLOAT + ",") + "\n"
|
||||
s += ' "population": ' + sgl.gen("population", max_tokens=8, regex=REGEX_INT + ",") + "\n"
|
||||
s += ' "top 3 landmarks": ' + sgl.gen( "landmarks", max_tokens=24, regex=REGEX_LIST) + "\n"
|
||||
s += "}\n"
|
||||
print(f'The warmp up json result is:\n{s["json_output"]}')
|
||||
# fmt: on
|
||||
|
||||
# fmt: off
|
||||
@sgl.function
|
||||
def json_decode(s, document):
|
||||
s += "Please extract the information of a city from the following wikipedia page.\n"
|
||||
s += "Page begin.\n" + document + "Page end.\n"
|
||||
s += "Here is the name, country, and symbol of the city in JSON format.\n"
|
||||
with s.var_scope("json_output"):
|
||||
s += "{\n"
|
||||
s += ' "name": ' + sgl.gen("name", max_tokens=8, regex=REGEX_STR + ",") + "\n"
|
||||
s += ' "country": ' + sgl.gen("country", max_tokens=8, regex=REGEX_STR + ",") + "\n"
|
||||
s += ' "latitude": ' + sgl.gen("latitude", max_tokens=8, regex=REGEX_FLOAT + ",") + "\n"
|
||||
s += ' "population": ' + sgl.gen("population", max_tokens=8, regex=REGEX_INT + ",") + "\n"
|
||||
s += ' "top 3 landmarks": ' + sgl.gen( "landmarks", max_tokens=24, regex=REGEX_LIST) + "\n"
|
||||
s += "}\n"
|
||||
# fmt: on
|
||||
|
||||
|
||||
def main(args):
|
||||
lines = read_jsonl(args.data_path)
|
||||
lines = list(lines)
|
||||
arguments = []
|
||||
for i in range(len(lines[: args.num_questions])):
|
||||
arguments.append(
|
||||
{
|
||||
"document": lines[i]["document"],
|
||||
}
|
||||
)
|
||||
|
||||
# Select backend
|
||||
backend = select_sglang_backend(args)
|
||||
sgl.set_default_backend(backend)
|
||||
|
||||
# Warm up
|
||||
json_warm_up.run().sync()
|
||||
|
||||
# Run requests
|
||||
tic = time.perf_counter()
|
||||
states = json_decode.run_batch(
|
||||
arguments, temperature=0, num_threads=args.parallel, progress_bar=True
|
||||
)
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
# Compute accuracy
|
||||
print(f"Latency: {latency:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(f"tmp_{args.backend}_json_results.txt", "w") as fout:
|
||||
for state in states:
|
||||
fout.write(state["json_output"] + "\n")
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "json_decode_regex",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"other": {
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="questions.jsonl")
|
||||
parser.add_argument("--num-questions", type=int, default=20)
|
||||
args = add_common_sglang_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,58 +0,0 @@
|
||||
import json
|
||||
|
||||
import transformers
|
||||
import wikipedia
|
||||
|
||||
model_path = "meta-llama/Llama-2-7b-chat-hf"
|
||||
t = transformers.AutoTokenizer.from_pretrained(model_path)
|
||||
city_names = [
|
||||
"los angles",
|
||||
"london",
|
||||
"tokyo",
|
||||
"beijing",
|
||||
"singapore",
|
||||
"paris",
|
||||
"dubai",
|
||||
"sydney",
|
||||
"moscow",
|
||||
"rome",
|
||||
"toronto",
|
||||
"rio de janeiro",
|
||||
"istanbul",
|
||||
"berlin",
|
||||
"auckland",
|
||||
"buenos aires",
|
||||
"mexico city",
|
||||
"mumbai",
|
||||
"seoul",
|
||||
"bangkok",
|
||||
"cairo",
|
||||
"athens",
|
||||
"jerusalem",
|
||||
]
|
||||
|
||||
|
||||
def get_content(city_name):
|
||||
content = str(wikipedia.page(city_name).content)
|
||||
content = content.replace("\n\n", "\n")
|
||||
|
||||
tokens = t.encode(content)
|
||||
|
||||
expected_tokens = 3000
|
||||
truncate_len = int((expected_tokens / len(tokens)) * len(content))
|
||||
truncate_content = content[:truncate_len]
|
||||
truncate_tokens = t.encode(truncate_content)
|
||||
|
||||
# Count token
|
||||
print(
|
||||
f"city_name: {city_name}, #tokens: {len(tokens)}, #truncate tokens: {len(truncate_tokens)}"
|
||||
)
|
||||
|
||||
return truncate_content
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
with open("questions.jsonl", "w") as fout:
|
||||
for city_name in city_names:
|
||||
truncate_content = get_content(city_name)
|
||||
fout.write(json.dumps({"document": truncate_content}) + "\n")
|
||||
@@ -1,88 +0,0 @@
|
||||
## Run benchmark
|
||||
|
||||
### Dependencies
|
||||
|
||||
```
|
||||
llama_cpp_python 0.2.38
|
||||
guidance 0.1.10
|
||||
vllm 0.2.7
|
||||
outlines 0.0.25
|
||||
```
|
||||
|
||||
### Build dataset
|
||||
|
||||
When benchmarking long document information retrieval, run the following command to build the dataset:
|
||||
|
||||
```bash
|
||||
pip install wikipedia
|
||||
python3 build_dataset.py
|
||||
```
|
||||
|
||||
### Benchmark sglang
|
||||
|
||||
Run Llama-7B
|
||||
|
||||
```bash
|
||||
python3 -m sglang.launch_server --model-path meta-llama/Llama-2-7b-chat-hf --port 30000
|
||||
```
|
||||
|
||||
Benchmark Character Generation
|
||||
|
||||
```bash
|
||||
python3 bench_sglang.py --mode character
|
||||
```
|
||||
|
||||
Benchmark City Information Retrieval
|
||||
|
||||
```bash
|
||||
python3 bench_sglang.py --mode city
|
||||
```
|
||||
|
||||
|
||||
### Benchmark Outlines + vLLM
|
||||
|
||||
Run Llama-7B
|
||||
|
||||
```bash
|
||||
python3 -m outlines.serve.serve --tokenizer-mode auto --model meta-llama/Llama-2-7b-chat-hf --disable-log-requests --port 21000
|
||||
```
|
||||
|
||||
Benchmark Character Generation
|
||||
|
||||
```bash
|
||||
python3 bench_other.py --mode character --backend outlines
|
||||
```
|
||||
|
||||
Benchmark City Information Retrieval
|
||||
|
||||
```bash
|
||||
python3 bench_other.py --mode city --backend outlines
|
||||
```
|
||||
|
||||
### Benchmark guidance
|
||||
|
||||
Run Llama-7B and benchmark character generation
|
||||
|
||||
```bash
|
||||
python3 bench_other.py --mode character --backend guidance --parallel 1 --n-ctx 4096 --model-path path/to/gguf
|
||||
```
|
||||
|
||||
Run Llama-7B and benchmark city information retrieval
|
||||
|
||||
```bash
|
||||
python3 bench_other.py --mode city --backend guidance --parallel 1 --n-ctx 4096 --model-path path/to/gguf
|
||||
```
|
||||
|
||||
### Benchmark lmql
|
||||
|
||||
Run Llama-7B and benchmark character generation
|
||||
|
||||
```
|
||||
python3 bench_other.py --mode character --backend lmql --parallel 1
|
||||
```
|
||||
|
||||
Run Llama-7B and benchmark city information retrieval
|
||||
|
||||
```
|
||||
python3 bench_other.py --mode city --backend lmql --parallel 1
|
||||
```
|
||||
@@ -1,288 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from functools import partial
|
||||
|
||||
import guidance
|
||||
from tqdm import tqdm
|
||||
|
||||
from sglang.test.test_utils import add_common_other_args_and_parse, get_call_generate
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
# there are some FSM bugs with json regex converted from pydantic model
|
||||
# here use a string regex instead
|
||||
# regex_string = build_regex_from_object(HarryPoterRole)
|
||||
character_regex = (
|
||||
r"""\{\n"""
|
||||
+ r""" "name": "[\w\d\s]{1,16}",\n"""
|
||||
+ r""" "house": "(Gryffindor|Slytherin|Ravenclaw|Hufflepuff)",\n"""
|
||||
+ r""" "blood status": "(Pure-blood|Half-blood|Muggle-born)",\n"""
|
||||
+ r""" "occupation": "(student|teacher|auror|ministry of magic|death eater|order of the phoenix)",\n"""
|
||||
+ r""" "wand": \{\n"""
|
||||
+ r""" "wood": "[\w\d\s]{1,16}",\n"""
|
||||
+ r""" "core": "[\w\d\s]{1,16}",\n"""
|
||||
+ r""" "length": [0-9]{1,2}\.[0-9]{0,2}\n"""
|
||||
+ r""" \},\n"""
|
||||
+ r""" "alive": "(Alive|Deceased)",\n"""
|
||||
+ r""" "patronus": "[\w\d\s]{1,16}",\n"""
|
||||
+ r""" "bogart": "[\w\d\s]{1,16}"\n"""
|
||||
+ r"""\}"""
|
||||
)
|
||||
|
||||
city_regex = (
|
||||
r"""\{\n"""
|
||||
+ r""" "name": "[\w\d\s]{1,16}",\n"""
|
||||
+ r""" "country": "[\w\d\s]{1,16}",\n"""
|
||||
+ r""" "latitude": [-+]?[0-9]*\.?[0-9]{0,2},\n"""
|
||||
+ r""" "population": [-+]?[0-9]{1,9},\n"""
|
||||
+ r""" "top 3 landmarks": \["[\w\d\s]{1,16}", "[\w\d\s]{1,16}", "[\w\d\s]{1,16}"\]\n"""
|
||||
+ r"""\}"""
|
||||
)
|
||||
|
||||
# fmt: off
|
||||
def character_gen(name, generate):
|
||||
s = name + " is a character in Harry Potter. Please fill in the following information about this character.\n"
|
||||
s += generate(s, max_tokens=256, regex=character_regex)
|
||||
return s
|
||||
# fmt: on
|
||||
|
||||
# fmt: off
|
||||
def city_gen(document, generate):
|
||||
s = "Please extract the information of a city from the following wikipedia page.\n"
|
||||
s += "Page begin.\n" + document + "Page end.\n"
|
||||
s += "Here is the name, country, and symbol of the city in JSON format.\n"
|
||||
s += generate(s, max_tokens=256, regex=city_regex)
|
||||
return s
|
||||
# fmt: on
|
||||
|
||||
|
||||
@guidance
|
||||
def character_maker(lm, name):
|
||||
regex_str_no_quote = r"[\w\d\s]+"
|
||||
regex_float = r"[0-9]+\.[0-9]+"
|
||||
lm += f"""\
|
||||
{name} is a character in Harry Potter. Please fill in the following information about this character.
|
||||
{{
|
||||
"name": "{guidance.gen("name", max_tokens=16, regex=regex_str_no_quote)}",
|
||||
"house": "{guidance.select(options=['Gryffindor', 'Slytherin', 'Ravenclaw', 'Hufflepuff'], name='house')}",
|
||||
"blood status": "{guidance.select(options=['Pure-blood', 'Half-blood', 'Muggle-born'], name='blood status')}",
|
||||
"occupation": "{guidance.select(options=['student', 'teacher', 'auror', 'ministry of magic', 'death eater', 'order of the phoenix'], name='occupation')}",
|
||||
"wand": {{
|
||||
"wood": "{guidance.gen("wood", max_tokens=16, regex=regex_str_no_quote)}",
|
||||
"core": "{guidance.gen('core', max_tokens=16, regex=regex_str_no_quote)}",
|
||||
"length": {guidance.gen('length', max_tokens=10, regex=regex_float)}
|
||||
}},
|
||||
"alive": "{guidance.select(options=['Alive', 'Deceased'], name='alive')}",
|
||||
"patronus": "{guidance.gen('patronus', max_tokens=16, regex=regex_str_no_quote)}",
|
||||
"bogart": "{guidance.gen('bogart', max_tokens=16, regex=regex_str_no_quote)}"
|
||||
}}
|
||||
"""
|
||||
|
||||
return lm
|
||||
|
||||
|
||||
async def call_generate_lmql(
|
||||
prompt, temperature, max_tokens, regex, max_len=4096, model=None, **kwargs
|
||||
):
|
||||
assert model is not None
|
||||
import lmql
|
||||
|
||||
@lmql.query(model=model)
|
||||
async def program(question, max_tokens, regex):
|
||||
'''lmql
|
||||
"""{question}[ANSWER]""" where len(TOKENS(ANSWER)) < max_tokens and REGEX(ANSWER, regex)
|
||||
return ANSWER
|
||||
'''
|
||||
|
||||
return await program(
|
||||
question=prompt,
|
||||
temperature=temperature,
|
||||
max_tokens=max_tokens,
|
||||
max_len=max_len,
|
||||
regex=regex,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
@guidance
|
||||
def city_maker(lm, document):
|
||||
regex_str_no_quote = r"[\w\d\s]+"
|
||||
regex_float = r"[0-9]+\.[0-9]+"
|
||||
lm += f"""\
|
||||
Please extract the information of a city from the following wikipedia page.
|
||||
Page begin.
|
||||
{document}
|
||||
Page end.
|
||||
Here is the name, country, and symbol of the city in JSON format.
|
||||
{{
|
||||
"name": "{guidance.gen("name", max_tokens=16, regex=regex_str_no_quote)}",
|
||||
"country": "{guidance.gen("country", max_tokens=16, regex=regex_str_no_quote)}",
|
||||
"latitude": {guidance.gen("latitude", max_tokens=10, regex=regex_float)},
|
||||
"population": {guidance.gen("population", max_tokens=10, regex=r"[0-9]+")},
|
||||
"top 3 landmarks": [
|
||||
"{guidance.gen("landmark1", max_tokens=16, regex=regex_str_no_quote)}", "{guidance.gen("landmark2", max_tokens=16, regex=regex_str_no_quote)}", "{guidance.gen("landmark3", max_tokens=16, regex=regex_str_no_quote)}"
|
||||
]
|
||||
}}
|
||||
"""
|
||||
|
||||
return lm
|
||||
|
||||
|
||||
def bench_character(args):
|
||||
arguments = []
|
||||
with open(args.data_path, "r") as f:
|
||||
for line in f:
|
||||
arguments.append({"name": line.strip()})
|
||||
arguments = arguments[: args.num_jsons]
|
||||
|
||||
states = [None] * len(arguments)
|
||||
|
||||
# Select backend
|
||||
if args.backend == "outlines":
|
||||
call_generate = partial(get_call_generate(args), temperature=0)
|
||||
|
||||
def get_one_answer(i):
|
||||
states[i] = character_gen(**arguments[i], generate=call_generate)
|
||||
|
||||
elif args.backend == "guidance":
|
||||
model = guidance.models.LlamaCpp(
|
||||
args.model_path,
|
||||
n_gpu_layers=-1,
|
||||
n_ctx=args.n_ctx,
|
||||
)
|
||||
|
||||
def get_one_answer(i):
|
||||
lm = model + character_maker(**arguments[i])
|
||||
states[i] = lm
|
||||
|
||||
elif args.backend == "lmql":
|
||||
import asyncio
|
||||
|
||||
import lmql
|
||||
|
||||
model = lmql.model(args.model_path, endpoint=f"{args.host}:{args.port}")
|
||||
call_generate = partial(
|
||||
call_generate_lmql,
|
||||
model=model,
|
||||
max_tokens=256,
|
||||
regex=character_regex,
|
||||
)
|
||||
|
||||
async def get_one_answer_async(i):
|
||||
states[i] = await call_generate(prompt=arguments[i]["name"], temperature=0)
|
||||
|
||||
else:
|
||||
raise ValueError(f"Invalid backend: {args.backend}")
|
||||
|
||||
tic = time.perf_counter()
|
||||
|
||||
if args.backend != "lmql":
|
||||
if args.parallel == 1:
|
||||
for i in tqdm(range(len(arguments))):
|
||||
get_one_answer(i)
|
||||
else:
|
||||
with ThreadPoolExecutor(args.parallel) as executor:
|
||||
rets = list(
|
||||
tqdm(
|
||||
executor.map(get_one_answer, list(range(len(arguments)))),
|
||||
total=len(arguments),
|
||||
)
|
||||
)
|
||||
for _ in rets:
|
||||
pass
|
||||
else:
|
||||
batches = []
|
||||
for i in range(0, len(arguments), args.parallel):
|
||||
batches.append(list(range(i, min(i + args.parallel, len(arguments)))))
|
||||
loop = asyncio.get_event_loop()
|
||||
|
||||
for bt in tqdm(batches):
|
||||
loop.run_until_complete(
|
||||
asyncio.gather(*[get_one_answer_async(i) for i in bt])
|
||||
)
|
||||
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
return states, latency
|
||||
|
||||
|
||||
def bench_city_doc(args):
|
||||
arguments = []
|
||||
for line in read_jsonl(args.data_path):
|
||||
arguments.append({"document": line["document"]})
|
||||
arguments = arguments[: args.num_jsons]
|
||||
|
||||
states = [None] * len(arguments)
|
||||
|
||||
# Select backend
|
||||
if args.backend == "outlines":
|
||||
call_generate = partial(get_call_generate(args), temperature=0)
|
||||
|
||||
def get_one_answer(i):
|
||||
states[i] = city_gen(**arguments[i], generate=call_generate)
|
||||
|
||||
elif args.backend == "guidance":
|
||||
model = guidance.models.LlamaCpp(
|
||||
args.model_path,
|
||||
n_gpu_layers=-1,
|
||||
n_ctx=args.n_ctx,
|
||||
)
|
||||
|
||||
def get_one_answer(i):
|
||||
lm = model + city_maker(**arguments[i])
|
||||
states[i] = lm
|
||||
|
||||
else:
|
||||
raise ValueError(f"Invalid backend: {args.backend}")
|
||||
|
||||
tic = time.perf_counter()
|
||||
if args.parallel == 1:
|
||||
for i in tqdm(range(len(arguments))):
|
||||
get_one_answer(i)
|
||||
else:
|
||||
with ThreadPoolExecutor(args.parallel) as executor:
|
||||
rets = executor.map(get_one_answer, list(range(len(arguments))))
|
||||
for _ in rets:
|
||||
pass
|
||||
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
return states, latency
|
||||
|
||||
|
||||
def main(args):
|
||||
if args.mode == "character":
|
||||
args.data_path = "dataset.txt"
|
||||
states, latency = bench_character(args)
|
||||
elif args.mode == "city":
|
||||
args.data_path = "questions.jsonl"
|
||||
states, latency = bench_city_doc(args)
|
||||
|
||||
# Compute accuracy
|
||||
print(f"Latency: {latency:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}_{args.mode}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "json_jump_forward",
|
||||
"backend": args.backend,
|
||||
"latency": round(latency, 3),
|
||||
"num_jsons": args.num_jsons,
|
||||
"mode": args.mode,
|
||||
"parallel": args.parallel,
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str)
|
||||
parser.add_argument("--num-jsons", type=int, default=50)
|
||||
parser.add_argument(
|
||||
"--mode", type=str, default="character", choices=["character", "city"]
|
||||
)
|
||||
args = add_common_other_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,143 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
|
||||
import sglang as sgl
|
||||
from sglang.test.test_utils import (
|
||||
add_common_sglang_args_and_parse,
|
||||
select_sglang_backend,
|
||||
)
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
# there are some FSM bugs with json regex converted from pydantic model
|
||||
# here use a string regex instead
|
||||
# regex_string = build_regex_from_object(HarryPoterRole)
|
||||
character_regex = (
|
||||
r"""\{\n"""
|
||||
+ r""" "name": "[\w\d\s]{1,16}",\n"""
|
||||
+ r""" "house": "(Gryffindor|Slytherin|Ravenclaw|Hufflepuff)",\n"""
|
||||
+ r""" "blood status": "(Pure-blood|Half-blood|Muggle-born)",\n"""
|
||||
+ r""" "occupation": "(student|teacher|auror|ministry of magic|death eater|order of the phoenix)",\n"""
|
||||
+ r""" "wand": \{\n"""
|
||||
+ r""" "wood": "[\w\d\s]{1,16}",\n"""
|
||||
+ r""" "core": "[\w\d\s]{1,16}",\n"""
|
||||
+ r""" "length": [0-9]{1,2}\.[0-9]{0,2}\n"""
|
||||
+ r""" \},\n"""
|
||||
+ r""" "alive": "(Alive|Deceased)",\n"""
|
||||
+ r""" "patronus": "[\w\d\s]{1,16}",\n"""
|
||||
+ r""" "bogart": "[\w\d\s]{1,16}"\n"""
|
||||
+ r"""\}"""
|
||||
)
|
||||
|
||||
city_regex = (
|
||||
r"""\{\n"""
|
||||
+ r""" "name": "[\w\d\s]{1,16}",\n"""
|
||||
+ r""" "country": "[\w\d\s]{1,16}",\n"""
|
||||
+ r""" "latitude": [-+]?[0-9]*\.?[0-9]{0,2},\n"""
|
||||
+ r""" "population": [-+]?[0-9]{1,9},\n"""
|
||||
+ r""" "top 3 landmarks": \["[\w\d\s]{1,16}", "[\w\d\s]{1,16}", "[\w\d\s]{1,16}"\]\n"""
|
||||
+ r"""\}"""
|
||||
)
|
||||
|
||||
# fmt: off
|
||||
@sgl.function
|
||||
def character_gen(s, name):
|
||||
s += name + " is a character in Harry Potter. Please fill in the following information about this character.\n"
|
||||
s += sgl.gen("json_output", max_tokens=256, regex=character_regex)
|
||||
# fmt: on
|
||||
|
||||
# fmt: off
|
||||
@sgl.function
|
||||
def city_gen(s, document):
|
||||
s += "Please extract the information of a city from the following wikipedia page.\n"
|
||||
s += "Page begin.\n" + document + "Page end.\n"
|
||||
s += "Here is the name, country, and symbol of the city in JSON format.\n"
|
||||
s += sgl.gen("json_output",max_tokens=256, regex=city_regex)
|
||||
# fmt: on
|
||||
|
||||
|
||||
def bench_city_doc(args):
|
||||
arguments = []
|
||||
for line in read_jsonl(args.data_path):
|
||||
arguments.append({"document": line["document"]})
|
||||
arguments = arguments[: args.num_jsons]
|
||||
|
||||
# Select backend
|
||||
backend = select_sglang_backend(args)
|
||||
sgl.set_default_backend(backend)
|
||||
|
||||
# Run requests
|
||||
tic = time.perf_counter()
|
||||
states = city_gen.run_batch(
|
||||
arguments,
|
||||
temperature=0,
|
||||
num_threads=args.parallel,
|
||||
progress_bar=True,
|
||||
)
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
return states, latency
|
||||
|
||||
|
||||
def bench_character(args):
|
||||
arguments = []
|
||||
with open(args.data_path, "r") as f:
|
||||
for line in f:
|
||||
arguments.append({"name": line.strip()})
|
||||
arguments = arguments[: args.num_jsons]
|
||||
|
||||
# Select backend
|
||||
backend = select_sglang_backend(args)
|
||||
sgl.set_default_backend(backend)
|
||||
|
||||
# Run requests
|
||||
tic = time.perf_counter()
|
||||
states = character_gen.run_batch(
|
||||
arguments,
|
||||
temperature=0,
|
||||
num_threads=args.parallel,
|
||||
progress_bar=True,
|
||||
)
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
return states, latency
|
||||
|
||||
|
||||
def main(args):
|
||||
if args.mode == "character":
|
||||
args.data_path = "dataset.txt"
|
||||
states, latency = bench_character(args)
|
||||
elif args.mode == "city":
|
||||
args.data_path = "questions.jsonl"
|
||||
states, latency = bench_city_doc(args)
|
||||
|
||||
# Compute accuracy
|
||||
print(f"Latency: {latency:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}_{args.mode}.txt", states)
|
||||
with open(f"{args.backend}_{args.mode}.json", "w") as fout:
|
||||
for state in states:
|
||||
fout.write(state["json_output"] + "\n")
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "json_jump_forward",
|
||||
"backend": args.backend,
|
||||
"latency": round(latency, 3),
|
||||
"num_jsons": args.num_jsons,
|
||||
"mode": args.mode,
|
||||
"parallel": args.parallel,
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str)
|
||||
parser.add_argument("--num-jsons", type=int, default=50)
|
||||
parser.add_argument(
|
||||
"--mode", type=str, default="character", choices=["character", "city"]
|
||||
)
|
||||
args = add_common_sglang_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,58 +0,0 @@
|
||||
import json
|
||||
|
||||
import transformers
|
||||
import wikipedia
|
||||
|
||||
model_path = "meta-llama/Llama-2-7b-chat-hf"
|
||||
t = transformers.AutoTokenizer.from_pretrained(model_path)
|
||||
city_names = [
|
||||
"los angles",
|
||||
"london",
|
||||
"tokyo",
|
||||
"beijing",
|
||||
"singapore",
|
||||
"paris",
|
||||
"dubai",
|
||||
"sydney",
|
||||
"moscow",
|
||||
"rome",
|
||||
"toronto",
|
||||
"rio de janeiro",
|
||||
"istanbul",
|
||||
"berlin",
|
||||
"auckland",
|
||||
"buenos aires",
|
||||
"mexico city",
|
||||
"mumbai",
|
||||
"seoul",
|
||||
"bangkok",
|
||||
"cairo",
|
||||
"athens",
|
||||
"jerusalem",
|
||||
]
|
||||
|
||||
|
||||
def get_content(city_name):
|
||||
content = str(wikipedia.page(city_name).content)
|
||||
content = content.replace("\n\n", "\n")
|
||||
|
||||
tokens = t.encode(content)
|
||||
|
||||
expected_tokens = 3000
|
||||
truncate_len = int((expected_tokens / len(tokens)) * len(content))
|
||||
truncate_content = content[:truncate_len]
|
||||
truncate_tokens = t.encode(truncate_content)
|
||||
|
||||
# Count token
|
||||
print(
|
||||
f"city_name: {city_name}, #tokens: {len(tokens)}, #truncate tokens: {len(truncate_tokens)}"
|
||||
)
|
||||
|
||||
return truncate_content
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
with open("questions.jsonl", "w") as fout:
|
||||
for city_name in city_names:
|
||||
truncate_content = get_content(city_name)
|
||||
fout.write(json.dumps({"document": truncate_content}) + "\n")
|
||||
@@ -1,50 +0,0 @@
|
||||
Harry Potter
|
||||
Hermione Granger
|
||||
Ron Weasley
|
||||
Albus Dumbledore
|
||||
Severus Snape
|
||||
Rubeus Hagrid
|
||||
Draco Malfoy
|
||||
Ginny Weasley
|
||||
Fred Weasley
|
||||
George Weasley
|
||||
Percy Weasley
|
||||
Sirius Black
|
||||
Remus Lupin
|
||||
Neville Longbottom
|
||||
Luna Lovegood
|
||||
Cedric Diggory
|
||||
Cho Chang
|
||||
Lord Voldemort
|
||||
Minerva McGonagall
|
||||
Filius Flitwick
|
||||
Dolores Umbridge
|
||||
Bellatrix Lestrange
|
||||
Lucius Malfoy
|
||||
Molly Weasley
|
||||
Arthur Weasley
|
||||
Nymphadora Tonks
|
||||
Dobby
|
||||
Moaning Myrtle
|
||||
Peter Pettigrew
|
||||
Alastor 'Mad-Eye' Moody
|
||||
Horace Slughorn
|
||||
Vernon Dursley
|
||||
Petunia Dursley
|
||||
Dudley Dursley
|
||||
Argus Filch
|
||||
Sybill Trelawney
|
||||
Gilderoy Lockhart
|
||||
Fleur Delacour
|
||||
Viktor Krum
|
||||
Bill Weasley
|
||||
Oliver Wood
|
||||
Cornelius Fudge
|
||||
Barty Crouch Sr.
|
||||
Barty Crouch Jr.
|
||||
Kingsley Shacklebolt
|
||||
Quirinus Quirrell
|
||||
Nearly Headless Nick
|
||||
Aunt Marge
|
||||
Griphook
|
||||
Ludo Bagman
|
||||
@@ -1,15 +0,0 @@
|
||||
## Run benchmark
|
||||
|
||||
### Benchmark sglang
|
||||
|
||||
Run Llama-8b
|
||||
|
||||
```bash
|
||||
python3 -m sglang.launch_server --model-path meta-llama/Llama-3.1-8B-Instruct --port 30000
|
||||
```
|
||||
|
||||
Benchmark
|
||||
|
||||
```bash
|
||||
python3 bench_sglang.py
|
||||
```
|
||||
@@ -1,146 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
from typing import List, Tuple
|
||||
|
||||
import jsonschema
|
||||
from datasets import load_dataset
|
||||
|
||||
import sglang as sgl
|
||||
from sglang.global_config import global_config
|
||||
from sglang.srt.utils.hf_transformers_utils import get_tokenizer
|
||||
from sglang.test.test_utils import (
|
||||
add_common_sglang_args_and_parse,
|
||||
select_sglang_backend,
|
||||
)
|
||||
from sglang.utils import dump_state_text
|
||||
|
||||
|
||||
@sgl.function
|
||||
def schema_gen(s, message: Tuple[str, str], json_schema: str):
|
||||
system, user = message
|
||||
s += sgl.system(system)
|
||||
s += sgl.user(user)
|
||||
s += sgl.assistant(
|
||||
sgl.gen("json_output", temperature=0, max_tokens=256, json_schema=json_schema)
|
||||
)
|
||||
|
||||
|
||||
def contains_formats(schema, formats: List[str]):
|
||||
if isinstance(schema, dict):
|
||||
if schema.get("format", None) in formats:
|
||||
return True
|
||||
for value in schema.values():
|
||||
if contains_formats(value, formats):
|
||||
return True
|
||||
elif isinstance(schema, list):
|
||||
for item in schema:
|
||||
if contains_formats(item, formats):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def convert_dataset(path: str):
|
||||
raw_dataset = load_dataset(path)
|
||||
dataset = []
|
||||
for data in raw_dataset["train"]:
|
||||
messages = data["prompt"]
|
||||
schema = data["schema"]
|
||||
obj = json.loads(schema)
|
||||
|
||||
# skip some corrupted examples
|
||||
if obj.get("type", None) is None:
|
||||
continue
|
||||
|
||||
# skip schema with format "email"
|
||||
# which is not supported by outlines for now
|
||||
if contains_formats(obj, ["email"]):
|
||||
continue
|
||||
|
||||
system = messages[0]
|
||||
user = messages[1]
|
||||
assert system["role"] == "system", "invalid role"
|
||||
assert user["role"] == "user", "invalid role"
|
||||
assert len(messages) == 2, "invalid message length"
|
||||
message = json.dumps(system["content"]), json.dumps(user["content"])
|
||||
dataset.append(
|
||||
{
|
||||
"message": message,
|
||||
"json_schema": schema,
|
||||
}
|
||||
)
|
||||
|
||||
return dataset
|
||||
|
||||
|
||||
def bench_schema(args):
|
||||
arguments = convert_dataset(args.data_path)
|
||||
|
||||
if args.num_jsons < 0 or args.num_jsons > len(arguments):
|
||||
args.num_jsons = len(arguments)
|
||||
arguments = arguments[: args.num_jsons]
|
||||
|
||||
# Select backend
|
||||
backend = select_sglang_backend(args)
|
||||
sgl.set_default_backend(backend)
|
||||
|
||||
# Run requests
|
||||
tic = time.perf_counter()
|
||||
states = schema_gen.run_batch(
|
||||
arguments,
|
||||
temperature=0,
|
||||
num_threads=args.parallel,
|
||||
progress_bar=True,
|
||||
)
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
# Check if the outputs are valid
|
||||
indexes = []
|
||||
for i, state in enumerate(states):
|
||||
try:
|
||||
schema = json.loads(arguments[i]["json_schema"])
|
||||
obj = json.loads(state["json_output"])
|
||||
assert jsonschema.validate(obj, schema) is None
|
||||
except Exception as e:
|
||||
print(e)
|
||||
indexes.append(i)
|
||||
|
||||
return states, latency
|
||||
|
||||
|
||||
def main(args):
|
||||
states, latency = bench_schema(args)
|
||||
|
||||
# Compute accuracy
|
||||
tokenizer = get_tokenizer(
|
||||
global_config.default_backend.get_server_info()["tokenizer_path"]
|
||||
)
|
||||
output_jsons = [state["json_output"] for state in states]
|
||||
num_output_tokens = sum(len(tokenizer.encode(x)) for x in output_jsons)
|
||||
print(f"Latency: {latency:.3f}")
|
||||
print(f"Output throughput: {num_output_tokens / latency:.3f} token/s")
|
||||
print(f"#output tokens: {num_output_tokens}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
with open(f"{args.backend}.jsonl", "w") as fout:
|
||||
for state in states:
|
||||
fout.write(state["json_output"] + "\n")
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "json_schema",
|
||||
"backend": args.backend,
|
||||
"latency": round(latency, 3),
|
||||
"num_jsons": args.num_jsons,
|
||||
"parallel": args.parallel,
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="NousResearch/json-mode-eval")
|
||||
parser.add_argument("--num-jsons", type=int, default=-1)
|
||||
args = add_common_sglang_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,37 +0,0 @@
|
||||
## Download data
|
||||
|
||||
```
|
||||
wget https://raw.githubusercontent.com/merrymercy/merrymercy.github.io/master/files/random_words.json
|
||||
python3 gen_data.py --number 1000
|
||||
```
|
||||
|
||||
## Run benchmark
|
||||
|
||||
### Benchmark sglang
|
||||
```
|
||||
python3 -m sglang.launch_server --model-path codellama/CodeLlama-7b-hf --port 30000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_sglang.py --src-index 600 --num-q 50 --parallel 1
|
||||
```
|
||||
|
||||
|
||||
###
|
||||
|
||||
```
|
||||
# original
|
||||
Accuracy: 0.940, latency: 332.83 s
|
||||
|
||||
# parallel encoding (no_adjust, offset = 1000)
|
||||
Accuracy: 0.760, latency: 238.46 s
|
||||
|
||||
# parallel encoding (no_adjust, offset = 3000)
|
||||
Accuracy: 0.760, latency: 238.46 s
|
||||
|
||||
# parallel encoding (no_adjust, offset = 0)
|
||||
Accuracy: 0.520, latency: 238.46 s
|
||||
|
||||
# parallel encoding (adjust_cache)
|
||||
Accuracy: 0.460, latency: 257.66 s
|
||||
```
|
||||
@@ -1,149 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import re
|
||||
import time
|
||||
|
||||
import numpy as np
|
||||
|
||||
import sglang as sgl
|
||||
from sglang.test.test_utils import (
|
||||
add_common_sglang_args_and_parse,
|
||||
select_sglang_backend,
|
||||
)
|
||||
from sglang.utils import dump_state_text
|
||||
|
||||
|
||||
@sgl.function
|
||||
def line_retrieval(s, prefix, suffix, body_0, body_1, body_2, body_3):
|
||||
s += prefix + "\n"
|
||||
|
||||
contexts = [body_0, body_1, body_2, body_3]
|
||||
position_ids_offset = [i * 1000 for i in range(len(contexts))]
|
||||
forks = s.fork(len(contexts), position_ids_offset)
|
||||
forks += lambda i: contexts[i] + "\n"
|
||||
forks.join(mode="concate_and_append")
|
||||
|
||||
s += "\n" + suffix
|
||||
s += sgl.gen("answer", max_tokens=16)
|
||||
|
||||
|
||||
def eval_model(args, line_obj, num_hoops, src_indices, dst_percents):
|
||||
arguments = []
|
||||
labels = []
|
||||
sum_src_indices = []
|
||||
sum_dst_indices = []
|
||||
|
||||
for i in range(len(src_indices)):
|
||||
for j in range(len(dst_percents)):
|
||||
src_index = src_indices[i]
|
||||
dst_percent = dst_percents[j]
|
||||
|
||||
query_indices = line_obj["group_by_num_hoops"][str(num_hoops)]
|
||||
query_indices = [
|
||||
q
|
||||
for q in query_indices
|
||||
if all(l <= src_index for l in line_obj["links"][q]) and q < src_index
|
||||
]
|
||||
dst_index = query_indices[
|
||||
min(int(len(query_indices) * dst_percent), len(query_indices) - 1)
|
||||
]
|
||||
label = line_obj["values"][dst_index]
|
||||
|
||||
body = line_obj["lines"][: src_index + 1]
|
||||
suffix = line_obj["suffix"].replace("???", line_obj["indices"][dst_index])
|
||||
body_part_len = len(body) // 4
|
||||
|
||||
arguments.append(
|
||||
{
|
||||
"prefix": line_obj["prefix"],
|
||||
"body_0": "\n".join(body[:body_part_len]),
|
||||
"body_1": "\n".join(body[body_part_len : 2 * body_part_len]),
|
||||
"body_2": "\n".join(body[2 * body_part_len : 3 * body_part_len]),
|
||||
"body_3": "\n".join(body[3 * body_part_len :]),
|
||||
"suffix": suffix,
|
||||
}
|
||||
)
|
||||
labels.append(label)
|
||||
sum_src_indices.append(src_index)
|
||||
sum_dst_indices.append(dst_index)
|
||||
|
||||
# Select backend
|
||||
backend = select_sglang_backend(args)
|
||||
|
||||
tic = time.perf_counter()
|
||||
states = line_retrieval.run_batch(
|
||||
arguments,
|
||||
temperature=0,
|
||||
backend=backend,
|
||||
num_threads=args.parallel,
|
||||
progress_bar=True,
|
||||
)
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
corrects = []
|
||||
for i in range(len(arguments)):
|
||||
output = states[i]["answer"]
|
||||
prompt_len = states[i].get_meta_info("answer").get("prompt_length", -1)
|
||||
label = labels[i]
|
||||
|
||||
# Try all numbers
|
||||
findall = re.findall("\d+", output)
|
||||
if not findall:
|
||||
response_number = output
|
||||
else:
|
||||
for response_number in findall:
|
||||
if response_number == label:
|
||||
break
|
||||
|
||||
correct = response_number == label
|
||||
corrects.append(correct)
|
||||
|
||||
# Log results
|
||||
summary = (
|
||||
f"Line index: {sum_src_indices[i]} -> {sum_dst_indices[i]}, "
|
||||
f"Prompt len: {prompt_len}, "
|
||||
f"Correct: {correct}, "
|
||||
f"Label: {label}, Predicted: {response_number}, "
|
||||
)
|
||||
print(summary)
|
||||
|
||||
accuracy = np.mean(corrects)
|
||||
print(f"Accuracy: {accuracy:.3f}, latency: {latency:.2f} s")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "line_retrieval",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"num_requests": len(arguments),
|
||||
"other": {
|
||||
"num_questions": len(arguments),
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
def main(args):
|
||||
line_obj = json.load(open(args.data_path, "r"))
|
||||
|
||||
num_hoops = args.num_hoops
|
||||
for src_index in args.src_index:
|
||||
src_indices = [src_index]
|
||||
num_queries = args.num_queries_per_src
|
||||
dst_percents = [i * (1 / (num_queries)) for i in range(num_queries)]
|
||||
eval_model(args, line_obj, num_hoops, src_indices, dst_percents)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="lines_1000_0.0.json")
|
||||
parser.add_argument("--src-index", type=int, nargs="+", default=[100])
|
||||
parser.add_argument("--num-queries-per-src", type=int, default=10)
|
||||
parser.add_argument("--num-hoops", type=int, default=1)
|
||||
args = add_common_sglang_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,139 +0,0 @@
|
||||
"""
|
||||
Generate line data for line retrieval task.
|
||||
|
||||
Usage:
|
||||
python3 gen_data.py --number 1000
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from collections import defaultdict
|
||||
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
def generate_lines(random_words, num_lines, redirect_ratio):
|
||||
prefix = "Here is a list of lines, each with its corresponding REGISTER_CONTENT value. Please memorize them. Be prepared to provide the REGISTER_CONTENT value for a specific line index when I ask."
|
||||
suffix = "The list has ended. Please give the final REGISTER_CONTENT value for a specific line after resolving the redirections and references. For example, the REGISTER_CONTENT of Line __idx0__ is __val0__. The REGISTER_CONTENT of Line __idx1__ is __val1__. The REGISTER_CONTENT of Line __idx2__ is __val2__. The REGISTER_CONTENT of Line ??? is"
|
||||
|
||||
# Raw lines
|
||||
visited_indices = set([None])
|
||||
visited_values = set([None])
|
||||
|
||||
lines = []
|
||||
redirects = []
|
||||
indices = []
|
||||
values = []
|
||||
for i in tqdm(range(num_lines)):
|
||||
line_index = None
|
||||
while line_index in visited_indices:
|
||||
line_index = "-".join(np.random.choice(random_words, size=(2,)))
|
||||
visited_indices.add(line_index)
|
||||
|
||||
line_value = np.random.randint(low=0, high=999999)
|
||||
line_value = f"{line_value:06}"
|
||||
|
||||
line = f"Line {line_index}: The REGISTER_CONTENT is {line_value}."
|
||||
lines.append(line)
|
||||
redirects.append(None)
|
||||
indices.append(line_index)
|
||||
values.append(line_value)
|
||||
|
||||
# Add redirect
|
||||
if redirect_ratio > 0:
|
||||
num_redirect_lines = int(len(lines) * redirect_ratio)
|
||||
redirect_indices = np.random.choice(
|
||||
np.arange(len(lines)), size=(num_redirect_lines,), replace=False
|
||||
)
|
||||
for i in redirect_indices:
|
||||
target_idx = np.random.choice(min(i * 2 + 100, num_lines))
|
||||
lines[i] = (
|
||||
f"Line {indices[i]}: The REGISTER_CONTENT is the same as Line {indices[target_idx]}."
|
||||
)
|
||||
redirects[i] = target_idx
|
||||
|
||||
# Build links and find sources
|
||||
links = [[] for _ in range(num_lines)]
|
||||
contains_ring = set()
|
||||
for i in range(num_lines):
|
||||
if redirects[i] is None:
|
||||
continue
|
||||
|
||||
tmp_link = []
|
||||
cur = i
|
||||
visited = set()
|
||||
while redirects[cur] is not None:
|
||||
visited.add(cur)
|
||||
tmp_link.append(redirects[cur])
|
||||
cur = redirects[cur]
|
||||
|
||||
if cur in visited:
|
||||
contains_ring.add(i)
|
||||
tmp_link = None
|
||||
break
|
||||
values[i] = values[cur]
|
||||
links[i] = tmp_link
|
||||
|
||||
# Group by num_links
|
||||
group_by_num_hoops = defaultdict(list)
|
||||
for i in range(num_lines):
|
||||
if i in contains_ring:
|
||||
continue
|
||||
group_by_num_hoops[len(links[i]) + 1].append(i)
|
||||
|
||||
keys = sorted(list(group_by_num_hoops.keys()))
|
||||
for num_links in keys:
|
||||
print(f"#links: {num_links}, #lines: {len(group_by_num_hoops[num_links])}")
|
||||
|
||||
# Append few-shot examples
|
||||
hoop1_candidates = list(group_by_num_hoops[1])
|
||||
hoop1_candidate_keys = {c: max([c] + links[c]) for c in hoop1_candidates}
|
||||
hoop1_candidates.sort(key=lambda c: hoop1_candidate_keys[c])
|
||||
hoop2_candidates = list(group_by_num_hoops[2])
|
||||
hoop2_candidate_keys = {c: max([c] + links[c]) for c in hoop2_candidates}
|
||||
hoop2_candidates.sort(key=lambda c: hoop2_candidate_keys[c])
|
||||
|
||||
i = hoop1_candidates[5]
|
||||
suffix = suffix.replace("__idx0__", indices[i]).replace("__val0__", values[i])
|
||||
if len(hoop2_candidates):
|
||||
i = hoop2_candidates[0]
|
||||
suffix = suffix.replace("__idx1__", indices[i]).replace("__val1__", values[i])
|
||||
i = hoop2_candidates[1]
|
||||
suffix = suffix.replace("__idx2__", indices[i]).replace("__val2__", values[i])
|
||||
else:
|
||||
i = hoop1_candidates[1]
|
||||
suffix = suffix.replace("__idx1__", indices[i]).replace("__val1__", values[i])
|
||||
i = hoop1_candidates[10]
|
||||
suffix = suffix.replace("__idx2__", indices[i]).replace("__val2__", values[i])
|
||||
|
||||
obj = {
|
||||
"prefix": prefix,
|
||||
"suffix": suffix,
|
||||
"lines": lines,
|
||||
"indices": indices,
|
||||
"values": values,
|
||||
"links": links,
|
||||
"group_by_num_hoops": group_by_num_hoops,
|
||||
"contains_ring": sorted(list(contains_ring)),
|
||||
}
|
||||
return obj
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--number", type=int)
|
||||
parser.add_argument("--redirect-ratio", type=float, default=0.0)
|
||||
args = parser.parse_args()
|
||||
|
||||
num_lines = args.number
|
||||
|
||||
random_words_filename = "random_words.json"
|
||||
random_words = json.load(open(random_words_filename, "r"))
|
||||
|
||||
np.random.seed(42)
|
||||
obj = generate_lines(random_words, num_lines, args.redirect_ratio)
|
||||
|
||||
fout = f"lines_{num_lines}_{args.redirect_ratio:.1f}.json"
|
||||
with open(fout, "w") as fout:
|
||||
json.dump(obj, fout, indent=2)
|
||||
@@ -1,33 +0,0 @@
|
||||
## Run benchmark
|
||||
|
||||
### Benchmark sglang
|
||||
```
|
||||
python -m sglang.launch_server --model-path meta-llama/Llama-2-7b-chat-hf --port 30000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_sglang.py --num-questions 25 --parallel 8
|
||||
python3 bench_sglang.py --num-questions 16 --parallel 1
|
||||
```
|
||||
|
||||
|
||||
### Benchmark vllm
|
||||
```
|
||||
python3 -m vllm.entrypoints.api_server --tokenizer-mode auto --model meta-llama/Llama-2-7b-chat-hf --disable-log-requests --port 21000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --backend vllm --num-questions 25
|
||||
```
|
||||
|
||||
|
||||
### Benchmark guidance
|
||||
```
|
||||
python3 bench_other.py --backend guidance --num-questions 25 --parallel 1 --n-ctx 4096 --model-path path/to/gguf
|
||||
```
|
||||
|
||||
### Benchmark lmql
|
||||
|
||||
```
|
||||
python3 bench_other.py --backend lmql --num-questions 25 --parallel 1
|
||||
```
|
||||
File diff suppressed because one or more lines are too long
@@ -1,151 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from functools import partial
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
from sglang.test.test_utils import add_common_other_args_and_parse, get_call_generate
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
system_prompt = "Please serve as an impartial judge and rigorously evaluate the quality of the following article. Apply the most stringent standards possible, showing no leniency."
|
||||
|
||||
dimension_prompts = [
|
||||
"Content: This refers to the essences of the essay. The substance should be well researched, accurate, relevant to the topic and should show a thorough understanding of the subject. The essay should also reflect a clear goal or purpose.",
|
||||
"Organization and Structure: An essay needs to be properly structured with a clear introduction, body, and conclusion. The essay should flow naturally, with one paragraph leading seamlessly into the next.",
|
||||
"Argument and Analysis: The argument made in the essay should be logical, coherent and clearly articulated. Each point made should be backed up by solid evidence and thorough analysis.",
|
||||
"Clarity and Precision: The essay should be written in a clear and concise manner. The points made should be easily understood by the reader. The language used should also be precise and unambiguous.",
|
||||
"Grammar and Punctuation: Proper use of grammar and punctuation is vital in an academic essay. Errors in grammar and punctuation not only distract the reader but can also negatively impact the meaning and interpretation of the content.",
|
||||
"Referencing and Citation: An essay should contain proper citations and references for all sources used. This not only prevents accusations of plagiarism but also gives credit to the authors of the works that have contributed to the essay. The citation should adhere to a specific format as required by the academic institution or specified by the professor.",
|
||||
]
|
||||
|
||||
|
||||
def multi_dimension_judge(article, generate):
|
||||
s = system_prompt
|
||||
s += "\n```\n" + article + "\n```\n\n"
|
||||
|
||||
judges = []
|
||||
for i in range(len(dimension_prompts)):
|
||||
comp = generate(
|
||||
s
|
||||
+ "USER: Please judge the quality based on the following metric. "
|
||||
+ dimension_prompts[i]
|
||||
+ " Please provide a single-paragraph judgement. "
|
||||
+ "Focus on the provided metric and do not say other things. "
|
||||
'End your judgement paragraph with the word "END"\nJUDGE:',
|
||||
max_tokens=256,
|
||||
stop="END",
|
||||
)
|
||||
judges.append(comp)
|
||||
|
||||
s += "I will judge the quality based on the following metrics.\n"
|
||||
for i in range(len(dimension_prompts)):
|
||||
s += dimension_prompts[i].split(":")[0] + ": " + judges[i].strip() + "\n"
|
||||
|
||||
s += "In summary, on a scale of 1 to 10, I would give the article a score of"
|
||||
s += generate(s, max_tokens=2, stop=None)
|
||||
|
||||
return s
|
||||
|
||||
|
||||
async def multi_dimension_judge_async(article, generate):
|
||||
s = system_prompt
|
||||
s += "\n```\n" + article + "\n```\n\n"
|
||||
|
||||
judges = []
|
||||
for i in range(len(dimension_prompts)):
|
||||
comp = await generate(
|
||||
s
|
||||
+ "USER: Please judge the quality based on the following metric. "
|
||||
+ dimension_prompts[i]
|
||||
+ " Please provide a single-paragraph judgement. "
|
||||
+ "Focus on the provided metric and do not say other things. "
|
||||
'End your judgement paragraph with the word "END"\nJUDGE:',
|
||||
max_tokens=256,
|
||||
stop="END",
|
||||
)
|
||||
judges.append(comp)
|
||||
|
||||
s += "I will judge the quality based on the following metrics.\n"
|
||||
for i in range(len(dimension_prompts)):
|
||||
s += dimension_prompts[i].split(":")[0] + ": " + judges[i].strip() + "\n"
|
||||
|
||||
s += "In summary, on a scale of 1 to 10, I would give the article a score of"
|
||||
s += await generate(s, max_tokens=2, stop=None)
|
||||
|
||||
return s
|
||||
|
||||
|
||||
def main(args):
|
||||
lines = read_jsonl(args.data_path)[: args.num_questions]
|
||||
states = [None] * len(lines)
|
||||
|
||||
# Select backend
|
||||
call_generate = partial(get_call_generate(args), temperature=0)
|
||||
|
||||
# Run requests
|
||||
tic = time.perf_counter()
|
||||
|
||||
if args.backend != "lmql":
|
||||
|
||||
def get_one_answer(i):
|
||||
states[i] = multi_dimension_judge(lines[i], call_generate)
|
||||
|
||||
if args.parallel == 1:
|
||||
for i in tqdm(range(len(lines))):
|
||||
get_one_answer(i)
|
||||
else:
|
||||
with ThreadPoolExecutor(args.parallel) as executor:
|
||||
list(
|
||||
tqdm(
|
||||
executor.map(get_one_answer, list(range(len(lines)))),
|
||||
total=len(lines),
|
||||
)
|
||||
)
|
||||
|
||||
else:
|
||||
import asyncio
|
||||
|
||||
async def get_one_answer_async(i):
|
||||
states[i] = await multi_dimension_judge_async(lines[i], call_generate)
|
||||
|
||||
batches = []
|
||||
for i in range(0, len(lines), args.parallel):
|
||||
batches.append(list(range(i, min(i + args.parallel, len(lines)))))
|
||||
|
||||
loop = asyncio.get_event_loop()
|
||||
for bt in tqdm(batches):
|
||||
loop.run_until_complete(
|
||||
asyncio.gather(*[get_one_answer_async(i) for i in bt])
|
||||
)
|
||||
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
# Compute accuracy
|
||||
print(f"Latency: {latency:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "llm_judge",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"other": {
|
||||
"num_questions": args.num_questions,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="articles.jsonl")
|
||||
parser.add_argument("--num-questions", type=int, default=20)
|
||||
args = add_common_other_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,97 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
|
||||
import sglang as sgl
|
||||
from sglang.test.test_utils import (
|
||||
add_common_sglang_args_and_parse,
|
||||
select_sglang_backend,
|
||||
)
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
system_prompt = "Please serve as an impartial judge and rigorously evaluate the quality of the following article. Apply the most stringent standards possible, showing no leniency."
|
||||
|
||||
dimension_prompts = [
|
||||
"Content: This refers to the essences of the essay. The substance should be well researched, accurate, relevant to the topic and should show a thorough understanding of the subject. The essay should also reflect a clear goal or purpose.",
|
||||
"Organization and Structure: An essay needs to be properly structured with a clear introduction, body, and conclusion. The essay should flow naturally, with one paragraph leading seamlessly into the next.",
|
||||
"Argument and Analysis: The argument made in the essay should be logical, coherent and clearly articulated. Each point made should be backed up by solid evidence and thorough analysis.",
|
||||
"Clarity and Precision: The essay should be written in a clear and concise manner. The points made should be easily understood by the reader. The language used should also be precise and unambiguous.",
|
||||
"Grammar and Punctuation: Proper use of grammar and punctuation is vital in an academic essay. Errors in grammar and punctuation not only distract the reader but can also negatively impact the meaning and interpretation of the content.",
|
||||
"Referencing and Citation: An essay should contain proper citations and references for all sources used. This not only prevents accusations of plagiarism but also gives credit to the authors of the works that have contributed to the essay. The citation should adhere to a specific format as required by the academic institution or specified by the professor.",
|
||||
]
|
||||
|
||||
|
||||
@sgl.function
|
||||
def multi_dimension_judge(s, article):
|
||||
s += system_prompt
|
||||
s += "\n```\n" + article + "\n```\n\n"
|
||||
|
||||
forks = s.fork(len(dimension_prompts))
|
||||
for i in range(len(dimension_prompts)):
|
||||
forks[i] += (
|
||||
"USER: Please judge the quality based on the following metric. "
|
||||
+ dimension_prompts[i]
|
||||
+ " Please provide a single-paragraph judgement. "
|
||||
+ "Focus on the provided metric and do not say other things. "
|
||||
'End your judgement paragraph with the word "END"\nJUDGE:'
|
||||
)
|
||||
forks[i] += sgl.gen("judgement", max_tokens=256, stop="END")
|
||||
forks.join()
|
||||
|
||||
s += "I will judge the quality based on the following metrics.\n"
|
||||
for i in range(len(dimension_prompts)):
|
||||
s += (
|
||||
dimension_prompts[i].split(":")[0]
|
||||
+ ": "
|
||||
+ forks[i]["judgement"].strip()
|
||||
+ "\n"
|
||||
)
|
||||
|
||||
s += "In summary, on a scale of 1 to 10, I would give the article a score of"
|
||||
s += sgl.gen("score", max_tokens=2)
|
||||
|
||||
|
||||
def main(args):
|
||||
lines = read_jsonl(args.data_path)[: args.num_questions]
|
||||
arguments = [{"article": l} for l in lines]
|
||||
|
||||
# Select backend
|
||||
backend = select_sglang_backend(args)
|
||||
|
||||
# Run requests
|
||||
tic = time.perf_counter()
|
||||
states = multi_dimension_judge.run_batch(
|
||||
arguments,
|
||||
temperature=0,
|
||||
backend=backend,
|
||||
num_threads=args.parallel,
|
||||
progress_bar=True,
|
||||
)
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
print(f"Latency: {latency:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "llm_judge",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"other": {
|
||||
"num_questions": args.num_questions,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="articles.jsonl")
|
||||
parser.add_argument("--num-questions", type=int, default=20)
|
||||
args = add_common_sglang_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,33 +0,0 @@
|
||||
## Run benchmark
|
||||
|
||||
### Benchmark sglang
|
||||
```
|
||||
python3 -m sglang.launch_server --model-path codellama/CodeLlama-7b-instruct-hf --port 30000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_sglang.py --num-questions 5 --parallel 1
|
||||
```
|
||||
|
||||
|
||||
### Benchmark vllm
|
||||
```
|
||||
python3 -m vllm.entrypoints.api_server --tokenizer-mode auto --model codellama/CodeLlama-7b-instruct-hf --disable-log-requests --port 21000 --gpu 0.97
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --backend vllm --num-questions 5
|
||||
```
|
||||
|
||||
|
||||
### Benchmark guidance
|
||||
```
|
||||
python3 bench_other.py --backend guidance --num-questions 5 --parallel 1 --n-ctx 11000 --model-path path/to/code-llama/gguf
|
||||
```
|
||||
|
||||
|
||||
### Build dataset
|
||||
```
|
||||
pip install wikipedia
|
||||
python3 build_dataset.py
|
||||
```
|
||||
@@ -1,89 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from functools import partial
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
from sglang.test.test_utils import add_common_other_args_and_parse, get_call_generate
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
|
||||
def json_decode(document, generate):
|
||||
s = "Please extract the information of a city from the following wikipedia page.\n"
|
||||
s += "Page begin.\n" + document + "Page end.\n"
|
||||
s += "Here is the name, country, and symbol of the city in JSON format.\n"
|
||||
s += "{\n"
|
||||
s += ' "name": "'
|
||||
s += generate(s, max_tokens=8, stop='"') + '",\n'
|
||||
s += ' "country": "'
|
||||
s += generate(s, max_tokens=8, stop='"') + '",\n'
|
||||
s += ' "air port code": "'
|
||||
s += generate(s, max_tokens=8, stop='"') + '",\n'
|
||||
s += ' "top 3 landmarks": "'
|
||||
s += generate(s, max_tokens=24, stop='"') + '",\n'
|
||||
s += "}\n"
|
||||
return s
|
||||
|
||||
|
||||
def main(args):
|
||||
lines = read_jsonl(args.data_path)
|
||||
arguments = []
|
||||
for i in range(len(lines[: args.num_questions])):
|
||||
arguments.append(
|
||||
{
|
||||
"document": lines[i]["document"],
|
||||
}
|
||||
)
|
||||
states = [None] * len(arguments)
|
||||
|
||||
# Select backend
|
||||
call_generate = partial(get_call_generate(args), temperature=0)
|
||||
|
||||
# Run requests
|
||||
def get_one_answer(i):
|
||||
states[i] = json_decode(generate=call_generate, **arguments[i])
|
||||
|
||||
tic = time.perf_counter()
|
||||
if args.parallel == 1:
|
||||
for i in tqdm(range(len(arguments))):
|
||||
get_one_answer(i)
|
||||
else:
|
||||
with ThreadPoolExecutor(args.parallel) as executor:
|
||||
list(
|
||||
tqdm(
|
||||
executor.map(get_one_answer, list(range(len(arguments)))),
|
||||
total=len(arguments),
|
||||
)
|
||||
)
|
||||
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
# Compute accuracy
|
||||
print(f"Latency: {latency:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "long_json_decode",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"other": {
|
||||
"num_questions": args.num_questions,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="questions.jsonl")
|
||||
parser.add_argument("--num-questions", type=int, default=100)
|
||||
args = add_common_other_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,81 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
|
||||
import sglang as sgl
|
||||
from sglang.test.test_utils import (
|
||||
add_common_sglang_args_and_parse,
|
||||
select_sglang_backend,
|
||||
)
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
|
||||
@sgl.function
|
||||
def json_decode(s, document):
|
||||
s += "Please extract the information of a city from the following wikipedia page.\n"
|
||||
s += "Page begin.\n" + document + "Page end.\n"
|
||||
s += "Here is the name, country, and symbol of the city in JSON format.\n"
|
||||
s += "{\n"
|
||||
s += ' "name": "' + sgl.gen("name", max_tokens=8, stop='"') + '",\n'
|
||||
s += ' "country": "' + sgl.gen("country", max_tokens=8, stop='"') + '",\n'
|
||||
s += (
|
||||
' "air port code": "'
|
||||
+ sgl.gen("air port code", max_tokens=8, stop='"')
|
||||
+ '",\n'
|
||||
)
|
||||
s += (
|
||||
' "top 3 landmarks": "'
|
||||
+ sgl.gen("landmarks", max_tokens=24, stop='"')
|
||||
+ '",\n'
|
||||
)
|
||||
s += "}\n"
|
||||
|
||||
|
||||
def main(args):
|
||||
lines = read_jsonl(args.data_path)
|
||||
arguments = []
|
||||
for i in range(len(lines[: args.num_questions])):
|
||||
arguments.append(
|
||||
{
|
||||
"document": lines[i]["document"],
|
||||
}
|
||||
)
|
||||
|
||||
# Select backend
|
||||
backend = select_sglang_backend(args)
|
||||
sgl.set_default_backend(backend)
|
||||
|
||||
# Run requests
|
||||
tic = time.perf_counter()
|
||||
states = json_decode.run_batch(
|
||||
arguments, temperature=0, num_threads=args.parallel, progress_bar=True
|
||||
)
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
# Compute accuracy
|
||||
print(f"Latency: {latency:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "long_json_decode",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"other": {
|
||||
"num_questions": args.num_questions,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="questions.jsonl")
|
||||
parser.add_argument("--num-questions", type=int, default=10)
|
||||
args = add_common_sglang_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,27 +0,0 @@
|
||||
import json
|
||||
|
||||
import transformers
|
||||
import wikipedia
|
||||
|
||||
name = "meta-llama/Llama-2-7b-chat-hf"
|
||||
t = transformers.AutoTokenizer.from_pretrained(name)
|
||||
city_names = ["los angles", "london", "tokyo", "beijing", "singapore"]
|
||||
|
||||
|
||||
for city_name in city_names:
|
||||
content = str(wikipedia.page(city_name).content)
|
||||
content = content.replace("\n\n", "\n")
|
||||
|
||||
tokens = t.encode(content)
|
||||
|
||||
truncate_len = int((10000 / len(tokens)) * len(content))
|
||||
truncate_content = content[:truncate_len]
|
||||
truncate_tokens = t.encode(truncate_content)
|
||||
|
||||
# Count token
|
||||
print(
|
||||
f"city_name: {city_name}, #tokens: {len(tokens)}, #truncate tokens: {len(truncate_tokens)}"
|
||||
)
|
||||
|
||||
with open("questions.jsonl", "a") as fout:
|
||||
fout.write(json.dumps({"document": truncate_content}) + "\n")
|
||||
@@ -18,42 +18,3 @@ python3 bench_sglang.py --nsub 10
|
||||
# OpenAI models
|
||||
python3 bench_sglang.py --backend gpt-3.5-turbo --parallel 8
|
||||
```
|
||||
|
||||
### Benchmark vllm
|
||||
```
|
||||
python3 -m vllm.entrypoints.api_server --tokenizer-mode auto --model meta-llama/Llama-2-7b-chat-hf --disable-log-requests --port 21000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --nsub 10 --backend vllm
|
||||
```
|
||||
|
||||
|
||||
### Benchmark lightllm
|
||||
```
|
||||
# A10G
|
||||
python -m lightllm.server.api_server --tokenizer_mode auto --model_dir ~/model_weights/llama-2-7b-chat-hf --max_total_token_num 16000 --port 22000
|
||||
|
||||
# V100
|
||||
python -m lightllm.server.api_server --tokenizer_mode auto --model_dir ~/model_weights/llama-2-7b-chat-hf --max_total_token_num 4500 --port 22000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --nsub 10 --backend lightllm
|
||||
```
|
||||
|
||||
|
||||
### Benchmark guidance
|
||||
```
|
||||
python3 bench_other.py --nsub 10 --backend guidance --parallel 1 --n-ctx 4096 --model-path path/to/gguf
|
||||
```
|
||||
|
||||
|
||||
### Benchmark lmql
|
||||
```
|
||||
CUDA_VISIBLE_DEVICES=0,1 lmql serve-model meta-llama/Llama-2-7b-chat-hf --cuda --port 23000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --nsub 10 --backend lmql --parallel 2
|
||||
```
|
||||
|
||||
@@ -1,173 +0,0 @@
|
||||
import argparse
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import tiktoken
|
||||
from tqdm import tqdm
|
||||
|
||||
from sglang.test.test_utils import add_common_other_args_and_parse, get_call_generate
|
||||
|
||||
choices = ["A", "B", "C", "D"]
|
||||
|
||||
tokenizer = tiktoken.encoding_for_model("gpt-3.5-turbo")
|
||||
|
||||
|
||||
def format_subject(subject):
|
||||
l = subject.split("_")
|
||||
s = ""
|
||||
for entry in l:
|
||||
s += " " + entry
|
||||
return s
|
||||
|
||||
|
||||
def format_example(df, idx, include_answer=True):
|
||||
prompt = df.iloc[idx, 0]
|
||||
k = df.shape[1] - 2
|
||||
for j in range(k):
|
||||
prompt += "\n{}. {}".format(choices[j], df.iloc[idx, j + 1])
|
||||
prompt += "\nAnswer:"
|
||||
if include_answer:
|
||||
prompt += " {}\n\n".format(df.iloc[idx, k + 1])
|
||||
return prompt
|
||||
|
||||
|
||||
def gen_prompt(train_df, subject, k=-1):
|
||||
prompt = "The following are multiple choice questions (with answers) about{}.\n\n".format(
|
||||
format_subject(subject)
|
||||
)
|
||||
if k == -1:
|
||||
k = train_df.shape[0]
|
||||
for i in range(k):
|
||||
prompt += format_example(train_df, i)
|
||||
return prompt
|
||||
|
||||
|
||||
def evaluate(args, subject, dev_df, test_df, call_generate):
|
||||
prompts = []
|
||||
labels = []
|
||||
|
||||
# Construct prompts
|
||||
k = args.ntrain
|
||||
train_prompt = gen_prompt(dev_df, subject, k)
|
||||
while len(tokenizer.encode(train_prompt)) > 1536:
|
||||
k -= 1
|
||||
train_prompt = gen_prompt(dev_df, subject, k)
|
||||
|
||||
for i in range(test_df.shape[0]):
|
||||
prompt_end = format_example(test_df, i, include_answer=False)
|
||||
prompt = train_prompt + prompt_end
|
||||
prompts.append(prompt)
|
||||
|
||||
label = test_df.iloc[i, test_df.shape[1] - 1]
|
||||
labels.append(label)
|
||||
|
||||
preds = [None] * len(prompts)
|
||||
max_tokens = 1
|
||||
|
||||
# Run requests
|
||||
if args.backend != "lmql":
|
||||
# Use thread pool
|
||||
def get_one_answer(i):
|
||||
pred = call_generate(prompts[i], temperature=0, max_tokens=max_tokens)
|
||||
preds[i] = pred.strip()[0]
|
||||
|
||||
tic = time.perf_counter()
|
||||
if args.parallel == 1:
|
||||
for i in range(len(prompts)):
|
||||
get_one_answer(i)
|
||||
else:
|
||||
with ThreadPoolExecutor(args.parallel) as executor:
|
||||
executor.map(get_one_answer, list(range(len(prompts))))
|
||||
else:
|
||||
# Use asyncio
|
||||
async def batched_call(batch_size):
|
||||
for i in range(0, len(prompts), batch_size):
|
||||
tasks = []
|
||||
for p in prompts[i : i + batch_size]:
|
||||
tasks.append(call_generate(p, temperature=0, max_tokens=max_tokens))
|
||||
rets = await asyncio.gather(*tasks)
|
||||
for j in range(len(rets)):
|
||||
preds[i + j] = rets[j].strip()[0]
|
||||
|
||||
tic = time.perf_counter()
|
||||
asyncio.run(batched_call(batch_size=args.parallel))
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
# Compute accuracy
|
||||
cors = [pred == label for pred, label in zip(preds, labels)]
|
||||
acc = np.mean(cors)
|
||||
cors = np.array(cors)
|
||||
|
||||
print(
|
||||
"Average accuracy {:.3f}, latency {:.2f}, #q: {} - {}".format(
|
||||
acc, latency, len(prompts), subject
|
||||
)
|
||||
)
|
||||
|
||||
return cors, acc, latency
|
||||
|
||||
|
||||
def main(args):
|
||||
subjects = sorted(
|
||||
[
|
||||
f.split("_test.csv")[0]
|
||||
for f in os.listdir(os.path.join(args.data_dir, "test"))
|
||||
if "_test.csv" in f
|
||||
]
|
||||
)
|
||||
|
||||
all_cors = []
|
||||
all_latencies = []
|
||||
num_requests = 0
|
||||
|
||||
# Select backend
|
||||
call_generate = get_call_generate(args)
|
||||
|
||||
for subject in tqdm(subjects[: args.nsub]):
|
||||
dev_df = pd.read_csv(
|
||||
os.path.join(args.data_dir, "dev", subject + "_dev.csv"), header=None
|
||||
)[: args.ntrain]
|
||||
test_df = pd.read_csv(
|
||||
os.path.join(args.data_dir, "test", subject + "_test.csv"), header=None
|
||||
)
|
||||
|
||||
cors, acc, latency = evaluate(args, subject, dev_df, test_df, call_generate)
|
||||
all_cors.append(cors)
|
||||
all_latencies.append(latency)
|
||||
num_requests += len(test_df)
|
||||
|
||||
total_latency = np.sum(all_latencies)
|
||||
print("Total latency: {:.3f}".format(total_latency))
|
||||
|
||||
weighted_acc = np.mean(np.concatenate(all_cors))
|
||||
print("Average accuracy: {:.3f}".format(weighted_acc))
|
||||
|
||||
# Write results
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "mmlu",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(total_latency, 3),
|
||||
"accuracy": round(weighted_acc, 3),
|
||||
"num_requests": num_requests,
|
||||
"other": {
|
||||
"nsub": args.nsub,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--ntrain", type=int, default=5)
|
||||
parser.add_argument("--data_dir", type=str, default="data")
|
||||
parser.add_argument("--nsub", type=int, default=60)
|
||||
args = add_common_other_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,48 +0,0 @@
|
||||
## Download Dataset
|
||||
|
||||
```sh
|
||||
wget -O question.jsonl https://raw.githubusercontent.com/lm-sys/FastChat/main/fastchat/llm_judge/data/mt_bench/question.jsonl
|
||||
```
|
||||
|
||||
## Run benchmark
|
||||
|
||||
### Benchmark sglang
|
||||
```
|
||||
python -m sglang.launch_server --model-path meta-llama/Llama-2-7b-chat-hf --port 30000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_sglang.py --num-questions 80
|
||||
```
|
||||
|
||||
### Benchmark sglang EAGLE
|
||||
```
|
||||
python3 -m sglang.launch_server --model meta-llama/Meta-Llama-3-8B-Instruct --speculative-algo EAGLE \
|
||||
--speculative-draft-model-path lmsys/sglang-EAGLE-LLaMA3-Instruct-8B --speculative-num-steps 5 \
|
||||
--speculative-eagle-topk 8 --speculative-num-draft-tokens 64 --dtype float16 --port 30000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_sglang_eagle.py --num-questions 80 --parallel 1
|
||||
```
|
||||
|
||||
|
||||
### Benchmark vllm
|
||||
```
|
||||
python3 -m vllm.entrypoints.api_server --tokenizer-mode auto --model meta-llama/Llama-2-7b-chat-hf --disable-log-requests --port 21000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --num-questions 80 --backend vllm
|
||||
```
|
||||
|
||||
|
||||
### Benchmark lightllm
|
||||
```
|
||||
# A10G
|
||||
python -m lightllm.server.api_server --tokenizer_mode auto --model_dir ~/model_weights/llama-2-7b-chat-hf --max_total_token_num 16000 --port 22000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --num-questions 80 --backend lightllm
|
||||
```
|
||||
@@ -1,118 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
from fastchat.model import get_conversation_template
|
||||
from tqdm import tqdm
|
||||
|
||||
from sglang.test.test_utils import add_common_other_args_and_parse, get_call_generate
|
||||
from sglang.utils import download_and_cache_file
|
||||
|
||||
|
||||
def load_questions(filename):
|
||||
questions = []
|
||||
with open(filename, "r") as fin:
|
||||
for line in fin:
|
||||
obj = json.loads(line)
|
||||
questions.append(obj)
|
||||
return questions
|
||||
|
||||
|
||||
def write_answers(filename, model_id, questions, answers):
|
||||
with open(os.path.expanduser(filename), "w") as fout:
|
||||
for i in range(len(answers)):
|
||||
ans_json = {
|
||||
"question_id": questions[i]["question_id"],
|
||||
"answer_id": uuid.uuid4().hex,
|
||||
"model_id": model_id,
|
||||
"choices": {
|
||||
"index": 0,
|
||||
"turns": [answers[i][0], answers[i][1]],
|
||||
},
|
||||
"tstamp": time.time(),
|
||||
}
|
||||
fout.write(json.dumps(ans_json) + "\n")
|
||||
|
||||
|
||||
def main(args):
|
||||
# Download question file if not exist
|
||||
question_file = args.question_file
|
||||
url = "https://raw.githubusercontent.com/lm-sys/FastChat/main/fastchat/llm_judge/data/mt_bench/question.jsonl"
|
||||
if not os.path.isfile(question_file):
|
||||
question_file = download_and_cache_file(url)
|
||||
|
||||
questions = load_questions(question_file)
|
||||
questions = (questions * 10)[: args.num_questions]
|
||||
max_tokens = 256
|
||||
model_id = "llama-2-chat"
|
||||
|
||||
conv_main = get_conversation_template(model_id)
|
||||
|
||||
# Select backend
|
||||
call_generate = get_call_generate(args)
|
||||
|
||||
answers = [None] * len(questions)
|
||||
|
||||
def get_answer(i):
|
||||
conv = conv_main.copy()
|
||||
cur_answers = []
|
||||
for j in range(2):
|
||||
q = questions[i]["turns"][j]
|
||||
conv.append_message(conv.roles[0], q)
|
||||
conv.append_message(conv.roles[1], None)
|
||||
|
||||
prompt = conv.get_prompt()
|
||||
output = call_generate(prompt, temperature=0, max_tokens=max_tokens).strip()
|
||||
|
||||
cur_answers.append(output)
|
||||
conv.update_last_message(output)
|
||||
|
||||
answers[i] = cur_answers
|
||||
|
||||
# Run requests
|
||||
tic = time.perf_counter()
|
||||
if args.parallel == 1:
|
||||
for i in tqdm(range(len(questions))):
|
||||
get_answer(i)
|
||||
else:
|
||||
with ThreadPoolExecutor(args.parallel) as executor:
|
||||
list(
|
||||
tqdm(
|
||||
executor.map(get_answer, list(range(len(questions)))),
|
||||
total=len(questions),
|
||||
)
|
||||
)
|
||||
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
print(f"#questions: {len(questions)}, Latency: {latency:.2f}")
|
||||
|
||||
# Write results
|
||||
answer_file = args.answer_file or f"tmp_output_{args.backend}.txt"
|
||||
write_answers(answer_file, model_id, questions, answers)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "mtbench",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"other": {
|
||||
"num_questions": args.num_questions,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--question-file", type=str, default="question.jsonl")
|
||||
parser.add_argument("--answer-file", type=str, default=None)
|
||||
parser.add_argument("--num-questions", type=int, default=80)
|
||||
args = add_common_other_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,106 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
|
||||
import sglang as sgl
|
||||
from sglang.test.test_utils import (
|
||||
add_common_sglang_args_and_parse,
|
||||
select_sglang_backend,
|
||||
)
|
||||
from sglang.utils import download_and_cache_file
|
||||
|
||||
|
||||
def load_questions(filename):
|
||||
questions = []
|
||||
with open(filename, "r") as fin:
|
||||
for line in fin:
|
||||
obj = json.loads(line)
|
||||
questions.append(obj)
|
||||
return questions
|
||||
|
||||
|
||||
def write_answers(filename, model_id, questions, answers):
|
||||
with open(os.path.expanduser(filename), "w") as fout:
|
||||
for i in range(len(answers)):
|
||||
ans_json = {
|
||||
"question_id": questions[i]["question_id"],
|
||||
"answer_id": uuid.uuid4().hex,
|
||||
"model_id": model_id,
|
||||
"choices": {
|
||||
"index": 0,
|
||||
"turns": [answers[i][0], answers[i][1]],
|
||||
},
|
||||
"tstamp": time.time(),
|
||||
}
|
||||
fout.write(json.dumps(ans_json) + "\n")
|
||||
|
||||
|
||||
@sgl.function
|
||||
def answer_mt_bench(s, question_1, question_2):
|
||||
s += sgl.system()
|
||||
s += sgl.user(question_1)
|
||||
s += sgl.assistant(sgl.gen("answer_1"))
|
||||
s += sgl.user(question_2)
|
||||
s += sgl.assistant(sgl.gen("answer_2"))
|
||||
|
||||
|
||||
def main(args):
|
||||
# Download question file if not exist
|
||||
question_file = args.question_file
|
||||
url = "https://raw.githubusercontent.com/lm-sys/FastChat/main/fastchat/llm_judge/data/mt_bench/question.jsonl"
|
||||
if not os.path.isfile(question_file):
|
||||
question_file = download_and_cache_file(url)
|
||||
|
||||
# Construct prompts
|
||||
questions = load_questions(question_file)[: args.num_questions]
|
||||
arguments = [
|
||||
{"question_1": q["turns"][0], "question_2": q["turns"][1]} for q in questions
|
||||
]
|
||||
|
||||
# Select backend
|
||||
backend = select_sglang_backend(args)
|
||||
sgl.set_default_backend(backend)
|
||||
|
||||
# Run requests
|
||||
tic = time.perf_counter()
|
||||
rets = answer_mt_bench.run_batch(
|
||||
arguments,
|
||||
temperature=0,
|
||||
max_new_tokens=256,
|
||||
num_threads=args.parallel,
|
||||
progress_bar=True,
|
||||
)
|
||||
answers = [[s["answer_1"], s["answer_2"]] for s in rets]
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
print(f"#questions: {len(questions)}, Latency: {latency:.2f}")
|
||||
|
||||
# Write results
|
||||
model_id = backend.model_info["model_path"]
|
||||
answer_file = args.answer_file or f"tmp_output_{args.backend}.txt"
|
||||
write_answers(answer_file, model_id, questions, answers)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "mtbench",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"other": {
|
||||
"num_questions": args.num_questions,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--question-file", type=str, default="question.jsonl")
|
||||
parser.add_argument("--answer-file", type=str, default=None)
|
||||
parser.add_argument("--num-questions", type=int, default=80)
|
||||
args = add_common_sglang_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,144 +0,0 @@
|
||||
"""
|
||||
Adapted from https://github.com/chromecast56/sglang/blob/6f145d2eadb93a116134f703358ce76f15381045/benchmark/mtbench/bench_sglang.py
|
||||
|
||||
Benchmark SGLang EAGLE/EAGLE3 Speculative Decoding
|
||||
|
||||
Usage:
|
||||
python3 benchmark/mtbench/bench_sglang_eagle.py --num-questions 80 --parallel 1
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
|
||||
import sglang as sgl
|
||||
from sglang.test.test_utils import (
|
||||
add_common_sglang_args_and_parse,
|
||||
select_sglang_backend,
|
||||
)
|
||||
from sglang.utils import download_and_cache_file
|
||||
|
||||
|
||||
def load_questions(filename):
|
||||
questions = []
|
||||
with open(filename, "r") as fin:
|
||||
for line in fin:
|
||||
obj = json.loads(line)
|
||||
questions.append(obj)
|
||||
return questions
|
||||
|
||||
|
||||
def write_answers(filename, model_id, questions, answers):
|
||||
with open(os.path.expanduser(filename), "w") as fout:
|
||||
for i in range(len(answers)):
|
||||
ans_json = {
|
||||
"question_id": questions[i]["question_id"],
|
||||
"answer_id": uuid.uuid4().hex,
|
||||
"model_id": model_id,
|
||||
"choices": {
|
||||
"index": 0,
|
||||
"turns": [answers[i][0], answers[i][1]],
|
||||
},
|
||||
"tstamp": time.time(),
|
||||
}
|
||||
fout.write(json.dumps(ans_json) + "\n")
|
||||
|
||||
|
||||
@sgl.function
|
||||
def answer_mt_bench(s, question_1, question_2):
|
||||
s += sgl.system(
|
||||
"You are a helpful, respectful and honest assistant. Always answer as helpfully as possible, while being safe. Your answers should not include any harmful, unethical, racist, sexist, toxic, dangerous, or illegal content. Please ensure that your responses are socially unbiased and positive in nature.\n\nIf a question does not make any sense, or is not factually coherent, explain why instead of answering something not correct. If you don't know the answer to a question, please don't share false information."
|
||||
)
|
||||
s += sgl.user(question_1)
|
||||
s += sgl.assistant(sgl.gen("answer_1"))
|
||||
s += sgl.user(question_2)
|
||||
s += sgl.assistant(sgl.gen("answer_2"))
|
||||
|
||||
|
||||
def main(args):
|
||||
# Download question file if not exist
|
||||
question_file = args.question_file
|
||||
url = "https://raw.githubusercontent.com/lm-sys/FastChat/main/fastchat/llm_judge/data/mt_bench/question.jsonl"
|
||||
if not os.path.isfile(question_file):
|
||||
question_file = download_and_cache_file(url)
|
||||
|
||||
# Construct prompts
|
||||
questions = load_questions(question_file)[: args.num_questions]
|
||||
arguments = [
|
||||
{"question_1": q["turns"][0], "question_2": q["turns"][1]} for q in questions
|
||||
]
|
||||
|
||||
# Select backend
|
||||
backend = select_sglang_backend(args)
|
||||
sgl.set_default_backend(backend)
|
||||
|
||||
# Run requests
|
||||
tic = time.perf_counter()
|
||||
rets = answer_mt_bench.run_batch(
|
||||
arguments,
|
||||
temperature=0,
|
||||
max_new_tokens=2048,
|
||||
num_threads=args.parallel,
|
||||
progress_bar=True,
|
||||
)
|
||||
answers = [[s["answer_1"], s["answer_2"]] for s in rets]
|
||||
|
||||
latency = time.perf_counter() - tic
|
||||
num_output_tokens = sum(
|
||||
s.get_meta_info("answer_1")["completion_tokens"]
|
||||
+ s.get_meta_info("answer_2")["completion_tokens"]
|
||||
for s in rets
|
||||
)
|
||||
|
||||
# NOTE: acceptance length is just completion_tokens / spec_verify_ct
|
||||
# {'id': '3bb9c5ead109488d8ed5ee9cbecaec29', 'finish_reason': {'type': 'length', 'length': 256}, 'prompt_tokens': 37, 'spec_verify_ct': 101, 'completion_tokens': 256, 'cached_tokens': 0}
|
||||
|
||||
output_throughput = num_output_tokens / latency
|
||||
|
||||
has_verify = "spec_verify_ct" in rets[0].get_meta_info("answer_1")
|
||||
if has_verify:
|
||||
num_verify_tokens = sum(
|
||||
s.get_meta_info("answer_1")["spec_verify_ct"]
|
||||
+ s.get_meta_info("answer_2")["spec_verify_ct"]
|
||||
for s in rets
|
||||
)
|
||||
|
||||
accept_length = num_output_tokens / num_verify_tokens
|
||||
else:
|
||||
accept_length = 1.0
|
||||
|
||||
print(
|
||||
f"#questions: {len(questions)}, Throughput: {output_throughput:.2f} token/s, Acceptance length: {accept_length:.2f}"
|
||||
)
|
||||
|
||||
# Write results
|
||||
model_id = backend.model_info["model_path"]
|
||||
answer_file = args.answer_file or f"tmp_output_{args.backend}.txt"
|
||||
write_answers(answer_file, model_id, questions, answers)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "mtbench",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"throughput": round(output_throughput, 3),
|
||||
"accept_length": round(accept_length, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"other": {
|
||||
"num_questions": args.num_questions,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--question-file", type=str, default="question.jsonl")
|
||||
parser.add_argument("--answer-file", type=str, default=None)
|
||||
parser.add_argument("--num-questions", type=int, default=80)
|
||||
args = add_common_sglang_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,49 +0,0 @@
|
||||
## Download data
|
||||
```
|
||||
wget https://raw.githubusercontent.com/openai/grade-school-math/master/grade_school_math/data/test.jsonl
|
||||
```
|
||||
|
||||
## Run benchmark
|
||||
|
||||
### Benchmark sglang
|
||||
```
|
||||
python -m sglang.launch_server --model-path meta-llama/Llama-2-7b-chat-hf --port 30000 --schedule-conservativeness 1.3
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_sglang.py --num-questions 64
|
||||
python3 bench_sglang.py --num-questions 32 --parallel 1
|
||||
```
|
||||
|
||||
|
||||
### Benchmark vllm
|
||||
```
|
||||
python3 -m vllm.entrypoints.api_server --tokenizer-mode auto --model meta-llama/Llama-2-7b-chat-hf --disable-log-requests --port 21000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --num-questions 64 --backend vllm
|
||||
```
|
||||
|
||||
|
||||
### Benchmark lightllm
|
||||
```
|
||||
# A10G
|
||||
python -m lightllm.server.api_server --tokenizer_mode auto --model_dir ~/model_weights/llama-2-7b-chat-hf --max_total_token_num 16000 --port 22000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --num-questions 64 --backend lightllm
|
||||
```
|
||||
|
||||
|
||||
### Benchmark guidance
|
||||
```
|
||||
python3 bench_other.py --num-questions 8 --backend guidance --parallel 1 --n-ctx 4096 --model-path path/to/gguf
|
||||
```
|
||||
|
||||
### Benchmark lmql
|
||||
|
||||
```
|
||||
python3 bench_other.py --num-questions 64 --backend lmql --parallel 1
|
||||
```
|
||||
@@ -1,186 +0,0 @@
|
||||
import argparse
|
||||
import ast
|
||||
import asyncio
|
||||
import json
|
||||
import re
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
|
||||
from sglang.test.test_utils import add_common_other_args_and_parse, get_call_generate
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
INVALID = -9999999
|
||||
|
||||
|
||||
def get_answer_value(answer_str):
|
||||
answer_str = answer_str.replace(",", "")
|
||||
numbers = re.findall(r"\d+", answer_str)
|
||||
if len(numbers) < 1:
|
||||
return INVALID
|
||||
try:
|
||||
return ast.literal_eval(numbers[-1])
|
||||
except SyntaxError:
|
||||
return INVALID
|
||||
|
||||
|
||||
prompt_lib = [
|
||||
"Let us think step by step.",
|
||||
"Approach this methodically. Let's dissect the problem into smaller, more manageable parts.",
|
||||
"It's important to proceed step by step, ensuring accuracy at each stage.",
|
||||
"Take a deep breath and break this down.",
|
||||
"A little bit of arithmetic and a logical approach will help us quickly arrive at the solution to this problem.",
|
||||
"I am extremely good at math.",
|
||||
]
|
||||
|
||||
|
||||
def multi_chain_gsm8k(question, num_chains, call_generate):
|
||||
s = "Question: " + question + "\n"
|
||||
# s += call_generate(s + "Answer: " + prompt_lib[0], max_tokens=256,
|
||||
# stop="Question", temperature=0)
|
||||
# return s
|
||||
|
||||
comps = []
|
||||
for i in range(num_chains):
|
||||
comps.append(
|
||||
call_generate(
|
||||
s + "Answer: " + prompt_lib[i % num_chains],
|
||||
max_tokens=256,
|
||||
temperature=0.3,
|
||||
stop="Question",
|
||||
)
|
||||
)
|
||||
|
||||
s += "Answer: To answer this question, here are some possible solutions. "
|
||||
s += "After considering all of them, I will do a majority vote.\n\n"
|
||||
for i in range(num_chains):
|
||||
s += f"Solution {i+1}: " + comps[i].strip() + "\n\n"
|
||||
s += "\nBy considering the above solutions and doing a majority vote, I think the final answer (a single integer number) is "
|
||||
s += call_generate(s, max_tokens=16, temperature=0, stop=None)
|
||||
return s
|
||||
|
||||
|
||||
async def multi_chain_gsm8k_async(question, num_chains, call_generate):
|
||||
s = "Question: " + question + "\n"
|
||||
# s += call_generate(s + "Answer: " + prompt_lib[0], max_tokens=256,
|
||||
# stop="Question", temperature=0)
|
||||
# return s
|
||||
|
||||
comps = []
|
||||
for i in range(num_chains):
|
||||
comps.append(
|
||||
await call_generate(
|
||||
s + "Answer: " + prompt_lib[i % num_chains],
|
||||
max_tokens=256,
|
||||
temperature=0.3,
|
||||
stop="Question",
|
||||
)
|
||||
)
|
||||
|
||||
s += "Answer: To answer this question, here are some possible solutions. "
|
||||
s += "After considering all of them, I will do a majority vote.\n\n"
|
||||
for i in range(num_chains):
|
||||
s += f"Solution {i+1}: " + comps[i].strip() + "\n\n"
|
||||
s += "\nBy considering the above solutions and doing a majority vote, I think the final answer (a single integer number) is "
|
||||
s += await call_generate(s, max_tokens=16, temperature=0, stop=None)
|
||||
return s
|
||||
|
||||
|
||||
def main(args):
|
||||
lines = list(read_jsonl(args.data_path))
|
||||
|
||||
# Construct prompts
|
||||
k = args.num_shot
|
||||
|
||||
questions = []
|
||||
labels = []
|
||||
for i in range(len(lines[: args.num_questions])):
|
||||
questions.append(lines[i]["question"])
|
||||
labels.append(get_answer_value(lines[i]["answer"]))
|
||||
assert all(l != INVALID for l in labels)
|
||||
|
||||
states = [None] * len(labels)
|
||||
|
||||
# Select backend
|
||||
call_generate = get_call_generate(args)
|
||||
|
||||
# Run requests
|
||||
if args.backend != "lmql":
|
||||
# Use thread pool
|
||||
def get_one_answer(i):
|
||||
answer = multi_chain_gsm8k(questions[i], args.num_chains, call_generate)
|
||||
states[i] = answer
|
||||
|
||||
tic = time.perf_counter()
|
||||
if args.parallel == 1:
|
||||
for i in tqdm(range(len(questions))):
|
||||
get_one_answer(i)
|
||||
else:
|
||||
with ThreadPoolExecutor(args.parallel) as executor:
|
||||
list(
|
||||
tqdm(
|
||||
executor.map(get_one_answer, list(range(len(questions)))),
|
||||
total=len(questions),
|
||||
)
|
||||
)
|
||||
|
||||
else:
|
||||
# Use asyncio
|
||||
async def get_one_answer_asyncio(i):
|
||||
answer = await multi_chain_gsm8k_async(
|
||||
questions[i], args.num_chains, call_generate
|
||||
)
|
||||
states[i] = answer
|
||||
|
||||
tic = time.perf_counter()
|
||||
loop = asyncio.get_event_loop()
|
||||
batches = [
|
||||
list(range(i, min(i + args.parallel, len(questions))))
|
||||
for i in range(0, len(questions), args.parallel)
|
||||
]
|
||||
for bt in tqdm(batches):
|
||||
tasks = [get_one_answer_asyncio(k) for k in bt]
|
||||
loop.run_until_complete(asyncio.gather(*tasks))
|
||||
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
preds = []
|
||||
for i in range(len(states)):
|
||||
preds.append(get_answer_value(states[i]))
|
||||
|
||||
# Compute accuracy
|
||||
acc = np.mean(np.array(preds) == np.array(labels))
|
||||
invalid = np.mean(np.array(preds) == INVALID)
|
||||
print(f"Latency: {latency:.3f}")
|
||||
print(f"Invalid: {invalid:.3f}")
|
||||
print(f"Accuracy: {acc:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "multi_chain_gsm8k",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"accuracy": round(acc, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"other": {
|
||||
"num_questions": args.num_questions,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--num-shot", type=int, default=0)
|
||||
parser.add_argument("--num-chains", type=int, default=5)
|
||||
parser.add_argument("--data-path", type=str, default="test.jsonl")
|
||||
parser.add_argument("--num-questions", type=int, default=50)
|
||||
args = add_common_other_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,140 +0,0 @@
|
||||
import argparse
|
||||
import ast
|
||||
import json
|
||||
import re
|
||||
import time
|
||||
|
||||
import numpy as np
|
||||
|
||||
from sglang.test.test_utils import (
|
||||
add_common_sglang_args_and_parse,
|
||||
select_sglang_backend,
|
||||
)
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
INVALID = -9999999
|
||||
|
||||
|
||||
def get_answer_value(answer_str):
|
||||
answer_str = answer_str.replace(",", "")
|
||||
numbers = re.findall(r"\d+", answer_str)
|
||||
if len(numbers) < 1:
|
||||
return INVALID
|
||||
try:
|
||||
return ast.literal_eval(numbers[-1])
|
||||
except SyntaxError:
|
||||
return INVALID
|
||||
|
||||
|
||||
prompt_lib = [
|
||||
"Let us think step by step.",
|
||||
"Approach this methodically. Let's dissect the problem into smaller, more manageable parts.",
|
||||
"It's important to proceed step by step, ensuring accuracy at each stage.",
|
||||
"Take a deep breath and break this down.",
|
||||
"A little bit of arithmetic and a logical approach will help us quickly arrive at the solution to this problem.",
|
||||
"I am extremely good at math.",
|
||||
]
|
||||
|
||||
|
||||
def main(args):
|
||||
lines = list(read_jsonl(args.data_path))
|
||||
|
||||
# Construct prompts
|
||||
# k = args.num_shot
|
||||
# few_shot_examples = get_few_shot_examples(lines, k)
|
||||
|
||||
questions = []
|
||||
labels = []
|
||||
for i in range(len(lines[: args.num_questions])):
|
||||
questions.append(lines[i]["question"])
|
||||
labels.append(get_answer_value(lines[i]["answer"]))
|
||||
assert all(l != INVALID for l in labels)
|
||||
arguments = [{"question": q} for q in questions]
|
||||
|
||||
num_chains = args.num_chains
|
||||
|
||||
#####################################
|
||||
######### SGL Program Begin #########
|
||||
#####################################
|
||||
|
||||
import sglang as sgl
|
||||
|
||||
@sgl.function
|
||||
def multi_chain_gsm8k(s, question):
|
||||
s += "Question: " + question + "\n"
|
||||
# s += "Answer: " + prompt_lib[0] + sgl.gen("answer", max_tokens=256, stop="Question",
|
||||
# temperature=0)
|
||||
# return
|
||||
|
||||
forks = s.fork(num_chains)
|
||||
for i in range(num_chains):
|
||||
forks[i] += (
|
||||
"Answer: "
|
||||
+ prompt_lib[i % num_chains]
|
||||
+ sgl.gen("chain", max_tokens=256, temperature=0.3, stop="Question")
|
||||
)
|
||||
forks.join()
|
||||
|
||||
s += "Answer: To answer this question, here are some possible solutions. "
|
||||
s += "After considering all of them, I will do a majority vote.\n\n"
|
||||
for i in range(num_chains):
|
||||
s += f"Solution {i+1}: " + forks[i]["chain"].strip() + "\n\n"
|
||||
s += "\nBy considering the above solutions and doing a majority vote, I think the final answer (a single integer number) is "
|
||||
s += sgl.gen("answer", max_tokens=16)
|
||||
|
||||
#####################################
|
||||
########## SGL Program End ##########
|
||||
#####################################
|
||||
|
||||
# Select backend
|
||||
backend = select_sglang_backend(args)
|
||||
|
||||
# Run requests
|
||||
tic = time.perf_counter()
|
||||
states = multi_chain_gsm8k.run_batch(
|
||||
arguments,
|
||||
temperature=0,
|
||||
backend=backend,
|
||||
num_threads=args.parallel,
|
||||
progress_bar=True,
|
||||
)
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
preds = []
|
||||
for i in range(len(states)):
|
||||
preds.append(get_answer_value(states[i]["answer"]))
|
||||
|
||||
# Compute accuracy
|
||||
acc = np.mean(np.array(preds) == np.array(labels))
|
||||
invalid = np.mean(np.array(preds) == INVALID)
|
||||
print(f"Latency: {latency:.3f}")
|
||||
print(f"Invalid: {invalid:.3f}")
|
||||
print(f"Accuracy: {acc:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "multi_chain_gsm8k",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"accuracy": round(acc, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"other": {
|
||||
"num_questions": args.num_questions,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--num-shot", type=int, default=0)
|
||||
parser.add_argument("--num-chains", type=int, default=5)
|
||||
parser.add_argument("--data-path", type=str, default="test.jsonl")
|
||||
parser.add_argument("--num-questions", type=int, default=50)
|
||||
args = add_common_sglang_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,47 +0,0 @@
|
||||
## Run benchmark
|
||||
|
||||
### Benchmark sglang
|
||||
```
|
||||
python3 -m sglang.launch_server --model-path codellama/CodeLlama-7b-instruct-hf --port 30000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_sglang.py --num-questions 10 --parallel 1
|
||||
```
|
||||
|
||||
|
||||
### Benchmark vllm
|
||||
```
|
||||
python3 -m vllm.entrypoints.api_server --tokenizer-mode auto --model codellama/CodeLlama-7b-instruct-hf --disable-log-requests --port 21000 --gpu 0.97
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --backend vllm --num-questions 64
|
||||
```
|
||||
|
||||
|
||||
### Benchmark guidance
|
||||
```
|
||||
python3 bench_other.py --backend guidance --num-questions 32 --parallel 1 --n-ctx 11000 --model-path path/to/code-llama/gguf
|
||||
```
|
||||
|
||||
|
||||
|
||||
### Build dataset
|
||||
|
||||
```
|
||||
pip install PyPDF2
|
||||
python3 build_dataset.py
|
||||
```
|
||||
|
||||
```python
|
||||
import PyPDF2
|
||||
|
||||
with open('llama2.pdf', 'rb') as file:
|
||||
reader = PyPDF2.PdfReader(file)
|
||||
text = ''
|
||||
for page_num in range(len(reader.pages)):
|
||||
text += reader.pages[page_num].extract_text()
|
||||
with open('output.txt', 'w') as text_file:
|
||||
text_file.write(text)
|
||||
```
|
||||
@@ -1,114 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from functools import partial
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
from sglang.test.test_utils import add_common_other_args_and_parse, get_call_generate
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
USER_PREFIX = "[INST] "
|
||||
USER_SUFFIX = " [/INST]"
|
||||
ASSISTANT_PREFIX = ""
|
||||
ASSISTANT_SUFFIX = " </s><s>"
|
||||
|
||||
|
||||
def multi_document_qa(docs, question, generate):
|
||||
s = USER_PREFIX
|
||||
s += "Please answer a question according to given documents.\n"
|
||||
s += "Question:" + question + "Documents begin.\n"
|
||||
|
||||
s += "".join(docs)
|
||||
|
||||
s += "\nDocuments end."
|
||||
s += (
|
||||
"\n\nBased on the above documents, please answer this question:\n"
|
||||
+ question
|
||||
+ "\nAnswer in three words or fewer."
|
||||
)
|
||||
s += USER_SUFFIX
|
||||
s += ASSISTANT_PREFIX
|
||||
answer = generate(s, max_tokens=16, stop=None)
|
||||
return answer
|
||||
|
||||
|
||||
def main(args):
|
||||
lines = read_jsonl(args.data_path)
|
||||
l = lines[0]
|
||||
arguments = []
|
||||
labels = []
|
||||
|
||||
num_docs = 10
|
||||
if args.backend == "guidance":
|
||||
num_docs = 7 # due to OOM
|
||||
|
||||
for i in range(len(l["questions"][: args.num_questions])):
|
||||
arguments.append(
|
||||
{
|
||||
"docs": l["documents"][:num_docs],
|
||||
"question": l["questions"][i],
|
||||
}
|
||||
)
|
||||
labels.append(l["answers"][i])
|
||||
states = [None] * len(arguments)
|
||||
|
||||
# Select backend
|
||||
call_generate = partial(get_call_generate(args), temperature=0)
|
||||
|
||||
# Run requests
|
||||
def get_one_answer(i):
|
||||
states[i] = multi_document_qa(generate=call_generate, **arguments[i])
|
||||
|
||||
tic = time.perf_counter()
|
||||
if args.parallel == 1:
|
||||
for i in tqdm(range(len(labels))):
|
||||
get_one_answer(i)
|
||||
else:
|
||||
with ThreadPoolExecutor(args.parallel) as executor:
|
||||
list(
|
||||
tqdm(
|
||||
executor.map(get_one_answer, list(range(len(labels)))),
|
||||
total=len(labels),
|
||||
)
|
||||
)
|
||||
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
# Compute accuracy
|
||||
print(states)
|
||||
correct = 0
|
||||
for s, label in zip(states, labels):
|
||||
answer = s.lower()
|
||||
if all(x in answer for x in label.lower().split(" ")):
|
||||
correct += 1
|
||||
accuracy = correct / len(labels)
|
||||
print(f"Accuracy: {accuracy:.3f}")
|
||||
print(f"Latency: {latency:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "multi_document_qa",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"accuracy": accuracy,
|
||||
"other": {
|
||||
"num_questions": args.num_questions,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="questions.jsonl")
|
||||
parser.add_argument("--num-questions", type=int, default=100)
|
||||
args = add_common_other_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,93 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
|
||||
import sglang as sgl
|
||||
from sglang.test.test_utils import (
|
||||
add_common_sglang_args_and_parse,
|
||||
select_sglang_backend,
|
||||
)
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
|
||||
@sgl.function
|
||||
def multi_document_qa(s, docs, question):
|
||||
s += sgl.user_begin()
|
||||
s += "Please answer a question according to given documents.\n"
|
||||
s += "Question:" + question + "Documents begin.\n"
|
||||
|
||||
forks = s.fork(len(docs))
|
||||
forks += lambda i: docs[i]
|
||||
forks.join("concate_and_append")
|
||||
|
||||
s += "\nDocuments end."
|
||||
s += (
|
||||
"\n\nBased on the above documents, please answer this question:\n"
|
||||
+ question
|
||||
+ "\nAnswer in three words or fewer."
|
||||
)
|
||||
s += sgl.user_end()
|
||||
s += sgl.assistant(sgl.gen("answer", max_tokens=16))
|
||||
|
||||
|
||||
def main(args):
|
||||
lines = read_jsonl(args.data_path)
|
||||
l = lines[0]
|
||||
arguments = []
|
||||
labels = []
|
||||
for i in range(len(l["questions"][: args.num_questions])):
|
||||
arguments.append(
|
||||
{
|
||||
"docs": l["documents"][:10],
|
||||
"question": l["questions"][i],
|
||||
}
|
||||
)
|
||||
labels.append(l["answers"][i])
|
||||
|
||||
# Select backend
|
||||
backend = select_sglang_backend(args)
|
||||
sgl.set_default_backend(backend)
|
||||
|
||||
# Run requests
|
||||
tic = time.perf_counter()
|
||||
states = multi_document_qa.run_batch(
|
||||
arguments, temperature=0, num_threads=args.parallel, progress_bar=True
|
||||
)
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
# Compute accuracy
|
||||
print([s["answer"] for s in states])
|
||||
correct = 0
|
||||
for s, label in zip(states, labels):
|
||||
answer = s["answer"].lower()
|
||||
if all(x in answer for x in label.lower().split(" ")):
|
||||
correct += 1
|
||||
accuracy = correct / len(labels)
|
||||
print(f"Accuracy: {accuracy:.3f}")
|
||||
print(f"Latency: {latency:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "multi_document_qa",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"accuracy": accuracy,
|
||||
"other": {
|
||||
"num_questions": args.num_questions,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="questions.jsonl")
|
||||
parser.add_argument("--num-questions", type=int, default=100)
|
||||
args = add_common_sglang_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,70 +0,0 @@
|
||||
import json
|
||||
|
||||
import transformers
|
||||
|
||||
content = "\n".join(
|
||||
open("llama2.txt", "r", encoding="utf-8", errors="ignore").readlines()
|
||||
)
|
||||
content = content.replace("\n\n", "\n")
|
||||
|
||||
# Count token
|
||||
name = "meta-llama/Llama-2-7b-chat-hf"
|
||||
t = transformers.AutoTokenizer.from_pretrained(name)
|
||||
print(f"num tokens: {len(t.encode(content))}")
|
||||
|
||||
# Segment
|
||||
SEP = "\n\n"
|
||||
parts = content.split(SEP)
|
||||
print(f"num segments: {len(parts)}")
|
||||
|
||||
segment_len = 1100
|
||||
|
||||
segments = []
|
||||
tmp = []
|
||||
tmp_len = 0
|
||||
for i in range(len(parts)):
|
||||
tmp.append(parts[i])
|
||||
tmp_len += len(t.encode(parts[i]))
|
||||
|
||||
if tmp_len > segment_len:
|
||||
segments.append(SEP.join(tmp))
|
||||
tmp = []
|
||||
tmp_len = 0
|
||||
|
||||
for i, s in enumerate(segments):
|
||||
print(i, len(t.encode(segments[i])))
|
||||
|
||||
# Dump
|
||||
with open("questions.jsonl", "w") as fout:
|
||||
fout.write(
|
||||
json.dumps(
|
||||
{
|
||||
"documents": segments[:30],
|
||||
"questions": [
|
||||
"What is the name of the fine-tuned LLMs?",
|
||||
"Which figure shows the helpfulness human evaluation results for Llama 2-Chat?",
|
||||
"What is the number of parameters in the largest Llama 2 model?",
|
||||
"What is the batch size of fine-tuning?",
|
||||
"Where can we find the details of potential data contamination?",
|
||||
"What is the full name of MPT?",
|
||||
"What is the power consumption of RSC in Watt?",
|
||||
"How many tokens of data do they train on?",
|
||||
"Which model's release is delayed due to a lack of time to sufficiently red team?",
|
||||
"Which activation function is used in Llama?",
|
||||
],
|
||||
"answers": [
|
||||
"Llama 2 Chat",
|
||||
"1",
|
||||
"70 B",
|
||||
"64",
|
||||
"A 6",
|
||||
"MosaicML",
|
||||
"400",
|
||||
"2 trillion",
|
||||
"34 B",
|
||||
"SwiGLU",
|
||||
],
|
||||
}
|
||||
)
|
||||
+ "\n"
|
||||
)
|
||||
@@ -1,66 +0,0 @@
|
||||
### Benchmark sglang
|
||||
|
||||
Run Llama-7B
|
||||
|
||||
```
|
||||
python3 -m sglang.launch_server --model-path meta-llama/Llama-2-7b-chat-hf --port 30000
|
||||
```
|
||||
|
||||
Run Mixtral-8x7B
|
||||
(When there is a CUDA out-of-memory error, try to reduce the `--mem-fraction-static`)
|
||||
|
||||
```
|
||||
python3 -m sglang.launch_server --model-path mistralai/Mixtral-8x7B-Instruct-v0.1 --port 30000 --tp-size 8
|
||||
```
|
||||
|
||||
Benchmark(short output)
|
||||
|
||||
```
|
||||
python3 bench_sglang.py --tokenizer meta-llama/Llama-2-7b-chat-hf
|
||||
```
|
||||
|
||||
Benchmark(long output)
|
||||
|
||||
```
|
||||
python3 bench_sglang.py --tokenizer meta-llama/Llama-2-7b-chat-hf --long
|
||||
```
|
||||
|
||||
### Benchmark vLLM
|
||||
|
||||
Run Llama-7B
|
||||
|
||||
```
|
||||
python3 -m vllm.entrypoints.api_server --tokenizer-mode auto --model meta-llama/Llama-2-7b-chat-hf --disable-log-requests --port 21000
|
||||
```
|
||||
|
||||
Run Mixtral-8x7B
|
||||
|
||||
```
|
||||
python3 -m vllm.entrypoints.api_server --tokenizer-mode auto --model mistralai/Mixtral-8x7B-Instruct-v0.1 --disable-log-requests --port 21000 --tensor-parallel-size 8
|
||||
```
|
||||
|
||||
Benchmark(short output)
|
||||
|
||||
```
|
||||
python3 bench_other.py --tokenizer meta-llama/Llama-2-7b-chat-hf --backend vllm
|
||||
```
|
||||
|
||||
Benchmark(long output)
|
||||
|
||||
```
|
||||
python3 bench_other.py --tokenizer meta-llama/Llama-2-7b-chat-hf --backend vllm --long
|
||||
```
|
||||
|
||||
### Benchmark guidance
|
||||
|
||||
Benchmark Llama-7B (short output)
|
||||
|
||||
```
|
||||
python3 bench_other.py --tokenizer meta-llama/Llama-2-7b-chat-hf --backend guidance --parallel 1 --n-ctx 4096 --model-path path/to/gguf
|
||||
```
|
||||
|
||||
Benchmark Llama-7B (long output)
|
||||
|
||||
```
|
||||
python3 bench_other.py --tokenizer meta-llama/Llama-2-7b-chat-hf --backend guidance --parallel 1 --n-ctx 4096 --model-path path/to/gguf --long
|
||||
```
|
||||
@@ -1,93 +0,0 @@
|
||||
import json
|
||||
import time
|
||||
from argparse import ArgumentParser
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from functools import partial
|
||||
|
||||
from data_gen import gen_arguments
|
||||
from tqdm import tqdm
|
||||
from vllm.transformers_utils.tokenizer import get_tokenizer
|
||||
|
||||
from sglang.test.test_utils import add_common_other_args_and_parse, get_call_generate
|
||||
from sglang.utils import dump_state_text
|
||||
|
||||
|
||||
def multi_turns(generate, qas):
|
||||
s = ""
|
||||
for qa in qas:
|
||||
s += qa["prompt"]
|
||||
s += generate(s, max_tokens=qa["new_tokens"])
|
||||
|
||||
return s
|
||||
|
||||
|
||||
def main(args):
|
||||
print(args)
|
||||
|
||||
tokenizer = get_tokenizer(args.tokenizer, trust_remote_code=args.trust_remote_code)
|
||||
|
||||
multi_qas = gen_arguments(args, tokenizer)
|
||||
|
||||
states = [None] * args.num_qa
|
||||
|
||||
call_generate = partial(get_call_generate(args), temperature=0)
|
||||
|
||||
def get_one_answer(i):
|
||||
states[i] = multi_turns(generate=call_generate, **multi_qas[i])
|
||||
|
||||
tic = time.perf_counter()
|
||||
if args.parallel == 1:
|
||||
for i in tqdm(range(len(multi_qas))):
|
||||
get_one_answer(i)
|
||||
else:
|
||||
with ThreadPoolExecutor(args.parallel) as executor:
|
||||
rets = list(
|
||||
tqdm(
|
||||
executor.map(get_one_answer, list(range(len(multi_qas)))),
|
||||
total=len(multi_qas),
|
||||
)
|
||||
)
|
||||
for _ in rets:
|
||||
pass
|
||||
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
# Compute accuracy
|
||||
print(f"Latency: {latency:.3f}")
|
||||
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "multi_turn_chat",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"num_requests": args.num_qa,
|
||||
"num_turns": args.turns,
|
||||
"other": {
|
||||
"parallel": args.parallel,
|
||||
"output_mode": "long" if args.long else "short",
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = ArgumentParser()
|
||||
parser.add_argument("--turns", type=int, default=4)
|
||||
parser.add_argument("--num-qa", type=int, default=20)
|
||||
parser.add_argument("--min-len-q", type=int, default=256)
|
||||
parser.add_argument("--max-len-q", type=int, default=512)
|
||||
parser.add_argument("--min-len-a", type=int, default=4)
|
||||
parser.add_argument("--max-len-a", type=int, default=8)
|
||||
parser.add_argument("--tokenizer", type=str, required=True)
|
||||
parser.add_argument("--trust-remote-code", action="store_true")
|
||||
parser.add_argument("--long", action="store_true")
|
||||
args = add_common_other_args_and_parse(parser)
|
||||
|
||||
if args.long:
|
||||
args.min_len_a = 256
|
||||
args.max_len_a = 512
|
||||
args.num_qa = 20
|
||||
main(args)
|
||||
@@ -1,79 +0,0 @@
|
||||
import json
|
||||
import time
|
||||
from argparse import ArgumentParser
|
||||
|
||||
from data_gen import gen_arguments
|
||||
from vllm.transformers_utils.tokenizer import get_tokenizer
|
||||
|
||||
import sglang as sgl
|
||||
from sglang.test.test_utils import (
|
||||
add_common_sglang_args_and_parse,
|
||||
select_sglang_backend,
|
||||
)
|
||||
from sglang.utils import dump_state_text
|
||||
|
||||
|
||||
@sgl.function
|
||||
def multi_turns(s, qas):
|
||||
for qa in qas:
|
||||
s += qa["prompt"]
|
||||
s += sgl.gen(max_tokens=qa["new_tokens"], ignore_eos=True)
|
||||
|
||||
|
||||
def main(args):
|
||||
tokenizer = get_tokenizer(args.tokenizer, trust_remote_code=args.trust_remote_code)
|
||||
|
||||
multi_qas = gen_arguments(args, tokenizer)
|
||||
|
||||
backend = select_sglang_backend(args)
|
||||
|
||||
tic = time.perf_counter()
|
||||
states = multi_turns.run_batch(
|
||||
multi_qas,
|
||||
temperature=0,
|
||||
backend=backend,
|
||||
num_threads=args.parallel,
|
||||
progress_bar=True,
|
||||
)
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
print(f"Latency: {latency:.3f}")
|
||||
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "multi_turn_chat",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"num_requests": args.num_qa,
|
||||
"num_turns": args.turns,
|
||||
"other": {
|
||||
"parallel": args.parallel,
|
||||
"output_mode": "long" if args.long else "short",
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = ArgumentParser()
|
||||
parser.add_argument("--turns", type=int, default=4)
|
||||
parser.add_argument("--num-qa", type=int, default=20)
|
||||
parser.add_argument("--min-len-q", type=int, default=256)
|
||||
parser.add_argument("--max-len-q", type=int, default=512)
|
||||
parser.add_argument("--min-len-a", type=int, default=4)
|
||||
parser.add_argument("--max-len-a", type=int, default=8)
|
||||
parser.add_argument("--tokenizer", type=str, required=True)
|
||||
parser.add_argument("--trust-remote-code", action="store_true")
|
||||
parser.add_argument("--long", action="store_true")
|
||||
args = add_common_sglang_args_and_parse(parser)
|
||||
|
||||
if args.long:
|
||||
args.min_len_a = 256
|
||||
args.max_len_a = 512
|
||||
args.num_qa = 20
|
||||
|
||||
print(args)
|
||||
main(args)
|
||||
@@ -1,29 +0,0 @@
|
||||
import random
|
||||
import string
|
||||
|
||||
random.seed(42)
|
||||
|
||||
|
||||
def gen_prompt(tokenizer, token_num):
|
||||
cha_set = string.ascii_letters + string.digits
|
||||
ret = "".join(random.choices(cha_set, k=token_num))
|
||||
while len(tokenizer(ret).input_ids) < token_num:
|
||||
ret += random.choice(cha_set)
|
||||
return ret
|
||||
|
||||
|
||||
def gen_arguments(args, tokenizer):
|
||||
multi_qas = [{"qas": []} for _ in range(args.num_qa)]
|
||||
for i in range(args.num_qa):
|
||||
qas = multi_qas[i]["qas"]
|
||||
for _ in range(args.turns):
|
||||
prompt_len = random.randint(args.min_len_q, args.max_len_q)
|
||||
new_tokens = random.randint(args.min_len_a, args.max_len_a)
|
||||
qas.append(
|
||||
{
|
||||
"prompt": gen_prompt(tokenizer, prompt_len),
|
||||
"new_tokens": new_tokens,
|
||||
}
|
||||
)
|
||||
|
||||
return multi_qas
|
||||
@@ -1,129 +0,0 @@
|
||||
import json
|
||||
import random
|
||||
import time
|
||||
from argparse import ArgumentParser
|
||||
from pathlib import Path
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
import sglang as sgl
|
||||
from sglang.srt.utils.hf_transformers_utils import get_tokenizer
|
||||
from sglang.test.test_utils import (
|
||||
add_common_sglang_args_and_parse,
|
||||
select_sglang_backend,
|
||||
)
|
||||
from sglang.utils import dump_state_text
|
||||
|
||||
|
||||
def gen_prompt(tokenizer, token_num):
|
||||
all_available_tokens = list(tokenizer.get_vocab().values())
|
||||
selected_tokens = random.choices(all_available_tokens, k=token_num)
|
||||
ret = tokenizer.decode(selected_tokens)
|
||||
return ret
|
||||
|
||||
|
||||
def get_cache_path(args):
|
||||
# Create cache directory under ~/.cache/sglang
|
||||
cache_dir = Path.home() / ".cache" / "sglang"
|
||||
|
||||
# Create a unique cache filename based on the arguments that affect generation
|
||||
cache_key = f"qa_{args.num_qa}_{args.turns}_{args.system_prompt_len}_{args.len_q}_{args.len_a}_{args.tokenizer.replace('/', '_')}.json"
|
||||
return cache_dir / cache_key
|
||||
|
||||
|
||||
def gen_arguments(args, tokenizer):
|
||||
cache_path = get_cache_path(args)
|
||||
|
||||
# Try to load from cache first
|
||||
if cache_path.exists():
|
||||
print(f"Loading cached arguments from {cache_path}")
|
||||
with open(cache_path, "r") as f:
|
||||
return json.load(f)
|
||||
|
||||
print("Generating new arguments...")
|
||||
# First progress bar for system prompts
|
||||
multi_qas = []
|
||||
for _ in tqdm(range(args.num_qa), desc="Generating system prompts"):
|
||||
multi_qas.append(
|
||||
{"system_prompt": gen_prompt(tokenizer, args.system_prompt_len), "qas": []}
|
||||
)
|
||||
|
||||
# Nested progress bars for QA pairs
|
||||
for i in tqdm(range(args.num_qa), desc="Generating QA pairs"):
|
||||
qas = multi_qas[i]["qas"]
|
||||
for j in range(args.turns):
|
||||
qas.append(
|
||||
{
|
||||
"prompt": gen_prompt(tokenizer, args.len_q),
|
||||
"new_tokens": args.len_a,
|
||||
}
|
||||
)
|
||||
|
||||
# Save to cache
|
||||
cache_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with open(cache_path, "w") as f:
|
||||
json.dump(multi_qas, f)
|
||||
print(f"Cached arguments saved to {cache_path}")
|
||||
|
||||
return multi_qas
|
||||
|
||||
|
||||
@sgl.function
|
||||
def multi_turns(s, system_prompt, qas):
|
||||
s += system_prompt
|
||||
|
||||
for i, qa in enumerate(qas):
|
||||
s += qa["prompt"]
|
||||
s += sgl.gen(max_tokens=qa["new_tokens"], ignore_eos=True)
|
||||
|
||||
|
||||
def main(args):
|
||||
tokenizer = get_tokenizer(args.tokenizer, trust_remote_code=args.trust_remote_code)
|
||||
|
||||
multi_qas = gen_arguments(args, tokenizer)
|
||||
|
||||
backend = select_sglang_backend(args)
|
||||
|
||||
tic = time.perf_counter()
|
||||
states = multi_turns.run_batch(
|
||||
multi_qas,
|
||||
temperature=0,
|
||||
backend=backend,
|
||||
num_threads="auto",
|
||||
progress_bar=True,
|
||||
)
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
print(f"Latency: {latency:.3f}")
|
||||
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "multi_turn_system_prompt_chat",
|
||||
"backend": args.backend,
|
||||
"latency": round(latency, 3),
|
||||
"num_requests": args.num_qa,
|
||||
"num_turns": args.turns,
|
||||
"other": {
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = ArgumentParser()
|
||||
parser.add_argument("--turns", type=int, default=8)
|
||||
parser.add_argument("--num-qa", type=int, default=128)
|
||||
parser.add_argument("--system-prompt-len", type=int, default=2048)
|
||||
parser.add_argument("--len-q", type=int, default=32)
|
||||
parser.add_argument("--len-a", type=int, default=128)
|
||||
parser.add_argument(
|
||||
"--tokenizer", type=str, default="meta-llama/Meta-Llama-3-8B-Instruct"
|
||||
)
|
||||
parser.add_argument("--trust-remote-code", action="store_true")
|
||||
args = add_common_sglang_args_and_parse(parser)
|
||||
|
||||
print(args)
|
||||
main(args)
|
||||
@@ -1,34 +0,0 @@
|
||||
## Run benchmark
|
||||
|
||||
NOTE: This is an implementation for replaying a given trace for throughput/latency benchmark purposes. It is not an actual ReAct agent implementation.
|
||||
|
||||
### Benchmark sglang
|
||||
```
|
||||
python -m sglang.launch_server --model-path meta-llama/Llama-2-7b-chat-hf --port 30000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_sglang.py --num-questions 100
|
||||
```
|
||||
|
||||
|
||||
### Benchmark vllm
|
||||
```
|
||||
python3 -m vllm.entrypoints.api_server --tokenizer-mode auto --model meta-llama/Llama-2-7b-chat-hf --disable-log-requests --port 21000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --num-questions 100 --backend vllm
|
||||
```
|
||||
|
||||
|
||||
### Benchmark guidance
|
||||
```
|
||||
python3 bench_other.py --num-questions 100 --backend guidance --parallel 1 --n-ctx 4096 --model-path path/to/gguf
|
||||
```
|
||||
|
||||
### Benchmark lmql
|
||||
|
||||
```
|
||||
python3 bench_other.py --num-questions 100 --backend lmql --parallel 1
|
||||
```
|
||||
@@ -1,202 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
from sglang.test.test_utils import add_common_other_args_and_parse, get_call_generate
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
|
||||
def get_prompt(question):
|
||||
prompt = (
|
||||
"""Solve a question answering task with interleaving Thought, Action, Observation steps. Thought can reason about the current situation, and Action can be three types:
|
||||
(1) Search[entity], which searches the exact entity on Wikipedia and returns the first paragraph if it exists. If not, it will return some similar entities to search.
|
||||
(2) Lookup[keyword], which returns the next sentence containing keyword in the current passage.
|
||||
(3) Finish[answer], which returns the answer and finishes the task.
|
||||
Here are some examples.
|
||||
Question: What is the elevation range for the area that the eastern sector of the Colorado orogeny extends into?
|
||||
Thought 1: I need to search Colorado orogeny, find the area that the eastern sector of the Colorado orogeny extends into, then find the elevation range of the area.
|
||||
Action 1: Search[Colorado orogeny]
|
||||
Observation 1: The Colorado orogeny was an episode of mountain building (an orogeny) in Colorado and surrounding areas.
|
||||
Thought 2: It does not mention the eastern sector. So I need to look up eastern sector.
|
||||
Action 2: Lookup[eastern sector]
|
||||
Observation 2: (Result 1 / 1) The eastern sector extends into the High Plains and is called the Central Plains orogeny.
|
||||
Thought 3: The eastern sector of Colorado orogeny extends into the High Plains. So I need to search High Plains and find its elevation range.
|
||||
Action 3: Search[High Plains]
|
||||
Observation 3: High Plains refers to one of two distinct land regions:
|
||||
Thought 4: I need to instead search High Plains (United States).
|
||||
Action 4: Search[High Plains (United States)]
|
||||
Observation 4: The High Plains are a subregion of the Great Plains. From east to west, the High Plains rise in elevation from around 1,800 to 7,000 ft (550 to 2,130 m).[3]
|
||||
Thought 5: High Plains rise in elevation from around 1,800 to 7,000 ft, so the answer is 1,800 to 7,000 ft.
|
||||
Action 5: Finish[1,800 to 7,000 ft]
|
||||
Question: Musician and satirist Allie Goertz wrote a song about the "The Simpsons" character Milhouse, who Matt Groening named after who?
|
||||
Thought 1: The question simplifies to "The Simpsons" character Milhouse is named after who. I only need to search Milhouse and find who it is named after.
|
||||
Action 1: Search[Milhouse]
|
||||
Observation 1: Milhouse Mussolini Van Houten is a recurring character in the Fox animated television series The Simpsons voiced by Pamela Hayden and created by Matt Groening.
|
||||
Thought 2: The paragraph does not tell who Milhouse is named after, maybe I can look up "named after".
|
||||
Action 2: Lookup[named after]
|
||||
Observation 2: (Result 1 / 1) Milhouse was named after U.S. president Richard Nixon, whose middle name was Milhous.
|
||||
Thought 3: Milhouse was named after U.S. president Richard Nixon, so the answer is Richard Nixon.
|
||||
Action 3: Finish[Richard Nixon]
|
||||
Question: Which documentary is about Finnish rock groups, Adam Clayton Powell or The Saimaa Gesture?
|
||||
Thought 1: I need to search Adam Clayton Powell and The Saimaa Gesture, and find which documentary is about Finnish rock groups.
|
||||
Action 1: Search[Adam Clayton Powell]
|
||||
Observation 1: Could not find [Adam Clayton Powell]. Similar: ['Adam Clayton Powell III', 'Seventh Avenue (Manhattan)', 'Adam Clayton Powell Jr. State Office Building', 'Isabel Washington Powell', 'Adam Powell', 'Adam Clayton Powell (film)', 'Giancarlo Esposito'].
|
||||
Thought 2: To find the documentary, I can search Adam Clayton Powell (film).
|
||||
Action 2: Search[Adam Clayton Powell (film)]
|
||||
Observation 2: Adam Clayton Powell is a 1989 American documentary film directed by Richard Kilberg.
|
||||
The film is about the rise and fall of influential African-American politician Adam Clayton Powell Jr.[3][4] It was later aired as part of the PBS series The American Experience.
|
||||
Thought 3: Adam Clayton Powell (film) is a documentary about an African-American politician, not Finnish rock groups. So the documentary about Finnish rock groups must instead be The Saimaa Gesture.
|
||||
Action 3: Finish[The Saimaa Gesture]
|
||||
Question: What profession does Nicholas Ray and Elia Kazan have in common?
|
||||
Thought 1: I need to search Nicholas Ray and Elia Kazan, find their professions, then find the profession they have in common.
|
||||
Action 1: Search[Nicholas Ray]
|
||||
Observation 1: Nicholas Ray (born Raymond Nicholas Kienzle Jr., August 7, 1911 – June 16, 1979) was an American film director, screenwriter, and actor best known for the 1955 film Rebel Without a Cause.
|
||||
Thought 2: Professions of Nicholas Ray are director, screenwriter, and actor. I need to search Elia Kazan next and find his professions.
|
||||
Action 2: Search[Elia Kazan]
|
||||
Observation 2: Elia Kazan was an American film and theatre director, producer, screenwriter and actor.
|
||||
Thought 3: Professions of Elia Kazan are director, producer, screenwriter, and actor. So profession Nicholas Ray and Elia Kazan have in common is director, screenwriter, and actor.
|
||||
Action 3: Finish[director, screenwriter, actor]
|
||||
Question: Which magazine was started first Arthur's Magazine or First for Women?
|
||||
Thought 1: I need to search Arthur's Magazine and First for Women, and find which was started first.
|
||||
Action 1: Search[Arthur's Magazine]
|
||||
Observation 1: Arthur's Magazine (1844-1846) was an American literary periodical published in Philadelphia in the 19th century.
|
||||
Thought 2: Arthur's Magazine was started in 1844. I need to search First for Women next.
|
||||
Action 2: Search[First for Women]
|
||||
Observation 2: First for Women is a woman's magazine published by Bauer Media Group in the USA.[1] The magazine was started in 1989.
|
||||
Thought 3: First for Women was started in 1989. 1844 (Arthur's Magazine) < 1989 (First for Women), so Arthur's Magazine was started first.
|
||||
Action 3: Finish[Arthur's Magazine]
|
||||
Question: Were Pavel Urysohn and Leonid Levin known for the same type of work?
|
||||
Thought 1: I need to search Pavel Urysohn and Leonid Levin, find their types of work, then find if they are the same.
|
||||
Action 1: Search[Pavel Urysohn]
|
||||
Observation 1: Pavel Samuilovich Urysohn (February 3, 1898 â August 17, 1924) was a Soviet mathematician who is best known for his contributions in dimension theory.
|
||||
Thought 2: Pavel Urysohn is a mathematician. I need to search Leonid Levin next and find its type of work.
|
||||
Action 2: Search[Leonid Levin]
|
||||
Observation 2: Leonid Anatolievich Levin is a Soviet-American mathematician and computer scientist.
|
||||
Thought 3: Leonid Levin is a mathematician and computer scientist. So Pavel Urysohn and Leonid Levin have the same type of work.
|
||||
Action 3: Finish[yes]
|
||||
"""
|
||||
+ question
|
||||
)
|
||||
return prompt
|
||||
|
||||
|
||||
def main(args):
|
||||
lines = read_jsonl(args.data_path)[: args.num_questions]
|
||||
arguments = [{"question": k, "triplets": v} for l in lines for k, v in l.items()]
|
||||
|
||||
states = []
|
||||
|
||||
# Select backend
|
||||
call_generate = get_call_generate(args)
|
||||
|
||||
def run_single_agent(argument):
|
||||
question = argument["question"]
|
||||
triplets = argument["triplets"]
|
||||
prompt = get_prompt(question)
|
||||
for i in range(1, len(triplets) + 2):
|
||||
prompt += "Thought " + str(i) + ":"
|
||||
states.append(prompt)
|
||||
answer = call_generate(
|
||||
prompt, max_tokens=200, temperature=0, stop="Observation"
|
||||
)
|
||||
if i > len(triplets):
|
||||
break
|
||||
prompt += (
|
||||
triplets[i - 1]["thought"]
|
||||
+ "\nAction "
|
||||
+ str(i)
|
||||
+ ":"
|
||||
+ triplets[i - 1]["action"]
|
||||
+ "\nObservation "
|
||||
+ str(i)
|
||||
+ ":"
|
||||
+ triplets[i - 1]["observation"]
|
||||
+ "\n"
|
||||
)
|
||||
|
||||
states.append(answer)
|
||||
|
||||
async def run_single_agent_async(argument):
|
||||
question = argument["question"]
|
||||
triplets = argument["triplets"]
|
||||
prompt = get_prompt(question)
|
||||
for i in range(1, len(triplets) + 2):
|
||||
prompt += "Thought " + str(i) + ":"
|
||||
states.append(prompt)
|
||||
answer = await call_generate(
|
||||
prompt, max_tokens=200, temperature=0, stop="Observation", max_len=4096
|
||||
)
|
||||
if i > len(triplets):
|
||||
break
|
||||
prompt += (
|
||||
triplets[i - 1]["thought"]
|
||||
+ "\nAction "
|
||||
+ str(i)
|
||||
+ ":"
|
||||
+ triplets[i - 1]["action"]
|
||||
+ "\nObservation "
|
||||
+ str(i)
|
||||
+ ":"
|
||||
+ triplets[i - 1]["observation"]
|
||||
+ "\n"
|
||||
)
|
||||
|
||||
states.append(answer)
|
||||
|
||||
tic = time.perf_counter()
|
||||
|
||||
if args.backend != "lmql":
|
||||
if args.parallel == 1:
|
||||
for arg in tqdm(arguments):
|
||||
run_single_agent(arg)
|
||||
else:
|
||||
with ThreadPoolExecutor(args.parallel) as executor:
|
||||
list(
|
||||
tqdm(
|
||||
executor.map(run_single_agent, arguments), total=len(arguments)
|
||||
)
|
||||
)
|
||||
|
||||
else:
|
||||
import asyncio
|
||||
|
||||
loop = asyncio.get_event_loop()
|
||||
batches = [
|
||||
[] for _ in range((len(arguments) + args.parallel - 1) // args.parallel)
|
||||
]
|
||||
for i, arg in enumerate(arguments):
|
||||
batches[i // args.parallel].append(arg)
|
||||
for bt in tqdm(batches):
|
||||
tasks = [run_single_agent_async(arg) for arg in bt]
|
||||
loop.run_until_complete(asyncio.gather(*tasks))
|
||||
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
print(f"Latency: {latency:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "ReAct Agents",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"num_requests": len(arguments),
|
||||
"other": {
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="hotpotqa_100.jsonl")
|
||||
parser.add_argument("--num-questions", type=int, default=10)
|
||||
args = add_common_other_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,153 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
|
||||
import sglang as sgl
|
||||
from sglang.test.test_utils import (
|
||||
add_common_sglang_args_and_parse,
|
||||
select_sglang_backend,
|
||||
)
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
|
||||
@sgl.function
|
||||
def webthink(s, question, triplets):
|
||||
s += (
|
||||
"""Solve a question answering task with interleaving Thought, Action, Observation steps. Thought can reason about the current situation, and Action can be three types:
|
||||
(1) Search[entity], which searches the exact entity on Wikipedia and returns the first paragraph if it exists. If not, it will return some similar entities to search.
|
||||
(2) Lookup[keyword], which returns the next sentence containing keyword in the current passage.
|
||||
(3) Finish[answer], which returns the answer and finishes the task.
|
||||
Here are some examples.
|
||||
Question: What is the elevation range for the area that the eastern sector of the Colorado orogeny extends into?
|
||||
Thought 1: I need to search Colorado orogeny, find the area that the eastern sector of the Colorado orogeny extends into, then find the elevation range of the area.
|
||||
Action 1: Search[Colorado orogeny]
|
||||
Observation 1: The Colorado orogeny was an episode of mountain building (an orogeny) in Colorado and surrounding areas.
|
||||
Thought 2: It does not mention the eastern sector. So I need to look up eastern sector.
|
||||
Action 2: Lookup[eastern sector]
|
||||
Observation 2: (Result 1 / 1) The eastern sector extends into the High Plains and is called the Central Plains orogeny.
|
||||
Thought 3: The eastern sector of Colorado orogeny extends into the High Plains. So I need to search High Plains and find its elevation range.
|
||||
Action 3: Search[High Plains]
|
||||
Observation 3: High Plains refers to one of two distinct land regions:
|
||||
Thought 4: I need to instead search High Plains (United States).
|
||||
Action 4: Search[High Plains (United States)]
|
||||
Observation 4: The High Plains are a subregion of the Great Plains. From east to west, the High Plains rise in elevation from around 1,800 to 7,000 ft (550 to 2,130 m).[3]
|
||||
Thought 5: High Plains rise in elevation from around 1,800 to 7,000 ft, so the answer is 1,800 to 7,000 ft.
|
||||
Action 5: Finish[1,800 to 7,000 ft]
|
||||
Question: Musician and satirist Allie Goertz wrote a song about the "The Simpsons" character Milhouse, who Matt Groening named after who?
|
||||
Thought 1: The question simplifies to "The Simpsons" character Milhouse is named after who. I only need to search Milhouse and find who it is named after.
|
||||
Action 1: Search[Milhouse]
|
||||
Observation 1: Milhouse Mussolini Van Houten is a recurring character in the Fox animated television series The Simpsons voiced by Pamela Hayden and created by Matt Groening.
|
||||
Thought 2: The paragraph does not tell who Milhouse is named after, maybe I can look up "named after".
|
||||
Action 2: Lookup[named after]
|
||||
Observation 2: (Result 1 / 1) Milhouse was named after U.S. president Richard Nixon, whose middle name was Milhous.
|
||||
Thought 3: Milhouse was named after U.S. president Richard Nixon, so the answer is Richard Nixon.
|
||||
Action 3: Finish[Richard Nixon]
|
||||
Question: Which documentary is about Finnish rock groups, Adam Clayton Powell or The Saimaa Gesture?
|
||||
Thought 1: I need to search Adam Clayton Powell and The Saimaa Gesture, and find which documentary is about Finnish rock groups.
|
||||
Action 1: Search[Adam Clayton Powell]
|
||||
Observation 1: Could not find [Adam Clayton Powell]. Similar: ['Adam Clayton Powell III', 'Seventh Avenue (Manhattan)', 'Adam Clayton Powell Jr. State Office Building', 'Isabel Washington Powell', 'Adam Powell', 'Adam Clayton Powell (film)', 'Giancarlo Esposito'].
|
||||
Thought 2: To find the documentary, I can search Adam Clayton Powell (film).
|
||||
Action 2: Search[Adam Clayton Powell (film)]
|
||||
Observation 2: Adam Clayton Powell is a 1989 American documentary film directed by Richard Kilberg.
|
||||
The film is about the rise and fall of influential African-American politician Adam Clayton Powell Jr.[3][4] It was later aired as part of the PBS series The American Experience.
|
||||
Thought 3: Adam Clayton Powell (film) is a documentary about an African-American politician, not Finnish rock groups. So the documentary about Finnish rock groups must instead be The Saimaa Gesture.
|
||||
Action 3: Finish[The Saimaa Gesture]
|
||||
Question: What profession does Nicholas Ray and Elia Kazan have in common?
|
||||
Thought 1: I need to search Nicholas Ray and Elia Kazan, find their professions, then find the profession they have in common.
|
||||
Action 1: Search[Nicholas Ray]
|
||||
Observation 1: Nicholas Ray (born Raymond Nicholas Kienzle Jr., August 7, 1911 – June 16, 1979) was an American film director, screenwriter, and actor best known for the 1955 film Rebel Without a Cause.
|
||||
Thought 2: Professions of Nicholas Ray are director, screenwriter, and actor. I need to search Elia Kazan next and find his professions.
|
||||
Action 2: Search[Elia Kazan]
|
||||
Observation 2: Elia Kazan was an American film and theatre director, producer, screenwriter and actor.
|
||||
Thought 3: Professions of Elia Kazan are director, producer, screenwriter, and actor. So profession Nicholas Ray and Elia Kazan have in common is director, screenwriter, and actor.
|
||||
Action 3: Finish[director, screenwriter, actor]
|
||||
Question: Which magazine was started first Arthur's Magazine or First for Women?
|
||||
Thought 1: I need to search Arthur's Magazine and First for Women, and find which was started first.
|
||||
Action 1: Search[Arthur's Magazine]
|
||||
Observation 1: Arthur's Magazine (1844-1846) was an American literary periodical published in Philadelphia in the 19th century.
|
||||
Thought 2: Arthur's Magazine was started in 1844. I need to search First for Women next.
|
||||
Action 2: Search[First for Women]
|
||||
Observation 2: First for Women is a woman's magazine published by Bauer Media Group in the USA.[1] The magazine was started in 1989.
|
||||
Thought 3: First for Women was started in 1989. 1844 (Arthur's Magazine) < 1989 (First for Women), so Arthur's Magazine was started first.
|
||||
Action 3: Finish[Arthur's Magazine]
|
||||
Question: Were Pavel Urysohn and Leonid Levin known for the same type of work?
|
||||
Thought 1: I need to search Pavel Urysohn and Leonid Levin, find their types of work, then find if they are the same.
|
||||
Action 1: Search[Pavel Urysohn]
|
||||
Observation 1: Pavel Samuilovich Urysohn (February 3, 1898 â August 17, 1924) was a Soviet mathematician who is best known for his contributions in dimension theory.
|
||||
Thought 2: Pavel Urysohn is a mathematician. I need to search Leonid Levin next and find its type of work.
|
||||
Action 2: Search[Leonid Levin]
|
||||
Observation 2: Leonid Anatolievich Levin is a Soviet-American mathematician and computer scientist.
|
||||
Thought 3: Leonid Levin is a mathematician and computer scientist. So Pavel Urysohn and Leonid Levin have the same type of work.
|
||||
Action 3: Finish[yes]
|
||||
"""
|
||||
+ question
|
||||
)
|
||||
for i in range(1, len(triplets) + 2):
|
||||
s += "Thought " + str(i) + ":"
|
||||
# NOTE: This is an implementation for replaying a given trace for benchmark purposes. It is not an actual ReAct agent implementation.
|
||||
ss = s.fork(1)
|
||||
ss[0] += sgl.gen(name="thought_action", max_tokens=200, stop="Observation")
|
||||
ss.join()
|
||||
# to verify the correctness of output, this should be collected
|
||||
# print(ss[0]["thought_action"])
|
||||
if i > len(triplets):
|
||||
break
|
||||
s += (
|
||||
triplets[i - 1]["thought"]
|
||||
+ "\nAction "
|
||||
+ str(i)
|
||||
+ ":"
|
||||
+ triplets[i - 1]["action"]
|
||||
+ "\nObservation "
|
||||
+ str(i)
|
||||
+ ":"
|
||||
+ triplets[i - 1]["observation"]
|
||||
+ "\n"
|
||||
)
|
||||
|
||||
|
||||
def main(args):
|
||||
lines = read_jsonl(args.data_path)[: args.num_questions]
|
||||
arguments = [{"question": k, "triplets": v} for l in lines for k, v in l.items()]
|
||||
|
||||
# Select backend
|
||||
backend = select_sglang_backend(args)
|
||||
sgl.set_default_backend(backend)
|
||||
|
||||
states = []
|
||||
tic = time.perf_counter()
|
||||
states = webthink.run_batch(
|
||||
arguments,
|
||||
temperature=0,
|
||||
num_threads=args.parallel,
|
||||
progress_bar=True,
|
||||
)
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
# Compute accuracy
|
||||
print(f"Latency: {latency:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "ReAct Agents",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"num_requests": len(arguments),
|
||||
"other": {
|
||||
"num_questions": args.num_questions,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="hotpotqa_100.jsonl")
|
||||
parser.add_argument("--num-questions", type=int, default=10)
|
||||
args = add_common_sglang_args_and_parse(parser)
|
||||
main(args)
|
||||
File diff suppressed because one or more lines are too long
@@ -1,77 +0,0 @@
|
||||
# Run benchmark
|
||||
|
||||
This benchmark is primarily intended to be used with reasoning models like `DeepSeek-R1` and its distilled models like `DeepSeek-R1-Distill-Qwen-1.5B`. Please use
|
||||
|
||||
```bash
|
||||
pip install antlr4-python3-runtime
|
||||
```
|
||||
|
||||
for `parse_latex` which we use for symbolic equality check.
|
||||
|
||||
## Benchmark sglang
|
||||
|
||||
1. Launch the Server
|
||||
```bash
|
||||
python3 -m sglang.launch_server --model-path deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B --port 30000
|
||||
```
|
||||
|
||||
Note that depending on the GPU this benchmark will take quiet some time. To employ data parallelism please use:
|
||||
|
||||
```bash
|
||||
python3 -m sglang_router.launch_server --model-path deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B --port 30000 --dp-size 4
|
||||
```
|
||||
|
||||
|
||||
2. Benchmarking
|
||||
|
||||
We use [suggested](https://github.com/deepseek-ai/DeepSeek-R1) parameters of `temperature=0.6`, `top_p=.95`, `max_new_tokens=32768`. The command line argument `num-tries` can be used to evaluate the model multiple times on the same question. We use the suggested `64` from the repo for AIME 2024. For LIMO, we use `8` as the number of tries due to the size of the dataset.
|
||||
|
||||
By default evaluate on LIMO dataset.
|
||||
|
||||
```bash
|
||||
python3 bench_sglang.py --parallel 256 --num-tries 64 --port 30000
|
||||
```
|
||||
|
||||
Evaluate on AIME 2024 dataset.
|
||||
|
||||
```bash
|
||||
python3 bench_sglang.py --parallel 256 --port 30000 --data-path Maxwell-Jia/AIME_2024 --question-key Problem --answer-key Answer --num-tries 64
|
||||
```
|
||||
|
||||
Evaluate on [AIME 2025 I dataset](https://huggingface.co/datasets/opencompass/AIME2025). For benchmark result see [here](https://matharena.ai/).
|
||||
|
||||
```bash
|
||||
python3 bench_sglang.py --parallel 256 --port 30000 --data-path opencompass/AIME2025 --question-key question --answer-key answer --num-tries 64
|
||||
```
|
||||
## Results
|
||||
|
||||
### Evaluation Results
|
||||
| Dataset | Num Tries | Accuracy | Reference | Standard Error |
|
||||
|------------|-----------|----------|-----------|-----------|
|
||||
| LIMO | 8 | 47.7% | ? | ? |
|
||||
| AIME 2024 | 64 | 33.2% | 28.9% | 3.4% |
|
||||
| AIME 2025 I| 64 | 29.9% | 25.0% | ? |
|
||||
|
||||
### Statistic Analysis Results
|
||||
Set up SGLang engine for statistic analysis, for high efficiency we use `--dp-size 8` for data parallelism:
|
||||
```bash
|
||||
python3 -m sglang_router.launch_server --model-path deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B --port 30000 --dp-size 8
|
||||
```
|
||||
**Experiment 1**:
|
||||
We fixed the number of attempts (num_tries) and conducted multiple runs to assess the consistency of the model's performance. The results show that all recorded accuracies lie within ± one standard error deviation from the mean. This suggests that **our metric serves as an effective upper bound for the deviation of reported accuracy**.
|
||||
|
||||
To collect the accuracy, run the following command 30 times:
|
||||
```bash
|
||||
python3 bench_sglang.py --parallel 64 --port 30000 --data-path Maxwell-Jia/AIME_2024 --question-key Problem --answer-key Answer --num-tries 64
|
||||
```
|
||||
|
||||

|
||||
|
||||
|
||||
**Experiment 2**: We explored the relationship between the number of attempts (num_tries) and the standard error (SE) by varying num_tries across a range (e.g., 8, 16, 32, ..., 256) and performing a single run for each value. The results demonstrate that as the number of attempts increases, the standard error decreases, leading to **greater stability in answer accuracy**.
|
||||
|
||||
To reveal the relationship, run the command 6 times and adjust the parameter `--num-tries` for each run:
|
||||
```bash
|
||||
python3 bench_sglang.py --parallel 64 --port 30000 --data-path Maxwell-Jia/AIME_2024 --question-key Problem --answer-key Answer --num-tries <num_tries>
|
||||
```
|
||||

|
||||
@@ -1,269 +0,0 @@
|
||||
# Adapted from https://github.com/deepseek-ai/DeepSeek-Math/blob/main/evaluation/data_processing/answer_extraction.py
|
||||
|
||||
import re
|
||||
|
||||
import regex
|
||||
|
||||
|
||||
def _fix_fracs(string):
|
||||
substrs = string.split("\\frac")
|
||||
new_str = substrs[0]
|
||||
if len(substrs) > 1:
|
||||
substrs = substrs[1:]
|
||||
for substr in substrs:
|
||||
new_str += "\\frac"
|
||||
if len(substr) > 0 and substr[0] == "{":
|
||||
new_str += substr
|
||||
else:
|
||||
try:
|
||||
assert len(substr) >= 2
|
||||
except:
|
||||
return string
|
||||
a = substr[0]
|
||||
b = substr[1]
|
||||
if b != "{":
|
||||
if len(substr) > 2:
|
||||
post_substr = substr[2:]
|
||||
new_str += "{" + a + "}{" + b + "}" + post_substr
|
||||
else:
|
||||
new_str += "{" + a + "}{" + b + "}"
|
||||
else:
|
||||
if len(substr) > 2:
|
||||
post_substr = substr[2:]
|
||||
new_str += "{" + a + "}" + b + post_substr
|
||||
else:
|
||||
new_str += "{" + a + "}" + b
|
||||
string = new_str
|
||||
return string
|
||||
|
||||
|
||||
def _fix_a_slash_b(string):
|
||||
if len(string.split("/")) != 2:
|
||||
return string
|
||||
a = string.split("/")[0]
|
||||
b = string.split("/")[1]
|
||||
try:
|
||||
if "sqrt" not in a:
|
||||
a = int(a)
|
||||
if "sqrt" not in b:
|
||||
b = int(b)
|
||||
assert string == "{}/{}".format(a, b)
|
||||
new_string = "\\frac{" + str(a) + "}{" + str(b) + "}"
|
||||
return new_string
|
||||
except:
|
||||
return string
|
||||
|
||||
|
||||
def _fix_sqrt(string):
|
||||
_string = re.sub(r"\\sqrt(-?[0-9.a-zA-Z]+)", r"\\sqrt{\1}", string)
|
||||
_string = re.sub(r"\\sqrt\s+(\w+)$", r"\\sqrt{\1}", _string)
|
||||
return _string
|
||||
|
||||
|
||||
def _fix_tan(string):
|
||||
_string = re.sub(r"\\tan(-?[0-9.a-zA-Z]+)", r"\\tan{\1}", string)
|
||||
_string = re.sub(r"\\tan\s+(\w+)$", r"\\tan{\1}", _string)
|
||||
return _string
|
||||
|
||||
|
||||
def strip_string(string):
|
||||
string = str(string).strip()
|
||||
# linebreaks
|
||||
string = string.replace("\n", "")
|
||||
|
||||
# right "."
|
||||
string = string.rstrip(".")
|
||||
|
||||
# remove inverse spaces
|
||||
string = string.replace("\\!", "")
|
||||
# string = string.replace("\\ ", "")
|
||||
|
||||
# replace \\ with \
|
||||
# string = string.replace("\\\\", "\\")
|
||||
# string = string.replace("\\\\", "\\")
|
||||
|
||||
if string.startswith("\\text{") and string.endswith("}"):
|
||||
string = string.split("{", 1)[1][:-1]
|
||||
|
||||
# replace tfrac and dfrac with frac
|
||||
string = string.replace("tfrac", "frac")
|
||||
string = string.replace("dfrac", "frac")
|
||||
string = string.replace("cfrac", "frac")
|
||||
|
||||
# remove \left and \right
|
||||
string = string.replace("\\left", "")
|
||||
string = string.replace("\\right", "")
|
||||
|
||||
# Remove unit: miles, dollars if after is not none
|
||||
_string = re.sub(r"\\text{.*?}$", "", string).strip()
|
||||
if _string != "" and _string != string:
|
||||
# print("Warning: unit not removed: '{}' -> '{}'".format(string, _string))
|
||||
string = _string
|
||||
|
||||
# Remove circ (degrees)
|
||||
string = string.replace("^{\\circ}", "").strip()
|
||||
string = string.replace("^\\circ", "").strip()
|
||||
|
||||
string = regex.sub(r"\{(c|m)?m\}(\^(2|3))?", "", string).strip()
|
||||
string = regex.sub(r"p\.m\.$", "", string).strip()
|
||||
string = regex.sub(r"(\d)\s*t$", r"\1", string).strip()
|
||||
|
||||
# remove dollar signs
|
||||
string = string.replace("\\$", "")
|
||||
string = string.replace("$", "")
|
||||
|
||||
# string = string.replace("\\text", "")
|
||||
string = string.replace("x\\in", "")
|
||||
|
||||
# remove percentage
|
||||
string = string.replace("\\%", "%")
|
||||
string = string.replace("\%", "%")
|
||||
# string = string.replace("%", "")
|
||||
|
||||
# " 0." equivalent to " ." and "{0." equivalent to "{." Alternatively, add "0" if "." is the start of the string
|
||||
string = string.replace(" .", " 0.")
|
||||
string = string.replace("{.", "{0.")
|
||||
|
||||
# cdot
|
||||
string = string.replace("\\cdot", "")
|
||||
|
||||
# inf
|
||||
string = string.replace("infinity", "\\infty")
|
||||
if "\\infty" not in string:
|
||||
string = string.replace("inf", "\\infty")
|
||||
string = string.replace("+\\inity", "\\infty")
|
||||
|
||||
# and
|
||||
# string = string.replace("and", "")
|
||||
string = string.replace("\\mathbf", "")
|
||||
string = string.replace("\\mathrm", "")
|
||||
|
||||
# use regex to remove \mbox{...}
|
||||
string = re.sub(r"\\mbox{.*?}", "", string)
|
||||
|
||||
# quote
|
||||
string.replace("'", "")
|
||||
string.replace('"', "")
|
||||
|
||||
# i, j
|
||||
if "j" in string and "i" not in string:
|
||||
string = string.replace("j", "i")
|
||||
|
||||
# replace a.000b where b is not number or b is end, with ab, use regex
|
||||
string = re.sub(r"(\d+)\.0+([^\d])", r"\1\2", string)
|
||||
string = re.sub(r"(\d+)\.0+$", r"\1", string)
|
||||
|
||||
# if empty, return empty string
|
||||
if len(string) == 0:
|
||||
return string
|
||||
if string[0] == ".":
|
||||
string = "0" + string
|
||||
|
||||
# to consider: get rid of e.g. "k = " or "q = " at beginning
|
||||
# if len(string.split("=")) == 2:
|
||||
# if len(string.split("=")[0]) <= 2:
|
||||
# string = string.split("=")[1]
|
||||
|
||||
string = _fix_sqrt(string)
|
||||
string = _fix_tan(string)
|
||||
string = string.replace(" ", "")
|
||||
|
||||
# \frac1b or \frac12 --> \frac{1}{b} and \frac{1}{2}, etc. Even works with \frac1{72} (but not \frac{72}1). Also does a/b --> \\frac{a}{b}
|
||||
string = _fix_fracs(string)
|
||||
|
||||
# NOTE: X/Y changed to \frac{X}{Y} in dataset, but in simple cases fix in case the model output is X/Y
|
||||
string = _fix_a_slash_b(string)
|
||||
|
||||
string = regex.sub(r"(\\|,|\.)+$", "", string)
|
||||
|
||||
return string
|
||||
|
||||
|
||||
def extract_boxed_answers(text):
|
||||
answers = []
|
||||
for piece in text.split("boxed{")[1:]:
|
||||
n = 0
|
||||
for i in range(len(piece)):
|
||||
if piece[i] == "{":
|
||||
n += 1
|
||||
elif piece[i] == "}":
|
||||
n -= 1
|
||||
if n < 0:
|
||||
if i + 1 < len(piece) and piece[i + 1] == "%":
|
||||
answers.append(piece[: i + 1])
|
||||
else:
|
||||
answers.append(piece[:i])
|
||||
break
|
||||
return answers
|
||||
|
||||
|
||||
def extract_program_output(pred_str):
|
||||
"""
|
||||
extract output between the last ```output\n...\n```
|
||||
"""
|
||||
if "```output" not in pred_str:
|
||||
return ""
|
||||
if "```output" in pred_str:
|
||||
pred_str = pred_str.split("```output")[-1]
|
||||
if "```" in pred_str:
|
||||
pred_str = pred_str.split("```")[0]
|
||||
output = pred_str.strip()
|
||||
return output
|
||||
|
||||
|
||||
def extract_answer(pred_str, exhaust=False):
|
||||
pred = []
|
||||
if "final answer is $" in pred_str and "$. I hope" in pred_str:
|
||||
tmp = pred_str.split("final answer is $", 1)[1]
|
||||
pred = [tmp.split("$. I hope", 1)[0].strip()]
|
||||
elif "boxed" in pred_str:
|
||||
pred = extract_boxed_answers(pred_str)
|
||||
elif "he answer is" in pred_str:
|
||||
pred = [pred_str.split("he answer is")[-1].strip()]
|
||||
else:
|
||||
program_output = extract_program_output(pred_str)
|
||||
if program_output != "":
|
||||
# fall back to program
|
||||
pred.append(program_output)
|
||||
else: # use the last number
|
||||
pattern = "-?\d*\.?\d+"
|
||||
answers = re.findall(pattern, pred_str.replace(",", ""))
|
||||
if len(answers) >= 1:
|
||||
last_ans = answers[-1]
|
||||
else:
|
||||
last_ans = ""
|
||||
if last_ans:
|
||||
pred.append(last_ans)
|
||||
|
||||
# multiple line
|
||||
_pred = []
|
||||
for each_ans in pred:
|
||||
each_ans = each_ans.strip().split("\n")[0]
|
||||
each_ans = each_ans.lstrip(":")
|
||||
each_ans = each_ans.rstrip(".")
|
||||
each_ans = each_ans.rstrip("/")
|
||||
each_ans = strip_string(each_ans)
|
||||
_pred.append(each_ans)
|
||||
if exhaust:
|
||||
return _pred
|
||||
else:
|
||||
return _pred[-1] if _pred else ""
|
||||
|
||||
|
||||
def extract_math_answer(question, reasoning, task):
|
||||
answer = []
|
||||
for ans in extract_answer(reasoning, exhaust=True):
|
||||
if "separated by commas" in question and all(ch not in ans for ch in "()[]"):
|
||||
answer.extend([a.strip() for a in ans.split(",")])
|
||||
elif regex.search(r"\\text\{\s*and\s*\}", ans):
|
||||
answer.extend(
|
||||
[
|
||||
a.strip()
|
||||
for a in regex.sub(r"\\text\{\s*and\s*\}", "[SEP]", ans).split(
|
||||
"[SEP]"
|
||||
)
|
||||
]
|
||||
)
|
||||
else:
|
||||
answer.append(ans.strip())
|
||||
return answer
|
||||
@@ -1,135 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
|
||||
import answer_extraction
|
||||
import eval_utils
|
||||
import numpy as np
|
||||
from datasets import load_dataset
|
||||
|
||||
import sglang as sgl
|
||||
from sglang.test.test_utils import (
|
||||
add_common_sglang_args_and_parse,
|
||||
select_sglang_backend,
|
||||
)
|
||||
from sglang.utils import dump_state_text
|
||||
|
||||
|
||||
@sgl.function
|
||||
def reasoning_gen(s, question: str):
|
||||
s += sgl.user(
|
||||
question
|
||||
+ "\nPlease reason step by step, and put your final answer within \boxed{}."
|
||||
)
|
||||
s += sgl.assistant(
|
||||
sgl.gen(
|
||||
"answer",
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def convert_dataset(path: str, question_key: str, answer_key: str, num_tries: int):
|
||||
raw_dataset = load_dataset(path)
|
||||
questions = []
|
||||
answers = []
|
||||
for data in raw_dataset["train"]:
|
||||
question = data[question_key]
|
||||
answer = data[answer_key]
|
||||
for _ in range(num_tries):
|
||||
questions.append({"question": question})
|
||||
answers.append({"answer": answer})
|
||||
return questions, answers
|
||||
|
||||
|
||||
def main(args):
|
||||
# Select backend
|
||||
sgl.set_default_backend(select_sglang_backend(args))
|
||||
|
||||
# Get dataset
|
||||
questions, answers = convert_dataset(
|
||||
args.data_path, args.question_key, args.answer_key, args.num_tries
|
||||
)
|
||||
|
||||
# Run requests
|
||||
tic = time.perf_counter()
|
||||
states = reasoning_gen.run_batch(
|
||||
questions,
|
||||
num_threads=args.parallel,
|
||||
progress_bar=True,
|
||||
temperature=0.6,
|
||||
max_new_tokens=32768,
|
||||
top_p=0.95,
|
||||
)
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
# Extract results and record outcomes in a list.
|
||||
outcomes = []
|
||||
for i, state in enumerate(states):
|
||||
try:
|
||||
pred_answer = answer_extraction.extract_math_answer(
|
||||
questions[i]["question"], state["answer"], "limo"
|
||||
)
|
||||
gt_answer = str(answers[i]["answer"])
|
||||
pred_answer = (
|
||||
pred_answer[-1] if isinstance(pred_answer, list) else pred_answer
|
||||
)
|
||||
is_correct = 1 if eval_utils.math_equal(pred_answer, gt_answer) else 0
|
||||
except Exception as e:
|
||||
print(f"Error extracting answer: {e}")
|
||||
is_correct = 0
|
||||
|
||||
outcomes.append(is_correct)
|
||||
|
||||
# Calculate overall accuracy using numpy
|
||||
overall_accuracy = np.mean(outcomes)
|
||||
print(f"Overall Accuracy: {overall_accuracy}")
|
||||
|
||||
# Calculate mean standard error over questions if num_tries >= 2
|
||||
if args.num_tries > 1:
|
||||
outcomes_np = np.array(outcomes).reshape(-1, args.num_tries)
|
||||
# Using sample standard deviation with ddof=1
|
||||
std_per_question = np.std(outcomes_np, axis=1, ddof=1)
|
||||
# Compute the standard error for each question: std / sqrt(num_tries)
|
||||
se_per_question = std_per_question / np.sqrt(args.num_tries)
|
||||
mean_se = se_per_question.mean()
|
||||
print(f"Mean Standard Error of Accuracy across questions: {mean_se}")
|
||||
else:
|
||||
mean_se = None
|
||||
print("Not enough samples per question to compute standard error.")
|
||||
|
||||
# Calculate output throughput
|
||||
num_output_tokens = sum(
|
||||
s.get_meta_info("answer")["completion_tokens"] for s in states
|
||||
)
|
||||
output_throughput = num_output_tokens / latency
|
||||
print(f"Output throughput: {output_throughput} token/s")
|
||||
|
||||
# Dump results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
# Write results
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "limo",
|
||||
"backend": args.backend,
|
||||
"latency": round(latency, 3),
|
||||
"overall_accuracy": round(overall_accuracy, 3),
|
||||
"mean_se_accuracy": round(mean_se, 3) if mean_se is not None else None,
|
||||
"num_requests": len(questions),
|
||||
"other": {
|
||||
"num_questions": len(questions),
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="GAIR/LIMO")
|
||||
parser.add_argument("--question-key", type=str, default="question")
|
||||
parser.add_argument("--answer-key", type=str, default="answer")
|
||||
parser.add_argument("--num-tries", type=int, default=1)
|
||||
add_common_sglang_args_and_parse(parser)
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -1,206 +0,0 @@
|
||||
# Adapted from https://github.com/deepseek-ai/DeepSeek-Math/blob/main/evaluation/eval/eval_utils.py
|
||||
|
||||
from math import isclose
|
||||
|
||||
import regex
|
||||
from sympy import N, simplify
|
||||
from sympy.parsing.latex import parse_latex
|
||||
from sympy.parsing.sympy_parser import parse_expr
|
||||
|
||||
|
||||
def parse_digits(num):
|
||||
# format: 234.23 || 23%
|
||||
num = regex.sub(",", "", str(num))
|
||||
try:
|
||||
return float(num)
|
||||
except:
|
||||
if num.endswith("%"):
|
||||
num = num[:-1]
|
||||
if num.endswith("\\"):
|
||||
num = num[:-1]
|
||||
try:
|
||||
return float(num) / 100
|
||||
except:
|
||||
pass
|
||||
return None
|
||||
|
||||
|
||||
def is_digit(num):
|
||||
# paired with parse_digits
|
||||
return parse_digits(num) is not None
|
||||
|
||||
|
||||
def symbolic_equal(a, b):
|
||||
def _parse(s):
|
||||
for f in [parse_latex, parse_expr]:
|
||||
try:
|
||||
return f(s)
|
||||
except:
|
||||
pass
|
||||
return s
|
||||
|
||||
a = _parse(a)
|
||||
b = _parse(b)
|
||||
|
||||
try:
|
||||
if simplify(a - b) == 0:
|
||||
return True
|
||||
except:
|
||||
pass
|
||||
|
||||
try:
|
||||
if isclose(N(a), N(b), abs_tol=1e-3):
|
||||
return True
|
||||
except:
|
||||
pass
|
||||
return False
|
||||
|
||||
|
||||
def math_equal(prediction, reference, include_percentage=True, is_close=True):
|
||||
"""
|
||||
Exact match of math if and only if:
|
||||
1. numerical equal: both can convert to float and are equal
|
||||
2. symbolic equal: both can convert to sympy expression and are equal
|
||||
"""
|
||||
if str(prediction) == str(reference):
|
||||
return True
|
||||
|
||||
try: # 1. numerical equal
|
||||
if is_digit(prediction) and is_digit(reference):
|
||||
prediction = parse_digits(prediction)
|
||||
reference = parse_digits(reference)
|
||||
# number questions
|
||||
if include_percentage:
|
||||
gt_result = [reference / 100, reference, reference * 100]
|
||||
else:
|
||||
gt_result = [reference]
|
||||
for item in gt_result:
|
||||
try:
|
||||
if is_close:
|
||||
if isclose(item, prediction, abs_tol=1e-3):
|
||||
return True
|
||||
else:
|
||||
if item == prediction:
|
||||
return True
|
||||
except Exception:
|
||||
continue
|
||||
return False
|
||||
except:
|
||||
pass
|
||||
|
||||
if not prediction and prediction not in [0, False]:
|
||||
return False
|
||||
|
||||
# 2. symbolic equal
|
||||
reference = str(reference).strip()
|
||||
prediction = str(prediction).strip()
|
||||
|
||||
if (
|
||||
regex.match(r"(\(|\[).+(\)|\])", prediction) is not None
|
||||
and regex.match(r"(\(|\[).+(\)|\])", reference) is not None
|
||||
):
|
||||
pred_parts = prediction[1:-1].split(",")
|
||||
ref_parts = reference[1:-1].split(",")
|
||||
if len(pred_parts) == len(ref_parts):
|
||||
if all(
|
||||
[
|
||||
math_equal(
|
||||
pred_parts[i], ref_parts[i], include_percentage, is_close
|
||||
)
|
||||
for i in range(len(pred_parts))
|
||||
]
|
||||
):
|
||||
return True
|
||||
|
||||
# Add back matrix comparison
|
||||
if (
|
||||
(
|
||||
prediction.startswith("\\begin{pmatrix}")
|
||||
or prediction.startswith("\\begin{bmatrix}")
|
||||
)
|
||||
and (
|
||||
prediction.endswith("\\end{pmatrix}")
|
||||
or prediction.endswith("\\end{bmatrix}")
|
||||
)
|
||||
and (
|
||||
reference.startswith("\\begin{pmatrix}")
|
||||
or reference.startswith("\\begin{bmatrix}")
|
||||
)
|
||||
and (
|
||||
reference.endswith("\\end{pmatrix}") or reference.endswith("\\end{bmatrix}")
|
||||
)
|
||||
):
|
||||
pred_lines = [
|
||||
line.strip()
|
||||
for line in prediction[
|
||||
len("\\begin{pmatrix}") : -len("\\end{pmatrix}")
|
||||
].split("\\\\")
|
||||
if line.strip()
|
||||
]
|
||||
ref_lines = [
|
||||
line.strip()
|
||||
for line in reference[
|
||||
len("\\begin{pmatrix}") : -len("\\end{pmatrix}")
|
||||
].split("\\\\")
|
||||
if line.strip()
|
||||
]
|
||||
matched = True
|
||||
if len(pred_lines) == len(ref_lines):
|
||||
for pred_line, ref_line in zip(pred_lines, ref_lines):
|
||||
pred_parts = pred_line.split("&")
|
||||
ref_parts = ref_line.split("&")
|
||||
if len(pred_parts) == len(ref_parts):
|
||||
if not all(
|
||||
[
|
||||
math_equal(
|
||||
pred_parts[i],
|
||||
ref_parts[i],
|
||||
include_percentage,
|
||||
is_close,
|
||||
)
|
||||
for i in range(len(pred_parts))
|
||||
]
|
||||
):
|
||||
matched = False
|
||||
break
|
||||
else:
|
||||
matched = False
|
||||
if not matched:
|
||||
break
|
||||
else:
|
||||
matched = False
|
||||
if matched:
|
||||
return True
|
||||
|
||||
# Add back equation comparison
|
||||
if prediction.count("=") == 1 and reference.count("=") == 1:
|
||||
pred = prediction.split("=")
|
||||
pred = f"{pred[0].strip()} - ({pred[1].strip()})"
|
||||
ref = reference.split("=")
|
||||
ref = f"{ref[0].strip()} - ({ref[1].strip()})"
|
||||
if symbolic_equal(pred, ref) or symbolic_equal(f"-({pred})", ref):
|
||||
return True
|
||||
elif (
|
||||
prediction.count("=") == 1
|
||||
and len(prediction.split("=")[0].strip()) <= 2
|
||||
and "=" not in reference
|
||||
):
|
||||
if math_equal(
|
||||
prediction.split("=")[1], reference, include_percentage, is_close
|
||||
):
|
||||
return True
|
||||
elif (
|
||||
reference.count("=") == 1
|
||||
and len(reference.split("=")[0].strip()) <= 2
|
||||
and "=" not in prediction
|
||||
):
|
||||
if math_equal(
|
||||
prediction, reference.split("=")[1], include_percentage, is_close
|
||||
):
|
||||
return True
|
||||
|
||||
# symbolic equal with sympy
|
||||
if symbolic_equal(prediction, reference):
|
||||
return True
|
||||
|
||||
return False
|
||||
Binary file not shown.
|
Before Width: | Height: | Size: 33 KiB |
Binary file not shown.
|
Before Width: | Height: | Size: 26 KiB |
@@ -1 +0,0 @@
|
||||
!topic.jsonl
|
||||
@@ -1,33 +0,0 @@
|
||||
## Run benchmark
|
||||
|
||||
### Benchmark sglang
|
||||
```
|
||||
python -m sglang.launch_server --model-path meta-llama/Llama-2-7b-chat-hf --port 30000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_sglang.py --num-questions 64
|
||||
python3 bench_sglang.py --num-questions 32 --parallel 1
|
||||
```
|
||||
|
||||
|
||||
### Benchmark vllm
|
||||
```
|
||||
python3 -m vllm.entrypoints.api_server --tokenizer-mode auto --model meta-llama/Llama-2-7b-chat-hf --disable-log-requests --port 21000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --backend vllm --num-questions 64
|
||||
```
|
||||
|
||||
|
||||
### Benchmark guidance
|
||||
```
|
||||
python3 bench_other.py --backend guidance --num-questions 32 --parallel 1 --n-ctx 4096 --model-path path/to/gguf
|
||||
```
|
||||
|
||||
### Benchmark lmql
|
||||
|
||||
```
|
||||
python3 bench_other.py --backend lmql --num-questions 32 --parallel 1
|
||||
```
|
||||
@@ -1,127 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from functools import partial
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
from sglang.test.test_utils import add_common_other_args_and_parse, get_call_generate
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
number = 5
|
||||
|
||||
|
||||
def expand_tip(topic, tip, generate):
|
||||
s = """Please expand a tip for a topic into a detailed paragraph.
|
||||
|
||||
Topic: staying healthy
|
||||
Tip: Regular Exercise
|
||||
Paragraph: Incorporate physical activity into your daily routine. This doesn't necessarily mean intense gym workouts; it can be as simple as walking, cycling, or yoga. Regular exercise helps in maintaining a healthy weight, improves cardiovascular health, boosts mental health, and can enhance cognitive function, which is crucial for fields that require intense intellectual engagement.
|
||||
|
||||
Topic: building a campfire
|
||||
Tip: Choose the Right Location
|
||||
Paragraph: Always build your campfire in a safe spot. This means selecting a location that's away from trees, bushes, and other flammable materials. Ideally, use a fire ring if available. If you're building a fire pit, it should be on bare soil or on a bed of stones, not on grass or near roots which can catch fire underground. Make sure the area above is clear of low-hanging branches.
|
||||
|
||||
Topic: writing a blog post
|
||||
Tip: structure your content effectively
|
||||
Paragraph: A well-structured post is easier to read and more enjoyable. Start with an engaging introduction that hooks the reader and clearly states the purpose of your post. Use headings and subheadings to break up the text and guide readers through your content. Bullet points and numbered lists can make information more digestible. Ensure each paragraph flows logically into the next, and conclude with a summary or call-to-action that encourages reader engagement.
|
||||
|
||||
Topic: """ + topic + "\nTip: " + tip + "\nParagraph:"
|
||||
return generate(s, max_tokens=128, stop=["\n\n"])
|
||||
|
||||
|
||||
def suggest_tips(topic, generate):
|
||||
s = "Please act as a helpful assistant. Your job is to provide users with useful tips on a specific topic.\n"
|
||||
s += "USER: Give some tips for " + topic + ".\n"
|
||||
s += (
|
||||
"ASSISTANT: Okay. Here are "
|
||||
+ str(number)
|
||||
+ " concise tips, each under 8 words:\n"
|
||||
)
|
||||
|
||||
tips = []
|
||||
for i in range(1, 1 + number):
|
||||
s += f"{i}."
|
||||
tip = generate(s, max_tokens=24, stop=[".", "\n"])
|
||||
s += tip + ".\n"
|
||||
tips.append(tip)
|
||||
|
||||
paragraphs = [expand_tip(topic, tip, generate=generate) for tip in tips]
|
||||
|
||||
for i in range(1, 1 + number):
|
||||
s += f"Tip {i}:" + paragraphs[i - 1] + "\n"
|
||||
return s
|
||||
|
||||
|
||||
def main(args):
|
||||
lines = read_jsonl(args.data_path)[: args.num_questions]
|
||||
states = [None] * len(lines)
|
||||
|
||||
# Select backend
|
||||
call_generate = partial(get_call_generate(args), temperature=0)
|
||||
|
||||
# Run requests
|
||||
tic = time.perf_counter()
|
||||
if args.backend != "lmql":
|
||||
|
||||
def get_one_answer(i):
|
||||
states[i] = suggest_tips(lines[i]["topic"], call_generate)
|
||||
|
||||
if args.parallel == 1:
|
||||
for i in tqdm(range(len(lines))):
|
||||
get_one_answer(i)
|
||||
else:
|
||||
with ThreadPoolExecutor(args.parallel) as executor:
|
||||
list(
|
||||
tqdm(
|
||||
executor.map(get_one_answer, list(range(len(lines)))),
|
||||
total=len(lines),
|
||||
)
|
||||
)
|
||||
|
||||
else:
|
||||
import asyncio
|
||||
|
||||
from lmql_funcs import suggest_tips_async
|
||||
|
||||
async def get_one_answer_async(i):
|
||||
states[i] = await suggest_tips_async(lines[i]["topic"], call_generate)
|
||||
|
||||
batches = []
|
||||
for i in range(0, len(lines), args.parallel):
|
||||
batches.append(list(range(i, min(i + args.parallel, len(lines)))))
|
||||
loop = asyncio.get_event_loop()
|
||||
for batch in tqdm(batches):
|
||||
loop.run_until_complete(
|
||||
asyncio.gather(*[get_one_answer_async(i) for i in batch])
|
||||
)
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
# Compute accuracy
|
||||
print(f"Latency: {latency:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "tip_suggestion",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"other": {
|
||||
"num_questions": args.num_questions,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="topic.jsonl")
|
||||
parser.add_argument("--num-questions", type=int, default=100)
|
||||
args = add_common_other_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,94 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import time
|
||||
|
||||
import sglang as sgl
|
||||
from sglang.test.test_utils import (
|
||||
add_common_sglang_args_and_parse,
|
||||
select_sglang_backend,
|
||||
)
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
number = 5
|
||||
|
||||
|
||||
@sgl.function
|
||||
def expand_tip(s, topic, tip):
|
||||
s += """Please expand a tip for a topic into a detailed paragraph.
|
||||
|
||||
Topic: staying healthy
|
||||
Tip: Regular Exercise
|
||||
Paragraph: Incorporate physical activity into your daily routine. This doesn't necessarily mean intense gym workouts; it can be as simple as walking, cycling, or yoga. Regular exercise helps in maintaining a healthy weight, improves cardiovascular health, boosts mental health, and can enhance cognitive function, which is crucial for fields that require intense intellectual engagement.
|
||||
|
||||
Topic: building a campfire
|
||||
Tip: Choose the Right Location
|
||||
Paragraph: Always build your campfire in a safe spot. This means selecting a location that's away from trees, bushes, and other flammable materials. Ideally, use a fire ring if available. If you're building a fire pit, it should be on bare soil or on a bed of stones, not on grass or near roots which can catch fire underground. Make sure the area above is clear of low-hanging branches.
|
||||
|
||||
Topic: writing a blog post
|
||||
Tip: structure your content effectively
|
||||
Paragraph: A well-structured post is easier to read and more enjoyable. Start with an engaging introduction that hooks the reader and clearly states the purpose of your post. Use headings and subheadings to break up the text and guide readers through your content. Bullet points and numbered lists can make information more digestible. Ensure each paragraph flows logically into the next, and conclude with a summary or call-to-action that encourages reader engagement.
|
||||
|
||||
Topic: """ + topic + "\nTip: " + tip + "\nParagraph:"
|
||||
s += sgl.gen("paragraph", max_tokens=128, stop=["\n\n"], temperature=0)
|
||||
|
||||
|
||||
@sgl.function
|
||||
def suggest_tips(s, topic):
|
||||
s += "Please act as a helpful assistant. Your job is to provide users with useful tips on a specific topic.\n"
|
||||
s += "USER: Give some tips for " + topic + ".\n"
|
||||
s += (
|
||||
"ASSISTANT: Okay. Here are "
|
||||
+ str(number)
|
||||
+ " concise tips, each under 8 words:\n"
|
||||
)
|
||||
|
||||
paragraphs = []
|
||||
for i in range(1, 1 + number):
|
||||
s += f"{i}." + sgl.gen(f"tip_{i}", max_tokens=24, stop=[".", "\n"]) + ".\n"
|
||||
paragraphs.append(expand_tip(topic=topic, tip=s[f"tip_{i}"]))
|
||||
|
||||
for i in range(1, 1 + number):
|
||||
s += f"Tip {i}:" + paragraphs[i - 1]["paragraph"] + "\n"
|
||||
|
||||
|
||||
def main(args):
|
||||
lines = read_jsonl(args.data_path)[: args.num_questions]
|
||||
arguments = [{"topic": l["topic"]} for l in lines]
|
||||
|
||||
# Select backend
|
||||
sgl.set_default_backend(select_sglang_backend(args))
|
||||
|
||||
# Run requests
|
||||
tic = time.perf_counter()
|
||||
states = suggest_tips.run_batch(
|
||||
arguments, temperature=0, num_threads=args.parallel, progress_bar=True
|
||||
)
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
# Compute accuracy
|
||||
print(f"Latency: {latency:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", states)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "tip_suggestion",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"other": {
|
||||
"num_questions": args.num_questions,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="topic.jsonl")
|
||||
parser.add_argument("--num-questions", type=int, default=100)
|
||||
args = add_common_sglang_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,44 +0,0 @@
|
||||
number = 5
|
||||
|
||||
|
||||
async def expand_tip_async(topic, tip, generate):
|
||||
s = """Please expand a tip for a topic into a detailed paragraph.
|
||||
|
||||
Topic: staying healthy
|
||||
Tip: Regular Exercise
|
||||
Paragraph: Incorporate physical activity into your daily routine. This doesn't necessarily mean intense gym workouts; it can be as simple as walking, cycling, or yoga. Regular exercise helps in maintaining a healthy weight, improves cardiovascular health, boosts mental health, and can enhance cognitive function, which is crucial for fields that require intense intellectual engagement.
|
||||
|
||||
Topic: building a campfire
|
||||
Tip: Choose the Right Location
|
||||
Paragraph: Always build your campfire in a safe spot. This means selecting a location that's away from trees, bushes, and other flammable materials. Ideally, use a fire ring if available. If you're building a fire pit, it should be on bare soil or on a bed of stones, not on grass or near roots which can catch fire underground. Make sure the area above is clear of low-hanging branches.
|
||||
|
||||
Topic: writing a blog post
|
||||
Tip: structure your content effectively
|
||||
Paragraph: A well-structured post is easier to read and more enjoyable. Start with an engaging introduction that hooks the reader and clearly states the purpose of your post. Use headings and subheadings to break up the text and guide readers through your content. Bullet points and numbered lists can make information more digestible. Ensure each paragraph flows logically into the next, and conclude with a summary or call-to-action that encourages reader engagement.
|
||||
|
||||
Topic: """ + topic + "\nTip: " + tip + "\nParagraph:"
|
||||
return await generate(s, max_tokens=128, stop="\n\n")
|
||||
|
||||
|
||||
async def suggest_tips_async(topic, generate):
|
||||
s = "Please act as a helpful assistant. Your job is to provide users with useful tips on a specific topic.\n"
|
||||
s += "USER: Give some tips for " + topic + ".\n"
|
||||
s += (
|
||||
"ASSISTANT: Okay. Here are "
|
||||
+ str(number)
|
||||
+ " concise tips, each under 8 words:\n"
|
||||
)
|
||||
|
||||
tips = []
|
||||
for i in range(1, 1 + number):
|
||||
s += f"{i}."
|
||||
# NOTE: stop is different due to lmql does not support a list of stop tokens
|
||||
tip = await generate(s, max_tokens=24, stop=".\n")
|
||||
s += tip + ".\n"
|
||||
tips.append(tip)
|
||||
|
||||
paragraphs = [await expand_tip_async(topic, tip, generate=generate) for tip in tips]
|
||||
|
||||
for i in range(1, 1 + number):
|
||||
s += f"Tip {i}:" + paragraphs[i - 1] + "\n"
|
||||
return s
|
||||
@@ -1,50 +0,0 @@
|
||||
{"topic": "organizing a successful charity event", "number": 6}
|
||||
{"topic": "improving personal credit scores", "number": 7}
|
||||
{"topic": "staying motivated during job searches", "number": 5}
|
||||
{"topic": "maintaining a work-life balance", "number": 9}
|
||||
{"topic": "reducing carbon footprint at home", "number": 8}
|
||||
{"topic": "starting a book club", "number": 5}
|
||||
{"topic": "learning to play a musical instrument", "number": 7}
|
||||
{"topic": "getting into freelance writing", "number": 6}
|
||||
{"topic": "beginner yoga poses", "number": 8}
|
||||
{"topic": "preparing for graduate school exams", "number": 5}
|
||||
{"topic": "exploring minimalist living", "number": 9}
|
||||
{"topic": "effective grocery shopping", "number": 7}
|
||||
{"topic": "winter camping", "number": 5}
|
||||
{"topic": "starting a podcast on a budget", "number": 8}
|
||||
{"topic": "creating a capsule wardrobe", "number": 6}
|
||||
{"topic": "improving your writing skills", "number": 7}
|
||||
{"topic": "learning a new software quickly", "number": 9}
|
||||
{"topic": "reducing anxiety before public speaking", "number": 5}
|
||||
{"topic": "planning a solo travel adventure", "number": 8}
|
||||
{"topic": "beginner skateboarders", "number": 6}
|
||||
{"topic": "studying abroad", "number": 7}
|
||||
{"topic": "planting a vegetable garden", "number": 5}
|
||||
{"topic": "adopting a shelter pet", "number": 9}
|
||||
{"topic": "learning to cook ethnic cuisines", "number": 8}
|
||||
{"topic": "effective conflict resolution", "number": 5}
|
||||
{"topic": "starting a vlog", "number": 7}
|
||||
{"topic": "keeping a daily journal", "number": 6}
|
||||
{"topic": "improving sleep hygiene", "number": 8}
|
||||
{"topic": "beginner mountain climbers", "number": 5}
|
||||
{"topic": "creating a mobile app", "number": 9}
|
||||
{"topic": "maintaining a saltwater aquarium", "number": 7}
|
||||
{"topic": "preparing for a baby's arrival", "number": 6}
|
||||
{"topic": "writing a fantasy novel", "number": 5}
|
||||
{"topic": "effective team leadership", "number": 8}
|
||||
{"topic": "making a documentary film", "number": 9}
|
||||
{"topic": "learning about historical events", "number": 7}
|
||||
{"topic": "baking gluten-free treats", "number": 6}
|
||||
{"topic": "improving mental arithmetic skills", "number": 5}
|
||||
{"topic": "building a treehouse", "number": 8}
|
||||
{"topic": "getting started with watercolor painting", "number": 9}
|
||||
{"topic": "creating a YouTube tutorial series", "number": 7}
|
||||
{"topic": "landscape photography", "number": 5}
|
||||
{"topic": "navigating cultural differences", "number": 6}
|
||||
{"topic": "preparing for a marathon", "number": 8}
|
||||
{"topic": "building an online business", "number": 9}
|
||||
{"topic": "learning to dance at home", "number": 5}
|
||||
{"topic": "self-publishing a book", "number": 7}
|
||||
{"topic": "starting an urban farm", "number": 6}
|
||||
{"topic": "improving your memory", "number": 8}
|
||||
{"topic": "creating a personal brand online", "number": 9}
|
||||
@@ -1,51 +0,0 @@
|
||||
## Download data
|
||||
```
|
||||
wget https://raw.githubusercontent.com/openai/grade-school-math/master/grade_school_math/data/test.jsonl
|
||||
```
|
||||
|
||||
## Run benchmark
|
||||
|
||||
NOTE: This is an implementation for throughput/latency benchmark purposes. The prompts are not tuned to achieve good accuracy on the GSM-8K tasks.
|
||||
|
||||
### Benchmark sglang
|
||||
```
|
||||
python -m sglang.launch_server --model-path meta-llama/Llama-2-7b-chat-hf --port 30000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_sglang.py --num-questions 32
|
||||
python3 bench_sglang.py --num-questions 16 --parallel 1
|
||||
```
|
||||
|
||||
|
||||
### Benchmark vllm
|
||||
```
|
||||
python3 -m vllm.entrypoints.api_server --tokenizer-mode auto --model meta-llama/Llama-2-7b-chat-hf --disable-log-requests --port 21000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --num-questions 32 --backend vllm
|
||||
```
|
||||
|
||||
|
||||
### Benchmark lightllm
|
||||
```
|
||||
# A10G
|
||||
python -m lightllm.server.api_server --tokenizer_mode auto --model_dir ~/model_weights/llama-2-7b-chat-hf --max_total_token_num 16000 --port 22000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --num-questions 32 --backend lightllm
|
||||
```
|
||||
|
||||
|
||||
### Benchmark guidance
|
||||
```
|
||||
python3 bench_other.py --num-questions 8 --backend guidance --parallel 1 --n-ctx 4096 --model-path path/to/gguf
|
||||
```
|
||||
|
||||
### Benchmark lmql
|
||||
|
||||
```
|
||||
python3 bench_other.py --num-questions 8 --backend lmql --parallel 1
|
||||
```
|
||||
@@ -1,222 +0,0 @@
|
||||
import argparse
|
||||
import ast
|
||||
import json
|
||||
import re
|
||||
import time
|
||||
from collections import Counter
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
|
||||
from sglang.test.test_utils import add_common_other_args_and_parse, get_call_generate
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
INVALID = -9999999
|
||||
|
||||
|
||||
def get_answer_value(answer_str):
|
||||
answer_str = answer_str.replace(",", "")
|
||||
numbers = re.findall(r"\d+", answer_str)
|
||||
if len(numbers) < 1:
|
||||
return INVALID
|
||||
try:
|
||||
return ast.literal_eval(numbers[-1])
|
||||
except SyntaxError:
|
||||
return INVALID
|
||||
|
||||
|
||||
def most_frequent_number(numbers):
|
||||
if not numbers:
|
||||
return None
|
||||
|
||||
frequency = Counter(numbers)
|
||||
most_frequent = max(frequency, key=frequency.get)
|
||||
return most_frequent
|
||||
|
||||
|
||||
USER_PREFIX = "[INST] "
|
||||
USER_SUFFIX = " [/INST]"
|
||||
ASSISTANT_PREFIX = ""
|
||||
ASSISTANT_SUFFIX = " </s><s>"
|
||||
|
||||
# Use a low temp to make the results more deterministic and the comparison more fair.
|
||||
temp = 0.001
|
||||
|
||||
|
||||
def propose_plan(s, question, num_branches, call_generate):
|
||||
s += (
|
||||
USER_PREFIX
|
||||
+ """Please generate a high-level plan for solving the following question. As the first step, just say what method and idea you will use to solve the question. You can reorganize the information in the question. Do not do the actual calculation. Keep your response concise and within 80 words. Question: """
|
||||
+ question
|
||||
+ USER_SUFFIX
|
||||
)
|
||||
|
||||
s += ASSISTANT_PREFIX
|
||||
comps = call_generate(
|
||||
s, max_tokens=256, temperature=temp, stop=None, n=num_branches
|
||||
)
|
||||
return [s + comp + ASSISTANT_SUFFIX for comp in comps]
|
||||
|
||||
|
||||
def execute_plan(s, num_branches, call_generate):
|
||||
s += (
|
||||
USER_PREFIX
|
||||
+ """The plan looks good! Now, use real numbers and do the calculation. Please solve the question step-by-step according to the high-level plan. Give me the final answer. Make your response short."""
|
||||
+ USER_SUFFIX
|
||||
)
|
||||
s += ASSISTANT_PREFIX
|
||||
comps = call_generate(
|
||||
s, max_tokens=256, temperature=temp, stop=None, n=num_branches
|
||||
)
|
||||
return [s + comp + ASSISTANT_SUFFIX for comp in comps]
|
||||
|
||||
|
||||
def reflect_solution(s, num_branches, call_generate):
|
||||
s += (
|
||||
USER_PREFIX
|
||||
+ """Okay. Now, evaluate your own solution and give it a score on a scale of 1 to 5. Please do rigorous check of the correctness."""
|
||||
+ USER_SUFFIX
|
||||
)
|
||||
s += ASSISTANT_PREFIX
|
||||
comps = call_generate(
|
||||
s, max_tokens=256, temperature=temp, stop=None, n=num_branches
|
||||
)
|
||||
return [s + comp + ASSISTANT_SUFFIX for comp in comps]
|
||||
|
||||
|
||||
def get_final_answer(s, num_branches, call_generate):
|
||||
s += (
|
||||
USER_PREFIX
|
||||
+ """Based on your reflection, do you change your mind? Now, give me the final answer after careful consideration."""
|
||||
+ USER_SUFFIX
|
||||
)
|
||||
s += ASSISTANT_PREFIX
|
||||
comps = call_generate(
|
||||
s, max_tokens=256, temperature=temp, stop=None, n=num_branches
|
||||
)
|
||||
return [s + comp + ASSISTANT_SUFFIX for comp in comps]
|
||||
|
||||
|
||||
def tree_search(question, num_branches, call_generate):
|
||||
plan_forks = propose_plan("", question, num_branches, call_generate)
|
||||
|
||||
sol_states = []
|
||||
for plan in plan_forks:
|
||||
forks = execute_plan(plan, num_branches, call_generate)
|
||||
sol_states.extend(forks)
|
||||
|
||||
ref_states = []
|
||||
for sol in sol_states:
|
||||
forks = reflect_solution(sol, num_branches, call_generate)
|
||||
ref_states.extend(forks)
|
||||
|
||||
solutions = []
|
||||
for sol in ref_states:
|
||||
ans = get_final_answer(sol, num_branches, call_generate)
|
||||
solutions.append(ans)
|
||||
|
||||
return solutions
|
||||
|
||||
|
||||
def main(args):
|
||||
lines = read_jsonl(args.data_path)
|
||||
|
||||
# Construct prompts
|
||||
num_branches = 2
|
||||
questions = []
|
||||
labels = []
|
||||
for i in range(len(lines[: args.num_questions])):
|
||||
questions.append(lines[i]["question"])
|
||||
labels.append(get_answer_value(lines[i]["answer"]))
|
||||
assert all(l != INVALID for l in labels)
|
||||
arguments = [{"question": q, "num_branches": num_branches} for q in questions]
|
||||
|
||||
# Select backend
|
||||
call_generate = get_call_generate(args)
|
||||
|
||||
# Run requests
|
||||
states = [None] * len(questions)
|
||||
|
||||
tic = time.perf_counter()
|
||||
if args.backend != "lmql":
|
||||
|
||||
def get_one_answer(i):
|
||||
states[i] = tree_search(**arguments[i], call_generate=call_generate)
|
||||
|
||||
if args.parallel == 1:
|
||||
for i in tqdm(range(len(questions))):
|
||||
get_one_answer(i)
|
||||
else:
|
||||
with ThreadPoolExecutor(args.parallel) as executor:
|
||||
list(
|
||||
tqdm(
|
||||
executor.map(get_one_answer, list(range(len(questions)))),
|
||||
total=len(questions),
|
||||
)
|
||||
)
|
||||
|
||||
else:
|
||||
import asyncio
|
||||
|
||||
from lmql_funcs import tree_search_async
|
||||
|
||||
async def get_one_answer_async(i):
|
||||
states[i] = await tree_search_async(
|
||||
**arguments[i], call_generate=call_generate
|
||||
)
|
||||
|
||||
batches = [
|
||||
[] for _ in range((len(questions) + args.parallel - 1) // args.parallel)
|
||||
]
|
||||
for i in range(len(questions)):
|
||||
batches[i // args.parallel].append(i)
|
||||
|
||||
loop = asyncio.get_event_loop()
|
||||
for bt in tqdm(batches):
|
||||
tasks = [get_one_answer_async(k) for k in bt]
|
||||
loop.run_until_complete(asyncio.gather(*tasks))
|
||||
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
answers_text = []
|
||||
for s in states:
|
||||
answers_text.append([x for xs in s for x in xs])
|
||||
|
||||
preds = []
|
||||
for i in range(len(states)):
|
||||
answers = [get_answer_value(v) for v in answers_text[i]]
|
||||
preds.append(most_frequent_number(answers))
|
||||
|
||||
# Compute accuracy
|
||||
acc = np.mean(np.array(preds) == np.array(labels))
|
||||
invalid = np.mean(np.array(preds) == INVALID)
|
||||
print(f"Latency: {latency:.3f}")
|
||||
print(f"Invalid: {invalid:.3f}")
|
||||
print(f"Accuracy: {acc:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", answers_text)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "tree_of_thought_gsm8k",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"accuracy": round(acc, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"other": {
|
||||
"num_questions": args.num_questions,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="test.jsonl")
|
||||
parser.add_argument("--num-questions", type=int, default=200)
|
||||
args = add_common_other_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,171 +0,0 @@
|
||||
import argparse
|
||||
import ast
|
||||
import json
|
||||
import re
|
||||
import time
|
||||
from collections import Counter
|
||||
|
||||
import numpy as np
|
||||
|
||||
import sglang as sgl
|
||||
from sglang.test.test_utils import (
|
||||
add_common_sglang_args_and_parse,
|
||||
select_sglang_backend,
|
||||
)
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
INVALID = -9999999
|
||||
|
||||
|
||||
def get_answer_value(answer_str):
|
||||
answer_str = answer_str.replace(",", "")
|
||||
numbers = re.findall(r"\d+", answer_str)
|
||||
if len(numbers) < 1:
|
||||
return INVALID
|
||||
try:
|
||||
return ast.literal_eval(numbers[-1])
|
||||
except SyntaxError:
|
||||
return INVALID
|
||||
|
||||
|
||||
def most_frequent_number(numbers):
|
||||
if not numbers:
|
||||
return None
|
||||
|
||||
frequency = Counter(numbers)
|
||||
most_frequent = max(frequency, key=frequency.get)
|
||||
return most_frequent
|
||||
|
||||
|
||||
# Use a low temp to make the results more deterministic and the comparison more fair.
|
||||
temp = 0.001
|
||||
|
||||
|
||||
def propose_plan(s, question, num_branches):
|
||||
s += sgl.user(
|
||||
"""Please generate a high-level plan for solving the following question. As the first step, just say what method and idea you will use to solve the question. You can reorganize the information in the question. Do not do the actual calculation. Keep your response concise and within 80 words. Question: """
|
||||
+ question
|
||||
)
|
||||
forks = s.fork(num_branches)
|
||||
forks += sgl.assistant(sgl.gen("plan", max_tokens=256, temperature=temp))
|
||||
return forks
|
||||
|
||||
|
||||
def execute_plan(s, num_branches):
|
||||
s += sgl.user(
|
||||
"""The plan looks good! Now, use real numbers and do the calculation. Please solve the question step-by-step according to the high-level plan. Give me the final answer. Make your response short."""
|
||||
)
|
||||
forks = s.fork(num_branches)
|
||||
forks += sgl.assistant(sgl.gen("answer", max_tokens=256, temperature=temp))
|
||||
return forks
|
||||
|
||||
|
||||
def reflect_solution(s, num_branches):
|
||||
s += sgl.user(
|
||||
"""Okay. Now, evaluate your own solution and give it a score on a scale of 1 to 5. Please do rigorous check of the correctness."""
|
||||
)
|
||||
forks = s.fork(num_branches)
|
||||
forks += sgl.assistant(sgl.gen("score", max_tokens=256, temperature=temp))
|
||||
return forks
|
||||
|
||||
|
||||
def get_final_answer(s, num_branches):
|
||||
s += sgl.user(
|
||||
"""Based on your reflection, do you change your mind? Now, give me the final answer after careful consideration."""
|
||||
)
|
||||
forks = s.fork(num_branches)
|
||||
forks += sgl.assistant(sgl.gen("final_answer", max_tokens=256, temperature=temp))
|
||||
return forks
|
||||
|
||||
|
||||
@sgl.function
|
||||
def tree_search(s, question, num_branches):
|
||||
plan_forks = propose_plan(s, question, num_branches)
|
||||
|
||||
sol_states = []
|
||||
for plan in plan_forks:
|
||||
forks = execute_plan(plan, num_branches)
|
||||
sol_states.extend(forks)
|
||||
|
||||
ref_states = []
|
||||
for sol in sol_states:
|
||||
forks = reflect_solution(sol, num_branches)
|
||||
ref_states.extend(forks)
|
||||
|
||||
solutions = []
|
||||
for sol in ref_states:
|
||||
forks = get_final_answer(sol, num_branches)
|
||||
solutions.append(forks)
|
||||
solutions = [[s.text() for s in forks] for forks in solutions]
|
||||
|
||||
return solutions
|
||||
|
||||
|
||||
def main(args):
|
||||
lines = read_jsonl(args.data_path)
|
||||
lines = list(lines)
|
||||
|
||||
# Construct prompts
|
||||
num_branches = 2
|
||||
questions = []
|
||||
labels = []
|
||||
for i in range(len(lines[: args.num_questions])):
|
||||
questions.append(lines[i]["question"])
|
||||
labels.append(get_answer_value(lines[i]["answer"]))
|
||||
assert all(l != INVALID for l in labels)
|
||||
arguments = [{"question": q, "num_branches": num_branches} for q in questions]
|
||||
|
||||
# Select backend
|
||||
backend = select_sglang_backend(args)
|
||||
|
||||
# Run requests
|
||||
tic = time.perf_counter()
|
||||
states = tree_search.run_batch(
|
||||
arguments,
|
||||
temperature=0,
|
||||
backend=backend,
|
||||
num_threads=args.parallel,
|
||||
progress_bar=True,
|
||||
)
|
||||
latency = time.perf_counter() - tic
|
||||
answers_text = []
|
||||
for s in states:
|
||||
answers_text.append([x for xs in s.ret_value for x in xs])
|
||||
|
||||
preds = []
|
||||
for i in range(len(states)):
|
||||
answers = [get_answer_value(v) for v in answers_text[i]]
|
||||
preds.append(most_frequent_number(answers))
|
||||
|
||||
# Compute accuracy
|
||||
acc = np.mean(np.array(preds) == np.array(labels))
|
||||
invalid = np.mean(np.array(preds) == INVALID)
|
||||
print(f"Latency: {latency:.3f}")
|
||||
print(f"Invalid: {invalid:.3f}")
|
||||
print(f"Accuracy: {acc:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", answers_text)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "tree_of_thought_gsm8k",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"accuracy": round(acc, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"other": {
|
||||
"num_questions": args.num_questions,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="test.jsonl")
|
||||
parser.add_argument("--num-questions", type=int, default=200)
|
||||
args = add_common_sglang_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,82 +0,0 @@
|
||||
from bench_other import (
|
||||
ASSISTANT_PREFIX,
|
||||
ASSISTANT_SUFFIX,
|
||||
USER_PREFIX,
|
||||
USER_SUFFIX,
|
||||
temp,
|
||||
)
|
||||
|
||||
|
||||
async def propose_plan_async(s, question, num_branches, call_generate):
|
||||
s += (
|
||||
USER_PREFIX
|
||||
+ """Please generate a high-level plan for solving the following question. As the first step, just say what method and idea you will use to solve the question. You can reorganize the information in the question. Do not do the actual calculation. Keep your response concise and within 80 words. Question: """
|
||||
+ question
|
||||
+ USER_SUFFIX
|
||||
)
|
||||
|
||||
s += ASSISTANT_PREFIX
|
||||
comps = await call_generate(
|
||||
s, max_tokens=256, temperature=temp, stop=None, n=num_branches
|
||||
)
|
||||
return [s + comp + ASSISTANT_SUFFIX for comp in comps]
|
||||
|
||||
|
||||
async def execute_plan_async(s, num_branches, call_generate):
|
||||
s += (
|
||||
USER_PREFIX
|
||||
+ """The plan looks good! Now, use real numbers and do the calculation. Please solve the question step-by-step according to the high-level plan. Give me the final answer. Make your response short."""
|
||||
+ USER_SUFFIX
|
||||
)
|
||||
s += ASSISTANT_PREFIX
|
||||
comps = await call_generate(
|
||||
s, max_tokens=256, temperature=temp, stop=None, n=num_branches
|
||||
)
|
||||
return [s + comp + ASSISTANT_SUFFIX for comp in comps]
|
||||
|
||||
|
||||
async def reflect_solution_async(s, num_branches, call_generate):
|
||||
s += (
|
||||
USER_PREFIX
|
||||
+ """Okay. Now, evaluate your own solution and give it a score on a scale of 1 to 5. Please do rigorous check of the correctness."""
|
||||
+ USER_SUFFIX
|
||||
)
|
||||
s += ASSISTANT_PREFIX
|
||||
comps = await call_generate(
|
||||
s, max_tokens=256, temperature=temp, stop=None, n=num_branches
|
||||
)
|
||||
return [s + comp + ASSISTANT_SUFFIX for comp in comps]
|
||||
|
||||
|
||||
async def get_final_answer_async(s, num_branches, call_generate):
|
||||
s += (
|
||||
USER_PREFIX
|
||||
+ """Based on your reflection, do you change your mind? Now, give me the final answer after careful consideration."""
|
||||
+ USER_SUFFIX
|
||||
)
|
||||
s += ASSISTANT_PREFIX
|
||||
comps = await call_generate(
|
||||
s, max_tokens=256, temperature=temp, stop=None, n=num_branches
|
||||
)
|
||||
return [s + comp + ASSISTANT_SUFFIX for comp in comps]
|
||||
|
||||
|
||||
async def tree_search_async(question, num_branches, call_generate):
|
||||
plan_forks = await propose_plan_async("", question, num_branches, call_generate)
|
||||
|
||||
sol_states = []
|
||||
for plan in plan_forks:
|
||||
forks = await execute_plan_async(plan, num_branches, call_generate)
|
||||
sol_states.extend(forks)
|
||||
|
||||
ref_states = []
|
||||
for sol in sol_states:
|
||||
forks = await reflect_solution_async(sol, num_branches, call_generate)
|
||||
ref_states.extend(forks)
|
||||
|
||||
solutions = []
|
||||
for sol in ref_states:
|
||||
ans = await get_final_answer_async(sol, num_branches, call_generate)
|
||||
solutions.append(ans)
|
||||
|
||||
return solutions
|
||||
@@ -1,43 +0,0 @@
|
||||
## Download data
|
||||
```
|
||||
wget https://raw.githubusercontent.com/openai/grade-school-math/master/grade_school_math/data/test.jsonl
|
||||
```
|
||||
|
||||
## Run benchmark
|
||||
|
||||
### Benchmark sglang
|
||||
```
|
||||
python -m sglang.launch_server --model-path meta-llama/Llama-2-7b-chat-hf --port 30000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_sglang.py --num-questions 32 --parallel 16
|
||||
python3 bench_sglang.py --num-questions 10 --parallel 1
|
||||
```
|
||||
|
||||
|
||||
### Benchmark vllm
|
||||
```
|
||||
python3 -m vllm.entrypoints.api_server --tokenizer-mode auto --model meta-llama/Llama-2-7b-chat-hf --disable-log-requests --port 21000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --num-questions 32 --backend vllm
|
||||
```
|
||||
|
||||
|
||||
### Benchmark lightllm
|
||||
```
|
||||
# A10G
|
||||
python -m lightllm.server.api_server --tokenizer_mode auto --model_dir ~/model_weights/llama-2-7b-chat-hf --max_total_token_num 16000 --port 22000
|
||||
```
|
||||
|
||||
```
|
||||
python3 bench_other.py --num-questions 32 --backend lightllm
|
||||
```
|
||||
|
||||
|
||||
### Benchmark guidance
|
||||
```
|
||||
python3 bench_other.py --num-questions 32 --backend guidance --parallel 1 --n-ctx 4096 --model-path path/to/gguf
|
||||
```
|
||||
@@ -1,179 +0,0 @@
|
||||
import argparse
|
||||
import ast
|
||||
import json
|
||||
import re
|
||||
import time
|
||||
from collections import Counter
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
|
||||
from sglang.test.test_utils import add_common_other_args_and_parse, get_call_generate
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
INVALID = -9999999
|
||||
|
||||
|
||||
def get_answer_value(answer_str):
|
||||
answer_str = answer_str.replace(",", "")
|
||||
numbers = re.findall(r"\d+", answer_str)
|
||||
if len(numbers) < 1:
|
||||
return INVALID
|
||||
try:
|
||||
return ast.literal_eval(numbers[-1])
|
||||
except SyntaxError:
|
||||
return INVALID
|
||||
|
||||
|
||||
def most_frequent_number(numbers):
|
||||
if not numbers:
|
||||
return None
|
||||
|
||||
frequency = Counter(numbers)
|
||||
most_frequent = max(frequency, key=frequency.get)
|
||||
return most_frequent
|
||||
|
||||
|
||||
USER_PREFIX = "[INST] "
|
||||
USER_SUFFIX = " [/INST]"
|
||||
ASSISTANT_PREFIX = ""
|
||||
ASSISTANT_SUFFIX = " </s><s>"
|
||||
|
||||
# Use a low temp to make the results more deterministic and the comparison more fair.
|
||||
temp = 0.3
|
||||
|
||||
|
||||
def propose_plan(s, question, num_branches, call_generate):
|
||||
s += (
|
||||
USER_PREFIX
|
||||
+ """Please generate a high-level plan for solving the following question. As the first step, just say what method and idea you will use to solve the question. You can reorganize the information in the question. Do not do the actual calculation. Keep your response concise and within 80 words. Question: """
|
||||
+ question
|
||||
+ USER_SUFFIX
|
||||
)
|
||||
|
||||
s += ASSISTANT_PREFIX
|
||||
comps = call_generate(
|
||||
s, max_tokens=256, temperature=temp, stop=None, n=num_branches
|
||||
)
|
||||
return [s + comp + ASSISTANT_SUFFIX for comp in comps]
|
||||
|
||||
|
||||
def execute_plan(s, num_branches, call_generate):
|
||||
s += (
|
||||
USER_PREFIX
|
||||
+ """The plan looks good! Now, use real numbers and do the calculation. Please solve the question step-by-step according to the high-level plan. Give me the final answer. Make your response short."""
|
||||
+ USER_SUFFIX
|
||||
)
|
||||
s += ASSISTANT_PREFIX
|
||||
comps = call_generate(
|
||||
s, max_tokens=256, temperature=temp, stop=None, n=num_branches
|
||||
)
|
||||
return [s + comp + ASSISTANT_SUFFIX for comp in comps]
|
||||
|
||||
|
||||
def reflect_solution(s, num_branches, call_generate):
|
||||
s += (
|
||||
USER_PREFIX
|
||||
+ """Okay. Now you evaluate your own solution and give it a score on a scale of 1 to 5. Please do rigorous check of the correctness."""
|
||||
+ USER_SUFFIX
|
||||
)
|
||||
s += ASSISTANT_PREFIX
|
||||
comps = call_generate(
|
||||
s, max_tokens=256, temperature=temp, stop=None, n=num_branches
|
||||
)
|
||||
return [s + comp + ASSISTANT_SUFFIX for comp in comps]
|
||||
|
||||
|
||||
def tree_search(question, num_branches, call_generate):
|
||||
s = ""
|
||||
solutions = []
|
||||
|
||||
plan_forks = propose_plan(s, question, num_branches, call_generate)
|
||||
for plan in plan_forks:
|
||||
sol_forks = execute_plan(plan, num_branches, call_generate)
|
||||
for sol in sol_forks:
|
||||
score_forks = reflect_solution(sol, num_branches, call_generate)
|
||||
solutions.append(sol_forks)
|
||||
|
||||
return solutions
|
||||
|
||||
|
||||
def main(args):
|
||||
lines = read_jsonl(args.data_path)
|
||||
|
||||
# Construct prompts
|
||||
num_branches = 3
|
||||
questions = []
|
||||
labels = []
|
||||
for i in range(len(lines[: args.num_questions])):
|
||||
questions.append(lines[i]["question"])
|
||||
labels.append(get_answer_value(lines[i]["answer"]))
|
||||
assert all(l != INVALID for l in labels)
|
||||
arguments = [{"question": q, "num_branches": num_branches} for q in questions]
|
||||
|
||||
# Select backend
|
||||
call_generate = get_call_generate(args)
|
||||
|
||||
# Run requests
|
||||
states = [None] * len(questions)
|
||||
|
||||
def get_one_answer(i):
|
||||
states[i] = tree_search(**arguments[i], call_generate=call_generate)
|
||||
|
||||
tic = time.perf_counter()
|
||||
if args.parallel == 1:
|
||||
for i in tqdm(range(len(questions))):
|
||||
get_one_answer(i)
|
||||
else:
|
||||
with ThreadPoolExecutor(args.parallel) as executor:
|
||||
list(
|
||||
tqdm(
|
||||
executor.map(get_one_answer, list(range(len(questions)))),
|
||||
total=len(questions),
|
||||
)
|
||||
)
|
||||
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
answers_text = []
|
||||
for s in states:
|
||||
answers_text.append([x for xs in s for x in xs])
|
||||
|
||||
preds = []
|
||||
for i in range(len(states)):
|
||||
answers = [get_answer_value(v) for v in answers_text[i]]
|
||||
preds.append(most_frequent_number(answers))
|
||||
|
||||
# Compute accuracy
|
||||
acc = np.mean(np.array(preds) == np.array(labels))
|
||||
invalid = np.mean(np.array(preds) == INVALID)
|
||||
print(f"Latency: {latency:.3f}")
|
||||
print(f"Invalid: {invalid:.3f}")
|
||||
print(f"Accuracy: {acc:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", answers_text)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "tree_of_thought_gsm8k",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"accuracy": round(acc, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"other": {
|
||||
"num_questions": args.num_questions,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="test.jsonl")
|
||||
parser.add_argument("--num-questions", type=int, default=200)
|
||||
args = add_common_other_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -1,159 +0,0 @@
|
||||
import argparse
|
||||
import ast
|
||||
import json
|
||||
import re
|
||||
import time
|
||||
from collections import Counter
|
||||
|
||||
import numpy as np
|
||||
|
||||
import sglang as sgl
|
||||
from sglang.test.test_utils import (
|
||||
add_common_sglang_args_and_parse,
|
||||
select_sglang_backend,
|
||||
)
|
||||
from sglang.utils import dump_state_text, read_jsonl
|
||||
|
||||
INVALID = -9999999
|
||||
|
||||
|
||||
def get_answer_value(answer_str):
|
||||
answer_str = answer_str.replace(",", "")
|
||||
numbers = re.findall(r"\d+", answer_str)
|
||||
if len(numbers) < 1:
|
||||
return INVALID
|
||||
try:
|
||||
return ast.literal_eval(numbers[-1])
|
||||
except SyntaxError:
|
||||
return INVALID
|
||||
|
||||
|
||||
def most_frequent_number(numbers):
|
||||
if not numbers:
|
||||
return None
|
||||
|
||||
frequency = Counter(numbers)
|
||||
most_frequent = max(frequency, key=frequency.get)
|
||||
return most_frequent
|
||||
|
||||
|
||||
# Use a low temp to make the results more deterministic and the comparison more fair.
|
||||
temp = 0.3
|
||||
|
||||
|
||||
def propose_plan(s, question, num_branches):
|
||||
s += sgl.user(
|
||||
"""Please generate a high-level plan for solving the following question. As the first step, just say what method and idea you will use to solve the question. You can reorganize the information in the question. Do not do the actual calculation. Keep your response concise and within 80 words. Question: """
|
||||
+ question
|
||||
)
|
||||
forks = s.fork(num_branches)
|
||||
forks += sgl.assistant(sgl.gen("plan", max_tokens=256, temperature=temp))
|
||||
return forks
|
||||
|
||||
|
||||
def execute_plan(s, num_branches):
|
||||
s += sgl.user(
|
||||
"""The plan looks good! Now, use real numbers and do the calculation. Please solve the question step-by-step according to the high-level plan. Give me the final answer. Make your response short."""
|
||||
)
|
||||
forks = s.fork(num_branches)
|
||||
forks += sgl.assistant(sgl.gen("answer", max_tokens=256, temperature=temp))
|
||||
return forks
|
||||
|
||||
|
||||
def reflect_solution(s, num_branches):
|
||||
s += sgl.user(
|
||||
"""Okay. Now you evaluate your own solution and give it a score on a scale of 1 to 5. Please do rigorous check of the correctness."""
|
||||
)
|
||||
forks = s.fork(num_branches)
|
||||
forks += sgl.assistant(sgl.gen("score", max_tokens=256, temperature=temp))
|
||||
return forks
|
||||
|
||||
|
||||
@sgl.function
|
||||
def tree_search(s, question, num_branches):
|
||||
forks_to_join = []
|
||||
|
||||
plan_forks = propose_plan(s, question, num_branches)
|
||||
forks_to_join.append(plan_forks)
|
||||
|
||||
sol_states = []
|
||||
for plan in plan_forks:
|
||||
forks = execute_plan(plan, num_branches)
|
||||
forks_to_join.append(forks)
|
||||
sol_states.extend(forks)
|
||||
|
||||
for sol in sol_states:
|
||||
forks = reflect_solution(sol, num_branches)
|
||||
forks_to_join.append(forks)
|
||||
|
||||
for f in reversed(forks_to_join):
|
||||
f.join()
|
||||
|
||||
|
||||
def main(args):
|
||||
lines = read_jsonl(args.data_path)
|
||||
|
||||
# Construct prompts
|
||||
num_branches = 3
|
||||
questions = []
|
||||
labels = []
|
||||
for i in range(len(lines[: args.num_questions])):
|
||||
questions.append(lines[i]["question"])
|
||||
labels.append(get_answer_value(lines[i]["answer"]))
|
||||
assert all(l != INVALID for l in labels)
|
||||
arguments = [{"question": q, "num_branches": num_branches} for q in questions]
|
||||
|
||||
# Select backend
|
||||
backend = select_sglang_backend(args)
|
||||
|
||||
# Run requests
|
||||
tic = time.perf_counter()
|
||||
states = tree_search.run_batch(
|
||||
arguments,
|
||||
temperature=0,
|
||||
backend=backend,
|
||||
num_threads=args.parallel,
|
||||
progress_bar=True,
|
||||
)
|
||||
latency = time.perf_counter() - tic
|
||||
answers_text = []
|
||||
for s in states:
|
||||
answers_text.append([x for xs in s["answer"] for x in xs])
|
||||
|
||||
preds = []
|
||||
for i in range(len(states)):
|
||||
answers = [get_answer_value(v) for v in answers_text[i]]
|
||||
preds.append(most_frequent_number(answers))
|
||||
|
||||
# Compute accuracy
|
||||
acc = np.mean(np.array(preds) == np.array(labels))
|
||||
invalid = np.mean(np.array(preds) == INVALID)
|
||||
print(f"Latency: {latency:.3f}")
|
||||
print(f"Invalid: {invalid:.3f}")
|
||||
print(f"Accuracy: {acc:.3f}")
|
||||
|
||||
# Write results
|
||||
dump_state_text(f"tmp_output_{args.backend}.txt", answers_text)
|
||||
|
||||
with open(args.result_file, "a") as fout:
|
||||
value = {
|
||||
"task": "tree_of_thought_gsm8k",
|
||||
"backend": args.backend,
|
||||
"num_gpus": 1,
|
||||
"latency": round(latency, 3),
|
||||
"accuracy": round(acc, 3),
|
||||
"num_requests": args.num_questions,
|
||||
"other": {
|
||||
"num_questions": args.num_questions,
|
||||
"parallel": args.parallel,
|
||||
},
|
||||
}
|
||||
fout.write(json.dumps(value) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--data-path", type=str, default="test.jsonl")
|
||||
parser.add_argument("--num-questions", type=int, default=200)
|
||||
args = add_common_sglang_args_and_parse(parser)
|
||||
main(args)
|
||||
@@ -19,7 +19,7 @@ import time
|
||||
import unittest
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from datetime import datetime
|
||||
from functools import partial, wraps
|
||||
from functools import wraps
|
||||
from io import BytesIO
|
||||
from pathlib import Path
|
||||
from types import ModuleType, SimpleNamespace
|
||||
@@ -46,7 +46,7 @@ from sglang.srt.utils import (
|
||||
)
|
||||
from sglang.srt.utils.network import is_port_available
|
||||
from sglang.test.run_eval import run_eval
|
||||
from sglang.utils import get_exception_traceback, normalize_base_url
|
||||
from sglang.utils import normalize_base_url
|
||||
|
||||
# General test models
|
||||
DEFAULT_MODEL_NAME_FOR_TEST = "meta-llama/Llama-3.1-8B-Instruct"
|
||||
@@ -264,23 +264,6 @@ if is_in_ci() and is_xpu():
|
||||
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH = 1800
|
||||
|
||||
|
||||
def call_generate_lightllm(prompt, temperature, max_tokens, stop=None, url=None):
|
||||
assert url is not None
|
||||
|
||||
data = {
|
||||
"inputs": prompt,
|
||||
"parameters": {
|
||||
"temperature": temperature,
|
||||
"max_new_tokens": max_tokens,
|
||||
"stop_sequences": stop,
|
||||
},
|
||||
}
|
||||
res = requests.post(url, json=data)
|
||||
assert res.status_code == 200
|
||||
pred = res.json()["generated_text"][0]
|
||||
return pred
|
||||
|
||||
|
||||
def find_available_port(base_port: int):
|
||||
port = base_port + random.randint(100, 1000)
|
||||
while True:
|
||||
@@ -292,174 +275,6 @@ def find_available_port(base_port: int):
|
||||
port -= 43
|
||||
|
||||
|
||||
def call_generate_vllm(prompt, temperature, max_tokens, stop=None, n=1, url=None):
|
||||
assert url is not None
|
||||
|
||||
data = {
|
||||
"prompt": prompt,
|
||||
"temperature": temperature,
|
||||
"max_tokens": max_tokens,
|
||||
"stop": stop,
|
||||
"n": n,
|
||||
}
|
||||
res = requests.post(url, json=data)
|
||||
assert res.status_code == 200
|
||||
if n == 1:
|
||||
pred = res.json()["text"][0][len(prompt) :]
|
||||
else:
|
||||
pred = [x[len(prompt) :] for x in res.json()["text"]]
|
||||
return pred
|
||||
|
||||
|
||||
def call_generate_outlines(
|
||||
prompt, temperature, max_tokens, stop=None, regex=None, n=1, url=None
|
||||
):
|
||||
assert url is not None
|
||||
|
||||
data = {
|
||||
"prompt": prompt,
|
||||
"temperature": temperature,
|
||||
"max_tokens": max_tokens,
|
||||
"stop": stop,
|
||||
"regex": regex,
|
||||
"n": n,
|
||||
}
|
||||
res = requests.post(url, json=data)
|
||||
assert res.status_code == 200
|
||||
if n == 1:
|
||||
pred = res.json()["text"][0][len(prompt) :]
|
||||
else:
|
||||
pred = [x[len(prompt) :] for x in res.json()["text"]]
|
||||
return pred
|
||||
|
||||
|
||||
def call_generate_srt_raw(prompt, temperature, max_tokens, stop=None, url=None):
|
||||
assert url is not None
|
||||
|
||||
data = {
|
||||
"text": prompt,
|
||||
"sampling_params": {
|
||||
"temperature": temperature,
|
||||
"max_new_tokens": max_tokens,
|
||||
"stop": stop,
|
||||
},
|
||||
}
|
||||
res = requests.post(url, json=data)
|
||||
assert res.status_code == 200
|
||||
obj = res.json()
|
||||
pred = obj["text"]
|
||||
return pred
|
||||
|
||||
|
||||
def call_generate_guidance(
|
||||
prompt, temperature, max_tokens, stop=None, n=1, regex=None, model=None
|
||||
):
|
||||
assert model is not None
|
||||
from guidance import gen
|
||||
|
||||
rets = []
|
||||
for _ in range(n):
|
||||
out = (
|
||||
model
|
||||
+ prompt
|
||||
+ gen(
|
||||
name="answer",
|
||||
max_tokens=max_tokens,
|
||||
temperature=temperature,
|
||||
stop=stop,
|
||||
regex=regex,
|
||||
)
|
||||
)
|
||||
rets.append(out["answer"])
|
||||
return rets if n > 1 else rets[0]
|
||||
|
||||
|
||||
def call_select_lightllm(context, choices, url=None):
|
||||
assert url is not None
|
||||
|
||||
scores = []
|
||||
for i in range(len(choices)):
|
||||
data = {
|
||||
"inputs": context + choices[i],
|
||||
"parameters": {
|
||||
"max_new_tokens": 1,
|
||||
},
|
||||
}
|
||||
res = requests.post(url, json=data)
|
||||
assert res.status_code == 200
|
||||
scores.append(0)
|
||||
return np.argmax(scores)
|
||||
|
||||
|
||||
def call_select_vllm(context, choices, url=None):
|
||||
assert url is not None
|
||||
|
||||
scores = []
|
||||
for i in range(len(choices)):
|
||||
data = {
|
||||
"prompt": context + choices[i],
|
||||
"max_tokens": 1,
|
||||
"prompt_logprobs": 1,
|
||||
}
|
||||
res = requests.post(url, json=data)
|
||||
assert res.status_code == 200
|
||||
scores.append(res.json().get("prompt_score", 0))
|
||||
return np.argmax(scores)
|
||||
|
||||
"""
|
||||
Modify vllm/entrypoints/api_server.py
|
||||
|
||||
if final_output.prompt_logprobs is not None:
|
||||
score = np.mean([prob[t_id] for t_id, prob in zip(final_output.prompt_token_ids[1:], final_output.prompt_logprobs[1:])])
|
||||
ret["prompt_score"] = score
|
||||
"""
|
||||
|
||||
|
||||
def call_select_guidance(context, choices, model=None):
|
||||
assert model is not None
|
||||
from guidance import select
|
||||
|
||||
out = model + context + select(choices, name="answer")
|
||||
return choices.index(out["answer"])
|
||||
|
||||
|
||||
def add_common_other_args_and_parse(parser: argparse.ArgumentParser):
|
||||
parser.add_argument("--parallel", type=int, default=64)
|
||||
parser.add_argument("--host", type=str, default="127.0.0.1")
|
||||
parser.add_argument("--port", type=int, default=None)
|
||||
parser.add_argument(
|
||||
"--backend",
|
||||
type=str,
|
||||
required=True,
|
||||
choices=[
|
||||
"vllm",
|
||||
"outlines",
|
||||
"lightllm",
|
||||
"gserver",
|
||||
"guidance",
|
||||
"srt-raw",
|
||||
"llama.cpp",
|
||||
],
|
||||
)
|
||||
parser.add_argument("--n-ctx", type=int, default=4096)
|
||||
parser.add_argument(
|
||||
"--model-path", type=str, default="meta-llama/Llama-2-7b-chat-hf"
|
||||
)
|
||||
parser.add_argument("--result-file", type=str, default="result.jsonl")
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.port is None:
|
||||
default_port = {
|
||||
"vllm": 21000,
|
||||
"outlines": 21000,
|
||||
"lightllm": 22000,
|
||||
"srt-raw": 30000,
|
||||
"gserver": 9988,
|
||||
}
|
||||
args.port = default_port.get(args.backend, None)
|
||||
return args
|
||||
|
||||
|
||||
def auto_config_device() -> str:
|
||||
"""Auto-config available device platform"""
|
||||
|
||||
@@ -506,71 +321,6 @@ def select_sglang_backend(args: argparse.Namespace):
|
||||
return backend
|
||||
|
||||
|
||||
def _get_call_generate(args: argparse.Namespace):
|
||||
base_url = normalize_base_url(args.host, args.port)
|
||||
if args.backend == "lightllm":
|
||||
return partial(call_generate_lightllm, url=f"{base_url}/generate")
|
||||
elif args.backend == "vllm":
|
||||
return partial(call_generate_vllm, url=f"{base_url}/generate")
|
||||
elif args.backend == "srt-raw":
|
||||
return partial(call_generate_srt_raw, url=f"{base_url}/generate")
|
||||
elif args.backend == "outlines":
|
||||
return partial(call_generate_outlines, url=f"{base_url}/generate")
|
||||
elif args.backend == "guidance":
|
||||
from guidance import models
|
||||
|
||||
model = models.LlamaCpp(args.model_path, n_gpu_layers=-1, n_ctx=args.n_ctx)
|
||||
call_generate = partial(call_generate_guidance, model=model)
|
||||
call_generate("Hello,", 1.0, 8, ".")
|
||||
return call_generate
|
||||
else:
|
||||
raise ValueError(f"Invalid backend: {args.backend}")
|
||||
|
||||
|
||||
def _get_call_select(args: argparse.Namespace):
|
||||
base_url = normalize_base_url(args.host, args.port)
|
||||
if args.backend == "lightllm":
|
||||
return partial(call_select_lightllm, url=f"{base_url}/generate")
|
||||
elif args.backend == "vllm":
|
||||
return partial(call_select_vllm, url=f"{base_url}/generate")
|
||||
elif args.backend == "guidance":
|
||||
from guidance import models
|
||||
|
||||
model = models.LlamaCpp(args.model_path, n_gpu_layers=-1, n_ctx=args.n_ctx)
|
||||
call_select = partial(call_select_guidance, model=model)
|
||||
|
||||
call_select("Hello,", ["world", "earth"])
|
||||
return call_select
|
||||
else:
|
||||
raise ValueError(f"Invalid backend: {args.backend}")
|
||||
|
||||
|
||||
def get_call_generate(args: argparse.Namespace):
|
||||
call_generate = _get_call_generate(args)
|
||||
|
||||
def func(*args, **kwargs):
|
||||
try:
|
||||
return call_generate(*args, **kwargs)
|
||||
except Exception:
|
||||
print("Exception in call_generate:\n" + get_exception_traceback())
|
||||
raise
|
||||
|
||||
return func
|
||||
|
||||
|
||||
def get_call_select(args: argparse.Namespace):
|
||||
call_select = _get_call_select(args)
|
||||
|
||||
def func(*args, **kwargs):
|
||||
try:
|
||||
return call_select(*args, **kwargs)
|
||||
except Exception:
|
||||
print("Exception in call_select:\n" + get_exception_traceback())
|
||||
raise
|
||||
|
||||
return func
|
||||
|
||||
|
||||
def _get_default_models():
|
||||
import inspect
|
||||
|
||||
|
||||
Reference in New Issue
Block a user