53 lines
1.9 KiB
Python
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
|