Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions doc/releases/changelog-dev.md
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,13 @@

<h3>New features since last release</h3>

* A new PennyLane operation :func:`~.fabricate` has been added to expose the PBC
``pbc.fabricate`` instruction from the frontend. The operation produces a new auxiliary
qubit in a logical factory state (``plus_i``, ``minus_i``, ``magic``, or ``magic_conj``)
and is lowered through the ``pbc.ref.fabricate`` reference-semantics op to
``pbc.fabricate`` during compilation.
[(#3019)](https://github.com/PennyLaneAI/catalyst/pull/3019)

* The `local-random` unitary folding option for :func:`~.mitigate_with_zne` is now implemented,
reproducing Mitiq's ``fold_gates_at_random``: every gate is folded ``floor((scale_factor-1)/2)``
times, then a random subset is folded once more (without replacement) to reach ``scale_factor * n``
Expand Down
4 changes: 4 additions & 0 deletions frontend/catalyst/api_extensions/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,8 +37,10 @@
MidCircuitPauliMeasure,
adjoint,
ctrl,
deallocate,
measure,
pauli_measure,
fabricate,
)

__all__ = (
Expand All @@ -58,6 +60,8 @@
"vmap",
"measure",
"pauli_measure",
"fabricate",
"deallocate",
"adjoint",
"ctrl",
)
59 changes: 58 additions & 1 deletion frontend/catalyst/api_extensions/quantum_operators.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,7 +46,14 @@
deduce_avals,
new_inner_tracer,
)
from catalyst.jax_primitives import AbstractQreg, adjoint_p, measure_p, pauli_measure_p
from catalyst.jax_primitives import (
AbstractQreg,
adjoint_p,
fabricate_p,
measure_p,
pauli_measure_p,
qdealloc_qb_p,
)
from catalyst.jax_tracer import (
HybridOp,
HybridOpRegion,
Expand Down Expand Up @@ -243,6 +250,56 @@ def pauli_measure(
return m


def fabricate(init_state: str) -> DynamicJaxprTracer:
r"""A :func:`qjit` compatible fabricate operation for PennyLane/Catalyst.

.. important::

Under the legacy tracing pathway, use :func:`catalyst.fabricate` or ensure
:func:`qp.fabricate() <pennylane.fabricate>` is called from within an active
:func:`~.qjit` compilation.

Args:
init_state (str): The logical state to fabricate. One of ``"plus_i"``,
``"minus_i"``, ``"magic"``, or ``"magic_conj"``.

Returns:
A JAX tracer for the fabricated qubit.
"""
EvaluationContext.check_is_tracing("catalyst.fabricate can only be used from within @qjit.")
EvaluationContext.check_is_quantum_tracing(
"catalyst.fabricate can only be used from within a qp.qnode."
)

valid_states = {"plus_i", "minus_i", "magic", "magic_conj"}
if init_state not in valid_states:
raise ValueError(
f'The init_state "{init_state}" is not allowed. '
f"Allowed values are {sorted(valid_states)}."
)

(qubit,) = fabricate_p.bind(init_state=init_state)
return qubit


def deallocate(*qubits) -> None:
r"""A :func:`qjit` compatible deallocate operation for standalone fabricated qubits.

Deallocates one or more standalone value-semantics qubits, such as those returned
by :func:`fabricate`.

Args:
*qubits: Standalone qubit tracers to deallocate.
"""
EvaluationContext.check_is_tracing("catalyst.deallocate can only be used from within @qjit.")
EvaluationContext.check_is_quantum_tracing(
"catalyst.deallocate can only be used from within a qp.qnode."
)

for qubit in qubits:
qdealloc_qb_p.bind(qubit)


def adjoint(f: Union[Callable, Operator], lazy=True) -> Union[Callable, Operator]:
"""A :func:`~.qjit` compatible adjoint transformer for PennyLane/Catalyst.

Expand Down
1 change: 1 addition & 0 deletions frontend/catalyst/device/qjit_device.py
Original file line number Diff line number Diff line change
Expand Up @@ -82,6 +82,7 @@
"MultiRZ",
"PauliRot",
"PauliMeasure",
"Fabricate",
"PauliX",
"PauliY",
"PauliZ",
Expand Down
36 changes: 27 additions & 9 deletions frontend/catalyst/from_plxpr/qfunc_interpreter.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@
from pennylane.capture import PlxprInterpreter, pause
from pennylane.capture.primitives import cond_prim as pl_cond_prim
from pennylane.capture.primitives import ctrl_transform_prim as plxpr_ctrl_transform_prim
from pennylane.capture.primitives import fabricate_prim as plxpr_fabricate_prim
from pennylane.capture.primitives import measure_prim as plxpr_measure_prim
from pennylane.capture.primitives import operator_p
from pennylane.capture.primitives import pauli_measure_prim as plxpr_pauli_measure_prim
Expand All @@ -40,6 +41,8 @@
qref_alloc_p,
qref_compbasis_p,
qref_dealloc_p,
qref_dealloc_qb_p,
qref_fabricate_p,
qref_get_p,
qref_gphase_p,
qref_hermitian_p,
Expand Down Expand Up @@ -427,17 +430,25 @@ def handle_allocate(self, *, num_wires, state=None, restored=False):
def handle_deallocate(self, *wires):
"""Handle the conversion from plxpr to Catalyst jaxpr for the qp.deallocate primitive"""
qregs = set()
standalone_qubits = []
for w in wires:
get_op = w.parent
parent_eqn = w.parent
if parent_eqn.primitive is qref_fabricate_p:
standalone_qubits.append(w)
elif parent_eqn.primitive is qref_get_p:
qreg = parent_eqn.in_tracers[0]
qregs.add(qreg)
else:
raise AssertionError(
"Manual deallocation is only supported for manually allocated or fabricated wires"
)
for qubit in standalone_qubits:
qref_dealloc_qb_p.bind(qubit)
if qregs:
assert (
get_op.primitive is qref_get_p
), "Manual deallocation is only supported for manually allocated wires"
qreg = get_op.in_tracers[0]
qregs.add(qreg)
assert (
len(qregs) == 1
), "Expected all wires to deallocate to come from the same allocation instruction"
qref_dealloc_p.bind(list(qregs)[0])
len(qregs) == 1
), "Expected all wires to deallocate to come from the same allocation instruction"
qref_dealloc_p.bind(list(qregs)[0])
return []


Expand Down Expand Up @@ -543,6 +554,13 @@ def wrapper(*args):
return ()


@PLxPRToQuantumJaxprInterpreter.register_primitive(plxpr_fabricate_prim)
def handle_fabricate(self, *_, init_state=""):
"""Handle the conversion from plxpr to Catalyst jaxpr for the fabricate primitive"""
(qubit,) = qref_fabricate_p.bind(init_state=init_state)
return (qubit,)


@PLxPRToQuantumJaxprInterpreter.register_primitive(plxpr_pauli_measure_prim)
def handle_pauli_measure(self, *wires_inval, pauli_word, **params):
"""Handle the conversion from plxpr to Catalyst jaxpr for the PauliMeasure primitive"""
Expand Down
43 changes: 42 additions & 1 deletion frontend/catalyst/from_plxpr/qref_jax_primitives.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@
from catalyst.jax_extras.patches import mock_attributes
from catalyst.jax_primitives import (
AbstractObs,
_logical_init_attr,
_named_obs_attribute,
extract_scalar,
safe_cast_to_f64,
Expand All @@ -61,13 +62,14 @@
),
):
from mlir_quantum.dialects.mbqc import RefMeasureInBasisOp
from mlir_quantum.dialects.pbc import RefPPMeasurementOp
from mlir_quantum.dialects.pbc import RefFabricateOp, RefPPMeasurementOp
from mlir_quantum.dialects.qref import (
AdjointOp,
AllocOp,
ComputationalBasisOp,
CustomOp,
DeallocOp,
DeallocQubitOp,
GetOp,
GlobalPhaseOp,
HermitianOp,
Expand Down Expand Up @@ -147,6 +149,8 @@ class MeasurementPlane(Enum):
qref_alloc_p = Primitive("qref_alloc")
qref_dealloc_p = Primitive("qref_dealloc")
qref_dealloc_p.multiple_results = True
qref_dealloc_qb_p = Primitive("qref_dealloc_qb")
qref_dealloc_qb_p.multiple_results = True
qref_get_p = Primitive("qref_get")
qref_set_state_p = Primitive("qref_state_prep")
qref_set_state_p.multiple_results = True
Expand All @@ -157,6 +161,8 @@ class MeasurementPlane(Enum):
qref_gphase_p = Primitive("qref_gphase")
qref_gphase_p.multiple_results = True
qref_pauli_measure_p = Primitive("pref_pauli_measure")
qref_fabricate_p = Primitive("qref_fabricate")
qref_fabricate_p.multiple_results = True
qref_pauli_rot_p = Primitive("qref_pauli_rot")
qref_pauli_rot_p.multiple_results = True
qref_unitary_p = Primitive("qref_unitary")
Expand Down Expand Up @@ -215,6 +221,21 @@ def _qref_dealloc_lowering(jax_ctx: mlir.LoweringRuleContext, qreg):
return ()


#
# qref_dealloc_qb_p
#
@qref_dealloc_qb_p.def_abstract_eval
def _qref_dealloc_qb_abstract_eval(qubit):
return ()


def _qref_dealloc_qb_lowering(jax_ctx: mlir.LoweringRuleContext, qubit):
ctx = jax_ctx.module_context.context
ctx.allow_unregistered_dialects = True
DeallocQubitOp(qubit=qubit)
return ()


#
# qref_get_p
#
Expand Down Expand Up @@ -542,6 +563,24 @@ def _qref_pauli_measure_lowering(
return (from_elements_op.results[0],)


#
# fabricate operation
#
@qref_fabricate_p.def_abstract_eval
def _qref_fabricate_abstract_eval(*_, init_state=""):
return (AbstractQubit(),)


def _qref_fabricate_lowering(jax_ctx: mlir.LoweringRuleContext, *_, init_state=""):
ctx = jax_ctx.module_context.context
ctx.allow_unregistered_dialects = True

qubit_type = ir.OpaqueType.get("qref", "bit", ctx)
return RefFabricateOp(
qubits=[qubit_type], init_state=_logical_init_attr(ctx, init_state)
).results


#
# qubit unitary operation
#
Expand Down Expand Up @@ -826,13 +865,15 @@ def _qref_hermitian_lowering(jax_ctx: mlir.LoweringRuleContext, matrix: ir.Value
(qref_operator_p, _qref_operator_p_lowering),
(qref_alloc_p, _qref_alloc_lowering),
(qref_dealloc_p, _qref_dealloc_lowering),
(qref_dealloc_qb_p, _qref_dealloc_qb_lowering),
(qref_get_p, _qref_get_lowering),
(qref_set_state_p, _qref_set_state_lowering),
(qref_set_basis_state_p, _qref_set_basis_state_lowering),
(qref_qinst_p, _qref_qinst_lowering),
(qref_gphase_p, _qref_gphase_lowering),
(qref_pauli_rot_p, _qref_pauli_rot_lowering),
(qref_pauli_measure_p, _qref_pauli_measure_lowering),
(qref_fabricate_p, _qref_fabricate_lowering),
(qref_unitary_p, _qref_unitary_lowering),
(qref_measure_p, _qref_measure_lowering),
(qref_measure_in_basis_p, _qref_measure_in_basis_lowering),
Expand Down
32 changes: 31 additions & 1 deletion frontend/catalyst/jax_primitives.py
Original file line number Diff line number Diff line change
Expand Up @@ -87,7 +87,7 @@
VJPOp,
)
from mlir_quantum.dialects.mitigation import ZneOp
from mlir_quantum.dialects.pbc import PPMeasurementOp
from mlir_quantum.dialects.pbc import FabricateOp, PPMeasurementOp
from mlir_quantum.dialects.quantum import (
AdjointOp,
AllocOp,
Expand Down Expand Up @@ -298,6 +298,8 @@ class Folding(Enum):
pauli_rot_p.multiple_results = True
pauli_measure_p = Primitive("pauli_measure")
pauli_measure_p.multiple_results = True
fabricate_p = Primitive("fabricate")
fabricate_p.multiple_results = True
measure_p = Primitive("measure")
measure_p.multiple_results = True
compbasis_p = Primitive("compbasis")
Expand Down Expand Up @@ -1660,6 +1662,29 @@ def _pauli_measure_lowering(
return (from_elements_op.results[0],) + tuple(out_qubits)


#
# fabricate operation
#
@fabricate_p.def_abstract_eval
def _fabricate_abstract_eval(*_, init_state=""):
return (AbstractQbit(),)


@fabricate_p.def_impl
def _fabricate_def_impl(*args, **kwargs): # pragma: no cover
raise NotImplementedError()


def _fabricate_lowering(jax_ctx: mlir.LoweringRuleContext, *_, init_state=""):
ctx = jax_ctx.module_context.context
ctx.allow_unregistered_dialects = True

qubit_type = ir.OpaqueType.get("quantum", "bit", ctx)
return FabricateOp(
out_qubits=[qubit_type], init_state=_logical_init_attr(ctx, init_state)
).results


#
# measure
#
Expand Down Expand Up @@ -1760,6 +1785,10 @@ def _namedobs_abstract_eval(qubit, kind):
return AbstractObs()


def _logical_init_attr(ctx, init_state: str):
return ir.Attribute.parse(f"#pbc<enum {init_state}>", context=ctx)


def _named_obs_attribute(ctx, kind: str):
return ir.OpaqueAttr.get(
"quantum", ("named_observable " + kind).encode("utf-8"), ir.NoneType.get(ctx), ctx
Expand Down Expand Up @@ -3080,6 +3109,7 @@ def subroutine_lowering(*args, **kwargs):
(unitary_p, _unitary_lowering),
(pauli_rot_p, _pauli_rot_lowering),
(pauli_measure_p, _pauli_measure_lowering),
(fabricate_p, _fabricate_lowering),
(measure_p, _measure_lowering),
(compbasis_p, _compbasis_lowering),
(namedobs_p, _named_obs_lowering),
Expand Down
Loading
Loading