Skip to content
Merged
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
10 changes: 9 additions & 1 deletion docs/source/simulation_architecture.rst
Original file line number Diff line number Diff line change
Expand Up @@ -216,7 +216,15 @@ Expansion does several things at once:
*not* resolved to a number. It is wrapped in a ``ParametricGate`` that, given
:math:`\theta`, constructs the gate matrix. This keeps gate construction
inside the traced/differentiated graph, which is what makes ``jax.grad`` with
respect to gate angles work.
respect to gate angles work. An angle may be any Quil arithmetic expression over memory
references -- ``+ - * / ^``, the functions ``SIN``, ``COS``, ``SQRT``, ``EXP`` and ``CIS``,
and real or complex literals -- so ``RX(theta[0]/2 + pi)``, the form quilc emits when it
compiles a parametric program, and ``RZ(2*SIN(phi[1]))`` both simulate directly. Gates whose
expressions have the same *shape* (``SIN(theta[0])`` and ``SIN(theta[1])``) are still built in
one vectorised batch. A complex-valued argument (``CIS``, a complex literal) is accepted only
by a ``DEFGATE`` gate, since the built-in gates take real angles; everything else is evaluated
in real arithmetic, so ``SQRT`` of a negative number or a fractional power of one gives ``nan``
where Quil would give a complex result.

* **DEFCIRCUIT and cycle expansion.** ``DEFCIRCUIT`` bodies are expanded with
formal-argument substitution. When a circuit invocation matches a
Expand Down
214 changes: 172 additions & 42 deletions pyquil/simulation/_resolver.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,20 @@
NoiseModelLike,
)
from pyquil.quil import Program
from pyquil.quilatom import MemoryReference, Qubit, _contained_mrefs, substitute
from pyquil.quilatom import (
Add,
BinaryExp,
Div,
Function,
MemoryReference,
Mul,
Parameter,
Pow,
Qubit,
Sub,
_contained_mrefs,
substitute,
)
from pyquil.quilbase import (
AbstractInstruction,
ArithmeticBinaryOp,
Expand Down Expand Up @@ -85,31 +98,146 @@
ParameterRef: TypeAlias = tuple[str, int]


#: The Quil arithmetic functions, as JAX functions. ``CIS(x)`` is ``exp(i x)``.
_FUNCTIONS: dict[str, Callable[[Array], Array]] = {
"SIN": jnp.sin,
"COS": jnp.cos,
"SQRT": jnp.sqrt,
"EXP": jnp.exp,
"CIS": lambda x: jnp.exp(1j * x),
}
_BINARY_OPERATORS: dict[type[BinaryExp], Callable[[Array, Array], Array]] = {
Add: jnp.add,
Sub: jnp.subtract,
Mul: jnp.multiply,
Div: jnp.divide,
Pow: jnp.power,
}


@dataclass(frozen=True, slots=True)
class ParameterExpression:
"""A gate argument given as a Quil arithmetic expression over memory references.

Any expression Quil allows is supported: ``+ - * / ^``, the functions ``SIN``, ``COS``,
``SQRT``, ``EXP`` and ``CIS``, and real or complex literals -- for example
``RX(theta[0]/2 + pi)``, ``RZ(2*SIN(phi[1]))``, or ``CPH(CIS(theta[0])) 0`` for a
``DEFGATE CPH(%z)`` taking a complex parameter. Evaluation goes through JAX, so an
expression can be jitted and differentiated through.

There are two ways in. :attr:`evaluate` takes the *narrowed* vector of only the values this
expression reads, which is what the simulator gathers once per vectorised batch; calling the
instance takes the whole circuit-wide parameter vector and gathers from it first.

:param slot_indices: Slots of the parameter vector the expression reads, in order of first
appearance.
:param key: The expression with its references replaced by ``%0``, ``%1``, ... in that
order. Two expressions with the same key have the same shape and constants and differ
only in which slots they read, which is what lets the simulator evaluate them together
in one vectorised operation.
:param is_complex: Whether the value may be complex (the expression contains ``CIS`` or a
complex literal). Everything else is evaluated in real arithmetic; note that Quil
would evaluate ``SQRT`` of a negative number or a fractional power of one as complex,
which real arithmetic reports as ``nan``.
:param evaluate: Evaluates the expression from the values of :attr:`slot_indices`, in that
order.
"""

slot_indices: tuple[int, ...]
key: str
is_complex: bool
evaluate: Callable[[Array], Array]

def __call__(self, params: Array) -> Array:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

You have params here and values above (in evaluate). In the former case you have circuit wide array of arguments; in the latter case you have a narrowed array of arguments. I think this could use clearer semantics - e.g. global or narrowed arguments (rather than params).

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Agreed

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Done: the narrowed vector is slot_values wherever it appears (ParameterExpression.evaluate, _GateBatch.single) and the two docstrings now say which is which — evaluate takes only the slots the expression reads, __call__ takes the circuit-wide vector and gathers from it. Kept params for the global one since that is what it is called everywhere else in _resolver/_simulator; say the word if you would rather it were renamed throughout.

"""Evaluate the expression from the whole circuit-wide parameter vector.

Gathers :attr:`slot_indices` out of ``params`` and hands the narrowed vector to
:attr:`evaluate`.
"""
return self.evaluate(params[jnp.asarray(self.slot_indices)])


def _literal_value(expression: Any) -> float | complex:
"""Evaluate a parameter-free expression (a number, ``pi/2``, ``SIN(pi/4)``, ``1i``) to a scalar."""
value = complex(np.asarray(expression, dtype=complex).item()) if not isinstance(expression, complex) else expression
return value.real if value.imag == 0 else value


def _build_expression(node: Any, references: list[MemoryReference]) -> tuple[Callable[[Array], Array], str, bool]:
"""Compile one node of a Quil expression tree.

:param node: The node to compile; recursion handles its operands.
:param references: The memory references seen so far, in first-appearance order. New ones are
appended, and a reference's position here is the index into the narrowed value vector that
the compiled closure reads.
:returns: The closure, the shape key, and whether the value may be complex.
:raises ValueError: If the node is an unbound ``DEFGATE`` parameter or an unknown function.
The message names only the offending node; :func:`expand_program` adds the instruction.
"""
if isinstance(node, MemoryReference):
if node not in references:
references.append(node)
index = references.index(node)
return (lambda slot_values, index=index: slot_values[index]), f"%{index}", False
if isinstance(node, Parameter):
raise ValueError(f"Unbound DEFGATE parameter {node}.")
if isinstance(node, BinaryExp):
left, left_key, left_complex = _build_expression(node.op1, references)
right, right_key, right_complex = _build_expression(node.op2, references)
operator = _BINARY_OPERATORS[type(node)]
return (
(lambda slot_values: operator(left(slot_values), right(slot_values))),
f"({left_key}{node.operator.strip()}{right_key})",
left_complex or right_complex,
)
if isinstance(node, Function):
if node.name not in _FUNCTIONS:
raise ValueError(f"Unknown Quil function {node.name!r}.")
function = _FUNCTIONS[node.name]
inner, inner_key, inner_complex = _build_expression(node.expression, references)
return (
(lambda slot_values: function(inner(slot_values))),
f"{node.name}({inner_key})",
inner_complex or node.name == "CIS",
)
literal = _literal_value(node)
return (lambda slot_values, literal=literal: jnp.asarray(literal)), repr(literal), isinstance(literal, complex)


def _compile_expression(expression: Any, slot_of: Callable[[MemoryReference], int]) -> ParameterExpression:
"""Compile a Quil expression over memory references into a :class:`ParameterExpression`.

:param expression: The gate parameter; must contain at least one memory reference.
:param slot_of: Maps a memory reference to its slot in the parameter vector.
:raises ValueError: If the expression contains an unbound ``DEFGATE`` parameter or an unknown
function.
"""
references: list[MemoryReference] = []
evaluate, key, is_complex = _build_expression(expression, references)
return ParameterExpression(tuple(slot_of(ref) for ref in references), key, is_complex, evaluate)


@dataclass(frozen=True, slots=True)
class ParametricGate:
"""A parametric gate whose matrix depends on runtime parameters.

Calling an instance with the flat parameter vector returns the gate's ``qx.Unitary``. The
constructor and parameter layout are exposed so that gates of the same kind can be built
constructor and argument layout are exposed so that gates of the same kind can be built
together in one vectorised operation.

:param gate_fn: The quax gate constructor (e.g. ``qx.gates.RX``), or a parametric
``DEFGATE`` callable.
:param param_indices: For each gate argument, its slot in the flat parameter vector, or
``-1`` when the argument is a literal number. Gates that read the same memory
reference share a slot; see :func:`expand_program`.
:param concrete_values: For each gate argument, its literal value (``nan`` for a slot).
:param arguments: One entry per gate argument: a literal number, or a
:class:`ParameterExpression` reading the parameter vector. Gates that read the same
memory reference share a slot; see :func:`expand_program`.
"""

gate_fn: Callable[..., qx.Operator]
param_indices: tuple[int, ...]
concrete_values: tuple[float, ...]

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Glad to see these go.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Likewise 🙂

arguments: tuple[float | complex | ParameterExpression, ...]

def __call__(self, params: Array) -> qx.Unitary:
"""Build the gate for one parameter vector."""
resolved: list[Any] = [
params[pi] if pi >= 0 else cv for pi, cv in zip(self.param_indices, self.concrete_values, strict=True)
]
resolved: list[Any] = [arg(params) if isinstance(arg, ParameterExpression) else arg for arg in self.arguments]
result = self.gate_fn(*resolved)
if not isinstance(result, qx.Unitary):
result = qx.Unitary.from_matrix(result.matrix, result.dims)
Expand Down Expand Up @@ -327,47 +455,49 @@ def _resolve_gate(inst: Gate) -> tuple[ExpandedOp, tuple[int, ...]]:
if any(_contained_mrefs(p) for p in inst.params): # type: ignore[arg-type]
gate_name = inst.name
if custom_gates is not None and gate_name in custom_gates:
gate_def = custom_gates[gate_name]
gate_def, is_builtin = custom_gates[gate_name], False
elif gate_name in qx.gates.QUANTUM_GATES:
gate_def = qx.gates.QUANTUM_GATES[gate_name]
gate_def, is_builtin = qx.gates.QUANTUM_GATES[gate_name], True
else:
raise KeyError(f"Unknown gate '{gate_name}'.")
if isinstance(gate_def, qx.Unitary):
raise ValueError(f"Gate '{gate_name}' is not parametric but {inst.out()!r} passes parameters.")

param_indices: list[int] = []
concrete_values: list[float] = []
arguments: list[float | complex | ParameterExpression] = []
for p in inst.params:
mrefs = _contained_mrefs(p) # type: ignore[arg-type]
if not mrefs:
# A concrete number: a compile-time constant for this gate.
param_indices.append(-1)
concrete_values.append(float(np.real(p)))
elif not isinstance(p, MemoryReference):
# An arithmetic expression over one or more memory regions, e.g.
# ``RX(theta[0] / 2) 0``. Each ParametricGate argument maps to a single
# slot of the flat parameter vector, which is what lets the simulator
# batch same-shaped gates under one ``jax.vmap``; an arbitrary
# expression would have to become part of that batching key.
raise ValueError(
f"Gate parameter {p} in {inst.out()!r} is an expression over memory "
f"region(s) {sorted(m.name for m in mrefs)}, which is not supported. "
"Pass the parameter directly (e.g. RX(theta[0]) with the division folded "
"into the value you bind), or substitute concrete values into the program "
"before simulating."
)
elif p.name in measure_regs:
# Classically-conditioned angle: the value is only known mid-circuit.
argument: float | complex | ParameterExpression
if not _contained_mrefs(p): # type: ignore[arg-type]
# A literal: a compile-time constant for this gate.
argument = _literal_value(p)
is_complex = isinstance(argument, complex)
else:
feed_forward = [m for m in _contained_mrefs(p) if m.name in measure_regs] # type: ignore[arg-type]
if feed_forward:
# Classically-conditioned angle: the value is only known mid-circuit.
raise ValueError(
f"Gate parameter {p} in {inst.out()!r} reads memory region "
f"'{feed_forward[0].name}', which is written by a MEASURE in this program. "
"Feed-forward (classically-conditioned) parameters are not supported."
)
try:
argument = _compile_expression(
p, lambda ref: slots.setdefault((ref.name, ref.offset), len(slots))
)
except ValueError as error:
# The compiler sees one expression; name the instruction it came from.
raise ValueError(f"Gate parameter {p} in {inst.out()!r}: {error}") from error
is_complex = argument.is_complex
# Checked for literals too: quax's built-in constructors take real angles, and a
# complex one otherwise surfaces much later as an opaque error from deep in quax.
if is_complex and is_builtin:
raise ValueError(
f"Gate parameter {p} in {inst.out()!r} reads memory region "
f"'{p.name}', which is written by a MEASURE in this program. "
"Feed-forward (classically-conditioned) parameters are not supported."
f"Gate parameter {p} in {inst.out()!r} is complex-valued (it contains CIS or a complex "
f"literal), but the built-in gate {gate_name} takes real angles. Complex arguments are "
"only supported for DEFGATE gates."
)
else:
param_indices.append(slots.setdefault((p.name, p.offset), len(slots)))
concrete_values.append(float("nan"))
arguments.append(argument)

return ParametricGate(gate_def, tuple(param_indices), tuple(concrete_values)), qubits
return ParametricGate(gate_def, tuple(arguments)), qubits

# Fixed gate → resolve to Unitary now.
unitary = get_instruction_unitary(inst, custom_gates=custom_gates)
Expand Down
Loading
Loading