diff --git a/tests/python/contrib/test_msc/test_manager.py b/tests/python/contrib/test_msc/test_manager.py index 393c8decbc79..c3e1583ef291 100644 --- a/tests/python/contrib/test_msc/test_manager.py +++ b/tests/python/contrib/test_msc/test_manager.py @@ -34,6 +34,7 @@ def _get_config(model_type, compile_type, inputs, outputs, atol=1e-2, rtol=1e-2): """Get msc config""" return { + "workspace": msc_utils.msc_dir(), "model_type": model_type, "inputs": inputs, "outputs": outputs, diff --git a/tests/python/contrib/test_msc/test_runner.py b/tests/python/contrib/test_msc/test_runner.py index a6005f5d41c2..cbc33d452846 100644 --- a/tests/python/contrib/test_msc/test_runner.py +++ b/tests/python/contrib/test_msc/test_runner.py @@ -82,7 +82,7 @@ def _test_from_torch(runner_cls, device, is_training=False, atol=1e-3, rtol=1e-3 torch_model = _get_torch_model("resnet50", is_training) if torch_model: - workspace = msc_utils.set_workspace() + workspace = msc_utils.set_workspace(msc_utils.msc_dir()) log_path = workspace.relpath("MSC_LOG", keep_history=False) msc_utils.set_global_logger("info", log_path) input_info = [([1, 3, 224, 224], "float32")] @@ -139,7 +139,7 @@ def test_tensorflow_runner(): tf_graph, graph_def = _get_tf_graph() if tf_graph and graph_def: - workspace = msc_utils.set_workspace() + workspace = msc_utils.set_workspace(msc_utils.msc_dir()) log_path = workspace.relpath("MSC_LOG", keep_history=False) msc_utils.set_global_logger("info", log_path) data = np.random.uniform(size=(1, 224, 224, 3)).astype("float32") diff --git a/tests/python/contrib/test_msc/test_tools.py b/tests/python/contrib/test_msc/test_tools.py index 4578e662acb8..6adf70605bfd 100644 --- a/tests/python/contrib/test_msc/test_tools.py +++ b/tests/python/contrib/test_msc/test_tools.py @@ -46,10 +46,10 @@ def _get_config( ): """Get msc config""" return { + "workspace": msc_utils.msc_dir(), "model_type": model_type, "inputs": inputs, "outputs": outputs, - "debug_level": 0, "dataset": {"loader": "from_random", "max_iter": 5}, "prepare": {"profile": {"benchmark": {"repeat": 10}}}, "baseline": {