diff --git a/src/yolo_webui/export_runner.py b/src/yolo_webui/export_runner.py index 16d6e0f..0eada5d 100755 --- a/src/yolo_webui/export_runner.py +++ b/src/yolo_webui/export_runner.py @@ -247,16 +247,19 @@ def main(argv: Sequence[str] | None = None) -> int: flush=True, ) - exported_path = model.export( - format=config.export_format, - imgsz=config.imgsz, - half=config.half, - int8=config.int8, - dynamic=config.dynamic, - simplify=config.simplify, - batch=config.batch, - workspace=config.workspace, - ) + export_kwargs: dict[str, Any] = { + "format": config.export_format, + "imgsz": config.imgsz, + "half": config.half, + "int8": config.int8, + "dynamic": config.dynamic, + "simplify": config.simplify, + "batch": config.batch, + } + if config.export_format in ("engine", "tensorrt", "trt"): + export_kwargs["workspace"] = config.workspace + + exported_path = model.export(**export_kwargs) result_path = _result_path(exported_path) print("Экспорт завершен успешно.", flush=True) diff --git a/tests/test_export_runner.py b/tests/test_export_runner.py index 9715869..3c6de01 100644 --- a/tests/test_export_runner.py +++ b/tests/test_export_runner.py @@ -88,12 +88,39 @@ def test_main_validates_config_and_reports_existing_export( "dynamic": True, "simplify": True, "batch": 2, - "workspace": 2.5, } ] assert f"__YOLO_WEBUI_RESULT__:{result}" in output.out +def test_main_passes_workspace_only_for_engine_format( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + model = _write_model(tmp_path, monkeypatch) + result = tmp_path / "models" / "model.engine" + result.write_bytes(b"engine") + _, export_calls = _install_fake_ultralytics(monkeypatch, result) + config = _write_config( + tmp_path, + { + "model": model.name, + "format": "engine", + "imgsz": 640, + "half": False, + "int8": False, + "dynamic": False, + "simplify": True, + "batch": 1, + "workspace": 4.0, + }, + ) + + return_code = export_runner.main([str(config)]) + assert return_code == 0 + assert export_calls[0]["workspace"] == 4.0 + + @pytest.mark.parametrize( "override", [