Improve: Add test code for ONNX inference with DmlExecutionProvider
This commit is contained in:
@@ -91,12 +91,14 @@ python = ["3.9", "3.10", "3.11"]
|
|||||||
|
|
||||||
# for ONNX inference (without PyTorch dependency)
|
# for ONNX inference (without PyTorch dependency)
|
||||||
[tool.hatch.envs.test-onnx]
|
[tool.hatch.envs.test-onnx]
|
||||||
dependencies = ["coverage[toml]>=6.5", "pytest", "scipy", "onnxruntime-gpu"]
|
dependencies = ["coverage[toml]>=6.5", "pytest", "scipy", "onnxruntime-gpu", "onnxruntime-directml; sys_platform == 'win32'"]
|
||||||
[tool.hatch.envs.test-onnx.scripts]
|
[tool.hatch.envs.test-onnx.scripts]
|
||||||
# Usage: `hatch run test-onnx:test`
|
# Usage: `hatch run test-onnx:test`
|
||||||
test = "pytest tests/test_main.py::test_synthesize_onnx_cpu"
|
test = "pytest tests/test_main.py::test_synthesize_onnx_cpu"
|
||||||
# Usage: `hatch run test-onnx:test-cuda`
|
# Usage: `hatch run test-onnx:test-cuda`
|
||||||
test-cuda = "pytest tests/test_main.py::test_synthesize_onnx_cuda"
|
test-cuda = "pytest tests/test_main.py::test_synthesize_onnx_cuda"
|
||||||
|
# Usage: `hatch run test-onnx:test-directml`
|
||||||
|
test-directml = "pytest tests/test_main.py::test_synthesize_onnx_directml"
|
||||||
# Usage: `hatch run test-onnx:coverage`
|
# Usage: `hatch run test-onnx:coverage`
|
||||||
test-cov = "coverage run -m pytest tests/test_main.py::test_synthesize_onnx_cpu"
|
test-cov = "coverage run -m pytest tests/test_main.py::test_synthesize_onnx_cpu"
|
||||||
# Usage: `hatch run test-onnx:cov-report`
|
# Usage: `hatch run test-onnx:cov-report`
|
||||||
|
|||||||
@@ -37,7 +37,9 @@ def synthesize(
|
|||||||
if f.endswith(".onnx") and not f.startswith(".")
|
if f.endswith(".onnx") and not f.startswith(".")
|
||||||
]
|
]
|
||||||
if len(model_files) == 0:
|
if len(model_files) == 0:
|
||||||
pytest.skip(f"音声合成モデル \"{model_info.name}\" のモデルファイルが見つかりませんでした。")
|
pytest.skip(
|
||||||
|
f'音声合成モデル "{model_info.name}" のモデルファイルが見つかりませんでした。'
|
||||||
|
)
|
||||||
|
|
||||||
# モデルをロード
|
# モデルをロード
|
||||||
model = model_holder.get_model(model_info.name, model_files[0])
|
model = model_holder.get_model(model_info.name, model_files[0])
|
||||||
@@ -100,3 +102,8 @@ def test_synthesize_onnx_cuda():
|
|||||||
("CUDAExecutionProvider", {"cudnn_conv_algo_search": "DEFAULT"})
|
("CUDAExecutionProvider", {"cudnn_conv_algo_search": "DEFAULT"})
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_synthesize_onnx_directml():
|
||||||
|
pytest.importorskip("onnxruntime_directml")
|
||||||
|
synthesize(inference_type="onnx", onnx_providers=["DmlExecutionProvider"])
|
||||||
|
|||||||
Reference in New Issue
Block a user