Fix IMA with flashinfer + spec + topk & Add radix attention test cases for eagle (#13740)
This commit is contained in:
@@ -183,6 +183,25 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
|||||||
kv_indices,
|
kv_indices,
|
||||||
req_to_token.size(1),
|
req_to_token.size(1),
|
||||||
)
|
)
|
||||||
|
mask_numel = (
|
||||||
|
paged_kernel_lens_sum * self.draft_token_num
|
||||||
|
+ (self.draft_token_num**2) * batch_size
|
||||||
|
)
|
||||||
|
if self.custom_mask.numel() < mask_numel:
|
||||||
|
# FIXME(attn): temporary fix for custom mask padding with cuda graph
|
||||||
|
self.custom_mask = torch.cat(
|
||||||
|
[
|
||||||
|
self.custom_mask,
|
||||||
|
torch.full(
|
||||||
|
(mask_numel - self.custom_mask.numel(),),
|
||||||
|
True,
|
||||||
|
dtype=torch.bool,
|
||||||
|
device=device,
|
||||||
|
),
|
||||||
|
],
|
||||||
|
dim=0,
|
||||||
|
)
|
||||||
|
|
||||||
return kv_indices, cum_kv_seq_len, qo_indptr, self.custom_mask
|
return kv_indices, cum_kv_seq_len, qo_indptr, self.custom_mask
|
||||||
|
|
||||||
def verify(
|
def verify(
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ import requests
|
|||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
from sglang.test.few_shot_gsm8k import run_eval as run_gsm8k_eval
|
from sglang.test.few_shot_gsm8k import run_eval as run_gsm8k_eval
|
||||||
|
from sglang.test.kits.radix_cache_server_kit import run_radix_attention_test
|
||||||
from sglang.test.server_fixtures.eagle_fixture import EagleServerBase
|
from sglang.test.server_fixtures.eagle_fixture import EagleServerBase
|
||||||
from sglang.test.test_utils import (
|
from sglang.test.test_utils import (
|
||||||
DEFAULT_EAGLE_TARGET_MODEL_FOR_TEST,
|
DEFAULT_EAGLE_TARGET_MODEL_FOR_TEST,
|
||||||
@@ -39,6 +40,10 @@ class TestEAGLEServerBasic(EagleServerBase):
|
|||||||
for p in threads:
|
for p in threads:
|
||||||
p.join()
|
p.join()
|
||||||
|
|
||||||
|
def test_radix_attention(self):
|
||||||
|
run_radix_attention_test(self.base_url)
|
||||||
|
self.assertIsNone(self.process.poll())
|
||||||
|
|
||||||
def test_max_token_one(self):
|
def test_max_token_one(self):
|
||||||
requests.get(self.base_url + "/flush_cache")
|
requests.get(self.base_url + "/flush_cache")
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user