diff --git a/src/mobius/tasks/_adapter.py b/src/mobius/tasks/_adapter.py index ccf5e5dea..f586beb55 100644 --- a/src/mobius/tasks/_adapter.py +++ b/src/mobius/tasks/_adapter.py @@ -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) diff --git a/src/mobius/tasks/_audio_feature_extraction.py b/src/mobius/tasks/_audio_feature_extraction.py index 647a7d8af..51a62b51f 100644 --- a/src/mobius/tasks/_audio_feature_extraction.py +++ b/src/mobius/tasks/_audio_feature_extraction.py @@ -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) diff --git a/src/mobius/tasks/_base.py b/src/mobius/tasks/_base.py index ff29723c6..b09276fcb 100644 --- a/src/mobius/tasks/_base.py +++ b/src/mobius/tasks/_base.py @@ -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 @@ -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, @@ -237,35 +238,36 @@ 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, @@ -273,9 +275,7 @@ def build_decoder_from_embeds( 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, @@ -284,12 +284,11 @@ 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 [], ) @@ -297,7 +296,7 @@ def build_decoder_from_embeds( _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) @@ -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) diff --git a/src/mobius/tasks/_cache_utils.py b/src/mobius/tasks/_cache_utils.py index 6a2d6875d..41cf5071c 100644 --- a/src/mobius/tasks/_cache_utils.py +++ b/src/mobius/tasks/_cache_utils.py @@ -14,6 +14,7 @@ from typing import NamedTuple import onnx_ir as ir +from onnxscript import GraphBuilder from mobius._configs import BaseModelConfig @@ -64,6 +65,7 @@ def linear_attention_dims(config: BaseModelConfig) -> LinearAttentionDims: def _make_kv_cache_inputs( + builder: GraphBuilder, num_layers: int, num_kv_heads: int, head_dim: int, @@ -74,42 +76,41 @@ def _make_kv_cache_inputs( prefix: str = "past_key_values", key_head_dim: int | None = None, value_head_dim: int | None = None, -) -> tuple[list[ir.Value], list[tuple[ir.Value, ir.Value]]]: +) -> list[tuple[ir.Value, ir.Value]]: """Create KV cache input values for ``num_layers`` layers. + Uses ``builder.input()`` to create and register graph inputs directly. + Args: + builder: The graph builder to register inputs on. key_head_dim: Head dim for keys. Defaults to ``head_dim``. For MLA attention, this is ``qk_nope_head_dim + qk_rope_head_dim``. value_head_dim: Head dim for values. Defaults to ``head_dim``. For MLA attention, this is ``v_head_dim``. Returns: - ``(flat_inputs, kv_pairs)`` where *flat_inputs* is a flat list - suitable for extending ``graph_inputs`` and *kv_pairs* is a list - of ``(key, value)`` tuples for passing to the module. + A list of ``(key, value)`` tuples for passing to the module. """ k_dim = key_head_dim if key_head_dim is not None else head_dim v_dim = value_head_dim if value_head_dim is not None else head_dim - flat: list[ir.Value] = [] pairs: list[tuple[ir.Value, ir.Value]] = [] for i in range(num_layers): - past_key = ir.Value( - name=f"{prefix}.{i}.key", - shape=ir.Shape([batch, num_kv_heads, past_seq_len, k_dim]), - type=ir.TensorType(dtype), + past_key = builder.input( + f"{prefix}.{i}.key", + dtype=dtype, + shape=[batch, num_kv_heads, past_seq_len, k_dim], ) - past_value = ir.Value( - name=f"{prefix}.{i}.value", - shape=ir.Shape([batch, num_kv_heads, past_seq_len, v_dim]), - type=ir.TensorType(dtype), + past_value = builder.input( + f"{prefix}.{i}.value", + dtype=dtype, + shape=[batch, num_kv_heads, past_seq_len, v_dim], ) - flat.extend([past_key, past_value]) pairs.append((past_key, past_value)) - return flat, pairs + return pairs def _register_kv_cache_outputs( - graph: ir.Graph, + builder: GraphBuilder, present_key_values: list[tuple[ir.Value, ir.Value]], *, prefix: str = "present", @@ -120,23 +121,23 @@ def _register_kv_cache_outputs( that runs during model optimization. """ for i, (present_key, present_value) in enumerate(present_key_values): - present_key.name = f"{prefix}.{i}.key" - present_value.name = f"{prefix}.{i}.value" - - graph.outputs.append(present_key) - graph.outputs.append(present_value) + builder.add_output(present_key, f"{prefix}.{i}.key") + builder.add_output(present_value, f"{prefix}.{i}.value") def _make_hybrid_cache_inputs( + builder: GraphBuilder, config: BaseModelConfig, dtype: ir.DataType, batch: ir.SymbolicDim, past_seq_len: ir.SymbolicDim, *, prefix: str = "past_key_values", -) -> tuple[list[ir.Value], list[StatePair]]: +) -> list[StatePair]: """Create cache inputs for hybrid models with mixed layer types. + Uses ``builder.input()`` to create and register graph inputs directly. + Supported layer types: ``"full_attention"`` — standard KV cache (key + value). ``"lightning_attention"`` — single recurrent state only; no conv_state. @@ -146,13 +147,10 @@ def _make_hybrid_cache_inputs( ``"mlp"`` — stateless, produces ``(None, None)`` pair. Returns: - ``(flat_inputs, state_pairs)`` — *flat_inputs* contains only - the ``ir.Value`` entries (no graph inputs for MLP layers); - *state_pairs* has one entry per layer, with ``(None, None)`` + A list of state pairs, one per layer, with ``(None, None)`` for stateless MLP layers. """ layer_types = config.layer_types or [] - flat: list[ir.Value] = [] pairs: list[StatePair] = [] # DeltaNet dimensions from config (computed once via shared helper) @@ -187,91 +185,79 @@ def _make_hybrid_cache_inputs( if ltype == "lightning_attention": # Lightning Attention: single recurrent state only (no conv_state) # State: (B, num_heads, head_dim, head_dim) — square matrix accumulator - rec_state = ir.Value( - name=f"{prefix}.{i}.recurrent_state", - shape=ir.Shape( - [batch, config.num_attention_heads, config.head_dim, config.head_dim] - ), - type=ir.TensorType(dtype), + rec_state = builder.input( + f"{prefix}.{i}.recurrent_state", + dtype=dtype, + shape=[batch, config.num_attention_heads, config.head_dim, config.head_dim], ) - flat.append(rec_state) pairs.append((rec_state,)) # 1-tuple: lightning has no conv_state elif ltype == "linear_attention": - conv_state = ir.Value( - name=f"{prefix}.{i}.conv_state", - shape=ir.Shape([batch, dims.conv_dim, dims.conv_kernel - 1]), - type=ir.TensorType(dtype), + conv_state = builder.input( + f"{prefix}.{i}.conv_state", + dtype=dtype, + shape=[batch, dims.conv_dim, dims.conv_kernel - 1], ) - rec_state = ir.Value( - name=f"{prefix}.{i}.recurrent_state", - shape=ir.Shape([batch, dims.num_v_heads, dims.head_k_dim, dims.head_v_dim]), - type=ir.TensorType(dtype), + rec_state = builder.input( + f"{prefix}.{i}.recurrent_state", + dtype=dtype, + shape=[batch, dims.num_v_heads, dims.head_k_dim, dims.head_v_dim], ) - flat.extend([conv_state, rec_state]) pairs.append((conv_state, rec_state)) elif ltype == "conv": # ShortConv layers: conv_state only (no SSM state) # State: (batch, hidden_size, short_conv_kernel - 1) short_conv_kernel = getattr(config, "short_conv_kernel", 3) - conv_state = ir.Value( - name=f"{prefix}.{i}.conv_state", - shape=ir.Shape([batch, config.hidden_size, short_conv_kernel - 1]), - type=ir.TensorType(dtype), + conv_state = builder.input( + f"{prefix}.{i}.conv_state", + dtype=dtype, + shape=[batch, config.hidden_size, short_conv_kernel - 1], ) - flat.append(conv_state) pairs.append((conv_state,)) # 1-tuple: conv has no second state elif ltype in ("mlp", "moe"): # MLP and MoE layers are stateless — no cache inputs needed pairs.append((None, None)) elif ltype == "mamba": - conv_state = ir.Value( - name=f"{prefix}.{i}.conv_state", - shape=ir.Shape([batch, mamba_d_inner, mamba_d_conv - 1]), - type=ir.TensorType(dtype), + conv_state = builder.input( + f"{prefix}.{i}.conv_state", + dtype=dtype, + shape=[batch, mamba_d_inner, mamba_d_conv - 1], ) - ssm_state = ir.Value( - name=f"{prefix}.{i}.ssm_state", - shape=ir.Shape([batch, mamba_d_inner, mamba_d_state]), - type=ir.TensorType(dtype), + ssm_state = builder.input( + f"{prefix}.{i}.ssm_state", + dtype=dtype, + shape=[batch, mamba_d_inner, mamba_d_state], ) - flat.extend([conv_state, ssm_state]) pairs.append((conv_state, ssm_state)) elif ltype == "mamba2": - conv_state = ir.Value( - name=f"{prefix}.{i}.conv_state", - shape=ir.Shape([batch, mamba2_conv_dim, mamba_d_conv - 1]), - type=ir.TensorType(dtype), + conv_state = builder.input( + f"{prefix}.{i}.conv_state", + dtype=dtype, + shape=[batch, mamba2_conv_dim, mamba_d_conv - 1], ) - ssm_state = ir.Value( - name=f"{prefix}.{i}.ssm_state", - shape=ir.Shape([batch, mamba2_n_heads, mamba2_d_state, mamba2_d_head]), - type=ir.TensorType(dtype), + ssm_state = builder.input( + f"{prefix}.{i}.ssm_state", + dtype=dtype, + shape=[batch, mamba2_n_heads, mamba2_d_state, mamba2_d_head], ) - flat.extend([conv_state, ssm_state]) pairs.append((conv_state, ssm_state)) else: - past_key = ir.Value( - name=f"{prefix}.{i}.key", - shape=ir.Shape( - [batch, config.num_key_value_heads, past_seq_len, config.head_dim] - ), - type=ir.TensorType(dtype), + past_key = builder.input( + f"{prefix}.{i}.key", + dtype=dtype, + shape=[batch, config.num_key_value_heads, past_seq_len, config.head_dim], ) - past_value = ir.Value( - name=f"{prefix}.{i}.value", - shape=ir.Shape( - [batch, config.num_key_value_heads, past_seq_len, config.head_dim] - ), - type=ir.TensorType(dtype), + past_value = builder.input( + f"{prefix}.{i}.value", + dtype=dtype, + shape=[batch, config.num_key_value_heads, past_seq_len, config.head_dim], ) - flat.extend([past_key, past_value]) pairs.append((past_key, past_value)) - return flat, pairs + return pairs def _register_hybrid_cache_outputs( - graph: ir.Graph, + builder: GraphBuilder, present_key_values: list[tuple[ir.Value, ...]], layer_types: list[str], *, @@ -295,27 +281,22 @@ def _register_hybrid_cache_outputs( if ltype == "lightning_attention": # Single recurrent state only (no conv_state for lightning) (state_a,) = states - state_a.name = f"{prefix}.{i}.recurrent_state" - graph.outputs.append(state_a) + builder.add_output(state_a, f"{prefix}.{i}.recurrent_state") elif ltype == "conv": # ShortConv: single conv_state only (state_a,) = states - state_a.name = f"{prefix}.{i}.conv_state" - graph.outputs.append(state_a) + builder.add_output(state_a, f"{prefix}.{i}.conv_state") else: state_a, state_b = states if ltype == "linear_attention": - state_a.name = f"{prefix}.{i}.conv_state" - state_b.name = f"{prefix}.{i}.recurrent_state" + builder.add_output(state_a, f"{prefix}.{i}.conv_state") + builder.add_output(state_b, f"{prefix}.{i}.recurrent_state") elif ltype in ("mamba", "mamba2"): - state_a.name = f"{prefix}.{i}.conv_state" - state_b.name = f"{prefix}.{i}.ssm_state" + builder.add_output(state_a, f"{prefix}.{i}.conv_state") + builder.add_output(state_b, f"{prefix}.{i}.ssm_state") else: - state_a.name = f"{prefix}.{i}.key" - state_b.name = f"{prefix}.{i}.value" - - graph.outputs.append(state_a) - graph.outputs.append(state_b) + builder.add_output(state_a, f"{prefix}.{i}.key") + builder.add_output(state_b, f"{prefix}.{i}.value") def _register_linear_attention_functions( diff --git a/src/mobius/tasks/_causal_lm.py b/src/mobius/tasks/_causal_lm.py index e6456c6b2..471d350b6 100644 --- a/src/mobius/tasks/_causal_lm.py +++ b/src/mobius/tasks/_causal_lm.py @@ -6,7 +6,7 @@ from __future__ import annotations import onnx_ir as ir -from onnxscript import nn +from onnxscript import GraphBuilder, nn from mobius._configs import ArchitectureConfig from mobius._model_package import ModelPackage @@ -110,23 +110,21 @@ def build( batch = ir.SymbolicDim("batch") seq_len = ir.SymbolicDim("sequence_len") + # --- Build graph first, then create inputs via builder --- + graph, builder = _make_graph() + op = builder.op + # --- Inputs common to both modes --- - input_ids = ir.Value( - name="input_ids", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), - ) - position_ids = ir.Value( - name="position_ids", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), - ) + input_ids = builder.input("input_ids", dtype=ir.DataType.INT64, shape=[batch, seq_len]) # --- Cache setup (static vs dynamic) --- if static: attention_mask = None - graph_inputs = [input_ids, position_ids] - cache_inputs, past_key_values = _make_static_cache_inputs( + position_ids = builder.input( + "position_ids", dtype=ir.DataType.INT64, shape=[batch, seq_len] + ) + past_key_values = _make_static_cache_inputs( + builder, config.num_hidden_layers, config.num_key_value_heads, config.head_dim, @@ -136,12 +134,14 @@ def build( ) else: past_seq_len = ir.SymbolicDim("past_sequence_len") - 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"], + ) + position_ids = builder.input( + "position_ids", dtype=ir.DataType.INT64, shape=[batch, seq_len] ) - graph_inputs = [input_ids, attention_mask, position_ids] # MLA attention: K/V heads equal q heads (no GQA reduction in # latent space). The ONNX Attention op is called with @@ -154,7 +154,8 @@ def build( config.num_attention_heads if use_mla else config.num_key_value_heads ) - cache_inputs, past_key_values = _make_kv_cache_inputs( + past_key_values = _make_kv_cache_inputs( + builder, config.num_hidden_layers, num_kv_cache_heads, config.head_dim, @@ -166,12 +167,6 @@ def build( value_head_dim=config.v_head_dim or None, ) - graph_inputs.extend(cache_inputs) - - # --- Build graph, invoke module, collect outputs --- - graph, builder = _make_graph(graph_inputs) - op = builder.op - logits, present_key_values = module( op, input_ids=input_ids, @@ -180,18 +175,17 @@ def build( past_key_values=past_key_values, ) - logits.name = "logits" - graph.outputs.append(logits) + builder.add_output(logits, "logits") # --- Output registration (static vs dynamic) --- if static: _register_static_cache_outputs( - graph, + builder, present_key_values, ) else: _register_kv_cache_outputs( - graph, + builder, present_key_values, ) @@ -228,34 +222,26 @@ def build( seq_len = ir.SymbolicDim("sequence_len") past_seq_len = ir.SymbolicDim("past_sequence_len") - input_ids = ir.Value( - name="input_ids", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), - ) - attention_mask = ir.Value( - name="attention_mask", - shape=ir.Shape([batch, "past_seq_len + seq_len"]), - type=ir.TensorType(ir.DataType.INT64), + graph, builder = _make_graph() + op = builder.op + + input_ids = builder.input("input_ids", dtype=ir.DataType.INT64, shape=[batch, seq_len]) + attention_mask = builder.input( + "attention_mask", + dtype=ir.DataType.INT64, + shape=[batch, "past_seq_len + seq_len"], ) - position_ids = ir.Value( - name="position_ids", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), + position_ids = builder.input( + "position_ids", dtype=ir.DataType.INT64, shape=[batch, seq_len] ) - graph_inputs = [input_ids, attention_mask, position_ids] - - cache_inputs, past_key_values = _make_hybrid_cache_inputs( + past_key_values = _make_hybrid_cache_inputs( + builder, config, config.dtype, batch, past_seq_len, ) - graph_inputs.extend(cache_inputs) - - graph, builder = _make_graph(graph_inputs) - op = builder.op logits, present_key_values = module( op, @@ -265,10 +251,9 @@ def build( past_key_values=past_key_values, ) - logits.name = "logits" - graph.outputs.append(logits) + builder.add_output(logits, "logits") _register_hybrid_cache_outputs( - graph, + builder, present_key_values, config.layer_types or [], ) @@ -279,51 +264,49 @@ def build( def _make_static_cache_inputs( + builder: GraphBuilder, num_layers: int, num_key_value_heads: int, head_dim: int, dtype: ir.DataType, batch: ir.SymbolicDim, max_seq_len: int, -) -> tuple[list[ir.Value], list[StaticCacheState]]: +) -> list[StaticCacheState]: """Create static KV cache inputs for ``num_layers`` layers. + Uses ``builder.input()`` to create and register graph inputs directly. + Returns: - ``(flat_inputs, static_caches)`` where *flat_inputs* is a flat - list suitable for extending ``graph_inputs``, and - *static_caches* is a list of :class:`StaticCacheState` tuples - for passing to the module via ``past_key_values``. + A list of :class:`StaticCacheState` tuples for passing to the + module via ``past_key_values``. """ kv_hidden = num_key_value_heads * head_dim - flat: list[ir.Value] = [] cache_pairs: list[tuple[ir.Value, ir.Value]] = [] for i in range(num_layers): - key_cache = ir.Value( - name=f"key_cache.{i}", - shape=ir.Shape([batch, max_seq_len, kv_hidden]), - type=ir.TensorType(dtype), + key_cache = builder.input( + f"key_cache.{i}", + dtype=dtype, + shape=[batch, max_seq_len, kv_hidden], ) - value_cache = ir.Value( - name=f"value_cache.{i}", - shape=ir.Shape([batch, max_seq_len, kv_hidden]), - type=ir.TensorType(dtype), + value_cache = builder.input( + f"value_cache.{i}", + dtype=dtype, + shape=[batch, max_seq_len, kv_hidden], ) - flat.extend([key_cache, value_cache]) cache_pairs.append((key_cache, value_cache)) # Shared inputs across all layers - write_indices = ir.Value( - name="write_indices", - shape=ir.Shape([batch]), - type=ir.TensorType(ir.DataType.INT64), + write_indices = builder.input( + "write_indices", + dtype=ir.DataType.INT64, + shape=[batch], ) - nonpad_kv_seqlen = ir.Value( - name="nonpad_kv_seqlen", - shape=ir.Shape([batch]), - type=ir.TensorType(ir.DataType.INT64), + nonpad_kv_seqlen = builder.input( + "nonpad_kv_seqlen", + dtype=ir.DataType.INT64, + shape=[batch], ) - flat.extend([write_indices, nonpad_kv_seqlen]) # Build StaticCacheState for each layer (shared indices) static_caches: list[StaticCacheState] = [] @@ -337,11 +320,11 @@ def _make_static_cache_inputs( ) ) - return flat, static_caches + return static_caches def _register_static_cache_outputs( - graph: ir.Graph, + builder: GraphBuilder, present_key_values: list[tuple[ir.Value, ir.Value]], ) -> None: """Name and register static cache outputs on the graph. @@ -350,10 +333,8 @@ def _register_static_cache_outputs( that runs during model optimization. """ for i, (updated_key, updated_value) in enumerate(present_key_values): - updated_key.name = f"updated_key_cache.{i}" - updated_value.name = f"updated_value_cache.{i}" - graph.outputs.append(updated_key) - graph.outputs.append(updated_value) + builder.add_output(updated_key, f"updated_key_cache.{i}") + builder.add_output(updated_value, f"updated_value_cache.{i}") def _validate_static_cache_support(module: nn.Module) -> None: diff --git a/src/mobius/tasks/_codec.py b/src/mobius/tasks/_codec.py index 604574297..ace4f3afb 100644 --- a/src/mobius/tasks/_codec.py +++ b/src/mobius/tasks/_codec.py @@ -65,17 +65,12 @@ def _build_decoder( num_q = ir.SymbolicDim("num_quantizers") seq_len = ir.SymbolicDim("sequence_len") - codes = ir.Value( - name="codes", - shape=ir.Shape([batch, num_q, seq_len]), - type=ir.TensorType(ir.DataType.INT64), - ) + graph, builder = _make_graph() + codes = builder.input("codes", dtype=ir.DataType.INT64, shape=[batch, num_q, seq_len]) - graph, builder = _make_graph([codes]) waveform = decoder(builder.op, codes) - waveform.name = "waveform" - graph.outputs.append(waveform) + builder.add_output(waveform, "waveform") return _make_model(graph) def _build_encoder( @@ -93,15 +88,12 @@ def _build_encoder( batch = ir.SymbolicDim("batch") audio_len = ir.SymbolicDim("audio_length") - waveform = ir.Value( - name="waveform", - shape=ir.Shape([batch, 1, audio_len]), - type=ir.TensorType(ir.DataType.FLOAT), + graph, builder = _make_graph(name="encoder") + waveform = builder.input( + "waveform", dtype=ir.DataType.FLOAT, shape=[batch, 1, audio_len] ) - graph, builder = _make_graph([waveform], name="encoder") codes = encoder(builder.op, waveform) - codes.name = "codes" - graph.outputs.append(codes) + builder.add_output(codes, "codes") return _make_model(graph) diff --git a/src/mobius/tasks/_controlnet.py b/src/mobius/tasks/_controlnet.py index c0c82bdd5..4354250fc 100644 --- a/src/mobius/tasks/_controlnet.py +++ b/src/mobius/tasks/_controlnet.py @@ -28,33 +28,25 @@ def build( module, config: ControlNetConfig, ) -> ModelPackage: - sample = ir.Value( - name="sample", - type=ir.TensorType(ir.DataType.FLOAT), - shape=ir.Shape(("batch", config.in_channels, "height", "width")), - ) - timestep = ir.Value( - name="timestep", - type=ir.TensorType(ir.DataType.INT64), - shape=ir.Shape(("batch",)), - ) - encoder_hidden_states = ir.Value( - name="encoder_hidden_states", - type=ir.TensorType(ir.DataType.FLOAT), - shape=ir.Shape(("batch", "sequence_length", config.cross_attention_dim)), + graph, builder = _make_graph() + op = builder.op + + sample = builder.input( + "sample", + dtype=ir.DataType.FLOAT, + shape=["batch", config.in_channels, "height", "width"], ) - controlnet_cond = ir.Value( - name="controlnet_cond", - type=ir.TensorType(ir.DataType.FLOAT), - shape=ir.Shape( - ("batch", config.conditioning_channels, "cond_height", "cond_width") - ), + timestep = builder.input("timestep", dtype=ir.DataType.INT64, shape=["batch"]) + encoder_hidden_states = builder.input( + "encoder_hidden_states", + dtype=ir.DataType.FLOAT, + shape=["batch", "sequence_length", config.cross_attention_dim], ) - - graph, builder = _make_graph( - [sample, timestep, encoder_hidden_states, controlnet_cond] + controlnet_cond = builder.input( + "controlnet_cond", + dtype=ir.DataType.FLOAT, + shape=["batch", config.conditioning_channels, "cond_height", "cond_width"], ) - op = builder.op down_outputs, mid_output = module( op, @@ -66,9 +58,7 @@ def build( # Register outputs for i, out in enumerate(down_outputs): - out.name = f"down_block_res_{i}" - graph.outputs.append(out) - mid_output.name = "mid_block_res" - graph.outputs.append(mid_output) + builder.add_output(out, f"down_block_res_{i}") + builder.add_output(mid_output, "mid_block_res") return ModelPackage({"model": _make_model(graph)}, config=config) diff --git a/src/mobius/tasks/_denoising.py b/src/mobius/tasks/_denoising.py index 96823f214..18c6f758a 100644 --- a/src/mobius/tasks/_denoising.py +++ b/src/mobius/tasks/_denoising.py @@ -28,25 +28,21 @@ def build( module, config: UNet2DConfig, ) -> ModelPackage: - sample = ir.Value( - name="sample", - type=ir.TensorType(ir.DataType.FLOAT), - shape=ir.Shape(("batch", config.in_channels, "height", "width")), - ) - timestep = ir.Value( - name="timestep", - type=ir.TensorType(ir.DataType.INT64), - shape=ir.Shape(("batch",)), + graph, builder = _make_graph() + op = builder.op + + sample = builder.input( + "sample", + dtype=ir.DataType.FLOAT, + shape=["batch", config.in_channels, "height", "width"], ) - encoder_hidden_states = ir.Value( - name="encoder_hidden_states", - type=ir.TensorType(ir.DataType.FLOAT), - shape=ir.Shape(("batch", "sequence_length", config.cross_attention_dim)), + timestep = builder.input("timestep", dtype=ir.DataType.INT64, shape=["batch"]) + encoder_hidden_states = builder.input( + "encoder_hidden_states", + dtype=ir.DataType.FLOAT, + shape=["batch", "sequence_length", config.cross_attention_dim], ) - graph, builder = _make_graph([sample, timestep, encoder_hidden_states]) - op = builder.op - noise_pred = module( op, sample=sample, @@ -54,7 +50,6 @@ def build( encoder_hidden_states=encoder_hidden_states, ) - noise_pred.name = "noise_pred" - graph.outputs.append(noise_pred) + builder.add_output(noise_pred, "noise_pred") return ModelPackage({"model": _make_model(graph)}, config=config) diff --git a/src/mobius/tasks/_feature_extraction.py b/src/mobius/tasks/_feature_extraction.py index 872181a11..0d32b55e7 100644 --- a/src/mobius/tasks/_feature_extraction.py +++ b/src/mobius/tasks/_feature_extraction.py @@ -37,25 +37,17 @@ def build( batch = ir.SymbolicDim("batch") seq_len = ir.SymbolicDim("sequence_len") - input_ids = ir.Value( - name="input_ids", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), - ) - attention_mask = ir.Value( - name="attention_mask", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), + graph, builder = _make_graph() + op = builder.op + + input_ids = builder.input("input_ids", dtype=ir.DataType.INT64, shape=[batch, seq_len]) + attention_mask = builder.input( + "attention_mask", dtype=ir.DataType.INT64, shape=[batch, seq_len] ) - token_type_ids = ir.Value( - name="token_type_ids", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), + token_type_ids = builder.input( + "token_type_ids", dtype=ir.DataType.INT64, shape=[batch, seq_len] ) - graph, builder = _make_graph([input_ids, attention_mask, token_type_ids]) - op = builder.op - last_hidden_state = module( op, input_ids=input_ids, @@ -63,7 +55,6 @@ def build( token_type_ids=token_type_ids, ) - 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) diff --git a/src/mobius/tasks/_gemma4.py b/src/mobius/tasks/_gemma4.py index ccb73e12f..532272b0c 100644 --- a/src/mobius/tasks/_gemma4.py +++ b/src/mobius/tasks/_gemma4.py @@ -22,7 +22,7 @@ from __future__ import annotations import onnx_ir as ir -from onnxscript import nn +from onnxscript import GraphBuilder, nn from mobius._configs import Gemma4Config from mobius._model_package import ModelPackage @@ -37,12 +37,15 @@ def _make_gemma4_kv_cache_inputs( + builder: GraphBuilder, config: Gemma4Config, batch: ir.SymbolicDim, past_seq_len: ir.SymbolicDim, -) -> tuple[list[ir.Value], list[tuple[ir.Value, ir.Value]]]: +) -> list[tuple[ir.Value, ir.Value]]: """Create per-layer KV cache inputs accounting for dual head_dim and KV sharing. + Uses ``builder.input()`` to create and register graph inputs directly. + Local (sliding_attention) layers use ``config.head_dim``; global (full_attention) layers use ``config.global_head_dim``. @@ -66,25 +69,23 @@ def _make_gemma4_kv_cache_inputs( f"num_hidden_layers ({config.num_hidden_layers})" ) - flat: list[ir.Value] = [] pairs: list[tuple[ir.Value, ir.Value]] = [] for i in range(num_kv_layers): layer_type = layer_types[i] if i < len(layer_types) else "sliding_attention" hd = global_head_dim if layer_type == "full_attention" else local_head_dim kv_heads = config.num_key_value_heads - past_key = ir.Value( - name=f"past_key_values.{i}.key", - shape=ir.Shape([batch, kv_heads, past_seq_len, hd]), - type=ir.TensorType(config.dtype), + past_key = builder.input( + f"past_key_values.{i}.key", + dtype=config.dtype, + shape=[batch, kv_heads, past_seq_len, hd], ) - past_value = ir.Value( - name=f"past_key_values.{i}.value", - shape=ir.Shape([batch, kv_heads, past_seq_len, hd]), - type=ir.TensorType(config.dtype), + past_value = builder.input( + f"past_key_values.{i}.value", + dtype=config.dtype, + shape=[batch, kv_heads, past_seq_len, hd], ) - flat.extend([past_key, past_value]) pairs.append((past_key, past_value)) - return flat, pairs + return pairs class Gemma4TextCausalLMTask(ModelTask): @@ -117,28 +118,26 @@ def build( seq_len = ir.SymbolicDim("sequence_len") past_seq_len = ir.SymbolicDim("past_sequence_len") - input_ids = ir.Value( - name="input_ids", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), + graph, builder = _make_graph() + op = builder.op + + input_ids = builder.input( + "input_ids", + dtype=ir.DataType.INT64, + shape=[batch, seq_len], ) - 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"], ) - position_ids = ir.Value( - name="position_ids", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), + position_ids = builder.input( + "position_ids", + dtype=ir.DataType.INT64, + shape=[batch, seq_len], ) - graph_inputs = [input_ids, attention_mask, position_ids] - kv_inputs, past_key_values = _make_gemma4_kv_cache_inputs(config, batch, past_seq_len) - graph_inputs.extend(kv_inputs) - - graph, graph_builder = _make_graph(graph_inputs) - op = graph_builder.op + past_key_values = _make_gemma4_kv_cache_inputs(builder, config, batch, past_seq_len) logits, present_key_values = module( op, @@ -147,9 +146,8 @@ def build( position_ids=position_ids, past_key_values=past_key_values, ) - logits.name = "logits" - graph.outputs.append(logits) - _register_kv_cache_outputs(graph, present_key_values) + builder.add_output(logits, "logits") + _register_kv_cache_outputs(builder, present_key_values) return ModelPackage({"model": _make_model(graph)}, config=config) @@ -226,34 +224,31 @@ def _build_decoder( 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(name="decoder") + op = builder.op + + 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"], ) - position_ids = ir.Value( - name="position_ids", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), + position_ids = builder.input( + "position_ids", + dtype=ir.DataType.INT64, + shape=[batch, seq_len], ) - input_ids = ir.Value( - name="input_ids", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), + input_ids = builder.input( + "input_ids", + dtype=ir.DataType.INT64, + shape=[batch, seq_len], ) - graph_inputs = [inputs_embeds, attention_mask, position_ids, input_ids] - - kv_inputs, past_key_values = _make_gemma4_kv_cache_inputs(config, batch, past_seq_len) - graph_inputs.extend(kv_inputs) - - graph, graph_builder = _make_graph(graph_inputs, name="decoder") - op = graph_builder.op + past_key_values = _make_gemma4_kv_cache_inputs(builder, config, batch, past_seq_len) logits, present_key_values = decoder( op, @@ -264,9 +259,8 @@ def _build_decoder( past_key_values=past_key_values, ) - logits.name = "logits" - graph.outputs.append(logits) - _register_kv_cache_outputs(graph, present_key_values) + builder.add_output(logits, "logits") + _register_kv_cache_outputs(builder, present_key_values) return _make_model(graph) @@ -293,30 +287,27 @@ def _build_vision( patch_size = config.vision.patch_size or 16 if config.vision else 16 pixel_dim = 3 * patch_size * patch_size - pixel_values = ir.Value( - name="pixel_values", - shape=ir.Shape([batch, num_patches, pixel_dim]), - type=ir.TensorType(config.dtype), + graph, builder = _make_graph(name="vision_encoder") + op = builder.op + + pixel_values = builder.input( + "pixel_values", + dtype=config.dtype, + shape=[batch, num_patches, pixel_dim], ) - pixel_position_ids = ir.Value( - name="pixel_position_ids", - shape=ir.Shape([batch, num_patches, 2]), - type=ir.TensorType(ir.DataType.INT64), + pixel_position_ids = builder.input( + "pixel_position_ids", + dtype=ir.DataType.INT64, + shape=[batch, num_patches, 2], ) - graph_inputs = [pixel_values, pixel_position_ids] - - graph, graph_builder = _make_graph(graph_inputs, name="vision_encoder") - op = graph_builder.op - image_features = vision( op, pixel_values=pixel_values, pixel_position_ids=pixel_position_ids, ) - image_features.name = "image_features" - graph.outputs.append(image_features) + builder.add_output(image_features, "image_features") return _make_model(graph) @@ -348,33 +339,29 @@ def _build_audio( time = ir.SymbolicDim("time") input_size = (config.audio.input_size if config.audio else None) or 128 - input_features = ir.Value( - name="input_features", - shape=ir.Shape([batch, time, input_size]), - type=ir.TensorType(config.dtype), - ) - input_features_mask = ir.Value( - name="input_features_mask", - shape=ir.Shape([batch, time]), - type=ir.TensorType(ir.DataType.BOOL), - ) + graph, builder = _make_graph(name="audio_encoder") + op = builder.op - graph, graph_builder = _make_graph( - [input_features, input_features_mask], name="audio_encoder" + input_features = builder.input( + "input_features", + dtype=config.dtype, + shape=[batch, time, input_size], + ) + input_features_mask = builder.input( + "input_features_mask", + dtype=ir.DataType.BOOL, + shape=[batch, time], ) - op = graph_builder.op audio_features, downsampled_mask = audio( op, input_features, input_features_mask=input_features_mask, ) - audio_features.name = "audio_features" - graph.outputs.append(audio_features) + builder.add_output(audio_features, "audio_features") if downsampled_mask is not None: - downsampled_mask.name = "audio_features_mask" - graph.outputs.append(downsampled_mask) + builder.add_output(downsampled_mask, "audio_features_mask") return _make_model(graph) @@ -388,31 +375,29 @@ def _build_embedding( seq_len = ir.SymbolicDim("sequence_len") num_image_tokens = ir.SymbolicDim("num_image_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") + op = builder.op + + input_ids = builder.input( + "input_ids", + dtype=ir.DataType.INT64, + shape=[batch, seq_len], ) - image_features = ir.Value( - name="image_features", - shape=ir.Shape([num_image_tokens, config.hidden_size]), - type=ir.TensorType(config.dtype), + image_features = builder.input( + "image_features", + dtype=config.dtype, + shape=[num_image_tokens, config.hidden_size], ) - graph_inputs = [input_ids, image_features] audio_features_val: ir.Value | None = None if config.audio is not None: num_audio_tokens = ir.SymbolicDim("num_audio_tokens") - audio_features_val = ir.Value( - name="audio_features", - shape=ir.Shape([num_audio_tokens, config.hidden_size]), - type=ir.TensorType(config.dtype), + audio_features_val = builder.input( + "audio_features", + dtype=config.dtype, + shape=[num_audio_tokens, config.hidden_size], ) - graph_inputs.append(audio_features_val) - - graph, graph_builder = _make_graph(graph_inputs, name="embedding") - op = graph_builder.op inputs_embeds = embedding( op, @@ -420,6 +405,5 @@ def _build_embedding( image_features=image_features, audio_features=audio_features_val, ) - inputs_embeds.name = "inputs_embeds" - graph.outputs.append(inputs_embeds) + builder.add_output(inputs_embeds, "inputs_embeds") return _make_model(graph) diff --git a/src/mobius/tasks/_image_classification.py b/src/mobius/tasks/_image_classification.py index 44a575a97..cb3b5dc57 100644 --- a/src/mobius/tasks/_image_classification.py +++ b/src/mobius/tasks/_image_classification.py @@ -37,18 +37,17 @@ def build( image_size = getattr(config, "image_size", 224) num_channels = getattr(config, "num_channels", 3) - pixel_values = ir.Value( - name="pixel_values", - shape=ir.Shape([batch, num_channels, image_size, image_size]), - type=ir.TensorType(ir.DataType.FLOAT), - ) - - graph, builder = _make_graph([pixel_values]) + graph, builder = _make_graph() op = builder.op + pixel_values = builder.input( + "pixel_values", + dtype=ir.DataType.FLOAT, + shape=[batch, num_channels, image_size, image_size], + ) + last_hidden_state = module(op, pixel_values=pixel_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) diff --git a/src/mobius/tasks/_multimodal.py b/src/mobius/tasks/_multimodal.py index d96dd3e8b..01f7d5e35 100644 --- a/src/mobius/tasks/_multimodal.py +++ b/src/mobius/tasks/_multimodal.py @@ -52,39 +52,41 @@ def build( seq_len = ir.SymbolicDim("sequence_len") past_seq_len = ir.SymbolicDim("past_sequence_len") - input_ids = ir.Value( - name="input_ids", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), + graph, builder = _make_graph() + op = builder.op + + input_ids = builder.input( + "input_ids", + dtype=ir.DataType.INT64, + shape=[batch, seq_len], ) - 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"], ) - position_ids = ir.Value( - name="position_ids", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), + position_ids = builder.input( + "position_ids", + dtype=ir.DataType.INT64, + shape=[batch, seq_len], ) image_size = config.vision.image_size or 224 if config.vision else 224 - pixel_values = ir.Value( - name="pixel_values", - shape=ir.Shape([batch, 3, image_size, image_size]), - type=ir.TensorType(config.dtype), + pixel_values = builder.input( + "pixel_values", + dtype=config.dtype, + shape=[batch, 3, image_size, image_size], ) audio_input_size = (config.audio.input_size if config.audio else None) or 80 - audio_features = ir.Value( - name="audio_features", - shape=ir.Shape([batch, "audio_seq_len", audio_input_size]), - type=ir.TensorType(config.dtype), + audio_features = builder.input( + "audio_features", + dtype=config.dtype, + shape=[batch, "audio_seq_len", audio_input_size], ) - graph_inputs = [input_ids, attention_mask, position_ids, pixel_values, audio_features] - - kv_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, @@ -92,10 +94,6 @@ def build( batch, past_seq_len, ) - graph_inputs.extend(kv_inputs) - - graph, builder = _make_graph(graph_inputs) - op = builder.op logits, present_key_values = module( op, @@ -107,8 +105,7 @@ def build( past_key_values=past_key_values, ) - logits.name = "logits" - graph.outputs.append(logits) - _register_kv_cache_outputs(graph, present_key_values) + builder.add_output(logits, "logits") + _register_kv_cache_outputs(builder, present_key_values) return ModelPackage({"model": _make_model(graph)}, config=config) diff --git a/src/mobius/tasks/_object_detection.py b/src/mobius/tasks/_object_detection.py index 2d06f690d..6bdf5c7ed 100644 --- a/src/mobius/tasks/_object_detection.py +++ b/src/mobius/tasks/_object_detection.py @@ -38,20 +38,18 @@ def build( image_size = getattr(config, "image_size", 224) num_channels = getattr(config, "num_channels", 3) - pixel_values = ir.Value( - name="pixel_values", - shape=ir.Shape([batch, num_channels, image_size, image_size]), - type=ir.TensorType(ir.DataType.FLOAT), - ) - - graph, builder = _make_graph([pixel_values]) + graph, builder = _make_graph() op = builder.op + pixel_values = builder.input( + "pixel_values", + dtype=ir.DataType.FLOAT, + shape=[batch, num_channels, image_size, image_size], + ) + logits, pred_boxes = module(op, pixel_values=pixel_values) - logits.name = "logits" - pred_boxes.name = "pred_boxes" - graph.outputs.append(logits) - graph.outputs.append(pred_boxes) + builder.add_output(logits, "logits") + builder.add_output(pred_boxes, "pred_boxes") return ModelPackage({"model": _make_model(graph)}, config=config) diff --git a/src/mobius/tasks/_phi4mm_multimodal.py b/src/mobius/tasks/_phi4mm_multimodal.py index d35ab3fa2..2e148a170 100644 --- a/src/mobius/tasks/_phi4mm_multimodal.py +++ b/src/mobius/tasks/_phi4mm_multimodal.py @@ -84,22 +84,22 @@ def _build_vision( num_images = ir.SymbolicDim("num_images") image_size = (config.vision.image_size if config.vision else None) or 448 - pixel_values = ir.Value( - name="pixel_values", - shape=ir.Shape([batch, 3, image_size, image_size]), - type=ir.TensorType(config.dtype), + graph, builder = _make_graph(name="vision_encoder") + + pixel_values = builder.input( + "pixel_values", + dtype=config.dtype, + shape=[batch, 3, image_size, image_size], ) - image_sizes = ir.Value( - name="image_sizes", - shape=ir.Shape([num_images, 2]), - type=ir.TensorType(ir.DataType.INT64), + image_sizes = builder.input( + "image_sizes", + dtype=ir.DataType.INT64, + shape=[num_images, 2], ) - graph, builder = _make_graph([pixel_values, image_sizes], name="vision_encoder") image_features = vision(builder.op, pixel_values, image_sizes=image_sizes) - image_features.name = "image_features" - graph.outputs.append(image_features) + builder.add_output(image_features, "image_features") return _make_model(graph) def _build_speech( @@ -117,26 +117,24 @@ def _build_speech( num_audio_clips = ir.SymbolicDim("num_audio_clips") input_size = (config.audio.input_size if config.audio else None) or 80 - audio_embeds = ir.Value( - name="audio_embeds", - shape=ir.Shape([batch, audio_seq_len, input_size]), - type=ir.TensorType(config.dtype), + graph, builder = _make_graph(name="audio_encoder") + + audio_embeds = builder.input( + "audio_embeds", + dtype=config.dtype, + shape=[batch, audio_seq_len, input_size], ) - audio_sizes = ir.Value( - name="audio_sizes", - shape=ir.Shape([num_audio_clips]), - type=ir.TensorType(ir.DataType.INT64), + audio_sizes = builder.input( + "audio_sizes", + dtype=ir.DataType.INT64, + shape=[num_audio_clips], ) - audio_projection_mode = ir.Value( - name="audio_projection_mode", - shape=ir.Shape([]), - type=ir.TensorType(ir.DataType.INT64), + audio_projection_mode = builder.input( + "audio_projection_mode", + dtype=ir.DataType.INT64, + shape=[], ) - graph, builder = _make_graph( - [audio_embeds, audio_sizes, audio_projection_mode], - name="audio_encoder", - ) speech_out = speech( builder.op, audio_embeds, @@ -144,8 +142,7 @@ def _build_speech( audio_projection_mode=audio_projection_mode, ) - speech_out.name = "audio_features" - graph.outputs.append(speech_out) + builder.add_output(speech_out, "audio_features") return _make_model(graph) def _build_embedding( @@ -159,26 +156,24 @@ def _build_embedding( num_image_tokens = ir.SymbolicDim("num_image_tokens") num_speech_tokens = ir.SymbolicDim("num_speech_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], ) - image_features = ir.Value( - name="image_features", - shape=ir.Shape([num_image_tokens, config.hidden_size]), - type=ir.TensorType(config.dtype), + image_features = builder.input( + "image_features", + dtype=config.dtype, + shape=[num_image_tokens, config.hidden_size], ) - audio_features = ir.Value( - name="audio_features", - shape=ir.Shape([num_speech_tokens, config.hidden_size]), - type=ir.TensorType(config.dtype), + audio_features = builder.input( + "audio_features", + dtype=config.dtype, + shape=[num_speech_tokens, config.hidden_size], ) - graph, builder = _make_graph( - [input_ids, image_features, audio_features], - name="embedding", - ) inputs_embeds = embedding( builder.op, input_ids=input_ids, @@ -186,6 +181,5 @@ def _build_embedding( audio_features=audio_features, ) - inputs_embeds.name = "inputs_embeds" - graph.outputs.append(inputs_embeds) + builder.add_output(inputs_embeds, "inputs_embeds") return _make_model(graph) diff --git a/src/mobius/tasks/_qwen_image_vae.py b/src/mobius/tasks/_qwen_image_vae.py index a03b09f74..3940753f5 100644 --- a/src/mobius/tasks/_qwen_image_vae.py +++ b/src/mobius/tasks/_qwen_image_vae.py @@ -40,20 +40,17 @@ def _build_encoder_graph( module, config: QwenImageVAEConfig, ) -> ir.Model: - sample = ir.Value( - name="sample", - type=ir.TensorType(ir.DataType.FLOAT), - shape=ir.Shape(("batch", 3, "frames", "height", "width")), - ) - - graph, builder = _make_graph([sample], name="vae_encoder") + graph, builder = _make_graph(name="vae_encoder") op = builder.op + sample = builder.input( + "sample", dtype=ir.DataType.FLOAT, shape=["batch", 3, "frames", "height", "width"] + ) + hidden_states = module.encoder(op, sample) hidden_states = module.quant_conv(op, hidden_states) - hidden_states.name = "latent_dist" - graph.outputs.append(hidden_states) + builder.add_output(hidden_states, "latent_dist") return _make_model(graph) @@ -62,19 +59,18 @@ def _build_decoder_graph( module, config: QwenImageVAEConfig, ) -> ir.Model: - latent_sample = ir.Value( - name="latent_sample", - type=ir.TensorType(ir.DataType.FLOAT), - shape=ir.Shape(("batch", config.z_dim, "frames", "height", "width")), - ) - - graph, builder = _make_graph([latent_sample], name="vae_decoder") + graph, builder = _make_graph(name="vae_decoder") op = builder.op + latent_sample = builder.input( + "latent_sample", + dtype=ir.DataType.FLOAT, + shape=["batch", config.z_dim, "frames", "height", "width"], + ) + hidden_states = module.post_quant_conv(op, latent_sample) hidden_states = module.decoder(op, hidden_states) - hidden_states.name = "sample" - graph.outputs.append(hidden_states) + builder.add_output(hidden_states, "sample") return _make_model(graph) diff --git a/src/mobius/tasks/_seq2seq.py b/src/mobius/tasks/_seq2seq.py index 147f176ee..0217fc9d3 100644 --- a/src/mobius/tasks/_seq2seq.py +++ b/src/mobius/tasks/_seq2seq.py @@ -18,9 +18,6 @@ _make_graph, _make_model, ) -from mobius.tasks._cache_utils import ( - _make_kv_cache_inputs, -) class Seq2SeqTask(ModelTask): @@ -58,26 +55,19 @@ def _build_encoder_graph( batch = ir.SymbolicDim("batch") seq_len = ir.SymbolicDim("sequence_len") - input_ids = ir.Value( - name="input_ids", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), - ) - attention_mask = ir.Value( - name="attention_mask", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), - ) - - graph, builder = _make_graph([input_ids, attention_mask], name="encoder") + graph, builder = _make_graph(name="encoder") op = builder.op + input_ids = builder.input("input_ids", dtype=ir.DataType.INT64, shape=[batch, seq_len]) + attention_mask = builder.input( + "attention_mask", dtype=ir.DataType.INT64, shape=[batch, seq_len] + ) + encoder_hidden_states = module.encoder( op, input_ids=input_ids, attention_mask=attention_mask ) - encoder_hidden_states.name = "last_hidden_state" - graph.outputs.append(encoder_hidden_states) + builder.add_output(encoder_hidden_states, "last_hidden_state") return _make_model(graph) @@ -91,63 +81,58 @@ def _build_decoder_graph( enc_seq_len = ir.SymbolicDim("encoder_sequence_len") past_seq_len = ir.SymbolicDim("past_sequence_len") - input_ids = ir.Value( - name="input_ids", - shape=ir.Shape([batch, dec_seq_len]), - type=ir.TensorType(ir.DataType.INT64), + graph, builder = _make_graph() + op = builder.op + + input_ids = builder.input( + "input_ids", + dtype=ir.DataType.INT64, + shape=[batch, dec_seq_len], ) - encoder_hidden_states = ir.Value( - name="encoder_hidden_states", - shape=ir.Shape([batch, enc_seq_len, config.hidden_size]), - type=ir.TensorType(config.dtype), + encoder_hidden_states = builder.input( + "encoder_hidden_states", + dtype=config.dtype, + shape=[batch, enc_seq_len, config.hidden_size], ) - attention_mask = ir.Value( - name="attention_mask", - shape=ir.Shape([batch, "past_seq_len + dec_seq_len"]), - type=ir.TensorType(ir.DataType.INT64), + attention_mask = builder.input( + "attention_mask", + dtype=ir.DataType.INT64, + shape=[batch, "past_seq_len + dec_seq_len"], ) - graph_inputs = [input_ids, encoder_hidden_states, attention_mask] - num_heads = config.num_attention_heads head_dim = config.head_dim num_decoder_layers = getattr(config, "num_decoder_layers", config.num_hidden_layers) - # Self-attention KV cache - self_kv_inputs, past_self_kvs = _make_kv_cache_inputs( - num_decoder_layers, - num_heads, - head_dim, - config.dtype, - batch, - past_seq_len, - prefix="past_key_values", - ) - # Use .self. naming for seq2seq self-attention KVs - for v in self_kv_inputs: - idx = v.name.split(".")[1] - kv_type = v.name.rsplit(".", 1)[-1] - v.name = f"past_key_values.{idx}.self.{kv_type}" - graph_inputs.extend(self_kv_inputs) - - # Cross-attention KV cache - cross_kv_inputs, cross_past_kvs = _make_kv_cache_inputs( - num_decoder_layers, - num_heads, - head_dim, - config.dtype, - batch, - enc_seq_len, - prefix="past_key_values", - ) - for v in cross_kv_inputs: - idx = v.name.split(".")[1] - kv_type = v.name.rsplit(".", 1)[-1] - v.name = f"past_key_values.{idx}.cross.{kv_type}" - graph_inputs.extend(cross_kv_inputs) - - graph, builder = _make_graph(graph_inputs) - op = builder.op + # Self-attention KV cache (named past_key_values.{i}.self.key/value) + past_self_kvs: list[tuple[ir.Value, ir.Value]] = [] + for i in range(num_decoder_layers): + past_key = builder.input( + f"past_key_values.{i}.self.key", + dtype=config.dtype, + shape=[batch, num_heads, past_seq_len, head_dim], + ) + past_value = builder.input( + f"past_key_values.{i}.self.value", + dtype=config.dtype, + shape=[batch, num_heads, past_seq_len, head_dim], + ) + past_self_kvs.append((past_key, past_value)) + + # Cross-attention KV cache (named past_key_values.{i}.cross.key/value) + cross_past_kvs: list[tuple[ir.Value, ir.Value]] = [] + for i in range(num_decoder_layers): + past_key = builder.input( + f"past_key_values.{i}.cross.key", + dtype=config.dtype, + shape=[batch, num_heads, enc_seq_len, head_dim], + ) + past_value = builder.input( + f"past_key_values.{i}.cross.value", + dtype=config.dtype, + shape=[batch, num_heads, enc_seq_len, head_dim], + ) + cross_past_kvs.append((past_key, past_value)) logits, present_self_kvs, present_cross_kvs = module.decoder( op, @@ -158,17 +143,14 @@ def _build_decoder_graph( cross_past_key_values=cross_past_kvs, ) - logits.name = "logits" - graph.outputs.append(logits) + builder.add_output(logits, "logits") for i, (k, v) in enumerate(present_self_kvs): - k.name = f"present.{i}.self.key" - v.name = f"present.{i}.self.value" - graph.outputs.extend([k, v]) + builder.add_output(k, f"present.{i}.self.key") + builder.add_output(v, f"present.{i}.self.value") for i, (k, v) in enumerate(present_cross_kvs): - k.name = f"present.{i}.cross.key" - v.name = f"present.{i}.cross.value" - graph.outputs.extend([k, v]) + builder.add_output(k, f"present.{i}.cross.key") + builder.add_output(v, f"present.{i}.cross.value") return _make_model(graph) diff --git a/src/mobius/tasks/_speech_language.py b/src/mobius/tasks/_speech_language.py index 9a0b79d53..855af2b73 100644 --- a/src/mobius/tasks/_speech_language.py +++ b/src/mobius/tasks/_speech_language.py @@ -82,15 +82,15 @@ def _build_audio_encoder( mel_seq = ir.SymbolicDim("mel_sequence_len") n_mels = (config.audio.num_mel_bins if config.audio else None) or 128 - input_features = ir.Value( - name="input_features", - shape=ir.Shape([batch, n_mels, mel_seq]), - type=ir.TensorType(config.dtype), + graph, builder = _make_graph(name="audio_encoder") + + input_features = builder.input( + "input_features", + dtype=config.dtype, + shape=[batch, n_mels, mel_seq], ) - graph, builder = _make_graph([input_features], name="audio_encoder") audio_features = audio_encoder(builder.op, input_features) - audio_features.name = "audio_features" - graph.outputs.append(audio_features) + builder.add_output(audio_features, "audio_features") return _make_model(graph) diff --git a/src/mobius/tasks/_speech_to_text.py b/src/mobius/tasks/_speech_to_text.py index ebed6edab..08854f035 100644 --- a/src/mobius/tasks/_speech_to_text.py +++ b/src/mobius/tasks/_speech_to_text.py @@ -69,19 +69,18 @@ def _build_encoder( batch = ir.SymbolicDim("batch") audio_seq_len = ir.SymbolicDim("audio_seq_len") - input_features = ir.Value( - name="input_features", - shape=ir.Shape([batch, config.num_mel_bins, audio_seq_len]), - type=ir.TensorType(ir.DataType.FLOAT), - ) - - graph, builder = _make_graph([input_features], name="encoder") + graph, builder = _make_graph(name="encoder") op = builder.op + input_features = builder.input( + "input_features", + dtype=ir.DataType.FLOAT, + shape=[batch, config.num_mel_bins, audio_seq_len], + ) + encoder_hidden_states = encoder(op, input_features=input_features) - encoder_hidden_states.name = "encoder_hidden_states" - graph.outputs.append(encoder_hidden_states) + builder.add_output(encoder_hidden_states, "encoder_hidden_states") return _make_model(graph) @@ -95,25 +94,27 @@ def _build_decoder( past_seq_len = ir.SymbolicDim("past_sequence_len") encoder_seq_len = ir.SymbolicDim("encoder_sequence_len") - decoder_input_ids = ir.Value( - name="decoder_input_ids", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), + graph, builder = _make_graph() + op = builder.op + + decoder_input_ids = builder.input( + "decoder_input_ids", + dtype=ir.DataType.INT64, + shape=[batch, seq_len], ) - encoder_hidden_states = ir.Value( - name="encoder_hidden_states", - shape=ir.Shape([batch, encoder_seq_len, config.hidden_size]), - type=ir.TensorType(config.dtype), + encoder_hidden_states = builder.input( + "encoder_hidden_states", + dtype=config.dtype, + shape=[batch, encoder_seq_len, config.hidden_size], ) - position_ids = ir.Value( - name="position_ids", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), + position_ids = builder.input( + "position_ids", + dtype=ir.DataType.INT64, + shape=[batch, seq_len], ) - graph_inputs = [decoder_input_ids, encoder_hidden_states, position_ids] - - kv_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, @@ -121,10 +122,6 @@ def _build_decoder( batch, past_seq_len, ) - graph_inputs.extend(kv_inputs) - - graph, builder = _make_graph(graph_inputs) - op = builder.op logits, present_key_values = decoder( op, @@ -134,8 +131,7 @@ def _build_decoder( past_key_values=past_key_values, ) - logits.name = "logits" - graph.outputs.append(logits) - _register_kv_cache_outputs(graph, present_key_values) + builder.add_output(logits, "logits") + _register_kv_cache_outputs(builder, present_key_values) return _make_model(graph) diff --git a/src/mobius/tasks/_ssm_causal_lm.py b/src/mobius/tasks/_ssm_causal_lm.py index 63f6f5046..efada9a7e 100644 --- a/src/mobius/tasks/_ssm_causal_lm.py +++ b/src/mobius/tasks/_ssm_causal_lm.py @@ -49,42 +49,34 @@ def _build_ssm_task( batch = ir.SymbolicDim("batch") seq_len = ir.SymbolicDim("sequence_len") - input_ids = ir.Value( - name="input_ids", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), - ) - graph_inputs: list[ir.Value] = [input_ids] + graph, builder = _make_graph() + + input_ids = builder.input("input_ids", dtype=ir.DataType.INT64, shape=[batch, seq_len]) past_states: list[tuple[ir.Value, ir.Value]] = [] for i in range(config.num_hidden_layers): - conv_state = ir.Value( - name=f"past_states.{i}.conv_state", - shape=ir.Shape([batch, *conv_state_shape]), - type=ir.TensorType(config.dtype), + conv_state = builder.input( + f"past_states.{i}.conv_state", + dtype=config.dtype, + shape=[batch, *conv_state_shape], ) - ssm_state = ir.Value( - name=f"past_states.{i}.ssm_state", - shape=ir.Shape([batch, *ssm_state_shape]), - type=ir.TensorType(config.dtype), + ssm_state = builder.input( + f"past_states.{i}.ssm_state", + dtype=config.dtype, + shape=[batch, *ssm_state_shape], ) - graph_inputs.extend([conv_state, ssm_state]) past_states.append((conv_state, ssm_state)) - graph, builder = _make_graph(graph_inputs) logits, present_states = module( builder.op, input_ids=input_ids, past_states=past_states, ) - logits.name = "logits" - graph.outputs.append(logits) + builder.add_output(logits, "logits") for i, (conv_state, ssm_state) in enumerate(present_states): - conv_state.name = f"present.{i}.conv_state" - ssm_state.name = f"present.{i}.ssm_state" - graph.outputs.append(conv_state) - graph.outputs.append(ssm_state) + builder.add_output(conv_state, f"present.{i}.conv_state") + builder.add_output(ssm_state, f"present.{i}.ssm_state") return ModelPackage({"model": _make_model(graph)}, config=config) diff --git a/src/mobius/tasks/_tts.py b/src/mobius/tasks/_tts.py index 729269954..d2a20d84f 100644 --- a/src/mobius/tasks/_tts.py +++ b/src/mobius/tasks/_tts.py @@ -90,26 +90,27 @@ def _build_talker( 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(name="talker") + + 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 (3, batch, seq_len) - position_ids = ir.Value( - name="position_ids", - shape=ir.Shape([3, batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), + position_ids = builder.input( + "position_ids", + dtype=ir.DataType.INT64, + shape=[3, batch, seq_len], ) - graph_inputs = [inputs_embeds, attention_mask, position_ids] - - kv_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, @@ -117,9 +118,7 @@ def _build_talker( batch, past_seq_len, ) - graph_inputs.extend(kv_inputs) - graph, builder = _make_graph(graph_inputs, name="talker") logits, last_hidden_state, present_key_values = talker( builder.op, inputs_embeds=inputs_embeds, @@ -128,11 +127,9 @@ def _build_talker( past_key_values=past_key_values, ) - logits.name = "logits" - last_hidden_state.name = "last_hidden_state" - graph.outputs.append(logits) - graph.outputs.append(last_hidden_state) - _register_kv_cache_outputs(graph, present_key_values) + builder.add_output(logits, "logits") + builder.add_output(last_hidden_state, "last_hidden_state") + _register_kv_cache_outputs(builder, present_key_values) return _make_model(graph) def _build_code_predictor( @@ -166,39 +163,35 @@ def _build_code_predictor( seq_len = ir.SymbolicDim("sequence_len") past_seq_len = ir.SymbolicDim("past_sequence_len") + graph, builder = _make_graph(name="code_predictor") + # Pre-embedded input in talker_hidden space (constructed by # generation loop). The model projects to cp_hidden internally. - inputs_embeds = ir.Value( - name="inputs_embeds", - shape=ir.Shape([batch, seq_len, config.hidden_size]), - type=ir.TensorType(config.dtype), + inputs_embeds = builder.input( + "inputs_embeds", + dtype=config.dtype, + shape=[batch, seq_len, config.hidden_size], ) # Step index: selects which lm_head to use (0..14) - step_index = ir.Value( - name="step_index", - shape=ir.Shape([]), - type=ir.TensorType(ir.DataType.INT64), + step_index = builder.input( + "step_index", + dtype=ir.DataType.INT64, + shape=[], ) - 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"], ) # 1D RoPE: 2D position_ids (batch, seq_len) - position_ids = ir.Value( - name="position_ids", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), + position_ids = builder.input( + "position_ids", + dtype=ir.DataType.INT64, + shape=[batch, seq_len], ) - graph_inputs = [ - inputs_embeds, - step_index, - attention_mask, - position_ids, - ] - - kv_inputs, past_key_values = _make_kv_cache_inputs( + past_key_values = _make_kv_cache_inputs( + builder, cp_num_hidden_layers, cp_num_key_value_heads, cp_head_dim, @@ -206,9 +199,7 @@ def _build_code_predictor( batch, past_seq_len, ) - graph_inputs.extend(kv_inputs) - graph, builder = _make_graph(graph_inputs, name="code_predictor") logits, present_key_values, codec_embeddings = code_predictor( builder.op, inputs_embeds=inputs_embeds, @@ -218,14 +209,12 @@ def _build_code_predictor( past_key_values=past_key_values, ) - logits.name = "logits" - graph.outputs.append(logits) + builder.add_output(logits, "logits") # Expose stacked codec embeddings for generation loop to extract. # The Identity node ensures renaming the output doesn't affect # the initializer name used for weight loading. - codec_embeddings.name = "codec_embeddings" - graph.outputs.append(codec_embeddings) - _register_kv_cache_outputs(graph, present_key_values) + builder.add_output(codec_embeddings, "codec_embeddings") + _register_kv_cache_outputs(builder, present_key_values) return _make_model(graph) def _build_embedding( @@ -238,28 +227,27 @@ def _build_embedding( text_seq = ir.SymbolicDim("text_sequence_len") codec_seq = ir.SymbolicDim("codec_sequence_len") - text_ids = ir.Value( - name="text_ids", - shape=ir.Shape([batch, text_seq]), - type=ir.TensorType(ir.DataType.INT64), + graph, builder = _make_graph(name="embedding") + + text_ids = builder.input( + "text_ids", + dtype=ir.DataType.INT64, + shape=[batch, text_seq], ) - codec_ids = ir.Value( - name="codec_ids", - shape=ir.Shape([batch, codec_seq]), - type=ir.TensorType(ir.DataType.INT64), + codec_ids = builder.input( + "codec_ids", + dtype=ir.DataType.INT64, + shape=[batch, codec_seq], ) - graph, builder = _make_graph([text_ids, codec_ids], name="embedding") text_embeds, codec_embeds = embedding( builder.op, text_ids=text_ids, codec_ids=codec_ids, ) - text_embeds.name = "text_embeds" - codec_embeds.name = "codec_embeds" - graph.outputs.append(text_embeds) - graph.outputs.append(codec_embeds) + builder.add_output(text_embeds, "text_embeds") + builder.add_output(codec_embeds, "codec_embeds") return _make_model(graph) def _build_speaker_encoder( @@ -274,15 +262,15 @@ def _build_speaker_encoder( se = tts.speaker_encoder if tts else None mel_dim = se.mel_dim if se else 128 - mel_input = ir.Value( - name="mel_input", - shape=ir.Shape([batch, mel_seq, mel_dim]), - type=ir.TensorType(config.dtype), + graph, builder = _make_graph(name="speaker_encoder") + + mel_input = builder.input( + "mel_input", + dtype=config.dtype, + shape=[batch, mel_seq, mel_dim], ) - graph, builder = _make_graph([mel_input], name="speaker_encoder") speaker_embedding = speaker_encoder(builder.op, mel_input) - speaker_embedding.name = "speaker_embedding" - graph.outputs.append(speaker_embedding) + builder.add_output(speaker_embedding, "speaker_embedding") return _make_model(graph) diff --git a/src/mobius/tasks/_vae.py b/src/mobius/tasks/_vae.py index 68f828884..c9dc62ee2 100644 --- a/src/mobius/tasks/_vae.py +++ b/src/mobius/tasks/_vae.py @@ -40,21 +40,20 @@ def _build_encoder_graph( module, config: VAEConfig, ) -> ir.Model: - sample = ir.Value( - name="sample", - type=ir.TensorType(ir.DataType.FLOAT), - shape=ir.Shape(("batch", config.in_channels, "height", "width")), - ) - - graph, builder = _make_graph([sample], name="vae_encoder") + graph, builder = _make_graph(name="vae_encoder") op = builder.op + sample = builder.input( + "sample", + dtype=ir.DataType.FLOAT, + shape=["batch", config.in_channels, "height", "width"], + ) + hidden_states = module.encoder(op, sample=sample) if module.quant_conv is not None: hidden_states = module.quant_conv(op, hidden_states) - hidden_states.name = "latent_dist" - graph.outputs.append(hidden_states) + builder.add_output(hidden_states, "latent_dist") return _make_model(graph) @@ -63,21 +62,20 @@ def _build_decoder_graph( module, config: VAEConfig, ) -> ir.Model: - latent_sample = ir.Value( - name="latent_sample", - type=ir.TensorType(ir.DataType.FLOAT), - shape=ir.Shape(("batch", config.latent_channels, "height", "width")), - ) - - graph, builder = _make_graph([latent_sample], name="vae_decoder") + graph, builder = _make_graph(name="vae_decoder") op = builder.op + latent_sample = builder.input( + "latent_sample", + dtype=ir.DataType.FLOAT, + shape=["batch", config.latent_channels, "height", "width"], + ) + hidden_states = latent_sample if module.post_quant_conv is not None: hidden_states = module.post_quant_conv(op, hidden_states) hidden_states = module.decoder(op, latent_sample=hidden_states) - hidden_states.name = "sample" - graph.outputs.append(hidden_states) + builder.add_output(hidden_states, "sample") return _make_model(graph) diff --git a/src/mobius/tasks/_video_denoising.py b/src/mobius/tasks/_video_denoising.py index 44a9a6e28..70645257f 100644 --- a/src/mobius/tasks/_video_denoising.py +++ b/src/mobius/tasks/_video_denoising.py @@ -29,33 +29,21 @@ def build( module, config: CogVideoXConfig, ) -> ModelPackage: - sample = ir.Value( - name="sample", - type=ir.TensorType(ir.DataType.FLOAT), - shape=ir.Shape( - ( - "batch", - "num_frames", - config.in_channels, - "height", - "width", - ) - ), - ) - timestep = ir.Value( - name="timestep", - type=ir.TensorType(ir.DataType.INT64), - shape=ir.Shape(("batch",)), + graph, builder = _make_graph() + op = builder.op + + sample = builder.input( + "sample", + dtype=ir.DataType.FLOAT, + shape=["batch", "num_frames", config.in_channels, "height", "width"], ) - encoder_hidden_states = ir.Value( - name="encoder_hidden_states", - type=ir.TensorType(ir.DataType.FLOAT), - shape=ir.Shape(("batch", "sequence_length", config.cross_attention_dim)), + timestep = builder.input("timestep", dtype=ir.DataType.INT64, shape=["batch"]) + encoder_hidden_states = builder.input( + "encoder_hidden_states", + dtype=ir.DataType.FLOAT, + shape=["batch", "sequence_length", config.cross_attention_dim], ) - graph, builder = _make_graph([sample, timestep, encoder_hidden_states]) - op = builder.op - noise_pred = module( op, sample=sample, @@ -63,7 +51,6 @@ def build( encoder_hidden_states=encoder_hidden_states, ) - noise_pred.name = "noise_pred" - graph.outputs.append(noise_pred) + builder.add_output(noise_pred, "noise_pred") return ModelPackage({"model": _make_model(graph)}, config=config) diff --git a/src/mobius/tasks/_vision_language.py b/src/mobius/tasks/_vision_language.py index 9be0ca28c..0fb530d27 100644 --- a/src/mobius/tasks/_vision_language.py +++ b/src/mobius/tasks/_vision_language.py @@ -47,49 +47,45 @@ def build( past_seq_len = ir.SymbolicDim("past_sequence_len") total_patches = ir.SymbolicDim("total_patches") - input_ids = ir.Value( - name="input_ids", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), + graph, builder = _make_graph() + op = builder.op + + input_ids = builder.input( + "input_ids", + dtype=ir.DataType.INT64, + shape=[batch, seq_len], ) - 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) - position_ids = ir.Value( - name="position_ids", - shape=ir.Shape([3, batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), + position_ids = builder.input( + "position_ids", + dtype=ir.DataType.INT64, + shape=[3, batch, seq_len], ) # Flattened image patches patch_size = config.vision.patch_size or 16 if config.vision else 16 temporal_patch_size = config.temporal_patch_size in_channels = config.vision.in_channels if config.vision else 3 pixel_dim = in_channels * temporal_patch_size * patch_size * patch_size - pixel_values = ir.Value( - name="pixel_values", - shape=ir.Shape([total_patches, pixel_dim]), - type=ir.TensorType(config.dtype), + pixel_values = builder.input( + "pixel_values", + dtype=config.dtype, + shape=[total_patches, pixel_dim], ) # Image grid dimensions for position embedding interpolation num_images = ir.SymbolicDim("num_images") - grid_thw = ir.Value( - name="grid_thw", - shape=ir.Shape([num_images, 3]), - type=ir.TensorType(ir.DataType.INT64), + grid_thw = builder.input( + "grid_thw", + dtype=ir.DataType.INT64, + shape=[num_images, 3], ) - graph_inputs = [ - input_ids, - attention_mask, - position_ids, - pixel_values, - grid_thw, - ] - - kv_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, @@ -97,10 +93,6 @@ def build( batch, past_seq_len, ) - graph_inputs.extend(kv_inputs) - - graph, builder = _make_graph(graph_inputs) - op = builder.op logits, present_key_values = module( op, @@ -112,8 +104,7 @@ def build( past_key_values=past_key_values, ) - logits.name = "logits" - graph.outputs.append(logits) - _register_kv_cache_outputs(graph, present_key_values) + builder.add_output(logits, "logits") + _register_kv_cache_outputs(builder, present_key_values) return ModelPackage({"model": _make_model(graph)}, config=config) diff --git a/src/mobius/tasks/_vision_language_3model.py b/src/mobius/tasks/_vision_language_3model.py index 735d7e787..00dae849e 100644 --- a/src/mobius/tasks/_vision_language_3model.py +++ b/src/mobius/tasks/_vision_language_3model.py @@ -82,17 +82,15 @@ def _build_vision( batch = ir.SymbolicDim("batch") image_size = (config.vision.image_size if config.vision else None) or 224 - pixel_values = ir.Value( - name="pixel_values", - shape=ir.Shape([batch, 3, image_size, image_size]), - type=ir.TensorType(config.dtype), + graph, builder = _make_graph(name="vision_encoder") + pixel_values = builder.input( + "pixel_values", + dtype=config.dtype, + shape=[batch, 3, image_size, image_size], ) + image_features = vision(builder.op, pixel_values=pixel_values) - graph, graph_builder = _make_graph([pixel_values], name="vision_encoder") - image_features = vision(graph_builder.op, pixel_values=pixel_values) - - image_features.name = "image_features" - graph.outputs.append(image_features) + builder.add_output(image_features, "image_features") return _make_model(graph) @@ -135,28 +133,25 @@ def _build_vision( in_channels = config.vision.in_channels if config.vision else 3 pixel_dim = in_channels * temporal_patch_size * patch_size * patch_size - pixel_values = ir.Value( - name="pixel_values", - shape=ir.Shape([total_patches, pixel_dim]), - type=ir.TensorType(config.dtype), + graph, builder = _make_graph(name="vision_encoder") + pixel_values = builder.input( + "pixel_values", + dtype=config.dtype, + shape=[total_patches, pixel_dim], ) - image_grid_thw = ir.Value( - name="image_grid_thw", - shape=ir.Shape([num_images, 3]), - type=ir.TensorType(ir.DataType.INT64), + image_grid_thw = builder.input( + "image_grid_thw", + dtype=ir.DataType.INT64, + shape=[num_images, 3], ) - graph, graph_builder = _make_graph( - [pixel_values, image_grid_thw], name="vision_encoder" - ) image_features = vision( - graph_builder.op, + builder.op, pixel_values=pixel_values, image_grid_thw=image_grid_thw, ) - image_features.name = "image_features" - graph.outputs.append(image_features) + builder.add_output(image_features, "image_features") return _make_model(graph) @@ -213,16 +208,13 @@ def _build_vision( height = ir.SymbolicDim("height") width = ir.SymbolicDim("width") - pixel_values = ir.Value( - name="pixel_values", - shape=ir.Shape([batch, 3, height, width]), - type=ir.TensorType(config.dtype), + graph, builder = _make_graph(name="vision_encoder") + pixel_values = builder.input( + "pixel_values", + dtype=config.dtype, + shape=[batch, 3, height, width], ) - - graph_inputs = [pixel_values] - - graph, graph_builder = _make_graph(graph_inputs, name="vision_encoder") - op = graph_builder.op + op = builder.op image_features = vision( op, @@ -234,8 +226,7 @@ def _build_vision( # vision encoder always processes one image at a time. image_features = op.Squeeze(image_features, [0]) - image_features.name = "image_features" - graph.outputs.append(image_features) + builder.add_output(image_features, "image_features") return _make_model(graph) @@ -294,60 +285,49 @@ def _build_decoder( cross_seq_len = ir.SymbolicDim("cross_sequence_len") cross_past_seq_len = ir.SymbolicDim("cross_past_seq_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"], ) - position_ids = ir.Value( - name="position_ids", - shape=ir.Shape([batch, seq_len]), - type=ir.TensorType(ir.DataType.INT64), + position_ids = builder.input( + "position_ids", + dtype=ir.DataType.INT64, + shape=[batch, seq_len], ) # Vision features: full on prefill, empty (0-length) on decode - cross_attention_states = ir.Value( - name="cross_attention_states", - shape=ir.Shape([batch, cross_seq_len, config.hidden_size]), - type=ir.TensorType(config.dtype), + cross_attention_states = builder.input( + "cross_attention_states", + dtype=config.dtype, + shape=[batch, cross_seq_len, config.hidden_size], ) - graph_inputs = [ - inputs_embeds, - attention_mask, - position_ids, - cross_attention_states, - ] - # Per-layer KV cache with separate dims for self-attention # (past_seq_len) and cross-attention (cross_past_seq_len) cross_attention_layers = set(config.cross_attention_layers or []) - flat_kv: list[ir.Value] = [] past_key_values: list[tuple[ir.Value, ir.Value]] = [] for i in range(config.num_hidden_layers): psl = cross_past_seq_len if i in cross_attention_layers else past_seq_len - past_key = ir.Value( - name=f"past_key_values.{i}.key", - shape=ir.Shape([batch, config.num_key_value_heads, psl, config.head_dim]), - type=ir.TensorType(config.dtype), + past_key = builder.input( + f"past_key_values.{i}.key", + dtype=config.dtype, + shape=[batch, config.num_key_value_heads, psl, config.head_dim], ) - past_value = ir.Value( - name=f"past_key_values.{i}.value", - shape=ir.Shape([batch, config.num_key_value_heads, psl, config.head_dim]), - type=ir.TensorType(config.dtype), + past_value = builder.input( + f"past_key_values.{i}.value", + dtype=config.dtype, + shape=[batch, config.num_key_value_heads, psl, config.head_dim], ) - flat_kv.extend([past_key, past_value]) past_key_values.append((past_key, past_value)) - graph_inputs.extend(flat_kv) - - graph, graph_builder = _make_graph(graph_inputs) - op = graph_builder.op + op = builder.op logits, present_key_values = decoder( op, @@ -358,8 +338,7 @@ def _build_decoder( past_key_values=past_key_values, ) - logits.name = "logits" - graph.outputs.append(logits) - _register_kv_cache_outputs(graph, present_key_values) + builder.add_output(logits, "logits") + _register_kv_cache_outputs(builder, present_key_values) return _make_model(graph)