Skip to content

Feature request: backend-dispatch registration hook for third-party packages #2427

Description

@cetagostini

Motivation

PyTensor's compile backends (numba, jax, mlx, pytensor) use singledispatch functions (numba_funcify, jax_funcify, etc.) defined inside pytensor.link.{backend}.dispatch.basic to find the implementation for each op type. Third-party packages that define custom ops must register their implementations in these singledispatch functions before any graph containing those ops is compiled.

PyTensor provides no plugin hook for this. There is no entry point, no register_backend_dispatch(), no __init_subclass__ on Op, and no import-time signal that a dispatch module has loaded.

Current workaround

pytensor_ml works around this with a sys.meta_path finder (_RegisterAfterImport) that intercepts imports of pytensor.link.{jax,mlx,numba,pytorch}.dispatch and loads the corresponding pytensor_ml.dispatch.{backend} registration module immediately after.

This mechanism:

  • Mutates process-global state (sys.meta_path), which can interact poorly with test frameworks that snapshot/restore sys.meta_path, import cleanup tools, and other libraries using the same technique.
  • Monkey-patches the module loader (spec.loader.exec_module) — a non-standard, CPython-specific hook.
  • Is invisible to PyTensor — PyTensor has no way to know that registrations happened or to verify they are complete.
  • Requires a hardcoded list of dispatch module paths that must be kept in sync with upstream renames or restructurings.

Evidence: what breaks

Without pytensor_ml.dispatch imported, numba_funcify has zero knowledge of pytensor_ml custom ops:

from pytensor.link.numba.dispatch.basic import numba_funcify
ml = [k for k in numba_funcify.registry if 'pytensor_ml' in getattr(k, '__module__', '')]
print(f'pytensor_ml registrations: {len(ml)}')  # → 0

After import pytensor_ml.dispatch, 4 ops are registered:

pytensor_ml registrations: 4
  PoolLayer, PoolLayerGrad, Im2Col, Col2Im

If the _RegisterAfterImport finder is removed from sys.meta_path between installation and the dispatch module loading (e.g. by a test framework restoring sys.meta_path), registrations for that backend are silently lost — no error, no warning, just a missing singledispatch implementation that surfaces as a confusing TypeError at compile time.

A committed reproducer lives at pymc-labs/pytensor.cpp#tests/test_dispatch_contract.py.

Proposed API

A single function in a new pytensor.registration module (or in pytensor.compile.mode) that third-party packages call to register their funcify implementations before or after the dispatch module loads:

# pytensor/registration.py  (new module)

from typing import Callable, Type

def register_funcify(backend: str, op_type: Type, func: Callable) -> None:
    """Register a funcify implementation for *op_type* on *backend*.

    If the backend's singledispatch function has not been created yet,
    the registration is deferred and applied when the dispatch module
    loads.  If it already exists, the registration is applied immediately.

    Parameters
    ----------
    backend : str
        One of "numba", "jax", "mlx", "pytorch".
    op_type : Type
        The Op subclass to register.
    func : Callable
        The implementation function.
    """

Semantics:

  • Can be called at any time — before or after the dispatch module loads.
  • If called before: stores the registration and applies it when the singledispatch function is first accessed.
  • If called after: applies immediately (same as dispatch.register).
  • Thread-safe.

Alternative: entry-point based

# pyproject.toml of third-party package
[project.entry-points."pytensor.dispatch"]
numba = "my_package.dispatch.numba"

PyTensor discovers entry points at startup and imports the registration modules after loading each backend's dispatch module.

Why the explicit API is preferred: Entry points add startup cost (scanning all installed packages) and require a specific pyproject.toml layout. The explicit register_funcify() call is lazy, explicit, and works with any import style.

Migration path

For pytensor_ml

  1. Replace the _RegisterAfterImport meta_path finder with calls to pytensor.registration.register_funcify("numba", PoolLayer, ...).
  2. Remove pytensor_ml.dispatch.__init__.py's sys.meta_path manipulation.
  3. Keep pytensor_ml.dispatch.{backend} modules as-is (they contain the actual implementation functions); only the trigger changes.

For downstream consumers (e.g. pytensor_cpp)

No change needed — the dependency is indirect (the singledispatch must be populated before compilation). Once the hook lands, the workaround comment becomes historical.

Reproducer

See tests/test_dispatch_contract.py in pymc-labs/pytensor.cpp. The test:

  1. Verifies that numba_funcify has exactly 4 pytensor_ml registrations after the shim runs.
  2. Verifies that the _RegisterAfterImport finder is installed on sys.meta_path.
  3. Demonstrates that removing the finder does not affect already-loaded backends but would silently lose registrations for backends not yet loaded.
  4. Compiles a graph containing a PoolLayer op with the numba backend, proving the end-to-end contract works today.

A standalone reproducer script is at scripts/repro_dispatch_ordering.py.

Activity

  1. ricardoV94 commented on Sep 19, 2026

    @ricardoV94
    Member

    Could you explain the problem without writing an essay? Is it that stuff like jax_funcify eagerly imports jax, so you can't import it without attempting to import jax eagerly? Or that you're trying to override a PyTensor dispatch?

    We can make the single dispatch elsewhere. The linkers are already importable without forcing the backend dependencies

  2. cetagostini commented on Sep 19, 2026

    @cetagostini
    ContributorAuthor

    Registering a funcify for our own Ops. The problem is that numba_funcify / jax_funcify are the registry, and they live in pytensor.link.<backend>.dispatch.basic, which imports the backend at module level (import jax, import numba, import mlx.core). PyTensor only imports that module lazily, at first compile, so the registry doesn't exist until then, and the only way to touch it earlier is to import the dispatch module yourself, which forces the backend dep. There's also no public register entry point; the two register_* helpers are numba-internal.

    So a third-party Op author has no supported way to say "here's my funcify" without eagerly importing a backend they may never use, or guessing your import order. If the registry lived somewhere that doesn't touch a backend, the linkers already import fine without the deps, we'd register at import time, there'd be nothing to defer, and no hook would be needed. Is that the shape you'd want?

  3. ricardoV94 commented on Sep 19, 2026

    @ricardoV94
    Member

    Yes, that and less llm text in my email notifications. Could have been one sentence

  4. cetagostini commented on Sep 19, 2026

    @cetagostini
    ContributorAuthor

    @ricardoV94 would be fine if I open a pr to solve as mentioned above?

  5. ricardoV94 commented on Sep 19, 2026

    @ricardoV94
    Member

    @ricardoV94 would be fine if I open a pr to solve as mentioned above?

    yes ofc

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions