Fix IMA with flashinfer + spec + topk & Add radix attention test cases for eagle (#13740)

This commit is contained in:
Liangsheng Yin
2025-12-13 23:04:41 +08:00
committed by GitHub
parent 0c23331e2e
commit d977dd2e07
2 changed files with 24 additions and 0 deletions

View File

@@ -183,6 +183,25 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
kv_indices,
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
def verify(

View File

@@ -13,6 +13,7 @@ import requests
from sglang.srt.environ import envs
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.kits.radix_cache_server_kit import run_radix_attention_test
from sglang.test.server_fixtures.eagle_fixture import EagleServerBase
from sglang.test.test_utils import (
DEFAULT_EAGLE_TARGET_MODEL_FOR_TEST,
@@ -39,6 +40,10 @@ class TestEAGLEServerBasic(EagleServerBase):
for p in threads:
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):
requests.get(self.base_url + "/flush_cache")