From 7b5a04d402961f5ab1e1b260625ff8985b135d15 Mon Sep 17 00:00:00 2001 From: Nikhil Sreekumar Date: Tue, 1 Sep 2026 17:40:03 -0400 Subject: [PATCH 1/9] draft implementation of adding decomposition rules based on backend gateset --- frontend/catalyst/from_plxpr/device_utils.py | 33 ++++++++++++++++++++ 1 file changed, 33 insertions(+) diff --git a/frontend/catalyst/from_plxpr/device_utils.py b/frontend/catalyst/from_plxpr/device_utils.py index 739734a4e4..b73ea67bae 100644 --- a/frontend/catalyst/from_plxpr/device_utils.py +++ b/frontend/catalyst/from_plxpr/device_utils.py @@ -47,6 +47,7 @@ verify_operations, ) from catalyst.utils.exceptions import CompileError +from catalyst.passes.builtin_passes import graph_decomposition_setup_inputs _named_obs_dict = { "PauliX": qp.X, @@ -99,6 +100,9 @@ def create_device_preprocessing_pipeline( _gradient_preprocessing( pipeline, unsupported_transforms, device, execution_config, shots, capabilities ) + _gateset_preprocessing( + pipeline, unsupported_transforms, device, execution_config, shots, capabilities + ) if unsupported_transforms and warn: warnings.warn( @@ -281,6 +285,35 @@ 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 `graph-decomposition` pass targetting the gateset + specified by the specific `device`""" + gate_set = [] + + # Go through capabilities to populate gate_set + for gate, properties in capabilities.operations.items(): + gate_set.append(gate) + + if properties.invertible: + gate_set.append(f"Adjoint({gate})") + if properties.controllable: + gate_set.append(f"C({gate})") + + # Get the default args/kwargs with the above gate_set + targs, tkwargs = graph_decomposition_setup_inputs(gate_set=gate_set) + t = qp.transform(pass_name="graph_decomposition") + + return BoundTransform(t, args=targs, kwargs=tkwargs) + + def _safe_create_bound_transform( transform: Transform, unsupported_transforms: list[str], warn=True, args=(), kwargs=None ) -> BoundTransform: From 36d43e13191636158f4ec24348c8ce5ef00f931a Mon Sep 17 00:00:00 2001 From: Nikhil Sreekumar Date: Wed, 9 Sep 2026 16:29:49 -0400 Subject: [PATCH 2/9] Added gateset preprocessing --- frontend/catalyst/from_plxpr/device_utils.py | 43 +++++++++----- frontend/catalyst/passes/builtin_passes.py | 27 +++++++++ .../from_plxpr/test_capture_integration.py | 8 +++ .../pytest/from_plxpr/test_preprocessing.py | 57 +++++++++++++++++++ 4 files changed, 121 insertions(+), 14 deletions(-) diff --git a/frontend/catalyst/from_plxpr/device_utils.py b/frontend/catalyst/from_plxpr/device_utils.py index b73ea67bae..880e7a8985 100644 --- a/frontend/catalyst/from_plxpr/device_utils.py +++ b/frontend/catalyst/from_plxpr/device_utils.py @@ -28,6 +28,8 @@ ) from pennylane.transforms.core import BoundTransform, Transform +from catalyst.passes.builtin_passes import adjoint_lowering, ctrl_lowering + from catalyst.device.decomposition import ( measurements_from_counts, measurements_from_samples, @@ -295,24 +297,37 @@ def _gateset_preprocessing( capabilities: DeviceCapabilities, ) -> None: """Insert a `graph-decomposition` pass targetting the gateset - specified by the specific `device`""" - gate_set = [] - - # Go through capabilities to populate gate_set - for gate, properties in capabilities.operations.items(): - gate_set.append(gate) - - if properties.invertible: - gate_set.append(f"Adjoint({gate})") - if properties.controllable: - gate_set.append(f"C({gate})") + specified by the specific `device` + + Note that `graph-decomposition` needs `adjoint-lowering` and `ctrl-lowering` to be run before it + """ + + # Insert adjoint-lowering + pipeline.append( + _safe_create_bound_transform( + adjoint_lowering, unsupported_transforms + ) + ) + + # Insert ctrl-lowering + pipeline.append( + _safe_create_bound_transform( + ctrl_lowering, unsupported_transforms + ) + ) + + gate_set = capabilities.gate_set() # Get the default args/kwargs with the above gate_set targs, tkwargs = graph_decomposition_setup_inputs(gate_set=gate_set) - t = qp.transform(pass_name="graph_decomposition") + t = qp.transform(pass_name="graph-decomposition") + + pipeline.append( + _safe_create_bound_transform( + t, unsupported_transforms, args=targs, kwargs=tkwargs + ) + ) - return BoundTransform(t, args=targs, kwargs=tkwargs) - def _safe_create_bound_transform( transform: Transform, unsupported_transforms: list[str], warn=True, args=(), kwargs=None diff --git a/frontend/catalyst/passes/builtin_passes.py b/frontend/catalyst/passes/builtin_passes.py index a43d754140..31ec9576d1 100644 --- a/frontend/catalyst/passes/builtin_passes.py +++ b/frontend/catalyst/passes/builtin_passes.py @@ -1855,6 +1855,31 @@ def rule_ref_name(rule): pass_name="graph-decomposition", setup_inputs=graph_decomposition_setup_inputs ) + +def adjoint_lowering_setup_inputs(): + r""" + The `adjoint-lowering` pass lowers the adjoint over the region. + """ + return (), {} + +adjoint_lowering = qp.transform( + pass_name="adjoint-lowering", setup_inputs=adjoint_lowering_setup_inputs +) + +def ctrl_lowering_setup_inputs(): + r""" + The `ctrl-lowering` pass distributes the controls over the region: every gate in the + region gains the control qubits/values (appended to any controls it already carries), + and the control qubits are threaded through the region. Structural ops (extract, insert, + alloc, dealloc) are passed through unchanged, and a nested `quantum.ctrl` region has its + controls merged. Measurements inside a `quantum.ctrl` region are rejected. + """ + return (), {} + +ctrl_lowering = qp.transform( + pass_name="ctrl-lowering", setup_inputs=ctrl_lowering_setup_inputs +) + __all__ = [ "cancel_inverses", "combine_global_phases", @@ -1875,4 +1900,6 @@ def rule_ref_name(rule): "decompose_arbitrary_ppr", "graph_decomposition", "diagonalize_measurements", + "adjoint_lowering", + "ctrl_lowering" ] diff --git a/frontend/test/pytest/from_plxpr/test_capture_integration.py b/frontend/test/pytest/from_plxpr/test_capture_integration.py index a6b05c594c..9f93b32834 100644 --- a/frontend/test/pytest/from_plxpr/test_capture_integration.py +++ b/frontend/test/pytest/from_plxpr/test_capture_integration.py @@ -347,6 +347,14 @@ def test_measure(self, backend, reset, op, expected): capture enabled. Hence, we only test that a simple example with a deterministic outcome returns correct results. """ + + if op == qp.I: + + pytest.xfail( + "Waiting for a fix on adjoint rule synthesis of Identity" + ) + + device = qp.device(backend, wires=1) @qjit(capture=True, collect_decomp_rules=False) diff --git a/frontend/test/pytest/from_plxpr/test_preprocessing.py b/frontend/test/pytest/from_plxpr/test_preprocessing.py index a3cff3e9ba..a9356f0c7f 100644 --- a/frontend/test/pytest/from_plxpr/test_preprocessing.py +++ b/frontend/test/pytest/from_plxpr/test_preprocessing.py @@ -585,6 +585,63 @@ def f(): ) +class TestGatesetPreprocessing: + """Tests for preprocessing related to invoking `graph-decomposition` pass with target gateset + as described by dveice TOML file.""" + + def test_gateset_obs_validation(self): + """Tests that the transforms for graph decomposition are added to the pipeline""" + dev = qp.device("null.qubit", wires=4) + + @qp.qnode(dev) + def f(): + return qp.expval(qp.Z(0)) + + device_pipelines = get_pipelines(f, skip_preprocess=False)[1][1] + + # AdjointLowering, CtrlLowering and GraphDecomposition must exist in that order, sequentially + # in the device pipeline + idx = -1 + for idx, pass_entry in enumerate(device_pipelines): + if pass_entry.pass_name == "adjoint-lowering": + break + + assert idx != -1 + assert device_pipelines[idx].pass_name == "adjoint-lowering" + assert device_pipelines[idx + 1].pass_name == "ctrl-lowering" + assert device_pipelines[idx + 2].pass_name == "graph-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": []}, + ) + + @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 == "graph-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.""" From 04d8c7faf42433f85cc87455a9c7d8f3438a0faa Mon Sep 17 00:00:00 2001 From: Nikhil Sreekumar Date: Thu, 10 Sep 2026 11:44:19 -0400 Subject: [PATCH 3/9] formatted + added parsing patches to fix graph-decomp failures --- frontend/catalyst/from_plxpr/device_utils.py | 26 +++++++------------ frontend/catalyst/passes/builtin_passes.py | 19 +++++++------- .../from_plxpr/test_capture_integration.py | 8 ------ .../pytest/from_plxpr/test_preprocessing.py | 22 +++++++++++----- .../graph_decomposition.cpp | 10 +++++++ 5 files changed, 44 insertions(+), 41 deletions(-) diff --git a/frontend/catalyst/from_plxpr/device_utils.py b/frontend/catalyst/from_plxpr/device_utils.py index 880e7a8985..6b60c9e49f 100644 --- a/frontend/catalyst/from_plxpr/device_utils.py +++ b/frontend/catalyst/from_plxpr/device_utils.py @@ -28,8 +28,6 @@ ) from pennylane.transforms.core import BoundTransform, Transform -from catalyst.passes.builtin_passes import adjoint_lowering, ctrl_lowering - from catalyst.device.decomposition import ( measurements_from_counts, measurements_from_samples, @@ -48,8 +46,12 @@ verify_no_state_variance_returns, verify_operations, ) +from catalyst.passes.builtin_passes import ( + adjoint_lowering, + ctrl_lowering, + graph_decomposition_setup_inputs, +) from catalyst.utils.exceptions import CompileError -from catalyst.passes.builtin_passes import graph_decomposition_setup_inputs _named_obs_dict = { "PauliX": qp.X, @@ -298,23 +300,15 @@ def _gateset_preprocessing( ) -> None: """Insert a `graph-decomposition` pass targetting the gateset specified by the specific `device` - + Note that `graph-decomposition` needs `adjoint-lowering` and `ctrl-lowering` to be run before it """ # Insert adjoint-lowering - pipeline.append( - _safe_create_bound_transform( - adjoint_lowering, unsupported_transforms - ) - ) + pipeline.append(_safe_create_bound_transform(adjoint_lowering, unsupported_transforms)) # Insert ctrl-lowering - pipeline.append( - _safe_create_bound_transform( - ctrl_lowering, unsupported_transforms - ) - ) + pipeline.append(_safe_create_bound_transform(ctrl_lowering, unsupported_transforms)) gate_set = capabilities.gate_set() @@ -323,9 +317,7 @@ def _gateset_preprocessing( t = qp.transform(pass_name="graph-decomposition") pipeline.append( - _safe_create_bound_transform( - t, unsupported_transforms, args=targs, kwargs=tkwargs - ) + _safe_create_bound_transform(t, unsupported_transforms, args=targs, kwargs=tkwargs) ) diff --git a/frontend/catalyst/passes/builtin_passes.py b/frontend/catalyst/passes/builtin_passes.py index 31ec9576d1..f229eb1c50 100644 --- a/frontend/catalyst/passes/builtin_passes.py +++ b/frontend/catalyst/passes/builtin_passes.py @@ -1862,23 +1862,24 @@ def adjoint_lowering_setup_inputs(): """ return (), {} + adjoint_lowering = qp.transform( pass_name="adjoint-lowering", setup_inputs=adjoint_lowering_setup_inputs ) + def ctrl_lowering_setup_inputs(): r""" - The `ctrl-lowering` pass distributes the controls over the region: every gate in the - region gains the control qubits/values (appended to any controls it already carries), - and the control qubits are threaded through the region. Structural ops (extract, insert, - alloc, dealloc) are passed through unchanged, and a nested `quantum.ctrl` region has its - controls merged. Measurements inside a `quantum.ctrl` region are rejected. + The `ctrl-lowering` pass distributes the controls over the region: every gate in the + region gains the control qubits/values (appended to any controls it already carries), + and the control qubits are threaded through the region. Structural ops (extract, insert, + alloc, dealloc) are passed through unchanged, and a nested `quantum.ctrl` region has its + controls merged. Measurements inside a `quantum.ctrl` region are rejected. """ return (), {} -ctrl_lowering = qp.transform( - pass_name="ctrl-lowering", setup_inputs=ctrl_lowering_setup_inputs -) + +ctrl_lowering = qp.transform(pass_name="ctrl-lowering", setup_inputs=ctrl_lowering_setup_inputs) __all__ = [ "cancel_inverses", @@ -1901,5 +1902,5 @@ def ctrl_lowering_setup_inputs(): "graph_decomposition", "diagonalize_measurements", "adjoint_lowering", - "ctrl_lowering" + "ctrl_lowering", ] diff --git a/frontend/test/pytest/from_plxpr/test_capture_integration.py b/frontend/test/pytest/from_plxpr/test_capture_integration.py index 9f93b32834..a6b05c594c 100644 --- a/frontend/test/pytest/from_plxpr/test_capture_integration.py +++ b/frontend/test/pytest/from_plxpr/test_capture_integration.py @@ -347,14 +347,6 @@ def test_measure(self, backend, reset, op, expected): capture enabled. Hence, we only test that a simple example with a deterministic outcome returns correct results. """ - - if op == qp.I: - - pytest.xfail( - "Waiting for a fix on adjoint rule synthesis of Identity" - ) - - device = qp.device(backend, wires=1) @qjit(capture=True, collect_decomp_rules=False) diff --git a/frontend/test/pytest/from_plxpr/test_preprocessing.py b/frontend/test/pytest/from_plxpr/test_preprocessing.py index a9356f0c7f..886cf9b111 100644 --- a/frontend/test/pytest/from_plxpr/test_preprocessing.py +++ b/frontend/test/pytest/from_plxpr/test_preprocessing.py @@ -631,15 +631,23 @@ def f(): 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 == "graph-decomposition" + t.kwargs["gate_set"] for t in device_pipelines if t.pass_name == "graph-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 + 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: diff --git a/mlir/lib/Quantum/Transforms/GraphDecomposition/graph_decomposition.cpp b/mlir/lib/Quantum/Transforms/GraphDecomposition/graph_decomposition.cpp index 0e8631cfc5..d0db1d86f0 100644 --- a/mlir/lib/Quantum/Transforms/GraphDecomposition/graph_decomposition.cpp +++ b/mlir/lib/Quantum/Transforms/GraphDecomposition/graph_decomposition.cpp @@ -310,6 +310,8 @@ struct GraphDecompositionPass : public impl::GraphDecompositionPassBase{params}{wires}{static}[uid]", where already carries * any name-wrapped op-level modifiers produced by `defaultGetGraphOpId`, e.g. * "C(Adjoint(RX)){0:[f64]}{wires:1}{}". + * + * Also remove the numeric prefix from , which is introduced in controllable gates, + * for gateset checking. + * e.g. "2C(RY){0:[f64]}{wires:1}{}" -> "C(RY)" and not "2C(RY)" */ OperatorNode parseOperator(llvm::StringRef raw) { OperatorNode node; + int index = 0; + // Consume the numeric prefix + raw.consumeInteger(10, index); + // Base op: either the graphOpId "Name{...}..." form or the legacy "Name(w,p)" form. if (raw.contains('[') || raw.contains('{')) { node.id = raw.str(); From 2119120ff8b6b5edd66aa8b547477b8b9b14e956 Mon Sep 17 00:00:00 2001 From: Nikhil Sreekumar Date: Thu, 10 Sep 2026 17:23:43 -0400 Subject: [PATCH 4/9] Add the device based decomposition pass and invoke with device preprocessing --- frontend/catalyst/from_plxpr/device_utils.py | 27 +++--- frontend/catalyst/from_plxpr/from_plxpr.py | 19 +++- frontend/catalyst/passes/builtin_passes.py | 31 ++----- mlir/include/Quantum/Transforms/Passes.td | 42 +++++++++ mlir/lib/Quantum/Transforms/CMakeLists.txt | 1 + .../Transforms/DeviceBasedDecomposition.cpp | 92 +++++++++++++++++++ 6 files changed, 171 insertions(+), 41 deletions(-) create mode 100644 mlir/lib/Quantum/Transforms/DeviceBasedDecomposition.cpp diff --git a/frontend/catalyst/from_plxpr/device_utils.py b/frontend/catalyst/from_plxpr/device_utils.py index 6b60c9e49f..b48f07f70c 100644 --- a/frontend/catalyst/from_plxpr/device_utils.py +++ b/frontend/catalyst/from_plxpr/device_utils.py @@ -62,9 +62,10 @@ 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.""" + """Create a pipeline of device preprocessing transforms for lowering QNodes. + """ shots_present = qp.math.is_abstract(shots) or shots != 0 raw_capabilities: DeviceCapabilities = get_qjit_device_capabilities( _load_device_capabilities(device) @@ -104,9 +105,11 @@ def create_device_preprocessing_pipeline( _gradient_preprocessing( pipeline, unsupported_transforms, device, execution_config, shots, capabilities ) - _gateset_preprocessing( - 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( @@ -298,23 +301,15 @@ def _gateset_preprocessing( shots: int, capabilities: DeviceCapabilities, ) -> None: - """Insert a `graph-decomposition` pass targetting the gateset + """Insert a `device-based-decomposition` pass targetting the gateset specified by the specific `device` - - Note that `graph-decomposition` needs `adjoint-lowering` and `ctrl-lowering` to be run before it """ - - # Insert adjoint-lowering - pipeline.append(_safe_create_bound_transform(adjoint_lowering, unsupported_transforms)) - - # Insert ctrl-lowering - pipeline.append(_safe_create_bound_transform(ctrl_lowering, unsupported_transforms)) - 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="graph-decomposition") + t = qp.transform(pass_name="device-based-decomposition") pipeline.append( _safe_create_bound_transform(t, unsupported_transforms, args=targs, kwargs=tkwargs) diff --git a/frontend/catalyst/from_plxpr/from_plxpr.py b/frontend/catalyst/from_plxpr/from_plxpr.py index 00155ea2c6..adc9cfd0bb 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,7 @@ 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 +408,8 @@ def handle_transform( ): """Handle the conversion from plxpr to Catalyst jaxpr for a PL transform.""" + + self.check = True 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 +428,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 f229eb1c50..37afa4a94c 100644 --- a/frontend/catalyst/passes/builtin_passes.py +++ b/frontend/catalyst/passes/builtin_passes.py @@ -1856,31 +1856,21 @@ def rule_ref_name(rule): ) -def adjoint_lowering_setup_inputs(): - r""" - The `adjoint-lowering` pass lowers the adjoint over the region. +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 (), {} -adjoint_lowering = qp.transform( - pass_name="adjoint-lowering", setup_inputs=adjoint_lowering_setup_inputs +device_based_decomposition = qp.transform( + pass_name="device-based-decomposition", setup_inputs=device_based_decomposition_setup_inputs ) - -def ctrl_lowering_setup_inputs(): - r""" - The `ctrl-lowering` pass distributes the controls over the region: every gate in the - region gains the control qubits/values (appended to any controls it already carries), - and the control qubits are threaded through the region. Structural ops (extract, insert, - alloc, dealloc) are passed through unchanged, and a nested `quantum.ctrl` region has its - controls merged. Measurements inside a `quantum.ctrl` region are rejected. - """ - return (), {} - - -ctrl_lowering = qp.transform(pass_name="ctrl-lowering", setup_inputs=ctrl_lowering_setup_inputs) - __all__ = [ "cancel_inverses", "combine_global_phases", @@ -1901,6 +1891,5 @@ def ctrl_lowering_setup_inputs(): "decompose_arbitrary_ppr", "graph_decomposition", "diagonalize_measurements", - "adjoint_lowering", - "ctrl_lowering", + "device_based_decomposition", ] 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..61073afb8b --- /dev/null +++ b/mlir/lib/Quantum/Transforms/DeviceBasedDecomposition.cpp @@ -0,0 +1,92 @@ +// 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 "llvm/Support/Debug.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::DeviceBasedDecompositionPassBase; + + void runOnOperation() final { + LLVM_DEBUG(dbgs() << "DeviceBasedDecompositionPass\n"); + + // This pass runs AdjointLowering -> CtrlLowering -> GraphDecomposition + // The options for this pass are handled in the frontend, and match the requirements + // as per the target device toml file. + + // Run the AdjointLoweringPass + OpPassManager adjointPM("builtin.module"); + adjointPM.addPass(createAdjointLoweringPass()); + if (failed(runPipeline(adjointPM, getOperation()))) { + return signalPassFailure(); + } + + // Run the CtrlLoweringPass + OpPassManager ctrlPM("builtin.module"); + ctrlPM.addPass(createCtrlLoweringPass()); + if (failed(runPipeline(ctrlPM, 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 From f465e170742925eb544b41665b23c5b6894b24ab Mon Sep 17 00:00:00 2001 From: Nikhil Sreekumar Date: Fri, 11 Sep 2026 12:10:56 -0400 Subject: [PATCH 5/9] tests updated to match optional nature of device based decomp pass --- .../pytest/from_plxpr/test_preprocessing.py | 38 ++++++++++--------- 1 file changed, 20 insertions(+), 18 deletions(-) diff --git a/frontend/test/pytest/from_plxpr/test_preprocessing.py b/frontend/test/pytest/from_plxpr/test_preprocessing.py index 886cf9b111..8414734505 100644 --- a/frontend/test/pytest/from_plxpr/test_preprocessing.py +++ b/frontend/test/pytest/from_plxpr/test_preprocessing.py @@ -31,6 +31,7 @@ from catalyst.from_plxpr import from_plxpr from catalyst.jax_primitives import quantum_kernel_p from catalyst.utils.exceptions import CompileError +from catalyst.passes.builtin_passes import device_based_decomposition pytestmark = pytest.mark.usefixtures("use_capture") from_plxpr_no_warn = partial(from_plxpr, _preprocess_warn=False) @@ -586,30 +587,30 @@ def f(): class TestGatesetPreprocessing: - """Tests for preprocessing related to invoking `graph-decomposition` pass with target gateset + """Tests for preprocessing related to invoking `device-based-decomposition` pass with target gateset as described by dveice TOML file.""" - def test_gateset_obs_validation(self): - """Tests that the transforms for graph decomposition are added to the pipeline""" + @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) - @qp.qnode(dev) - def f(): - return qp.expval(qp.Z(0)) + 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) - # AdjointLowering, CtrlLowering and GraphDecomposition must exist in that order, sequentially - # in the device pipeline - idx = -1 - for idx, pass_entry in enumerate(device_pipelines): - if pass_entry.pass_name == "adjoint-lowering": - break - - assert idx != -1 - assert device_pipelines[idx].pass_name == "adjoint-lowering" - assert device_pipelines[idx + 1].pass_name == "ctrl-lowering" - assert device_pipelines[idx + 2].pass_name == "graph-decomposition" + 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 @@ -625,13 +626,14 @@ def test_gateset_matches_device_capabilities(self): 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 == "graph-decomposition" + t.kwargs["gate_set"] for t in device_pipelines if t.pass_name == "device-based-decomposition" ) assert ( From 5d77d4d485100153eed31543510ea6b03848deed Mon Sep 17 00:00:00 2001 From: Nikhil Sreekumar Date: Fri, 11 Sep 2026 18:56:48 -0400 Subject: [PATCH 6/9] corrected parsing., only consumeinteger for name cos thats used for gateset matching --- frontend/catalyst/from_plxpr/device_utils.py | 2 -- .../graph_decomposition.cpp | 24 ++++++++++--------- 2 files changed, 13 insertions(+), 13 deletions(-) diff --git a/frontend/catalyst/from_plxpr/device_utils.py b/frontend/catalyst/from_plxpr/device_utils.py index b48f07f70c..66fdbde509 100644 --- a/frontend/catalyst/from_plxpr/device_utils.py +++ b/frontend/catalyst/from_plxpr/device_utils.py @@ -47,8 +47,6 @@ verify_operations, ) from catalyst.passes.builtin_passes import ( - adjoint_lowering, - ctrl_lowering, graph_decomposition_setup_inputs, ) from catalyst.utils.exceptions import CompileError diff --git a/mlir/lib/Quantum/Transforms/GraphDecomposition/graph_decomposition.cpp b/mlir/lib/Quantum/Transforms/GraphDecomposition/graph_decomposition.cpp index d0db1d86f0..15f560365a 100644 --- a/mlir/lib/Quantum/Transforms/GraphDecomposition/graph_decomposition.cpp +++ b/mlir/lib/Quantum/Transforms/GraphDecomposition/graph_decomposition.cpp @@ -317,6 +317,10 @@ 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) { @@ -564,18 +568,10 @@ struct GraphDecompositionPass : public impl::GraphDecompositionPassBase{params}{wires}{static}[uid]", where already carries * any name-wrapped op-level modifiers produced by `defaultGetGraphOpId`, e.g. * "C(Adjoint(RX)){0:[f64]}{wires:1}{}". - * - * Also remove the numeric prefix from , which is introduced in controllable gates, - * for gateset checking. - * e.g. "2C(RY){0:[f64]}{wires:1}{}" -> "C(RY)" and not "2C(RY)" */ OperatorNode parseOperator(llvm::StringRef raw) { OperatorNode node; - int index = 0; - // Consume the numeric prefix - raw.consumeInteger(10, index); - // Base op: either the graphOpId "Name{...}..." form or the legacy "Name(w,p)" form. if (raw.contains('[') || raw.contains('{')) { node.id = raw.str(); @@ -584,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; } From 93f28e3ade75555082d6bbccd78ff2780067b2b8 Mon Sep 17 00:00:00 2001 From: Nikhil Sreekumar Date: Mon, 14 Sep 2026 09:46:43 -0400 Subject: [PATCH 7/9] Modify pass to run successfully for test. but address nested adj/ctrl --- .../TestDeviceBasedDecomposition.mlir | 44 +++++++++++++++++++ .../Transforms/DeviceBasedDecomposition.cpp | 14 +++--- 2 files changed, 51 insertions(+), 7 deletions(-) create mode 100644 frontend/test/lit/DeviceBasedDecomposition/TestDeviceBasedDecomposition.mlir 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/mlir/lib/Quantum/Transforms/DeviceBasedDecomposition.cpp b/mlir/lib/Quantum/Transforms/DeviceBasedDecomposition.cpp index 61073afb8b..3b1b9016ad 100644 --- a/mlir/lib/Quantum/Transforms/DeviceBasedDecomposition.cpp +++ b/mlir/lib/Quantum/Transforms/DeviceBasedDecomposition.cpp @@ -43,13 +43,6 @@ struct DeviceBasedDecompositionPass : public impl::DeviceBasedDecompositionPassB // The options for this pass are handled in the frontend, and match the requirements // as per the target device toml file. - // Run the AdjointLoweringPass - OpPassManager adjointPM("builtin.module"); - adjointPM.addPass(createAdjointLoweringPass()); - if (failed(runPipeline(adjointPM, getOperation()))) { - return signalPassFailure(); - } - // Run the CtrlLoweringPass OpPassManager ctrlPM("builtin.module"); ctrlPM.addPass(createCtrlLoweringPass()); @@ -57,6 +50,13 @@ struct DeviceBasedDecompositionPass : public impl::DeviceBasedDecompositionPassB 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 From 77005e96aa3771735cd02f71508db188b87f72c1 Mon Sep 17 00:00:00 2001 From: Nikhil Sreekumar Date: Mon, 14 Sep 2026 10:53:12 -0400 Subject: [PATCH 8/9] formatting --- frontend/catalyst/from_plxpr/device_utils.py | 9 ++++++--- frontend/catalyst/from_plxpr/from_plxpr.py | 8 ++++++-- frontend/catalyst/passes/builtin_passes.py | 2 +- .../pytest/from_plxpr/test_preprocessing.py | 11 ++++++++--- .../Transforms/DeviceBasedDecomposition.cpp | 19 ++++++++++--------- 5 files changed, 31 insertions(+), 18 deletions(-) diff --git a/frontend/catalyst/from_plxpr/device_utils.py b/frontend/catalyst/from_plxpr/device_utils.py index 66fdbde509..c0b89044dc 100644 --- a/frontend/catalyst/from_plxpr/device_utils.py +++ b/frontend/catalyst/from_plxpr/device_utils.py @@ -60,10 +60,13 @@ def create_device_preprocessing_pipeline( - device: qp.devices.Device, execution_config: ExecutionConfig, shots: int, warn: bool = True, needs_gateset_preprocessing: 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. - """ + """Create a pipeline of device preprocessing transforms for lowering QNodes.""" shots_present = qp.math.is_abstract(shots) or shots != 0 raw_capabilities: DeviceCapabilities = get_qjit_device_capabilities( _load_device_capabilities(device) diff --git a/frontend/catalyst/from_plxpr/from_plxpr.py b/frontend/catalyst/from_plxpr/from_plxpr.py index adc9cfd0bb..30814fc18a 100644 --- a/frontend/catalyst/from_plxpr/from_plxpr.py +++ b/frontend/catalyst/from_plxpr/from_plxpr.py @@ -330,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, needs_gateset_preprocessing=self.needs_gateset_preprocessing + qnode.device, + execution_config, + shots, + warn=self._preprocess_warn, + needs_gateset_preprocessing=self.needs_gateset_preprocessing, ) pipelines += (("device", device_preprocessing_pipeline),) @@ -428,7 +432,7 @@ def handle_transform( # Apply the corresponding Catalyst pass counterpart next_eval = copy(self) - + 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. diff --git a/frontend/catalyst/passes/builtin_passes.py b/frontend/catalyst/passes/builtin_passes.py index 37afa4a94c..68f42bae75 100644 --- a/frontend/catalyst/passes/builtin_passes.py +++ b/frontend/catalyst/passes/builtin_passes.py @@ -1861,7 +1861,7 @@ def device_based_decomposition_setup_inputs(): 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 (), {} diff --git a/frontend/test/pytest/from_plxpr/test_preprocessing.py b/frontend/test/pytest/from_plxpr/test_preprocessing.py index 8414734505..972f444af4 100644 --- a/frontend/test/pytest/from_plxpr/test_preprocessing.py +++ b/frontend/test/pytest/from_plxpr/test_preprocessing.py @@ -30,8 +30,8 @@ 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.utils.exceptions import CompileError from catalyst.passes.builtin_passes import device_based_decomposition +from catalyst.utils.exceptions import CompileError pytestmark = pytest.mark.usefixtures("use_capture") from_plxpr_no_warn = partial(from_plxpr, _preprocess_warn=False) @@ -592,17 +592,20 @@ class TestGatesetPreprocessing: @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 + """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)) @@ -633,7 +636,9 @@ def f(): 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" + t.kwargs["gate_set"] + for t in device_pipelines + if t.pass_name == "device-based-decomposition" ) assert ( diff --git a/mlir/lib/Quantum/Transforms/DeviceBasedDecomposition.cpp b/mlir/lib/Quantum/Transforms/DeviceBasedDecomposition.cpp index 3b1b9016ad..f2c6a27377 100644 --- a/mlir/lib/Quantum/Transforms/DeviceBasedDecomposition.cpp +++ b/mlir/lib/Quantum/Transforms/DeviceBasedDecomposition.cpp @@ -18,11 +18,10 @@ #include "mlir/Dialect/Func/IR/FuncOps.h" #include "mlir/IR/Operation.h" #include "mlir/IR/PatternMatch.h" -#include "llvm/Support/Debug.h" #include "mlir/Pass/Pass.h" #include "mlir/Pass/PassManager.h" -#include "Quantum/Transforms/Passes.h" +#include "Quantum/Transforms/Passes.h" using namespace mlir; using namespace llvm; @@ -32,9 +31,11 @@ namespace quantum { #define GEN_PASS_DEF_DEVICEBASEDDECOMPOSITIONPASS #include "Quantum/Transforms/Passes.h.inc" - -struct DeviceBasedDecompositionPass : public impl::DeviceBasedDecompositionPassBase { - using impl::DeviceBasedDecompositionPassBase::DeviceBasedDecompositionPassBase; + +struct DeviceBasedDecompositionPass + : public impl::DeviceBasedDecompositionPassBase { + using impl::DeviceBasedDecompositionPassBase< + DeviceBasedDecompositionPass>::DeviceBasedDecompositionPassBase; void runOnOperation() final { LLVM_DEBUG(dbgs() << "DeviceBasedDecompositionPass\n"); @@ -63,15 +64,15 @@ struct DeviceBasedDecompositionPass : public impl::DeviceBasedDecompositionPassB // Copy all the options GraphDecompositionPassOptions GDOptions; - for(auto& targetGate : targetGateSetOption) { + for (auto &targetGate : targetGateSetOption) { GDOptions.targetGateSetOption.push_back(targetGate); } - for(auto& fixedDecomp : fixedDecompsOption) { + for (auto &fixedDecomp : fixedDecompsOption) { GDOptions.fixedDecompsOption.push_back(fixedDecomp); } - for(auto& altDecomp : altDecompsOption) { + for (auto &altDecomp : altDecompsOption) { GDOptions.altDecompsOption.push_back(altDecomp); } @@ -87,6 +88,6 @@ struct DeviceBasedDecompositionPass : public impl::DeviceBasedDecompositionPassB } } }; - + } // namespace quantum } // namespace catalyst From 80c5d30a79f27a02f75f365d9214ad5f3b107c10 Mon Sep 17 00:00:00 2001 From: Nikhil Sreekumar Date: Mon, 14 Sep 2026 11:22:33 -0400 Subject: [PATCH 9/9] minor corrections --- frontend/catalyst/from_plxpr/from_plxpr.py | 1 - mlir/lib/Quantum/Transforms/DeviceBasedDecomposition.cpp | 2 +- 2 files changed, 1 insertion(+), 2 deletions(-) diff --git a/frontend/catalyst/from_plxpr/from_plxpr.py b/frontend/catalyst/from_plxpr/from_plxpr.py index 30814fc18a..3736e0eb02 100644 --- a/frontend/catalyst/from_plxpr/from_plxpr.py +++ b/frontend/catalyst/from_plxpr/from_plxpr.py @@ -413,7 +413,6 @@ def handle_transform( """Handle the conversion from plxpr to Catalyst jaxpr for a PL transform.""" - self.check = True consts = args[_tuple_to_slice(consts_slice)] non_const_args = args[_tuple_to_slice(args_slice)] targs = args[_tuple_to_slice(targs_slice)] diff --git a/mlir/lib/Quantum/Transforms/DeviceBasedDecomposition.cpp b/mlir/lib/Quantum/Transforms/DeviceBasedDecomposition.cpp index f2c6a27377..50d11d9ca6 100644 --- a/mlir/lib/Quantum/Transforms/DeviceBasedDecomposition.cpp +++ b/mlir/lib/Quantum/Transforms/DeviceBasedDecomposition.cpp @@ -40,7 +40,7 @@ struct DeviceBasedDecompositionPass void runOnOperation() final { LLVM_DEBUG(dbgs() << "DeviceBasedDecompositionPass\n"); - // This pass runs AdjointLowering -> CtrlLowering -> GraphDecomposition + // 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.