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:
@@ -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