[Spec][trtllm] use decode kernel for draft extend (#24566)
This commit is contained in:
@@ -848,7 +848,10 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
|||||||
|
|
||||||
page_table = self._get_layer_page_table(layer, forward_batch)
|
page_table = self._get_layer_page_table(layer, forward_batch)
|
||||||
|
|
||||||
if forward_batch.forward_mode.is_target_verify():
|
if (
|
||||||
|
forward_batch.forward_mode.is_target_verify()
|
||||||
|
or forward_batch.forward_mode.is_draft_extend_v2()
|
||||||
|
):
|
||||||
o = flashinfer.decode.trtllm_batch_decode_with_kv_cache(
|
o = flashinfer.decode.trtllm_batch_decode_with_kv_cache(
|
||||||
query=q,
|
query=q,
|
||||||
kv_cache=kv_cache,
|
kv_cache=kv_cache,
|
||||||
|
|||||||
@@ -464,6 +464,18 @@ def _warmup_cache(
|
|||||||
print("Cache warmup completed")
|
print("Cache warmup completed")
|
||||||
|
|
||||||
|
|
||||||
|
def _flush_cache_with_retry(url: str, endpoint: str, max_retries: int = 3):
|
||||||
|
"""Post to a cache flush endpoint with retries on failure."""
|
||||||
|
for attempt in range(max_retries):
|
||||||
|
response = requests.post(url + endpoint, timeout=DEFAULT_TIMEOUT)
|
||||||
|
if response.status_code == 200:
|
||||||
|
return
|
||||||
|
if attempt < max_retries - 1:
|
||||||
|
time.sleep(2)
|
||||||
|
else:
|
||||||
|
response.raise_for_status()
|
||||||
|
|
||||||
|
|
||||||
def run_one_case(
|
def run_one_case(
|
||||||
url: str,
|
url: str,
|
||||||
batch_size: int,
|
batch_size: int,
|
||||||
@@ -500,11 +512,9 @@ def run_one_case(
|
|||||||
):
|
):
|
||||||
if backend == "vllm":
|
if backend == "vllm":
|
||||||
# You need to have export VLLM_SERVER_DEV_MODE=1 in your environment to use this endpoint.
|
# You need to have export VLLM_SERVER_DEV_MODE=1 in your environment to use this endpoint.
|
||||||
response = requests.post(url + "/reset_prefix_cache", timeout=DEFAULT_TIMEOUT)
|
_flush_cache_with_retry(url, "/reset_prefix_cache")
|
||||||
response.raise_for_status()
|
|
||||||
else:
|
else:
|
||||||
response = requests.post(url + "/flush_cache", timeout=DEFAULT_TIMEOUT)
|
_flush_cache_with_retry(url, "/flush_cache")
|
||||||
response.raise_for_status()
|
|
||||||
|
|
||||||
# Load input token ids via bench_serving.get_dataset
|
# Load input token ids via bench_serving.get_dataset
|
||||||
supported_datasets = ("random", "mmmu", "generated-shared-prefix")
|
supported_datasets = ("random", "mmmu", "generated-shared-prefix")
|
||||||
|
|||||||
Reference in New Issue
Block a user