mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
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:
parent
89723c8612
commit
32fabb5555
2 changed files with 27 additions and 1 deletions
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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['']}.")
|
||||
|
|
|
|||
Loading…
Reference in a new issue