From 712eaf689cafe81b442c65d4320bc462af2ec5a9 Mon Sep 17 00:00:00 2001 From: positr0nium Date: Mon, 14 Sep 2026 12:02:58 +0200 Subject: [PATCH] fixed an issue that arose due to custom_control/custom_inverse using two different tracing mechanisms (jit vs make_jaxpr) --- .../general/changelog/changelog-dev.rst | 19 ++- .../custom_control_environment.py | 26 +++- .../custom_inversion_environment.py | 28 +++- .../interpreter_tools/abstract_interpreter.py | 10 +- src/qrisp/jasp/jasp_expression/centerclass.py | 8 +- .../jasp/jasp_expression/control_transform.py | 50 +++---- .../jasp/jasp_expression/inv_transform.py | 117 ++++------------- src/qrisp/jasp/jasp_expression/jaxpr_utils.py | 123 ++++++++++++++++++ tests/jax_tests/test_custom_inverse.py | 108 +++++++++++++++ 9 files changed, 347 insertions(+), 142 deletions(-) diff --git a/documentation/source/general/changelog/changelog-dev.rst b/documentation/source/general/changelog/changelog-dev.rst index 09f48e753..6e2fd6b3e 100644 --- a/documentation/source/general/changelog/changelog-dev.rst +++ b/documentation/source/general/changelog/changelog-dev.rst @@ -110,13 +110,18 @@ Bug Fixes * Fixed a failure when a function decorated with :func:`custom_inversion ` was inverted twice, which raised ``Automatic loop inversion is only supported for jrange-based loops``. - An inverted Jaspr keeps a back-pointer to the Jaspr it inverts, and the second - inversion follows it instead of deriving an inverse. Inverting can reclassify - arguments as constants, and folding them back rewrapped the Jaspr without - carrying the back-pointer over, so the second inversion fell back to inverting - the body, which is the derivation ``custom_inversion`` exists to avoid. This - affected, for instance, inverting a :func:`prepare ` whose - amplitudes are only known at run time, as produced by a + The custom inverse and the custom controlled version registered by + :func:`custom_inversion ` and + :func:`custom_control ` are now brought into the calling + convention of the function they stand in for at the point where they are + traced, instead of being patched up every time an inversion or a control is + applied. Patching them up late meant rewrapping the Jaspr, and a rewrapped + Jaspr lost the very registrations these decorators exist to provide, so a + second inversion fell back to deriving an inverse from the body. The + signatures are now also checked against each other when the pair is + registered, so a mismatch is reported where it is introduced. This affected, + for instance, inverting a :func:`prepare ` whose amplitudes are + only known at run time, as produced by a :class:`~qrisp.block_encodings.BlockEncoding` simulation with traced coefficients. diff --git a/src/qrisp/environments/custom_control_environment.py b/src/qrisp/environments/custom_control_environment.py index 9ed95929e..873702508 100644 --- a/src/qrisp/environments/custom_control_environment.py +++ b/src/qrisp/environments/custom_control_environment.py @@ -28,7 +28,9 @@ from qrisp.environments.quantum_environments import QuantumEnvironment from qrisp.jasp import ( AbstractQubit, + check_aval_equivalence, check_for_tracing_mode, + closure_convert_jaspr, get_last_equation, make_jaspr, qache, @@ -249,7 +251,9 @@ def adaptive_control_function(*args, **kwargs): # Retrieve the pjit equation jit_eqn = get_last_equation() - if not jit_eqn.params["jaxpr"].ctrl_jaspr: + forward_jaspr = jit_eqn.params["jaxpr"] + + if not forward_jaspr.ctrl_jaspr: # Trace the controlled version # Make sure the inv keyword argument is treated as a static argument @@ -275,8 +279,26 @@ def ammended_func(*ammended_args, **kwargs): controlled_jaspr = make_jaspr(ammended_func, **cusc_kwargs)(*ammended_args, **kwargs) + # The uncontrolled version was traced by qache, i.e. through + # jax.jit, so Jax closure converted whatever the function + # captured into leading invars of forward_jaspr. make_jaspr + # leaves those in constvars/consts instead, so bring the + # controlled version into the same calling convention before + # caching it - it has to stand in for the uncontrolled version + # at a jit equation that supplies [ctrl_qubit] + forward_jaspr's + # arguments. The control qubit is invars[0] (see ammended_args + # above), so the folded invars go right behind it. + controlled_jaspr = closure_convert_jaspr(controlled_jaspr, insert_at=1) + + if not check_aval_equivalence(controlled_jaspr.invars[1:], forward_jaspr.invars): + raise Exception( + f"Custom control of {func.__name__} does not take the same arguments as the " + f"function itself (control: {[var.aval for var in controlled_jaspr.invars[1:]]}, " + f"function: {[var.aval for var in forward_jaspr.invars]})." + ) + # Store controlled version - jit_eqn.params["jaxpr"].ctrl_jaspr = controlled_jaspr + forward_jaspr.ctrl_jaspr = controlled_jaspr return res diff --git a/src/qrisp/environments/custom_inversion_environment.py b/src/qrisp/environments/custom_inversion_environment.py index 5e5bf99fb..066ae98b0 100644 --- a/src/qrisp/environments/custom_inversion_environment.py +++ b/src/qrisp/environments/custom_inversion_environment.py @@ -19,7 +19,9 @@ import jax.numpy as jnp from qrisp.jasp import ( + check_aval_equivalence, check_for_tracing_mode, + closure_convert_jaspr, get_last_equation, make_jaspr, qache, @@ -163,7 +165,9 @@ def adaptive_inversion_function(*args, **kwargs): # Retrieve the pjit equation jit_eqn = get_last_equation() - if not jit_eqn.params["jaxpr"].inv_jaspr: + forward_jaspr = jit_eqn.params["jaxpr"] + + if not forward_jaspr.inv_jaspr: # Trace the inverted version def ammended_func(*args, **kwargs): @@ -172,9 +176,25 @@ def ammended_func(*args, **kwargs): inverted_jaspr = make_jaspr(ammended_func)(*args, **kwargs) - # Store controlled version - jit_eqn.params["jaxpr"].inv_jaspr = inverted_jaspr - inverted_jaspr.inv_jaspr = jit_eqn.params["jaxpr"] + # The forward version was traced by qache, i.e. through jax.jit, + # so Jax closure converted whatever the function captured into + # leading invars of forward_jaspr. make_jaspr leaves those in + # constvars/consts instead, so bring the inverse into the same + # calling convention before caching it - it has to stand in for + # the forward version at a jit equation that supplies exactly + # forward_jaspr's arguments. + inverted_jaspr = closure_convert_jaspr(inverted_jaspr) + + if not check_aval_equivalence(inverted_jaspr.invars, forward_jaspr.invars): + raise Exception( + f"Custom inverse of {func.__name__} does not take the same arguments as the " + f"function itself (inverse: {[var.aval for var in inverted_jaspr.invars]}, " + f"function: {[var.aval for var in forward_jaspr.invars]})." + ) + + # Store inverted version + forward_jaspr.inv_jaspr = inverted_jaspr + inverted_jaspr.inv_jaspr = forward_jaspr return res diff --git a/src/qrisp/jasp/interpreter_tools/abstract_interpreter.py b/src/qrisp/jasp/interpreter_tools/abstract_interpreter.py index 1165a5836..374c151dc 100644 --- a/src/qrisp/jasp/interpreter_tools/abstract_interpreter.py +++ b/src/qrisp/jasp/interpreter_tools/abstract_interpreter.py @@ -199,7 +199,8 @@ def reinterpret(jaxpr: Jaxpr | ClosedJaxpr, eqn_evaluator: Callable = exec_eqn): evaluator = eval_jaxpr(inter_jaxpr, eqn_evaluator=eqn_evaluator) jaxpr_input_avals = [var.aval for var in inter_jaxpr.constvars + inter_jaxpr.invars] - res = make_jaxpr(evaluator)(*jaxpr_input_avals).jaxpr + retraced = make_jaxpr(evaluator)(*jaxpr_input_avals) + res = retraced.jaxpr res.constvars.extend(res.invars[: len(inter_jaxpr.constvars)]) temp = list(res.invars[len(inter_jaxpr.constvars) :]) @@ -207,7 +208,12 @@ def reinterpret(jaxpr: Jaxpr | ClosedJaxpr, eqn_evaluator: Callable = exec_eqn): res.invars.extend(temp) if isinstance(jaxpr, ClosedJaxpr): - res = ClosedJaxpr(res, jaxpr.consts) + # The retrace can hoist values that it closes over into constvars of its + # own. Those sit in front of the input's own constvars (which were + # appended behind them just above), so their consts have to go in front + # here as well - otherwise the result is a ClosedJaxpr whose consts no + # longer line up with its constvars. + res = ClosedJaxpr(res, list(retraced.consts) + list(jaxpr.consts)) return res diff --git a/src/qrisp/jasp/jasp_expression/centerclass.py b/src/qrisp/jasp/jasp_expression/centerclass.py index ac9768573..289a53e92 100644 --- a/src/qrisp/jasp/jasp_expression/centerclass.py +++ b/src/qrisp/jasp/jasp_expression/centerclass.py @@ -1523,8 +1523,12 @@ def amended_function(*args, **kwargs): return jaspr_creator -def check_aval_equivalence(invars_1, invars_2) -> bool: - """Return True if every paired invar has the same abstract-value type.""" +def check_aval_equivalence(invars_1: list[Var], invars_2: list[Var]) -> bool: + """Return True if both signatures have the same length and every paired + invar has the same abstract-value type. + """ + if len(invars_1) != len(invars_2): + return False return all(type(v1.aval) is type(v2.aval) for v1, v2 in zip(invars_1, invars_2)) diff --git a/src/qrisp/jasp/jasp_expression/control_transform.py b/src/qrisp/jasp/jasp_expression/control_transform.py index 5c615689e..39297424a 100644 --- a/src/qrisp/jasp/jasp_expression/control_transform.py +++ b/src/qrisp/jasp/jasp_expression/control_transform.py @@ -16,6 +16,8 @@ """Implements ControlledJaspr and the transformations that add quantum control to Jaspr equations.""" +from collections.abc import Callable + import numpy as np from jax.extend.core import JaxprEqn, Var @@ -40,7 +42,7 @@ class ControlledJaspr(Jaspr): __slots__ = ("base_jaspr", "ctrl_state") - def __init__(self, base_jaspr, ctrl_state, stop_recursion=False): + def __init__(self, base_jaspr: Jaspr, ctrl_state: int | str, stop_recursion: bool = False) -> None: self.base_jaspr = base_jaspr self.ctrl_state = str(ctrl_state) @@ -56,7 +58,7 @@ def __init__(self, base_jaspr, ctrl_state, stop_recursion=False): if self.base_jaspr.inv_jaspr and not stop_recursion: self.inv_jaspr = ControlledJaspr(base_jaspr.inv_jaspr, ctrl_state, stop_recursion=True) - def control(self, num_ctrl, ctrl_state=-1): + def control(self, num_ctrl: int, ctrl_state: int | str = -1) -> "ControlledJaspr": if isinstance(ctrl_state, int): if ctrl_state < 0: @@ -68,20 +70,20 @@ def control(self, num_ctrl, ctrl_state=-1): return ControlledJaspr.from_cache(self.base_jaspr, ctrl_state + self.ctrl_state) - def inverse(self): + def inverse(self) -> "ControlledJaspr": return ControlledJaspr.from_cache(self.base_jaspr.inverse(), self.ctrl_state) # LRU cache controlled by QRISP_COMPILATION_CACHE_SIZE env var @classmethod @qrisp_lru_compilation_cache - def from_cache(cls, base_jaspr, ctrl_state): + def from_cache(cls, base_jaspr: Jaspr, ctrl_state: str) -> "ControlledJaspr": return ControlledJaspr(base_jaspr, ctrl_state) control_var_count = np.zeros(1) -def control_eqn(eqn, ctrl_qubit_var): +def control_eqn(eqn: JaxprEqn, ctrl_qubit_var: Var) -> JaxprEqn: """Receives and equation that describes either an operation or a pjit primitive and returns an equation that describes the inverse. @@ -103,30 +105,14 @@ def control_eqn(eqn, ctrl_qubit_var): invars = list(eqn.invars) if isinstance(eqn.params["jaxpr"], Jaspr): - orig_jaxpr = eqn.params["jaxpr"] - controlled_jaxpr = orig_jaxpr.control(1) - - # Jaspr.control() may retrieve a pre-cached ctrl_jaspr (populated - # by the custom_control decorator's own retrace via make_jaspr), - # or build one via multi_control_jaspr. Either way, values that - # are only used inside nested jit/cond/while sub-equations (e.g. - # arrays closed over by q_switch case functions) can end up - # reclassified as unexpected constvars during that retrace. Fold - # any such constvars back into genuine invars so the wrapping - # equation's invars line up with the controlled callee's real - # invars, see fold_extra_constvars_into_invars. - from qrisp.jasp.jasp_expression.inv_transform import fold_extra_constvars_into_invars - - # controlled_jaxpr.invars starts with the newly added control - # qubit (see custom_control_environment.ammended_func / - # multi_control_jaspr), so any newly introduced constvars must be - # folded back in right after it to line up with the wrapping - # equation's [ctrl_qubit_var] + eqn.invars ordering. - normalized = fold_extra_constvars_into_invars(controlled_jaxpr, len(orig_jaxpr.constvars), insert_at=1) - if normalized is not controlled_jaxpr: - controlled_jaxpr = Jaspr(normalized) - - new_params["jaxpr"] = controlled_jaxpr + # The controlled version takes [ctrl_qubit] + the arguments of the + # equation it replaces: a cached ctrl_jaspr was brought into pjit's + # calling convention when custom_control created it (see + # closure_convert_jaspr), and a derived one is built by + # multi_control_jaspr from this Jaspr's own signature. Not + # rewrapping here is what lets a ControlledJaspr stay one, keeping + # its efficient nested control and its custom inverse. + new_params["jaxpr"] = eqn.params["jaxpr"].control(1) new_params["name"] = "c" + new_params["name"] invars = [ctrl_qubit_var] + eqn.invars @@ -227,7 +213,7 @@ def control_eqn(eqn, ctrl_qubit_var): # LRU cache controlled by QRISP_COMPILATION_CACHE_SIZE env var @qrisp_lru_compilation_cache -def control_jaspr(jaspr): +def control_jaspr(jaspr: Jaspr) -> Jaspr: """Takes a Jaspr and returns a Jaspr that has an additional Qubit argument (located behind the QuantumState argument). The returned Jaspr is controlled on that Qubit argument. @@ -276,7 +262,7 @@ def control_jaspr(jaspr): ) -def multi_control_jaspr(jaspr, num_ctrl, ctrl_state): +def multi_control_jaspr(jaspr: Jaspr, num_ctrl: int, ctrl_state: str) -> Jaspr: """Similar to control_jaspr but allows specification of more than one control and a control state @@ -311,7 +297,7 @@ def multi_control_jaspr(jaspr, num_ctrl, ctrl_state): ) -def exec_multi_controlled_jaspr(jaspr, num_ctrls, ctrl_state): +def exec_multi_controlled_jaspr(jaspr: Jaspr, num_ctrls: int, ctrl_state: str) -> Callable: def multi_controlled_jaspr_executor(*args): diff --git a/src/qrisp/jasp/jasp_expression/inv_transform.py b/src/qrisp/jasp/jasp_expression/inv_transform.py index 14aa3c32a..6774844b2 100644 --- a/src/qrisp/jasp/jasp_expression/inv_transform.py +++ b/src/qrisp/jasp/jasp_expression/inv_transform.py @@ -16,83 +16,29 @@ """Implements Jaspr/equation inversion (daggering), including while-loop inversion for jrange loops.""" +from typing import TYPE_CHECKING + import numpy as np from jax import make_jaxpr -from jax.extend.core import ClosedJaxpr, Jaxpr, JaxprEqn, Var +from jax.extend.core import ClosedJaxpr, JaxprEqn, Var from jax.lax import add_p, sub_p from sympy import lambdify from qrisp._cache_config import qrisp_lru_compilation_cache from qrisp.jasp.interpreter_tools import copy_jaxpr_eqn, extract_invalues, insert_outvalues, reinterpret -from qrisp.jasp.jasp_expression.jaxpr_utils import rebuild_closed_jaxpr +from qrisp.jasp.jasp_expression.jaxpr_utils import ( + fold_extra_constvars_into_invars, + rebuild_closed_jaxpr, +) from qrisp.jasp.primitives import AbstractQuantumState, greek_letters, quantum_gate_p -qc_var_count = np.zeros(1, dtype=np.int64) - - -def fold_extra_constvars_into_invars(closed_jaxpr, n_expected_const, insert_at=0): - """Fold unexpected leading constvars of a (re)traced ClosedJaxpr back into - genuine invars. - - Retracing a Jaspr (e.g. during inversion) evaluates nested jit/cond/while - sub-equations through Jax's native primitive binding. As a side effect, - values that are only used inside such nested sub-equations (e.g. arrays - closed over by q_switch case functions, see prepare_qswitch) can get - reclassified from regular invars into self-contained constvars/consts of - the retraced jaxpr - Jax "closure converts" them relative to the retrace, - even though they were explicitly passed in as invars. If left as consts, - the resulting Jaspr would carry live tracers baked into ``consts``, which - breaks as soon as the Jaspr is reused in a different tracing context - (e.g. via LRU/custom-inversion caching or a subsequent re-trace such as - terminal_sampling's own sampling pass), producing either an arity - mismatch or a "leaked tracer" error. - - This function folds any such newly introduced constvars back into being - genuine invars (in the same order in which they were introduced, which - corresponds to a prefix of the original invars). - - Parameters - ---------- - closed_jaxpr : jax.extend.core.ClosedJaxpr - The (re)traced ClosedJaxpr to normalize. - n_expected_const : int - The number of constvars that were already present/expected before - the retrace (usually 0). - insert_at : int, optional - The position (within the resulting invars list) at which to - re-insert the folded constvars. Use this when the retrace prepends - its own leading invars (e.g. a control qubit) that must stay in - front of the folded ones so the wrapping equation's argument order - matches. Default is 0 (folded invars go first). +if TYPE_CHECKING: + from qrisp.jasp.jasp_expression.centerclass import Jaspr - Returns - ------- - jax.extend.core.ClosedJaxpr - A ClosedJaxpr with at most ``n_expected_const`` constvars. - - """ - core = closed_jaxpr.jaxpr - n_extra = len(core.constvars) - n_expected_const - if n_extra <= 0: - return closed_jaxpr - - folded_invars = list(core.constvars[:n_extra]) - new_constvars = list(core.constvars[n_extra:]) - new_consts = list(closed_jaxpr.consts[n_extra:]) - new_invars = list(core.invars) - new_invars[insert_at:insert_at] = folded_invars - new_core = Jaxpr( - constvars=new_constvars, - invars=new_invars, - outvars=list(core.outvars), - eqns=list(core.eqns), - effects=core.effects, - debug_info=core.debug_info, - ) - return ClosedJaxpr(new_core, new_consts) +qc_var_count = np.zeros(1, dtype=np.int64) -def invert_eqn(eqn): +def invert_eqn(eqn: JaxprEqn) -> JaxprEqn: """Receives and equation that describes either an operation or a pjit primitive and returns an equation that describes the inverse. @@ -109,31 +55,12 @@ def invert_eqn(eqn): """ if eqn.primitive.name == "jit": params = dict(eqn.params) - orig_jaxpr = eqn.params["jaxpr"] - inv_jaxpr = orig_jaxpr.inverse() - - # The inverted Jaspr may have been produced by a generic retrace or - # returned from a custom_inversion cache (inv_jaspr); either way, it - # must expose exactly the same number of (non-const) invars as - # `eqn.invars` supplies. Normalize away any unexpectedly introduced - # constvars, see fold_extra_constvars_into_invars. - from qrisp.jasp import Jaspr - - normalized = fold_extra_constvars_into_invars(inv_jaxpr, len(orig_jaxpr.constvars)) - if normalized is not inv_jaxpr: - # Wrapping the normalized jaxpr creates a fresh Jaspr, which starts out - # without the inv_jaspr back-pointer custom_inversion registered on the - # one being replaced. Carry it over: dropping it makes a second - # inversion fall back to inverting the body structurally, which for a - # custom_inversion user is exactly the derivation that does not apply. - # Folding restores the reclassified constvars to invars, so the - # normalized Jaspr's signature matches the wrapping equation and - # the original Jaspr referenced by the back-pointer. - preserved_inv_jaspr = inv_jaxpr.inv_jaspr - inv_jaxpr = Jaspr(normalized) - inv_jaxpr.inv_jaspr = preserved_inv_jaspr - - params["jaxpr"] = inv_jaxpr + + # The inverse takes the same arguments as the equation it replaces: a + # cached inv_jaspr was brought into pjit's calling convention when + # custom_inversion created it (see closure_convert_jaspr), and a derived + # one is built from this Jaspr's own signature. + params["jaxpr"] = eqn.params["jaxpr"].inverse() name = params["name"] if name[-3:] == "_dg": @@ -180,7 +107,7 @@ def invert_eqn(eqn): # LRU cache controlled by QRISP_COMPILATION_CACHE_SIZE env var @qrisp_lru_compilation_cache -def invert_jaspr(jaspr): +def invert_jaspr(jaspr: "Jaspr") -> "Jaspr": """Takes a Jaspr and returns a Jaspr, which performs the inverted quantum operation Parameters @@ -272,6 +199,10 @@ def eqn_evaluator(eqn, context_dic): processed_jaxpr = reinterpret(temp_jaxpr, eqn_evaluator) + # The retrace above can hoist values into constvars of its own. Keep the + # derived inverse in the same calling convention as everything else by + # folding those back into invars, leaving only the Jaspr's genuine + # constvars behind. processed_jaxpr = fold_extra_constvars_into_invars(processed_jaxpr, len(jaspr.constvars)) from qrisp.jasp import Jaspr @@ -302,7 +233,7 @@ def eqn_evaluator(eqn, context_dic): # the comparison to determine the break condition needs to be adjusted. -def invert_loop_body(jaxpr): +def invert_loop_body(jaxpr: ClosedJaxpr) -> "Jaspr": # This function treats the loop body # This list will contain the equations with the loop index decrementation @@ -376,7 +307,7 @@ def invert_loop_body(jaxpr): # This function performs the above mentioned step 2 to treat the loop primitive. -def invert_loop_eqn(eqn): +def invert_loop_eqn(eqn: JaxprEqn) -> JaxprEqn: # Process the loop body body_jaxpr = eqn.params["body_jaxpr"] diff --git a/src/qrisp/jasp/jasp_expression/jaxpr_utils.py b/src/qrisp/jasp/jasp_expression/jaxpr_utils.py index 5e9b049d9..9061f5f8c 100644 --- a/src/qrisp/jasp/jasp_expression/jaxpr_utils.py +++ b/src/qrisp/jasp/jasp_expression/jaxpr_utils.py @@ -92,3 +92,126 @@ def rebuild_closed_jaxpr(base: "ClosedJaxpr | Jaspr", *, eqns=None, outvars=None """ return ClosedJaxpr(rebuild_jaxpr(base.jaxpr, eqns=eqns, outvars=outvars), base.consts) + + +def fold_extra_constvars_into_invars( + closed_jaxpr: ClosedJaxpr, n_expected_const: int, insert_at: int = 0 +) -> ClosedJaxpr: + """Fold the leading constvars of a ClosedJaxpr back into genuine invars. + + Qrisp builds the very same quantum function in two ways, and the two disagree + on how closed-over values are passed. + + ``qache`` traces the forward version through ``jax.jit``, so Jax performs its + own closure conversion: values that the traced function closes over (e.g. the + angle arrays captured by ``prepare_qswitch``'s case functions) become + *leading invars* of the callee, and the wrapping ``jit`` equation supplies + them as extra leading operands. + + ``make_jaspr`` traces through ``jax.make_jaxpr``, which does not closure + convert: those same values stay behind in ``constvars``/``consts``. A Jaspr + traced that way holds live tracers in ``consts`` - which breaks as soon as it + is reused in a different tracing context - and exposes fewer invars than the + ``jit`` equation supplies. + + This function rewrites the second shape into the first, so that a Jaspr + produced by ``make_jaspr`` can stand in as the callee of a ``jit`` equation + that was built for a ``qache``-produced one. + + Note that the *values* of the folded constvars (``consts[:n_extra]``) are + dropped: the folded vars become plain invars, and it is the caller's + responsibility to supply equivalent values positionally. That holds precisely + because Jax's own closure conversion put the same values in the same leading + positions on the forward path - see ``closure_convert_jaspr``, which pairs + this rewrite with a signature check against the forward Jaspr. + + Parameters + ---------- + closed_jaxpr : jax.extend.core.ClosedJaxpr + The ClosedJaxpr to normalize. + n_expected_const : int + The number of constvars that are genuine and must be left alone + (usually 0). + insert_at : int, optional + The position (within the resulting invars list) at which to re-insert the + folded constvars. Use this when the signature starts with invars that + must stay in front of the folded ones (e.g. the control qubit that + ``custom_control`` prepends). Default is 0 (folded invars go first). + + Returns + ------- + jax.extend.core.ClosedJaxpr + A ClosedJaxpr with exactly ``n_expected_const`` constvars. The input is + returned unchanged if there was nothing to fold. + + """ + core = closed_jaxpr.jaxpr + n_extra = len(core.constvars) - n_expected_const + if n_extra <= 0: + return closed_jaxpr + + folded_invars = list(core.constvars[:n_extra]) + new_constvars = list(core.constvars[n_extra:]) + new_consts = list(closed_jaxpr.consts[n_extra:]) + new_invars = list(core.invars) + new_invars[insert_at:insert_at] = folded_invars + new_core = Jaxpr( + constvars=new_constvars, + invars=new_invars, + outvars=list(core.outvars), + eqns=list(core.eqns), + effects=core.effects, + debug_info=core.debug_info, + ) + return ClosedJaxpr(new_core, new_consts) + + +def closure_convert_jaspr(jaspr: "Jaspr", insert_at: int = 0) -> "Jaspr": + """Return ``jaspr`` in the calling convention that pjit uses for its callees. + + Every constvar is folded into an invar (see + ``fold_extra_constvars_into_invars``), so the result takes all of its inputs + through its signature and carries no ``consts``. Apply this to any Jaspr that + was traced with ``make_jaspr`` but is going to be used as the callee of a + ``jit`` equation, i.e. the cached ``inv_jaspr``/``ctrl_jaspr`` of the + ``custom_inversion``/``custom_control`` decorators. + + Doing this once, where the Jaspr is created, is what keeps the transformation + passes (``invert_eqn``, ``control_eqn``) free of rewrapping: a rewrap there + would have to reattach every piece of Jaspr metadata by hand, and could not + preserve a ``ControlledJaspr`` at all. + + Parameters + ---------- + jaspr : Jaspr + The Jaspr to convert. + insert_at : int, optional + Where to insert the folded invars, see + ``fold_extra_constvars_into_invars``. Default is 0. + + Returns + ------- + Jaspr + The converted Jaspr, or ``jaspr`` itself if there was nothing to fold. + + """ + from qrisp.jasp.jasp_expression.centerclass import Jaspr + + normalized = fold_extra_constvars_into_invars(jaspr, 0, insert_at=insert_at) + if normalized is jaspr: + return jaspr + + # Rewrapping produces a fresh Jaspr, so every attribute that is not part of + # the Jaxpr itself has to be carried over explicitly. The variable objects + # are shared with the input, so the permeability dict keyed by them still + # applies. + res = Jaspr( + normalized, + permeability=jaspr.permeability, + isqfree=jaspr.isqfree, + ctrl_jaspr=jaspr.ctrl_jaspr, + inv_jaspr=jaspr.inv_jaspr, + ) + res.envs_flattened = jaspr.envs_flattened + + return res diff --git a/tests/jax_tests/test_custom_inverse.py b/tests/jax_tests/test_custom_inverse.py index b5e08b508..fe51ed190 100644 --- a/tests/jax_tests/test_custom_inverse.py +++ b/tests/jax_tests/test_custom_inverse.py @@ -15,7 +15,10 @@ ******************************************************************************** """ +from collections.abc import Iterator + import jax.numpy as jnp +from jax.extend.core import Jaxpr, JaxprEqn from qrisp import * from qrisp.alg_primitives.state_preparation import prepare @@ -302,3 +305,108 @@ def body(): assert "measure" in gate_names(1), "a single inversion must uncompute via measurement" assert "measure" not in forward assert gate_names(2) == forward + + +def _walk_jit_eqns(jaxpr: Jaxpr, seen: set[int] | None = None) -> Iterator[JaxprEqn]: + """Yield every jit equation reachable from jaxpr, including nested ones.""" + if seen is None: + seen = set() + if id(jaxpr) in seen: + return + seen.add(id(jaxpr)) + + for eqn in jaxpr.eqns: + if eqn.primitive.name == "jit": + yield eqn + yield from _walk_jit_eqns(eqn.params["jaxpr"].jaxpr, seen) + elif eqn.primitive.name == "while": + yield from _walk_jit_eqns(eqn.params["body_jaxpr"].jaxpr, seen) + yield from _walk_jit_eqns(eqn.params["cond_jaxpr"].jaxpr, seen) + elif eqn.primitive.name == "cond": + for branch in eqn.params["branches"]: + yield from _walk_jit_eqns(branch.jaxpr, seen) + + +def test_registered_inverse_matches_forward_signature(): + """The cached custom inverse must be usable as-is at the forward equation. + + custom_inversion traces the inverse with make_jaspr, which leaves closed-over + values in constvars/consts, while the forward version goes through qache, + i.e. jax.jit, which closure converts them into leading invars. The inverse is + brought into that same convention where it is registered, so that applying an + inversion is a plain swap of the callee and never has to rewrap (and thereby + lose) the Jaspr. This pins that: every registered inverse takes exactly the + arguments of the function it inverts and carries no consts of its own. + + prepare with run-time amplitudes is the case that has closed-over values to + begin with: its q_switch case functions capture the angle arrays. + """ + + def main(scale): + qv = QuantumFloat(2) + weights = jnp.arange(1, 5, dtype=float) * scale + prepare(qv, weights / jnp.linalg.norm(weights)) + return qv + + jaspr = make_jaspr(main)(1.0) + + checked = 0 + for eqn in _walk_jit_eqns(jaspr.jaxpr): + callee = eqn.params["jaxpr"] + inv_jaspr = getattr(callee, "inv_jaspr", None) + if inv_jaspr is None: + continue + checked += 1 + + assert not inv_jaspr.constvars, "a registered inverse must not carry constvars" + assert not inv_jaspr.consts, "a registered inverse must not carry consts" + assert check_aval_equivalence(inv_jaspr.invars, callee.invars), ( + f"registered inverse takes {[var.aval for var in inv_jaspr.invars]}, " + f"but the function it inverts takes {[var.aval for var in callee.invars]}" + ) + + assert checked, "found no custom_inversion function to check" + + +def test_double_inversion_of_prepare_under_control_round_trips(): + """Two inversions must cancel when they are applied to a controlled equation. + + Controlling an equation swaps in the controlled callee the same way inverting + swaps in the inverted one, and neither may rewrap the Jaspr: for a + custom_control/custom_inversion function such as q_switch, a rewrap drops both + registrations at once. Here the control sits outside both inversions, so the + outer inversion is applied to an already controlled equation. + """ + size = 4 + + def amplitudes(scale): + weights = jnp.arange(1, size + 1, dtype=float) * scale + return weights / jnp.linalg.norm(weights) + + @terminal_sampling + def main(scale): + condition = QuantumBool() + condition.flip() + qv = QuantumFloat(2) + with invert(): + with control(condition[0]): + with invert(): + with invert(): + prepare(qv, amplitudes(scale)) + return qv + + @terminal_sampling + def reference(scale): + qv = QuantumFloat(2) + with invert(): + prepare(qv, amplitudes(scale)) + return qv + + # The control qubit is held at |1>, so the controlled body acts + # unconditionally and the inner pair of inversions cancels, leaving the + # single outer inversion. + res = main(1.0) + expected = reference(1.0) + + for state in range(size): + assert abs(expected.get(state, 0) - res.get(state, 0)) < 1e-4, f"{expected} vs {res}"