diff --git a/frontend/catalyst/from_plxpr/device_utils.py b/frontend/catalyst/from_plxpr/device_utils.py index 739734a4e4..c0b89044dc 100644 --- a/frontend/catalyst/from_plxpr/device_utils.py +++ b/frontend/catalyst/from_plxpr/device_utils.py @@ -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 = { @@ -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 @@ -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 " @@ -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: diff --git a/frontend/catalyst/from_plxpr/from_plxpr.py b/frontend/catalyst/from_plxpr/from_plxpr.py index 00155ea2c6..3736e0eb02 100644 --- a/frontend/catalyst/from_plxpr/from_plxpr.py +++ b/frontend/catalyst/from_plxpr/from_plxpr.py @@ -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): @@ -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 @@ -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),) @@ -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)] @@ -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) diff --git a/frontend/catalyst/passes/builtin_passes.py b/frontend/catalyst/passes/builtin_passes.py index a43d754140..68f42bae75 100644 --- a/frontend/catalyst/passes/builtin_passes.py +++ b/frontend/catalyst/passes/builtin_passes.py @@ -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", @@ -1875,4 +1891,5 @@ def rule_ref_name(rule): "decompose_arbitrary_ppr", "graph_decomposition", "diagonalize_measurements", + "device_based_decomposition", ] diff --git a/frontend/test/lit/DeviceBasedDecomposition/TestDeviceBasedDecomposition.mlir b/frontend/test/lit/DeviceBasedDecomposition/TestDeviceBasedDecomposition.mlir new file mode 100644 index 0000000000..f5593a9466 --- /dev/null +++ b/frontend/test/lit/DeviceBasedDecomposition/TestDeviceBasedDecomposition.mlir @@ -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 +} diff --git a/frontend/test/pytest/from_plxpr/test_preprocessing.py b/frontend/test/pytest/from_plxpr/test_preprocessing.py index a3cff3e9ba..972f444af4 100644 --- a/frontend/test/pytest/from_plxpr/test_preprocessing.py +++ b/frontend/test/pytest/from_plxpr/test_preprocessing.py @@ -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") @@ -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.""" diff --git a/mlir/include/Quantum/Transforms/Passes.td b/mlir/include/Quantum/Transforms/Passes.td index b1d7724f1a..0f8e4d68b7 100644 --- a/mlir/include/Quantum/Transforms/Passes.td +++ b/mlir/include/Quantum/Transforms/Passes.td @@ -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."; diff --git a/mlir/lib/Quantum/Transforms/CMakeLists.txt b/mlir/lib/Quantum/Transforms/CMakeLists.txt index a7513b75ce..a973c3dde1 100644 --- a/mlir/lib/Quantum/Transforms/CMakeLists.txt +++ b/mlir/lib/Quantum/Transforms/CMakeLists.txt @@ -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 diff --git a/mlir/lib/Quantum/Transforms/DeviceBasedDecomposition.cpp b/mlir/lib/Quantum/Transforms/DeviceBasedDecomposition.cpp new file mode 100644 index 0000000000..50d11d9ca6 --- /dev/null +++ b/mlir/lib/Quantum/Transforms/DeviceBasedDecomposition.cpp @@ -0,0 +1,93 @@ +// 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. + +#define DEBUG_TYPE "remove-global-phases" + +#include "llvm/Support/Debug.h" +#include "mlir/Dialect/Func/IR/FuncOps.h" +#include "mlir/IR/Operation.h" +#include "mlir/IR/PatternMatch.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Pass/PassManager.h" + +#include "Quantum/Transforms/Passes.h" + +using namespace mlir; +using namespace llvm; + +namespace catalyst { +namespace quantum { + +#define GEN_PASS_DEF_DEVICEBASEDDECOMPOSITIONPASS +#include "Quantum/Transforms/Passes.h.inc" + +struct DeviceBasedDecompositionPass + : public impl::DeviceBasedDecompositionPassBase { + using impl::DeviceBasedDecompositionPassBase< + DeviceBasedDecompositionPass>::DeviceBasedDecompositionPassBase; + + void runOnOperation() final { + LLVM_DEBUG(dbgs() << "DeviceBasedDecompositionPass\n"); + + // This pass runs CtrlLowering -> AdjointLowering -> GraphDecomposition + // The options for this pass are handled in the frontend, and match the requirements + // as per the target device toml file. + + // Run the CtrlLoweringPass + OpPassManager ctrlPM("builtin.module"); + ctrlPM.addPass(createCtrlLoweringPass()); + if (failed(runPipeline(ctrlPM, getOperation()))) { + return signalPassFailure(); + } + + // Run the AdjointLoweringPass + OpPassManager adjointPM("builtin.module"); + adjointPM.addPass(createAdjointLoweringPass()); + if (failed(runPipeline(adjointPM, getOperation()))) { + return signalPassFailure(); + } + + // Populate the options for the GraphDecompositionPass + + // DeviceBasedDecompositionPass's options are the same as GraphDecompositionPassOptions + // Copy all the options + GraphDecompositionPassOptions GDOptions; + + for (auto &targetGate : targetGateSetOption) { + GDOptions.targetGateSetOption.push_back(targetGate); + } + + for (auto &fixedDecomp : fixedDecompsOption) { + GDOptions.fixedDecompsOption.push_back(fixedDecomp); + } + + for (auto &altDecomp : altDecompsOption) { + GDOptions.altDecompsOption.push_back(altDecomp); + } + + GDOptions.bytecodeRulesFile = bytecodeRulesFile; + GDOptions.libQPDPath = libQPDPath; + GDOptions.libpythonPath = libpythonPath; + + // Run the GraphDecompositionPass + OpPassManager gdPM("builtin.module"); + gdPM.addPass(createGraphDecompositionPass(GDOptions)); + if (failed(runPipeline(gdPM, getOperation()))) { + return signalPassFailure(); + } + } +}; + +} // namespace quantum +} // namespace catalyst diff --git a/mlir/lib/Quantum/Transforms/GraphDecomposition/graph_decomposition.cpp b/mlir/lib/Quantum/Transforms/GraphDecomposition/graph_decomposition.cpp index 0e8631cfc5..15f560365a 100644 --- a/mlir/lib/Quantum/Transforms/GraphDecomposition/graph_decomposition.cpp +++ b/mlir/lib/Quantum/Transforms/GraphDecomposition/graph_decomposition.cpp @@ -310,11 +310,17 @@ struct GraphDecompositionPass : public impl::GraphDecompositionPassBase in the case of controllable gates + int index = 0; + opName.consumeInteger(10, index); + bool success = to_float(cost, targetGateSet.ops[opName.str()]); if (!success) { @@ -574,10 +580,10 @@ struct GraphDecompositionPass : public impl::GraphDecompositionPassBase in the case of controllable gates + StringRef name = node.name; + int index = 0; + name.consumeInteger(10, index); + node.name = name.str(); + return node; }