.. DO NOT EDIT. .. THIS FILE WAS AUTOMATICALLY GENERATED BY SPHINX-GALLERY. .. TO MAKE CHANGES, EDIT THE SOURCE PYTHON FILE: .. "auto_examples_compute/plot_graph_builder_benchmark.py" .. LINE NUMBERS ARE GIVEN BELOW. .. only:: html .. note:: :class: sphx-glr-download-link-note :ref:`Go to the end ` to download the full example code. .. rst-class:: sphx-glr-example-title .. _sphx_glr_auto_examples_compute_plot_graph_builder_benchmark.py: Benchmark GraphBuilder against onnxscript GraphBuilder ====================================================== This example builds the same attention-style graph with 100, 200, and 500 nodes using :class:`onnx_light.onnx_core.graph_builder.GraphBuilder` and :class:`onnxscript.GraphBuilder`. It checks that both models produce the same outputs before measuring model construction with and without final protobuf serialization. The timings include model finalization (``to_onnx`` or ``onnx_ir.to_proto``), but exclude imports, input generation, and execution. The example also reports the operator-type distribution and plots both timing modes against the graph size. Run this example with the ``docs`` optional dependencies installed. Each 20-node block splits a symbolic width into two heads using runtime ``Shape``, ``Gather``, ``Div``, and ``Concat`` operations. It reshapes and transposes the input, computes scaled dot-product self-attention, then merges the heads and applies an averaged residual connection and ``Relu``. Batch size, sequence length, and width are all dynamic; width must be divisible by two. The same exported models are checked on several shapes, including a singleton sequence. Both builders receive identical operators and initializers; their default shape-inference behavior is included in the construction timings. .. GENERATED FROM PYTHON SOURCE LINES 25-253 .. image-sg:: /auto_examples_compute/images/sphx_glr_plot_graph_builder_benchmark_001.png :alt: Model construction, Model construction and serialization :srcset: /auto_examples_compute/images/sphx_glr_plot_graph_builder_benchmark_001.png :class: sphx-glr-single-img .. rst-class:: sphx-glr-script-out .. code-block:: none nodes builder model (ms) model + serialization (ms) 100 onnx-light 29.50 29.76 100 onnxscript 26.62 26.28 200 onnx-light 53.64 54.83 200 onnxscript 51.14 51.17 500 onnx-light 124.54 125.92 500 onnxscript 126.34 126.38 Node-type distribution (identical for both builders): operator 100 200 500 ------------------------------------ Add 5 10 25 Cast 5 10 25 Concat 5 10 25 Div 10 20 50 Gather 15 30 75 MatMul 10 20 50 Mul 5 10 25 Relu 5 10 25 Reshape 10 20 50 Shape 5 10 25 Softmax 5 10 25 Sqrt 5 10 25 Transpose 15 30 75 ------------------------------------ Total 100 200 500 | .. code-block:: Python from __future__ import annotations from collections import Counter import gc import statistics import time import matplotlib.pyplot import numpy import onnx import onnx_ir import onnxscript from onnx.reference import ReferenceEvaluator from onnx_light.onnx import TensorProto from onnx_light.onnx_core.graph_builder import GraphBuilder NODE_COUNTS = (100, 200, 500) OPSET = 18 SHAPE = ["batch", "sequence", "width"] INPUT_SHAPES = ((1, 1, 4), (2, 3, 6), (3, 5, 8)) BLOCK_SIZE = 20 def attention_blocks(op, value, constants, node_count: int): """Builds attention blocks with runtime-derived head and output shapes.""" if node_count <= 0 or node_count % BLOCK_SIZE: raise ValueError(f"node_count must be a positive multiple of {BLOCK_SIZE}.") batch_index, sequence_index, width_index, heads, half = constants for _ in range(node_count // BLOCK_SIZE): shape = op.Shape(value) batch = op.Gather(shape, batch_index, axis=0) sequence = op.Gather(shape, sequence_index, axis=0) width = op.Gather(shape, width_index, axis=0) head_width = op.Div(width, heads) head_shape = op.Concat(batch, sequence, heads, head_width, axis=0) query = op.Reshape(value, head_shape) query = op.Transpose(query, perm=[0, 2, 1, 3]) key = op.Transpose(query, perm=[0, 1, 3, 2]) scores = op.MatMul(query, key) scale = op.Cast(head_width, to=TensorProto.FLOAT) scale = op.Sqrt(scale) scores = op.Div(scores, scale) weights = op.Softmax(scores, axis=-1) context = op.MatMul(weights, query) context = op.Transpose(context, perm=[0, 2, 1, 3]) context = op.Reshape(context, shape) value = op.Add(value, context) value = op.Mul(value, half) value = op.Relu(value) return value def constants(): """Returns the shared shape indices, head count, and residual scale.""" return ( numpy.array([0], dtype=numpy.int64), numpy.array([1], dtype=numpy.int64), numpy.array([2], dtype=numpy.int64), numpy.array([2], dtype=numpy.int64), numpy.array(0.5, dtype=numpy.float32), ) def build_light(node_count: int): """Builds dynamic attention blocks with onnx-light and returns its model.""" builder = GraphBuilder("attention") builder.set_opset_version("", OPSET) value = builder.inp("X", TensorProto.FLOAT, SHAPE) initializers = [ builder.init(array, name=f"c{index}") for index, array in enumerate(constants()) ] value = attention_blocks(builder.op, value, initializers, node_count) builder.out(value, TensorProto.FLOAT, SHAPE) return builder.to_onnx("model") def build_onnxscript(node_count: int): """Builds the same dynamic attention blocks with onnxscript.""" graph = onnx_ir.Graph( inputs=[], outputs=[], nodes=[], opset_imports={"": OPSET}, name="attention" ) builder = onnxscript.GraphBuilder(graph) value = builder.input("X", dtype=onnx_ir.DataType.FLOAT, shape=SHAPE) initializers = [ builder.initializer(onnx_ir.tensor(array), name=f"c{index}") for index, array in enumerate(constants()) ] value = attention_blocks(builder.op, value, initializers, node_count) value.shape = onnx_ir.Shape(SHAPE) value.type = onnx_ir.TensorType(onnx_ir.DataType.FLOAT) builder.add_output(value, None) return onnx_ir.to_proto(onnx_ir.Model(graph, ir_version=10)) def check_models(node_count: int) -> None: """Checks both models and compares outputs across dynamic input shapes.""" light = onnx.load_from_string(build_light(node_count).SerializeToString()) scripted = build_onnxscript(node_count) # ONNX's checker and evaluator expect "" rather than "ai.onnx". for opset in light.opset_import: if opset.domain == "ai.onnx": opset.domain = "" for model in (light, scripted): onnx.checker.check_model(model) assert len(model.graph.node) == node_count light_session = ReferenceEvaluator(light) scripted_session = ReferenceEvaluator(scripted) rng = numpy.random.default_rng(0) for shape in INPUT_SHAPES: feeds = {"X": rng.standard_normal(shape).astype(numpy.float32)} light_outputs = light_session.run(None, feeds) scripted_outputs = scripted_session.run(None, feeds) assert len(light_outputs) == len(scripted_outputs) == 1 assert light_outputs[0].shape == scripted_outputs[0].shape == shape numpy.testing.assert_allclose(light_outputs[0], scripted_outputs[0], rtol=0, atol=0) def measure(build, node_count: int, serialize: bool, repeats: int = 3) -> float: """Returns the median construction time in milliseconds.""" samples = [] gc_was_enabled = gc.isenabled() gc.disable() try: for _ in range(repeats): gc.collect() start = time.perf_counter() model = build(node_count) if serialize: model.SerializeToString() elapsed = (time.perf_counter() - start) * 1000 samples.append(elapsed) del model finally: gc.collect() if gc_was_enabled: gc.enable() return statistics.median(samples) def node_type_distribution(model) -> Counter: """Counts nodes by operator type.""" return Counter(node.op_type for node in model.graph.node) def format_node_type_table(distributions: dict[int, Counter]) -> str: """Formats node-type counts for every benchmark graph size.""" node_counts = sorted(distributions) operator_types = sorted( { operator_type for distribution in distributions.values() for operator_type in distribution } ) header = f"{'operator':<12}" + "".join(f"{node_count:>8}" for node_count in node_counts) separator = "-" * len(header) rows = [header, separator] for operator_type in operator_types: rows.append( f"{operator_type:<12}" + "".join( f"{distributions[node_count].get(operator_type, 0):>8}" for node_count in node_counts ) ) rows.extend( [ separator, f"{'Total':<12}" + "".join( f"{sum(distributions[node_count].values()):>8}" for node_count in node_counts ), ] ) return "\n".join(rows) def plot_benchmark(results: list[dict]): """Plots construction times with and without serialization.""" figure, axes = matplotlib.pyplot.subplots(1, 2, figsize=(11, 4), sharex=True) for axis, key, title in ( (axes[0], "model", "Model construction"), (axes[1], "serialized", "Model construction and serialization"), ): for builder_name in ("onnx-light", "onnxscript"): rows = [row for row in results if row["builder"] == builder_name] axis.plot( [row["nodes"] for row in rows], [row[key] for row in rows], "o-", label=builder_name, ) axis.set_title(title) axis.set_xlabel("number of nodes") axis.set_ylabel("median time (ms)") axis.grid(True, alpha=0.3) axis.legend() figure.tight_layout() return figure if __name__ == "__main__": print("nodes builder model (ms) model + serialization (ms)") results = [] distributions = {} for count in NODE_COUNTS: check_models(count) distribution_model = build_light(count) distributions[count] = node_type_distribution(distribution_model) del distribution_model for name, build in (("onnx-light", build_light), ("onnxscript", build_onnxscript)): model_time = measure(build, count, False) serialized_time = measure(build, count, True) results.append( { "nodes": count, "builder": name, "model": model_time, "serialized": serialized_time, } ) print(f"{count:5} {name:12} {model_time:10.2f} {serialized_time:26.2f}") print("\nNode-type distribution (identical for both builders):") print(format_node_type_table(distributions)) plot_benchmark(results) matplotlib.pyplot.show() .. rst-class:: sphx-glr-timing **Total running time of the script:** (0 minutes 8.959 seconds) .. _sphx_glr_download_auto_examples_compute_plot_graph_builder_benchmark.py: .. only:: html .. container:: sphx-glr-footer sphx-glr-footer-example .. container:: sphx-glr-download sphx-glr-download-jupyter :download:`Download Jupyter notebook: plot_graph_builder_benchmark.ipynb ` .. container:: sphx-glr-download sphx-glr-download-python :download:`Download Python source code: plot_graph_builder_benchmark.py ` .. container:: sphx-glr-download sphx-glr-download-zip :download:`Download zipped: plot_graph_builder_benchmark.zip ` .. include:: plot_graph_builder_benchmark.recommendations .. only:: html .. rst-class:: sphx-glr-signature `Gallery generated by Sphinx-Gallery `_