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
28 changes: 13 additions & 15 deletions src/mobius/tasks/_adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,33 +23,31 @@ def build(
module,
config,
) -> ModelPackage:
graph, builder = _make_graph()
op = builder.op

# Determine input shape based on adapter type
if hasattr(config, "in_channels"):
# T2I-Adapter: conditioning image input
condition = ir.Value(
name="condition",
type=ir.TensorType(ir.DataType.FLOAT),
shape=ir.Shape(("batch", config.in_channels, "height", "width")),
condition = builder.input(
"condition",
dtype=ir.DataType.FLOAT,
shape=["batch", config.in_channels, "height", "width"],
)
else:
# IP-Adapter: image embedding input
condition = ir.Value(
name="image_embeds",
type=ir.TensorType(ir.DataType.FLOAT),
shape=ir.Shape(("batch", config.image_embed_dim)),
condition = builder.input(
"image_embeds",
dtype=ir.DataType.FLOAT,
shape=["batch", config.image_embed_dim],
)

graph, builder = _make_graph([condition])
op = builder.op

outputs = module(op, condition)

if isinstance(outputs, list):
for i, out in enumerate(outputs):
out.name = f"feature_{i}"
graph.outputs.append(out)
builder.add_output(out, f"feature_{i}")
else:
outputs.name = "adapter_output"
graph.outputs.append(outputs)
builder.add_output(outputs, "adapter_output")

return ModelPackage({"model": _make_model(graph)}, config=config)
15 changes: 6 additions & 9 deletions src/mobius/tasks/_audio_feature_extraction.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,18 +28,15 @@ def build(
module,
config: ArchitectureConfig,
) -> ModelPackage:
input_values = ir.Value(
name="input_values",
type=ir.TensorType(ir.DataType.FLOAT),
shape=ir.Shape(("batch", "time")),
)

graph, builder = _make_graph([input_values])
graph, builder = _make_graph()
op = builder.op

input_values = builder.input(
"input_values", dtype=ir.DataType.FLOAT, shape=["batch", "time"]
)

last_hidden_state = module(op, input_values=input_values)

last_hidden_state.name = "last_hidden_state"
graph.outputs.append(last_hidden_state)
builder.add_output(last_hidden_state, "last_hidden_state")

return ModelPackage({"model": _make_model(graph)}, config=config)
72 changes: 35 additions & 37 deletions src/mobius/tasks/_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,8 +9,7 @@
from typing import ClassVar

import onnx_ir as ir
from onnxscript import nn
from onnxscript._internal.builder import GraphBuilder
from onnxscript import GraphBuilder, nn

import mobius
from mobius._configs import BaseModelConfig
Expand Down Expand Up @@ -103,16 +102,18 @@ def __repr__(self) -> str:


def _make_graph(
inputs: list[ir.Value],
name: str = "main_graph",
) -> tuple[ir.Graph, GraphBuilder]:
"""Create an empty graph and its builder.

Inputs should be added after creation via ``builder.input()``.
Outputs should be registered via ``builder.add_output()``.

Returns:
``(graph, builder)`` — call ``builder.op`` to get the op handle.
"""
graph = ir.Graph(
inputs,
[],
[],
nodes=[],
name=name,
Expand Down Expand Up @@ -237,45 +238,44 @@ def build_decoder_from_embeds(
seq_len = ir.SymbolicDim("sequence_len")
past_seq_len = ir.SymbolicDim("past_sequence_len")

inputs_embeds = ir.Value(
name="inputs_embeds",
shape=ir.Shape([batch, seq_len, config.hidden_size]),
type=ir.TensorType(config.dtype),
graph, builder = _make_graph()
inputs_embeds = builder.input(
"inputs_embeds",
dtype=config.dtype,
shape=[batch, seq_len, config.hidden_size],
)
attention_mask = ir.Value(
name="attention_mask",
shape=ir.Shape([batch, "past_seq_len + seq_len"]),
type=ir.TensorType(ir.DataType.INT64),
attention_mask = builder.input(
"attention_mask",
dtype=ir.DataType.INT64,
shape=[batch, "past_seq_len + seq_len"],
)
# MRoPE: 3D position IDs (temporal, height, width) — shape [3, batch, seq_len]
# Standard: shape [batch, seq_len]
position_ids = ir.Value(
name="position_ids",
shape=ir.Shape([3, batch, seq_len] if mrope else [batch, seq_len]),
type=ir.TensorType(ir.DataType.INT64),
position_ids = builder.input(
"position_ids",
dtype=ir.DataType.INT64,
shape=[3, batch, seq_len] if mrope else [batch, seq_len],
)

graph_inputs = [inputs_embeds, attention_mask, position_ids]

if hybrid:
cache_inputs, past_key_values = _make_hybrid_cache_inputs(
past_key_values = _make_hybrid_cache_inputs(
builder,
config,
config.dtype,
batch,
past_seq_len,
)
else:
cache_inputs, past_key_values = _make_kv_cache_inputs(
past_key_values = _make_kv_cache_inputs(
builder,
config.num_hidden_layers,
config.num_key_value_heads,
config.head_dim,
config.dtype,
batch,
past_seq_len,
)
graph_inputs.extend(cache_inputs)

graph, builder = _make_graph(graph_inputs)
logits, present_key_values = decoder(
builder.op,
inputs_embeds=inputs_embeds,
Expand All @@ -284,20 +284,19 @@ def build_decoder_from_embeds(
past_key_values=past_key_values,
)

logits.name = "logits"
graph.outputs.append(logits)
builder.add_output(logits, "logits")

if hybrid:
_register_hybrid_cache_outputs(
graph,
builder,
present_key_values,
config.layer_types or [],
)
model = _make_model(graph)
_register_linear_attention_functions(model, config)
return model
else:
_register_kv_cache_outputs(graph, present_key_values)
_register_kv_cache_outputs(builder, present_key_values)
return _make_model(graph)


Expand Down Expand Up @@ -328,24 +327,23 @@ def build_embedding_from_features(
seq_len = ir.SymbolicDim("sequence_len")
num_feature_tokens = ir.SymbolicDim("num_feature_tokens")

input_ids = ir.Value(
name="input_ids",
shape=ir.Shape([batch, seq_len]),
type=ir.TensorType(ir.DataType.INT64),
graph, builder = _make_graph(name="embedding")
input_ids = builder.input(
"input_ids",
dtype=ir.DataType.INT64,
shape=[batch, seq_len],
)
features = ir.Value(
name=feature_name,
shape=ir.Shape([num_feature_tokens, feature_dim]),
type=ir.TensorType(config.dtype),
features = builder.input(
feature_name,
dtype=config.dtype,
shape=[num_feature_tokens, feature_dim],
)

graph, builder = _make_graph([input_ids, features], name="embedding")
inputs_embeds = embedding(
builder.op,
input_ids=input_ids,
**{feature_name: features},
)

inputs_embeds.name = "inputs_embeds"
graph.outputs.append(inputs_embeds)
builder.add_output(inputs_embeds, "inputs_embeds")
return _make_model(graph)
Loading
Loading