2026-01-28 06:16:04 +00:00

53 lines
1.9 KiB
Python

import onnx.version_converter
import logging
from .helper import meta
from .helper.update_opset_8 import convert_opset_8_to_9
from .helper.update_opset_9 import convert_opset_9_to_10
from .helper.update_opset_10 import convert_opset_10_to_11
logger = logging.getLogger("kneronnxopt.version_updater")
custom_opset_converter = {
# 8: convert_opset_8_to_9,
# 9: convert_opset_9_to_10,
# 10: convert_opset_10_to_11,
}
def update_version(model, target_opset=meta.latest_opset):
"""Update opset version of the model
:model: onnx model
:target_opset: opset version
:returns: updated model
"""
# Check if the current version is available for updating.
current_opset = meta.get_opset(model)
if current_opset == target_opset:
return model
if current_opset > target_opset:
logger.error(
f"Current opset version ({current_opset}) is newer than the target opset version ({target_opset})."
)
raise ValueError(
f"Current opset version ({current_opset}) is newer than the target opset version ({target_opset})."
)
if current_opset not in meta.supported_opset:
logger.error(f"Current opset version ({current_opset}) is not supported.")
raise ValueError(f"Current opset version ({current_opset}) is not supported.")
# Update the opset version to the latest version.
logger.info(f"Updating opset version from {current_opset} to {target_opset}...")
if target_opset != meta.latest_opset:
logger.warning(
f"Target opset version ({target_opset}) is not the latest opset version ({meta.latest_opset})."
)
while current_opset != target_opset:
model = onnx.version_converter.convert_version(model, current_opset + 1)
if current_opset in custom_opset_converter:
model = custom_opset_converter[current_opset](model)
current_opset += 1
return model