Note
Go to the end to download the full example code.
Benchmark GraphBuilder against onnxscript GraphBuilder#
This example builds the same attention-style graph with 100, 200, and
500 nodes using onnx_light.onnx_core.graph_builder.GraphBuilder and
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.

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
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()
Total running time of the script: (0 minutes 8.959 seconds)
Related examples
Gallery generated by Sphinx-Gallery
Example last updated
- Date:
2026-10-05