[SPEC] feat: add adaptive speculative decoding metrics (#25940)
Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com> Co-authored-by: Jarrod Barnes <jbarnes850@gmail.com>
This commit is contained in:
co-authored by
Claude Opus 4.7
Jarrod Barnes
parent
106092123f
commit
a0670b5ba3
@@ -126,6 +126,12 @@ sglang:gen_throughput{model_name="meta-llama/Llama-3.1-8B-Instruct"} 86.50814177
|
|||||||
# HELP sglang:num_queue_reqs The number of requests in the waiting queue
|
# HELP sglang:num_queue_reqs The number of requests in the waiting queue
|
||||||
# TYPE sglang:num_queue_reqs gauge
|
# TYPE sglang:num_queue_reqs gauge
|
||||||
sglang:num_queue_reqs{model_name="meta-llama/Llama-3.1-8B-Instruct"} 2826.0
|
sglang:num_queue_reqs{model_name="meta-llama/Llama-3.1-8B-Instruct"} 2826.0
|
||||||
|
# HELP sglang:spec_num_steps Currently active speculative_num_steps.
|
||||||
|
# TYPE sglang:spec_num_steps gauge
|
||||||
|
sglang:spec_num_steps{model_name="meta-llama/Llama-3.1-8B-Instruct"} 3.0
|
||||||
|
# HELP sglang:spec_num_draft_tokens Currently active speculative_num_draft_tokens (decouples from steps under topk>1).
|
||||||
|
# TYPE sglang:spec_num_draft_tokens gauge
|
||||||
|
sglang:spec_num_draft_tokens{model_name="meta-llama/Llama-3.1-8B-Instruct"} 4.0
|
||||||
```
|
```
|
||||||
|
|
||||||
## Setup Guide
|
## Setup Guide
|
||||||
|
|||||||
@@ -315,6 +315,31 @@ class SchedulerMetricsReporter:
|
|||||||
var_decode_kv_tokens=decode_q.variance(),
|
var_decode_kv_tokens=decode_q.variance(),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _active_spec_config_snapshot(self) -> dict[str, int]:
|
||||||
|
"""Read the currently active speculative decoding configuration."""
|
||||||
|
draft_worker = self.scheduler.draft_worker
|
||||||
|
if draft_worker is None:
|
||||||
|
return {
|
||||||
|
"num_steps": 0,
|
||||||
|
"num_draft_tokens": 0,
|
||||||
|
}
|
||||||
|
|
||||||
|
# Fallback to server_args if draft_worker does not have the attributes.
|
||||||
|
server_args = self.scheduler.server_args
|
||||||
|
num_steps = getattr(
|
||||||
|
draft_worker, "speculative_num_steps", server_args.speculative_num_steps
|
||||||
|
)
|
||||||
|
num_draft_tokens = getattr(
|
||||||
|
draft_worker,
|
||||||
|
"speculative_num_draft_tokens",
|
||||||
|
server_args.speculative_num_draft_tokens,
|
||||||
|
)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"num_steps": num_steps or 0,
|
||||||
|
"num_draft_tokens": num_draft_tokens or 0,
|
||||||
|
}
|
||||||
|
|
||||||
def update_spec_metrics(self, bs: int, num_correct_drafts: int):
|
def update_spec_metrics(self, bs: int, num_correct_drafts: int):
|
||||||
self.spec_num_accept_tokens += num_correct_drafts + bs
|
self.spec_num_accept_tokens += num_correct_drafts + bs
|
||||||
self.spec_num_forward_ct += bs
|
self.spec_num_forward_ct += bs
|
||||||
@@ -666,6 +691,8 @@ class SchedulerMetricsReporter:
|
|||||||
iter_msg = f" [{batch_iter}]" if LOG_FORWARD_ITERS else ""
|
iter_msg = f" [{batch_iter}]" if LOG_FORWARD_ITERS else ""
|
||||||
msg = f"Decode batch{iter_msg}, #running-req: {num_running_reqs}, {token_usage_msg}"
|
msg = f"Decode batch{iter_msg}, #running-req: {num_running_reqs}, {token_usage_msg}"
|
||||||
|
|
||||||
|
spec_num_steps = 0
|
||||||
|
spec_num_draft_tokens = 0
|
||||||
if self.scheduler.spec_algorithm.is_none():
|
if self.scheduler.spec_algorithm.is_none():
|
||||||
spec_accept_length = 0
|
spec_accept_length = 0
|
||||||
spec_accept_rate = 0
|
spec_accept_rate = 0
|
||||||
@@ -686,6 +713,12 @@ class SchedulerMetricsReporter:
|
|||||||
self.spec_total_num_forward_ct += self.spec_num_forward_ct
|
self.spec_total_num_forward_ct += self.spec_num_forward_ct
|
||||||
self.spec_num_accept_tokens = self.spec_num_forward_ct = 0
|
self.spec_num_accept_tokens = self.spec_num_forward_ct = 0
|
||||||
msg += f"accept len: {spec_accept_length:.2f}, accept rate: {spec_accept_rate:.2f}, "
|
msg += f"accept len: {spec_accept_length:.2f}, accept rate: {spec_accept_rate:.2f}, "
|
||||||
|
|
||||||
|
if self.current_scheduler_metrics_enabled:
|
||||||
|
spec_snapshot = self._active_spec_config_snapshot()
|
||||||
|
spec_num_steps = spec_snapshot["num_steps"]
|
||||||
|
spec_num_draft_tokens = spec_snapshot["num_draft_tokens"]
|
||||||
|
|
||||||
cache_hit_rate = 0.0
|
cache_hit_rate = 0.0
|
||||||
|
|
||||||
if self.scheduler.disaggregation_mode == DisaggregationMode.DECODE:
|
if self.scheduler.disaggregation_mode == DisaggregationMode.DECODE:
|
||||||
@@ -751,6 +784,8 @@ class SchedulerMetricsReporter:
|
|||||||
# Speculative decoding
|
# Speculative decoding
|
||||||
self.stats.spec_accept_length = spec_accept_length
|
self.stats.spec_accept_length = spec_accept_length
|
||||||
self.stats.spec_accept_rate = spec_accept_rate
|
self.stats.spec_accept_rate = spec_accept_rate
|
||||||
|
self.stats.spec_num_steps = spec_num_steps
|
||||||
|
self.stats.spec_num_draft_tokens = spec_num_draft_tokens
|
||||||
|
|
||||||
# Retract
|
# Retract
|
||||||
self.stats.num_retracted_reqs = self.num_retracted_reqs
|
self.stats.num_retracted_reqs = self.num_retracted_reqs
|
||||||
|
|||||||
@@ -109,6 +109,9 @@ class SchedulerStats:
|
|||||||
# Speculative decoding
|
# Speculative decoding
|
||||||
spec_accept_length: float = 0.0
|
spec_accept_length: float = 0.0
|
||||||
spec_accept_rate: float = 0.0
|
spec_accept_rate: float = 0.0
|
||||||
|
# Adaptive speculative decoding (currently active tier).
|
||||||
|
spec_num_steps: int = 0
|
||||||
|
spec_num_draft_tokens: int = 0
|
||||||
|
|
||||||
# Retract
|
# Retract
|
||||||
num_retracted_reqs: int = 0
|
num_retracted_reqs: int = 0
|
||||||
@@ -405,6 +408,18 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
|
|||||||
labelnames=labels.keys(),
|
labelnames=labels.keys(),
|
||||||
multiprocess_mode="mostrecent",
|
multiprocess_mode="mostrecent",
|
||||||
)
|
)
|
||||||
|
self.spec_num_steps = Gauge(
|
||||||
|
name="sglang:spec_num_steps",
|
||||||
|
documentation="Currently active speculative_num_steps.",
|
||||||
|
labelnames=labels.keys(),
|
||||||
|
multiprocess_mode="mostrecent",
|
||||||
|
)
|
||||||
|
self.spec_num_draft_tokens = Gauge(
|
||||||
|
name="sglang:spec_num_draft_tokens",
|
||||||
|
documentation="Currently active speculative_num_draft_tokens (decouples from steps under topk>1).",
|
||||||
|
labelnames=labels.keys(),
|
||||||
|
multiprocess_mode="mostrecent",
|
||||||
|
)
|
||||||
|
|
||||||
# =================================================================
|
# =================================================================
|
||||||
# Retract
|
# Retract
|
||||||
@@ -1248,6 +1263,8 @@ class SchedulerMetricsCollector(_StatLoggerDIMixin):
|
|||||||
# Speculative decoding
|
# Speculative decoding
|
||||||
self._log_gauge(self.spec_accept_length, stats.spec_accept_length)
|
self._log_gauge(self.spec_accept_length, stats.spec_accept_length)
|
||||||
self._log_gauge(self.spec_accept_rate, stats.spec_accept_rate)
|
self._log_gauge(self.spec_accept_rate, stats.spec_accept_rate)
|
||||||
|
self._log_gauge(self.spec_num_steps, stats.spec_num_steps)
|
||||||
|
self._log_gauge(self.spec_num_draft_tokens, stats.spec_num_draft_tokens)
|
||||||
|
|
||||||
# Retract
|
# Retract
|
||||||
self._log_gauge(self.num_retracted_reqs, stats.num_retracted_reqs)
|
self._log_gauge(self.num_retracted_reqs, stats.num_retracted_reqs)
|
||||||
|
|||||||
@@ -79,6 +79,7 @@ class TestAdaptiveSpeculativeServer(CustomTestCase):
|
|||||||
"--speculative-adaptive",
|
"--speculative-adaptive",
|
||||||
"--speculative-adaptive-config",
|
"--speculative-adaptive-config",
|
||||||
cls.adaptive_config_path,
|
cls.adaptive_config_path,
|
||||||
|
"--enable-metrics",
|
||||||
"--skip-server-warmup",
|
"--skip-server-warmup",
|
||||||
"--mem-fraction-static",
|
"--mem-fraction-static",
|
||||||
"0.7",
|
"0.7",
|
||||||
@@ -100,6 +101,24 @@ class TestAdaptiveSpeculativeServer(CustomTestCase):
|
|||||||
self.assertEqual(response.status_code, 200, response.text)
|
self.assertEqual(response.status_code, 200, response.text)
|
||||||
return response.json()["internal_states"][0]
|
return response.json()["internal_states"][0]
|
||||||
|
|
||||||
|
def _scrape_metric(self, name: str, **label_filter) -> float | None:
|
||||||
|
"""Return the value of a Prometheus sample line, or None if absent.
|
||||||
|
|
||||||
|
Matches a line whose metric name is exactly *name* (next char is '{'
|
||||||
|
or whitespace) and whose labels include every key=value in
|
||||||
|
*label_filter*.
|
||||||
|
"""
|
||||||
|
text = requests.get(self.base_url + "/metrics", timeout=30).text
|
||||||
|
for line in text.splitlines():
|
||||||
|
if line.startswith("#") or not line.startswith(name):
|
||||||
|
continue
|
||||||
|
rest = line[len(name) :]
|
||||||
|
if rest and rest[0] not in "{ ":
|
||||||
|
continue
|
||||||
|
if all(f'{k}="{v}"' in line for k, v in label_filter.items()):
|
||||||
|
return float(line.rsplit(" ", 1)[1])
|
||||||
|
return None
|
||||||
|
|
||||||
def _generate(self, prompt: str, max_new_tokens: int = 64) -> dict:
|
def _generate(self, prompt: str, max_new_tokens: int = 64) -> dict:
|
||||||
response = requests.post(
|
response = requests.post(
|
||||||
self.base_url + "/generate",
|
self.base_url + "/generate",
|
||||||
@@ -165,6 +184,21 @@ class TestAdaptiveSpeculativeServer(CustomTestCase):
|
|||||||
avg_accept_len = server_info["internal_states"][0]["avg_spec_accept_length"]
|
avg_accept_len = server_info["internal_states"][0]["avg_spec_accept_length"]
|
||||||
print(f"avg_spec_accept_length={avg_accept_len:.4f}")
|
print(f"avg_spec_accept_length={avg_accept_len:.4f}")
|
||||||
|
|
||||||
|
def test_adaptive_metrics_exposed(self):
|
||||||
|
"""After an upshift, the adaptive current-state gauges are scrapeable."""
|
||||||
|
state = self._drive_upshift()
|
||||||
|
self.assertEqual(state["speculative_num_steps"], 3, f"Never upshifted: {state}")
|
||||||
|
# One more decode so the reporter emits a fresh logging interval.
|
||||||
|
self._generate(HIGH_ACCEPT_PROMPT)
|
||||||
|
|
||||||
|
steps = self._scrape_metric("sglang:spec_num_steps")
|
||||||
|
draft_tokens = self._scrape_metric("sglang:spec_num_draft_tokens")
|
||||||
|
|
||||||
|
self.assertEqual(steps, 3.0, "spec_num_steps gauge missing or wrong")
|
||||||
|
self.assertEqual(
|
||||||
|
draft_tokens, 4.0, "spec_num_draft_tokens gauge missing or wrong"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user