diff --git a/bindings/python/src/smg/serve.py b/bindings/python/src/smg/serve.py index 0603fd2e15..ee37bed8aa 100644 --- a/bindings/python/src/smg/serve.py +++ b/bindings/python/src/smg/serve.py @@ -363,13 +363,7 @@ def _add_trtllm_stub_args(parser: argparse.ArgumentParser) -> None: """ group = parser.add_argument_group("TensorRT-LLM Options") group.add_argument("--model", type=str, help="Model path (HuggingFace ID or local path)") - group.add_argument("--tp-size", type=int, help="Tensor parallel size (overrides config file)") - group.add_argument( - "--config", - type=str, - required=False, - help="Config file path (YAML, optional - must contain tensor_parallel_size if provided)", - ) + group.add_argument("--tp_size", type=int, help="Tensor parallel size (overrides config file)") BACKEND_ARG_ADDERS = { @@ -482,7 +476,10 @@ def parse_serve_args( _import_backend_args(backend, parser) RouterArgs.add_cli_args(parser, use_router_prefix=True, exclude_host_port=True) - args = parser.parse_args(argv) + if backend == "trtllm": + args, _ = parser.parse_known_args(argv) + else: + args = parser.parse_args(argv) return backend, args, backend_args diff --git a/bindings/python/tests/test_serve.py b/bindings/python/tests/test_serve.py index 40ced714b5..c53875c0e2 100644 --- a/bindings/python/tests/test_serve.py +++ b/bindings/python/tests/test_serve.py @@ -158,9 +158,12 @@ class TestImportBackendArgs: def test_trtllm_adds_model_arg(self): parser = argparse.ArgumentParser() _import_backend_args("trtllm", parser) - args = parser.parse_args(["--model", "/path/to/model", "--config", "/path/to/config.yml"]) + args, backend_args = parser.parse_known_args( + ["--model", "/path/to/model", "--config", "/path/to/config.yml"] + ) assert args.model == "/path/to/model" - assert args.config == "/path/to/config.yml" + assert "--config" in backend_args + assert "/path/to/config.yml" in backend_args def test_sglang_import_error(self): """sglang is not installed in test env, so parser.error should be called.""" @@ -304,7 +307,7 @@ def test_two_pass_extracts_backend_first(self): def test_unknown_arg_rejected_in_pass2(self): """Unknown args should be rejected by the full parser in pass 2.""" with pytest.raises(SystemExit): - parse_serve_args(["--backend", "trtllm", "--totally-unknown-flag"]) + parse_serve_args(["--backend", "sglang", "--totally-unknown-flag"]) # ---------------------------------------------------------------------------