Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
37 changes: 26 additions & 11 deletions src/mobius/_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -201,14 +201,17 @@ def build_from_module(

def _maybe_apply_opset_lowering(package: ModelPackage, execution_provider: str) -> None:
"""Lower default-domain opset 24 to 23 for sub-models where it is safe."""
if not flags.ort_lower_opset_for_ep:
if not flags.ort_lower_opset_for_ep and execution_provider != "openvino":
return
if execution_provider in ("default", "cpu"):
return
for name, model in package.items():
if "" not in model.graph.opset_imports:
continue
if _graph_requires_opset24(model.graph):
functions = list(model.functions.values())
if _graph_requires_opset24(model.graph) or any(
_graph_requires_opset24(function) for function in functions
):
logger.info(
"Skipped opset→23 lowering for '%s' (EP=%s): graph uses "
"opset-24-only ops (TensorScatter / Attention nonpad_kv_seqlen). "
Expand All @@ -219,15 +222,27 @@ def _maybe_apply_opset_lowering(package: ModelPackage, execution_provider: str)
continue
original = model.graph.opset_imports[""]
model.graph.opset_imports[""] = 23
logger.warning(
"Lowered opset %d→23 for '%s' (EP=%s). "
"ORT does not yet register opset %d kernels for this EP. "
"Track https://github.com/microsoft/onnxruntime/issues/27729",
original,
name,
execution_provider,
original,
)
for function in functions:
if "" in function.opset_imports:
function.opset_imports[""] = 23
if execution_provider == "openvino":
logger.info(
"Lowered opset %d→23 for '%s' (EP=%s) to avoid the OpenVINO "
"opset-24 Attention mask Pad.",
original,
name,
execution_provider,
)
else:
logger.warning(
"Lowered opset %d→23 for '%s' (EP=%s). "
"ORT does not yet register opset %d kernels for this EP. "
"Track https://github.com/microsoft/onnxruntime/issues/27729",
original,
name,
execution_provider,
original,
)


def _graph_requires_opset24(graph: ir.Graph) -> bool:
Expand Down
11 changes: 11 additions & 0 deletions src/mobius/_builder_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -183,3 +183,14 @@ def test_maybe_apply_opset_lowering_skipped_when_flag_disabled(
_maybe_apply_opset_lowering(pkg, execution_provider="cuda")

assert pkg["embedding"].graph.opset_imports[""] == 24


def test_maybe_apply_opset_lowering_required_by_openvino(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(flags, "ort_lower_opset_for_ep", False)
pkg = ModelPackage({"decoder": _model_with(_standard_nodes())})

_maybe_apply_opset_lowering(pkg, execution_provider="openvino")

assert pkg["decoder"].graph.opset_imports[""] == 23
Loading
Loading