[Auto Sync] Update scheduler.py (20260104) (#16424)
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com> Co-authored-by: Hanming Lu <69857889+hanming-lu@users.noreply.github.com> Co-authored-by: Hanming Lu <hanming@x.ai> Co-authored-by: Jie Ren <no-replay@github.com>
This commit is contained in:
co-authored by
github-actions[bot]
Hanming Lu
Hanming Lu
Jie Ren
parent
7f6a678f8f
commit
012dc5866d
@@ -67,6 +67,7 @@ class BenchArgs:
|
|||||||
pydantic_result_filename: Optional[str] = None
|
pydantic_result_filename: Optional[str] = None
|
||||||
append_to_github_summary: bool = True
|
append_to_github_summary: bool = True
|
||||||
seed: int = 42
|
seed: int = 42
|
||||||
|
cache_hit_rate: float = 0.0
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def add_cli_args(parser: argparse.ArgumentParser):
|
def add_cli_args(parser: argparse.ArgumentParser):
|
||||||
@@ -143,6 +144,13 @@ class BenchArgs:
|
|||||||
help="Disable appending the output of this run to github ci summary",
|
help="Disable appending the output of this run to github ci summary",
|
||||||
)
|
)
|
||||||
parser.add_argument("--seed", type=int, default=BenchArgs.seed)
|
parser.add_argument("--seed", type=int, default=BenchArgs.seed)
|
||||||
|
parser.add_argument(
|
||||||
|
"--cache-hit-rate",
|
||||||
|
type=float,
|
||||||
|
default=BenchArgs.cache_hit_rate,
|
||||||
|
help="Cache hit rate for benchmarking (0.0-1.0). "
|
||||||
|
"0.0 means no cache hits (flush all), 0.4 means 40%% of input tokens are cached.",
|
||||||
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_cli_args(cls, args: argparse.Namespace):
|
def from_cli_args(cls, args: argparse.Namespace):
|
||||||
@@ -212,6 +220,57 @@ def launch_server_process(server_args: ServerArgs):
|
|||||||
raise TimeoutError("Server failed to start within the timeout period.")
|
raise TimeoutError("Server failed to start within the timeout period.")
|
||||||
|
|
||||||
|
|
||||||
|
def _warmup_cache(
|
||||||
|
url: str,
|
||||||
|
input_ids: List[List[int]],
|
||||||
|
input_len: int,
|
||||||
|
cache_hit_rate: float,
|
||||||
|
dataset_name: str = "random",
|
||||||
|
image_data: Optional[List] = None,
|
||||||
|
):
|
||||||
|
"""Warm up the cache by sending prefix tokens to populate the radix cache.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
url: Server URL
|
||||||
|
input_ids: List of input token id lists
|
||||||
|
input_len: Length of input tokens
|
||||||
|
cache_hit_rate: Fraction of input tokens to cache (0.0-1.0)
|
||||||
|
dataset_name: Name of the dataset (used to determine if image data should be included)
|
||||||
|
image_data: Optional image data for VLM models
|
||||||
|
"""
|
||||||
|
cached_token_len = int(input_len * cache_hit_rate)
|
||||||
|
if cached_token_len <= 0:
|
||||||
|
return
|
||||||
|
|
||||||
|
print(
|
||||||
|
f"Warming up cache with {cache_hit_rate*100:.1f}% hit rate "
|
||||||
|
f"({cached_token_len} tokens per request)"
|
||||||
|
)
|
||||||
|
# Create prefix input_ids for cache warming
|
||||||
|
cache_warmup_input_ids = [ids[:cached_token_len] for ids in input_ids]
|
||||||
|
cache_warmup_payload = {
|
||||||
|
"input_ids": cache_warmup_input_ids,
|
||||||
|
"sampling_params": {
|
||||||
|
"temperature": 0.0,
|
||||||
|
"max_new_tokens": 1, # Minimal output, just to populate cache
|
||||||
|
"ignore_eos": True,
|
||||||
|
},
|
||||||
|
"stream": False,
|
||||||
|
}
|
||||||
|
if dataset_name == "mmmu" and image_data is not None:
|
||||||
|
# include image data in cache warmup
|
||||||
|
cache_warmup_payload["image_data"] = image_data
|
||||||
|
|
||||||
|
warmup_response = requests.post(
|
||||||
|
url + "/generate",
|
||||||
|
json=cache_warmup_payload,
|
||||||
|
)
|
||||||
|
if warmup_response.status_code != 200:
|
||||||
|
print(f"Warning: Cache warmup request failed: {warmup_response.text}")
|
||||||
|
else:
|
||||||
|
print("Cache warmup completed")
|
||||||
|
|
||||||
|
|
||||||
def run_one_case(
|
def run_one_case(
|
||||||
url: str,
|
url: str,
|
||||||
batch_size: int,
|
batch_size: int,
|
||||||
@@ -232,6 +291,7 @@ def run_one_case(
|
|||||||
dataset_name: str = BenchArgs.dataset_name,
|
dataset_name: str = BenchArgs.dataset_name,
|
||||||
dataset_path: str = BenchArgs.dataset_path,
|
dataset_path: str = BenchArgs.dataset_path,
|
||||||
parallel_batch: bool = False,
|
parallel_batch: bool = False,
|
||||||
|
cache_hit_rate: float = BenchArgs.cache_hit_rate,
|
||||||
):
|
):
|
||||||
requests.post(url + "/flush_cache")
|
requests.post(url + "/flush_cache")
|
||||||
|
|
||||||
@@ -296,6 +356,17 @@ def run_one_case(
|
|||||||
|
|
||||||
payload["input_ids"] = input_ids
|
payload["input_ids"] = input_ids
|
||||||
|
|
||||||
|
# Warm up cache if cache_hit_rate > 0.0
|
||||||
|
if cache_hit_rate > 0.0:
|
||||||
|
_warmup_cache(
|
||||||
|
url=url,
|
||||||
|
input_ids=input_ids,
|
||||||
|
input_len=input_len,
|
||||||
|
cache_hit_rate=cache_hit_rate,
|
||||||
|
dataset_name=dataset_name,
|
||||||
|
image_data=payload.get("image_data"),
|
||||||
|
)
|
||||||
|
|
||||||
# Turn on profiler
|
# Turn on profiler
|
||||||
profile_link = None
|
profile_link = None
|
||||||
if profile:
|
if profile:
|
||||||
@@ -394,12 +465,28 @@ def should_skip_due_to_token_capacity(
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def should_skip_due_to_max_running_requests(
|
||||||
|
batch_size, skip_max_running_requests_threshold
|
||||||
|
):
|
||||||
|
if batch_size > skip_max_running_requests_threshold:
|
||||||
|
print(
|
||||||
|
"=" * 8
|
||||||
|
+ f"Skip benchmark {batch_size=} > {skip_max_running_requests_threshold=} due to max running requests limit."
|
||||||
|
+ "=" * 8
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
def get_report_summary(
|
def get_report_summary(
|
||||||
results: List[BenchOneCaseResult], bench_args: BenchArgs, server_args: ServerArgs
|
results: List[BenchOneCaseResult], bench_args: BenchArgs, server_args: ServerArgs
|
||||||
):
|
):
|
||||||
summary = (
|
summary = (
|
||||||
f"\nInput lens: {bench_args.input_len}. Output lens: {bench_args.output_len}.\n"
|
f"\nInput lens: {bench_args.input_len}. Output lens: {bench_args.output_len}."
|
||||||
)
|
)
|
||||||
|
if bench_args.cache_hit_rate > 0.0:
|
||||||
|
summary += f" Cache hit rate: {bench_args.cache_hit_rate*100:.1f}%."
|
||||||
|
summary += "\n"
|
||||||
|
|
||||||
if is_blackwell():
|
if is_blackwell():
|
||||||
hourly_cost_per_gpu = 4 # $4/hour for one B200
|
hourly_cost_per_gpu = 4 # $4/hour for one B200
|
||||||
@@ -479,6 +566,16 @@ def run_benchmark(server_args: ServerArgs, bench_args: BenchArgs):
|
|||||||
internal_state[0].get("memory_usage", {}).get("token_capacity", 1000000000)
|
internal_state[0].get("memory_usage", {}).get("token_capacity", 1000000000)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Get effective max running requests
|
||||||
|
max_running_requests_per_dp = internal_state[0].get(
|
||||||
|
"effective_max_running_requests_per_dp", -1
|
||||||
|
)
|
||||||
|
dp_size = server_info.get("dp_size", None) or 1
|
||||||
|
assert (
|
||||||
|
max_running_requests_per_dp > 0
|
||||||
|
), f"effective_max_running_requests_per_dp is not set, {max_running_requests_per_dp=}"
|
||||||
|
skip_max_running_requests_threshold = max_running_requests_per_dp * dp_size
|
||||||
|
|
||||||
# Warmup
|
# Warmup
|
||||||
if not bench_args.skip_warmup:
|
if not bench_args.skip_warmup:
|
||||||
print("=" * 8 + " Warmup Begin " + "=" * 8)
|
print("=" * 8 + " Warmup Begin " + "=" * 8)
|
||||||
@@ -509,7 +606,9 @@ def run_benchmark(server_args: ServerArgs, bench_args: BenchArgs):
|
|||||||
for bs, il, ol in itertools.product(
|
for bs, il, ol in itertools.product(
|
||||||
bench_args.batch_size, bench_args.input_len, bench_args.output_len
|
bench_args.batch_size, bench_args.input_len, bench_args.output_len
|
||||||
):
|
):
|
||||||
if should_skip_due_to_token_capacity(
|
if should_skip_due_to_max_running_requests(
|
||||||
|
bs, skip_max_running_requests_threshold
|
||||||
|
) or should_skip_due_to_token_capacity(
|
||||||
bs, il, ol, skip_token_capacity_threshold
|
bs, il, ol, skip_token_capacity_threshold
|
||||||
):
|
):
|
||||||
continue
|
continue
|
||||||
@@ -529,6 +628,7 @@ def run_benchmark(server_args: ServerArgs, bench_args: BenchArgs):
|
|||||||
dataset_name=bench_args.dataset_name,
|
dataset_name=bench_args.dataset_name,
|
||||||
dataset_path=bench_args.dataset_path,
|
dataset_path=bench_args.dataset_path,
|
||||||
parallel_batch=bench_args.parallel_batch,
|
parallel_batch=bench_args.parallel_batch,
|
||||||
|
cache_hit_rate=bench_args.cache_hit_rate,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -538,7 +638,9 @@ def run_benchmark(server_args: ServerArgs, bench_args: BenchArgs):
|
|||||||
for bs, il, ol in itertools.product(
|
for bs, il, ol in itertools.product(
|
||||||
bench_args.batch_size, bench_args.input_len, bench_args.output_len
|
bench_args.batch_size, bench_args.input_len, bench_args.output_len
|
||||||
):
|
):
|
||||||
if should_skip_due_to_token_capacity(
|
if should_skip_due_to_max_running_requests(
|
||||||
|
bs, skip_max_running_requests_threshold
|
||||||
|
) or should_skip_due_to_token_capacity(
|
||||||
bs, il, ol, skip_token_capacity_threshold
|
bs, il, ol, skip_token_capacity_threshold
|
||||||
):
|
):
|
||||||
continue
|
continue
|
||||||
@@ -561,6 +663,7 @@ def run_benchmark(server_args: ServerArgs, bench_args: BenchArgs):
|
|||||||
dataset_name=bench_args.dataset_name,
|
dataset_name=bench_args.dataset_name,
|
||||||
dataset_path=bench_args.dataset_path,
|
dataset_path=bench_args.dataset_path,
|
||||||
parallel_batch=bench_args.parallel_batch,
|
parallel_batch=bench_args.parallel_batch,
|
||||||
|
cache_hit_rate=bench_args.cache_hit_rate,
|
||||||
profile=bench_args.profile,
|
profile=bench_args.profile,
|
||||||
profile_steps=bench_args.profile_steps,
|
profile_steps=bench_args.profile_steps,
|
||||||
profile_by_stage=bench_args.profile_by_stage,
|
profile_by_stage=bench_args.profile_by_stage,
|
||||||
|
|||||||
@@ -2510,6 +2510,7 @@ class Scheduler(
|
|||||||
"token_capacity": int(self.max_total_num_tokens),
|
"token_capacity": int(self.max_total_num_tokens),
|
||||||
"graph": round(self.tp_worker.model_runner.graph_mem_usage, 2),
|
"graph": round(self.tp_worker.model_runner.graph_mem_usage, 2),
|
||||||
}
|
}
|
||||||
|
ret["effective_max_running_requests_per_dp"] = self.max_running_requests
|
||||||
|
|
||||||
if not self.spec_algorithm.is_none() and self.spec_total_num_forward_ct > 0:
|
if not self.spec_algorithm.is_none() and self.spec_total_num_forward_ct > 0:
|
||||||
ret["avg_spec_accept_length"] = (
|
ret["avg_spec_accept_length"] = (
|
||||||
|
|||||||
Reference in New Issue
Block a user