From 24a3fc1a9ce22a25956268fc8a44615f8f6afabe Mon Sep 17 00:00:00 2001 From: zhangxingzhi Date: Wed, 16 Nov 2022 14:39:00 +0800 Subject: [PATCH] ci: enable test_dump_distribution --- .../unittests/test_pipeline_job_schema.py | 11 ++++++----- 1 file changed, 6 insertions(+), 5 deletions(-) diff --git a/sdk/ml/azure-ai-ml/tests/pipeline_job/unittests/test_pipeline_job_schema.py b/sdk/ml/azure-ai-ml/tests/pipeline_job/unittests/test_pipeline_job_schema.py index 4450dd1072f1..7fc6b2a458d1 100644 --- a/sdk/ml/azure-ai-ml/tests/pipeline_job/unittests/test_pipeline_job_schema.py +++ b/sdk/ml/azure-ai-ml/tests/pipeline_job/unittests/test_pipeline_job_schema.py @@ -1121,9 +1121,9 @@ def test_pipeline_job_with_environment_variables(self) -> None: "abc": "def", } - @pytest.mark.skip("Pipeline: discuss how to refactor _to_dict in PipelineJob & CommandComponent later.") def test_dump_distribution(self): - from azure.ai.ml._restclient.v2021_10_01.models import TensorFlow + # pipeline level test is in test_pipeline_job_create_with_distribution_component + from azure.ai.ml import TensorFlowDistribution from azure.ai.ml._schema.job.distribution import PyTorchDistributionSchema, TensorFlowDistributionSchema distribution_dict = { @@ -1132,11 +1132,12 @@ def test_dump_distribution(self): "parameter_server_count": 0, "worker_count": 5, } - distribution_obj = TensorFlow.from_dict(distribution_dict) + # msrest has been removed from public interface + distribution_obj = TensorFlowDistribution(**distribution_dict) - with pytest.raises(ValidationError, match=r"Value passed is not in set \['pytorch']"): + with pytest.raises(ValidationError, match=r"Cannot dump non-PyTorchDistribution object into PyTorchDistributionSchema"): _ = PyTorchDistributionSchema(context={"base_path": "./"}).dump(distribution_dict) - with pytest.raises(ValidationError, match=r"Value passed is not in set \['pytorch']"): + with pytest.raises(ValidationError, match=r"Cannot dump non-PyTorchDistribution object into PyTorchDistributionSchema"): _ = PyTorchDistributionSchema(context={"base_path": "./"}).dump(distribution_obj) after_dump_correct = TensorFlowDistributionSchema(context={"base_path": "./"}).dump(distribution_obj)