Skip to content
Merged
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
22 changes: 22 additions & 0 deletions src/mobius/integrations/onnx_genai/auto_export.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,12 +16,15 @@
import os
from typing import Any

import yaml

from mobius.integrations.onnx_genai.decoder_metadata import (
decoder_metadata_from_config,
write_decoder_metadata,
)
from mobius.integrations.onnx_genai.inference_metadata import (
SchedulerConfig,
add_explicit_package_io,
load_diffusers_scheduler_config,
write_audio_codec_pipeline_metadata,
write_diffusion_pipeline_metadata,
Expand All @@ -35,6 +38,21 @@
_DENOISER_KEYS = ("denoiser", "transformer", "unet")


def _add_explicit_io_to_file(path: str, pkg: Any, config: Any) -> None:
"""Augment an emitted sidecar with roles derived from the actual ONNX ports."""
try:
models = list(pkg.values())
except AttributeError:
return
if not models or any(not hasattr(model, "graph") for model in models):
return
with open(path, encoding="utf-8") as handle:
metadata = yaml.safe_load(handle)
add_explicit_package_io(metadata, pkg, config)
with open(path, "w", encoding="utf-8") as handle:
yaml.safe_dump(metadata, handle, sort_keys=False)


def _write_clip_tokenizer(output_dir: str, source: str | None) -> str | None:
"""Emit ``tokenizer.json`` for a text-conditioned diffusion package.

Expand Down Expand Up @@ -451,6 +469,7 @@ def write_onnx_genai_config(
activation_dtype=_activation_dtype_tag(resolved_config),
**kwargs,
)
_add_explicit_io_to_file(path, pkg, resolved_config)
artifacts = {"inference_metadata": path}
tokenizer_path = _write_hf_tokenizer(output_dir, source)
if tokenizer_path is not None:
Expand All @@ -467,6 +486,7 @@ def write_onnx_genai_config(
activation_dtype=_activation_dtype_tag(resolved_config),
**kwargs,
)
_add_explicit_io_to_file(path, pkg, resolved_config)
artifacts = {"inference_metadata": path}
tokenizer_path = _write_hf_tokenizer(output_dir, source)
if tokenizer_path is not None:
Expand Down Expand Up @@ -495,6 +515,7 @@ def write_onnx_genai_config(
decoder_metadata=decoder_metadata,
**_tts_component_kwargs(pkg, resolved_config),
)
_add_explicit_io_to_file(path, pkg, resolved_config)
artifacts = {"inference_metadata": path}
tokenizer_path = _write_hf_tokenizer(output_dir, source)
if tokenizer_path is not None:
Expand All @@ -520,6 +541,7 @@ def write_onnx_genai_config(
path = write_decoder_metadata(
output_dir, config=resolved_config, kv_native_dtype=kv_native_dtype
)
_add_explicit_io_to_file(path, pkg, resolved_config)
artifacts = {"inference_metadata": path}
tokenizer_path = _write_hf_tokenizer(output_dir, source)
if tokenizer_path is not None:
Expand Down
21 changes: 11 additions & 10 deletions src/mobius/integrations/onnx_genai/auto_export_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -297,7 +297,7 @@ def test_dispatch_speech_to_text_pipeline(tmp_path):
pkg = _EncoderDecoderPkg(
{
"encoder": _FakeModel(["input_features"], ["encoder_hidden_states"]),
"decoder": _FakeModel(["decoder_input_ids", "encoder_hidden_states"]),
"decoder": _FakeModel(["decoder_input_ids", "encoder_hidden_states"], ["logits"]),
}
)
artifacts = write_onnx_genai_config(pkg, str(tmp_path), kv_native_dtype="bf16")
Expand All @@ -307,14 +307,14 @@ def test_dispatch_speech_to_text_pipeline(tmp_path):

assert metadata["kv_cache"] == {"native_dtype": "bfloat16"}
pipeline = metadata["pipeline"]
assert pipeline["models"] == {
"encoder": {"filename": "encoder/model.onnx", "type": "encoder"},
"decoder": {
"filename": "decoder/model.onnx",
"type": "decoder",
"tokenizer": "tokenizer.json",
},
}
assert pipeline["models"]["encoder"]["filename"] == "encoder/model.onnx"
assert pipeline["models"]["encoder"]["type"] == "encoder"
decoder_model = pipeline["models"]["decoder"]
assert decoder_model["filename"] == "decoder/model.onnx"
assert decoder_model["type"] == "decoder"
assert decoder_model["tokenizer"] == "tokenizer.json"
assert decoder_model["io"]["logits_output"] == "logits"
assert decoder_model["io"]["kv_ownership"] == "owned"
assert pipeline["dataflow"] == [
{
"from": "encoder.encoder_hidden_states",
Expand Down Expand Up @@ -403,7 +403,7 @@ def test_dispatch_multi_decoder_tts_with_pre_embedder(tmp_path):
pkg = _TTSPkg(
{
"talker": _FakeModel(["inputs_embeds"], ["logits", "last_hidden_state"]),
"code_predictor": _FakeModel(["inputs_embeds"], ["logits"]),
"code_predictor": _FakeModel(["inputs_embeds"], ["logits", "codec_embeddings"]),
"talker_step_embedder": _FakeModel(["frame_codes"], ["inputs_embeds"]),
"talker_prefill_embedder": _FakeModel(
["text_ids"], ["prefill_embeds", "trailing_text_embeds"]
Expand All @@ -425,6 +425,7 @@ def test_dispatch_multi_decoder_tts_with_pre_embedder(tmp_path):
}
stage = pipeline["strategy"]["stages"][0]["strategy"]
assert stage["kind"] == "nested_autoregressive"
assert stage["inner_embedding_output"] == "codec_embeddings"
assert stage["pre_embedder"]["component"] == "talker_step_embedder"
assert stage["prefill_embedder"]["component"] == "talker_prefill_embedder"
assert stage["num_code_groups"] == 16
Expand Down
Loading
Loading