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.
This commit is contained in:
Xavier Dupré 2023-11-22 18:15:11 +01:00 committed by GitHub
parent 89723c8612
commit 32fabb5555
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
2 changed files with 27 additions and 1 deletions

View file

@ -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()

View file

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