Skip to content
Open
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
19 changes: 12 additions & 7 deletions documentation/source/general/changelog/changelog-dev.rst
Original file line number Diff line number Diff line change
Expand Up @@ -110,13 +110,18 @@ Bug Fixes
* Fixed a failure when a function decorated with
:func:`custom_inversion <qrisp.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 <qrisp.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 <qrisp.custom_inversion>` and
:func:`custom_control <qrisp.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 <qrisp.prepare>` whose amplitudes are
only known at run time, as produced by a
:class:`~qrisp.block_encodings.BlockEncoding` simulation with traced
coefficients.

Expand Down
26 changes: 24 additions & 2 deletions src/qrisp/environments/custom_control_environment.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand All @@ -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

Expand Down
28 changes: 24 additions & 4 deletions src/qrisp/environments/custom_inversion_environment.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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):
Expand All @@ -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

Expand Down
10 changes: 8 additions & 2 deletions src/qrisp/jasp/interpreter_tools/abstract_interpreter.py
Original file line number Diff line number Diff line change
Expand Up @@ -199,15 +199,21 @@ 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) :])
res.invars.clear()
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))
Comment on lines +211 to +216

return res

Expand Down
8 changes: 6 additions & 2 deletions src/qrisp/jasp/jasp_expression/centerclass.py
Original file line number Diff line number Diff line change
Expand Up @@ -1523,8 +1523,12 @@
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

Check warning on line 1527 in src/qrisp/jasp/jasp_expression/centerclass.py

View workflow job for this annotation

GitHub Actions / ruff lint (reviewdog)

[rdjson] reported by reviewdog 🐶 1 blank line required between summary line and description Raw Output: message:"1 blank line required between summary line and description" location:{path:"/home/runner/work/Qrisp/Qrisp/src/qrisp/jasp/jasp_expression/centerclass.py" range:{start:{line:1527 column:5} end:{line:1529 column:8}}} severity:WARNING source:{name:"ruff" url:"https://docs.astral.sh/ruff"} code:{value:"D205" url:"https://docs.astral.sh/ruff/rules/missing-blank-line-after-summary"}
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))


Expand Down
50 changes: 18 additions & 32 deletions src/qrisp/jasp/jasp_expression/control_transform.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -40,7 +42,7 @@

__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:

Check warning on line 45 in src/qrisp/jasp/jasp_expression/control_transform.py

View workflow job for this annotation

GitHub Actions / ruff lint (reviewdog)

[rdjson] reported by reviewdog 🐶 Missing docstring in `__init__` Raw Output: message:"Missing docstring in `__init__`" location:{path:"/home/runner/work/Qrisp/Qrisp/src/qrisp/jasp/jasp_expression/control_transform.py" range:{start:{line:45 column:9} end:{line:45 column:17}}} severity:WARNING source:{name:"ruff" url:"https://docs.astral.sh/ruff"} code:{value:"D107" url:"https://docs.astral.sh/ruff/rules/undocumented-public-init"}

self.base_jaspr = base_jaspr
self.ctrl_state = str(ctrl_state)
Expand All @@ -56,7 +58,7 @@
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":

Check warning on line 61 in src/qrisp/jasp/jasp_expression/control_transform.py

View workflow job for this annotation

GitHub Actions / ruff lint (reviewdog)

[rdjson] reported by reviewdog 🐶 Missing docstring in public method Raw Output: message:"Missing docstring in public method" location:{path:"/home/runner/work/Qrisp/Qrisp/src/qrisp/jasp/jasp_expression/control_transform.py" range:{start:{line:61 column:9} end:{line:61 column:16}}} severity:WARNING source:{name:"ruff" url:"https://docs.astral.sh/ruff"} code:{value:"D102" url:"https://docs.astral.sh/ruff/rules/undocumented-public-method"}

if isinstance(ctrl_state, int):
if ctrl_state < 0:
Expand All @@ -68,20 +70,20 @@

return ControlledJaspr.from_cache(self.base_jaspr, ctrl_state + self.ctrl_state)

def inverse(self):
def inverse(self) -> "ControlledJaspr":

Check warning on line 73 in src/qrisp/jasp/jasp_expression/control_transform.py

View workflow job for this annotation

GitHub Actions / ruff lint (reviewdog)

[rdjson] reported by reviewdog 🐶 Missing docstring in public method Raw Output: message:"Missing docstring in public method" location:{path:"/home/runner/work/Qrisp/Qrisp/src/qrisp/jasp/jasp_expression/control_transform.py" range:{start:{line:73 column:9} end:{line:73 column:16}}} severity:WARNING source:{name:"ruff" url:"https://docs.astral.sh/ruff"} code:{value:"D102" url:"https://docs.astral.sh/ruff/rules/undocumented-public-method"}
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":

Check warning on line 79 in src/qrisp/jasp/jasp_expression/control_transform.py

View workflow job for this annotation

GitHub Actions / ruff lint (reviewdog)

[rdjson] reported by reviewdog 🐶 Missing docstring in public method Raw Output: message:"Missing docstring in public method" location:{path:"/home/runner/work/Qrisp/Qrisp/src/qrisp/jasp/jasp_expression/control_transform.py" range:{start:{line:79 column:9} end:{line:79 column:19}}} severity:WARNING source:{name:"ruff" url:"https://docs.astral.sh/ruff"} code:{value:"D102" url:"https://docs.astral.sh/ruff/rules/undocumented-public-method"}
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:

Check warning on line 86 in src/qrisp/jasp/jasp_expression/control_transform.py

View workflow job for this annotation

GitHub Actions / ruff lint (reviewdog)

[rdjson] reported by reviewdog 🐶 Missing argument description in the docstring for `control_eqn`: `ctrl_qubit_var` Raw Output: message:"Missing argument description in the docstring for `control_eqn`: `ctrl_qubit_var`" location:{path:"/home/runner/work/Qrisp/Qrisp/src/qrisp/jasp/jasp_expression/control_transform.py" range:{start:{line:86 column:5} end:{line:86 column:16}}} severity:WARNING source:{name:"ruff" url:"https://docs.astral.sh/ruff"} code:{value:"D417" url:"https://docs.astral.sh/ruff/rules/undocumented-param"}
"""Receives and equation that describes either an operation or a pjit primitive
and returns an equation that describes the inverse.

Expand All @@ -103,30 +105,14 @@

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
Expand Down Expand Up @@ -227,7 +213,7 @@

# LRU cache controlled by QRISP_COMPILATION_CACHE_SIZE env var
@qrisp_lru_compilation_cache
def control_jaspr(jaspr):
def control_jaspr(jaspr: Jaspr) -> Jaspr:

Check warning on line 216 in src/qrisp/jasp/jasp_expression/control_transform.py

View workflow job for this annotation

GitHub Actions / ruff lint (reviewdog)

[rdjson] reported by reviewdog 🐶 Missing argument description in the docstring for `control_jaspr`: `jaspr` Raw Output: message:"Missing argument description in the docstring for `control_jaspr`: `jaspr`" location:{path:"/home/runner/work/Qrisp/Qrisp/src/qrisp/jasp/jasp_expression/control_transform.py" range:{start:{line:216 column:5} end:{line:216 column:18}}} severity:WARNING source:{name:"ruff" url:"https://docs.astral.sh/ruff"} code:{value:"D417" url:"https://docs.astral.sh/ruff/rules/undocumented-param"}
"""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.
Expand Down Expand Up @@ -276,7 +262,7 @@
)


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

Expand Down Expand Up @@ -311,7 +297,7 @@
)


def exec_multi_controlled_jaspr(jaspr, num_ctrls, ctrl_state):
def exec_multi_controlled_jaspr(jaspr: Jaspr, num_ctrls: int, ctrl_state: str) -> Callable:

Check warning on line 300 in src/qrisp/jasp/jasp_expression/control_transform.py

View workflow job for this annotation

GitHub Actions / ruff lint (reviewdog)

[rdjson] reported by reviewdog 🐶 Missing docstring in public function Raw Output: message:"Missing docstring in public function" location:{path:"/home/runner/work/Qrisp/Qrisp/src/qrisp/jasp/jasp_expression/control_transform.py" range:{start:{line:300 column:5} end:{line:300 column:32}}} severity:WARNING source:{name:"ruff" url:"https://docs.astral.sh/ruff"} code:{value:"D103" url:"https://docs.astral.sh/ruff/rules/undocumented-public-function"}

def multi_controlled_jaspr_executor(*args):

Expand Down
Loading
Loading