Expose NSA index sharing as launch-time overrides
Index-topk sharing previously required editing model config or passing raw JSON overrides. Add first-class server args that merge into json_model_override_args before ModelConfig is read, preserving existing JSON overrides while letting launch scripts toggle index_topk_freq directly. Treat offset 0 as unset so wrappers can use 0 for default behavior without injecting a model-config override. Constraint: Prefill and decode launch commands need the same effective model config without mutating /ssd model files. Rejected: Require editing config.json | operationally fragile across g0034/g0035/g0036 model copies. Rejected: Only use --json-model-override-args | works but is too error-prone for frequent launch-command tuning. Confidence: high Scope-risk: narrow Directive: Keep these flags as model-config override shortcuts; apply them before any get_model_config() call. Tested: Remote pytest in g0034 container for NSA index override parser tests: 2 passed. Tested: py_compile for server_args.py and test_server_args.py. Not-tested: Full prefill/decode ETE launch with --nsa-index-topk-freq enabled.
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user