benchmark/lora: make number of LoRA adapters configurable (#25363)

This commit is contained in:
Pai Liu
2026-05-20 18:23:13 -07:00
committed by GitHub
parent 643d44d699
commit cf1fd26d16
2 changed files with 68 additions and 14 deletions
+37 -10
View File
@@ -1,22 +1,25 @@
import argparse
import os
NUM_LORAS = 4
LORA_PATH = {
"base": "meta-llama/Llama-2-7b-hf",
"lora": "winddude/wizardLM-LlaMA-LoRA-7B",
}
DEFAULT_BASE_MODEL_PATH = "meta-llama/Llama-2-7b-hf"
DEFAULT_LORA_PATH = "winddude/wizardLM-LlaMA-LoRA-7B"
DEFAULT_NUM_LORAS = 4
def launch_server(args):
base_path = LORA_PATH["base"]
lora_path = LORA_PATH["lora"]
base_path = args.base_model_path
lora_path = args.lora_path
if args.base_only:
cmd = f"python3 -m sglang.launch_server --model {base_path} "
cmd = f"python3 -m sglang.launch_server --model-path {base_path} "
else:
cmd = f"python3 -m sglang.launch_server --model {base_path} --lora-paths "
for i in range(NUM_LORAS):
if args.num_loras <= 0:
raise ValueError(
"--num-loras must be greater than 0 unless --base-only is set"
)
cmd = f"python3 -m sglang.launch_server --model-path {base_path} --lora-paths "
for i in range(args.num_loras):
lora_name = f"lora{i}"
cmd += f"{lora_name}={lora_path} "
cmd += f"--disable-radix "
@@ -36,6 +39,30 @@ def launch_server(args):
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument(
"--base-model-path",
type=str,
default=DEFAULT_BASE_MODEL_PATH,
help="Base model path or Hugging Face model ID.",
)
parser.add_argument(
"--lora-path",
type=str,
default=DEFAULT_LORA_PATH,
help=(
"LoRA adapter path or Hugging Face model ID used for all registered "
"LoRA adapters."
),
)
parser.add_argument(
"--num-loras",
type=int,
default=DEFAULT_NUM_LORAS,
help=(
"Number of LoRA adapters to register. For example, 4 registers "
"lora0, lora1, lora2, and lora3."
),
)
parser.add_argument(
"--base-only",
action="store_true",
+31 -4
View File
@@ -25,7 +25,6 @@ from datetime import datetime
from typing import Any, Dict, List, Optional, Tuple
import numpy as np
from launch_server import LORA_PATH, NUM_LORAS
from tqdm.asyncio import tqdm
from transformers import PreTrainedTokenizerBase
@@ -39,6 +38,9 @@ from sglang.bench_serving import (
from sglang.benchmark.datasets.random import sample_random_requests
from sglang.benchmark.utils import get_tokenizer, remove_prefix
DEFAULT_BASE_MODEL_PATH = "meta-llama/Llama-2-7b-hf"
DEFAULT_NUM_LORAS = 4
global args
@@ -75,7 +77,7 @@ async def async_request_openai_completions(
payload = {
"text": prompt,
"sampling_params": {"max_new_tokens": request_func_input.output_len},
"lora_path": f"lora{random.randint(0, NUM_LORAS - 1)}",
"lora_path": f"lora{random.randint(0, args.num_loras - 1)}",
}
headers = {"Authorization": ""}
@@ -278,6 +280,9 @@ async def benchmark(
result = {
"backend": args.backend,
"request_rate": request_rate,
"base_model_path": args.base_model_path,
"base_only": args.base_only,
"num_loras": args.num_loras,
"total_input_tokens": metrics.total_input,
"total_output_tokens": metrics.total_output,
"total_output_tokens_retokenized": metrics.total_output_retokenized,
@@ -308,6 +313,11 @@ async def benchmark(
file.write(json.dumps(result) + "\n")
result = {
"backend": args.backend,
"request_rate": request_rate,
"base_model_path": args.base_model_path,
"base_only": args.base_only,
"num_loras": args.num_loras,
"duration": benchmark_duration,
"completed": metrics.completed,
"total_input_tokens": metrics.total_input,
@@ -348,6 +358,8 @@ def run_benchmark(args_: argparse.Namespace):
set_ulimit()
random.seed(args.seed)
np.random.seed(args.seed)
if not args.base_only and args.num_loras <= 0:
raise ValueError("--num-loras must be greater than 0 unless --base-only is set")
# Set url
if args.port is None:
@@ -370,8 +382,8 @@ def run_benchmark(args_: argparse.Namespace):
# Read dataset
backend = args.backend
model_id = args.model = LORA_PATH["base"]
tokenizer_id = args.model
model_id = args.base_model_path
tokenizer_id = args.base_model_path
tokenizer = get_tokenizer(tokenizer_id)
@@ -411,6 +423,21 @@ def set_ulimit(target_soft_limit=65535):
if __name__ == "__main__":
parser = ArgumentParser(description="Benchmark the online lora serving throughput.")
parser.add_argument(
"--base-model-path",
type=str,
default=DEFAULT_BASE_MODEL_PATH,
help="Base model path or Hugging Face model ID.",
)
parser.add_argument(
"--num-loras",
type=int,
default=DEFAULT_NUM_LORAS,
help=(
"Number of LoRA adapters used by the benchmark. Must match the "
"server launcher."
),
)
parser.add_argument(
"--backend",
type=str,