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:
laoyao0822
2026-06-10 03:24:33 +08:00
parent 24da983ff5
commit a32050d1d8
2 changed files with 113 additions and 0 deletions
@@ -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")