diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 27d0adeeb..db1a3a433 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -450,6 +450,8 @@ class ServerArgs: # Model override args in JSON json_model_override_args: str = "{}" + nsa_index_topk_freq: Optional[int] = None + nsa_index_skip_topk_offset: Optional[int] = None preferred_sampling_params: Optional[str] = None # LoRA @@ -755,6 +757,10 @@ class ServerArgs: # Validate SSL arguments early (before dummy-model short-circuit). self._handle_ssl_validation() + # Apply convenience model-config override flags before any code reads + # hf_config through get_model_config(). + self._handle_nsa_index_model_override_args() + # Validate CP shared KV constraints early (before dummy-model short-circuit). self._handle_cp_shared_kv_validation() @@ -915,6 +921,43 @@ class ServerArgs: "to be specified." ) + def _handle_nsa_index_model_override_args(self): + """Merge first-class NSA index-cache flags into model config overrides.""" + + overrides: Dict[str, int] = {} + if self.nsa_index_topk_freq is not None: + if self.nsa_index_topk_freq < 1: + raise ValueError("--nsa-index-topk-freq must be a positive integer.") + overrides["index_topk_freq"] = self.nsa_index_topk_freq + + if self.nsa_index_skip_topk_offset is not None: + if self.nsa_index_skip_topk_offset < 0: + raise ValueError( + "--nsa-index-skip-topk-offset must be non-negative; use 0 " + "or omit the flag to keep the model-config default." + ) + if self.nsa_index_skip_topk_offset > 0: + overrides["index_skip_topk_offset"] = self.nsa_index_skip_topk_offset + + if not overrides: + return + + try: + model_override_args = json.loads(self.json_model_override_args or "{}") + except json.JSONDecodeError as exc: + raise ValueError( + "--json-model-override-args must be valid JSON when combined " + "with --nsa-index-topk-freq or --nsa-index-skip-topk-offset." + ) from exc + + if not isinstance(model_override_args, dict): + raise ValueError("--json-model-override-args must decode to a JSON object.") + + model_override_args.update(overrides) + self.json_model_override_args = json.dumps( + model_override_args, separators=(",", ":") + ) + def _handle_cp_shared_kv_validation(self): assert not ( self.enable_nsa_prefill_cp_shared_kv @@ -4648,6 +4691,29 @@ class ServerArgs: help="A dictionary in JSON string format used to override default model configurations.", default=ServerArgs.json_model_override_args, ) + parser.add_argument( + "--nsa-index-topk-freq", + type=int, + default=ServerArgs.nsa_index_topk_freq, + help=( + "Shortcut for setting model config index_topk_freq without " + "editing config.json. Values greater than 1 enable NSA/DSA " + "indexer top-k sharing across layers. Equivalent to adding " + "'index_topk_freq' to --json-model-override-args; this flag " + "takes precedence if both are specified." + ), + ) + parser.add_argument( + "--nsa-index-skip-topk-offset", + type=int, + default=ServerArgs.nsa_index_skip_topk_offset, + help=( + "Shortcut for setting model config index_skip_topk_offset " + "without editing config.json. Equivalent to adding " + "'index_skip_topk_offset' to --json-model-override-args; this " + "flag takes precedence if both are specified." + ), + ) parser.add_argument( "--preferred-sampling-params", type=json.loads, diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 55cb63d09..1c9f98c4d 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -92,6 +92,53 @@ def test_hicache_mem_layout_parser_accepts_layer_page_first(): assert args.hicache_mem_layout == "layer_page_first" +def test_nsa_index_topk_cli_overrides_model_config_json(): + import argparse + + parser = argparse.ArgumentParser() + ServerArgs.add_cli_args(parser) + raw_args = parser.parse_args( + [ + "--model-path", + "dummy", + "--json-model-override-args", + '{"rope_scaling": {"factor": 2.0}, "index_topk_freq": 2}', + "--nsa-index-topk-freq", + "4", + "--nsa-index-skip-topk-offset", + "1", + ] + ) + args = ServerArgs.from_cli_args(raw_args) + override_args = json.loads(args.json_model_override_args) + assert args.nsa_index_topk_freq == 4 + assert args.nsa_index_skip_topk_offset == 1 + assert override_args["rope_scaling"] == {"factor": 2.0} + assert override_args["index_topk_freq"] == 4 + assert override_args["index_skip_topk_offset"] == 1 + + +def test_nsa_index_skip_topk_offset_zero_means_unset(): + import argparse + + parser = argparse.ArgumentParser() + ServerArgs.add_cli_args(parser) + raw_args = parser.parse_args( + [ + "--model-path", + "dummy", + "--nsa-index-topk-freq", + "4", + "--nsa-index-skip-topk-offset", + "0", + ] + ) + args = ServerArgs.from_cli_args(raw_args) + override_args = json.loads(args.json_model_override_args) + assert args.nsa_index_skip_topk_offset == 0 + assert override_args == {"index_topk_freq": 4} + + class TestLoadBalanceMethod(unittest.TestCase): def test_non_pd_defaults_to_round_robin(self): server_args = ServerArgs(model_path="dummy", disaggregation_mode="null")