From e958f4561f93c1607f07fa83d27adfa042cf7ace Mon Sep 17 00:00:00 2001 From: roikoren755 <26850796+roikoren755@users.noreply.github.com> Date: Tue, 26 May 2026 11:03:29 +0300 Subject: [PATCH] [feat] Support `extra_buffer` in Mamba2-based models (#15829) Signed-off-by: Roi Koren --- docs/advanced_features/server_arguments.md | 2 +- .../sglang/srt/arg_groups/nemotron_h_hook.py | 3 +- .../attention/hybrid_linear_attn_backend.py | 98 +++++++---- .../srt/layers/attention/mamba/mamba.py | 16 +- .../layers/attention/mamba/mamba2_metadata.py | 10 ++ python/sglang/srt/managers/schedule_batch.py | 17 +- .../sglang/srt/mem_cache/mamba_radix_cache.py | 15 +- python/sglang/srt/models/falcon_h1.py | 1 + python/sglang/srt/models/granitemoehybrid.py | 1 + python/sglang/srt/models/nemotron_h.py | 1 + python/sglang/srt/server_args.py | 28 ++-- test/manual/models/test_granite_moe_hybrid.py | 33 ++++ .../models/test_nvidia_nemotron_nano_v2.py | 117 +++++++------ .../test_disaggregation_hybrid_attention.py | 154 +++++++++++++++++- .../unit/mem_cache/test_mamba_unittest.py | 10 +- .../test_unified_radix_cache_bench.py | 17 +- .../test_unified_radix_cache_unittest.py | 12 +- 17 files changed, 405 insertions(+), 130 deletions(-) create mode 100644 test/manual/models/test_granite_moe_hybrid.py diff --git a/docs/advanced_features/server_arguments.md b/docs/advanced_features/server_arguments.md index 067865d3a..8a88b0ff2 100644 --- a/docs/advanced_features/server_arguments.md +++ b/docs/advanced_features/server_arguments.md @@ -347,7 +347,7 @@ Please consult the documentation below and [server_args.py](https://github.com/s | `--max-mamba-cache-size` | The maximum size of the mamba cache. | `None` | Type: int | | `--mamba-ssm-dtype` | The data type of the SSM states in mamba cache. | `float32` | `float32`, `bfloat16`, `float16` | | `--mamba-full-memory-ratio` | The ratio of mamba state memory to full kv cache memory. | `0.9` | Type: float | -| `--mamba-scheduler-strategy` | The strategy to use for mamba scheduler. `auto` currently defaults to `no_buffer`. 1. `no_buffer` does not support overlap scheduler due to not allocating extra mamba state buffers. Branching point caching support is feasible but not implemented. 2. `extra_buffer` supports overlap schedule by allocating extra mamba state buffers to track mamba state for caching (mamba state usage per running req becomes `2x` for non-spec; `1+(1/(2+speculative_num_draft_tokens))x` for spec dec (e.g. 1.16x if speculative_num_draft_tokens==4)). 2a. `extra_buffer` is strictly better for non-KV-cache-bound cases; for KV-cache-bound cases, the tradeoff depends on whether enabling overlap outweighs reduced max running requests. 2b. mamba caching at radix cache branching point is strictly better than non-branch but requires kernel support (currently only FLA backend), currently only extra_buffer supports branching. | `auto` | `auto`, `no_buffer`, `extra_buffer` | +| `--mamba-scheduler-strategy` | The strategy to use for mamba scheduler. `auto` currently defaults to `no_buffer`. 1. `no_buffer` does not support overlap scheduler due to not allocating extra mamba state buffers. Branching point caching support is feasible but not implemented. 2. `extra_buffer` supports overlap schedule by allocating extra mamba state buffers to track mamba state for caching (mamba state usage per running req becomes `2x` for non-spec; `1+(1/(2+speculative_num_draft_tokens))x` for spec dec (e.g. 1.16x if speculative_num_draft_tokens==4)). 2a. `extra_buffer` is strictly better for non-KV-cache-bound cases; for KV-cache-bound cases, the tradeoff depends on whether enabling overlap outweighs reduced max running requests. 2b. mamba caching at radix cache branching point is strictly better than non-branch but requires kernel support, currently only extra_buffer supports branching. | `auto` | `auto`, `no_buffer`, `extra_buffer` | | `--mamba-track-interval` | The interval (in tokens) to track the mamba state during decode. Only used when `--mamba-scheduler-strategy` is `extra_buffer`. Must be divisible by page_size if set, and must be >= speculative_num_draft_tokens when using speculative decoding. | `256` | Type: int | ## Hierarchical cache diff --git a/python/sglang/srt/arg_groups/nemotron_h_hook.py b/python/sglang/srt/arg_groups/nemotron_h_hook.py index b0955a7aa..87b568f3c 100644 --- a/python/sglang/srt/arg_groups/nemotron_h_hook.py +++ b/python/sglang/srt/arg_groups/nemotron_h_hook.py @@ -41,9 +41,8 @@ def apply_nemotron_h_defaults(server_args: "ServerArgs", model_arch: str) -> Non server_args._handle_mamba_radix_cache( model_arch=model_arch, - support_mamba_cache=True, - support_mamba_cache_extra_buffer=False, sm100_default_attention_backend="flashinfer", + fallback_attention_backend="flashinfer", ) assert server_args.attention_backend != "triton", ( "NemotronHForCausalLM does not support triton attention backend," diff --git a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py index a388e9d33..a1876944d 100644 --- a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py @@ -23,12 +23,6 @@ from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.server_args import get_global_server_args from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput from sglang.srt.speculative.spec_info import SpecInput -from sglang.srt.utils import is_cpu - -if not is_cpu(): - from sglang.srt.layers.attention.fla.chunk_delta_h import ( - CHUNK_SIZE as FLA_CHUNK_SIZE, - ) logger = logging.getLogger(__name__) @@ -278,9 +272,9 @@ class MambaAttnBackendBase(AttentionBackend): After processing a prefill chunk, we need to save the last `conv_state_len` tokens of the processed region for prefix caching. - The key insight is that FLA (Flash Linear Attention) processes sequences in chunks - of FLA_CHUNK_SIZE. We only track the conv state up to the last complete chunk boundary - (aligned_len). + The key insight is that FLA (Flash Linear Attention) and Mamba2 processes sequences in chunks + of the chunk size (FLA_CHUNK_SIZE=64 for FLA, mamba_chunk_size for Mamba2). + We only track the conv state up to the last complete chunk boundary (aligned_len). start_indices is the starting token index of the conv state to track in this extend batch. indices include all pos to track in this extend batch, conv_state_len for each req that @@ -317,28 +311,33 @@ class MambaAttnBackendBase(AttentionBackend): Compute source and destination indices for tracking SSM states for prefix caching. After processing a prefill, we need to save the SSM recurrent state for prefix caching. - The FLA kernel outputs intermediate hidden states `h` at each chunk boundary, + The kernel outputs intermediate hidden states `h` at each chunk boundary, plus a `last_recurrent_state` at the end of the chunked prefill size. + The chunk size varies by model type: + - FLA models: FLA_CHUNK_SIZE (64) + - Mamba2 models: mamba_chunk_size (256) + The challenge is that sequences may or may not end on a chunk boundary: - - Aligned case (len % FLA_CHUNK_SIZE == 0): In this case, FLA will store the to-cache - state in the last_recurrent_state. - - Unaligned case (len % FLA_CHUNK_SIZE != 0): The last_recurrent_state includes the + - Aligned case (len % chunk_size == 0): The to-cache state is stored in + the last_recurrent_state. + - Unaligned case (len % chunk_size != 0): The last_recurrent_state includes the unaligned position, but we only want state up to the last chunk boundary. We must extract from the intermediate `h` tensor at the appropriate chunk index. We compute the src and dst indices for all requests that need to be cached (i.e. mamba_track_mask is True) based on the rule above. - For example: - 1. If chunked prefill length is < 64, then only final state has value. In this case we - cache `final` state. - 2. if chunked prefill length == 64, then only final state has value. In this case we - cache pos 64, from `final` state - 3. if chunked prefill length >64 and < 128, then both h and final state have value. - We cache pos 64 from `h` state - 4. if chunked prefill length ==128, then both h and final state have value. We cache - pos 128 from `final` state. Note `h` doesn't include the pos 128. + For example (assuming chunk_size=64): + 1. If chunked prefill length is < chunk_size, then only final state has value. + In this case we cache `final` state. + 2. If chunked prefill length == chunk_size, then only final state has value. + In this case we cache pos chunk_size, from `final` state. + 3. If chunked prefill length > chunk_size and < 2 * chunk_size, then both h and + final state have value. We cache pos chunk_size from `h` state. + 4. If chunked prefill length == 2 * chunk_size, then both h and final state have + value. We cache pos 2 * chunk_size from `final` state. Note `h` doesn't include + the final position. Returns: track_ssm_h_src: Source indices into the packed `h` tensor (for unaligned seqs) @@ -346,6 +345,7 @@ class MambaAttnBackendBase(AttentionBackend): track_ssm_final_src: Source indices into last_recurrent_state buffer (for aligned seqs) track_ssm_final_dst: Destination cache slot indices (for aligned seqs) """ + mamba_cache_chunk_size = get_global_server_args().mamba_cache_chunk_size # Move to CPU to avoid kernel launches for masking operations mamba_track_mask = forward_batch.mamba_track_mask.cpu() extend_seq_lens = forward_batch.extend_seq_lens.cpu() @@ -355,7 +355,10 @@ class MambaAttnBackendBase(AttentionBackend): prefix_lens = forward_batch.extend_prefix_lens.cpu() # Calculate the number of hidden states per request - num_h_states = (extend_seq_lens - 1) // FLA_CHUNK_SIZE + 1 + if isinstance(self, Mamba2AttnBackend): + num_h_states = extend_seq_lens // mamba_cache_chunk_size + else: + num_h_states = (extend_seq_lens - 1) // mamba_cache_chunk_size + 1 # Calculate the starting offset for each sequence in the packed batch track_ssm_src_offset = torch.zeros_like(num_h_states) @@ -368,17 +371,17 @@ class MambaAttnBackendBase(AttentionBackend): dst_masked = mamba_track_indices[mamba_track_mask] # Determine if the sequence ends at a chunk boundary - is_aligned = (lens_masked % FLA_CHUNK_SIZE) == 0 + is_aligned = (lens_masked % mamba_cache_chunk_size) == 0 # Case 1: Aligned. Use last_recurrent_state from ssm_states. track_ssm_final_src = mamba_cache_indices[mamba_track_mask][is_aligned] track_ssm_final_dst = dst_masked[is_aligned] # Case 2: Unaligned. Use intermediate state from h. - # TODO: if support FLA_CHUNK_SIZE % page size != 0, then need to modify this + # TODO: if support mamba_cache_chunk_size % page size != 0, then need to modify this not_aligned = ~is_aligned track_ssm_h_src = offset_masked[not_aligned] + ( - lens_masked[not_aligned] // FLA_CHUNK_SIZE + lens_masked[not_aligned] // mamba_cache_chunk_size ) track_ssm_h_dst = dst_masked[not_aligned] @@ -639,10 +642,10 @@ class MambaAttnBackendBase(AttentionBackend): """ Track and copy SSM states during extend for prefix caching. - After the FLA chunked prefill kernel runs, we need to save the SSM recurrent + After the chunked prefill kernel runs, we need to save the SSM recurrent state at the last chunk boundary so it can be reused for prefix caching. The source of the state depends on whether the sequence length is aligned - to FLA_CHUNK_SIZE. See `_init_track_ssm_indices` for more details on how + to the chunk size. See `_init_track_ssm_indices` for more details on how the source and destination indices are computed. Note: Conv state tracking for extend is handled separately via gather operations @@ -669,6 +672,17 @@ class Mamba2AttnBackend(MambaAttnBackendBase): config = model_runner.mamba2_config assert config is not None self.mamba_chunk_size = config.mamba_chunk_size + self.conv_states_shape = ( + model_runner.req_to_token_pool.mamba_pool.mamba_cache.conv[0].shape + ) + + if model_runner.server_args.enable_mamba_extra_buffer(): + assert ( + self.conv_states_shape[-1] < self.mamba_chunk_size + ), f"{self.conv_states_shape[-1]=} should be less than {self.mamba_chunk_size}" + assert ( + model_runner.server_args.mamba_track_interval >= self.mamba_chunk_size + ), f"mamba_track_interval ({model_runner.server_args.mamba_track_interval}) must be >= mamba_chunk_size ({self.mamba_chunk_size})" def init_forward_metadata(self, forward_batch: ForwardBatch): self._execute_deferred_mamba_cow_and_clear(forward_batch) @@ -726,20 +740,46 @@ class Mamba2AttnBackend(MambaAttnBackendBase): hidden_states: torch.Tensor, output: torch.Tensor, layer_id: int, + forward_batch: ForwardBatch, mup_vector: Optional[torch.Tensor] = None, use_triton_causal_conv: bool = False, ): assert isinstance(self.forward_metadata, Mamba2Metadata) layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer_id) - return mixer.forward( + intermediate_states = mixer.forward( hidden_states=hidden_states, output=output, layer_cache=layer_cache, metadata=self.forward_metadata, + forward_batch=forward_batch, mup_vector=mup_vector, use_triton_causal_conv=use_triton_causal_conv, ) + if forward_batch.mamba_track_mask is not None: + if ( + intermediate_states is not None + and forward_batch.mamba_track_mask is not None + and forward_batch.mamba_track_mask.any() + ): + self._track_mamba_state_extend( + forward_batch, + intermediate_states, + layer_cache.temporal, + self.forward_metadata, + ) + + if self.forward_metadata.num_decodes > 0: + num_decodes = self.forward_metadata.num_decodes + track_mamba_states_if_needed( + layer_cache.conv[0], + layer_cache.temporal, + self.forward_metadata.mamba_cache_indices[-num_decodes:], + forward_batch.mamba_track_mask[-num_decodes:], + forward_batch.mamba_track_indices[-num_decodes:], + num_decodes, + ) + def forward_decode(self, *args, **kwargs): raise NotImplementedError( "Mamba2AttnBackend's forward is called directly instead of through HybridLinearAttnBackend, as it supports mixed prefill and decode" diff --git a/python/sglang/srt/layers/attention/mamba/mamba.py b/python/sglang/srt/layers/attention/mamba/mamba.py index 1d48809ca..1c46a35d4 100644 --- a/python/sglang/srt/layers/attention/mamba/mamba.py +++ b/python/sglang/srt/layers/attention/mamba/mamba.py @@ -26,6 +26,7 @@ from sglang.srt.layers.linear import ( ) from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.mem_cache.memory_pool import MambaPool +from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_loader.weight_utils import ( composed_weight_loader, sharded_weight_loader, @@ -410,6 +411,7 @@ class MambaMixer2(torch.nn.Module): output: torch.Tensor, layer_cache: MambaPool.State, metadata: Mamba2Metadata, + forward_batch: ForwardBatch, mup_vector: Optional[torch.Tensor] = None, use_triton_causal_conv: bool = False, ): @@ -420,6 +422,7 @@ class MambaMixer2(torch.nn.Module): state_indices_tensor = metadata.mamba_cache_indices conv_state = layer_cache.conv[0] ssm_state = layer_cache.temporal + intermediate_states = None query_start_loc = metadata.query_start_loc @@ -517,6 +520,14 @@ class MambaMixer2(torch.nn.Module): x = hidden_states_B_C_p.transpose( 0, 1 ) # this is the form that causal-conv see + if ( + forward_batch.mamba_track_mask is not None + and forward_batch.mamba_track_mask.any() + and metadata.track_conv_indices is not None + ): + x_to_track = x[:, metadata.track_conv_indices].transpose(0, 1) + mask_indices = forward_batch.mamba_track_mask.nonzero(as_tuple=True)[0] + conv_state[forward_batch.mamba_track_indices[mask_indices]] = x_to_track ccfn = ( causal_conv1d_fn if not use_triton_causal_conv @@ -546,7 +557,7 @@ class MambaMixer2(torch.nn.Module): ) # NOTE: final output is an in-place update of out tensor - varlen_state = mamba_chunk_scan_combined( + intermediate_states, varlen_state = mamba_chunk_scan_combined( hidden_states_p.view( 1, num_prefill_tokens, self.num_heads // self.tp_size, self.head_dim ), @@ -565,6 +576,7 @@ class MambaMixer2(torch.nn.Module): initial_states=initial_states, return_varlen_states=True, return_final_states=False, + return_intermediate_states=True, dt_softplus=True, dt_limit=(0.0, float("inf")), out=preallocated_ssm_out_p.view( @@ -708,6 +720,8 @@ class MambaMixer2(torch.nn.Module): # 5. Final linear projection output[:num_actual_tokens], _ = self.out_proj(hidden_states) + return intermediate_states + @property def mamba_type(self) -> str: return "mamba2" diff --git a/python/sglang/srt/layers/attention/mamba/mamba2_metadata.py b/python/sglang/srt/layers/attention/mamba/mamba2_metadata.py index 7a805b56a..044cb8b14 100644 --- a/python/sglang/srt/layers/attention/mamba/mamba2_metadata.py +++ b/python/sglang/srt/layers/attention/mamba/mamba2_metadata.py @@ -171,6 +171,11 @@ class Mamba2Metadata(ForwardMetadata): retrieve_next_token=forward_metadata.retrieve_next_token, retrieve_next_sibling=forward_metadata.retrieve_next_sibling, retrieve_parent_token=forward_metadata.retrieve_parent_token, + track_conv_indices=forward_metadata.track_conv_indices, + track_ssm_h_src=forward_metadata.track_ssm_h_src, + track_ssm_h_dst=forward_metadata.track_ssm_h_dst, + track_ssm_final_src=forward_metadata.track_ssm_final_src, + track_ssm_final_dst=forward_metadata.track_ssm_final_dst, num_decodes=len(seq_lens), num_prefills=0, num_prefill_tokens=0, @@ -239,6 +244,11 @@ class Mamba2Metadata(ForwardMetadata): retrieve_next_token=forward_metadata.retrieve_next_token, retrieve_next_sibling=forward_metadata.retrieve_next_sibling, retrieve_parent_token=forward_metadata.retrieve_parent_token, + track_conv_indices=forward_metadata.track_conv_indices, + track_ssm_h_src=forward_metadata.track_ssm_h_src, + track_ssm_h_dst=forward_metadata.track_ssm_h_dst, + track_ssm_final_src=forward_metadata.track_ssm_final_src, + track_ssm_final_dst=forward_metadata.track_ssm_final_dst, num_prefills=num_prefills, num_prefill_tokens=num_prefill_tokens, num_decodes=num_decodes, diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index cda0d4513..de2585ea7 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -68,7 +68,6 @@ from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST, DisaggregationM from sglang.srt.distributed.parallel_state import get_tensor_model_parallel_rank from sglang.srt.dllm.mixin.req import ReqDllmMixin from sglang.srt.environ import envs -from sglang.srt.layers.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE from sglang.srt.managers.embed_types import PositionalEmbeds from sglang.srt.managers.scheduler_components.new_token_ratio_tracker import ( NewTokenRatioTracker, @@ -2090,18 +2089,19 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): self, req: Req, ) -> "_MambaRadixCacheV2TrackEntry": + mamba_cache_chunk_size = get_global_server_args().mamba_cache_chunk_size + def _force_track_h(i: int) -> int: - assert i % FLA_CHUNK_SIZE == 0 + assert i % mamba_cache_chunk_size == 0 # There are 3 cases for mamba_track_seqlen passed to mamba_track_seqlens_cpu: - # 1) aligned with FLA_CHUNK_SIZE-> retrieve from last_recurrent_state + # 1) aligned with mamba_cache_chunk_size-> retrieve from last_recurrent_state # a) is the last position -> retrieve from last_recurrent_state # b) is NOT the last position -> retrieve from h - # 2) unaligned with FLA_CHUNK_SIZE -> retrieve from h + # 2) unaligned with mamba_cache_chunk_size -> retrieve from h # Currently, the math calculation only supports case 1a and 2. So for 1b, we need to add 1 # to force the math calculation to retrieve the correct mamba state from h. return i + 1 - mamba_cache_chunk_size = get_global_server_args().mamba_cache_chunk_size mask = req.extend_input_len >= mamba_cache_chunk_size track_index = req.mamba_ping_pong_track_buffer[req.mamba_next_track_idx].item() mamba_track_seqlen = -1 @@ -2123,13 +2123,14 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): * mamba_cache_chunk_size ) - # mamba_track_fla_chunk_aligned is the aligned seqlen based on FLA_CHUNK_SIZE + # mamba_track_fla_chunk_aligned is the aligned seqlen based on mamba_cache_chunk_size # If mamba_track_fla_chunk_aligned != mamba_track_seqlen_aligned, which can be true when - # page_size > FLA_CHUNK_SIZE, we need to force the math calculation to retrieve the correct mamba state from h + # page_size > mamba_cache_chunk_size, we need to force the math calculation to retrieve the correct mamba state from h # by _force_track_h() mamba_track_fla_chunk_aligned = ( len(req.prefix_indices) - + (req.extend_input_len // FLA_CHUNK_SIZE) * FLA_CHUNK_SIZE + + (req.extend_input_len // mamba_cache_chunk_size) + * mamba_cache_chunk_size ) if mamba_track_fla_chunk_aligned != mamba_track_seqlen_aligned: # We want to track mamba_track_seqlen_aligned, and it's not the last position, diff --git a/python/sglang/srt/mem_cache/mamba_radix_cache.py b/python/sglang/srt/mem_cache/mamba_radix_cache.py index be209c6a9..5a5c7a7f3 100644 --- a/python/sglang/srt/mem_cache/mamba_radix_cache.py +++ b/python/sglang/srt/mem_cache/mamba_radix_cache.py @@ -28,7 +28,6 @@ import torch from numpy import float64 from sglang.srt.distributed import get_tensor_model_parallel_rank -from sglang.srt.layers.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE from sglang.srt.mem_cache.allocator import ( PagedTokenToKVPoolAllocator, TokenToKVPoolAllocator, @@ -426,6 +425,7 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): ) or isinstance(params.token_to_kv_pool_allocator, PagedTokenToKVPoolAllocator) self.req_to_token_pool: HybridReqToTokenPool = params.req_to_token_pool self.token_to_kv_pool_allocator = params.token_to_kv_pool_allocator + self.mamba_cache_chunk_size = get_global_server_args().mamba_cache_chunk_size self.page_size = params.page_size self.disable = params.disable @@ -637,7 +637,7 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): assert page_aligned_len == len( kv_indices - ), f"page_aligned_len != len(kv_indices), {page_aligned_len=}, {len(kv_indices)=}, {cache_len=}, {self.page_size=}, {FLA_CHUNK_SIZE=}" + ), f"page_aligned_len != len(kv_indices), {page_aligned_len=}, {len(kv_indices)=}, {cache_len=}, {self.page_size=}, {self.mamba_cache_chunk_size=}" page_aligned_token_ids = token_ids[:page_aligned_len] @@ -1041,14 +1041,11 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache): # Calculate the branching point. It is defined as the last aligned position that # does not have a mamba value. if len(value) > best_value_len: - mamba_cache_chunk_size = get_global_server_args().mamba_cache_chunk_size - mamba_cache_chunk_aligned_seqlen = ( - sum(len(v) for v in value) // mamba_cache_chunk_size - ) * mamba_cache_chunk_size + chunk_aligned_seqlen = ( + sum(len(v) for v in value) // self.mamba_cache_chunk_size + ) * self.mamba_cache_chunk_size mamba_branching_seqlen = ( - mamba_cache_chunk_aligned_seqlen - if mamba_cache_chunk_aligned_seqlen > 0 - else None + chunk_aligned_seqlen if chunk_aligned_seqlen > 0 else None ) else: mamba_branching_seqlen = None diff --git a/python/sglang/srt/models/falcon_h1.py b/python/sglang/srt/models/falcon_h1.py index b12986658..3be39824e 100644 --- a/python/sglang/srt/models/falcon_h1.py +++ b/python/sglang/srt/models/falcon_h1.py @@ -349,6 +349,7 @@ class FalconH1HybridAttentionDecoderLayer(nn.Module): hidden_states * self.ssm_in_multiplier, mamba_hidden_states, layer_id=self.layer_id, + forward_batch=forward_batch, mup_vector=self.mup_vector, ) mamba_hidden_states = mamba_hidden_states * self.ssm_out_multiplier diff --git a/python/sglang/srt/models/granitemoehybrid.py b/python/sglang/srt/models/granitemoehybrid.py index a26e71601..85385b4fd 100644 --- a/python/sglang/srt/models/granitemoehybrid.py +++ b/python/sglang/srt/models/granitemoehybrid.py @@ -148,6 +148,7 @@ class GraniteMoeHybridMambaDecoderLayer(nn.Module): layer_id=self.layer_idx, hidden_states=hidden_states, output=output, + forward_batch=forward_batch, use_triton_causal_conv=True, ) diff --git a/python/sglang/srt/models/nemotron_h.py b/python/sglang/srt/models/nemotron_h.py index b60ffa372..1840b0727 100644 --- a/python/sglang/srt/models/nemotron_h.py +++ b/python/sglang/srt/models/nemotron_h.py @@ -423,6 +423,7 @@ class NemotronHMambaDecoderLayer(nn.Module): layer_id=self.layer_id, hidden_states=hidden_states, output=output, + forward_batch=forward_batch, use_triton_causal_conv=True, ) return output diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 422a67f33..eb4041cef 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -2456,8 +2456,6 @@ class ServerArgs: ]: self._handle_mamba_radix_cache( model_arch=model_arch, - support_mamba_cache=True, - support_mamba_cache_extra_buffer=False, sm100_default_attention_backend="triton", ) @@ -2470,8 +2468,7 @@ class ServerArgs: if has_mamba: self._handle_mamba_radix_cache( model_arch=model_arch, - support_mamba_cache_extra_buffer=False, - sm100_default_attention_backend="triton", + sm100_default_attention_backend="flashinfer", ) elif model_arch in ["Lfm2ForCausalLM"]: @@ -2556,6 +2553,7 @@ class ServerArgs: support_mamba_cache: bool = True, support_mamba_cache_extra_buffer: bool = True, sm100_default_attention_backend: str = None, + fallback_attention_backend: str = "triton", ): if ( is_sm100_supported() @@ -2582,7 +2580,7 @@ class ServerArgs: if self.enable_mamba_extra_buffer(): # extra_buffer if self.disable_radix_cache: raise ValueError( - "mamba extra_buffer is not compatible with --disable-radix-cache " + "mamba extra_buffer is not compatible with --disable-radix-cache. " "Overlap scheduling is already supported with no_buffer + disable_radix_cache. " "Please use --mamba-scheduler-strategy no_buffer instead." ) @@ -2599,11 +2597,7 @@ class ServerArgs: assert ( self.mamba_track_interval % self.page_size == 0 ), f"mamba_track_interval {self.mamba_track_interval} must be divisible by page_size {self.page_size}" - assert ( - max(FLA_CHUNK_SIZE, self.page_size) - % min(FLA_CHUNK_SIZE, self.page_size) - == 0 - ), f"For SSM models with extra buffer, either FLA_CHUNK_SIZE or page_size must be divisible by the other, got {FLA_CHUNK_SIZE=}, {self.page_size=}" + assert self.mamba_cache_chunk_size is not None elif not self.disable_radix_cache: # no_buffer if self.page_size is not None and self.page_size != 1: logger.warning( @@ -2621,7 +2615,7 @@ class ServerArgs: if self.attention_backend == "trtllm_mha": logger.warning( "Disabling radix cache since trtllm_mha does not support page_size = 1, which is required by MambaRadixCache. " - "Try to use --attention-backend triton if radix cache is necessary." + f"Try to use --attention-backend {fallback_attention_backend} if radix cache is necessary." ) self.disable_radix_cache = True self.disable_overlap_schedule = False @@ -7047,9 +7041,17 @@ class ServerArgs: @property def mamba_cache_chunk_size(self) -> int: - # For mamba cache with extra buffer, the chunk size is the max of FLA_CHUNK_SIZE and page_size. + # For mamba cache with extra buffer, the chunk size is the max of FLA_CHUNK_SIZE + # (or mamba_chunk_size if it is defined in the model's config) and page_size. # It is used to determine the caching point in a sequence during prefill. - return max(FLA_CHUNK_SIZE, self.page_size) + if not hasattr(self, "_mamba_cache_chunk_size"): + hf_config = self.get_model_config().hf_config + chunk_size = getattr(hf_config, "mamba_chunk_size", FLA_CHUNK_SIZE) + assert ( + max(chunk_size, self.page_size) % min(chunk_size, self.page_size) == 0 + ), f"For SSM models, either chunk_size or page_size must be divisible by the other, got {chunk_size=}, {self.page_size=}" + self._mamba_cache_chunk_size = max(chunk_size, self.page_size) + return self._mamba_cache_chunk_size def check_server_args(self): # Check parallel size constraints diff --git a/test/manual/models/test_granite_moe_hybrid.py b/test/manual/models/test_granite_moe_hybrid.py new file mode 100644 index 000000000..a8e498523 --- /dev/null +++ b/test/manual/models/test_granite_moe_hybrid.py @@ -0,0 +1,33 @@ +import unittest + +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.server_fixtures.default_fixture import DefaultServerBase + +GRANITE_MOE_HYBRID_MODEL = "ibm-granite/granite-4.0-h-micro" + + +class TestGraniteMoeHybrid(GSM8KMixin, DefaultServerBase): + model = GRANITE_MOE_HYBRID_MODEL + gsm8k_accuracy_thres = 0.78 + + +class TestGraniteMoeHybridExtraBuffer( + GSM8KMixin, KLDivergenceMixin, PrefixCacheBranchingMixin, DefaultServerBase +): + model = GRANITE_MOE_HYBRID_MODEL + cache_chunk_size = 256 + gsm8k_accuracy_thres = 0.78 + kl_div_thres = 0.002 + kl_div_thres_prefill = 0.02 + other_args = [ + "--mem-fraction-static", + "0.8", + "--mamba-scheduler-strategy", + "extra_buffer", + ] + + +if __name__ == "__main__": + unittest.main() diff --git a/test/manual/models/test_nvidia_nemotron_nano_v2.py b/test/manual/models/test_nvidia_nemotron_nano_v2.py index 7a93af006..770f4f544 100644 --- a/test/manual/models/test_nvidia_nemotron_nano_v2.py +++ b/test/manual/models/test_nvidia_nemotron_nano_v2.py @@ -2,91 +2,104 @@ import unittest from sglang.srt.utils import is_blackwell 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.server_fixtures.default_fixture import DefaultServerBase +NVIDIA_NEMOTRON_NANO_V2_MODEL = "nvidia/NVIDIA-Nemotron-Nano-9B-v2" + class TestNvidiaNemotronNanoV2BF16(GSM8KMixin, DefaultServerBase): - model = "nvidia/NVIDIA-Nemotron-Nano-9B-v2" + model = NVIDIA_NEMOTRON_NANO_V2_MODEL gsm8k_accuracy_thres = 0.87 other_args = ["--max-mamba-cache-size", "256"] class TestNvidiaNemotronNanoV2BF16PP(GSM8KMixin, DefaultServerBase): - model = "nvidia/NVIDIA-Nemotron-Nano-9B-v2" + model = NVIDIA_NEMOTRON_NANO_V2_MODEL gsm8k_accuracy_thres = 0.87 other_args = ["--max-mamba-cache-size", "256", "--pp-size", "2"] -class TestNvidiaNemotronNanoV2FP8(GSM8KMixin, DefaultServerBase): +class TestNvidiaNemotronNanoV2BF16ExtraBuffer( + GSM8KMixin, KLDivergenceMixin, PrefixCacheBranchingMixin, DefaultServerBase +): + model = NVIDIA_NEMOTRON_NANO_V2_MODEL + cache_chunk_size = 256 gsm8k_accuracy_thres = 0.87 + kl_div_thres = 0.002 + kl_div_thres_prefill = 0.01 + other_args = [ + "--max-mamba-cache-size", + "256", + "--mem-fraction-static", + "0.8", + "--mamba-scheduler-strategy", + "extra_buffer", + ] + + +class TestNvidiaNemotronNanoV2FP8(GSM8KMixin, DefaultServerBase): model = "nvidia/NVIDIA-Nemotron-Nano-9B-v2-FP8" + gsm8k_accuracy_thres = 0.87 other_args = ["--max-mamba-cache-size", "256"] @unittest.skipIf(not is_blackwell(), "NVFP4 only supported on blackwell") class TestNvidiaNemotronNanoV2NVFP4(GSM8KMixin, DefaultServerBase): - gsm8k_accuracy_thres = 0.855 model = "nvidia/NVIDIA-Nemotron-Nano-9B-v2-NVFP4" + gsm8k_accuracy_thres = 0.855 other_args = ["--max-mamba-cache-size", "256"] -@unittest.skip( - "STANDALONE speculative decoding does not yet support target and draft models " - "with different hidden sizes (Nemotron-9B: 4480, Llama-3.2-1B: 2048)" -) +SPECULATIVE_DECODING_OTHER_ARGS = [ + "--speculative-algorithm", + "STANDALONE", + "--speculative-num-steps", + "2", + "--speculative-eagle-topk", + "3", + "--speculative-num-draft-tokens", + "5", + "--speculative-draft-model-path", + "meta-llama/Llama-3.2-1B", + "--speculative-draft-load-format", + "dummy", + "--max-running-requests", + "8", + "--max-total-tokens", + "2048", + "--json-model-override-args", + '{"vocab_size": 131072, "hidden_size": 4480}', +] + + class TestNvidiaNemotronNanoV2SpeculativeDecoding(GSM8KMixin, DefaultServerBase): + model = NVIDIA_NEMOTRON_NANO_V2_MODEL gsm8k_accuracy_thres = 0.87 - model = "nvidia/NVIDIA-Nemotron-Nano-9B-v2" - other_args = [ - "--speculative-algorithm", - "STANDALONE", - "--speculative-num-steps", - "2", - "--speculative-eagle-topk", - "3", - "--speculative-num-draft-tokens", - "5", - "--speculative-draft-model-path", - "meta-llama/Llama-3.2-1B", - "--speculative-draft-load-format", - "dummy", - "--max-running-requests", - "8", - "--max-total-tokens", - "2048", - "--json-model-override-args", - '{"vocab_size": 131072}', + other_args = SPECULATIVE_DECODING_OTHER_ARGS + [ + "--disable-radix-cache", + ] + + +class TestNvidiaNemotronNanoV2SpeculativeDecodingExtraBuffer( + GSM8KMixin, DefaultServerBase +): + model = NVIDIA_NEMOTRON_NANO_V2_MODEL + gsm8k_accuracy_thres = 0.87 + other_args = SPECULATIVE_DECODING_OTHER_ARGS + [ + "--mamba-scheduler-strategy", + "extra_buffer", ] -@unittest.skip( - "STANDALONE speculative decoding does not yet support target and draft models " - "with different hidden sizes (Nemotron-9B: 4480, Llama-3.2-1B: 2048)" -) class TestNvidiaNemotronNanoV2SpeculativeDecodingBF16Cache( GSM8KMixin, DefaultServerBase ): + model = NVIDIA_NEMOTRON_NANO_V2_MODEL gsm8k_accuracy_thres = 0.87 - model = "nvidia/NVIDIA-Nemotron-Nano-9B-v2" - other_args = [ - "--speculative-algorithm", - "STANDALONE", - "--speculative-num-steps", - "2", - "--speculative-eagle-topk", - "3", - "--speculative-num-draft-tokens", - "5", - "--speculative-draft-model-path", - "meta-llama/Llama-3.2-1B", - "--speculative-draft-load-format", - "dummy", - "--max-running-requests", - "8", - "--max-total-tokens", - "2048", - "--json-model-override-args", - '{"vocab_size": 131072}', + other_args = SPECULATIVE_DECODING_OTHER_ARGS + [ + "--disable-radix-cache", "--mamba-ssm-dtype", "bfloat16", ] diff --git a/test/registered/disaggregation/test_disaggregation_hybrid_attention.py b/test/registered/disaggregation/test_disaggregation_hybrid_attention.py index 240f6506e..a993f8b58 100644 --- a/test/registered/disaggregation/test_disaggregation_hybrid_attention.py +++ b/test/registered/disaggregation/test_disaggregation_hybrid_attention.py @@ -16,7 +16,7 @@ register_cuda_ci(est_time=310, stage="extra-b", runner_config="8-gpu-h200") @unittest.skipIf(is_in_ci(), "Temporarily disable the flaky test.") -class TestDisaggregationHybridAttentionMamba(PDDisaggregationServerBase): +class TestDisaggregationHybridAttentionGDN(PDDisaggregationServerBase): @classmethod def setUpClass(cls): super().setUpClass() @@ -88,7 +88,7 @@ class TestDisaggregationHybridAttentionMamba(PDDisaggregationServerBase): self.assertGreater(metrics["score"], 0.93) -class TestDisaggregationHybridAttentionMambaExtraBuffer(PDDisaggregationServerBase): +class TestDisaggregationHybridAttentionGDNExtraBuffer(PDDisaggregationServerBase): @classmethod def setUpClass(cls): super().setUpClass() @@ -165,7 +165,7 @@ class TestDisaggregationHybridAttentionMambaExtraBuffer(PDDisaggregationServerBa self.assertGreater(metrics["score"], 0.90) -class TestDisaggregationHybridAttentionMambaDPDecode(PDDisaggregationServerBase): +class TestDisaggregationHybridAttentionGDNDPDecode(PDDisaggregationServerBase): """Test with prefill tp=2 and decode tp=2/dp=2 with dp-attention enabled.""" @classmethod @@ -244,5 +244,153 @@ class TestDisaggregationHybridAttentionMambaDPDecode(PDDisaggregationServerBase) self.assertGreater(metrics["score"], 0.90) +class TestDisaggregationHybridAttentionMamba(PDDisaggregationServerBase): + @classmethod + def setUpClass(cls): + super().setUpClass() + cls.model = "nvidia/NVIDIA-Nemotron-Nano-9B-v2" + + # Non blocking start servers + cls.start_prefill() + cls.start_decode() + + # Block until both + cls.wait_server_ready(cls.prefill_url + "/health", process=cls.process_prefill) + cls.wait_server_ready(cls.decode_url + "/health", process=cls.process_decode) + + cls.launch_lb() + + @classmethod + def start_prefill(cls): + prefill_args = [ + "--trust-remote-code", + "--disaggregation-mode", + "prefill", + "--disaggregation-bootstrap-port", + cls.bootstrap_port, + "--tp", + "4", + ] + prefill_args += cls.transfer_backend + cls.rdma_devices + cls.process_prefill = popen_launch_pd_server( + cls.model, + cls.prefill_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=prefill_args, + ) + + @classmethod + def start_decode(cls): + decode_args = [ + "--trust-remote-code", + "--disaggregation-mode", + "decode", + "--disaggregation-bootstrap-port", + cls.bootstrap_port, + "--tp", + "4", + "--base-gpu-id", + "4", + ] + decode_args += cls.transfer_backend + cls.rdma_devices + cls.process_decode = popen_launch_pd_server( + cls.model, + cls.decode_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=decode_args, + ) + + def test_gsm8k(self): + args = SimpleNamespace( + base_url=self.base_url, + model=self.model, + eval_name="gsm8k", + api="completion", + max_tokens=512, + num_examples=200, + num_threads=128, + ) + metrics = run_eval(args) + print(f"Evaluation metrics: {metrics}") + + self.assertGreater(metrics["score"], 0.87) + + +class TestDisaggregationHybridAttentionMambaExtraBuffer(PDDisaggregationServerBase): + @classmethod + def setUpClass(cls): + super().setUpClass() + cls.model = "nvidia/NVIDIA-Nemotron-Nano-9B-v2" + + # Non blocking start servers + cls.start_prefill() + cls.start_decode() + + # Block until both + cls.wait_server_ready(cls.prefill_url + "/health", process=cls.process_prefill) + cls.wait_server_ready(cls.decode_url + "/health", process=cls.process_decode) + + cls.launch_lb() + + @classmethod + def start_prefill(cls): + prefill_args = [ + "--trust-remote-code", + "--disaggregation-mode", + "prefill", + "--disaggregation-bootstrap-port", + cls.bootstrap_port, + "--tp", + "4", + "--mamba-scheduler-strategy", + "extra_buffer", + ] + prefill_args += cls.transfer_backend + cls.rdma_devices + cls.process_prefill = popen_launch_pd_server( + cls.model, + cls.prefill_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=prefill_args, + ) + + @classmethod + def start_decode(cls): + decode_args = [ + "--trust-remote-code", + "--disaggregation-mode", + "decode", + "--disaggregation-bootstrap-port", + cls.bootstrap_port, + "--tp", + "4", + "--base-gpu-id", + "4", + "--mamba-scheduler-strategy", + "extra_buffer", + ] + decode_args += cls.transfer_backend + cls.rdma_devices + cls.process_decode = popen_launch_pd_server( + cls.model, + cls.decode_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=decode_args, + ) + + def test_gsm8k(self): + args = SimpleNamespace( + base_url=self.base_url, + model=self.model, + eval_name="gsm8k", + api="completion", + max_tokens=512, + num_examples=200, + num_threads=128, + ) + metrics = run_eval(args) + print(f"Evaluation metrics: {metrics}") + + self.assertGreater(metrics["score"], 0.87) + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/mem_cache/test_mamba_unittest.py b/test/registered/unit/mem_cache/test_mamba_unittest.py index df69be2b1..315fb69d3 100755 --- a/test/registered/unit/mem_cache/test_mamba_unittest.py +++ b/test/registered/unit/mem_cache/test_mamba_unittest.py @@ -6,6 +6,7 @@ import torch from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape from sglang.srt.disaggregation.kv_events import BlockRemoved, BlockStored from sglang.srt.environ import envs +from sglang.srt.layers.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE from sglang.srt.managers.schedule_batch import Req from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator from sglang.srt.mem_cache.base_prefix_cache import ( @@ -410,9 +411,12 @@ class TestMamba(unittest.TestCase): def _setup_tree_and_allocator(self, enable_kv_cache_events=False): """Helper to create a MambaRadixCache with allocator for testing.""" - set_global_server_args_for_scheduler( - ServerArgs(model_path="dummy", page_size=1) - ) + server_args = ServerArgs(model_path="dummy", page_size=1) + # MambaRadixCache reads mamba_cache_chunk_size, whose property otherwise + # loads the HF config for self.model_path — impossible for the dummy model. + # Mirror the property's default for a dummy HF config: FLA_CHUNK_SIZE. + server_args._mamba_cache_chunk_size = FLA_CHUNK_SIZE + set_global_server_args_for_scheduler(server_args) size = 128 dtype = torch.bfloat16 head_num = 2 diff --git a/test/registered/unit/mem_cache/test_unified_radix_cache_bench.py b/test/registered/unit/mem_cache/test_unified_radix_cache_bench.py index 0cd815718..7516bfa48 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_cache_bench.py +++ b/test/registered/unit/mem_cache/test_unified_radix_cache_bench.py @@ -21,6 +21,7 @@ import torch from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape from sglang.srt.environ import envs +from sglang.srt.layers.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator from sglang.srt.mem_cache.base_prefix_cache import ( DecLockRefParams, @@ -691,9 +692,11 @@ def run_all_benchmarks( if benchmarks is None or "all" in benchmarks: benchmarks = list(ALL_BENCHMARKS.keys()) - set_global_server_args_for_scheduler( - ServerArgs(model_path="dummy", page_size=page_size) - ) + server_args = ServerArgs(model_path="dummy", page_size=page_size) + # MambaRadixCache reads mamba_cache_chunk_size, whose property otherwise + # loads the HF config for self.model_path — impossible for the dummy model. + server_args._mamba_cache_chunk_size = max(FLA_CHUNK_SIZE, page_size) + set_global_server_args_for_scheduler(server_args) impl_name = (tree_cls or UnifiedRadixCache).__name__ results = [] @@ -780,9 +783,11 @@ class _BenchSuite: @classmethod def setUpClass(cls): - set_global_server_args_for_scheduler( - ServerArgs(model_path="dummy", page_size=cls.bench_cfg["page_size"]) - ) + page_size = cls.bench_cfg["page_size"] + server_args = ServerArgs(model_path="dummy", page_size=page_size) + # See run_all_benchmarks for why _mamba_cache_chunk_size is preset. + server_args._mamba_cache_chunk_size = max(FLA_CHUNK_SIZE, page_size) + set_global_server_args_for_scheduler(server_args) def _run(self, bench_fn): cfg = self.bench_cfg diff --git a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py index 1e7e51da1..819f2c9e6 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py +++ b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py @@ -10,6 +10,7 @@ import torch from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape from sglang.srt.environ import envs +from sglang.srt.layers.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE from sglang.srt.managers.schedule_batch import Req from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator from sglang.srt.mem_cache.base_prefix_cache import ( @@ -114,9 +115,12 @@ class CacheConfig: def build_fixture(cfg: CacheConfig): """Create (tree, allocator, req_to_token_pool) from a CacheConfig.""" - set_global_server_args_for_scheduler( - ServerArgs(model_path="dummy", page_size=cfg.page_size) - ) + server_args = ServerArgs(model_path="dummy", page_size=cfg.page_size) + # MambaRadixCache reads mamba_cache_chunk_size, whose property otherwise + # loads the HF config for self.model_path — impossible for the dummy model. + # Mirror the property's default for a dummy HF config: FLA_CHUNK_SIZE. + server_args._mamba_cache_chunk_size = max(FLA_CHUNK_SIZE, cfg.page_size) + set_global_server_args_for_scheduler(server_args) device = get_device() mamba2_cache_params = None @@ -1336,6 +1340,8 @@ class UnifiedRadixCacheSuite: hicache_io_backend="direct", hicache_write_policy=write_policy, ) + # See build_fixture for why _mamba_cache_chunk_size is preset. + server_args._mamba_cache_chunk_size = max(FLA_CHUNK_SIZE, self.cfg.page_size) set_global_server_args_for_scheduler(server_args) tree.init_hicache(server_args, tree.cache_init_params) tree.write_through_threshold = 1 << 30