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
@@ -1042,6 +1042,7 @@ def chunk_kda_fwd(
A_log: Optional[torch.Tensor] = None, A_log: Optional[torch.Tensor] = None,
dt_bias: Optional[torch.Tensor] = None, dt_bias: Optional[torch.Tensor] = None,
lower_bound: Optional[float] = None, lower_bound: Optional[float] = None,
output_intermediate_states: bool = False,
): ):
chunk_size = 64 chunk_size = 64
# Pre-compute chunk indices once and thread through all downstream kernels. # Pre-compute chunk indices once and thread through all downstream kernels.
@@ -1128,8 +1129,11 @@ def chunk_kda_fwd(
cu_seqlens=cu_seqlens, cu_seqlens=cu_seqlens,
chunk_indices=chunk_indices, chunk_indices=chunk_indices,
) )
del Aqk, v_new, h del Aqk, v_new
if output_intermediate_states:
return o, h
del h
return o return o
@@ -1147,6 +1151,7 @@ def chunk_kda(
A_log: Optional[torch.Tensor] = None, A_log: Optional[torch.Tensor] = None,
dt_bias: Optional[torch.Tensor] = None, dt_bias: Optional[torch.Tensor] = None,
lower_bound: Optional[float] = None, lower_bound: Optional[float] = None,
output_intermediate_states: bool = False,
**kwargs, **kwargs,
): ):
if scale is None: if scale is None:
@@ -1156,7 +1161,8 @@ def chunk_kda(
q = l2norm_fwd(q.contiguous()) q = l2norm_fwd(q.contiguous())
k = l2norm_fwd(k.contiguous()) k = l2norm_fwd(k.contiguous())
o = chunk_kda_fwd( # Returns o [B, T, H, V] when output_intermediate_states=False, or (o, h [B, NT, H, V, K]) when output_intermediate_states=True.
return chunk_kda_fwd(
q=q, q=q,
k=k, k=k,
v=v.contiguous(), v=v.contiguous(),
@@ -1169,5 +1175,5 @@ def chunk_kda(
A_log=A_log, A_log=A_log,
dt_bias=dt_bias, dt_bias=dt_bias,
lower_bound=lower_bound, lower_bound=lower_bound,
output_intermediate_states=output_intermediate_states,
) )
return o
@@ -1081,6 +1081,7 @@ _MAMBA_RADIX_CACHE_ARCHS = frozenset(
# delegates here. # delegates here.
_MAMBA_EXTRA_BUFFER_ARCHS = frozenset( _MAMBA_EXTRA_BUFFER_ARCHS = frozenset(
{ {
"KimiLinearForCausalLM",
"Qwen3_5ForConditionalGeneration", "Qwen3_5ForConditionalGeneration",
"Qwen3_5MoeForConditionalGeneration", "Qwen3_5MoeForConditionalGeneration",
"Qwen3NextForCausalLM", "Qwen3NextForCausalLM",
@@ -820,6 +820,10 @@ class HybridLinearAttnBackend(AttentionBackend):
or linear_attn_backend.needs_cpu_seq_lens or linear_attn_backend.needs_cpu_seq_lens
) )
@property
def data_type(self):
return self.full_attn_backend.data_type
def _is_full_attn( def _is_full_attn(
self, layer: Optional[RadixAttention], layer_id: Optional[int] = None self, layer: Optional[RadixAttention], layer_id: Optional[int] = None
) -> bool: ) -> bool:
@@ -254,6 +254,12 @@ class KDAAttnBackend(MambaAttnBackendBase):
def __init__(self, model_runner: ModelRunner): def __init__(self, model_runner: ModelRunner):
super().__init__(model_runner) super().__init__(model_runner)
# mamba_cache.conv is [..., kernel-1, dim] while conv_states_shape expects the window length (kernel-1) at shape[-1], hence the transpose.
self.conv_states_shape = (
model_runner.req_to_token_pool.mamba_pool.mamba_cache.conv[0]
.transpose(-1, -2)
.shape
)
decode_backend = get_linear_attn_decode_backend() decode_backend = get_linear_attn_decode_backend()
prefill_backend = get_linear_attn_prefill_backend() prefill_backend = get_linear_attn_prefill_backend()
# KDA FlashInfer speculative decode (target_verify) is linear-chain only -- # KDA FlashInfer speculative decode (target_verify) is linear-chain only --
@@ -274,9 +280,22 @@ class KDAAttnBackend(MambaAttnBackendBase):
self.req_to_token_pool.size, dtype=torch.int32, device=model_runner.device self.req_to_token_pool.size, dtype=torch.int32, device=model_runner.device
) )
def init_forward_metadata(self, forward_batch: ForwardBatch):
super().init_forward_metadata(forward_batch)
if self.forward_metadata.has_mamba_track_mask:
self.forward_metadata.mamba_track_mask_indices = (
forward_batch.mamba_track_mask.nonzero(as_tuple=True)[0]
)
self.forward_metadata.conv_states_mask_indices = (
forward_batch.mamba_track_indices[
self.forward_metadata.mamba_track_mask_indices
]
)
def forward_decode( def forward_decode(
self, self,
layer: RadixLinearAttention, layer: RadixLinearAttention,
forward_batch: ForwardBatch,
mixed_qkv: Union[torch.Tensor, Tuple[torch.Tensor, ...]], mixed_qkv: Union[torch.Tensor, Tuple[torch.Tensor, ...]],
a: torch.Tensor, a: torch.Tensor,
b: torch.Tensor, b: torch.Tensor,
@@ -316,7 +335,7 @@ class KDAAttnBackend(MambaAttnBackendBase):
"KDA packed decode requires one token per sequence (T=1): " "KDA packed decode requires one token per sequence (T=1): "
f"got {qkv.shape[0]} tokens for {cache_indices.shape[0]} requests." f"got {qkv.shape[0]} tokens for {cache_indices.shape[0]} requests."
) )
return self.kernel_dispatcher.packed_decode( core_attn_out = self.kernel_dispatcher.packed_decode(
mixed_qkv=qkv, mixed_qkv=qkv,
a=a, a=a,
b=b, b=b,
@@ -333,13 +352,17 @@ class KDAAttnBackend(MambaAttnBackendBase):
replayssm_write_pos=replayssm_write_pos, replayssm_write_pos=replayssm_write_pos,
replayssm_force_flush=replayssm_force_flush, replayssm_force_flush=replayssm_force_flush,
) )
self._track_mamba_state_decode(
forward_batch, conv_states, ssm_states, cache_indices
)
return core_attn_out
q, k, v = qkv.split([layer.q_dim, layer.k_dim, layer.v_dim], dim=-1) q, k, v = qkv.split([layer.q_dim, layer.k_dim, layer.v_dim], dim=-1)
q = q.unflatten(-1, (-1, layer.head_q_dim)).unsqueeze(0) # n (h d) -> 1 n h d q = q.unflatten(-1, (-1, layer.head_q_dim)).unsqueeze(0) # n (h d) -> 1 n h d
k = k.unflatten(-1, (-1, layer.head_k_dim)).unsqueeze(0) # n (h d) -> 1 n h d k = k.unflatten(-1, (-1, layer.head_k_dim)).unsqueeze(0) # n (h d) -> 1 n h d
v = v.unflatten(-1, (-1, layer.head_v_dim)).unsqueeze(0) # n (h d) -> 1 n h d v = v.unflatten(-1, (-1, layer.head_v_dim)).unsqueeze(0) # n (h d) -> 1 n h d
return self.kernel_dispatcher.decode( core_attn_out = self.kernel_dispatcher.decode(
q=q, q=q,
k=k, k=k,
v=v, v=v,
@@ -352,6 +375,12 @@ class KDAAttnBackend(MambaAttnBackendBase):
query_start_loc=query_start_loc, query_start_loc=query_start_loc,
) )
self._track_mamba_state_decode(
forward_batch, conv_states, ssm_states, cache_indices
)
return core_attn_out
def forward_extend( def forward_extend(
self, self,
layer: RadixLinearAttention, layer: RadixLinearAttention,
@@ -376,6 +405,11 @@ class KDAAttnBackend(MambaAttnBackendBase):
has_initial_state = forward_batch.extend_prefix_lens > 0 has_initial_state = forward_batch.extend_prefix_lens > 0
if self.forward_metadata.has_mamba_track_mask:
mamba_cache_params.conv[0][
self.forward_metadata.conv_states_mask_indices
] = mixed_qkv[self.forward_metadata.track_conv_indices]
splits = [layer.q_dim, layer.k_dim, layer.v_dim] splits = [layer.q_dim, layer.k_dim, layer.v_dim]
q, k, v = mixed_qkv.transpose(0, 1).split(splits, dim=0) q, k, v = mixed_qkv.transpose(0, 1).split(splits, dim=0)
q_conv_weight, k_conv_weight, v_conv_weight = layer.conv_weights.split( q_conv_weight, k_conv_weight, v_conv_weight = layer.conv_weights.split(
@@ -425,6 +459,7 @@ class KDAAttnBackend(MambaAttnBackendBase):
k = k.unflatten(-1, (-1, layer.head_k_dim)).unsqueeze(0) # n (h d) -> 1 n h d k = k.unflatten(-1, (-1, layer.head_k_dim)).unsqueeze(0) # n (h d) -> 1 n h d
v = v.unflatten(-1, (-1, layer.head_v_dim)).unsqueeze(0) # n (h d) -> 1 n h d v = v.unflatten(-1, (-1, layer.head_v_dim)).unsqueeze(0) # n (h d) -> 1 n h d
track_ssm = self.forward_metadata.has_mamba_track_mask
core_attn_out = self.kernel_dispatcher.extend( core_attn_out = self.kernel_dispatcher.extend(
q=q, q=q,
k=k, k=k,
@@ -441,7 +476,13 @@ class KDAAttnBackend(MambaAttnBackendBase):
# draft_extend_v2 must stay rollback-able, so kernels that commit state # draft_extend_v2 must stay rollback-able, so kernels that commit state
# in place (e.g. FlashKDA) must not run for it. # in place (e.g. FlashKDA) must not run for it.
is_spec_decode=forward_batch.forward_mode.is_draft_extend_v2(), is_spec_decode=forward_batch.forward_mode.is_draft_extend_v2(),
return_intermediate_states=track_ssm,
) )
if track_ssm:
core_attn_out, h = core_attn_out
self._track_mamba_state_extend(
forward_batch, h, ssm_states, self.forward_metadata
)
return core_attn_out return core_attn_out
@@ -105,6 +105,13 @@ class CuteDSLKDAKernel(LinearAttnKernelBase):
lower_bound: Optional[float] = None, lower_bound: Optional[float] = None,
**kwargs, **kwargs,
) -> torch.Tensor: ) -> torch.Tensor:
if kwargs.get("return_intermediate_states"):
raise NotImplementedError(
"CuteDSLKDAKernel.extend cannot return intermediate chunk "
"states required by mamba_radix_cache_strategy=extra_buffer; "
"use --linear-attn-prefill-backend triton or "
"--mamba-radix-cache-strategy no_buffer."
)
head_k_dim = k.shape[-1] head_k_dim = k.shape[-1]
self._ensure_extend_loaded(head_k_dim) self._ensure_extend_loaded(head_k_dim)
@@ -39,13 +39,15 @@ def _triton_fallback(
A_log=None, A_log=None,
dt_bias=None, dt_bias=None,
lower_bound=None, lower_bound=None,
return_intermediate_states=False,
): ):
"""Fall back to the Triton chunk_kda kernel (handles all preprocessing). """Fall back to the Triton chunk_kda kernel (handles all preprocessing).
`g` is the RAW gate; chunk_kda applies the gate activation internally when `g` is the RAW gate; chunk_kda applies the gate activation internally when
A_log is provided, so A_log/dt_bias/lower_bound must be threaded through too A_log is provided, so A_log/dt_bias/lower_bound must be threaded through too
-- otherwise the fallback silently skips activation. chunk_kda updates the -- otherwise the fallback silently skips activation. chunk_kda updates the
ssm state in-place via cache_indices and returns only the output tensor. ssm state in-place via cache_indices and returns only the output tensor
(or (output, h) when return_intermediate_states is set).
""" """
from sglang.kernels.ops.attention.fla.kda import chunk_kda from sglang.kernels.ops.attention.fla.kda import chunk_kda
@@ -62,6 +64,7 @@ def _triton_fallback(
A_log=A_log, A_log=A_log,
dt_bias=dt_bias, dt_bias=dt_bias,
lower_bound=lower_bound, lower_bound=lower_bound,
output_intermediate_states=return_intermediate_states,
) )
@@ -111,9 +114,10 @@ class FlashKDAKernel(LinearAttnKernelBase):
lower_bound: Optional[float] = None, lower_bound: Optional[float] = None,
extend_seq_lens_cpu: Optional[list] = None, extend_seq_lens_cpu: Optional[list] = None,
is_spec_decode: bool = False, is_spec_decode: bool = False,
return_intermediate_states: bool = False,
**kwargs, **kwargs,
) -> torch.Tensor: ) -> torch.Tensor:
if self._should_fall_back( if return_intermediate_states or self._should_fall_back(
lower_bound, is_spec_decode, query_start_loc, extend_seq_lens_cpu lower_bound, is_spec_decode, query_start_loc, extend_seq_lens_cpu
): ):
return _triton_fallback( return _triton_fallback(
@@ -128,6 +132,7 @@ class FlashKDAKernel(LinearAttnKernelBase):
A_log=A_log, A_log=A_log,
dt_bias=dt_bias, dt_bias=dt_bias,
lower_bound=lower_bound, lower_bound=lower_bound,
return_intermediate_states=return_intermediate_states,
) )
return self._flashkda_extend( return self._flashkda_extend(
@@ -202,6 +202,7 @@ class TritonKDAKernel(LinearAttnKernelBase):
A_log: Optional[torch.Tensor] = None, A_log: Optional[torch.Tensor] = None,
dt_bias: Optional[torch.Tensor] = None, dt_bias: Optional[torch.Tensor] = None,
lower_bound: Optional[float] = None, lower_bound: Optional[float] = None,
return_intermediate_states: bool = False,
**kwargs, **kwargs,
) -> torch.Tensor: ) -> torch.Tensor:
return chunk_kda( return chunk_kda(
@@ -217,4 +218,5 @@ class TritonKDAKernel(LinearAttnKernelBase):
A_log=A_log, A_log=A_log,
dt_bias=dt_bias, dt_bias=dt_bias,
lower_bound=lower_bound, lower_bound=lower_bound,
output_intermediate_states=return_intermediate_states,
) )
@@ -417,7 +417,9 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
"""Get the fill value for sequence lengths in CUDA graph.""" """Get the fill value for sequence lengths in CUDA graph."""
return 1 return 1
def init_mha_chunk_metadata(self, forward_batch: ForwardBatch) -> None: def init_mha_chunk_metadata(
self, forward_batch: ForwardBatch, disable_flashinfer_ragged: bool = False
) -> None:
has_prefix = any(forward_batch.extend_prefix_lens_cpu) has_prefix = any(forward_batch.extend_prefix_lens_cpu)
fallback_to_flashinfer_impl = ( fallback_to_flashinfer_impl = (
self.disable_chunked_prefix_cache and has_prefix self.disable_chunked_prefix_cache and has_prefix
@@ -11,6 +11,7 @@ class KLDivergenceMixin:
kl_div_max_samples: int = 32 kl_div_max_samples: int = 32
kl_div_prefill_max_new_tokens: int = 512 kl_div_prefill_max_new_tokens: int = 512
kl_div_decode_max_new_tokens: int = 512 kl_div_decode_max_new_tokens: int = 512
kl_div_trust_remote_code: bool = False
@classmethod @classmethod
def _build_acc_thresholds(cls, threshold): def _build_acc_thresholds(cls, threshold):
@@ -27,6 +28,7 @@ class KLDivergenceMixin:
model_name=cls.model, model_name=cls.model,
max_samples=cls.kl_div_max_samples, max_samples=cls.kl_div_max_samples,
max_new_tokens=cls.kl_div_prefill_max_new_tokens, max_new_tokens=cls.kl_div_prefill_max_new_tokens,
trust_remote_code=cls.kl_div_trust_remote_code,
) )
@classmethod @classmethod
@@ -39,4 +41,5 @@ class KLDivergenceMixin:
model_name=cls.model, model_name=cls.model,
max_samples=cls.kl_div_max_samples, max_samples=cls.kl_div_max_samples,
max_new_tokens=cls.kl_div_decode_max_new_tokens, 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( 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.""" """Get input_ids from LongBench V2 dataset with local caching."""
# Create cache key based on parameters # Create cache key based on parameters
@@ -67,7 +70,7 @@ def get_input_ids(
"Please install the 'datasets' package: pip install datasets" "Please install the 'datasets' package: pip install datasets"
) from exc ) 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)...") print(f"Downloading {num_samples} samples from LongBench V2 (streaming)...")
dataset = load_dataset( dataset = load_dataset(
@@ -183,12 +186,21 @@ def _extract_output_logprobs(result):
def test_input_output_logprobs_match_helper( 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 num_samples = DEFAULT_NUM_SAMPLES
if max_samples is not None and max_samples > num_samples: if max_samples is not None and max_samples > num_samples:
num_samples = max_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: if max_samples is not None:
input_ids = input_ids[:max_samples] input_ids = input_ids[:max_samples]
print(f"Running test_input_output_logprobs_match with {len(input_ids)} prompts") 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( 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() server_info = requests.get(base_url + "/server_info").json()
if server_info["disable_radix_cache"]: 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 num_samples = DEFAULT_NUM_SAMPLES
if max_samples is not None and max_samples > num_samples: if max_samples is not None and max_samples > num_samples:
num_samples = max_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: if max_samples is not None:
input_ids = input_ids[:max_samples] input_ids = input_ids[:max_samples]
print( 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( 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() server_info = requests.get(base_url + "/server_info").json()
if server_info["disable_radix_cache"]: 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: if max_samples is not None and max_samples > num_samples:
num_samples = max_samples num_samples = max_samples
first_turn_input_ids = get_input_ids( 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: if max_samples is not None:
first_turn_input_ids = first_turn_input_ids[:max_samples] 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) 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(",") comma_token_id = tokenizer.encode(",")
second_turn_input_ids = [ second_turn_input_ids = [
@@ -3,7 +3,11 @@ from types import SimpleNamespace
from sglang.srt.utils import kill_process_tree from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.eval_accuracy_kit import GSM8KMixin
from sglang.test.kits.kl_divergence_kit import KLDivergenceMixin
from sglang.test.kits.prefix_cache_branching_kit import PrefixCacheBranchingMixin
from sglang.test.run_eval import run_eval from sglang.test.run_eval import run_eval
from sglang.test.server_fixtures.default_fixture import DefaultServerBase
from sglang.test.test_utils import ( from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST, DEFAULT_URL_FOR_TEST,
@@ -11,13 +15,15 @@ from sglang.test.test_utils import (
popen_launch_server, popen_launch_server,
) )
register_cuda_ci(est_time=178, stage="base-b", runner_config="2-gpu-large") register_cuda_ci(est_time=600, stage="base-b", runner_config="2-gpu-large")
KIMI_LINEAR_MODEL = "moonshotai/Kimi-Linear-48B-A3B-Instruct"
class TestKimiLinear(CustomTestCase): class TestKimiLinear(CustomTestCase):
@classmethod @classmethod
def setUpClass(cls): def setUpClass(cls):
cls.model = "moonshotai/Kimi-Linear-48B-A3B-Instruct" cls.model = KIMI_LINEAR_MODEL
cls.base_url = DEFAULT_URL_FOR_TEST cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server( cls.process = popen_launch_server(
cls.model, cls.model,
@@ -45,5 +51,32 @@ class TestKimiLinear(CustomTestCase):
self.assertGreater(metrics["score"], 0.88) self.assertGreater(metrics["score"], 0.88)
class TestKimiLinearExtraBuffer(
GSM8KMixin, KLDivergenceMixin, PrefixCacheBranchingMixin, DefaultServerBase
):
"""Regression guard: KDA never wrote mamba track snapshots, so states
donated to the radix cache under mamba_radix_cache_strategy=extra_buffer
were garbage and prefix-cache hits restored wrong KDA state (GSM8K
0.150 pre-fix vs 0.895 post-fix). Pre-fix, launching KimiLinear with
extra_buffer also fails the arch allowlist assert."""
model = KIMI_LINEAR_MODEL
cache_chunk_size = 64
gsm8k_score_threshold = 0.88
kl_div_thres = 0.02
kl_div_trust_remote_code = True
other_args = [
"--trust-remote-code",
"--tp-size",
"2",
"--chunked-prefill-size",
"2048",
"--mamba-radix-cache-strategy",
"extra_buffer",
"--mamba-track-interval",
"2",
]
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()