From 500adafe9f6561545776106b8f067791ce3272f4 Mon Sep 17 00:00:00 2001 From: Nathaniel Tornow Date: Fri, 18 Sep 2026 17:36:51 +0200 Subject: [PATCH] analysis: reference analysis for squin qubits and jeff wires --- src/bloqade/analysis/reference/__init__.py | 25 ++ src/bloqade/analysis/reference/analysis.py | 263 +++++++++++++ src/bloqade/analysis/reference/impls.py | 156 ++++++++ src/bloqade/analysis/reference/lattice.py | 192 ++++++++++ src/bloqade/{jeff => }/constants.py | 6 +- src/bloqade/jeff/analysis/wire.py | 259 +++++++++++++ src/bloqade/jeff/dialects/stmts/call.py | 14 +- src/bloqade/jeff/dialects/stmts/wire.py | 10 +- src/bloqade/jeff/emit/classical.py | 6 +- src/bloqade/qubit/analysis/__init__.py | 4 +- src/bloqade/qubit/analysis/reference_impl.py | 21 ++ src/bloqade/squin/analysis/reference.py | 53 +++ test/analysis/reference/test_lattice.py | 237 ++++++++++++ test/analysis/reference/test_reference.py | 371 +++++++++++++++++++ test/jeff/test_dialect.py | 6 +- test/jeff/test_reference_wire.py | 285 ++++++++++++++ test/jeff/test_statements.py | 7 - 17 files changed, 1889 insertions(+), 26 deletions(-) create mode 100644 src/bloqade/analysis/reference/__init__.py create mode 100644 src/bloqade/analysis/reference/analysis.py create mode 100644 src/bloqade/analysis/reference/impls.py create mode 100644 src/bloqade/analysis/reference/lattice.py rename src/bloqade/{jeff => }/constants.py (67%) create mode 100644 src/bloqade/jeff/analysis/wire.py create mode 100644 src/bloqade/qubit/analysis/reference_impl.py create mode 100644 src/bloqade/squin/analysis/reference.py create mode 100644 test/analysis/reference/test_lattice.py create mode 100644 test/analysis/reference/test_reference.py create mode 100644 test/jeff/test_reference_wire.py diff --git a/src/bloqade/analysis/reference/__init__.py b/src/bloqade/analysis/reference/__init__.py new file mode 100644 index 000000000..37b2a598f --- /dev/null +++ b/src/bloqade/analysis/reference/__init__.py @@ -0,0 +1,25 @@ +"""This package holds the analysis that states which root each value refers to. + +`ReferenceAnalysis` holds the rules for kirin's own dialects. A subclass adds the +roots of one kind of tracked state, such as squin qubits or jeff wires. +""" + +from . import impls as impls +from .lattice import ( + CARRIED as CARRIED, + UNTRACKED as UNTRACKED, + Ref as Ref, + Slot as Slot, + Items as Items, + Whole as Whole, + Bottom as Bottom, + Unknown as Unknown, + Register as Register, + Returned as Returned, + Positions as Positions, + Untracked as Untracked, +) +from .analysis import ( + KEY as KEY, + ReferenceAnalysis as ReferenceAnalysis, +) diff --git a/src/bloqade/analysis/reference/analysis.py b/src/bloqade/analysis/reference/analysis.py new file mode 100644 index 000000000..40d192928 --- /dev/null +++ b/src/bloqade/analysis/reference/analysis.py @@ -0,0 +1,263 @@ +"""This module holds the forward analysis that states what each value refers to.""" + +from abc import ABC, abstractmethod +from dataclasses import dataclass + +from kirin import ir, types +from kirin.analysis.forward import Forward, ForwardFrame + +from bloqade.constants import constant_int + +from .lattice import ( + CARRIED, + UNTRACKED, + Ref, + Root, + Slot, + Whole, + Members, + Unknown, + Register, + Returned, + Positions, +) + +KEY = "reference" +"""The registry key of the rules for kirin's own dialects.""" + + +def _inside(root: Root, code: ir.Statement) -> bool: + """Return True if the function body `code` allocates or receives `root`.""" + return code.is_ancestor(root.call if isinstance(root, Returned) else root.owner) + + +@dataclass +class ReferenceAnalysis(Forward[Ref], ABC): + """A forward analysis that states which root each value refers to.""" + + keys = (KEY, "absint") + lattice = Ref + + @abstractmethod + def kind(self, type_: types.TypeAttribute) -> type[Whole] | type[Register] | None: + """Return what a value of type `type_` refers to. + + The result is `Whole` for one item of tracked state, `Register` for a + register of items, and None for a value without tracked state. + """ + + @abstractmethod + def register_length(self, root: ir.SSAValue, call: Returned | None) -> int | None: + """Return the static length of the register root `root`, or None. + + If `root` lies inside a callee, `call` names the call of that callee, so + the length can come from a constant argument of the call. + """ + + def run( + self, method: ir.Method, *args: Ref, **kwargs: Ref + ) -> tuple[ForwardFrame[Ref], Ref]: + """Analyze `method` and every function that it calls. + + Without `args`, each parameter of a tracked type is a root, a tuple + parameter with a tracked member is `Unknown`, and every other parameter is + `Untracked`. The returned frame holds the reference of each value of + `method`, and the returned value is the result of `method`. + """ + if not args and not kwargs: + params: list[Ref] = [] + for arg in method.callable_region.blocks[0].args[1:]: + if (kind := self.kind(arg.type)) is not None: + params.append(kind(arg)) + elif ( + isinstance(arg.type, types.Generic) + and arg.type.is_subseteq(types.Tuple) + and any(self.kind(member) is not None for member in arg.type.vars) + ): + params.append(Unknown("a member of a tuple parameter")) + else: + params.append(UNTRACKED) + args = tuple(params) + return super().run(method, *args, **kwargs) + + def recursion_limit_reached(self) -> Ref: + """Return `Unknown` for the deepest call of a recursion.""" + return Unknown("the result of a recursive call") + + def method_self(self, method: ir.Method) -> Ref: + """Return `Untracked` for the method object.""" + return UNTRACKED + + def eval_fallback( + self, frame: ForwardFrame[Ref], node: ir.Statement + ) -> tuple[Ref, ...]: + """Give each tracked result of a statement without a rule `Unknown`.""" + return self.unknown_results(node, f"a value computed by '{node.name}'") + + def unknown_results(self, stmt: ir.Statement, reason: str) -> tuple[Ref, ...]: + """Return `Unknown` for each tracked result of `stmt` and `Untracked` else.""" + return tuple( + Unknown(reason) if self.kind(r.type) is not None else UNTRACKED + for r in stmt.results + ) + + def origin(self, root: Root) -> tuple[ir.SSAValue, Returned | None]: + """Return the SSA value that allocates or receives `root`, and its call. + + A `Returned` root leads to the root inside the callee, and the innermost + `Returned` on that path names the call. A root of the analyzed function + has no call. + """ + call = None + while isinstance(root, Returned): + call = root + root = root.inner + return root, call + + def items(self, ref: Ref) -> tuple[Ref, ...] | None: + """Return the tracked items that `ref` holds, or None if they are unknown. + + One item is itself. A literal list holds its members. A register of static + length holds one slot per index. + """ + match ref: + case Whole() | Slot(): + return (ref,) + case Members(members): + return members + case Register(root): + size = self.register_length(*self.origin(root)) + if size is not None: + return tuple(Slot(root, i) for i in range(size)) + return None + + def index(self, ref: Ref, index: int | ir.SSAValue) -> Ref: + """Return the reference of the item at `index` of the list or register `ref`. + + A register gives a `Slot`. If the register has a static length, a constant + index becomes an index in the range 0 to the length minus 1. A literal list + gives its member at a constant index. + """ + constant = index if isinstance(index, int) else constant_int(index) + match ref: + case Register(root): + if constant is None: + return Slot(root, index) + size = self.register_length(*self.origin(root)) + if size is None: + return Slot(root, constant) + if not -size <= constant < size: + return Unknown("a constant index out of range") + return Slot(root, constant % size) + case Members(members): + if constant is None: + return Unknown("a list or tuple read at a runtime index") + if not -len(members) <= constant < len(members): + return Unknown("a constant index out of range") + return members[constant] + case Unknown(): + return ref + return Unknown("an index into a value that is not a register or a list") + + def call_result(self, frame: ForwardFrame[Ref], call: ir.Statement) -> Ref: + """Return the reference of the result of `call` in the terms of the caller. + + The analysis runs the callee with the caller's references as arguments, like + kirin's type inference. Kirin gives a call that reaches itself with the same + references `Bottom` until the result settles, and a recursion that runs to + kirin's `max_depth` returns `Unknown` from its deepest call. + """ + callee = call.get_present_trait(ir.StaticCall).get_callee(call) + args = frame.get_values(call.args) + _, result = self.call(callee.code, self.method_self(callee), *args) + return self._returned(call, callee, args, result) + + def _returned( + self, + call: ir.Statement, + callee: ir.Method, + args: tuple[Ref, ...], + result: Ref, + ) -> Ref: + """Return `result` of `callee` with each root that the callee owns renamed. + + A root that the callee allocates and returns whole at position `p` becomes + `Returned(call, p)`, where `p` counts the members of a `Positions` result. + A root inside a list or a nested tuple has no position and is `Unknown`. A + parameter + of the callee is a root only when the callee is the method under analysis, + in a recursion, and it becomes the root of the argument at its position. + Every other root that the callee owns is `Unknown`. A slot at an index + that is a parameter of the callee takes the argument of `call` as its + index, and a slot at an index that the callee computes is `Unknown`. + """ + code = callee.code + params = callee.callable_region.blocks[0].args[1:] + arguments: dict[ir.SSAValue, ir.SSAValue] = dict(zip(params, call.args)) + passed: dict[ir.SSAValue, Ref] = dict(zip(params, args)) + returned: dict[Root, Root] = {} + outputs = result.refs if isinstance(result, Positions) else (result,) + for position, ref in enumerate(outputs): + if isinstance(ref, (Whole, Register)) and _inside(ref.root, code): + returned.setdefault(ref.root, Returned(call, position, ref.root)) + + def root_in_caller(root: Root) -> Root | None: + """Return the root as the caller names it, or None if it cannot.""" + if not _inside(root, code): + return root + if isinstance(root, ir.SSAValue): + match passed.get(root): + case Whole(named) | Register(named): + return named + return returned.get(root) + + def rename(ref: Ref, nested: bool) -> Ref: + owned = ( + "a root returned inside a list or a nested tuple" + if nested + else "a root that the callee owns" + ) + match ref: + case Whole(root) | Register(root): + if (named := root_in_caller(root)) is None: + return Unknown(owned) + return type(ref)(named) + case Slot(root, index): + if (named := root_in_caller(root)) is None: + return Unknown(owned) + if index in arguments: + return self.index(Register(named), arguments[index]) + if isinstance(index, ir.SSAValue) and _inside(index, code): + return Unknown("an item at an index that the callee computes") + return Slot(named, index) + case Members(members): + return type(ref)(tuple(rename(m, True) for m in members)) + return ref + + if isinstance(result, Positions): + return Positions(tuple(rename(m, False) for m in result.refs)) + return rename(result, False) + + def run_loop( + self, + frame: ForwardFrame[Ref], + stmt: ir.Statement, + body: ir.Region, + carried: tuple[Ref, ...], + ) -> tuple[Ref, ...]: + """Run a loop body until the references that it carries stop changing. + + The body gets `Untracked` for the loop index and the carried references. + """ + # A join moves a carried reference up, and each one can move up once. + for _ in range(len(carried) + 1): + with self.new_frame(stmt, has_parent_access=True) as inner: + yielded = self.frame_call_region(inner, stmt, body, UNTRACKED, *carried) + if not isinstance(yielded, tuple) or len(yielded) != len(carried): + return self.unknown_results(stmt, CARRIED) + joined = tuple(c.join(y) for c, y in zip(carried, yielded, strict=True)) + if joined == carried: + frame.entries.update(inner.entries) + return carried + carried = joined + raise AssertionError("a carried reference moved up the lattice twice") diff --git a/src/bloqade/analysis/reference/impls.py b/src/bloqade/analysis/reference/impls.py new file mode 100644 index 000000000..688a7804c --- /dev/null +++ b/src/bloqade/analysis/reference/impls.py @@ -0,0 +1,156 @@ +"""This module holds the reference rules for kirin's own dialects.""" + +from kirin import interp +from kirin.dialects import py, scf, func, ilist +from kirin.analysis.forward import ForwardFrame + +from bloqade.constants import constant_int + +from .lattice import ( + UNTRACKED, + Ref, + Items, + Members, + Unknown, + Register, + Positions, +) +from .analysis import KEY, ReferenceAnalysis + + +@py.assign.dialect.register(key=KEY) +class _Assign(interp.MethodTable): + """A method table that passes a reference through an alias.""" + + @interp.impl(py.assign.Alias) + def alias( + self, + analysis: ReferenceAnalysis, + frame: ForwardFrame[Ref], + stmt: py.assign.Alias, + ) -> tuple[Ref, ...]: + """Return the reference of the aliased value.""" + return (frame.get(stmt.value),) + + +@ilist.dialect.register(key=KEY) +class _IList(interp.MethodTable): + """A method table that builds the references of literal lists and maps.""" + + @interp.impl(ilist.New) + def new( + self, analysis: ReferenceAnalysis, frame: ForwardFrame[Ref], stmt: ilist.New + ) -> tuple[Ref, ...]: + """Return `Items` if the list holds tracked state, and `Untracked` otherwise.""" + members = frame.get_values(stmt.values) + if analysis.kind(stmt.result.type) is not None or any( + m != UNTRACKED for m in members + ): + return (Items(members),) + return (UNTRACKED,) + + @interp.impl(ilist.Map) + def map_( + self, analysis: ReferenceAnalysis, frame: ForwardFrame[Ref], stmt: ilist.Map + ) -> tuple[Ref, ...]: + """Return the result as a register root if the map produces tracked values.""" + if analysis.kind(stmt.result.type) is Register: + return (Register(stmt.result),) + return (UNTRACKED,) + + +@py.binop.dialect.register(key=KEY) +class _BinOp(interp.MethodTable): + """A method table that concatenates literal lists.""" + + @interp.impl(py.binop.Add) + def add( + self, analysis: ReferenceAnalysis, frame: ForwardFrame[Ref], stmt: py.binop.Add + ) -> tuple[Ref, ...]: + """Join the items of two literal lists, and make other sums `Unknown`.""" + left, right = frame.get(stmt.lhs), frame.get(stmt.rhs) + if isinstance(left, Items) and isinstance(right, Items): + return (Items(left.refs + right.refs),) + return analysis.unknown_results(stmt, "a concatenation of registers") + + +@py.tuple.dialect.register(key=KEY) +class _Tuple(interp.MethodTable): + """A method table that builds the references of tuples.""" + + @interp.impl(py.tuple.New) + def new( + self, analysis: ReferenceAnalysis, frame: ForwardFrame[Ref], stmt: py.tuple.New + ) -> tuple[Ref, ...]: + """Return the references of the tuple members as `Positions`.""" + return (Positions(frame.get_values(stmt.args)),) + + +@py.indexing.dialect.register(key=KEY) +class _Indexing(interp.MethodTable): + """A method table that reads items of registers, lists and tuples.""" + + @interp.impl(py.indexing.GetItem) + def getitem( + self, + analysis: ReferenceAnalysis, + frame: ForwardFrame[Ref], + stmt: py.indexing.GetItem, + ) -> tuple[Ref, ...]: + """Return the reference of the item at the index.""" + obj = frame.get(stmt.obj) + kind = analysis.kind(stmt.result.type) + if constant_int(stmt.index) is None: + # A register read with a slice gives a register, which has no root. + if isinstance(obj, Register) and kind is Register: + return (Unknown("a slice of a register"),) + # A tuple read at a runtime index is fine if the item is untracked. + if isinstance(obj, Positions) and kind is None: + return (UNTRACKED,) + # An untracked item of a list or tuple is looked up. Any other untracked + # read, such as a classical list, needs no reference. + if kind is None and not isinstance(obj, Members): + return (UNTRACKED,) + return (analysis.index(obj, stmt.index),) + + +@func.dialect.register(key=KEY) +class _Func(interp.MethodTable): + """A method table that runs the callee at each call with the caller's references.""" + + @interp.impl(func.Return) + def return_( + self, analysis: ReferenceAnalysis, frame: ForwardFrame[Ref], stmt: func.Return + ) -> interp.ReturnValue[Ref]: + """Return the reference of the returned value.""" + return interp.ReturnValue(frame.get(stmt.value)) + + @interp.impl(func.Invoke) + def invoke( + self, analysis: ReferenceAnalysis, frame: ForwardFrame[Ref], stmt: func.Invoke + ) -> tuple[Ref, ...]: + """Return the reference of the result in the terms of the caller.""" + return (analysis.call_result(frame, stmt),) + + @interp.impl(func.Call) + def call( + self, analysis: ReferenceAnalysis, frame: ForwardFrame[Ref], stmt: func.Call + ) -> tuple[Ref, ...]: + """Give each tracked result of a dynamic call `Unknown`.""" + return analysis.unknown_results(stmt, "the result of a dynamic call") + + +@scf.dialect.register(key=KEY) +class _Scf(interp.MethodTable): + """A method table that joins the references that an `scf.for` loop carries. + + A carried value keeps its reference if the body hands back the same one. + """ + + @interp.impl(scf.For) + def for_( + self, analysis: ReferenceAnalysis, frame: ForwardFrame[Ref], stmt: scf.For + ) -> tuple[Ref, ...]: + """Run the body until the carried references stop changing.""" + carried = frame.get_values(stmt.initializers) + return analysis.run_loop(frame, stmt, stmt.body, carried) diff --git a/src/bloqade/analysis/reference/lattice.py b/src/bloqade/analysis/reference/lattice.py new file mode 100644 index 000000000..4e14939a0 --- /dev/null +++ b/src/bloqade/analysis/reference/lattice.py @@ -0,0 +1,192 @@ +"""This module holds the lattice of references that the reference analysis computes.""" + +from dataclasses import field, dataclass + +from kirin import ir, types +from kirin.lattice import ( + SingletonMeta, + BoundedLattice, + SimpleJoinMixin, + SimpleMeetMixin, +) + + +@dataclass(frozen=True, repr=False) +class Returned: + """A root that a call allocates and returns at one position of its result. + + `call` is a statement with kirin's `StaticCall` trait. `inner` is the root + inside the callee that the call returns at `position`. + """ + + call: ir.Statement + position: int + inner: "Root" + + @property + def callee(self) -> ir.Method: + """Return the method that the call invokes, through kirin's `StaticCall`.""" + return self.call.get_present_trait(ir.StaticCall).get_callee(self.call) + + def __repr__(self) -> str: + """Return the call result that holds the root and its origin in the callee. + + The form is `%x = make()[0]`. A result that packs several outputs in a tuple + prints as `%pair[1] = make()[1]`. A result without a name prints as + `make()[0]`. + """ + origin = f"{self.callee.sym_name}()[{self.position}]" + results = self.call.results + packed = len(results) == 1 + value = results[0] if packed else results[self.position] + if not value.name: + return origin + if packed and self.callee.return_type.is_subseteq(types.Tuple): + return f"%{value.name}[{self.position}] = {origin}" + return f"%{value.name} = {origin}" + + +Root = ir.SSAValue | Returned +"""An owner of tracked state: an allocation, a parameter, or a `Returned` root.""" + +CARRIED = "a value carried by a loop or branch" +"""The reason of an `Unknown` that a join of two different references produces.""" + + +class Ref(SimpleJoinMixin["Ref"], SimpleMeetMixin["Ref"], BoundedLattice["Ref"]): + """A lattice element that states what a value refers to. + + The join of two different references is `Unknown`. + """ + + @classmethod + def top(cls) -> "Ref": + """Return `Unknown` with the reason `CARRIED`.""" + return Unknown(CARRIED) + + @classmethod + def bottom(cls) -> "Ref": + """Return the element of a value that the analysis never reached.""" + return Bottom() + + def is_subseteq(self, other: "Ref") -> bool: + """Return True if `other` is `Unknown` or equal to `self`.""" + return isinstance(other, Unknown) or self == other + + +def _name(value: ir.SSAValue) -> str: + """Return the printed name of an SSA value, such as `%qs`.""" + return f"%{value.name}" if value.name else "%?" + + +def _owner(root: Root) -> str: + """Return the printed name of a root.""" + return repr(root) if isinstance(root, Returned) else _name(root) + + +@dataclass(frozen=True, repr=False) +class Whole(Ref): + """A reference to one item of tracked state, such as one qubit.""" + + root: Root + + def __repr__(self) -> str: + """Return the short form, such as `Whole(%q)`.""" + return f"Whole({_owner(self.root)})" + + +@dataclass(frozen=True, repr=False) +class Register(Ref): + """A reference to a whole register, which holds items of tracked state.""" + + root: Root + + def __repr__(self) -> str: + """Return the short form, such as `Register(%qs)`.""" + return f"Register({_owner(self.root)})" + + +@dataclass(frozen=True, repr=False) +class Slot(Ref): + """A reference to one item of a register root. + + If the register has a static length, a constant index lies in the range 0 to the + length minus 1. + """ + + root: Root + index: int | ir.SSAValue + + def __repr__(self) -> str: + """Return the short form, such as `Slot(%qs, 2)`.""" + at = self.index if isinstance(self.index, int) else _name(self.index) + return f"Slot({_owner(self.root)}, {at})" + + +@dataclass(frozen=True, repr=False) +class Members(Ref): + """A value that holds other values, with one reference for each member. + + The join of two different member tuples is `Unknown`, even if some members are + equal. + """ + + refs: tuple[Ref, ...] + + +@dataclass(frozen=True, repr=False) +class Items(Members): + """The references of the items of a literal list.""" + + def __repr__(self) -> str: + """Return the short form, such as `[Whole(%a), Whole(%b)]`.""" + return "[" + ", ".join(map(repr, self.refs)) + "]" + + +@dataclass(frozen=True, repr=False) +class Positions(Members): + """The references of the positions of a tuple.""" + + def __repr__(self) -> str: + """Return the short form, such as `(Whole(%a), Untracked)`.""" + return "(" + ", ".join(map(repr, self.refs)) + ")" + + +@dataclass(frozen=True, repr=False) +class Untracked(Ref, metaclass=SingletonMeta): + """A value that refers to no tracked state.""" + + def __repr__(self) -> str: + """Return `Untracked`.""" + return "Untracked" + + +@dataclass(frozen=True, repr=False) +class Unknown(Ref): + """A value that may refer to tracked state that the analysis cannot name. + + `Unknown` is the top element. The reason explains the lost information, and it + does not take part in equality. + """ + + reason: str = field(compare=False, hash=False) + + def __repr__(self) -> str: + """Return the reason, such as `Unknown: a slice of a register`.""" + return f"Unknown: {self.reason}" + + +@dataclass(frozen=True, repr=False) +class Bottom(Ref, metaclass=SingletonMeta): + """A value that the analysis never reached, the bottom element.""" + + def is_subseteq(self, other: Ref) -> bool: + """Return True, because bottom lies below every element.""" + return True + + def __repr__(self) -> str: + """Return `Bottom`.""" + return "Bottom" + + +UNTRACKED = Untracked() diff --git a/src/bloqade/jeff/constants.py b/src/bloqade/constants.py similarity index 67% rename from src/bloqade/jeff/constants.py rename to src/bloqade/constants.py index e865b3d16..d39e31483 100644 --- a/src/bloqade/jeff/constants.py +++ b/src/bloqade/constants.py @@ -3,10 +3,10 @@ from kirin import ir -def const_int(value: ir.SSAValue) -> int | None: - """Return the integer that a constant value holds, or None. +def constant_int(value: ir.SSAValue) -> int | None: + """Return the integer that the constant `value` holds, or None. - If the constant holds a `bool`, the function returns None. + The owner of `value` must carry the `ConstantLike` trait. A `bool` gives None. """ owner = value.owner if not isinstance(owner, ir.Statement) or not owner.has_trait(ir.ConstantLike): diff --git a/src/bloqade/jeff/analysis/wire.py b/src/bloqade/jeff/analysis/wire.py new file mode 100644 index 000000000..9878948a6 --- /dev/null +++ b/src/bloqade/jeff/analysis/wire.py @@ -0,0 +1,259 @@ +"""This module holds the reference analysis for jeff wires and registers.""" + +from dataclasses import dataclass + +from kirin import ir, types, interp +from kirin.analysis.forward import ForwardFrame + +from bloqade.constants import constant_int +from bloqade.jeff.types import WireType, QuregType, qureg_length +from bloqade.jeff.dialects import stmts +from bloqade.analysis.reference import ( + KEY, + CARRIED, + UNTRACKED, + Ref, + Whole, + Bottom, + Unknown, + Register, + Returned, + Positions, + ReferenceAnalysis, +) + +WIRE_KEY = "jeff.reference" +"""The registry key of the rules for jeff statements.""" + + +@dataclass +class WireReferenceAnalysis(ReferenceAnalysis): + """A reference analysis whose roots are jeff wires and registers. + + A gate, a reset and a measurement hand each wire on to a result, so the result + refers to the root of the wire. + """ + + keys = (WIRE_KEY, KEY, "absint") + + def kind(self, type_: types.TypeAttribute) -> type[Whole] | type[Register] | None: + """Return `Whole` for a wire, `Register` for a register, else None.""" + if isinstance(type_, types.BottomType): + return None + if type_.is_subseteq(QuregType): + return Register + if type_.is_subseteq(WireType): + return Whole + return None + + def register_length(self, root: ir.SSAValue, call: Returned | None) -> int | None: + """Return the constant allocation size or the typed length of a register.""" + if isinstance(root, ir.ResultValue) and isinstance(root.stmt, stmts.RegAlloc): + return constant_int(root.stmt.size) + return qureg_length(root.type) + + +@stmts.wire.dialect.register(key=WIRE_KEY) +class _Wire(interp.MethodTable): + """A method table that follows wires through allocation, extraction and insert.""" + + @interp.impl(stmts.Alloc) + def alloc( + self, + analysis: WireReferenceAnalysis, + frame: ForwardFrame[Ref], + stmt: stmts.Alloc, + ) -> tuple[Ref, ...]: + """Return the new wire as a root.""" + return (Whole(stmt.result),) + + @interp.impl(stmts.RegAlloc) + def reg_alloc( + self, + analysis: WireReferenceAnalysis, + frame: ForwardFrame[Ref], + stmt: stmts.RegAlloc, + ) -> tuple[Ref, ...]: + """Return the new register as a root.""" + return (Register(stmt.result),) + + @interp.impl(stmts.Reset) + def reset( + self, + analysis: WireReferenceAnalysis, + frame: ForwardFrame[Ref], + stmt: stmts.Reset, + ) -> tuple[Ref, ...]: + """Hand the wire on.""" + return (frame.get(stmt.wire),) + + @interp.impl(stmts.MeasureNd) + def measure_nd( + self, + analysis: WireReferenceAnalysis, + frame: ForwardFrame[Ref], + stmt: stmts.MeasureNd, + ) -> tuple[Ref, ...]: + """Hand the wire on and make the bit `Untracked`.""" + return (frame.get(stmt.wire), UNTRACKED) + + @interp.impl(stmts.RegLength) + def length( + self, + analysis: WireReferenceAnalysis, + frame: ForwardFrame[Ref], + stmt: stmts.RegLength, + ) -> tuple[Ref, ...]: + """Hand the register on and make the length `Untracked`.""" + return (frame.get(stmt.reg), UNTRACKED) + + @interp.impl(stmts.Extract) + def extract( + self, + analysis: WireReferenceAnalysis, + frame: ForwardFrame[Ref], + stmt: stmts.Extract, + ) -> tuple[Ref, ...]: + """Hand the register on and make the wire a `Slot` of the register.""" + register = frame.get(stmt.reg) + return (register, analysis.index(register, stmt.index)) + + @interp.impl(stmts.Insert) + def insert( + self, + analysis: WireReferenceAnalysis, + frame: ForwardFrame[Ref], + stmt: stmts.Insert, + ) -> tuple[Ref, ...]: + """Hand the register on if the wire returns to the slot that it came from.""" + register, wire = frame.get(stmt.reg), frame.get(stmt.wire) + if isinstance(register, Register) and wire == analysis.index( + register, stmt.index + ): + return (register,) + return (Unknown("a register that holds a wire from another slot"),) + + +@stmts.gate.dialect.register(key=WIRE_KEY) +class _Gate(interp.MethodTable): + """A method table that hands each wire of a gate on to its result.""" + + @interp.impl(stmts.Gate) + @interp.impl(stmts.Ppr) + def gate( + self, + analysis: WireReferenceAnalysis, + frame: ForwardFrame[Ref], + stmt: stmts.Gate | stmts.Ppr, + ) -> tuple[Ref, ...]: + """Return the references of the targets and then of the controls.""" + return frame.get_values((*stmt.targets, *stmt.controls)) + + +@stmts.scf.dialect.register(key=WIRE_KEY) +class _Scf(interp.MethodTable): + """A method table that joins the references that jeff control flow carries.""" + + @interp.impl(stmts.For) + def for_( + self, + analysis: WireReferenceAnalysis, + frame: ForwardFrame[Ref], + stmt: stmts.For, + ) -> tuple[Ref, ...]: + """Run the body until the carried references stop changing.""" + carried = frame.get_values(stmt.state) + return analysis.run_loop(frame, stmt, stmt.body, carried) + + @interp.impl(stmts.Switch) + def switch( + self, + analysis: WireReferenceAnalysis, + frame: ForwardFrame[Ref], + stmt: stmts.Switch, + ) -> tuple[Ref, ...]: + """Join the references that the branches and the default yield.""" + inputs = frame.get_values(stmt.inputs) + joined = tuple(Bottom() for _ in stmt.results) + for region in stmt.regions: + with analysis.new_frame(stmt, has_parent_access=True) as inner: + yielded = analysis.frame_call_region(inner, stmt, region, *inputs) + frame.entries.update(inner.entries) + if not isinstance(yielded, tuple) or len(yielded) != len(stmt.results): + return analysis.unknown_results(stmt, CARRIED) + joined = tuple(a.join(b) for a, b in zip(joined, yielded, strict=True)) + return joined + + @interp.impl(stmts.While) + def while_( + self, + analysis: WireReferenceAnalysis, + frame: ForwardFrame[Ref], + stmt: stmts.While, + ) -> tuple[Ref, ...]: + """Run both regions until the carried references stop changing. + + The before region yields the condition and the outputs. The after region + takes the outputs and yields the next inputs. + """ + carried = frame.get_values(stmt.inputs) + # A join moves a carried reference up, and each one can move up once. + for _ in range(len(carried) + 1): + with analysis.new_frame(stmt, has_parent_access=True) as before: + yielded = analysis.frame_call_region( + before, stmt, stmt.before, *carried + ) + if not isinstance(yielded, tuple) or len(yielded) != len(carried) + 1: + return analysis.unknown_results(stmt, CARRIED) + outputs = yielded[1:] + with analysis.new_frame(stmt, has_parent_access=True) as after: + again = analysis.frame_call_region(after, stmt, stmt.after, *outputs) + if not isinstance(again, tuple) or len(again) != len(carried): + return analysis.unknown_results(stmt, CARRIED) + joined = tuple(c.join(n) for c, n in zip(carried, again, strict=True)) + if joined == carried: + frame.entries.update(before.entries) + frame.entries.update(after.entries) + return outputs + carried = joined + raise AssertionError("a carried reference moved up the lattice twice") + + @interp.impl(stmts.Yield) + def yield_( + self, + analysis: WireReferenceAnalysis, + frame: ForwardFrame[Ref], + stmt: stmts.Yield, + ) -> interp.YieldValue[Ref]: + """End the region with the yielded references.""" + return interp.YieldValue(frame.get_values(stmt.values)) + + +@stmts.call.dialect.register(key=WIRE_KEY) +class _Call(interp.MethodTable): + """A method table that runs the callee at each call with the caller's references.""" + + @interp.impl(stmts.Call) + def call( + self, + analysis: WireReferenceAnalysis, + frame: ForwardFrame[Ref], + stmt: stmts.Call, + ) -> tuple[Ref, ...]: + """Return the references of the outputs in the terms of the caller.""" + match analysis.call_result(frame, stmt): + case Positions(refs) if len(refs) == len(stmt.results): + return refs + case Unknown() | Bottom() as unsettled: + return tuple(unsettled for _ in stmt.results) + return analysis.unknown_results(stmt, "a return that does not match the call") + + @interp.impl(stmts.Return) + def return_( + self, + analysis: WireReferenceAnalysis, + frame: ForwardFrame[Ref], + stmt: stmts.Return, + ) -> interp.ReturnValue[Ref]: + """Return the references of the returned values as `Positions`.""" + return interp.ReturnValue(Positions(frame.get_values(stmt.values))) diff --git a/src/bloqade/jeff/dialects/stmts/call.py b/src/bloqade/jeff/dialects/stmts/call.py index 539eec99c..3bbfe1ca4 100644 --- a/src/bloqade/jeff/dialects/stmts/call.py +++ b/src/bloqade/jeff/dialects/stmts/call.py @@ -81,6 +81,15 @@ def verify_type(self) -> None: ) +class CallCallee(ir.StaticCall["Call"]): + """A trait that gives kirin the callee of a jeff call.""" + + @classmethod + def get_callee(cls, stmt: "Call") -> ir.Method: + """Return the method that `stmt` calls.""" + return stmt.callee + + @statement(dialect=dialect, init=False) class Call(ir.Statement): """A statement that calls a jeff function. @@ -89,6 +98,7 @@ class Call(ir.Statement): """ name = "call" + traits = frozenset({CallCallee()}) callee: ir.Method = info.attribute() inputs: tuple[ir.SSAValue, ...] = info.argument() @@ -109,10 +119,6 @@ def __init__( def verify(self) -> None: """Check that the call has one input per parameter and one result per output.""" super().verify() - if not isinstance(self.callee, ir.Method): - raise ir.ValidationError( - self, f"a call's callee {self.callee!r} is not a method" - ) if len(self.callee.self_type.params_type) != len(self.inputs): raise ir.ValidationError( self, "a call's argument count differs from the callee's" diff --git a/src/bloqade/jeff/dialects/stmts/wire.py b/src/bloqade/jeff/dialects/stmts/wire.py index 62bb72a79..4f5c2f485 100644 --- a/src/bloqade/jeff/dialects/stmts/wire.py +++ b/src/bloqade/jeff/dialects/stmts/wire.py @@ -6,8 +6,8 @@ from kirin import ir, types from kirin.decl import info, statement +from bloqade.constants import constant_int from bloqade.jeff.types import WireType, QuregType, is_subtype, qureg_length -from bloqade.jeff.constants import const_int dialect = ir.Dialect("jeff.wire") @@ -25,7 +25,7 @@ def _register_length(reg: ir.SSAValue) -> int | None: return length owner = reg.owner if isinstance(owner, RegAlloc): - return const_int(owner.size) + return constant_int(owner.size) if isinstance(owner, RegCreate): return len(owner.wires) if isinstance(owner, (Extract, Insert)): @@ -40,7 +40,7 @@ def _within(stmt: ir.Statement, reg: ir.SSAValue, index: ir.SSAValue) -> None: A negative constant index is always outside. The upper bound applies when the register length is known. """ - slot = const_int(index) + slot = constant_int(index) if slot is None: return length = _register_length(reg) @@ -151,14 +151,14 @@ class RegAlloc(ir.Statement): def verify(self) -> None: """Check that a constant size is zero or positive.""" super().verify() - size = const_int(self.size) + size = constant_int(self.size) if size is not None and size < 0: raise ir.ValidationError(self, f"a register of {size} qubits") def verify_type(self) -> None: """Check the operand types and that a typed length equals a constant size.""" super().verify_type() - size = const_int(self.size) + size = constant_int(self.size) typed = qureg_length(self.result.type) if size is not None and typed not in (None, size): raise ir.TypeCheckError( diff --git a/src/bloqade/jeff/emit/classical.py b/src/bloqade/jeff/emit/classical.py index cbb6a3d82..3f8695c32 100644 --- a/src/bloqade/jeff/emit/classical.py +++ b/src/bloqade/jeff/emit/classical.py @@ -12,9 +12,9 @@ IntArrayType, FloatArrayType, ) +from bloqade.constants import constant_int from bloqade.jeff.names import subkind from bloqade.jeff.dialects import stmts -from bloqade.jeff.constants import const_int from .base import EmitJeff, JeffFrame @@ -35,7 +35,7 @@ class _Classical(interp.MethodTable): """Emit the `jeff.classical` statements.""" @interp.impl(stmts.ConstInt) - def const_int( + def constant_int( self, emit: EmitJeff, frame: JeffFrame, stmt: stmts.ConstInt ) -> tuple[JeffValue, ...]: """Emit an integer constant.""" @@ -283,7 +283,7 @@ def int_array_zero( self, emit: EmitJeff, frame: JeffFrame, stmt: stmts.IntArrayZero ) -> tuple[JeffValue, ...]: """Emit an integer array that holds zeros.""" - length = const_int(stmt.size) + length = constant_int(stmt.size) op = frame.push( JeffOp( "intArray", diff --git a/src/bloqade/qubit/analysis/__init__.py b/src/bloqade/qubit/analysis/__init__.py index dea64d4e8..4d87d5dc5 100644 --- a/src/bloqade/qubit/analysis/__init__.py +++ b/src/bloqade/qubit/analysis/__init__.py @@ -1 +1,3 @@ -from . import address_impl as address_impl +"""This package holds the analysis rules for the qubit dialect.""" + +from . import address_impl as address_impl, reference_impl as reference_impl diff --git a/src/bloqade/qubit/analysis/reference_impl.py b/src/bloqade/qubit/analysis/reference_impl.py new file mode 100644 index 000000000..4cd72ae7a --- /dev/null +++ b/src/bloqade/qubit/analysis/reference_impl.py @@ -0,0 +1,21 @@ +"""This module holds the reference analysis rule for the qubit dialect.""" + +from kirin import interp +from kirin.analysis import ForwardFrame + +from bloqade.analysis.reference import KEY, Ref, Whole, ReferenceAnalysis + +from .. import stmts +from .._dialect import dialect + + +@dialect.register(key=KEY) +class _Qubit(interp.MethodTable): + """A method table that makes each new qubit a root.""" + + @interp.impl(stmts.New) + def new( + self, analysis: ReferenceAnalysis, frame: ForwardFrame[Ref], stmt: stmts.New + ) -> tuple[Ref, ...]: + """Return a reference to the new qubit as a whole root.""" + return (Whole(stmt.result),) diff --git a/src/bloqade/squin/analysis/reference.py b/src/bloqade/squin/analysis/reference.py new file mode 100644 index 000000000..253533845 --- /dev/null +++ b/src/bloqade/squin/analysis/reference.py @@ -0,0 +1,53 @@ +"""This module holds the reference analysis for squin kernels.""" + +from dataclasses import dataclass + +from kirin import ir, types +from kirin.dialects import ilist + +from bloqade import squin +from bloqade.types import QubitType +from bloqade.constants import constant_int +from bloqade.analysis.reference import ( + KEY, + Whole, + Register, + Returned, + ReferenceAnalysis, +) + + +@dataclass +class QubitReferenceAnalysis(ReferenceAnalysis): + """A reference analysis whose roots are squin qubits and registers. + + `qubit.new` allocates a qubit root, and `squin.qalloc` allocates a register root. + """ + + keys = (KEY, "absint") + + def kind(self, type_: types.TypeAttribute) -> type[Whole] | type[Register] | None: + """Return `Whole` for a qubit, `Register` for a list of qubits, else None.""" + if isinstance(type_, types.BottomType): + return None + if type_.is_subseteq(ilist.IListType[QubitType, types.Any]): + return Register + if type_.is_subseteq(QubitType): + return Whole + return None + + def register_length(self, root: ir.SSAValue, call: Returned | None) -> int | None: + """Return the static length of the register root `root`, or None. + + A register that `squin.qalloc` returns has the constant size of the call, + and a negative size allocates no qubit. A parameter has the length in its + type. + """ + if call is not None and call.callee is squin.qalloc: + size = constant_int(call.call.args[0]) + return None if size is None else max(0, size) + kind = root.type + length = kind.vars[1] if isinstance(kind, types.Generic) else None + if isinstance(length, types.Literal) and isinstance(length.data, int): + return length.data + return None diff --git a/test/analysis/reference/test_lattice.py b/test/analysis/reference/test_lattice.py new file mode 100644 index 000000000..89fe94ce0 --- /dev/null +++ b/test/analysis/reference/test_lattice.py @@ -0,0 +1,237 @@ +"""Test reference lattice operations, formatting, and queries.""" + +from typing import Any, Literal + +from kirin.dialects import func, ilist + +from bloqade import squin +from bloqade.types import Qubit, MeasurementResult +from bloqade.analysis.reference import ( + CARRIED, + UNTRACKED, + Slot, + Items, + Whole, + Bottom, + Unknown, + Register, + Positions, +) +from bloqade.squin.analysis.reference import ( + QubitReferenceAnalysis, +) + + +def analyzed(kernel): + mt = kernel.similar() + frame, _ = QubitReferenceAnalysis(mt.dialects).run(mt) + return mt, frame.entries + + +def returned(refs, call, position=0): + """Return the root that `call` returns at `position`, as the analysis names it.""" + ref = refs[call.result] + ref = ref.refs[position] if isinstance(ref, Positions) else ref + return ref.root + + +def allocations(mt, refs): + """Return the register roots of the `qalloc` calls of `mt`, in order.""" + return [returned(refs, n) for n in calls(mt, "qalloc")] + + +def calls(mt, name): + return [ + n + for n in mt.callable_region.walk() + if isinstance(n, func.Invoke) and n.callee.sym_name == name + ] + + +# -- kernels ------------------------------------------------------------------- + + +@squin.kernel +def two_registers(n: int, qs: ilist.IList[Qubit, Any]): + a = squin.qalloc(n) + b = squin.qalloc(n) + c = squin.qalloc(2) + d = squin.qalloc(2) + e = squin.qalloc(-1) + squin.cx(a[0], b[0]) + squin.cx(c[0], d[0]) + squin.cx(e[0], qs[0]) + + +@squin.kernel +def maker() -> tuple[Qubit, MeasurementResult]: + a = squin.qubit.new() + return a, squin.qubit.measure(a) + + +@squin.kernel +def returns_list() -> ilist.IList[Qubit, Any]: + a = squin.qubit.new() + b = squin.qubit.new() + return [a, b] + + +@squin.kernel +def returns_element() -> Qubit: + qs = squin.qalloc(2) + return qs[0] + + +@squin.kernel +def declares_three() -> ilist.IList[Qubit, Literal[3]]: + return squin.qalloc(2) + + +@squin.kernel +def reads_declared_three(): + qs = declares_three() + squin.h(qs[0]) + + +@squin.kernel +def reads_returned_list(): + qs = returns_list() + squin.h(qs[0]) + + +@squin.kernel +def reads_returned_element(): + squin.h(returns_element()) + + +@squin.kernel +def reads_pair_at_runtime(i: int): + made = maker() + squin.h(made[0]) + a = squin.qubit.new() + b = squin.qubit.new() + squin.h([a, b][i]) + squin.h([a, b][2]) + + +@squin.kernel(fold=False, typeinfer=False) +def takes_tuple(pair: tuple[Qubit, Qubit]): + squin.h(pair[0]) + + +# -- order and printing ---------------------------------------------------------- + + +def test_order(): + mt, refs = analyzed(two_registers) + a, b = allocations(mt, refs)[:2] + whole, slot = Whole(a), Slot(a, 0) + assert Bottom().join(whole) is whole and whole.join(Bottom()) is whole + assert whole.join(slot) == Unknown(CARRIED) + assert isinstance(whole.join(Whole(b)), Unknown) + assert isinstance(UNTRACKED.join(whole), Unknown) + assert whole.meet(Unknown("x")) is whole and Unknown("x").meet(whole) is whole + assert whole.meet(slot) is Bottom() and whole.meet(whole) is whole + assert Bottom().is_subseteq(whole) and whole.is_subseteq(Unknown("x")) + assert not whole.is_subseteq(slot) + + +def test_printed_forms(): + mt, refs = analyzed(two_registers) + a = allocations(mt, refs)[0] + n = mt.callable_region.blocks[0].args[1] + assert repr(Whole(a)) == "Whole(%a = qalloc()[0])" + assert repr(Slot(a, 0)) == "Slot(%a = qalloc()[0], 0)" + assert repr(Slot(a, n)) == "Slot(%a = qalloc()[0], %n)" + assert repr(Whole(n)) == "Whole(%n)" + assert repr(Items((Whole(n), Slot(n, 1)))) == "[Whole(%n), Slot(%n, 1)]" + assert repr(Positions((UNTRACKED, Unknown("why")))) == "(Untracked, Unknown: why)" + assert repr(Bottom()) == "Bottom" + mt, refs = analyzed(reads_pair_at_runtime) + (made,) = calls(mt, "maker") + assert repr(Whole(returned(refs, made))) == "Whole(%made[0] = maker()[0])" + + +# -- lengths --------------------------------------------------------------------- + + +def test_lengths_of_allocations_and_parameters(): + mt, refs = analyzed(two_registers) + analysis = QubitReferenceAnalysis(mt.dialects) + analysis.run(mt) + a, b, c, d, e = allocations(mt, refs) + qs = mt.callable_region.blocks[0].args[2] + + def static_length(root): + return analysis.register_length(*analysis.origin(root)) + + assert static_length(a) is None and static_length(c) == 2 + assert static_length(e) == 0 # a negative size allocates nothing + assert static_length(qs) is None + + +def reason(ref) -> str | None: + return ref.reason if isinstance(ref, Unknown) else None + + +def test_returns_the_lattice_cannot_root(): + mt, refs = analyzed(reads_returned_list) + (h,) = calls(mt, "h") + assert ( + reason(refs[h.inputs[0]]) == "a root returned inside a list or a nested tuple" + ) + mt, refs = analyzed(reads_returned_element) + (h,) = calls(mt, "h") + assert reason(refs[h.inputs[0]]) == "a root that the callee owns" + + +def test_the_allocated_length_is_what_a_returned_register_bears(): + mt = reads_declared_three.similar() + analysis = QubitReferenceAnalysis(mt.dialects) + frame, _ = analysis.run(mt) + (h,) = calls(mt, "h") + (call,) = calls(mt, "declares_three") + ref = frame.entries[h.inputs[0]] + assert ref == Slot(returned(frame.entries, call), 0) + assert analysis.register_length(*analysis.origin(ref.root)) == 2 + + +def test_reads_of_tuples_and_lists(): + mt, refs = analyzed(reads_pair_at_runtime) + from_tuple, runtime_list, out_of_range = calls(mt, "h") + (made,) = calls(mt, "maker") + assert refs[from_tuple.inputs[0]] == Whole(returned(refs, made)) + assert ( + reason(refs[runtime_list.inputs[0]]) + == "a list or tuple read at a runtime index" + ) + assert reason(refs[out_of_range.inputs[0]]) == "a constant index out of range" + + +def test_a_tuple_parameter_has_no_convention(): + mt, refs = analyzed(takes_tuple) + pair = mt.callable_region.blocks[0].args[1] + assert reason(refs[pair]) == "a member of a tuple parameter" + + +def test_only_a_register_has_slots(): + mt = two_registers.similar() + analysis = QubitReferenceAnalysis(mt.dialects) + frame, _ = analysis.run(mt) + a, *_ = allocations(mt, frame.entries) + qubit = Whole(mt.callable_region.blocks[0].args[1]) + assert analysis.index(Register(a), 1) == Slot(a, 1) + assert isinstance(analysis.index(qubit, 1), Unknown) + + +def test_the_items_of_a_reference(): + mt = two_registers.similar() + analysis = QubitReferenceAnalysis(mt.dialects) + frame, _ = analysis.run(mt) + a, b, c, *_ = allocations(mt, frame.entries) + qubit = Whole(mt.callable_region.blocks[0].args[1]) + assert analysis.items(qubit) == (qubit,) + assert analysis.items(Items((qubit, Slot(a, 0)))) == (qubit, Slot(a, 0)) + assert analysis.items(Register(c)) == (Slot(c, 0), Slot(c, 1)) + assert analysis.items(Register(a)) is None + assert analysis.items(UNTRACKED) is None diff --git a/test/analysis/reference/test_reference.py b/test/analysis/reference/test_reference.py new file mode 100644 index 000000000..736a3517f --- /dev/null +++ b/test/analysis/reference/test_reference.py @@ -0,0 +1,371 @@ +"""Test reference tracking through allocations, calls, indexing, and loops.""" + +import math +from typing import Any + +import pytest +from kirin import types +from kirin.dialects import py, scf, func, ilist + +from bloqade import squin +from bloqade.types import Qubit, QubitType, MeasurementResult +from bloqade.analysis.reference import ( + UNTRACKED, + Slot, + Items, + Whole, + Unknown, + Register, + Positions, +) +from bloqade.squin.analysis.reference import ( + QubitReferenceAnalysis, +) + +QUBIT_LIST = ilist.IListType[QubitType, types.Any] + +FIXPOINT = pytest.mark.skip( + reason="kirin solves a call that reaches itself as a fixpoint from 0.23" +) + + +def refs_of(kernel): + mt = kernel.similar() + frame, _ = QubitReferenceAnalysis(mt.dialects).run(mt) + return mt, frame.entries + + +def statements(mt, kind): + return [n for n in mt.callable_region.walk() if isinstance(n, kind)] + + +def returned(refs, call, position=0): + """Return the root that `call` returns at `position`, as the analysis names it.""" + ref = refs[call.result] + ref = ref.refs[position] if isinstance(ref, Positions) else ref + return ref.root + + +def allocations(mt, refs): + """Return the register roots of the `qalloc` calls of `mt`, in order.""" + return [returned(refs, n) for n in calls(mt, "qalloc")] + + +def calls(mt, name): + return [n for n in statements(mt, func.Invoke) if n.callee.sym_name == name] + + +# -- kernels ------------------------------------------------------------------- + + +@squin.kernel +def prep(q: Qubit) -> tuple[Qubit, MeasurementResult]: + squin.h(q) + r = q + return r, squin.qubit.measure(r) + + +@squin.kernel +def maker() -> tuple[Qubit, MeasurementResult]: + a = squin.qubit.new() + return a, squin.qubit.measure(a) + + +@squin.kernel +def through(q: Qubit) -> Qubit: + return prep(q)[0] + + +@squin.kernel +def twice() -> tuple[Qubit, Qubit]: + a = squin.qubit.new() + return a, a + + +@squin.kernel +def elements(qs: ilist.IList[Qubit, Any], i: int): + q = squin.qubit.new() + made = maker() + squin.cx(qs[1], qs[-1]) + squin.cx(qs[i], q) + squin.cx(made[0], q) + prep(q) + squin.h(q) + for _ in range(2): + squin.h(qs[0]) + return made[1] + + +@squin.kernel +def sized(): + qs = squin.qalloc(3) + squin.cx(qs[-1], qs[0]) + squin.broadcast.x([qs[0], qs[2]]) + return squin.qubit.measure(qs[2]) + + +@squin.kernel +def bad_reads(): + qs = squin.qalloc(3) + squin.broadcast.x(qs[0:2]) + squin.h(qs[5]) + + +@squin.kernel +def swaps_in_loop(): + q = squin.qubit.new() + p = squin.qubit.new() + for _ in range(2): + squin.h(q) + q = p + return squin.qubit.measure(q) + + +@squin.kernel +def rotation(theta: float): + q = squin.qubit.new() + squin.rx(theta, q) + squin.rz(math.pi / 2, q) + return squin.qubit.measure(q) + + +# -- the rules ----------------------------------------------------------------- + + +def test_parameters_allocations_and_elements(): + mt, refs = refs_of(elements) + _, qs, i = mt.callable_region.blocks[0].args + assert refs[qs] == Register(qs) and refs[i] == UNTRACKED + cx = calls(mt, "cx") + assert [refs[v] for v in cx[0].inputs] == [Slot(qs, 1), Slot(qs, -1)] + (new,) = calls(mt, "new") + fresh = Whole(returned(refs, new)) + assert [refs[v] for v in cx[1].inputs] == [Slot(qs, i), fresh] + (made,) = calls(mt, "maker") + assert [refs[v] for v in cx[2].inputs] == [Whole(returned(refs, made)), fresh] + assert refs[made.result] == Positions((Whole(returned(refs, made)), UNTRACKED)) + h_prepped, h_loop = calls(mt, "h") + assert refs[h_prepped.inputs[0]] == fresh # the qubit lent to prep stays q + assert refs[h_loop.inputs[0]] == Slot(qs, 0) + (ret,) = statements(mt, func.Return) + assert refs[ret.value] == UNTRACKED + + +def test_static_length_normalizes_constant_indices(): + mt, refs = refs_of(sized) + (qs,) = allocations(mt, refs) + (cx,) = calls(mt, "cx") + assert [refs[v] for v in cx.inputs] == [Slot(qs, 2), Slot(qs, 0)] + (operands,) = statements(mt, ilist.New) # the inlined broadcast's operand + assert refs[operands.result] == Items((Slot(qs, 0), Slot(qs, 2))) + + +def reason(ref) -> str | None: + return ref.reason if isinstance(ref, Unknown) else None + + +def test_unknowns_carry_their_reason(): + mt, refs = refs_of(bad_reads) + (x,) = calls(mt, "x") + assert reason(refs[x.inputs[0]]) == "a slice of a register" + (h,) = calls(mt, "h") + assert reason(refs[h.inputs[0]]) == "a constant index out of range" + + +def test_a_qubit_swapped_by_a_loop_is_unknown(): + """A loop-carried reference becomes unknown when the body changes its root.""" + mt, refs = refs_of(swaps_in_loop) + (loop,) = statements(mt, scf.For) + carried = [a for a in loop.body.blocks[0].args[1:] if a.type.is_subseteq(QubitType)] + assert "a value carried by a loop or branch" in [reason(refs[a]) for a in carried] + assert all(isinstance(refs[a], (Whole, Unknown)) for a in carried) + + +# -- what calls hand back ---------------------------------------------------------- + + +def test_a_lent_qubit_handed_back_is_the_callers(): + """Nested calls preserve the root of a returned argument.""" + mt, refs = refs_of(through) + (q,) = mt.callable_region.blocks[0].args[1:] + (ret,) = statements(mt, func.Return) + assert refs[ret.value] == Whole(q) + (call,) = calls(mt, "prep") + assert refs[call.result] == Positions((Whole(q), UNTRACKED)) + + +def test_an_allocation_made_by_a_call_is_rooted_at_the_call(): + mt, refs = refs_of(elements) + (made,) = calls(mt, "maker") + assert refs[made.result] == Positions((Whole(returned(refs, made)), UNTRACKED)) + + +@squin.kernel +def reads_twice(): + pair = twice() + squin.cx(pair[0], pair[1]) + + +def test_an_allocation_at_two_positions_is_one_root(): + """The caller sees the same qubit twice, rooted at its first position.""" + mt, refs = refs_of(reads_twice) + (call,) = calls(mt, "twice") + (cx,) = calls(mt, "cx") + first, second = (refs[v] for v in cx.inputs) + assert first == second == Whole(returned(refs, call)) + + +@squin.kernel +def make_pair() -> ilist.IList[Qubit, Any]: + return squin.qalloc(2) + + +@squin.kernel +def reads_made_pair(): + qs = make_pair() + squin.cx(qs[1], qs[-1]) + + +def test_a_returned_register_keeps_its_static_length(): + """The length of the register inside the callee normalizes a negative index.""" + mt, refs = refs_of(reads_made_pair) + (cx,) = calls(mt, "cx") + first, second = (refs[v] for v in cx.inputs) + assert isinstance(first, Slot) and isinstance(second, Slot) + assert first.root == second.root + assert first.index == second.index == 1 + assert repr(first) == "Slot(%qs = make_pair()[0], 1)" + + +@FIXPOINT +def test_a_recursive_call_hands_its_qubit_back(): + @squin.kernel + def rec(q: Qubit, n: int) -> Qubit: + squin.h(q) + if n > 0: + r = rec(q, n - 1) + else: + r = q + return r + + mt, refs = refs_of(rec) + (call,) = calls(mt, "rec") + assert refs[call.result] == Whole(mt.callable_region.blocks[0].args[1]) + + +def test_unknown_is_one_top_element(): + """Unknown reasons do not affect equality, and a join keeps an unknown's reason.""" + a, b = Unknown("a"), Unknown("b") + assert a == b and a.is_subseteq(b) and b.is_subseteq(a) + assert a.join(b) is b and b.join(a) is a + assert UNTRACKED.join(a) is a and a.join(UNTRACKED) is a + + +@squin.kernel +def branches(flag: bool): + a = squin.qubit.new() + b = squin.qubit.new() + if flag: + q = a + r = a + else: + q = a + r = b + squin.h(q) + squin.h(r) + + +def test_a_branch_keeps_a_reference_that_both_paths_yield(): + mt, refs = refs_of(branches) + a, _ = calls(mt, "new") + same, different = (refs[call.inputs[0]] for call in calls(mt, "h")) + assert same == Whole(returned(refs, a)) + assert reason(different) == "a value carried by a loop or branch" + + +@squin.kernel +def concatenated(): + a = squin.qubit.new() + b = squin.qubit.new() + squin.broadcast.h([a] + [b]) + + +def test_literal_lists_concatenate_item_by_item(): + mt, refs = refs_of(concatenated) + a, b = (Whole(returned(refs, call)) for call in calls(mt, "new")) + (add,) = statements(mt, py.binop.Add) + assert refs[add.result] == Items((a, b)) + + +@squin.kernel +def pick(qs: ilist.IList[Qubit, Any], i: int) -> Qubit: + return qs[i] + + +@squin.kernel +def pick_next(qs: ilist.IList[Qubit, Any], i: int) -> Qubit: + return qs[i + 1] + + +@squin.kernel +def picks(): + qs = squin.qalloc(3) + a = squin.qubit.new() + b = squin.qubit.new() + squin.h(pick(qs, -1)) + squin.h(pick(qs, 0)) + squin.h(pick([a, b], 1)) + squin.h(pick_next(qs, 0)) + + +def test_a_slot_at_a_parameter_index_takes_the_argument(): + """A slot keeps a parameter index, so the caller's constant resolves it. + + A literal list has no slots, so a read at a parameter index stays unknown. + """ + mt, refs = refs_of(picks) + (qs,) = allocations(mt, refs) + last, first, literal, computed = (refs[h.inputs[0]] for h in calls(mt, "h")) + assert last == Slot(qs, 2) and first == Slot(qs, 0) + assert reason(literal) == "a list or tuple read at a runtime index" + assert reason(computed) == "an item at an index that the callee computes" + + +def test_a_function_runs_in_terms_of_its_own_parameters(): + mt = pick.similar() + _, result = QubitReferenceAnalysis(mt.dialects).run(mt) + qs, i = mt.callable_region.blocks[0].args[1:] + assert result == Slot(qs, i) + + +@squin.kernel +def nested_pair() -> tuple[tuple[Qubit, Qubit], int]: + a = squin.qubit.new() + b = squin.qubit.new() + return (a, b), 1 + + +@squin.kernel +def reads_nested_pair() -> Qubit: + made = nested_pair() + inner = made[0] + return inner[0] + + +def test_a_root_inside_a_nested_tuple_is_unknown(): + """`Returned` names a position of the result, and a nested tuple has none.""" + mt, refs = refs_of(reads_nested_pair) + (call,) = calls(mt, "nested_pair") + ref = refs[call.result] + assert isinstance(ref, Positions) + assert ( + reason(ref.refs[0].refs[0]) == "a root returned inside a list or a nested tuple" + ) + assert ref.refs[1] == UNTRACKED + + +def test_a_bottom_type_is_not_tracked(): + analysis = QubitReferenceAnalysis(squin.kernel) + assert analysis.kind(types.Bottom) is None + assert analysis.kind(QubitType) is Whole + assert analysis.kind(QUBIT_LIST) is Register diff --git a/test/jeff/test_dialect.py b/test/jeff/test_dialect.py index c656d0c46..c96e76199 100644 --- a/test/jeff/test_dialect.py +++ b/test/jeff/test_dialect.py @@ -5,9 +5,9 @@ from kirin.dialects import func from bloqade import jeff +from bloqade.constants import constant_int from bloqade.jeff.types import qureg, family, is_bit, is_linear, qureg_length from bloqade.jeff.dialects import stmts -from bloqade.jeff.constants import const_int from .build import add, entry, method @@ -25,7 +25,7 @@ def test_a_type_outside_the_families_has_none(): def test_a_block_argument_is_no_constant(): _, (n,) = entry(types.Int) - assert const_int(n) is None + assert constant_int(n) is None def test_region_statements_require_a_yield(): @@ -55,7 +55,7 @@ def test_array_statements_build_and_verify(): def test_a_constant_without_a_value_is_no_constant(): block, _ = entry() none = add(block, func.ConstantNone()).result - assert const_int(none) is None + assert constant_int(none) is None def test_bottom_belongs_to_no_type(): diff --git a/test/jeff/test_reference_wire.py b/test/jeff/test_reference_wire.py new file mode 100644 index 000000000..36851df0d --- /dev/null +++ b/test/jeff/test_reference_wire.py @@ -0,0 +1,285 @@ +"""Test the reference analysis on jeff wires, registers, loops, switches and calls.""" + +import pytest +from kirin import types + +from bloqade import jeff +from bloqade.jeff.types import qureg +from bloqade.jeff.dialects import stmts +from bloqade.analysis.reference import ( + UNTRACKED, + Slot, + Bottom, + Whole, + Unknown, + Register, + Returned, + Positions, +) +from bloqade.jeff.analysis.wire import WireReferenceAnalysis + +from .build import add, entry, method, switch, for_loop, while_loop + +FIXPOINT = pytest.mark.skip( + reason="kirin solves a call that reaches itself as a fixpoint from 0.23" +) + + +def returned(mt): + """Return the references of the values that `mt` returns.""" + _, result = WireReferenceAnalysis(jeff.kernel).run(mt) + assert isinstance(result, Positions) + return result.refs + + +def reason(ref) -> str | None: + return ref.reason if isinstance(ref, Unknown) else None + + +def constants(block, *values): + return [add(block, stmts.ConstInt(value=v)).result for v in values] + + +def gate(block, name, *wires): + return add(block, stmts.Gate(tuple(wires), (), (), gate_name=name)).results + + +def test_a_gate_a_measurement_and_a_loop_hand_their_wires_on(): + block, (q0, q1) = entry(jeff.WireType, jeff.WireType) + (gated,) = gate(block, "h", q0) + measured = add(block, stmts.MeasureNd(gated)) + lo, hi, one = constants(block, 0, 2, 1) + loop = for_loop(block, lo, hi, one, (q1,), lambda b, i, s: gate(b, "x", s)) + mt = method( + block, + (measured.bit, measured.result_wire, loop.results[0]), + types.Generic(tuple, types.Bool, jeff.WireType, jeff.WireType), + inputs=(jeff.WireType, jeff.WireType), + ) + assert returned(mt) == (UNTRACKED, Whole(q0), Whole(q1)) + + +def test_a_loop_that_swaps_its_wires_loses_both(): + block, (q0, q1) = entry(jeff.WireType, jeff.WireType) + lo, hi, one = constants(block, 0, 2, 1) + loop = for_loop(block, lo, hi, one, (q0, q1), lambda b, i, a, c: (c, a)) + mt = method( + block, + tuple(loop.results), + types.Generic(tuple, jeff.WireType, jeff.WireType), + inputs=(jeff.WireType, jeff.WireType), + ) + assert [reason(ref) for ref in returned(mt)] == [ + "a value carried by a loop or branch" + ] * 2 + + +def flip_second_if(): + """Build a function that measures its first wire and may flip its second.""" + block, (w0, w1, c) = entry(jeff.WireType, jeff.WireType, types.Int) + picked = switch(block, c, (w1,), [lambda b, s: (s,)], lambda b, s: gate(b, "x", s)) + bit = add(block, stmts.Measure(w0)).bit + return method( + block, + (bit, picked.results[0]), + types.Generic(tuple, types.Bool, jeff.WireType), + inputs=(jeff.WireType, jeff.WireType, types.Int), + name="flip_second_if", + ) + + +def test_a_switch_hands_on_a_wire_that_every_branch_yields(): + callee = flip_second_if() + w1 = callee.callable_region.blocks[0].args[2] + assert returned(callee) == (UNTRACKED, Whole(w1)) + + +def test_a_call_hands_back_the_wire_of_the_caller(): + callee = flip_second_if() + block, (a, b) = entry(jeff.WireType, jeff.WireType) + (c,) = constants(block, 0) + call = add(block, stmts.Call(callee, (a, b, c), (types.Bool, jeff.WireType))) + mt = method( + block, + (call.results[1],), + jeff.WireType, + inputs=(jeff.WireType, jeff.WireType), + ) + assert returned(mt) == (Whole(b),) + + +def test_a_wire_that_a_callee_allocates_is_rooted_at_the_call(): + inner, _ = entry() + fresh = add(inner, stmts.Alloc()).result + callee = method(inner, (fresh,), jeff.WireType, name="make") + block, _ = entry() + call = add(block, stmts.Call(callee, (), (jeff.WireType,))) + mt = method(block, (call.results[0],), jeff.WireType) + (ref,) = returned(mt) + assert ref == Whole(Returned(call, 0, fresh)) + assert repr(ref) == "Whole(make()[0])" + + +def test_a_wire_inserted_into_its_own_slot_hands_the_register_on(): + block, (reg,) = entry(qureg(2)) + zero, one = constants(block, 0, 1) + extracted = add(block, stmts.Extract(reg, zero)) + (flipped,) = gate(block, "x", extracted.wire) + back = add(block, stmts.Insert(extracted.result_reg, zero, flipped)) + moved = add(block, stmts.Extract(back.result, zero)) + elsewhere = add(block, stmts.Insert(moved.result_reg, one, moved.wire)) + mt = method( + block, + (back.result, elsewhere.result), + types.Generic(tuple, jeff.QuregType, jeff.QuregType), + inputs=(qureg(2),), + ) + frame, _ = WireReferenceAnalysis(jeff.kernel).run(mt) + assert frame.entries[extracted.wire] == Slot(reg, 0) + assert frame.entries[back.result] == Register(reg) + assert reason(frame.entries[elsewhere.result]) == ( + "a register that holds a wire from another slot" + ) + + +@FIXPOINT +def test_a_call_that_only_returns_itself_has_no_reference(): + block, (q,) = entry(jeff.WireType) + mt = method(block, (q,), jeff.WireType, inputs=(jeff.WireType,), name="rec") + ret = block.last_stmt + call = stmts.Call(mt, (q,), (jeff.WireType,)) + call.insert_before(ret) + ret.replace_by(stmts.Return(call.results[0])) + (ref,) = returned(mt) + assert ref == Bottom() + + +@FIXPOINT +def test_a_recursion_with_a_base_case_hands_its_wire_back(): + """`rec(n, w)` flips `w` and calls itself with `n - 1` until `n` is 0.""" + block, (n, w) = entry(types.Int, jeff.WireType) + (zero,) = constants(block, 0) + done = add(block, stmts.IntLteS(n, zero)).result + mt = method( + block, + (w, n), + types.Generic(tuple, jeff.WireType, types.Int), + inputs=(types.Int, jeff.WireType), + name="rec", + ) + + def base(body, n_, w_): + return [w_, n_] + + def step(body, n_, w_): + (one,) = constants(body, 1) + less = add(body, stmts.IntSub(n_, one)).result + (flipped,) = gate(body, "x", w_) + call = add(body, stmts.Call(mt, (less, flipped), (jeff.WireType, types.Int))) + return [call.results[0], call.results[1]] + + ret = block.last_stmt + chosen = switch(block, done, (n, w), [step], base) + chosen.detach() + chosen.insert_before(ret) + ret.replace_by(stmts.Return(chosen.results[0], chosen.results[1])) + assert returned(mt) == (Whole(w), UNTRACKED) + + +@FIXPOINT +def test_a_mutually_recursive_call_that_only_returns_itself_has_no_reference(): + """`ping` calls `pong`, which calls `ping` again.""" + pong_block, (p,) = entry(jeff.WireType) + pong = method(pong_block, (p,), jeff.WireType, inputs=(jeff.WireType,), name="pong") + ping_block, (q,) = entry(jeff.WireType) + call_pong = add(ping_block, stmts.Call(pong, (q,), (jeff.WireType,))) + ping = method( + ping_block, + (call_pong.results[0],), + jeff.WireType, + inputs=(jeff.WireType,), + name="ping", + ) + ret = pong_block.last_stmt + call_ping = stmts.Call(ping, (p,), (jeff.WireType,)) + call_ping.insert_before(ret) + ret.replace_by(stmts.Return(call_ping.results[0])) + (ref,) = returned(ping) + assert ref == Bottom() + + +def test_a_returned_register_keeps_its_allocated_length(): + inner, _ = entry() + (three,) = constants(inner, 3) + made = add(inner, stmts.RegAlloc(three)).result + measured = add(inner, stmts.RegLength(made)) + callee = method(inner, (measured.result_reg,), jeff.QuregType, name="make") + block, (q,) = entry(jeff.WireType) + reset = add(block, stmts.Reset(q)).result + call = add(block, stmts.Call(callee, (), (jeff.QuregType,))) + (last,) = constants(block, -1) + extracted = add(block, stmts.Extract(call.results[0], last)) + mt = method( + block, + (reset, extracted.result_reg, extracted.wire), + types.Generic(tuple, jeff.WireType, jeff.QuregType, jeff.WireType), + inputs=(jeff.WireType,), + ) + root = Returned(call, 0, made) + assert returned(mt) == (Whole(q), Register(root), Slot(root, 2)) + + +def test_a_callee_runs_with_the_references_of_the_caller(): + callee = flip_second_if() + block, (a, b) = entry(jeff.WireType, jeff.WireType) + (c,) = constants(block, 0) + call = add(block, stmts.Call(callee, (a, b, c), (types.Bool, jeff.WireType))) + mt = method( + block, + (call.results[1],), + jeff.WireType, + inputs=(jeff.WireType, jeff.WireType), + ) + frame, _ = WireReferenceAnalysis(jeff.kernel).run(mt) + assert frame.entries[call.results[1]] == Whole(b) + + +def test_a_while_loop_hands_on_the_wire_that_it_carries(): + block, (q,) = entry(jeff.WireType) + + def before(b, w): + measured = add(b, stmts.MeasureNd(w)) + return (measured.bit, measured.result_wire) + + loop = while_loop(block, (q,), before, lambda b, w: gate(b, "x", w)) + mt = method(block, tuple(loop.results), jeff.WireType, inputs=(jeff.WireType,)) + frame, result = WireReferenceAnalysis(jeff.kernel).run(mt) + assert result == Positions((Whole(q),)) + inner = loop.after.blocks[0].args[0] + assert frame.entries[inner] == Whole(q) + + +def test_a_while_loop_that_swaps_its_wires_loses_both(): + block, (q0, q1) = entry(jeff.WireType, jeff.WireType) + + def before(b, a, c): + measured = add(b, stmts.MeasureNd(a)) + return (measured.bit, c, measured.result_wire) + + loop = while_loop(block, (q0, q1), before, lambda b, a, c: (a, c)) + mt = method( + block, + tuple(loop.results), + types.Generic(tuple, jeff.WireType, jeff.WireType), + inputs=(jeff.WireType, jeff.WireType), + ) + assert [reason(ref) for ref in returned(mt)] == [ + "a value carried by a loop or branch" + ] * 2 + + +def test_a_bottom_type_is_not_tracked(): + analysis = WireReferenceAnalysis(jeff.kernel) + assert analysis.kind(types.Bottom) is None + assert analysis.kind(jeff.WireType) is Whole + assert analysis.kind(jeff.QuregType) is Register diff --git a/test/jeff/test_statements.py b/test/jeff/test_statements.py index 7a5fa2cf8..3d49cd9f7 100644 --- a/test/jeff/test_statements.py +++ b/test/jeff/test_statements.py @@ -455,13 +455,6 @@ def test_rejects_a_constant_typed_against_its_bitwidth(): check(method(block, None, types.NoneType)) -def test_rejects_a_call_to_something_that_is_not_a_method(): - block, _ = entry() - add(block, stmts.Call(42, (), ())) # type: ignore[arg-type] - with pytest.raises(ir.ValidationError, match="is not a method"): - check(method(block, None, types.NoneType)) - - def test_rejects_a_negative_register_size(): block, _ = entry() size = add(block, stmts.ConstInt(value=-3)).result