From b9c59525c4ea880fa4b6ab0b39f79765e9b798b5 Mon Sep 17 00:00:00 2001 From: albi3ro Date: Fri, 10 Jul 2026 15:12:24 -0400 Subject: [PATCH 1/2] add state and restored to AllocOp --- .../catalyst/from_plxpr/qfunc_interpreter.py | 2 +- .../from_plxpr/qref_jax_primitives.py | 20 +++++++++++++++---- mlir/include/QRef/IR/QRefOps.td | 4 +++- mlir/include/Quantum/IR/QuantumOps.td | 4 +++- .../Transforms/value_semantics_conversion.cpp | 6 ++++-- .../reference_semantics_conversion.cpp | 7 +++++-- 6 files changed, 32 insertions(+), 11 deletions(-) diff --git a/frontend/catalyst/from_plxpr/qfunc_interpreter.py b/frontend/catalyst/from_plxpr/qfunc_interpreter.py index f644fe1337..834b3936ef 100644 --- a/frontend/catalyst/from_plxpr/qfunc_interpreter.py +++ b/frontend/catalyst/from_plxpr/qfunc_interpreter.py @@ -419,7 +419,7 @@ def handle_allocate(self, *, num_wires, state=None, restored=False): ), "number of dynamically allocated qubits must be statically known" self.has_dynamic_allocation = True - new_qreg = qref_alloc_p.bind(static_num_qubits=num_wires) + new_qreg = qref_alloc_p.bind(static_num_qubits=num_wires, state=state, restored=restored) return [qref_get_p.bind(new_qreg, i) for i in range(num_wires)] diff --git a/frontend/catalyst/from_plxpr/qref_jax_primitives.py b/frontend/catalyst/from_plxpr/qref_jax_primitives.py index 903942dfe3..678ce0163e 100644 --- a/frontend/catalyst/from_plxpr/qref_jax_primitives.py +++ b/frontend/catalyst/from_plxpr/qref_jax_primitives.py @@ -31,6 +31,8 @@ from pennylane.capture.primitives import adjoint_transform_prim as plxpr_adjoint_transform_prim from pennylane.wires import AbstractQubit +from catalyst.jax_extras.lowering import get_mlir_attribute_from_pyval + # TODO: remove after jax v0.7.2 upgrade # Mock _ods_cext.globals.register_traceback_file_exclusion due to API conflicts between # Catalyst's MLIR version and the MLIR version used by JAX. The current JAX version has not @@ -172,7 +174,9 @@ class MeasurementPlane(Enum): # qref_alloc_p # @qref_alloc_p.def_abstract_eval -def _qref_alloc_abstract_eval(*dynamic_num_qubits, static_num_qubits=None): +def _qref_alloc_abstract_eval( + *dynamic_num_qubits, static_num_qubits=None, state=None, restored=False +): static_num_qubits_present = static_num_qubits is not None assert bool(dynamic_num_qubits) ^ static_num_qubits_present if static_num_qubits_present: @@ -182,21 +186,29 @@ def _qref_alloc_abstract_eval(*dynamic_num_qubits, static_num_qubits=None): def _qref_alloc_lowering( - jax_ctx: mlir.LoweringRuleContext, *dynamic_num_qubits, static_num_qubits=None + jax_ctx: mlir.LoweringRuleContext, + *dynamic_num_qubits, + static_num_qubits=None, + state=None, + restored=False, ): static_num_qubits_present = static_num_qubits is not None assert bool(dynamic_num_qubits) ^ static_num_qubits_present ctx = jax_ctx.module_context.context ctx.allow_unregistered_dialects = True + state = str(state.value) if state else "zero" + state = get_mlir_attribute_from_pyval(state) + restored = get_mlir_attribute_from_pyval(restored) + if static_num_qubits_present: size_attr = ir.IntegerAttr.get(ir.IntegerType.get_signless(64, ctx), static_num_qubits) qreg_type = ir.OpaqueType.get("qref", "reg<" + str(static_num_qubits) + ">", ctx) - return AllocOp(qreg_type, nqubits_attr=size_attr).results + return AllocOp(qreg_type, nqubits_attr=size_attr, state=state, restored=restored).results else: size_value = extract_scalar(dynamic_num_qubits[0], "qref_alloc") qreg_type = ir.OpaqueType.get("qref", "reg", ctx) - return AllocOp(qreg_type, nqubits=size_value).results + return AllocOp(qreg_type, nqubits=size_value, state=state, restored=restored).results # diff --git a/mlir/include/QRef/IR/QRefOps.td b/mlir/include/QRef/IR/QRefOps.td index 9d5ad72f95..fe8bf9d060 100644 --- a/mlir/include/QRef/IR/QRefOps.td +++ b/mlir/include/QRef/IR/QRefOps.td @@ -43,7 +43,9 @@ def AllocOp : Memory_Op<"alloc"> { let arguments = (ins Optional:$nqubits, - OptionalAttr>:$nqubits_attr + OptionalAttr>:$nqubits_attr, + DefaultValuedAttr:$state, + DefaultValuedAttr:$restored ); let results = (outs diff --git a/mlir/include/Quantum/IR/QuantumOps.td b/mlir/include/Quantum/IR/QuantumOps.td index f29b216dbd..9038596976 100644 --- a/mlir/include/Quantum/IR/QuantumOps.td +++ b/mlir/include/Quantum/IR/QuantumOps.td @@ -118,7 +118,9 @@ def AllocOp : Memory_Op<"alloc"> { let arguments = (ins Optional:$nqubits, - OptionalAttr>:$nqubits_attr + OptionalAttr>:$nqubits_attr, + DefaultValuedAttr:$state, + DefaultValuedAttr:$restored ); let results = (outs diff --git a/mlir/lib/QRef/Transforms/value_semantics_conversion.cpp b/mlir/lib/QRef/Transforms/value_semantics_conversion.cpp index 68cb97ea67..676c60a57d 100644 --- a/mlir/lib/QRef/Transforms/value_semantics_conversion.cpp +++ b/mlir/lib/QRef/Transforms/value_semantics_conversion.cpp @@ -1051,10 +1051,12 @@ void handleAlloc(IRRewriter &builder, qref::AllocOp rAllocOp, QubitValueTracker std::optional nqubitsAttr = rAllocOp.getNqubitsAttr(); if (nqubitsAttr.has_value()) { vAllocOp = quantum::AllocOp::create(builder, loc, qregType, {}, - IntegerAttr::get(i64Type, *nqubitsAttr)); + IntegerAttr::get(i64Type, *nqubitsAttr), + rAllocOp.getStateAttr(), rAllocOp.getRestoredAttr()); } else { - vAllocOp = quantum::AllocOp::create(builder, loc, qregType, rAllocOp.getNqubits(), nullptr); + vAllocOp = quantum::AllocOp::create(builder, loc, qregType, rAllocOp.getNqubits(), nullptr, + rAllocOp.getStateAttr(), rAllocOp.getRestoredAttr()); } tracker.setCurrentVQreg(rAllocOp.getQreg(), vAllocOp.getQreg()); } diff --git a/mlir/lib/Quantum/Transforms/reference_semantics_conversion.cpp b/mlir/lib/Quantum/Transforms/reference_semantics_conversion.cpp index d89082f4ed..e1dd7c0ba4 100644 --- a/mlir/lib/Quantum/Transforms/reference_semantics_conversion.cpp +++ b/mlir/lib/Quantum/Transforms/reference_semantics_conversion.cpp @@ -238,11 +238,14 @@ void handleAlloc(IRRewriter &builder, quantum::AllocOp vAllocOp, QubitValueTrack Type qregType; if (vAllocOp.getNqubitsAttr().has_value()) { qregType = qref::QuregType::get(ctx, vAllocOp.getNqubitsAttrAttr()); - rAllocOp = qref::AllocOp::create(builder, loc, qregType, {}, vAllocOp.getNqubitsAttrAttr()); + rAllocOp = qref::AllocOp::create(builder, loc, qregType, {}, vAllocOp.getNqubitsAttrAttr(), + vAllocOp.getStateAttr(), vAllocOp.getRestoredAttr()); } else { qregType = qref::QuregType::get(ctx, builder.getI64IntegerAttr(ShapedType::kDynamic)); - rAllocOp = qref::AllocOp::create(builder, loc, qregType, vAllocOp.getNqubits(), nullptr); + rAllocOp = qref::AllocOp::create(builder, loc, qregType, vAllocOp.getNqubits(), nullptr, + vAllocOp.getNqubitsAttrAttr(), vAllocOp.getStateAttr(), + vAllocOp.getRestoredAttr()); } tracker.setRQreg(vAllocOp.getQreg(), rAllocOp.getQreg()); From a83f052e11ec8af20dae1927356459b9ca4d7969 Mon Sep 17 00:00:00 2001 From: albi3ro Date: Mon, 13 Jul 2026 17:28:05 -0400 Subject: [PATCH 2/2] fixed by Cursor, accidentally left an extra argument. --- mlir/lib/Quantum/Transforms/reference_semantics_conversion.cpp | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/mlir/lib/Quantum/Transforms/reference_semantics_conversion.cpp b/mlir/lib/Quantum/Transforms/reference_semantics_conversion.cpp index e1dd7c0ba4..916b60477a 100644 --- a/mlir/lib/Quantum/Transforms/reference_semantics_conversion.cpp +++ b/mlir/lib/Quantum/Transforms/reference_semantics_conversion.cpp @@ -244,8 +244,7 @@ void handleAlloc(IRRewriter &builder, quantum::AllocOp vAllocOp, QubitValueTrack else { qregType = qref::QuregType::get(ctx, builder.getI64IntegerAttr(ShapedType::kDynamic)); rAllocOp = qref::AllocOp::create(builder, loc, qregType, vAllocOp.getNqubits(), nullptr, - vAllocOp.getNqubitsAttrAttr(), vAllocOp.getStateAttr(), - vAllocOp.getRestoredAttr()); + vAllocOp.getStateAttr(), vAllocOp.getRestoredAttr()); } tracker.setRQreg(vAllocOp.getQreg(), rAllocOp.getQreg());