Fix KDA prefix caching under mamba extra_buffer and enable it for kimi_linear (#31474)

This commit is contained in:
Yuhao Yang
2026-07-19 20:09:03 +08:00
committed by GitHub
parent 7a03d30149
commit a03ca46a28
11 changed files with 153 additions and 19 deletions
@@ -11,6 +11,7 @@ class KLDivergenceMixin:
kl_div_max_samples: int = 32
kl_div_prefill_max_new_tokens: int = 512
kl_div_decode_max_new_tokens: int = 512
kl_div_trust_remote_code: bool = False
@classmethod
def _build_acc_thresholds(cls, threshold):
@@ -27,6 +28,7 @@ class KLDivergenceMixin:
model_name=cls.model,
max_samples=cls.kl_div_max_samples,
max_new_tokens=cls.kl_div_prefill_max_new_tokens,
trust_remote_code=cls.kl_div_trust_remote_code,
)
@classmethod
@@ -39,4 +41,5 @@ class KLDivergenceMixin:
model_name=cls.model,
max_samples=cls.kl_div_max_samples,
max_new_tokens=cls.kl_div_decode_max_new_tokens,
trust_remote_code=cls.kl_div_trust_remote_code,
)
+39 -9
View File
@@ -28,7 +28,10 @@ def format_longbench_v2_example(example):
def get_input_ids(
tokenizer_path, max_prompt_tokens=DEFAULT_PROMPT_TOKENS, num_samples=None
tokenizer_path,
max_prompt_tokens=DEFAULT_PROMPT_TOKENS,
num_samples=None,
trust_remote_code=False,
):
"""Get input_ids from LongBench V2 dataset with local caching."""
# Create cache key based on parameters
@@ -67,7 +70,7 @@ def get_input_ids(
"Please install the 'datasets' package: pip install datasets"
) from exc
tokenizer = get_tokenizer(tokenizer_path)
tokenizer = get_tokenizer(tokenizer_path, trust_remote_code=trust_remote_code)
print(f"Downloading {num_samples} samples from LongBench V2 (streaming)...")
dataset = load_dataset(
@@ -183,12 +186,21 @@ def _extract_output_logprobs(result):
def test_input_output_logprobs_match_helper(
base_url, ACC_THRESHOLDS, model_name, max_samples=None, max_new_tokens=16000
base_url,
ACC_THRESHOLDS,
model_name,
max_samples=None,
max_new_tokens=16000,
trust_remote_code=False,
):
num_samples = DEFAULT_NUM_SAMPLES
if max_samples is not None and max_samples > num_samples:
num_samples = max_samples
input_ids = get_input_ids(tokenizer_path=model_name, num_samples=num_samples)
input_ids = get_input_ids(
tokenizer_path=model_name,
num_samples=num_samples,
trust_remote_code=trust_remote_code,
)
if max_samples is not None:
input_ids = input_ids[:max_samples]
print(f"Running test_input_output_logprobs_match with {len(input_ids)} prompts")
@@ -217,7 +229,12 @@ def test_input_output_logprobs_match_helper(
def test_input_output_logprobs_match_prefill_cache_hit_helper(
base_url, ACC_THRESHOLDS, model_name, max_samples=None, max_new_tokens=8192
base_url,
ACC_THRESHOLDS,
model_name,
max_samples=None,
max_new_tokens=8192,
trust_remote_code=False,
):
server_info = requests.get(base_url + "/server_info").json()
if server_info["disable_radix_cache"]:
@@ -227,7 +244,11 @@ def test_input_output_logprobs_match_prefill_cache_hit_helper(
num_samples = DEFAULT_NUM_SAMPLES
if max_samples is not None and max_samples > num_samples:
num_samples = max_samples
input_ids = get_input_ids(tokenizer_path=model_name, num_samples=num_samples)
input_ids = get_input_ids(
tokenizer_path=model_name,
num_samples=num_samples,
trust_remote_code=trust_remote_code,
)
if max_samples is not None:
input_ids = input_ids[:max_samples]
print(
@@ -271,7 +292,12 @@ def test_input_output_logprobs_match_prefill_cache_hit_helper(
def test_input_output_logprobs_match_decode_cache_hit_helper(
base_url, ACC_THRESHOLDS, model_name, max_samples=None, max_new_tokens=8192
base_url,
ACC_THRESHOLDS,
model_name,
max_samples=None,
max_new_tokens=8192,
trust_remote_code=False,
):
server_info = requests.get(base_url + "/server_info").json()
if server_info["disable_radix_cache"]:
@@ -282,7 +308,9 @@ def test_input_output_logprobs_match_decode_cache_hit_helper(
if max_samples is not None and max_samples > num_samples:
num_samples = max_samples
first_turn_input_ids = get_input_ids(
tokenizer_path=model_name, num_samples=num_samples
tokenizer_path=model_name,
num_samples=num_samples,
trust_remote_code=trust_remote_code,
)
if max_samples is not None:
first_turn_input_ids = first_turn_input_ids[:max_samples]
@@ -298,7 +326,9 @@ def test_input_output_logprobs_match_decode_cache_hit_helper(
)
assert len(results) == len(first_turn_input_ids)
tokenizer = get_tokenizer(tokenizer_name=model_name)
tokenizer = get_tokenizer(
tokenizer_name=model_name, trust_remote_code=trust_remote_code
)
comma_token_id = tokenizer.encode(",")
second_turn_input_ids = [