Skip to content
38 changes: 37 additions & 1 deletion frontend/catalyst/from_plxpr/device_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,9 @@
verify_no_state_variance_returns,
verify_operations,
)
from catalyst.passes.builtin_passes import (
graph_decomposition_setup_inputs,
)
from catalyst.utils.exceptions import CompileError

_named_obs_dict = {
Expand All @@ -57,7 +60,11 @@


def create_device_preprocessing_pipeline(
device: qp.devices.Device, execution_config: ExecutionConfig, shots: int, warn: bool = True
device: qp.devices.Device,
execution_config: ExecutionConfig,
shots: int,
warn: bool = True,
needs_gateset_preprocessing: bool = True,
) -> list[BoundTransform]:
"""Create a pipeline of device preprocessing transforms for lowering QNodes."""
shots_present = qp.math.is_abstract(shots) or shots != 0
Expand Down Expand Up @@ -100,6 +107,11 @@ def create_device_preprocessing_pipeline(
pipeline, unsupported_transforms, device, execution_config, shots, capabilities
)

if needs_gateset_preprocessing:
_gateset_preprocessing(
pipeline, unsupported_transforms, device, execution_config, shots, capabilities
)

if unsupported_transforms and warn:
warnings.warn(
"The following device-preprocessing transforms are currently not supported with "
Expand Down Expand Up @@ -281,6 +293,30 @@ def _gradient_preprocessing(
)


# pylint: disable=unused-argument
def _gateset_preprocessing(
pipeline: list[BoundTransform],
unsupported_transforms: list[str],
device: qp.devices.Device,
execution_config: ExecutionConfig,
shots: int,
capabilities: DeviceCapabilities,
) -> None:
"""Insert a `device-based-decomposition` pass targetting the gateset
specified by the specific `device`
"""
gate_set = capabilities.gate_set()

# Get the default args/kwargs with the above gate_set
# `device-based-decomposition` uses the same arguments as `graph-decomposition`
targs, tkwargs = graph_decomposition_setup_inputs(gate_set=gate_set)
t = qp.transform(pass_name="device-based-decomposition")

pipeline.append(
_safe_create_bound_transform(t, unsupported_transforms, args=targs, kwargs=tkwargs)
)


def _safe_create_bound_transform(
transform: Transform, unsupported_transforms: list[str], warn=True, args=(), kwargs=None
) -> BoundTransform:
Expand Down
22 changes: 18 additions & 4 deletions frontend/catalyst/from_plxpr/from_plxpr.py
Original file line number Diff line number Diff line change
Expand Up @@ -245,6 +245,7 @@ def __copy__(self):
new_version.init_qreg = self.init_qreg
new_version.requires_decompose_lowering = self.requires_decompose_lowering
new_version.decompose_tkwargs = copy(self.decompose_tkwargs)
new_version.needs_gateset_preprocessing = self.needs_gateset_preprocessing
return new_version

def __init__(self, skip_preprocess=False, _preprocess_warn=True, collect_decomp_rules=True):
Expand All @@ -253,6 +254,7 @@ def __init__(self, skip_preprocess=False, _preprocess_warn=True, collect_decomp_
self._skip_preprocess = skip_preprocess
self._preprocess_warn = _preprocess_warn
self._collect_decomp_rules = collect_decomp_rules
self.needs_gateset_preprocessing = False

# Compiler options for the new decomposition system
self.requires_decompose_lowering = False
Expand Down Expand Up @@ -328,7 +330,11 @@ def calling_convention(*args):
pipelines = (("main", tuple(self._pass_pipeline) + device_pass_pipeline(qnode.device)),)
if not self._skip_preprocess:
device_preprocessing_pipeline = create_device_preprocessing_pipeline(
qnode.device, execution_config, shots, warn=self._preprocess_warn
qnode.device,
execution_config,
shots,
warn=self._preprocess_warn,
needs_gateset_preprocessing=self.needs_gateset_preprocessing,
)
pipelines += (("device", device_preprocessing_pipeline),)

Expand Down Expand Up @@ -406,6 +412,7 @@ def handle_transform(
):
"""Handle the conversion from plxpr to Catalyst jaxpr for a
PL transform."""

consts = args[_tuple_to_slice(consts_slice)]
non_const_args = args[_tuple_to_slice(args_slice)]
targs = args[_tuple_to_slice(targs_slice)]
Expand All @@ -424,9 +431,16 @@ def handle_transform(

# Apply the corresponding Catalyst pass counterpart
next_eval = copy(self)
t = qp.transform(pass_name=transform.pass_name)
bound_pass = qp.transforms.core.BoundTransform(t, args=targs, kwargs=pl_tkwargs)
next_eval._pass_pipeline.insert(0, bound_pass)

if transform.pass_name == "device-based-decomposition":
# device-based-decomposition is not applied here, but is delayed to the device preprocessing pipeline
# notify that this needs to be done rather than applying the pass here.
next_eval.needs_gateset_preprocessing = True
else:
t = qp.transform(pass_name=transform.pass_name)
bound_pass = qp.transforms.core.BoundTransform(t, args=targs, kwargs=pl_tkwargs)
next_eval._pass_pipeline.insert(0, bound_pass)

return next_eval.eval(inner_jaxpr, consts, *non_const_args)


Expand Down
17 changes: 17 additions & 0 deletions frontend/catalyst/passes/builtin_passes.py
Original file line number Diff line number Diff line change
Expand Up @@ -1855,6 +1855,22 @@ def rule_ref_name(rule):
pass_name="graph-decomposition", setup_inputs=graph_decomposition_setup_inputs
)


def device_based_decomposition_setup_inputs():
R"""
Specify that the ``-device-based-decomposition`` MLIR compiler pass for applying the graph-based
decomposition should be applied to the decorated QNode during :func:`~.qjit` compilation, using
the gatseset automatically detected from the backend toml file.

Runs `adjoint-lowering` -> `ctrl-lowering` -> `graph-decomposition` with derived gateset
"""
return (), {}


device_based_decomposition = qp.transform(
pass_name="device-based-decomposition", setup_inputs=device_based_decomposition_setup_inputs
)

__all__ = [
"cancel_inverses",
"combine_global_phases",
Expand All @@ -1875,4 +1891,5 @@ def rule_ref_name(rule):
"decompose_arbitrary_ppr",
"graph_decomposition",
"diagonalize_measurements",
"device_based_decomposition",
]
Original file line number Diff line number Diff line change
@@ -0,0 +1,44 @@
// Copyright 2026 Xanadu Quantum Technologies Inc.

// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at

// http://www.apache.org/licenses/LICENSE-2.0

// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.

// RUN: catalyst --tool=opt --pass-pipeline='builtin.module(device-based-decomposition{gate-set=C(Adjoint(H))=1.0,C(H)=1.0 alt-decomps=C(Adjoint(U)){}{wires:1}{}=ctrl_adj_u,C(U){}{wires:1}{}=ctrl_u})' %s | FileCheck %s

// CHECK-LABEL: func.func @controlled_adjoint_region(
// CHECK-SAME: %[[C:.*]]: !quantum.bit, %[[Q:.*]]: !quantum.bit
func.func @controlled_adjoint_region(%c: !quantum.bit, %q: !quantum.bit) -> (!quantum.bit, !quantum.bit) {
%true = arith.constant true
%a:2 = quantum.adjoint(%c, %q) : !quantum.bit, !quantum.bit {
^bb0(%ac: !quantum.bit, %aq: !quantum.bit):
%oc, %or = quantum.ctrl(%ac) ctrlvals(%true) (%aq) : !quantum.bit -> !quantum.bit {
^bb1(%iq: !quantum.bit):
%u = quantum.custom "U"() %iq : !quantum.bit
quantum.yield %u : !quantum.bit
}
quantum.yield %oc, %or : !quantum.bit, !quantum.bit
}
// CHECK: %[[A:.*]], %[[AC:.*]] = quantum.custom "H"() %[[Q]] adj ctrls(%[[C]]) ctrlvals(%{{.*}}) : !quantum.bit ctrls !quantum.bit
// CHECK: %[[B:.*]], %[[BC:.*]] = quantum.custom "H"() %[[A]] adj ctrls(%[[AC]]) ctrlvals(%{{.*}}) : !quantum.bit ctrls !quantum.bit
// CHECK: return %[[BC]], %[[B]]
return %a#0, %a#1 : !quantum.bit, !quantum.bit
}

// C(Adjoint(U)) -> two C(Adjoint(H)).
func.func private @ctrl_adj_u(%q: !quantum.bit, %ctrl: !quantum.bit) -> (!quantum.bit, !quantum.bit) attributes {
target_gate = "C(Adjoint(U)){}{wires:1}{}",
resources = {operations = {"C(Adjoint(H)){}{wires:1}{}" = 2 : i64}} } {
%true = arith.constant true
%a, %ac = quantum.custom "H"() %q adj ctrls(%ctrl) ctrlvals(%true) : !quantum.bit ctrls !quantum.bit
%b, %bc = quantum.custom "H"() %a adj ctrls(%ac) ctrlvals(%true) : !quantum.bit ctrls !quantum.bit
return %b, %bc : !quantum.bit, !quantum.bit
}
72 changes: 72 additions & 0 deletions frontend/test/pytest/from_plxpr/test_preprocessing.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@
from catalyst.device.decomposition import measurements_from_counts, measurements_from_samples
from catalyst.from_plxpr import from_plxpr
from catalyst.jax_primitives import quantum_kernel_p
from catalyst.passes.builtin_passes import device_based_decomposition
from catalyst.utils.exceptions import CompileError

pytestmark = pytest.mark.usefixtures("use_capture")
Expand Down Expand Up @@ -585,6 +586,77 @@ def f():
)


class TestGatesetPreprocessing:
"""Tests for preprocessing related to invoking `device-based-decomposition` pass with target gateset
as described by dveice TOML file."""

@pytest.mark.parametrize("apply_device_based_decomposition", [True, False])
def test_device_based_decomposition_added_to_pipeline(self, apply_device_based_decomposition):
"""Tests that the device-based-decomposition pass is added to the pipeline
when the @device_based_decomposition decorator is applied."""

dev = qp.device("null.qubit", wires=4)

if apply_device_based_decomposition:

@device_based_decomposition
@qp.qnode(dev)
def f():
return qp.expval(qp.Z(0))

else:

@qp.qnode(dev)
def f():
return qp.expval(qp.Z(0))

device_pipelines = get_pipelines(f, skip_preprocess=False)[1][1]
pass_exists = any(t.pass_name == "device-based-decomposition" for t in device_pipelines)

assert pass_exists == apply_device_based_decomposition

def test_gateset_matches_device_capabilities(self):
"""Test that device operations and their C/Adjoint expansions are in the
graph-decomposition gate_set."""
dev = CapabilitiesDevice(wires=4)
dev.capabilities = DeviceCapabilities(
operations={
"PauliX": OperatorProperties(invertible=False, controllable=False),
"PauliY": OperatorProperties(invertible=False, controllable=True),
"PauliZ": OperatorProperties(invertible=True, controllable=False),
"Hadamard": OperatorProperties(invertible=True, controllable=True),
},
measurement_processes={"ExpectationMP": [], "SampleMP": [], "CountsMP": []},
)

@device_based_decomposition
@qp.qnode(dev, shots=1)
def f():
qp.expval(qp.Z(0))

device_pipelines = get_pipelines(f, skip_preprocess=False)[1][1]
gate_set = next(
t.kwargs["gate_set"]
for t in device_pipelines
if t.pass_name == "device-based-decomposition"
)

assert (
"PauliX" in gate_set
and "Adjoint(PauliX)" not in gate_set
and "C(PauliX)" not in gate_set
)
assert (
"PauliY" in gate_set and "Adjoint(PauliY)" not in gate_set and "C(PauliY)" in gate_set
)
assert (
"PauliZ" in gate_set and "Adjoint(PauliZ)" in gate_set and "C(PauliZ)" not in gate_set
)
assert (
"Hadamard" in gate_set and "Adjoint(Hadamard)" in gate_set and "C(Hadamard)" in gate_set
)


class TestIntegration:
"""Integration tests for device preprocessing with program capture."""

Expand Down
42 changes: 42 additions & 0 deletions mlir/include/Quantum/Transforms/Passes.td
Original file line number Diff line number Diff line change
Expand Up @@ -150,6 +150,48 @@ def MergeRotationsPass : Pass<"merge-rotations"> {
let dependentDialects = ["math::MathDialect"];
}

def DeviceBasedDecompositionPass: Pass<"device-based-decomposition"> {
let summary = "Invoke the graph decomposition pass with the gatseset automatically detected from the backend toml file.";

let options = [
ListOption<
/*C++ name*/"targetGateSetOption",
/*CLI name*/"gate-set",
/*Type*/"std::string",
/*Description*/"The accepted gates to decompose to.">,
ListOption<
/*C++ name*/"fixedDecompsOption",
/*CLI name*/"fixed-decomps",
/*Type*/"std::string",
/*Description*/"Maps an operator to a decomposition rule that will be applied.">,
ListOption<
/*C++ name*/"altDecompsOption",
/*CLI name*/"alt-decomps",
/*Type*/"std::string",
/*Description*/
"Maps an operator to a list of alternative decomposition rules that will be considered"
" alongside any built-in rules for the operator.">,
Option<
/*C++ name*/"bytecodeRulesFile",
/*CLI name*/"bytecode-rules",
/*Type*/"std::string",
/*Default*/[{ "" }],
/*Description*/"A path to a bytecode file of compiled decomposition rules.">,
Option<
/*C++ name*/"libQPDPath",
/*CLI name*/"libQPD-path",
/*Type*/"std::string",
/*Default*/[{ "" }],
/*Description*/"A path to the QuantumPythonDecompositions dynamic library.">,
Option<
/*C++ name*/"libpythonPath",
/*CLI name*/"libpython-path",
/*Type*/"std::string",
/*Default*/[{ "" }],
/*Description*/"A path to the python shared library.">
];
}

def GraphDecompositionPass : Pass<"graph-decomposition", "mlir::ModuleOp"> {
let summary = "Decompose gates using an MLIR-native graph-based framework.";

Expand Down
1 change: 1 addition & 0 deletions mlir/lib/Quantum/Transforms/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@ file(GLOB SRC
CtrlLowering/CtrlLowering.cpp
RemoveGlobalPhasesPass.cpp
ConversionPatterns.cpp
DeviceBasedDecomposition.cpp
GraphDecomposition/DecompUtils.cpp
GraphDecomposition/DecomposeLoweringPatterns.cpp
GraphDecomposition/decompose_lowering.cpp
Expand Down
Loading
Loading