Skip to content

Commit 81089da

Browse files
committed
Versions
Signed-off-by: Peter Jung <peter@jung.ninja>
1 parent bf9795b commit 81089da

6 files changed

Lines changed: 97 additions & 43 deletions

File tree

.gitignore

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,3 @@
1+
# Test files
2+
info.txt
3+
run_tests.log

mlflow_export_import/click_doc.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,8 @@
99

1010
model_stages = "Stages to export (comma seperated). Default is all stages. Values are Production, Staging, Archived and None."
1111

12+
model_versions = "Versions to export (comma seperated). Default is all versions. Values are valid integer numbers."
13+
1214
delete_model = "If the model exists, first delete the model and all its versions."
1315

1416
use_threads = "Process export/import in parallel using threads."

mlflow_export_import/model/export_model.py

Lines changed: 33 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -11,18 +11,20 @@
1111
from mlflow_export_import import utils, click_doc
1212

1313
class ModelExporter():
14-
def __init__(self, mlflow_client, export_source_tags=False, notebook_formats=None, stages=None, export_run=True):
14+
def __init__(self, mlflow_client, export_source_tags=False, notebook_formats=None, stages=None, versions=None, export_run=True):
1515
"""
1616
:param mlflow_client: MLflow client or if None create default client.
1717
:param export_source_tags: Export source run metadata tags.
1818
:param notebook_formats: List of notebook formats to export. Values are SOURCE, HTML, JUPYTER or DBC.
1919
:param stages: Stages to export. Default is all stages. Values are Production, Staging, Archived and None.
20+
:param versions: Versions to export. Default is all versions. Values are valid integer numbers.
2021
:param export_run: Export the run that generated a registered model's version.
2122
"""
2223
self.mlflow_client = mlflow_client
2324
self.http_client = MlflowHttpClient()
2425
self.run_exporter = RunExporter(self.mlflow_client, export_source_tags=export_source_tags, notebook_formats=notebook_formats)
2526
self.stages = self._normalize_stages(stages)
27+
self.versions = self._normalize_versions(versions)
2628
self.export_run = export_run
2729

2830
def export_model(self, model_name, output_dir):
@@ -50,6 +52,8 @@ def _export_model(self, model_name, output_dir):
5052
for vr in versions:
5153
if len(self.stages) > 0 and not vr.current_stage.lower() in self.stages:
5254
continue
55+
if len(self.versions) > 0 and vr.version not in self.versions:
56+
continue
5357
run_id = vr.run_id
5458
opath = os.path.join(output_dir,run_id)
5559
opath = opath.replace("dbfs:", "/dbfs")
@@ -89,14 +93,26 @@ def _normalize_stages(self, stages):
8993
print(f"WARNING: stage '{stage}' must be one of: {model_version_stages.ALL_STAGES}")
9094
return stages
9195

96+
def _normalize_versions(self, versions):
97+
if versions is None:
98+
return []
99+
if isinstance(versions, str):
100+
versions = versions.split(",")
101+
for version in versions:
102+
try:
103+
int(version)
104+
except ValueError:
105+
print(f"WARNING: version '{version}' must be a valid number")
106+
return versions
107+
92108
@click.command()
93-
@click.option("--model",
94-
help="Registered model name.",
109+
@click.option("--model",
110+
help="Registered model name.",
95111
type=str,
96112
required=True
97113
)
98-
@click.option("--output-dir",
99-
help="Output directory.",
114+
@click.option("--output-dir",
115+
help="Output directory.",
100116
type=str,
101117
required=True
102118
)
@@ -106,24 +122,29 @@ def _normalize_stages(self, stages):
106122
default=False,
107123
show_default=True
108124
)
109-
@click.option("--notebook-formats",
110-
help=click_doc.notebook_formats,
125+
@click.option("--notebook-formats",
126+
help=click_doc.notebook_formats,
111127
type=str,
112-
default="",
128+
default="",
113129
show_default=True
114130
)
115-
@click.option("--stages",
116-
help=click_doc.model_stages,
131+
@click.option("--stages",
132+
help=click_doc.model_stages,
133+
type=str,
134+
required=False
135+
)
136+
@click.option("--versions",
137+
help=click_doc.model_versions,
117138
type=str,
118139
required=False
119140
)
120141

121-
def main(model, output_dir, export_source_tags, notebook_formats, stages):
142+
def main(model, output_dir, export_source_tags, notebook_formats, stages, versions):
122143
print("Options:")
123144
for k,v in locals().items():
124145
print(f" {k}: {v}")
125146
client = mlflow.tracking.MlflowClient()
126-
exporter = ModelExporter(client, export_source_tags=export_source_tags, notebook_formats=utils.string_to_list(notebook_formats), stages=stages)
147+
exporter = ModelExporter(client, export_source_tags=export_source_tags, notebook_formats=utils.string_to_list(notebook_formats), stages=stages, versions=versions)
127148
exporter.export_model(model, output_dir)
128149

129150
if __name__ == "__main__":

mlflow_export_import/model/import_model.py

Lines changed: 25 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -19,9 +19,9 @@ def __init__(self, mlflow_client, run_importer=None, await_creation_for=None):
1919
:param run_importer: RunImporter instance.
2020
:param await_creation_for: Seconds to wait for model version crreation.
2121
"""
22-
self.mlflow_client = mlflow_client
22+
self.mlflow_client = mlflow_client
2323
self.run_importer = run_importer if run_importer else RunImporter(self.mlflow_client, mlmodel_fix=True)
24-
self.await_creation_for = await_creation_for
24+
self.await_creation_for = await_creation_for
2525

2626
def _import_version(self, model_name, src_vr, dst_run_id, dst_source, sleep_time):
2727
"""
@@ -33,10 +33,10 @@ def _import_version(self, model_name, src_vr, dst_run_id, dst_source, sleep_time
3333
"""
3434
src_current_stage = src_vr["current_stage"]
3535
dst_source = dst_source.replace("file://","") # OSS MLflow
36-
if not dst_source.startswith("dbfs:") and not os.path.exists(dst_source):
36+
if not dst_source.startswith("dbfs:") and not dst_source.startswith("s3:") and not os.path.exists(dst_source):
3737
raise MlflowExportImportException(f"'source' argument for MLflowClient.create_model_version does not exist: {dst_source}")
3838
kwargs = {"await_creation_for": self.await_creation_for } if self.await_creation_for else {}
39-
version = self.mlflow_client.create_model_version(model_name, dst_source, dst_run_id, **kwargs)
39+
version = self.mlflow_client.create_model_version(model_name, dst_source, dst_run_id, src_vr["tags"], **kwargs)
4040
model_utils.wait_until_version_is_ready(self.mlflow_client, model_name, version, sleep_time=sleep_time)
4141
if src_current_stage != "None":
4242
self.mlflow_client.transition_model_version_stage(model_name, version.version, src_current_stage)
@@ -173,46 +173,46 @@ def _path_join(x,y):
173173
""" Account for DOS backslash """
174174
path = os.path.join(x,y)
175175
if path.startswith("dbfs:"):
176-
path = path.replace("\\","/")
176+
path = path.replace("\\","/")
177177
return path
178178

179179
@click.command()
180-
@click.option("--input-dir",
181-
help="Input directory produced by export_model.py.",
180+
@click.option("--input-dir",
181+
help="Input directory produced by export_model.py.",
182182
type=str,
183183
required=True
184184
)
185-
@click.option("--model",
186-
help="New registered model name.",
185+
@click.option("--model",
186+
help="New registered model name.",
187187
type=str,
188-
required=True,
188+
required=True,
189189
)
190-
@click.option("--experiment-name",
191-
help="Destination experiment name - will be created if it does not exist.",
190+
@click.option("--experiment-name",
191+
help="Destination experiment name - will be created if it does not exist.",
192192
type=str,
193193
required=True
194194
)
195-
@click.option("--delete-model",
196-
help=click_doc.delete_model,
195+
@click.option("--delete-model",
196+
help=click_doc.delete_model,
197197
type=bool,
198-
default=False,
198+
default=False,
199199
show_default=True
200200
)
201-
@click.option("--await-creation-for",
202-
help="Await creation for specified seconds.",
203-
type=int,
204-
default=None,
201+
@click.option("--await-creation-for",
202+
help="Await creation for specified seconds.",
203+
type=int,
204+
default=None,
205205
show_default=True
206206
)
207-
@click.option("--sleep-time",
208-
help="Sleep time for polling until version.status==READY.",
207+
@click.option("--sleep-time",
208+
help="Sleep time for polling until version.status==READY.",
209209
type=int,
210210
default=5,
211211
)
212-
@click.option("--verbose",
213-
help="Verbose.",
214-
type=bool,
215-
default=False,
212+
@click.option("--verbose",
213+
help="Verbose.",
214+
type=bool,
215+
default=False,
216216
show_default=True
217217
)
218218

tests/run_tests.sh

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -27,13 +27,13 @@ export MLFLOW_TRACKING_URI_SRC=http://localhost:${PORT_SRC}
2727
export MLFLOW_TRACKING_URI_DST=http://localhost:${PORT_DST}
2828

2929
message() {
30-
echo
30+
echo
3131
echo "******************************************************"
3232
echo "*"
3333
echo "* $*"
3434
echo "*"
3535
echo "******************************************************"
36-
echo
36+
echo
3737
}
3838

3939
run_tests() {
@@ -50,7 +50,7 @@ launch_server() {
5050
mlflow server \
5151
--host localhost --port ${port} \
5252
--backend-store-uri sqlite:///mlflow_${port}.db \
53-
--default-artifact-root $PWD/mlruns_${port}
53+
--default-artifact-root "$(PWD)/mlruns_${port}"
5454
}
5555

5656
kill_server() {

tests/test_models.py

Lines changed: 31 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,7 @@
11
from mlflow_export_import.model.export_model import ModelExporter
22
from mlflow_export_import.model.import_model import ModelImporter
33
from mlflow_export_import import utils
4-
import utils_test
4+
import utils_test
55
import compare_utils
66
from init_tests import mlflow_context
77

@@ -57,6 +57,34 @@ def test_export_import_model_stages(mlflow_context):
5757
[vr_prod_src, vr_staging_src],
5858
[vr_prod_dst, vr_staging_dst])
5959

60+
def test_export_import_model_versions(mlflow_context):
61+
model_name_src = utils_test.mk_test_object_name_default()
62+
model_src = mlflow_context.client_src.create_registered_model(model_name_src)
63+
64+
vr_staging_src = _create_version(mlflow_context.client_src, model_name_src, "Staging")
65+
vr_prod_src = _create_version(mlflow_context.client_src, model_name_src, "Production")
66+
67+
exporter = ModelExporter(mlflow_context.client_src, versions=[vr_staging_src.version, vr_prod_src.version])
68+
exporter.export_model(model_name_src, mlflow_context.output_dir)
69+
70+
model_name_dst = utils_test.create_dst_model_name(model_name_src)
71+
experiment_name = model_name_dst
72+
importer = ModelImporter(mlflow_context.client_dst)
73+
importer.import_model(model_name_dst, mlflow_context.output_dir, experiment_name, delete_model=True, verbose=False, sleep_time=10)
74+
75+
model_dst = mlflow_context.client_dst.get_registered_model(model_name_dst)
76+
assert len(model_dst.latest_versions) == 2
77+
78+
versions = mlflow_context.client_dst.get_latest_versions(model_name_dst)
79+
vr_prod_dst = [vr for vr in versions if vr.version == vr_prod_src.version][0]
80+
vr_staging_dst = [vr for vr in versions if vr.version == vr_staging_src.version][0]
81+
82+
_compare_models(model_src, model_dst)
83+
_compare_version_lists(
84+
mlflow_context, mlflow_context.output_dir,
85+
[vr_prod_src, vr_staging_src],
86+
[vr_prod_dst, vr_staging_dst])
87+
6088

6189
def _create_version(client, model_name, stage=None):
6290
run = _create_run(client)
@@ -69,7 +97,7 @@ def _create_version(client, model_name, stage=None):
6997
def _create_run(client):
7098
_, run = utils_test.create_simple_run(client)
7199
return client.get_run(run.info.run_id)
72-
100+
73101
def _compare_models(model_src, model_dst):
74102
assert model_src.description == model_dst.description
75103
assert model_src.tags == model_dst.tags
@@ -85,7 +113,7 @@ def _compare_versions(mlflow_context, output_dir, vr_src, vr_dst):
85113
assert vr_src.status == vr_dst.status
86114
assert vr_src.status_message == vr_dst.status_message
87115
if not utils.importing_into_databricks():
88-
assert vr_src.user_id == vr_dst.user_id
116+
assert vr_src.user_id == vr_dst.user_id
89117
assert vr_src.run_id != vr_dst.run_id
90118
run_src = mlflow_context.client_src.get_run(vr_src.run_id)
91119
run_dst = mlflow_context.client_dst.get_run(vr_dst.run_id)

0 commit comments

Comments
 (0)