How to use a custom kernel#

ReferenceEvaluator dispatches every NodeProto against the static C++ onnx_light::core::runtime::KernelDispatchTable(). An operator that is not built in — typically an operator from a user-defined domain, an experimental op, or a stand-in for one that is not yet implemented — would otherwise fail with unsupported op_type.

This page shows how to plug a custom kernel into the runtime so such a graph runs, in Python and in C++. The hook is exposed at three layers; pick the one that matches your use case. The per-session layers share the single C++ entry point onnx_light::core::runtime::RuntimeContext::RegisterCustomKernel(); a kernel can also be registered globally (see Register globally or per session).

Register globally or per session#

The examples above register a kernel on a single ReferenceEvaluator (equivalently, on one RuntimeContext) — the kernel is only visible to that object. onnx-light also supports global (process-wide) registration: a global kernel is picked up by every RuntimeContext created afterwards, so you install it once instead of on every evaluator.

Both scopes are supported, and a per-session registration always overrides a global one for the same (domain, op_type). Resolution precedence, from highest to lowest, is: model-local functions, the built-in control-flow operators (If / Loop / Scan / SequenceMap), per-session custom kernels, global custom kernels, then the built-in onnx_light::core::runtime::KernelDispatchTable().

Because an evaluator caches its runtime sessions on first run(), register a global kernel before running the evaluators that should use it.

from onnx_light.onnx.reference import ReferenceEvaluator

def square(node, x):
    return x * x

# Registered once; visible to every evaluator created afterwards.
ReferenceEvaluator.register_custom_kernel_global("my.domain", "Square", square)

sess = ReferenceEvaluator(model)  # no per-session registration needed
(y,) = sess.run(None, {"x": np.array([1.0, 2.0, 3.0], dtype=np.float32)})

# Remove the global registration when done.
ReferenceEvaluator.unregister_custom_kernel_global("my.domain", "Square")

The low-level counterparts live on the runtime submodule: runtime.register_custom_kernel(domain, op_type, fn), runtime.unregister_custom_kernel(domain, op_type) and runtime.clear_custom_kernels() (module-level, i.e. global), as opposed to the identically named methods on RuntimeContext (per session).

#include "onnx_core/runtime/kernels/kernel_dispatch_table.h"

using namespace onnx_light::core::runtime;

// Global: picked up by every RuntimeContext.
RegisterGlobalCustomKernel(
    "my.domain", "Scale",
    [](const NodeProto &node, RuntimeContext &c) {
      const Tensor &x = c.Get(node.input(0));
      // ...
      c.Put(node.output(0), /* Tensor */ ...);
    });

// Per session: only this context (overrides the global one above).
RuntimeContext ctx(KernelContext(/*opset=*/18));
ctx.RegisterCustomKernel("my.domain", "Scale", /* ... */);

UnregisterGlobalCustomKernel("my.domain", "Scale");  // remove global

Override a built-in kernel#

The empty default ONNX domain is normalised to "ai.onnx", so registering a kernel under the default domain takes precedence over the entry that onnx_light::core::runtime::KernelDispatchTable() would otherwise dispatch. This is convenient to instrument or replace a specific kernel without patching the C++ runtime.

sess = ReferenceEvaluator(
    parser.parse_model(
        '<ir_version: 10, opset_import: ["" : 18]>'
        "agraph (float[3] x) => (float[3] y) { y = Abs(x) }"
    )
)
sess.register_custom_kernel("", "Abs", lambda node, x: -x)
(y,) = sess.run(None, {"x": np.array([-1.0, -2.0, -3.0], dtype=np.float32)})
# Abs replaced by negation: y == [1., 2., 3.]

Unregister a kernel and restore the original#

unregister_custom_kernel() removes a previously registered custom kernel. Because custom kernels are consulted before the built-in onnx_light::core::runtime::KernelDispatchTable(), unregistering one that overrode a built-in operator restores the original built-in kernel on the next run(). It returns True when a custom kernel was removed and False otherwise; the empty domain is normalised to "ai.onnx" just like when registering.

sess.register_custom_kernel("", "Abs", lambda node, x: -x)
# ... use the negated override ...
sess.unregister_custom_kernel("", "Abs")  # restores the built-in Abs
(y,) = sess.run(None, {"x": np.array([-1.0, -2.0, -3.0], dtype=np.float32)})
# y == [1., 2., 3.] (built-in Abs again)

At the C++ / low-level binding layer the same is achieved with onnx_light::core::runtime::RuntimeContext::UnregisterCustomKernel(), which erases the custom entry so RunNode() falls back to the built-in kernel.

using namespace onnx_light::core::runtime;

RuntimeContext rt(KernelContext(/*opset=*/18));
rt.Set("x", Tensor::FromFloat("x", {3}, {-1.0f, -2.0f, -3.0f}));

// The empty domain is normalised to "ai.onnx", so this overrides
// the built-in Abs entry with a negation.
rt.RegisterCustomKernel(
    "", "Abs", [](const NodeProto &node, RuntimeContext &ctx) {
      const Tensor &in = ctx.Get(node.input(0));
      std::vector<float> out(static_cast<size_t>(in.element_count()));
      const float *src = in.AsFloat();
      for (size_t i = 0; i < out.size(); ++i) {
        out[i] = -src[i];
      }
      ctx.Put(node.output(0),
              Tensor::FromFloat(node.output(0), in.shape, out));
    });
// Abs replaced by negation: y == [1., 2., 3.]

Use the low-level context binding#

For kernels that need direct access to the runtime context — for example to read sequences or to participate in the event log — use the low-level RuntimeContext binding. The callback receives the raw NodeProto and RuntimeContext and is responsible for any tensor encoding/decoding. This Python binding mirrors the C++ onnx_light::core::runtime::RuntimeContext one-to-one.

from onnx_light.onnx_py._onnxpykernels import runtime as rt

ctx = rt.RuntimeContext(rt.KernelContext(rt.default_opset(18)))
ctx.set("x", ...)

def scale(node, c):
    x = c.get(str(node.input[0]))
    ...
    c.put(str(node.output[0]), ..., "output")

ctx.register_custom_kernel("my.domain", "Scale", scale)
rt.register_model_functions(model, ctx)
plan = ctx.get_execution_plan(model.graph)
rt.RuntimeSession(plan).run(ctx)
#include "onnx_core/runtime/kernels/run_nodes.h"
#include "onnx_core/runtime/runtime_context.h"

using namespace onnx_light::core::runtime;

RuntimeContext ctx(KernelContext(/*opset=*/18));
ctx.Set("x", /* Tensor */ ...);

ctx.RegisterCustomKernel(
    "my.domain", "Scale",
    [](const NodeProto &node, RuntimeContext &c) {
      const Tensor &x = c.Get(node.input(0));
      // ...
      c.Put(node.output(0), /* Tensor */ ...);
    });
RegisterModelFunctions(model, ctx);
const auto &plan = ctx.GetExecutionPlan(model.graph());
RuntimeSession(plan).Run(ctx);

The low-level binding and the C++ tab above are two faces of the same onnx_light::core::runtime::RuntimeContext::RegisterCustomKernel() entry point, which is also how C++ extension modules ship additional kernels without rebuilding lib_onnx_kernels.

See also#