diff --git a/tests/config/test_config_utils.py b/tests/config/test_config_utils.py index 24ef1b52a95d..678705f2a3db 100644 --- a/tests/config/test_config_utils.py +++ b/tests/config/test_config_utils.py @@ -221,6 +221,21 @@ def test_cache_config_hash_ignores_prefix_cache_retention_interval(): assert CacheConfig(prefix_cache_retention_interval=64).compute_hash() == base_hash +def test_swa_page_sizes_have_distinct_compile_cache_hashes(): + configs = ( + CacheConfig(swa_block_size=32), + CacheConfig(swa_block_size=64), + CacheConfig(swa_block_size=128), + ) + assert len({config.compute_hash() for config in configs}) == 3 + + +@pytest.mark.parametrize("block_size", [0, -32, 16, 96, 256]) +def test_swa_page_size_rejects_unsupported_kernel_geometry(block_size): + with pytest.raises(ValueError, match="swa_block_size"): + CacheConfig(swa_block_size=block_size) + + def test_envs_compile_factors_relocation_invariant(tmp_path): """Relocating HOME or the XDG roots must not change the compile-cache env hash. diff --git a/tests/engine/test_arg_utils.py b/tests/engine/test_arg_utils.py index 6ae65179788e..89e68b79172a 100644 --- a/tests/engine/test_arg_utils.py +++ b/tests/engine/test_arg_utils.py @@ -490,6 +490,22 @@ def test_attention_config(): engine_args.create_engine_config() +def test_swa_block_size_cli_choices(tmp_path): + parser = EngineArgs.add_cli_args(FlexibleArgumentParser()) + model_args = ["--model", str(tmp_path)] + assert ( + EngineArgs.from_cli_args(parser.parse_args(model_args)).swa_block_size is None + ) + args = parser.parse_args([*model_args, "--swa-block-size", "None"]) + assert EngineArgs.from_cli_args(args).swa_block_size is None + for size in (32, 64, 128): + args = parser.parse_args([*model_args, "--swa-block-size", str(size)]) + assert EngineArgs.from_cli_args(args).swa_block_size == size + parser.exit_on_error = False + with pytest.raises(ArgumentError): + parser.parse_args(["--swa-block-size", "96"]) + + def test_prefix_cache_default(): parser = EngineArgs.add_cli_args(FlexibleArgumentParser()) args = parser.parse_args([]) diff --git a/tests/models/test_deepseek_v4_1_ced_model.py b/tests/models/test_deepseek_v4_1_ced_model.py index f2f1f33cddd8..7477c284a2dc 100644 --- a/tests/models/test_deepseek_v4_1_ced_model.py +++ b/tests/models/test_deepseek_v4_1_ced_model.py @@ -181,6 +181,7 @@ def test_indexer_short_scan_reservation_preserves_long_scan_and_decode( capacity=4096, n_local_heads=16, swa_width=128, + swa_cache_layer=SimpleNamespace(block_size=64), is_ced_decoder=ced, layer_id=2, candidate_source_layer=20, @@ -306,7 +307,8 @@ def test_ced_global_context_and_full_row_output_abi(monkeypatch, compact): positions = torch.arange(100, 108) inputs = torch.arange(16, dtype=torch.float32).reshape(8, 2) / 8 indices = torch.tensor([2, 3, 6, 7, -1, -1]) if compact else None - shared, row_work = {}, [] + shared: dict[str, torch.Tensor | int] = {} + row_work: list[tuple[str, int, int]] = [] buffer = torch.empty(8, 8) target = SimpleNamespace( embed_input_ids=lambda _: inputs, @@ -345,6 +347,7 @@ def test_ced_global_context_and_full_row_output_abi(monkeypatch, compact): if index >= 2: expected_aux.append(expected.clone()) if compact: + assert indices is not None keep = torch.zeros(8, dtype=torch.bool) keep[indices[indices >= 0]] = True expected[~keep] = 0 @@ -360,7 +363,7 @@ def test_ced_global_context_and_full_row_output_abi(monkeypatch, compact): atol=1e-6, ) assert shared["global_rows"] == 8 - decoder_rows = len(indices) if compact else 8 + decoder_rows = len(indices) if indices is not None else 8 assert row_work == [ (kind, i, 8 if i < 2 else decoder_rows) for i in range(4) diff --git a/tests/test_config.py b/tests/test_config.py index df7e814d13b1..c9ef8bbe2406 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -17,6 +17,7 @@ import vllm.envs as envs from vllm.compilation.backends import VllmBackend from vllm.config import ( + CacheConfig, CompilationConfig, KernelConfig, ModelConfig, @@ -41,6 +42,39 @@ DEVICE_TYPE = current_platform.device_type +@pytest.mark.parametrize("swa_size,prefix_unit", [(None, 128), (32, 64), (128, 96)]) +def test_swa_page_size_rejects_incompatible_prefix_matching(swa_size, prefix_unit): + config = SimpleNamespace( + model_config=SimpleNamespace(architecture="DeepseekV41ForCausalLM"), + speculative_config=None, + cache_config=CacheConfig( + swa_block_size=swa_size, prefix_match_unit=prefix_unit + ), + ) + with pytest.raises(ValueError, match="must be divisible by --prefix-match-unit"): + VllmConfig.validate_swa_block_size(config) + config.cache_config.prefix_match_unit = 32 + VllmConfig.validate_swa_block_size(config) + + +def test_swa_page_size_is_scoped_to_v41_target_and_draft(): + target = SimpleNamespace(architecture="DeepseekV41ForCausalLM") + draft = SimpleNamespace(architecture="DSparkDeepseekV4ForCausalLM") + config = SimpleNamespace( + model_config=draft, + speculative_config=SimpleNamespace( + target_model_config=target, draft_model_config=draft + ), + cache_config=CacheConfig(swa_block_size=128, prefix_match_unit=32), + ) + VllmConfig.validate_swa_block_size(config) + target.architecture = "DeepseekV4ForCausalLM" + with pytest.raises(ValueError, match="only supported by native DeepSeek V4.1"): + VllmConfig.validate_swa_block_size(config) + config.cache_config.swa_block_size = None + VllmConfig.validate_swa_block_size(config) + + @pytest.mark.parametrize( "configured,expected", [(None, "auto"), ("b12x", "b12x"), ("FLASHINFER-CUTLASS", "flashinfer_cutlass")], diff --git a/tests/v1/attention/test_b12x_v41_workspace.py b/tests/v1/attention/test_b12x_v41_workspace.py index 9a7ad484ccff..2020c3081c43 100644 --- a/tests/v1/attention/test_b12x_v41_workspace.py +++ b/tests/v1/attention/test_b12x_v41_workspace.py @@ -76,7 +76,7 @@ def test_indexer_loads_complete_projections_on_tp_rank(native_workspace, monkeyp torch.testing.assert_close(projection.weight, source, rtol=0, atol=0) -def _layer(attention, layer_id=0): +def _layer(attention, layer_id=0, swa_page=32): # Avoid checkpoint/model construction: exercise the real planning and # attention methods with the same TP4 head geometry and serving capacity. layer = attention.DeepseekV4Attention.__new__(attention.DeepseekV4Attention) @@ -105,13 +105,63 @@ def _layer(attention, layer_id=0): layer.indexer = SimpleNamespace( heads=32, k_cache=SimpleNamespace(prefix=layer.prefix + ".indexer.k_cache") ) - layer.swa_cache_layer = SimpleNamespace(prefix=layer.prefix + ".swa_cache") + layer.swa_cache_layer = SimpleNamespace( + prefix=layer.prefix + ".swa_cache", block_size=swa_page + ) layer.compressor = None layer._context = {layer.prefix: layer} layer._ready = False return layer +@pytest.mark.parametrize("swa_page", [32, 64, 128]) +def test_context_cache_write_preserves_page_boundaries_and_padding( + native_workspace, swa_page +): + from b12x.attention._shared.mla.compressed_reference import ( + pack_deepseek_v41_cache_reference, + ) + + attention, _, _ = native_workspace + device = torch.device("cuda") + torch.manual_seed(145) + layer = _layer(attention, swa_page=swa_page) + rows = 2 * swa_page + 3 + layer.rotary_emb = SimpleNamespace( + cos_sin_cache=torch.cat( + ( + torch.ones((rows, 32), device=device), + torch.zeros((rows, 32), device=device), + ), + dim=-1, + ) + ) + page_bytes = attention.mla.page_nbytes( + swa_page, cache_kind="swa", cache_format="deepseek_v41" + ) + storage = torch.full((6, 113920), 0xA5, dtype=torch.uint8, device=device) + layer.swa_cache_layer.kv_cache = storage[:, :page_bytes] + expected = storage.clone() + logical = torch.arange(rows, device=device) + swa_page - 1 + pages = torch.tensor([3, 1, 4, 2], device=device) + slots = pages[logical // swa_page] * swa_page + logical % swa_page + slots[-2:] = -1 + positions = torch.arange(rows, device=device) + kv = torch.randn((rows, 512), device=device, dtype=torch.bfloat16) + records = pack_deepseek_v41_cache_reference( + kv, page_size=swa_page, cache_kind="swa" + ).view(-1, 528) + valid = slots >= 0 + columns = (slots[valid] % swa_page)[:, None] * 528 + torch.arange( + 528, device=device + ) + expected[(slots[valid] // swa_page)[:, None], columns] = records[:rows][valid] + + layer.insert_context_kv(kv, positions, slots) + + torch.testing.assert_close(storage, expected, rtol=0, atol=0) + + def test_prepare_memory_is_metadata_not_capacity_activations(native_workspace): attention, manager, _ = native_workspace device = torch.device("cuda", torch.accelerator.current_device_index()) @@ -287,6 +337,7 @@ def record(device, capacity, hidden): assert bool(torch.isfinite(y).all()) +@pytest.mark.parametrize("main_page,swa_page", [(64, 32), (128, 64), (256, 128)]) @pytest.mark.parametrize( "is_decode,rows,live_rows", [ @@ -299,14 +350,15 @@ def record(device, capacity, hidden): ], ) def test_attention_shared_scratch_graph_replay( - native_workspace, monkeypatch, is_decode, rows, live_rows + native_workspace, monkeypatch, is_decode, rows, live_rows, main_page, swa_page ): attention, manager, workspace = native_workspace # A TP4 rank must score all replicated heads without obtaining a TP group. monkeypatch.setattr(attention, "get_tensor_model_parallel_world_size", lambda: 4) device = torch.device("cuda", torch.accelerator.current_device_index()) torch.manual_seed(142) - layer = _layer(attention) + layer = _layer(attention, swa_page=swa_page) + layer.config.cache_config.block_size = main_page layer.max_model_len = 4096 if rows == 36: layer.config.speculative_config.parallel_drafting = True @@ -346,14 +398,14 @@ def metadata(page): context = SimpleNamespace( cudagraph_runtime_mode=CUDAGraphMode.NONE, attn_metadata={ - layer.swa_cache_layer.prefix: metadata(32), - layer.prefix: metadata(64), - layer.indexer.k_cache.prefix: metadata(64), + layer.swa_cache_layer.prefix: metadata(swa_page), + layer.prefix: metadata(main_page), + layer.indexer.k_cache.prefix: metadata(main_page), }, ) monkeypatch.setattr(attention, "get_forward_context", lambda: context) kv = torch.randn((length, 512), device=device, dtype=torch.bfloat16) - for kind, page in (("swa", 32), ("indexed", 64)): + for kind, page in (("swa", swa_page), ("indexed", main_page)): cache = torch.empty( ( length // page + 1, @@ -377,7 +429,10 @@ def metadata(page): else: layer.kv_cache = cache layer.indexer.k_cache.kv_cache = torch.empty( - (length // 64 + 1, attention.dsa_indexer.MXFP4_INDEX_PAGE_BYTES), + ( + length // main_page + 1, + attention.dsa_indexer.index_mxfp4_page_bytes(main_page), + ), dtype=torch.uint8, device=device, ) @@ -391,7 +446,8 @@ def metadata(page): attention.dsa_indexer.quantize_write_index_k_mxfp4( index_keys, index_k_cache=layer.indexer.k_cache.kv_cache, - slot_mapping=torch.arange(length, device=device) + 64, + slot_mapping=torch.arange(length, device=device) + main_page, + page_size=main_page, ) layer.attn_sink = torch.zeros(layer.n_local_heads, device=device) q = torch.randn( @@ -780,9 +836,10 @@ def test_ced_global_preparation_preserves_full_row_cache_bytes( ) +@pytest.mark.parametrize("main_page,swa_page", [(64, 32), (128, 64), (256, 128)]) @torch.inference_mode() def test_ced_compact_attention_bounded_oracle_and_frozen_replay( - native_workspace, monkeypatch + native_workspace, monkeypatch, main_page, swa_page ): from b12x import freeze_kernel_resolution, unfreeze_kernel_resolution from b12x.attention._shared.mla.compressed_reference import ( @@ -792,7 +849,8 @@ def test_ced_compact_attention_bounded_oracle_and_frozen_replay( attention, manager, workspace = native_workspace torch.manual_seed(713) device = torch.device("cuda") - layer = _layer(attention, 20) + layer = _layer(attention, 20, swa_page=swa_page) + layer.config.cache_config.block_size = main_page layer.is_ced_decoder = True layer._prepare(device) rows, length, boundary = 128, 256, 128 @@ -825,14 +883,14 @@ def metadata(page): context = SimpleNamespace( attn_metadata={ - layer.swa_cache_layer.prefix: metadata(32), - layer.prefix: metadata(64), - layer.indexer.k_cache.prefix: metadata(64), + layer.swa_cache_layer.prefix: metadata(swa_page), + layer.prefix: metadata(main_page), + layer.indexer.k_cache.prefix: metadata(main_page), } ) monkeypatch.setattr(attention, "get_forward_context", lambda: context) kv = torch.randn((length, 512), device=device, dtype=torch.bfloat16) - for kind, page in (("swa", 32), ("indexed", 64)): + for kind, page in (("swa", swa_page), ("indexed", main_page)): cache = torch.zeros( ( length // page + 1, @@ -865,7 +923,10 @@ def metadata(page): cache, page_size=page, cache_kind=kind )[page : page + length].clone() layer.indexer.k_cache.kv_cache = torch.zeros( - (length // 64 + 1, attention.dsa_indexer.index_mxfp4_page_bytes(64)), + ( + length // main_page + 1, + attention.dsa_indexer.index_mxfp4_page_bytes(main_page), + ), device=device, dtype=torch.uint8, ) @@ -875,8 +936,8 @@ def metadata(page): attention.dsa_indexer.quantize_write_index_k_mxfp4( keys, index_k_cache=layer.indexer.k_cache.kv_cache, - slot_mapping=torch.arange(length, device=device) + 64, - page_size=64, + slot_mapping=torch.arange(length, device=device) + main_page, + page_size=main_page, ) q = torch.randn((rows, 16, 512), device=device, dtype=torch.bfloat16) heads = layer.indexer.heads @@ -955,9 +1016,10 @@ def oracle(live): del graph, resources +@pytest.mark.parametrize("swa_page", [32, 64, 128]) @torch.inference_mode() def test_ced_replay_window_high_page_stride_and_invalid_rows( - native_workspace, monkeypatch + native_workspace, monkeypatch, swa_page ): from b12x.attention._shared.mla.compressed_reference import ( unpack_deepseek_v41_cache_reference, @@ -965,7 +1027,7 @@ def test_ced_replay_window_high_page_stride_and_invalid_rows( attention, _, _ = native_workspace device = torch.device("cuda") - layer = _layer(attention, 20) + layer = _layer(attention, 20, swa_page=swa_page) layer.is_ced_decoder = True layer.compress_ratio = 0 layer.indexer = None @@ -975,14 +1037,16 @@ def test_ced_replay_window_high_page_stride_and_invalid_rows( storage = torch.empty((high_pid + 1, stride), device=device, dtype=torch.uint8) cache = storage[ :, - : attention.mla.page_nbytes(32, cache_kind="swa", cache_format="deepseek_v41"), + : attention.mla.page_nbytes( + swa_page, cache_kind="swa", cache_format="deepseek_v41" + ), ] kv = torch.randn((1, 512), device=device, dtype=torch.bfloat16) attention.mla.write_cache( kv, cache, - torch.tensor([high_pid * 32], device=device), - page_size=32, + torch.tensor([high_pid * swa_page], device=device), + page_size=swa_page, cache_kind="swa", cache_format="deepseek_v41", ) @@ -992,7 +1056,7 @@ def test_ced_replay_window_high_page_stride_and_invalid_rows( positions=positions, req_id_per_token=torch.tensor([0, 0, -1], device=device, dtype=torch.int32), block_table=torch.tensor( - [[-1, -1, -1, -1, high_pid]], device=device, dtype=torch.int32 + [[-1] * (128 // swa_page) + [high_pid]], device=device, dtype=torch.int32 ), query_start_loc=torch.tensor([0, 1], device=device, dtype=torch.int32), request_positions=positions[:1], @@ -1012,7 +1076,9 @@ def test_ced_replay_window_high_page_stride_and_invalid_rows( layer.attn_sink = torch.zeros(16, device=device) layer.forward_mqa(q, None, positions, out) value = unpack_deepseek_v41_cache_reference( - cache[high_pid : high_pid + 1].contiguous(), page_size=32, cache_kind="swa" + cache[high_pid : high_pid + 1].contiguous(), + page_size=swa_page, + cache_kind="swa", )[0] score = torch.einsum("hd,d->h", q[0].float(), value) * 512**-0.5 expected = score.sigmoid()[:, None] * value diff --git a/tests/v1/core/test_contiguous_kv_packing.py b/tests/v1/core/test_contiguous_kv_packing.py index d7ef975a61ec..e4185692885a 100644 --- a/tests/v1/core/test_contiguous_kv_packing.py +++ b/tests/v1/core/test_contiguous_kv_packing.py @@ -192,7 +192,8 @@ def test_v41_mixed_cache_pages_preserve_request_partial_states(monkeypatch): assert (views["index"][blocks["index"]] == 23).all() -def test_v41_full_context_packs_shared_global_cache_without_page_inflation(): +@pytest.mark.parametrize("swa_size", [None, 32, 64, 128]) +def test_v41_full_context_packs_shared_global_cache_without_page_inflation(swa_size): from types import SimpleNamespace from vllm.models.deepseek_v4_1.attention import DeepseekV4Attention, _Cache @@ -201,6 +202,7 @@ def test_v41_full_context_packs_shared_global_cache_without_page_inflation(): config = _mock_vllm_config("BLHNC") config.cache_config.block_size = 256 + config.cache_config.swa_block_size = swa_size config.kv_transfer_config = None config.compilation_config.static_forward_context = {} config.speculative_config = None @@ -211,13 +213,14 @@ def test_v41_full_context_packs_shared_global_cache_without_page_inflation(): config.parallel_config.decode_context_parallel_size = 1 config.parallel_config.prefill_context_parallel_size = 1 specs = {} - global_names = [] + global_names: list[str] = [] for layer in range(43): prefix = f"model.layers.{layer}.self_attn" swa = _Cache( config, prefix + ".swa_cache", kind="swa", window=128, draft=layer >= 40 ) specs[swa.prefix] = swa.get_kv_cache_spec(config) + assert specs[swa.prefix].block_size == (64 if swa_size is None else swa_size) if layer not in (2, 8, 14, 20): continue ratio = 1 if layer == 20 else 2 diff --git a/vllm/config/cache.py b/vllm/config/cache.py index 4c7b4f2840c5..4eaca113f64e 100644 --- a/vllm/config/cache.py +++ b/vllm/config/cache.py @@ -83,6 +83,11 @@ class CacheConfig: block_size: int = Field(default=None, gt=0) # type: ignore[assignment] """Size of a contiguous cache block in number of tokens. Accepts None (meaning "use default"). After construction, always int.""" + swa_block_size: Literal[32, 64, 128] | None = None + """Tokens per sliding-window cache page for native DeepSeek V4.1 B12X. + None uses 64 tokens. Independent of the logical attention window and + the main/index cache's block_size. Requires a server restart; supported + values are 32, 64 and 128. Other models do not support this override.""" user_specified_block_size: bool = field(default=False, init=False) """Whether block_size was explicitly provided. Derived automatically.""" user_specified_mamba_block_size: bool = field(default=False, init=False) diff --git a/vllm/config/vllm.py b/vllm/config/vllm.py index f0b7c321225a..9ffc819950a1 100644 --- a/vllm/config/vllm.py +++ b/vllm/config/vllm.py @@ -2801,6 +2801,34 @@ def validate_nvfp4_kv_cache_with_mla(self) -> "VllmConfig": ) return self + @model_validator(mode="after") + def validate_swa_block_size(self) -> "VllmConfig": + model_config = self.model_config + if model_config is None: + return self + speculative = self.speculative_config + if speculative is not None and model_config is speculative.draft_model_config: + model_config = speculative.target_model_config + cache_config = self.cache_config + if model_config.architecture != "DeepseekV41ForCausalLM": + if cache_config.swa_block_size is not None: + raise ValueError( + "--swa-block-size is only supported by native DeepSeek V4.1 B12X" + ) + return self + swa_block_size = cache_config.swa_block_size or 64 + prefix_unit = cache_config.prefix_match_unit + if ( + cache_config.enable_prefix_caching + and prefix_unit is not None + and swa_block_size % prefix_unit != 0 + ): + raise ValueError( + f"SWA block size ({swa_block_size}) must be divisible by " + f"--prefix-match-unit ({prefix_unit})" + ) + return self + @model_validator(mode="after") def validate_mamba_block_size(self) -> "VllmConfig": if self.model_config is None: diff --git a/vllm/engine/arg_utils.py b/vllm/engine/arg_utils.py index 0b6229011165..c1bc0e0bf515 100644 --- a/vllm/engine/arg_utils.py +++ b/vllm/engine/arg_utils.py @@ -584,6 +584,7 @@ class EngineArgs: ParallelConfig.max_parallel_loading_workers ) block_size: int | None = None + swa_block_size: Literal[32, 64, 128] | None = CacheConfig.swa_block_size enable_prefix_caching: bool | None = None prefix_caching_hash_algo: PrefixCachingHashAlgo = ( CacheConfig.prefix_caching_hash_algo @@ -1306,6 +1307,12 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: description=CacheConfig.__doc__, ) cache_group.add_argument("--block-size", **cache_kwargs["block_size"]) + swa_kwargs = cache_kwargs["swa_block_size"].copy() + # argparse checks choices after optional_type converts "None" to None. + swa_kwargs["choices"] = [ + None if value == "None" else value for value in swa_kwargs["choices"] + ] + cache_group.add_argument("--swa-block-size", **swa_kwargs) cache_group.add_argument( "--gpu-memory-utilization", **cache_kwargs["gpu_memory_utilization"] ) @@ -2154,6 +2161,7 @@ def create_engine_config( cache_config = CacheConfig( block_size=self.block_size, # type: ignore[arg-type] + swa_block_size=self.swa_block_size, gpu_memory_utilization=self.gpu_memory_utilization, kv_cache_memory_bytes=self.kv_cache_memory_bytes, cache_dtype=resolved_cache_dtype, # type: ignore[arg-type] diff --git a/vllm/models/deepseek_v4_1/attention.py b/vllm/models/deepseek_v4_1/attention.py index 73778c45ea82..ae892b1366d3 100644 --- a/vllm/models/deepseek_v4_1/attention.py +++ b/vllm/models/deepseek_v4_1/attention.py @@ -97,7 +97,11 @@ def __init__(self, config, prefix, *, kind, ratio=1, window=0, draft=False): super().__init__() self.prefix, self.kind, self.ratio, self.window = prefix, kind, ratio, window self.draft = draft - self.block_size = 32 if kind == "swa" else config.cache_config.block_size + self.block_size = ( + config.cache_config.swa_block_size or 64 + if kind == "swa" + else config.cache_config.block_size + ) self.kv_cache = torch.tensor([]) context = config.compilation_config.static_forward_context if prefix in context: @@ -498,7 +502,7 @@ def alloc(shape, dtype=torch.bfloat16): max_width=self.swa_width + (512 if self.compress_ratio else 0), swa_width=self.swa_width, indexed_width=512 if self.compress_ratio else 0, - swa_page_size=32, + swa_page_size=self.swa_cache_layer.block_size, indexed_page_size=self._main_page, max_page_table_width=self._main_width, mode=mode, @@ -596,7 +600,7 @@ def insert_context_kv(self, kv, positions, slot_mapping): rotated, self.swa_cache_layer.kv_cache, slot_mapping, - page_size=32, + page_size=self.swa_cache_layer.block_size, cache_kind="swa", cache_format="deepseek_v41", ) @@ -831,7 +835,7 @@ def forward_mqa(self, q, kv, positions, output, *, index_query=None): swa.request_positions, 0, swa.block_table.stride(0), - 32, + self.swa_cache_layer.block_size, self.window_size, self.swa_width, self.is_draft, @@ -871,7 +875,7 @@ def forward_mqa(self, q, kv, positions, output, *, index_query=None): binding=binding, swa_k_cache=self.swa_cache_layer.kv_cache, indexed_k_cache=self._owner().kv_cache if main is not None else None, - swa_page_size=32, + swa_page_size=self.swa_cache_layer.block_size, indexed_page_size=self._main_page, sm_scale=512**-0.5, attn_sink=self.attn_sink,