Skip to content

Commit 3861631

Browse files
justinchubyCopilot
andauthored
WP1: native VLM package emission (typed metadata) (#420)
Implements native VLM package emission aligned with the typed onnx-genai metadata schema. Independent review (Resch): - Model-name gate: no VLM production logic branches on model identity; the sole grep hit is an unrelated pre-existing Qwen3-TTS docstring. - `python3 -m pytest src/mobius/integrations/onnx_genai/inference_metadata_test.py -q`: 40 passed. (`python` is unavailable on this host, so the equivalent `python3` interpreter was used.) - `ruff check src/mobius/integrations/onnx_genai/inference_metadata.py`: passed. - Spot-checked schema fields/enums/capabilities against onnx-genai Rust and JSON schemas. Authored by Dave; independently reviewed by Resch. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
1 parent 54be48a commit 3861631

2 files changed

Lines changed: 35 additions & 14 deletions

File tree

src/mobius/integrations/onnx_genai/inference_metadata.py

Lines changed: 20 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -686,6 +686,7 @@ def _state_and_kv_pairs(
686686
class _PositionProgram:
687687
rank: int
688688
axes: tuple[str, ...]
689+
generation: str
689690
continuation: str
690691
matches: Callable[[Any], bool]
691692
sections_attribute: str | None = None
@@ -695,13 +696,15 @@ class _PositionProgram:
695696
_PositionProgram(
696697
rank=1,
697698
axes=("sequence",),
699+
generation="linear",
698700
continuation="linear_increment",
699701
matches=lambda config: True,
700702
),
701703
_PositionProgram(
702704
rank=3,
703705
axes=("temporal", "height", "width"),
704-
continuation="from_grid",
706+
generation="processor_coordinates",
707+
continuation="carry_max",
705708
matches=lambda config: (
706709
bool(getattr(config, "mrope_interleaved", False))
707710
and bool(getattr(config, "mrope_section", None))
@@ -729,7 +732,9 @@ def _positions_from_registry(position: _Port, config: Any) -> dict[str, Any]:
729732
positions: dict[str, Any] = {
730733
"input": position.name,
731734
"rank": program.rank,
735+
"tensor_rank": position.rank,
732736
"dtype": position.dtype,
737+
"generation": program.generation,
733738
"continuation": program.continuation,
734739
"axes": list(program.axes),
735740
}
@@ -975,12 +980,12 @@ def validate_executable_closure(pkg: Any, metadata: dict[str, Any]) -> None:
975980
f"but the consumer requires {target.dtype}/{target.rank}.",
976981
"insert an explicit typed transform or remove the incompatible edge.",
977982
)
978-
if edge.get("dtype") != source.dtype or edge.get("rank") != source.rank:
983+
if edge.get("dtype") != source.dtype:
979984
raise _closure_error(
980985
target_endpoint,
981-
f"the edge declares dtype/rank {edge.get('dtype')}/{edge.get('rank')}, "
982-
f"but the ONNX ports are {source.dtype}/{source.rank}.",
983-
"derive edge dtype and rank directly from the matched graph ports.",
986+
f"the edge declares dtype {edge.get('dtype')}, "
987+
f"but the ONNX ports use {source.dtype}.",
988+
"derive the edge dtype directly from the matched graph ports.",
984989
)
985990
incoming.setdefault(target_endpoint, []).append(edge)
986991

@@ -1182,7 +1187,6 @@ def build_native_vlm_package_metadata(
11821187
"from": f"{source_name}.{source_port.name}",
11831188
"to": f"{target_name}.{target_port.name}",
11841189
"dtype": source_port.dtype,
1185-
"rank": source_port.rank,
11861190
"device_transfer": False,
11871191
}
11881192
)
@@ -1314,17 +1318,24 @@ def build_native_vlm_package_metadata(
13141318
vision_config = {key: value for key, value in vision_config.items() if value is not None}
13151319

13161320
metadata = dict(decoder_metadata or {})
1321+
metadata.setdefault("schema_version", "v1")
13171322
capabilities = list(metadata.get("required_capabilities", []))
13181323
for capability in (
1319-
"multimodal_image_preprocessing",
1324+
"image_preprocessing_program",
13201325
"autoregressive_every_step_components",
13211326
):
13221327
if capability not in capabilities:
13231328
capabilities.append(capability)
1329+
if len(preprocessing_outputs) > 1 and "packed_image_outputs" not in capabilities:
1330+
capabilities.append("packed_image_outputs")
1331+
if positions is not None and "position_program" not in capabilities:
1332+
capabilities.append("position_program")
13241333
if positions is not None and positions["rank"] > 1:
1325-
capabilities.append("multiaxis_positions")
1334+
capabilities.append("multi_axis_positions")
13261335
if decoder_io.get("state_pairs"):
1327-
capabilities.append("loop_state")
1336+
capabilities.append("loop_carried_state")
1337+
if decoder_io.get("token_input") and decoder_io.get("inputs_embeds_input"):
1338+
capabilities.append("dual_sequence_inputs")
13281339
metadata["required_capabilities"] = capabilities
13291340
metadata.setdefault("model", {})["io"] = decoder_io
13301341
metadata["preprocessing"] = {

src/mobius/integrations/onnx_genai/inference_metadata_test.py

Lines changed: 15 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -420,26 +420,30 @@ def test_gemma4_routes_all_embedding_outputs(self, tmp_path):
420420
validate_executable_closure(package, metadata)
421421
self._validate(metadata)
422422
_assert_all_graph_ports_declared(package, metadata)
423+
assert metadata["schema_version"] == "v1"
424+
assert {
425+
"image_preprocessing_program",
426+
"packed_image_outputs",
427+
"position_program",
428+
"dual_sequence_inputs",
429+
} <= set(metadata["required_capabilities"])
423430
flow = metadata["pipeline"]["dataflow"]
424431
assert {
425432
"from": "embedding.inputs_embeds",
426433
"to": "decoder.inputs_embeds",
427434
"dtype": "fp32",
428-
"rank": 3,
429435
"device_transfer": False,
430436
} in flow
431437
assert {
432438
"from": "embedding.per_layer_inputs",
433439
"to": "decoder.per_layer_inputs",
434440
"dtype": "fp32",
435-
"rank": 3,
436441
"device_transfer": False,
437442
} in flow
438443
assert {
439444
"from": "audio_encoder.audio_features",
440445
"to": "embedding.audio_features",
441446
"dtype": "fp32",
442-
"rank": 2,
443447
"device_transfer": False,
444448
} in flow
445449
embedding_audio = next(
@@ -621,12 +625,19 @@ def test_qwen_packed_grid_rank3_positions_sparse_and_fixed_state(self, tmp_path)
621625
)
622626
self._validate(metadata)
623627
_assert_all_graph_ports_declared(package, metadata)
628+
assert {
629+
"position_program",
630+
"multi_axis_positions",
631+
"loop_carried_state",
632+
} <= set(metadata["required_capabilities"])
624633
positions = metadata["pipeline"]["positions"]
625634
assert positions == {
626635
"input": "position_ids",
627636
"rank": 3,
637+
"tensor_rank": 3,
628638
"dtype": "int64",
629-
"continuation": "from_grid",
639+
"generation": "processor_coordinates",
640+
"continuation": "carry_max",
630641
"axes": ["temporal", "height", "width"],
631642
"sections": [16, 24, 24],
632643
"processor_summaries": ["vision_encoder.image_grid_thw"],
@@ -764,7 +775,6 @@ def test_phi_routes_both_modality_gates_and_mask_processor(self, tmp_path):
764775
"from": f"embedding.{gate}",
765776
"to": f"decoder.{gate}",
766777
"dtype": "fp32",
767-
"rank": 0,
768778
"device_transfer": False,
769779
} in flow
770780
outputs = metadata["preprocessing"]["image"]["outputs"]

0 commit comments

Comments
 (0)