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['']}.")