From 32fabb555501a020751b6123de94c7fc14086f2b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Xavier=20Dupr=C3=A9?= Date: Wed, 22 Nov 2023 18:15:11 +0100 Subject: [PATCH] Fix opset version of the optimizer in function generate_artifacts (#18300) ### Description `generate_artifacts` generates 4 graphs for training. All graphs should share the same opset version, the one coming from the model to train, but the optimizer is left undefined. onnxruntime is using the latest version defined by onnx but onnxruntime does not necessarily support it. ### Motivation and Context The code does not let the user change it. --- .../orttraining/python/training/artifacts.py | 10 +++++++++- .../orttraining_test_ort_apis_onnxblock.py | 18 ++++++++++++++++++ 2 files changed, 27 insertions(+), 1 deletion(-) diff --git a/orttraining/orttraining/python/training/artifacts.py b/orttraining/orttraining/python/training/artifacts.py index 549614de49..a57105545e 100644 --- a/orttraining/orttraining/python/training/artifacts.py +++ b/orttraining/orttraining/python/training/artifacts.py @@ -53,6 +53,8 @@ def generate_artifacts( 3. Checkpoint (directory): Contains the model parameters. 4. Optimizer model (onnx.ModelProto): Model containing the optimizer graph. + All generated ModelProtos will use the same opsets defined by *model*. + Args: model: The base model to be used for gradient graph generation. requires_grad: List of names of model parameters that require gradient computation @@ -207,11 +209,17 @@ def generate_artifacts( logging.info("Optimizer enum provided: %s", optimizer.name) + opset_version = None + for domain in model.opset_import: + if domain.domain == "" or domain.domain == "ai.onnx": + opset_version = domain.version + break + optim_model = None optim_blocks = {OptimType.AdamW: onnxblock.optim.AdamW, OptimType.SGD: onnxblock.optim.SGD} optim_block = optim_blocks[optimizer]() - with onnxblock.empty_base(): + with onnxblock.empty_base(opset_version=opset_version): _ = optim_block(model_params) optim_model = optim_block.to_model_proto() diff --git a/orttraining/orttraining/test/python/orttraining_test_ort_apis_onnxblock.py b/orttraining/orttraining/test/python/orttraining_test_ort_apis_onnxblock.py index f7a7220dd6..6e5d54cbb9 100644 --- a/orttraining/orttraining/test/python/orttraining_test_ort_apis_onnxblock.py +++ b/orttraining/orttraining/test/python/orttraining_test_ort_apis_onnxblock.py @@ -17,6 +17,14 @@ from onnxruntime.training import artifacts # PyTorch Module definitions +def get_opsets_model(filename): + if isinstance(filename, onnx.ModelProto): + onx = filename + else: + onx = onnx.load(filename) + return {d.domain: d.version for d in onx.opset_import} + + class SimpleNet(torch.nn.Module): def __init__(self, input_size, hidden_size, num_classes): super().__init__() @@ -999,3 +1007,13 @@ def test_save_ort_format(): assert os.path.exists(os.path.join(temp_dir, "eval_model.ort")) assert os.path.exists(os.path.join(temp_dir, "optimizer_model.onnx")) assert os.path.exists(os.path.join(temp_dir, "optimizer_model.ort")) + base_opsets = get_opsets_model(base_model) + training_opsets = get_opsets_model(os.path.join(temp_dir, "training_model.onnx")) + eval_opsets = get_opsets_model(os.path.join(temp_dir, "eval_model.onnx")) + optimizer_opsets = get_opsets_model(os.path.join(temp_dir, "optimizer_model.onnx")) + if base_opsets[""] != training_opsets[""]: + raise AssertionError(f"Opsets mismatch {base_opsets['']} != {training_opsets['']}.") + if base_opsets[""] != eval_opsets[""]: + raise AssertionError(f"Opsets mismatch {base_opsets['']} != {eval_opsets['']}.") + if base_opsets[""] != optimizer_opsets[""]: + raise AssertionError(f"Opsets mismatch {base_opsets['']} != {optimizer_opsets['']}.")